mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-13 15:32:24 -06:00
Compare commits
39 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 7ea150fa71 | |||
| 4d665a5f62 | |||
| 3bc3250869 | |||
| db937486cf | |||
| 554257ac4d | |||
| 5f0004dc91 | |||
| 6cc1b3a5bd | |||
| cc9afe94cd | |||
| 136b75fdef | |||
| 4d1107839b | |||
| 7d66bc2159 | |||
| 165cbb2d29 | |||
| c79c47b940 | |||
| 660c273e8e | |||
| c7586abd0a | |||
| 14d57176ce | |||
| 96084ca5f3 | |||
| 8b92302247 | |||
| e195ca54a6 | |||
| 50277cd4de | |||
| fb190f8977 | |||
| 25b5e32089 | |||
| 339981a258 | |||
| 06de9ff83b | |||
| 924b976f1f | |||
| d5db817391 | |||
| fc8ceb4c72 | |||
| 07234dec4d | |||
| dd4cc0b30d | |||
| e7fe8fca9d | |||
| 42b9f89988 | |||
| 77c0a7736b | |||
| 872e1770e6 | |||
| a6e929b0a0 | |||
| 047680d669 | |||
| 0fd0ad3b2d | |||
| f3dba836dd | |||
| 7adda343fc | |||
| a20a058c59 |
@@ -22,6 +22,7 @@ OPENAI_API_KEY=sk-...
|
||||
# -- Authentication ------------------------------------------------------------
|
||||
# TURNSTONE_AUTH_ENABLED=true
|
||||
# TURNSTONE_AUTH_TOKEN=your-secret-token
|
||||
# TURNSTONE_JWT_SECRET=python -c "import secrets; print(secrets.token_hex(32))"
|
||||
|
||||
# -- Ports ---------------------------------------------------------------------
|
||||
# SERVER_PORT=8080
|
||||
|
||||
@@ -17,3 +17,5 @@ venv/
|
||||
.plan.md
|
||||
.plan-*.md
|
||||
.hypothesis/
|
||||
PROGRESS.md
|
||||
.coverage
|
||||
|
||||
+1
-1
@@ -34,7 +34,7 @@ RUN useradd --create-home --shell /bin/bash turnstone
|
||||
|
||||
# Install the wheel with all optional extras
|
||||
COPY --from=builder /build/wheels/*.whl /tmp/wheels/
|
||||
RUN pip install --no-cache-dir "$(ls /tmp/wheels/*.whl)[mq,console,sim,postgres]" \
|
||||
RUN pip install --no-cache-dir "$(ls /tmp/wheels/*.whl)[mq,console,sim,postgres,discord]" \
|
||||
&& rm -rf /tmp/wheels
|
||||
|
||||
# Health check script (stdlib only, no pip deps needed)
|
||||
|
||||
@@ -11,7 +11,7 @@ Named after the [Ruddy Turnstone](https://en.wikipedia.org/wiki/Ruddy_turnstone)
|
||||
|
||||
## What it does
|
||||
|
||||
Turnstone gives LLMs tools — shell, files, search, web, planning — and orchestrates multi-turn conversations where the model investigates, acts, and reports. It runs as:
|
||||
Turnstone gives LLMs tools — shell, files, search, web, planning — and orchestrates multi-turn conversations where the model investigates, acts, and reports. Native deferred tool loading for Anthropic and OpenAI APIs reduces token overhead and improves tool selection accuracy when MCP servers expose many tools; local models (vLLM, llama.cpp) get a transparent client-side BM25 fallback. 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
|
||||
@@ -19,13 +19,9 @@ Turnstone gives LLMs tools — shell, files, search, web, planning — and orche
|
||||
- **Cluster dashboard** — real-time view of all nodes and workstreams, workstream creation with node targeting, reverse proxy for server UIs (only the console port needs network access)
|
||||
- **Cluster simulator** — test the stack at scale (up to 1000 nodes) without an LLM backend
|
||||
|
||||
```
|
||||
External System → Message Queue → Bridge (per node) → Turnstone Server → LLM + Tools
|
||||
↓
|
||||
Pub/Sub → Progress Events → External System
|
||||
↓
|
||||
turnstone-console → Cluster Dashboard (browser)
|
||||
```
|
||||
<p align="center">
|
||||
<img src="docs/diagrams/architecture-overview.svg" alt="Turnstone system architecture — data flow from clients through gateways, Redis MQ, cluster nodes, to LLM providers" width="960"/>
|
||||
</p>
|
||||
|
||||
## Quickstart
|
||||
|
||||
@@ -111,69 +107,7 @@ All frontends connect to any OpenAI-compatible API (vLLM, NVIDIA NIM/NGC, llama.
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
turnstone/
|
||||
├── core/ # UI-agnostic engine
|
||||
│ ├── session.py # ChatSession — multi-turn loop, tool dispatch, agents
|
||||
│ ├── providers/ # LLM provider adapters (OpenAI, Anthropic)
|
||||
│ │ ├── _protocol.py # LLMProvider protocol, ModelCapabilities, StreamChunk
|
||||
│ │ ├── _openai.py # OpenAI-compatible (OpenAI, vLLM, llama.cpp)
|
||||
│ │ └── _anthropic.py # Anthropic Messages API (native streaming, thinking)
|
||||
│ ├── tools.py # Tool definitions (auto-loaded from JSON)
|
||||
│ ├── workstream.py # WorkstreamManager — parallel independent sessions
|
||||
│ ├── mcp_client.py # MCP client manager (external tool servers)
|
||||
│ ├── model_registry.py # ModelRegistry — named models, fallback routing, per-workstream selection
|
||||
│ ├── config.py # Unified TOML config (~/.config/turnstone/config.toml)
|
||||
│ ├── memory.py # Persistence facade (delegates to storage/)
|
||||
│ ├── storage/ # Pluggable storage backend (SQLite + PostgreSQL)
|
||||
│ ├── metrics.py # Prometheus-compatible metrics collector
|
||||
│ ├── healthcheck.py # Backend health monitor + circuit breaker
|
||||
│ ├── ratelimit.py # Per-IP token-bucket rate limiter
|
||||
│ ├── edit.py # File editing (fuzzy match, indentation)
|
||||
│ ├── safety.py # Path validation, sandbox checks
|
||||
│ ├── sandbox.py # Command sandboxing
|
||||
│ └── web.py # Web fetch/search helpers
|
||||
├── mq/ # Message queue integration
|
||||
│ ├── protocol.py # Typed message dataclasses (JSON serialization)
|
||||
│ ├── broker.py # Abstract MessageBroker + RedisBroker
|
||||
│ ├── bridge.py # Bridge service (queue ↔ HTTP API, multi-node routing)
|
||||
│ └── client.py # TurnstoneClient — Python API for external systems
|
||||
├── console/ # Cluster dashboard
|
||||
│ ├── collector.py # ClusterCollector — aggregates all nodes via Redis + HTTP
|
||||
│ ├── server.py # Dashboard Starlette/ASGI server + SSE
|
||||
│ └── static/ # Cluster dashboard web UI
|
||||
├── tools/ # Tool schemas (one JSON file per tool)
|
||||
├── ui/ # Frontend assets and terminal rendering
|
||||
│ └── static/ # Web UI (HTML, CSS, JS)
|
||||
├── sim/ # Cluster simulator
|
||||
│ ├── cluster.py # SimCluster — orchestrates N nodes + dispatchers
|
||||
│ ├── node.py # SimNode + SimWorkstream — protocol-compatible node
|
||||
│ ├── engine.py # LLM + tool execution simulation
|
||||
│ ├── scenario.py # 5 workload scenarios (steady, burst, node_failure, …)
|
||||
│ ├── metrics.py # Latency, throughput, utilization collection
|
||||
│ └── cli.py # CLI entry point (turnstone-sim)
|
||||
├── cli.py # Terminal frontend (+ /cluster commands for console)
|
||||
├── server.py # Web frontend (Starlette/ASGI + SSE)
|
||||
└── eval.py # Evaluation and prompt optimization harness
|
||||
├── api/ # OpenAPI spec generation (Pydantic v2 models)
|
||||
├── sdk/ # Client SDKs (sync + async, Python)
|
||||
docs/
|
||||
├── architecture.md # System architecture and threading model
|
||||
├── api-reference.md # Web server API and SSE event reference
|
||||
├── sdk.md # Client SDK reference (Python + TypeScript)
|
||||
├── console.md # Cluster dashboard service (turnstone-console)
|
||||
├── docker.md # Docker Compose deployment and configuration
|
||||
├── simulator.md # Cluster simulator usage and scenarios
|
||||
├── tools.md # Tool schemas, execution pipeline, approval flow
|
||||
├── eval.md # Evaluation harness internals
|
||||
└── diagrams/ # UML architecture diagrams (PlantUML sources + PNGs)
|
||||
└── png/ # Pre-rendered diagram images
|
||||
deploy/
|
||||
├── helm/turnstone/ # Helm chart for Kubernetes
|
||||
└── terraform/ # Terraform modules (AWS ECS/Fargate)
|
||||
```
|
||||
|
||||
### Architecture Diagrams
|
||||
### Diagrams
|
||||
|
||||
Detailed UML diagrams are available in [`docs/diagrams/`](docs/diagrams/):
|
||||
|
||||
@@ -217,12 +151,12 @@ Bridges BLPOP from their per-node queue (priority) then the shared queue. Direct
|
||||
|
||||
## Tools
|
||||
|
||||
14 built-in tools, 2 agent tools, plus external tools via MCP:
|
||||
15 built-in tools, 2 agent tools, plus external tools via MCP:
|
||||
|
||||
| Tool | Description | Auto-approved |
|
||||
|------|-------------|:---:|
|
||||
| `bash` | Execute shell commands | |
|
||||
| `read_file` | Read file contents | yes |
|
||||
| `read_file` | Read file contents (text or images with vision models) | yes |
|
||||
| `write_file` | Write/create files | |
|
||||
| `edit_file` | Fuzzy-match file editing | |
|
||||
| `search` | Search files by name/content | yes |
|
||||
@@ -233,13 +167,16 @@ Bridges BLPOP from their per-node queue (priority) then the shared queue. Direct
|
||||
| `remember` | Save persistent facts | yes |
|
||||
| `recall` | Search memories and history | yes |
|
||||
| `forget` | Remove a memory | yes |
|
||||
| `notify` | Send notifications to linked channels | yes |
|
||||
| `task` | Spawn autonomous sub-agent | |
|
||||
| `plan` | Explore codebase, write .plan.md | |
|
||||
| `mcp__*` | External tools from MCP servers | |
|
||||
|
||||
When the total tool count exceeds a configurable threshold (default 20), MCP tools are automatically deferred using native `defer_loading` on Anthropic and OpenAI APIs, or a transparent client-side BM25 search for local models. The LLM discovers deferred tools on demand via a `tool_search` capability — no configuration needed beyond `--tool-search auto` (the default).
|
||||
|
||||
### MCP Tool Servers
|
||||
|
||||
Turnstone supports the [Model Context Protocol](https://modelcontextprotocol.io/) (MCP) for connecting external tool servers. MCP tools are discovered at startup, converted to OpenAI function-calling format, and merged with built-in tools. Each MCP tool is prefixed with `mcp__{server}__{tool}` to avoid name collisions.
|
||||
Turnstone supports the [Model Context Protocol](https://modelcontextprotocol.io/) (MCP) for connecting external tool servers. MCP tools are discovered at startup, converted to OpenAI function-calling format, and merged with built-in tools. Each MCP tool is prefixed with `mcp__{server}__{tool}` to avoid name collisions. Tool lists stay fresh via push notifications (`tools.listChanged`), periodic polling for servers without push, and manual `/mcp refresh`.
|
||||
|
||||
Configure via `config.toml` or `--mcp-config`:
|
||||
|
||||
@@ -259,7 +196,7 @@ turnstone --mcp-config ~/.config/turnstone/mcp.json
|
||||
turnstone-server --mcp-config ~/.config/turnstone/mcp.json
|
||||
```
|
||||
|
||||
Use `/mcp` in the REPL to list connected tools. MCP tools require user approval by default (overridden by `--skip-permissions` or UI auto-approve).
|
||||
Use `/mcp` in the REPL to list connected tools, `/mcp refresh` to re-fetch tool lists from servers. MCP tools require user approval by default (overridden by `--skip-permissions` or UI auto-approve).
|
||||
|
||||
### Multi-Model and Multi-Provider Support
|
||||
|
||||
@@ -314,6 +251,9 @@ agent_model = "" # model alias for plan/task sub-agents
|
||||
[tools]
|
||||
timeout = 30
|
||||
skip_permissions = false
|
||||
search = "auto" # "auto" (enable when >threshold tools), "on", "off"
|
||||
search_threshold = 20 # min tools before tool search activates
|
||||
search_max_results = 5 # max tools returned per search query
|
||||
|
||||
[server]
|
||||
host = "0.0.0.0"
|
||||
@@ -354,6 +294,7 @@ path = ".turnstone.db" # SQLite file path (relative to working directory)
|
||||
|
||||
[mcp]
|
||||
config_path = "" # path to MCP JSON config file (alternative to TOML sections)
|
||||
refresh_interval = 14400 # periodic refresh for servers without push notifications (seconds, 0 to disable)
|
||||
|
||||
[mcp.servers.example] # one section per MCP server
|
||||
command = "npx"
|
||||
@@ -412,6 +353,7 @@ Per-workstream metrics are labeled by `ws_id` (bounded to 10 max workstreams).
|
||||
- Redis (for message queue bridge — `pip install turnstone[mq]`)
|
||||
- Anthropic provider (optional — `pip install turnstone[anthropic]`)
|
||||
- PostgreSQL (optional, for production — `pip install turnstone[postgres]`)
|
||||
- [Git LFS](https://git-lfs.com/) (for cloning — diagram PNGs are stored in LFS)
|
||||
|
||||
## License
|
||||
|
||||
|
||||
+262
-1
@@ -5,8 +5,8 @@
|
||||
# Default (SQLite): docker compose up
|
||||
# Production (PG): DB_BACKEND=postgresql docker compose --profile production up
|
||||
# (or set DB_BACKEND=postgresql in .env)
|
||||
# 10-node cluster: docker compose --profile cluster up
|
||||
# With simulator: docker compose --profile sim up
|
||||
# Scale bridges: docker compose up --scale bridge=3
|
||||
# =============================================================================
|
||||
|
||||
name: turnstone
|
||||
@@ -28,6 +28,7 @@ services:
|
||||
image: postgres:17-alpine
|
||||
profiles:
|
||||
- production
|
||||
- cluster
|
||||
environment:
|
||||
POSTGRES_DB: turnstone
|
||||
POSTGRES_USER: ${POSTGRES_USER:-turnstone}
|
||||
@@ -109,9 +110,11 @@ services:
|
||||
- SKIP_PERMISSIONS=${SKIP_PERMISSIONS:-}
|
||||
- TURNSTONE_AUTH_ENABLED=${TURNSTONE_AUTH_ENABLED:-}
|
||||
- TURNSTONE_AUTH_TOKEN=${TURNSTONE_AUTH_TOKEN:-}
|
||||
- TURNSTONE_JWT_SECRET=${TURNSTONE_JWT_SECRET:-}
|
||||
- MODEL=${MODEL:-}
|
||||
- TURNSTONE_DB_BACKEND=${DB_BACKEND:-sqlite}
|
||||
- TURNSTONE_DB_URL=${DATABASE_URL:-}
|
||||
- TURNSTONE_NODE_ID=${TURNSTONE_NODE_ID:-}
|
||||
extra_hosts:
|
||||
- "host.docker.internal:host-gateway"
|
||||
networks:
|
||||
@@ -148,6 +151,7 @@ services:
|
||||
environment:
|
||||
- REDIS_PASSWORD=${REDIS_PASSWORD:-}
|
||||
- TURNSTONE_AUTH_TOKEN=${TURNSTONE_AUTH_TOKEN:-}
|
||||
- TURNSTONE_JWT_SECRET=${TURNSTONE_JWT_SECRET:-}
|
||||
networks:
|
||||
- turnstone-net
|
||||
depends_on:
|
||||
@@ -177,6 +181,9 @@ services:
|
||||
- REDIS_PASSWORD=${REDIS_PASSWORD:-}
|
||||
- TURNSTONE_AUTH_ENABLED=${TURNSTONE_AUTH_ENABLED:-}
|
||||
- TURNSTONE_AUTH_TOKEN=${TURNSTONE_AUTH_TOKEN:-}
|
||||
- TURNSTONE_JWT_SECRET=${TURNSTONE_JWT_SECRET:-}
|
||||
- TURNSTONE_DB_BACKEND=${DB_BACKEND:-sqlite}
|
||||
- TURNSTONE_DB_URL=${DATABASE_URL:-}
|
||||
networks:
|
||||
- turnstone-net
|
||||
depends_on:
|
||||
@@ -190,6 +197,44 @@ services:
|
||||
start_period: 10s
|
||||
restart: unless-stopped
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# turnstone-channel — Channel gateway (Discord, Slack, etc.)
|
||||
# Requires TURNSTONE_DISCORD_TOKEN to enable Discord adapter
|
||||
# -------------------------------------------------------------------
|
||||
channel:
|
||||
build:
|
||||
context: .
|
||||
dockerfile: Dockerfile
|
||||
profiles:
|
||||
- production
|
||||
- cluster
|
||||
command:
|
||||
- sh
|
||||
- -c
|
||||
- >-
|
||||
turnstone-channel
|
||||
--redis-host=redis
|
||||
--redis-port=6379
|
||||
--http-host=0.0.0.0
|
||||
$${TURNSTONE_DISCORD_GUILD:+--discord-guild $$TURNSTONE_DISCORD_GUILD}
|
||||
environment:
|
||||
- TURNSTONE_DISCORD_TOKEN=${TURNSTONE_DISCORD_TOKEN:-}
|
||||
- TURNSTONE_DISCORD_GUILD=${TURNSTONE_DISCORD_GUILD:-0}
|
||||
- REDIS_PASSWORD=${REDIS_PASSWORD:-}
|
||||
- TURNSTONE_JWT_SECRET=${TURNSTONE_JWT_SECRET:-}
|
||||
- TURNSTONE_DB_BACKEND=${DB_BACKEND:-sqlite}
|
||||
- TURNSTONE_DB_URL=${DATABASE_URL:-}
|
||||
- TURNSTONE_CHANNEL_ADVERTISE_URL=http://channel:8091
|
||||
networks:
|
||||
- turnstone-net
|
||||
depends_on:
|
||||
redis:
|
||||
condition: service_healthy
|
||||
postgres:
|
||||
condition: service_healthy
|
||||
required: false
|
||||
restart: unless-stopped
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# turnstone-sim — Multi-node cluster simulator (no LLM needed)
|
||||
# Start with: docker compose --profile sim up
|
||||
@@ -229,3 +274,219 @@ services:
|
||||
redis:
|
||||
condition: service_healthy
|
||||
restart: "no"
|
||||
|
||||
# ===================================================================
|
||||
# 10-node cluster (profile: cluster)
|
||||
#
|
||||
# Each node is a server + bridge pair. All share the same PostgreSQL
|
||||
# and Redis instances. Access via console at :8090.
|
||||
#
|
||||
# Start: docker compose --profile cluster up
|
||||
# ===================================================================
|
||||
|
||||
# -- cluster servers ------------------------------------------------
|
||||
|
||||
server-1: &cluster-server
|
||||
build: { context: ., dockerfile: Dockerfile }
|
||||
profiles: [cluster]
|
||||
command: &cluster-server-cmd
|
||||
- sh
|
||||
- -c
|
||||
- >-
|
||||
turnstone-server
|
||||
--host 0.0.0.0
|
||||
--port 8080
|
||||
--base-url "$${LLM_BASE_URL}"
|
||||
--api-key "$${OPENAI_API_KEY}"
|
||||
$${MODEL:+--model $$MODEL}
|
||||
$${SKIP_PERMISSIONS:+--skip-permissions}
|
||||
volumes: [turnstone-data:/data]
|
||||
environment: &cluster-server-env
|
||||
LLM_BASE_URL: ${LLM_BASE_URL:-http://host.docker.internal:8000/v1}
|
||||
OPENAI_API_KEY: ${OPENAI_API_KEY:-dummy}
|
||||
TAVILY_API_KEY: ${TAVILY_API_KEY:-}
|
||||
SKIP_PERMISSIONS: ${SKIP_PERMISSIONS:-}
|
||||
TURNSTONE_AUTH_ENABLED: ${TURNSTONE_AUTH_ENABLED:-}
|
||||
TURNSTONE_AUTH_TOKEN: ${TURNSTONE_AUTH_TOKEN:-}
|
||||
TURNSTONE_JWT_SECRET: ${TURNSTONE_JWT_SECRET:-}
|
||||
MODEL: ${MODEL:-}
|
||||
TURNSTONE_DB_BACKEND: ${DB_BACKEND:-postgresql}
|
||||
TURNSTONE_DB_URL: ${DATABASE_URL:-postgresql://${POSTGRES_USER:-turnstone}:${POSTGRES_PASSWORD:?}@postgres:5432/turnstone}
|
||||
TURNSTONE_NODE_ID: node-1
|
||||
extra_hosts: ["host.docker.internal:host-gateway"]
|
||||
networks: [turnstone-net]
|
||||
depends_on:
|
||||
redis: { condition: service_healthy }
|
||||
postgres: { condition: service_healthy }
|
||||
healthcheck:
|
||||
test: ["CMD", "python", "/usr/local/bin/healthcheck.py", "http://127.0.0.1:8080/health"]
|
||||
interval: 10s
|
||||
timeout: 5s
|
||||
retries: 3
|
||||
start_period: 15s
|
||||
deploy:
|
||||
resources:
|
||||
limits: { memory: 384M, cpus: '0.5' }
|
||||
restart: unless-stopped
|
||||
|
||||
server-2:
|
||||
<<: *cluster-server
|
||||
environment: { <<: *cluster-server-env, TURNSTONE_NODE_ID: node-2 }
|
||||
server-3:
|
||||
<<: *cluster-server
|
||||
environment: { <<: *cluster-server-env, TURNSTONE_NODE_ID: node-3 }
|
||||
server-4:
|
||||
<<: *cluster-server
|
||||
environment: { <<: *cluster-server-env, TURNSTONE_NODE_ID: node-4 }
|
||||
server-5:
|
||||
<<: *cluster-server
|
||||
environment: { <<: *cluster-server-env, TURNSTONE_NODE_ID: node-5 }
|
||||
server-6:
|
||||
<<: *cluster-server
|
||||
environment: { <<: *cluster-server-env, TURNSTONE_NODE_ID: node-6 }
|
||||
server-7:
|
||||
<<: *cluster-server
|
||||
environment: { <<: *cluster-server-env, TURNSTONE_NODE_ID: node-7 }
|
||||
server-8:
|
||||
<<: *cluster-server
|
||||
environment: { <<: *cluster-server-env, TURNSTONE_NODE_ID: node-8 }
|
||||
server-9:
|
||||
<<: *cluster-server
|
||||
environment: { <<: *cluster-server-env, TURNSTONE_NODE_ID: node-9 }
|
||||
server-10:
|
||||
<<: *cluster-server
|
||||
environment: { <<: *cluster-server-env, TURNSTONE_NODE_ID: node-10 }
|
||||
|
||||
# -- cluster bridges ------------------------------------------------
|
||||
|
||||
bridge-1: &cluster-bridge
|
||||
build: { context: ., dockerfile: Dockerfile }
|
||||
profiles: [cluster]
|
||||
command:
|
||||
- turnstone-bridge
|
||||
- --server-url=http://server-1:8080
|
||||
- --redis-host=redis
|
||||
- --redis-port=6379
|
||||
- --heartbeat-ttl=${HEARTBEAT_TTL:-60}
|
||||
- --approval-timeout=${APPROVAL_TIMEOUT:-3600}
|
||||
environment: &cluster-bridge-env
|
||||
REDIS_PASSWORD: ${REDIS_PASSWORD:-}
|
||||
TURNSTONE_AUTH_TOKEN: ${TURNSTONE_AUTH_TOKEN:-}
|
||||
TURNSTONE_JWT_SECRET: ${TURNSTONE_JWT_SECRET:-}
|
||||
networks: [turnstone-net]
|
||||
depends_on:
|
||||
server-1: { condition: service_healthy }
|
||||
redis: { condition: service_healthy }
|
||||
deploy:
|
||||
resources:
|
||||
limits: { memory: 256M, cpus: '0.25' }
|
||||
restart: unless-stopped
|
||||
|
||||
bridge-2:
|
||||
<<: *cluster-bridge
|
||||
command:
|
||||
- turnstone-bridge
|
||||
- --server-url=http://server-2:8080
|
||||
- --redis-host=redis
|
||||
- --redis-port=6379
|
||||
- --heartbeat-ttl=${HEARTBEAT_TTL:-60}
|
||||
- --approval-timeout=${APPROVAL_TIMEOUT:-3600}
|
||||
depends_on:
|
||||
server-2: { condition: service_healthy }
|
||||
redis: { condition: service_healthy }
|
||||
bridge-3:
|
||||
<<: *cluster-bridge
|
||||
command:
|
||||
- turnstone-bridge
|
||||
- --server-url=http://server-3:8080
|
||||
- --redis-host=redis
|
||||
- --redis-port=6379
|
||||
- --heartbeat-ttl=${HEARTBEAT_TTL:-60}
|
||||
- --approval-timeout=${APPROVAL_TIMEOUT:-3600}
|
||||
depends_on:
|
||||
server-3: { condition: service_healthy }
|
||||
redis: { condition: service_healthy }
|
||||
bridge-4:
|
||||
<<: *cluster-bridge
|
||||
command:
|
||||
- turnstone-bridge
|
||||
- --server-url=http://server-4:8080
|
||||
- --redis-host=redis
|
||||
- --redis-port=6379
|
||||
- --heartbeat-ttl=${HEARTBEAT_TTL:-60}
|
||||
- --approval-timeout=${APPROVAL_TIMEOUT:-3600}
|
||||
depends_on:
|
||||
server-4: { condition: service_healthy }
|
||||
redis: { condition: service_healthy }
|
||||
bridge-5:
|
||||
<<: *cluster-bridge
|
||||
command:
|
||||
- turnstone-bridge
|
||||
- --server-url=http://server-5:8080
|
||||
- --redis-host=redis
|
||||
- --redis-port=6379
|
||||
- --heartbeat-ttl=${HEARTBEAT_TTL:-60}
|
||||
- --approval-timeout=${APPROVAL_TIMEOUT:-3600}
|
||||
depends_on:
|
||||
server-5: { condition: service_healthy }
|
||||
redis: { condition: service_healthy }
|
||||
bridge-6:
|
||||
<<: *cluster-bridge
|
||||
command:
|
||||
- turnstone-bridge
|
||||
- --server-url=http://server-6:8080
|
||||
- --redis-host=redis
|
||||
- --redis-port=6379
|
||||
- --heartbeat-ttl=${HEARTBEAT_TTL:-60}
|
||||
- --approval-timeout=${APPROVAL_TIMEOUT:-3600}
|
||||
depends_on:
|
||||
server-6: { condition: service_healthy }
|
||||
redis: { condition: service_healthy }
|
||||
bridge-7:
|
||||
<<: *cluster-bridge
|
||||
command:
|
||||
- turnstone-bridge
|
||||
- --server-url=http://server-7:8080
|
||||
- --redis-host=redis
|
||||
- --redis-port=6379
|
||||
- --heartbeat-ttl=${HEARTBEAT_TTL:-60}
|
||||
- --approval-timeout=${APPROVAL_TIMEOUT:-3600}
|
||||
depends_on:
|
||||
server-7: { condition: service_healthy }
|
||||
redis: { condition: service_healthy }
|
||||
bridge-8:
|
||||
<<: *cluster-bridge
|
||||
command:
|
||||
- turnstone-bridge
|
||||
- --server-url=http://server-8:8080
|
||||
- --redis-host=redis
|
||||
- --redis-port=6379
|
||||
- --heartbeat-ttl=${HEARTBEAT_TTL:-60}
|
||||
- --approval-timeout=${APPROVAL_TIMEOUT:-3600}
|
||||
depends_on:
|
||||
server-8: { condition: service_healthy }
|
||||
redis: { condition: service_healthy }
|
||||
bridge-9:
|
||||
<<: *cluster-bridge
|
||||
command:
|
||||
- turnstone-bridge
|
||||
- --server-url=http://server-9:8080
|
||||
- --redis-host=redis
|
||||
- --redis-port=6379
|
||||
- --heartbeat-ttl=${HEARTBEAT_TTL:-60}
|
||||
- --approval-timeout=${APPROVAL_TIMEOUT:-3600}
|
||||
depends_on:
|
||||
server-9: { condition: service_healthy }
|
||||
redis: { condition: service_healthy }
|
||||
bridge-10:
|
||||
<<: *cluster-bridge
|
||||
command:
|
||||
- turnstone-bridge
|
||||
- --server-url=http://server-10:8080
|
||||
- --redis-host=redis
|
||||
- --redis-port=6379
|
||||
- --heartbeat-ttl=${HEARTBEAT_TTL:-60}
|
||||
- --approval-timeout=${APPROVAL_TIMEOUT:-3600}
|
||||
depends_on:
|
||||
server-10: { condition: service_healthy }
|
||||
redis: { condition: service_healthy }
|
||||
|
||||
@@ -1,221 +0,0 @@
|
||||
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 860 520" font-family="ui-monospace,SFMono-Regular,Menlo,Monaco,Consolas,monospace" font-size="13">
|
||||
<style>
|
||||
@keyframes pulse-green { 0%,100% { opacity:0.5 } 50% { opacity:1 } }
|
||||
@keyframes pulse-yellow { 0%,100% { opacity:0.4 } 50% { opacity:1 } }
|
||||
@keyframes pulse-blue { 0%,100% { opacity:0.3 } 50% { opacity:1 } }
|
||||
@keyframes fadein { from { opacity:0 } to { opacity:1 } }
|
||||
.pg { animation: pulse-green 2s infinite }
|
||||
.py { animation: pulse-yellow 1.8s infinite }
|
||||
.pb { animation: pulse-blue 2.2s infinite }
|
||||
.f1 { animation: fadein 0.4s 0.2s both }
|
||||
.f2 { animation: fadein 0.4s 0.4s both }
|
||||
.f3 { animation: fadein 0.4s 0.6s both }
|
||||
.f4 { animation: fadein 0.4s 0.8s both }
|
||||
.f5 { animation: fadein 0.4s 1.0s both }
|
||||
.f6 { animation: fadein 0.4s 1.3s both }
|
||||
.f7 { animation: fadein 0.4s 1.5s both }
|
||||
.f8 { animation: fadein 0.4s 1.7s both }
|
||||
.f9 { animation: fadein 0.4s 1.9s both }
|
||||
.f10 { animation: fadein 0.4s 2.1s both }
|
||||
.f11 { animation: fadein 0.4s 2.3s both }
|
||||
.f12 { animation: fadein 0.4s 2.5s both }
|
||||
</style>
|
||||
|
||||
<!-- Window chrome -->
|
||||
<rect rx="10" width="860" height="520" fill="#1a1b26"/>
|
||||
<rect width="860" height="36" rx="10" fill="#16161e"/>
|
||||
<rect y="26" width="860" height="10" fill="#16161e"/>
|
||||
<circle cx="20" cy="18" r="6" fill="#f7768e"/>
|
||||
<circle cx="40" cy="18" r="6" fill="#e0af68"/>
|
||||
<circle cx="60" cy="18" r="6" fill="#9ece6a"/>
|
||||
<text x="430" y="22" text-anchor="middle" fill="#565f89" font-size="12">turnstone — console</text>
|
||||
|
||||
<!-- Header -->
|
||||
<rect y="36" width="860" height="30" fill="#24283b"/>
|
||||
<rect y="66" width="860" height="1" fill="#3b4261"/>
|
||||
<text x="16" y="56" fill="#7aa2f7" font-size="14" font-weight="bold">turnstone console</text>
|
||||
<text x="200" y="56" fill="#565f89" font-size="12">6 nodes · 10 workstreams</text>
|
||||
|
||||
<!-- ====== State cards ====== -->
|
||||
<g transform="translate(16, 78)" class="f1" opacity="0">
|
||||
<!-- RUN card -->
|
||||
<rect x="0" y="0" width="156" height="64" rx="6" fill="#24283b" stroke="#3b4261"/>
|
||||
<rect x="0" y="0" width="156" height="3" rx="6" fill="#9ece6a"/>
|
||||
<text x="78" y="30" text-anchor="middle" fill="#a9b1d6" font-size="22" font-weight="bold">3</text>
|
||||
<text x="78" y="50" text-anchor="middle" fill="#565f89" font-size="10">▸ RUN</text>
|
||||
|
||||
<!-- THINK card -->
|
||||
<rect x="168" y="0" width="156" height="64" rx="6" fill="#24283b" stroke="#3b4261"/>
|
||||
<rect x="168" y="0" width="156" height="3" rx="6" fill="#7aa2f7"/>
|
||||
<text x="246" y="30" text-anchor="middle" fill="#a9b1d6" font-size="22" font-weight="bold">2</text>
|
||||
<text x="246" y="50" text-anchor="middle" fill="#565f89" font-size="10">◌ THINK</text>
|
||||
|
||||
<!-- ATTN card -->
|
||||
<rect x="336" y="0" width="156" height="64" rx="6" fill="#24283b" stroke="#3b4261"/>
|
||||
<rect x="336" y="0" width="156" height="3" rx="6" fill="#e0af68"/>
|
||||
<text x="414" y="30" text-anchor="middle" fill="#a9b1d6" font-size="22" font-weight="bold">1</text>
|
||||
<text x="414" y="50" text-anchor="middle" fill="#565f89" font-size="10">◆ ATTN</text>
|
||||
|
||||
<!-- ERR card -->
|
||||
<rect x="504" y="0" width="156" height="64" rx="6" fill="#24283b" stroke="#3b4261"/>
|
||||
<rect x="504" y="0" width="156" height="3" rx="6" fill="#f7768e"/>
|
||||
<text x="582" y="30" text-anchor="middle" fill="#a9b1d6" font-size="22" font-weight="bold">0</text>
|
||||
<text x="582" y="50" text-anchor="middle" fill="#565f89" font-size="10">✖ ERR</text>
|
||||
|
||||
<!-- IDLE card -->
|
||||
<rect x="672" y="0" width="156" height="64" rx="6" fill="#24283b" stroke="#3b4261"/>
|
||||
<rect x="672" y="0" width="156" height="3" rx="6" fill="#565f89"/>
|
||||
<text x="750" y="30" text-anchor="middle" fill="#a9b1d6" font-size="22" font-weight="bold">4</text>
|
||||
<text x="750" y="50" text-anchor="middle" fill="#565f89" font-size="10">· IDLE</text>
|
||||
</g>
|
||||
|
||||
<!-- Aggregate bar -->
|
||||
<text x="16" y="160" fill="#565f89" font-size="11" class="f2" opacity="0">197k tokens · 42 tool calls</text>
|
||||
|
||||
<!-- ====== NODES section ====== -->
|
||||
<text x="16" y="182" fill="#7aa2f7" font-size="12" font-weight="bold" class="f3" opacity="0">NODES</text>
|
||||
|
||||
<!-- Node column headers -->
|
||||
<g transform="translate(0, 190)" class="f4" opacity="0">
|
||||
<rect width="860" height="20" fill="#24283b"/>
|
||||
<rect y="20" width="860" height="1" fill="#3b4261"/>
|
||||
<text y="14" fill="#565f89" font-size="10" letter-spacing="0.5">
|
||||
<tspan x="36">NODE</tspan>
|
||||
<tspan x="560">WS</tspan>
|
||||
<tspan x="610">RUN</tspan>
|
||||
<tspan x="660">ATTN</tspan>
|
||||
<tspan x="710">TOKENS</tspan>
|
||||
<tspan x="790">LOAD</tspan>
|
||||
</text>
|
||||
</g>
|
||||
|
||||
<!-- Node rows -->
|
||||
<g transform="translate(0, 214)">
|
||||
|
||||
<!-- Node 1: db-west-04 — 3 ws, 1 running, has-running bar -->
|
||||
<g class="f5" opacity="0">
|
||||
<rect y="0" width="860" height="38" fill="#1a1b26"/>
|
||||
<rect y="0" width="3" height="38" fill="#9ece6a"/>
|
||||
<circle cx="22" cy="19" r="4" fill="#9ece6a"/>
|
||||
<text x="36" y="23" fill="#a9b1d6" font-size="12" font-weight="bold">db-west-04</text>
|
||||
<text x="566" y="23" fill="#a9b1d6" font-size="11">3</text>
|
||||
<text x="616" y="23" fill="#a9b1d6" font-size="11">1</text>
|
||||
<text x="666" y="23" fill="#565f89" font-size="11">0</text>
|
||||
<text x="710" y="23" fill="#565f89" font-size="11">57.6k</text>
|
||||
<!-- Load bar: 3/10 = 30% -->
|
||||
<rect x="770" y="15" width="60" height="6" rx="3" fill="#292e42"/>
|
||||
<rect x="770" y="15" width="18" height="6" rx="3" fill="#9ece6a"/>
|
||||
<text x="838" y="23" fill="#565f89" font-size="11">30%</text>
|
||||
</g>
|
||||
|
||||
<!-- Node 2: api-east-01 — 3 ws, 1 attention, has-attention bar -->
|
||||
<g class="f6" opacity="0">
|
||||
<rect y="40" width="860" height="38" fill="#24283b"/>
|
||||
<rect y="40" width="3" height="38" fill="#e0af68"/>
|
||||
<circle cx="22" cy="59" r="4" fill="#9ece6a"/>
|
||||
<text x="36" y="63" fill="#a9b1d6" font-size="12" font-weight="bold">api-east-01</text>
|
||||
<text x="566" y="63" fill="#a9b1d6" font-size="11">3</text>
|
||||
<text x="616" y="63" fill="#565f89" font-size="11">0</text>
|
||||
<text x="666" y="63" fill="#a9b1d6" font-size="11">1</text>
|
||||
<text x="710" y="63" fill="#565f89" font-size="11">109k</text>
|
||||
<!-- Load bar: 3/10 = 30% -->
|
||||
<rect x="770" y="55" width="60" height="6" rx="3" fill="#292e42"/>
|
||||
<rect x="770" y="55" width="18" height="6" rx="3" fill="#9ece6a"/>
|
||||
<text x="838" y="63" fill="#565f89" font-size="11">30%</text>
|
||||
</g>
|
||||
|
||||
<!-- Node 3: sre-node-03 — 2 ws, 1 running, has-running bar -->
|
||||
<g class="f7" opacity="0">
|
||||
<rect y="80" width="860" height="38" fill="#1a1b26"/>
|
||||
<rect y="80" width="3" height="38" fill="#9ece6a"/>
|
||||
<circle cx="22" cy="99" r="4" fill="#9ece6a"/>
|
||||
<text x="36" y="103" fill="#a9b1d6" font-size="12" font-weight="bold">sre-node-03</text>
|
||||
<text x="566" y="103" fill="#a9b1d6" font-size="11">2</text>
|
||||
<text x="616" y="103" fill="#a9b1d6" font-size="11">1</text>
|
||||
<text x="666" y="103" fill="#565f89" font-size="11">0</text>
|
||||
<text x="710" y="103" fill="#565f89" font-size="11">64.4k</text>
|
||||
<!-- Load bar: 2/10 = 20% -->
|
||||
<rect x="770" y="95" width="60" height="6" rx="3" fill="#292e42"/>
|
||||
<rect x="770" y="95" width="12" height="6" rx="3" fill="#9ece6a"/>
|
||||
<text x="838" y="103" fill="#565f89" font-size="11">20%</text>
|
||||
</g>
|
||||
|
||||
<!-- Node 4: analytics-02 — 1 ws, thinking, has-thinking bar -->
|
||||
<g class="f8" opacity="0">
|
||||
<rect y="120" width="860" height="38" fill="#24283b"/>
|
||||
<rect y="120" width="3" height="38" fill="#7aa2f7"/>
|
||||
<circle cx="22" cy="139" r="4" fill="#9ece6a"/>
|
||||
<text x="36" y="143" fill="#a9b1d6" font-size="12" font-weight="bold">analytics-02</text>
|
||||
<text x="566" y="143" fill="#a9b1d6" font-size="11">1</text>
|
||||
<text x="616" y="143" fill="#565f89" font-size="11">0</text>
|
||||
<text x="666" y="143" fill="#565f89" font-size="11">0</text>
|
||||
<text x="710" y="143" fill="#565f89" font-size="11">18.3k</text>
|
||||
<!-- Load bar: 1/10 = 10% -->
|
||||
<rect x="770" y="135" width="60" height="6" rx="3" fill="#292e42"/>
|
||||
<rect x="770" y="135" width="6" height="6" rx="3" fill="#9ece6a"/>
|
||||
<text x="838" y="143" fill="#565f89" font-size="11">10%</text>
|
||||
</g>
|
||||
|
||||
<!-- Node 5: data-ops-05 — 1 ws, thinking, has-thinking bar -->
|
||||
<g class="f9" opacity="0">
|
||||
<rect y="160" width="860" height="38" fill="#1a1b26"/>
|
||||
<rect y="160" width="3" height="38" fill="#7aa2f7"/>
|
||||
<circle cx="22" cy="179" r="4" fill="#9ece6a"/>
|
||||
<text x="36" y="183" fill="#a9b1d6" font-size="12" font-weight="bold">data-ops-05</text>
|
||||
<text x="566" y="183" fill="#a9b1d6" font-size="11">1</text>
|
||||
<text x="616" y="183" fill="#565f89" font-size="11">0</text>
|
||||
<text x="666" y="183" fill="#565f89" font-size="11">0</text>
|
||||
<text x="710" y="183" fill="#565f89" font-size="11">8.7k</text>
|
||||
<!-- Load bar: 1/10 = 10% -->
|
||||
<rect x="770" y="175" width="60" height="6" rx="3" fill="#292e42"/>
|
||||
<rect x="770" y="175" width="6" height="6" rx="3" fill="#9ece6a"/>
|
||||
<text x="838" y="183" fill="#565f89" font-size="11">10%</text>
|
||||
</g>
|
||||
|
||||
<!-- Node 6: ml-gpu-07 — 0 ws, empty, no bar -->
|
||||
<g class="f10" opacity="0">
|
||||
<rect y="200" width="860" height="38" fill="#24283b"/>
|
||||
<rect y="200" width="3" height="38" fill="transparent"/>
|
||||
<circle cx="22" cy="219" r="4" fill="#9ece6a"/>
|
||||
<text x="36" y="223" fill="#a9b1d6" font-size="12" font-weight="bold">ml-gpu-07</text>
|
||||
<text x="566" y="223" fill="#565f89" font-size="11">0</text>
|
||||
<text x="616" y="223" fill="#565f89" font-size="11">0</text>
|
||||
<text x="666" y="223" fill="#565f89" font-size="11">0</text>
|
||||
<text x="710" y="223" fill="#565f89" font-size="11">0</text>
|
||||
<!-- Load bar: 0/10 = 0% (empty track) -->
|
||||
<rect x="770" y="215" width="60" height="6" rx="3" fill="#292e42"/>
|
||||
<text x="842" y="223" fill="#565f89" font-size="11">0%</text>
|
||||
</g>
|
||||
|
||||
</g>
|
||||
|
||||
<!-- ====== Footer ====== -->
|
||||
<g transform="translate(0, 468)" class="f12" opacity="0">
|
||||
<rect width="860" height="1" fill="#3b4261"/>
|
||||
<rect y="1" width="860" height="24" fill="#16161e"/>
|
||||
|
||||
<circle cx="20" cy="14" r="3" fill="#9ece6a"/>
|
||||
<text x="28" y="18" fill="#565f89" font-size="10">db-west-04</text>
|
||||
|
||||
<circle cx="120" cy="14" r="3" fill="#9ece6a"/>
|
||||
<text x="128" y="18" fill="#565f89" font-size="10">api-east-01</text>
|
||||
|
||||
<circle cx="225" cy="14" r="3" fill="#9ece6a"/>
|
||||
<text x="233" y="18" fill="#565f89" font-size="10">sre-node-03</text>
|
||||
|
||||
<circle cx="335" cy="14" r="3" fill="#9ece6a"/>
|
||||
<text x="343" y="18" fill="#565f89" font-size="10">analytics-02</text>
|
||||
|
||||
<circle cx="450" cy="14" r="3" fill="#9ece6a"/>
|
||||
<text x="458" y="18" fill="#565f89" font-size="10">data-ops-05</text>
|
||||
|
||||
<circle cx="560" cy="14" r="3" fill="#9ece6a"/>
|
||||
<text x="568" y="18" fill="#565f89" font-size="10">ml-gpu-07</text>
|
||||
|
||||
<text x="680" y="18" fill="#3b4261" font-size="10">258k tokens · 42 calls · 12m</text>
|
||||
</g>
|
||||
|
||||
<!-- Bottom edge -->
|
||||
<rect y="493" width="860" height="27" fill="#16161e"/>
|
||||
<rect y="510" width="860" height="10" rx="10" fill="#16161e"/>
|
||||
</svg>
|
||||
|
Before Width: | Height: | Size: 11 KiB |
+190
-22
@@ -56,6 +56,170 @@ console.log(result.content);
|
||||
|
||||
---
|
||||
|
||||
## Authentication
|
||||
|
||||
When auth is enabled (`[auth].enabled = true` or `TURNSTONE_AUTH_ENABLED=1`), all API endpoints except public paths require a valid token.
|
||||
|
||||
### Sending Credentials
|
||||
|
||||
Include a token in one of two ways:
|
||||
|
||||
- **Bearer header**: `Authorization: Bearer <token>`
|
||||
- **Cookie**: `turnstone_auth=<token>` (set automatically by the login endpoint)
|
||||
|
||||
The server accepts three token types:
|
||||
|
||||
| Type | Format | Example |
|
||||
|------|--------|---------|
|
||||
| JWT | Base64 segments separated by dots | `eyJhbG...` |
|
||||
| API token | `ts_` prefix + 64 hex chars | `ts_a1b2c3d4...` |
|
||||
| Config token | Arbitrary string from `config.toml` | `my-secret-token` |
|
||||
|
||||
JWTs are the recommended credential for browser sessions. API tokens are suitable for programmatic access and CI/CD. Config tokens are a simple option for single-node deployments.
|
||||
|
||||
### `POST /v1/api/auth/login`
|
||||
|
||||
Authenticate with credentials and receive a JWT. Accepts two credential formats:
|
||||
|
||||
**Username + password:**
|
||||
|
||||
```json
|
||||
{"username": "alice", "password": "hunter2"}
|
||||
```
|
||||
|
||||
**API token:**
|
||||
|
||||
```json
|
||||
{"token": "ts_a1b2c3d4e5f6..."}
|
||||
```
|
||||
|
||||
**Response (success):** `200`
|
||||
|
||||
```json
|
||||
{
|
||||
"status": "ok",
|
||||
"role": "full",
|
||||
"scopes": "approve,read,write",
|
||||
"jwt": "eyJhbGciOiJIUzI1NiIs...",
|
||||
"user_id": "u_abc123"
|
||||
}
|
||||
```
|
||||
|
||||
The response also sets a `turnstone_auth` HttpOnly cookie containing the JWT.
|
||||
|
||||
**Response (failure):** `401`
|
||||
|
||||
```json
|
||||
{"error": "Invalid credentials"}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### `POST /v1/api/auth/logout`
|
||||
|
||||
Clears the `turnstone_auth` cookie. No request body required.
|
||||
|
||||
**Response:** `200`
|
||||
|
||||
```json
|
||||
{"status": "ok"}
|
||||
```
|
||||
|
||||
The response includes a `Set-Cookie` header that expires the auth cookie.
|
||||
|
||||
---
|
||||
|
||||
### `GET /v1/api/auth/status`
|
||||
|
||||
Returns the current authentication state. Works with or without a valid token.
|
||||
|
||||
**Response (authenticated):** `200`
|
||||
|
||||
```json
|
||||
{
|
||||
"authenticated": true,
|
||||
"user_id": "u_abc123",
|
||||
"scopes": ["approve", "read", "write"],
|
||||
"source": "jwt"
|
||||
}
|
||||
```
|
||||
|
||||
**Response (not authenticated):** `200`
|
||||
|
||||
```json
|
||||
{
|
||||
"authenticated": false,
|
||||
"user_id": null,
|
||||
"scopes": [],
|
||||
"source": null
|
||||
}
|
||||
```
|
||||
|
||||
**Response (auth disabled):** `200`
|
||||
|
||||
```json
|
||||
{
|
||||
"authenticated": false,
|
||||
"auth_enabled": false
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### `POST /v1/api/auth/setup`
|
||||
|
||||
Creates the first admin user when no users exist in the database. This is a
|
||||
public endpoint (no authentication required) that only succeeds when auth is
|
||||
enabled and the user database is empty. Both the server and console expose
|
||||
this endpoint.
|
||||
|
||||
**Request body:**
|
||||
|
||||
```json
|
||||
{
|
||||
"username": "admin",
|
||||
"display_name": "Admin",
|
||||
"password": "strongpass"
|
||||
}
|
||||
```
|
||||
|
||||
| Field | Type | Required | Validation |
|
||||
|----------------|--------|----------|-----------------------------|
|
||||
| `username` | string | yes | 1-64 ASCII characters |
|
||||
| `display_name` | string | yes | Non-empty |
|
||||
| `password` | string | yes | Minimum 8 characters |
|
||||
|
||||
**Response (success):** `200`
|
||||
|
||||
```json
|
||||
{
|
||||
"status": "ok",
|
||||
"user_id": "u_abc123",
|
||||
"username": "admin",
|
||||
"role": "full",
|
||||
"scopes": "approve,read,write",
|
||||
"jwt": "eyJhbGciOiJIUzI1NiIs..."
|
||||
}
|
||||
```
|
||||
|
||||
The response also sets a `turnstone_auth` HttpOnly cookie containing the JWT.
|
||||
|
||||
**Response (already set up):** `409`
|
||||
|
||||
```json
|
||||
{"error": "Setup already completed"}
|
||||
```
|
||||
|
||||
Returned when one or more users already exist in the database.
|
||||
|
||||
**Response (auth disabled):** `400`
|
||||
|
||||
```json
|
||||
{"error": "Auth is not enabled"}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Endpoints
|
||||
|
||||
### `GET /`
|
||||
@@ -351,8 +515,8 @@ Returns a list of all active workstreams.
|
||||
```json
|
||||
{
|
||||
"workstreams": [
|
||||
{"id": "abc123", "name": "default", "state": "idle", "session_id": "a1b2c3d4e5f6"},
|
||||
{"id": "def456", "name": "hacker-news", "state": "thinking", "session_id": "c5d6e7f8a9b0"}
|
||||
{"id": "abc123", "name": "default", "state": "idle"},
|
||||
{"id": "def456", "name": "hacker-news", "state": "thinking"}
|
||||
]
|
||||
}
|
||||
```
|
||||
@@ -364,22 +528,21 @@ Each workstream object:
|
||||
| `id` | string | Unique workstream routing identifier |
|
||||
| `name` | string | Display name (alias if set, otherwise `ws-xxxx`) |
|
||||
| `state` | string | Current state (see state values above) |
|
||||
| `session_id` | string/null | Session ID of the workstream's `ChatSession`, used for deduplication against `/v1/api/sessions` |
|
||||
|
||||
---
|
||||
|
||||
### `GET /v1/api/sessions`
|
||||
### `GET /v1/api/workstreams/saved`
|
||||
|
||||
Returns a list of saved sessions from the database, ordered by most recently
|
||||
Returns a list of saved workstreams from the database, ordered by most recently
|
||||
updated.
|
||||
|
||||
**Response:**
|
||||
|
||||
```json
|
||||
{
|
||||
"sessions": [
|
||||
"workstreams": [
|
||||
{
|
||||
"session_id": "a1b2c3d4e5f6",
|
||||
"ws_id": "a1b2c3d4e5f6",
|
||||
"alias": "refactor",
|
||||
"title": "JWT Authentication Refactor",
|
||||
"created": "2026-03-01 10:00:00",
|
||||
@@ -390,16 +553,16 @@ updated.
|
||||
}
|
||||
```
|
||||
|
||||
Each session object:
|
||||
Each saved workstream object:
|
||||
|
||||
| Field | Type | Description |
|
||||
|-----------------|-------------|--------------------------------------------|
|
||||
| `session_id` | string | Unique 12-char hex session identifier |
|
||||
| `ws_id` | string | Unique workstream identifier |
|
||||
| `alias` | string/null | User-assigned short name |
|
||||
| `title` | string/null | LLM-generated title |
|
||||
| `created` | string | ISO timestamp of session creation |
|
||||
| `created` | string | ISO timestamp of workstream creation |
|
||||
| `updated` | string | ISO timestamp of last message |
|
||||
| `message_count` | int | Number of messages in the session |
|
||||
| `message_count` | int | Number of messages in the workstream |
|
||||
|
||||
---
|
||||
|
||||
@@ -550,22 +713,25 @@ Creates a new workstream. The server supports up to 10 concurrent workstreams.
|
||||
|
||||
All fields are optional. The body can be empty or an empty JSON object.
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|----------------|--------|---------|------------------------------------------------|
|
||||
| `name` | string | auto | Workstream display name |
|
||||
| `model` | string | default | Model alias from the registry (`[models.*]`) |
|
||||
| `auto_approve` | bool | false | Auto-approve all tool calls for this workstream |
|
||||
| Field | Type | Default | Description |
|
||||
|------------------|--------|---------|----------------------------------------------------------------|
|
||||
| `name` | string | auto | Workstream display name |
|
||||
| `model` | string | default | Model alias from the registry (`[models.*]`) |
|
||||
| `auto_approve` | bool | false | Auto-approve all tool calls for this workstream |
|
||||
| `resume_ws` | string | "" | Workstream ID to resume atomically during creation (empty = fresh)|
|
||||
|
||||
**Response (success):**
|
||||
|
||||
```json
|
||||
{"ws_id": "ghi789", "name": "ws-3"}
|
||||
{"ws_id": "ghi789", "name": "ws-3", "resumed": false, "message_count": 0}
|
||||
```
|
||||
|
||||
| Field | Type | Description |
|
||||
|---------|--------|------------------------------------|
|
||||
| `ws_id` | string | Unique ID of the new workstream |
|
||||
| `name` | string | Auto-generated workstream name |
|
||||
| Field | Type | Description |
|
||||
|-----------------|--------|-----------------------------------------------------|
|
||||
| `ws_id` | string | Unique ID of the new workstream |
|
||||
| `name` | string | Auto-generated workstream name |
|
||||
| `resumed` | bool | Whether a previous session was successfully resumed |
|
||||
| `message_count` | int | Number of messages in the resumed session (0 if fresh) |
|
||||
|
||||
**Error (limit reached):**
|
||||
|
||||
@@ -692,7 +858,8 @@ liveness probes.
|
||||
```json
|
||||
{
|
||||
"status": "ok",
|
||||
"version": "0.3.0",
|
||||
"version": "0.4.0",
|
||||
"node_id": "worker-01_a3f2",
|
||||
"uptime_seconds": 3614.72,
|
||||
"model": "llama-3.1-70b-instruct",
|
||||
"workstreams": {
|
||||
@@ -714,6 +881,7 @@ liveness probes.
|
||||
|-------|------|-------------|
|
||||
| `status` | string | `"ok"` or `"degraded"` (degraded when backend unreachable) |
|
||||
| `version` | string | turnstone server version |
|
||||
| `node_id` | string | Server-generated node identity (`{hostname}_{4hex}`) |
|
||||
| `uptime_seconds` | number | Seconds since the server process started |
|
||||
| `model` | string | Model name detected or configured at startup |
|
||||
| `workstreams.total` | integer | Total active workstreams |
|
||||
|
||||
+248
-61
@@ -21,6 +21,8 @@ plugs in.
|
||||
| `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-admin` | `turnstone.core.admin_cli` | — | Offline user and API token management |
|
||||
|
||||
---
|
||||
|
||||
@@ -40,7 +42,8 @@ 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, async-sync bridge
|
||||
mcp_client.py MCPClientManager — MCP server connections, tool discovery, dynamic refresh, async-sync bridge
|
||||
tool_search.py Dynamic tool search — BM25 index, session-scoped tool visibility
|
||||
model_registry.py ModelRegistry — named model configs, lazy client creation, fallback routing
|
||||
memory.py Persistence facade (delegates to storage backend)
|
||||
storage/ Pluggable storage: StorageBackend protocol, SQLite + PostgreSQL
|
||||
@@ -73,8 +76,15 @@ turnstone/
|
||||
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
|
||||
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
|
||||
_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)
|
||||
ui/
|
||||
colors.py ANSI color constants with NO_COLOR support
|
||||
@@ -85,7 +95,7 @@ turnstone/
|
||||
style.css Page-specific UI styles (dashboard layout, approval blocks)
|
||||
app.js Page-specific client-side JavaScript (SSE, workstreams, markdown)
|
||||
tools/
|
||||
*.json 14 tool schemas (OpenAI function-calling format + turnstone metadata)
|
||||
*.json 15 tool schemas (OpenAI function-calling format + turnstone metadata)
|
||||
```
|
||||
|
||||
Both UIs share a common design system extracted into `turnstone/shared_static/`: design tokens, login overlay, toast notifications, theme toggle, keyboard shortcuts, and utility functions. Each UI imports `base.css` and the shared JS modules at `/shared/`, then adds only page-specific code at `/static/`.
|
||||
@@ -463,7 +473,7 @@ independently, then returns the final content as the tool result.
|
||||
|
||||
- **task**: uses `self._task_tools` (`TASK_AGENT_TOOLS` + MCP tools)
|
||||
- **plan**: uses `self._agent_tools` (`AGENT_TOOLS` + MCP tools). Writes output
|
||||
to `.plan-<session_id>.md` — unique per `ChatSession` so concurrent workstreams
|
||||
to `.plan-<ws_id>.md` — unique per `ChatSession` so concurrent workstreams
|
||||
don't collide. On repeat invocations the prior `plan` tool call and its result
|
||||
are forwarded from `self.messages` so the agent refines the existing plan rather
|
||||
than starting over. Planning instructions are injected as a developer message
|
||||
@@ -488,17 +498,32 @@ bridges this with a background asyncio event loop in a daemon thread.
|
||||
1. `create_mcp_client()` reads server configs from TOML or JSON
|
||||
2. `MCPClientManager.start()` launches the background event loop thread
|
||||
3. `_connect_all()` connects to each server (stdio subprocess or HTTP), runs
|
||||
`initialize()` + `list_tools()`, converts schemas to OpenAI format
|
||||
4. `ChatSession.__init__` receives the manager and builds `self._tools` (built-in + MCP)
|
||||
`initialize()` + `list_tools()`, converts schemas to OpenAI format, detects
|
||||
`tools.listChanged` capability for push notification support
|
||||
4. `ChatSession.__init__` receives the manager, builds `self._tools` (built-in + MCP),
|
||||
and registers a listener callback for tool-change notifications
|
||||
5. `_prepare_tool()` routes MCP tools to `_prepare_mcp_tool()` / `_exec_mcp_tool()`
|
||||
6. `_exec_mcp_tool()` calls `call_tool_sync()` which dispatches to the async loop
|
||||
via `asyncio.run_coroutine_threadsafe()`
|
||||
|
||||
**Tool refresh:** Three mechanisms keep tools up-to-date without restart:
|
||||
- **Push:** Servers declaring `tools.listChanged` send `ToolListChangedNotification`;
|
||||
the registered `message_handler` triggers immediate single-server refresh.
|
||||
- **Periodic:** Servers without push support are polled on a staggered interval
|
||||
(default 4 h, configurable via `[mcp] refresh_interval` or `--mcp-refresh-interval`).
|
||||
- **Manual:** `/mcp refresh [server]` calls `refresh_sync()` for on-demand refresh
|
||||
(also attempts reconnection for disconnected servers).
|
||||
|
||||
When tools change, `_rebuild_tools()` creates new `_tools`/`_tool_map` objects
|
||||
(copy-on-write for thread safety) and notifies listener callbacks. Each `ChatSession`
|
||||
rebuilds its merged tool lists and reconstructs `ToolSearchManager` (preserving
|
||||
expanded tools).
|
||||
|
||||
**Tool naming:** `mcp__{server}__{tool}` — double underscore delimiter, validated
|
||||
at connection time (server names with `__` are rejected).
|
||||
|
||||
**Error isolation:** Per-server connection failures are caught and logged; other
|
||||
servers still connect. Tool execution errors return error strings to the LLM
|
||||
**Error isolation:** Per-server connection/refresh failures are caught and logged; other
|
||||
servers are unaffected. Tool execution errors return error strings to the LLM
|
||||
rather than crashing the session.
|
||||
|
||||
### Provider Adapter Layer
|
||||
@@ -535,21 +560,23 @@ LLMProvider (protocol)
|
||||
|------|--------|
|
||||
| `StreamChunk` | `content_delta`, `reasoning_delta`, `tool_call_deltas`, `info_delta`, `usage`, `finish_reason` |
|
||||
| `CompletionResult` | `content`, `tool_calls`, `finish_reason`, `usage` |
|
||||
| `ModelCapabilities` | `context_window`, `max_output_tokens`, `supports_temperature`, `token_param`, `thinking_mode`, `supports_effort`, `supports_web_search` |
|
||||
| `ModelCapabilities` | `context_window`, `max_output_tokens`, `supports_temperature`, `token_param`, `thinking_mode`, `supports_effort`, `supports_web_search`, `supports_tool_search`, `supports_vision` |
|
||||
| `UsageInfo` | `prompt_tokens`, `completion_tokens`, `total_tokens` |
|
||||
|
||||
**OpenAIProvider** (`_openai.py`): passes messages through unchanged (they are
|
||||
already in OpenAI format). Model capability lookup table covers
|
||||
GPT-5/5.1/5.2, O-series, and search models (`gpt-5-search-api`).
|
||||
already in OpenAI format), including multi-part content blocks (text + images)
|
||||
in tool results. Model capability lookup table covers GPT-5/5.1/5.2/5.3/5.4,
|
||||
O-series, and search models (`gpt-5-search-api`) — all with `supports_vision`.
|
||||
For search models, injects `web_search_options` and removes the `web_search`
|
||||
function tool (the model always searches). Citations from `url_citation`
|
||||
annotations are formatted as footnotes. Unknown models (local servers) get
|
||||
permissive defaults and use Tavily for web search.
|
||||
permissive defaults with `supports_vision=False` and use Tavily for web search.
|
||||
|
||||
**AnthropicProvider** (`_anthropic.py`): converts OpenAI-format messages to
|
||||
Anthropic content blocks, maps `system`/`developer` roles to the `system`
|
||||
parameter, groups consecutive `tool` result messages into user-role content
|
||||
blocks, and translates tool schemas from OpenAI function-calling format to
|
||||
blocks (converting `image_url` parts to Anthropic's `image` source format),
|
||||
and translates tool schemas from OpenAI function-calling format to
|
||||
Anthropic's `input_schema` format. Supports both manual and adaptive thinking
|
||||
modes, with effort parameter support for models like Claude Opus 4.6 and
|
||||
Sonnet 4.6. Replaces the `web_search` function tool with Anthropic's native
|
||||
@@ -595,6 +622,18 @@ agent_model = "claude"
|
||||
|
||||
Each `[models.*]` entry produces a `ModelConfig` with a `provider` field
|
||||
(default: `"openai"`). Supported values: `"openai"` and `"anthropic"`.
|
||||
An optional `[models.*.capabilities]` sub-table overrides per-model
|
||||
`ModelCapabilities` flags (useful for local models whose capabilities
|
||||
cannot be detected programmatically):
|
||||
|
||||
```toml
|
||||
[models.qwen-vl]
|
||||
base_url = "http://localhost:8000/v1"
|
||||
model = "qwen-3.5-vl"
|
||||
|
||||
[models.qwen-vl.capabilities]
|
||||
supports_vision = true
|
||||
```
|
||||
|
||||
**Lifecycle:**
|
||||
1. `load_model_registry()` reads `[models.*]` sections from config.toml and
|
||||
@@ -679,8 +718,11 @@ memories
|
||||
created TEXT NOT NULL
|
||||
updated TEXT NOT NULL
|
||||
|
||||
sessions
|
||||
session_id TEXT PRIMARY KEY
|
||||
workstreams
|
||||
ws_id TEXT PRIMARY KEY
|
||||
node_id TEXT NOT NULL
|
||||
name TEXT NOT NULL
|
||||
state TEXT NOT NULL DEFAULT 'idle'
|
||||
alias TEXT UNIQUE -- user-assigned short name (nullable)
|
||||
title TEXT -- LLM-generated title (nullable)
|
||||
created TEXT NOT NULL
|
||||
@@ -688,7 +730,7 @@ sessions
|
||||
|
||||
conversations
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT
|
||||
session_id TEXT NOT NULL
|
||||
ws_id TEXT NOT NULL
|
||||
timestamp TEXT NOT NULL
|
||||
role TEXT NOT NULL -- user | assistant | tool_call | tool_result
|
||||
content TEXT
|
||||
@@ -697,8 +739,8 @@ conversations
|
||||
tool_call_id TEXT -- links tool_call ↔ tool_result for resume
|
||||
provider_data TEXT -- raw provider content (e.g. Anthropic encrypted)
|
||||
|
||||
session_config
|
||||
session_id TEXT NOT NULL -- composite PK with key
|
||||
workstream_config
|
||||
ws_id TEXT NOT NULL -- composite PK with key
|
||||
key TEXT NOT NULL
|
||||
value TEXT
|
||||
|
||||
@@ -713,18 +755,21 @@ and are the single source of truth for both backends and Alembic migrations.
|
||||
|
||||
| Method | Purpose |
|
||||
|--------|---------|
|
||||
| `register_session(session_id, title)` | Create a sessions row (no-op if exists) |
|
||||
| `save_message(session_id, role, content, ...)` | Log a message to conversations |
|
||||
| `load_session_messages(session_id)` | Reconstruct OpenAI message format from DB rows |
|
||||
| `list_sessions(limit)` | List sessions with >=1 message, ordered by updated DESC |
|
||||
| `delete_session(session_id)` | Delete session and all its messages |
|
||||
| `prune_sessions(retention_days)` | Remove empty sessions and old unnamed sessions |
|
||||
| `resolve_session(alias_or_id)` | Resolve alias, exact id, or id prefix to full session_id |
|
||||
| `save_session_config(session_id, config)` | Persist session configuration key/value pairs |
|
||||
| `load_session_config(session_id)` | Retrieve session configuration |
|
||||
| `set_session_alias(session_id, alias)` | Set user-friendly alias (returns False if taken) |
|
||||
| `get_session_name(session_id)` | Return alias if set, else title, else None |
|
||||
| `update_session_title(session_id, title)` | Set/update LLM-generated title |
|
||||
| `register_workstream(ws_id, node_id, name, state)` | Create a workstreams row (no-op if exists) |
|
||||
| `save_message(ws_id, role, content, ...)` | Log a message to conversations |
|
||||
| `load_messages(ws_id)` | Reconstruct OpenAI message format from DB rows |
|
||||
| `list_workstreams_with_history(limit)` | List workstreams with >=1 message, ordered by updated DESC |
|
||||
| `delete_workstream(ws_id)` | Delete workstream and cascade conversations + config |
|
||||
| `prune_workstreams(retention_days)` | Remove empty workstreams and old unnamed workstreams |
|
||||
| `resolve_workstream(alias_or_id)` | Resolve alias, exact id, or id prefix to full ws_id |
|
||||
| `save_workstream_config(ws_id, config)` | Persist workstream configuration key/value pairs |
|
||||
| `load_workstream_config(ws_id)` | Retrieve workstream configuration |
|
||||
| `set_workstream_alias(ws_id, alias)` | Set user-friendly alias (returns False if taken) |
|
||||
| `get_workstream_display_name(ws_id)` | Return alias if set, else title, else None |
|
||||
| `update_workstream_title(ws_id, title)` | Set/update LLM-generated title |
|
||||
| `update_workstream_state(ws_id, state)` | Update workstream state and bump timestamp |
|
||||
| `update_workstream_name(ws_id, name)` | Update workstream display name |
|
||||
| `list_workstreams(node_id, limit)` | List workstreams, optionally by node |
|
||||
| `kv_get(key)` / `kv_set(key, value)` / `kv_delete(key)` | Generic key-value store (backs memories table) |
|
||||
| `kv_list()` / `kv_search(query)` | List or search key-value pairs |
|
||||
| `search_history(query, limit)` | Full-text search (FTS5 on SQLite, tsvector on PostgreSQL) |
|
||||
@@ -743,57 +788,59 @@ pool_size = 5 # PostgreSQL connection pool size
|
||||
|
||||
Environment variables: `TURNSTONE_DB_BACKEND`, `TURNSTONE_DB_URL`, `TURNSTONE_DB_PATH`.
|
||||
|
||||
### Session Persistence and Resume
|
||||
### Persistence and Resume
|
||||
|
||||
Each `ChatSession` generates a 12-char hex `_session_id` on creation and
|
||||
registers it in the `sessions` table. Messages are saved to `conversations`
|
||||
as they happen via `save_message()`.
|
||||
`ws_id` is the sole persistent identity for both routing and conversation
|
||||
history. There is no separate `session_id` — the `workstreams` table holds
|
||||
alias, title, and state alongside the routing fields (`node_id`, `name`).
|
||||
Messages are saved to `conversations` (keyed by `ws_id`) as they happen
|
||||
via `save_message()`. Workstream state changes are tracked via
|
||||
`update_workstream_state()`.
|
||||
|
||||
**Auto-titling:** After the first complete exchange (user message + assistant
|
||||
response), a background thread calls the LLM with a title-generation prompt
|
||||
(`reasoning_effort: "low"`, `max_completion_tokens: 200`). The generated
|
||||
title (3-8 words) is stored in `sessions.title`.
|
||||
title (3-8 words) is stored in `workstreams.title`.
|
||||
|
||||
**Resume flow:** `ChatSession.resume_session(session_id)` calls
|
||||
`load_session_messages()` which reconstructs the OpenAI message format from
|
||||
database rows:
|
||||
**Resume flow:** `ChatSession.resume(ws_id)` calls `load_messages()` which
|
||||
reconstructs the OpenAI message format from database rows:
|
||||
|
||||
- `user` and `assistant` rows map directly
|
||||
- Consecutive `tool_call` rows are grouped into one assistant message's
|
||||
`tool_calls` array, paired with subsequent `tool_result` rows via
|
||||
`tool_call_id` (or positional matching for legacy data)
|
||||
- **Interrupted session repair:** If the last assistant message has
|
||||
`tool_calls` but fewer tool results than expected (session was
|
||||
- **Interrupted conversation repair:** If the last assistant message has
|
||||
`tool_calls` but fewer tool results than expected (conversation was
|
||||
interrupted mid-execution), the incomplete turn is stripped so the
|
||||
LLM can re-generate cleanly
|
||||
- The session adopts the old `_session_id`, so new messages continue in
|
||||
the same session
|
||||
- The `ChatSession` adopts the resumed `_ws_id`, so new messages continue
|
||||
in the same workstream
|
||||
|
||||
**Config persistence:** LLM-affecting parameters (`temperature`,
|
||||
`reasoning_effort`, `max_tokens`, `instructions`, `creative_mode`) are
|
||||
persisted to the `session_config` table on creation and whenever changed
|
||||
via slash commands. `resume_session()` restores these values so resumed
|
||||
sessions behave identically to the original.
|
||||
persisted to the `workstream_config` table on creation and whenever changed
|
||||
via slash commands. `resume()` restores these values so resumed workstreams
|
||||
behave identically to the original.
|
||||
|
||||
**`/clear` vs `/new`:** `/clear` wipes in-memory context but preserves
|
||||
messages in the database for future resume. `/new` starts a fresh session
|
||||
(new `_session_id`), leaving the old session resumable.
|
||||
messages in the database for future resume. `/new` starts a fresh workstream
|
||||
(new `_ws_id`), leaving the old workstream resumable.
|
||||
|
||||
**Resolution:** `resolve_session()` accepts aliases, exact session IDs, or
|
||||
session ID prefixes, enabling `turnstone --resume refactor` or `/resume abc12`.
|
||||
**Resolution:** `resolve_workstream()` accepts aliases, exact workstream IDs,
|
||||
or ID prefixes, enabling `turnstone --resume refactor` or `/resume abc12`.
|
||||
|
||||
**Session listing:** `list_sessions()` only returns sessions that have at
|
||||
least one saved message (`WHERE EXISTS` on `conversations`). Sessions
|
||||
registered but never used (e.g., from process startup) are invisible until
|
||||
a message is sent.
|
||||
**Workstream listing:** `list_workstreams_with_history()` only returns
|
||||
workstreams that have at least one saved message (`WHERE EXISTS` on
|
||||
`conversations`). Workstreams registered but never used (e.g., from process
|
||||
startup) are invisible until a message is sent.
|
||||
|
||||
**Session pruning:** `prune_sessions(retention_days, log_fn)` runs once at
|
||||
startup (CLI and server). It removes:
|
||||
- Sessions with no messages (orphaned registrations)
|
||||
- Unnamed sessions (`alias IS NULL`) older than `retention_days` days (default 90)
|
||||
**Workstream pruning:** `prune_workstreams(retention_days, log_fn)` runs once
|
||||
at startup (CLI and server). It removes:
|
||||
- Workstreams with no messages (orphaned registrations)
|
||||
- Unnamed workstreams (`alias IS NULL`) older than `retention_days` days (default 90)
|
||||
|
||||
Named (aliased) sessions are never age-pruned. Configure with
|
||||
`--session-retention-days N` (0 = disable age pruning).
|
||||
Named (aliased) workstreams are never age-pruned. Configure with
|
||||
`--retention-days N` (0 = disable age pruning).
|
||||
|
||||
---
|
||||
|
||||
@@ -910,6 +957,98 @@ limits using a token-bucket algorithm. Each IP gets a `TokenBucket` with
|
||||
|
||||
---
|
||||
|
||||
## User Identity and Authentication
|
||||
|
||||
Turnstone supports three authentication mechanisms, unified behind an
|
||||
`AuthResult` dataclass that carries `user_id`, `scopes`, and `token_source`:
|
||||
|
||||
1. **Config-file tokens** — static secrets in `config.toml` `[[auth.tokens]]`
|
||||
or the `TURNSTONE_AUTH_TOKEN` env var. Validated in-memory via
|
||||
`hmac.compare_digest`. Map to scopes through their role (`read` or `full`).
|
||||
2. **API tokens** — database-backed, prefixed `ts_`, stored as SHA-256 hashes
|
||||
in the `api_tokens` table. Can be exchanged for JWTs via
|
||||
`POST /v1/api/auth/login`.
|
||||
3. **JWTs** — short-lived HMAC-SHA256 session tokens (default 24h) issued after
|
||||
successful credential validation. Contain `sub` (user_id), `scopes`, and
|
||||
`src` (origin) in claims.
|
||||
|
||||
### Scope Model
|
||||
|
||||
Three hierarchical scopes control endpoint access:
|
||||
|
||||
| Scope | Grants | Endpoints |
|
||||
|-------|--------|-----------|
|
||||
| `read` | SSE streams, workstream listing, history | GET endpoints |
|
||||
| `write` | `read` + send, command, workstream create/close | POST to `/api/send`, `/api/command`, etc. |
|
||||
| `approve` | `write` + tool approval, admin operations | POST to `/api/approve`, `/api/admin/*` |
|
||||
|
||||
### Middleware Flow
|
||||
|
||||
`AuthMiddleware` (ASGI) intercepts every request:
|
||||
|
||||
1. **Public path check** — `/`, `/static/*`, `/shared/*`, `/health`,
|
||||
`/metrics`, `/openapi.json`, `/docs`, `/api/auth/*`, and `/api/auth/setup`
|
||||
are always allowed.
|
||||
2. **Token extraction** — `Authorization: Bearer <token>` header first, then
|
||||
`turnstone_auth` cookie as fallback.
|
||||
3. **Token type detection** — dots in the token indicate JWT; `ts_` prefix
|
||||
indicates API token; otherwise config-file token.
|
||||
4. **Validation** — JWT signature check, API token hash lookup in storage, or
|
||||
config-token hmac comparison.
|
||||
5. **Scope check** — `required_scope(method, path)` determines the minimum
|
||||
scope; the request is rejected with 403 if the token lacks it.
|
||||
6. **Context propagation** — on success, `ctx_user_id` is set so structured
|
||||
logging includes the authenticated identity on every log event.
|
||||
|
||||
### Architecture Split
|
||||
|
||||
- **Console** is the auth management hub — it hosts the admin endpoints for
|
||||
creating users, issuing API tokens, and managing channel mappings. User
|
||||
records and token hashes live in the shared storage backend. The console
|
||||
dashboard includes an **admin panel** (Users and Tokens tabs) for managing
|
||||
credentials through the browser.
|
||||
- **Server** is a JWT validator only — it validates tokens on each request but
|
||||
never creates users or tokens. Both processes share the same `jwt_secret`
|
||||
(via `TURNSTONE_JWT_SECRET` env var or `[auth].jwt_secret` config).
|
||||
- **First-time setup** — both server and console expose
|
||||
`POST /v1/api/auth/setup`, a public endpoint that creates the initial admin
|
||||
user when no users exist. This avoids the chicken-and-egg problem of needing
|
||||
`approve` scope to create the first user via `/api/admin/users`.
|
||||
|
||||
### Auth Storage Tables
|
||||
|
||||
Three tables in `storage/_schema.py` support identity:
|
||||
|
||||
```sql
|
||||
users
|
||||
user_id TEXT PRIMARY KEY
|
||||
username TEXT NOT NULL UNIQUE
|
||||
display_name TEXT NOT NULL
|
||||
password_hash TEXT NOT NULL -- bcrypt
|
||||
created TEXT NOT NULL
|
||||
|
||||
api_tokens
|
||||
token_id TEXT PRIMARY KEY
|
||||
token_hash TEXT NOT NULL UNIQUE -- SHA-256 of raw token
|
||||
token_prefix TEXT NOT NULL -- first 8 chars for display
|
||||
user_id TEXT NOT NULL
|
||||
name TEXT NOT NULL -- human-readable label
|
||||
scopes TEXT NOT NULL -- comma-separated
|
||||
created TEXT NOT NULL
|
||||
expires TEXT -- optional expiry timestamp
|
||||
|
||||
channel_users
|
||||
channel_type TEXT NOT NULL -- e.g. "slack", "discord"
|
||||
channel_user_id TEXT NOT NULL -- platform-specific user ID
|
||||
user_id TEXT NOT NULL -- FK to users
|
||||
PRIMARY KEY (channel_type, channel_user_id)
|
||||
```
|
||||
|
||||
See [docs/security.md](security.md) for full security details including token
|
||||
lifecycle, password hashing, and deployment hardening.
|
||||
|
||||
---
|
||||
|
||||
## Threading Model
|
||||
|
||||
### CLI
|
||||
@@ -1037,13 +1176,19 @@ a response or the approval timeout (default 3600s / 1 hour) expires.
|
||||
`ws_id` for active sends. When the global SSE reports `ws_state → idle` for a tracked
|
||||
workstream, the bridge emits a synthetic `TurnCompleteEvent` with the correlation ID.
|
||||
|
||||
**Multi-node routing:** Each bridge has a `node_id` (defaults to hostname) and BLPOPs
|
||||
**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
|
||||
|
||||
@@ -1069,7 +1214,10 @@ 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`
|
||||
for parallel HTTP polling.
|
||||
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).
|
||||
|
||||
The console has two write-path capabilities:
|
||||
|
||||
@@ -1128,7 +1276,7 @@ typed event dataclasses.
|
||||
|
||||
**Two client pairs** (sync + async):
|
||||
|
||||
- `TurnstoneServer` / `AsyncTurnstoneServer` — server API (workstreams, chat, streaming, sessions)
|
||||
- `TurnstoneServer` / `AsyncTurnstoneServer` — server API (workstreams, chat, streaming)
|
||||
- `TurnstoneConsole` / `AsyncTurnstoneConsole` — console API (cluster overview, nodes, workstreams)
|
||||
|
||||
**Design**: async-first with thin sync wrappers. `_BaseClient` provides httpx
|
||||
@@ -1152,3 +1300,42 @@ with TurnstoneServer("http://localhost:8080", token="tok_xxx") as client:
|
||||
result = client.send_and_wait("Hello!", ws.ws_id)
|
||||
print(result.content)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Channel Integrations
|
||||
|
||||
> 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
|
||||
platform adapter implements the `ChannelAdapter` protocol and translates
|
||||
between platform-native events and turnstone MQ messages.
|
||||
|
||||
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
|
||||
the old workstream's conversation during creation in a single HTTP
|
||||
request, eliminating ordering fragility. The bridge emits a
|
||||
`WorkstreamResumedEvent` to confirm success.
|
||||
|
||||
Discord ships as the first adapter. See [channels.md](channels.md) for
|
||||
setup instructions, configuration reference, and the adapter development
|
||||
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
|
||||
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`),
|
||||
and POSTs to `POST /v1/api/notify` on the first healthy gateway. The
|
||||
gateway validates the JWT, resolves the target (username lookup via
|
||||
`channel_users` or direct `channel_type`+`channel_id`), and delegates to
|
||||
the appropriate `ChannelAdapter.send()`. Delivery retries up to 3 times
|
||||
with backoff, re-querying the service registry on each attempt. See
|
||||
[Notification Flow diagram](diagrams/png/17-notify-flow.png).
|
||||
|
||||
@@ -0,0 +1,346 @@
|
||||
# Channel Integrations
|
||||
|
||||
The `turnstone-channel` gateway connects external messaging platforms to
|
||||
turnstone workstreams via Redis MQ. Each platform adapter translates
|
||||
platform-native events (messages, button clicks, slash commands) into
|
||||
turnstone MQ messages, and renders workstream output back into the
|
||||
platform's UI.
|
||||
|
||||
Discord ships as the first adapter. The adapter protocol is designed for
|
||||
future Slack and Teams integrations.
|
||||
|
||||
---
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
Discord Gateway
|
||||
|
|
||||
v
|
||||
turnstone-channel (Discord adapter)
|
||||
|
|
||||
v
|
||||
Redis MQ
|
||||
|
|
||||
v
|
||||
turnstone-bridge ──> turnstone-server
|
||||
```
|
||||
|
||||
Key components:
|
||||
|
||||
- **ChannelAdapter protocol** (`turnstone/channels/_protocol.py`) — generic
|
||||
interface for any messaging platform. Defines `start()`, `stop()`,
|
||||
`send()`, `edit_message()`, `send_approval_request()`,
|
||||
`send_plan_review()`, and `create_thread()`.
|
||||
- **ChannelRouter** (`turnstone/channels/_routing.py`) — maps
|
||||
channel/thread IDs to turnstone workstream IDs. Handles workstream
|
||||
creation via MQ, stale route detection, and user identity resolution.
|
||||
- **AsyncRedisBroker** (`turnstone/mq/async_broker.py`) — async Redis
|
||||
client compatible with discord.py's event loop. Used by the router for
|
||||
pub/sub and queue operations.
|
||||
- **channel_users table** — maps `(channel_type, channel_user_id)` to a
|
||||
turnstone `user_id`. Messages from unlinked users are silently dropped.
|
||||
- **channel_routes table** — persistent channel-to-workstream mappings.
|
||||
Survives bot restarts. Stale routes (evicted workstreams) are detected
|
||||
and refreshed on the next message.
|
||||
|
||||
---
|
||||
|
||||
## Discord Setup
|
||||
|
||||
### 1. Create a Discord Application
|
||||
|
||||
1. Go to https://discord.com/developers/applications
|
||||
2. Click **New Application** and give it a name
|
||||
3. Navigate to the **Bot** tab and click **Reset Token** to generate a
|
||||
bot token. Copy it immediately — it is shown only once.
|
||||
4. On the same **Bot** tab, scroll down to **Privileged Gateway Intents**
|
||||
and enable **MESSAGE CONTENT INTENT**
|
||||
5. Navigate to **OAuth2 > URL Generator**
|
||||
6. Under **Scopes**, check `bot` and `applications.commands`
|
||||
7. Under **Bot Permissions**, check:
|
||||
- View Channels
|
||||
- Send Messages
|
||||
- Send Messages in Threads
|
||||
- Create Public Threads
|
||||
- Read Message History
|
||||
- Add Reactions
|
||||
- Embed Links
|
||||
8. Copy the generated URL, open it in a browser, and add the bot to your
|
||||
Discord server
|
||||
|
||||
### 2. Configure Turnstone
|
||||
|
||||
**Environment variables** (recommended for Docker):
|
||||
|
||||
```bash
|
||||
TURNSTONE_DISCORD_TOKEN=your-bot-token-here
|
||||
TURNSTONE_DISCORD_GUILD=123456789 # optional, restrict to one guild
|
||||
```
|
||||
|
||||
**CLI flags** (bare-metal):
|
||||
|
||||
```bash
|
||||
turnstone-channel \
|
||||
--discord-token "your-bot-token" \
|
||||
--discord-guild 123456789 \
|
||||
--redis-host localhost \
|
||||
--redis-port 6379
|
||||
```
|
||||
|
||||
**Docker Compose** (production profile):
|
||||
|
||||
```bash
|
||||
# In .env file:
|
||||
TURNSTONE_DISCORD_TOKEN=your-bot-token
|
||||
TURNSTONE_DISCORD_GUILD=123456789
|
||||
```
|
||||
|
||||
Then start the stack:
|
||||
|
||||
```bash
|
||||
docker compose --profile production up
|
||||
```
|
||||
|
||||
The `channel` service starts automatically when
|
||||
`TURNSTONE_DISCORD_TOKEN` is set.
|
||||
|
||||
### 3. Link User Accounts
|
||||
|
||||
Discord users must link their account to a turnstone user before they can
|
||||
interact with the bot. Unlinked users' messages are silently ignored.
|
||||
|
||||
1. The user must have a turnstone API token — created via the admin panel
|
||||
or `turnstone-admin create-token`
|
||||
2. In Discord, the user runs `/link`. A modal appears prompting for the
|
||||
API token (the token is never visible in Discord audit logs because it
|
||||
is submitted via modal, not as a slash command argument).
|
||||
3. The token is validated against the database. If valid, a
|
||||
`channel_users` mapping is created.
|
||||
4. The user can now @mention the bot or use slash commands.
|
||||
|
||||
An admin can also force-link or unlink users via the console admin panel
|
||||
(Admin > Channels tab).
|
||||
|
||||
---
|
||||
|
||||
## Usage
|
||||
|
||||
### Conversations
|
||||
|
||||
- **@mention** the bot in any allowed channel to start a new conversation.
|
||||
The bot creates a Discord thread from the message and a turnstone
|
||||
workstream behind it.
|
||||
- All subsequent messages in the thread are routed to the same workstream.
|
||||
- The bot streams responses via message edits, updated approximately every
|
||||
1.5 seconds.
|
||||
- If the workstream is evicted for capacity, the next message in the
|
||||
thread auto-creates a new workstream and atomically resumes the
|
||||
previous workstream via the `resume_ws` field on
|
||||
`CreateWorkstreamMessage`. The server resumes the workstream during
|
||||
creation (same HTTP request), and the bridge emits a
|
||||
`WorkstreamResumedEvent` back to the channel. The thread receives a
|
||||
*"Resumed: {name} ({count} messages restored)"* confirmation.
|
||||
|
||||
### Slash Commands
|
||||
|
||||
| Command | Description |
|
||||
|---------|-------------|
|
||||
| `/link` | Link Discord account to turnstone (opens modal for API token) |
|
||||
| `/unlink` | Unlink Discord account |
|
||||
| `/ask <message>` | Create a new thread and workstream with an initial message |
|
||||
| `/status` | Show workstream info for the current thread (ephemeral) |
|
||||
| `/close` | Close the workstream, delete the route, and archive the thread |
|
||||
|
||||
### Tool Approvals
|
||||
|
||||
When manual approval is enabled (the default), tool calls are displayed as
|
||||
an orange embed with:
|
||||
|
||||
- Tool name and argument preview
|
||||
- **Approve** (green), **Reject** (red), **Always Approve** (gray) buttons
|
||||
- Only linked users can interact with approval buttons
|
||||
- The approval decision is forwarded through MQ to the bridge, which
|
||||
relays it to the server
|
||||
|
||||
Buttons use static `custom_id` values so they survive bot restarts.
|
||||
Correlation data (`ws_id`, `correlation_id`) is stored in the embed footer.
|
||||
|
||||
**Auto-approval:** When `auto_approve` is true (via `--auto-approve`), or when
|
||||
all tools in the request match the `auto_approve_tools` list in the adapter
|
||||
config, the bot auto-responds with approval and posts a
|
||||
"*Tool auto-approved.*" notice to the thread instead of showing buttons. The
|
||||
`auto_approve_tools` list is set via the `ChannelConfig.auto_approve_tools`
|
||||
field (useful for allowing specific tools like `bash` or `read_file` while
|
||||
still requiring manual approval for others).
|
||||
|
||||
### Plan Reviews
|
||||
|
||||
Plan review requests are displayed as a blue embed with:
|
||||
|
||||
- **Approve Plan** (green) button — approves the plan with empty feedback
|
||||
- **Request Changes** (gray) button — opens a modal for feedback text
|
||||
(up to 2000 characters)
|
||||
- Feedback is forwarded through MQ as a `PlanFeedbackMessage`
|
||||
|
||||
---
|
||||
|
||||
## Configuration Reference
|
||||
|
||||
| CLI Flag | Env Var | Default | Description |
|
||||
|----------|---------|---------|-------------|
|
||||
| `--discord-token` | `TURNSTONE_DISCORD_TOKEN` | — | Bot token (required to enable Discord) |
|
||||
| `--discord-guild` | — | `0` (all guilds) | Restrict to a single Discord guild |
|
||||
| `--discord-channels` | — | empty (all) | Comma-separated channel IDs to allow |
|
||||
| `--redis-host` | `REDIS_HOST` | `localhost` | Redis host |
|
||||
| `--redis-port` | — | `6379` | Redis port |
|
||||
| `--redis-password` | `REDIS_PASSWORD` | — | Redis password |
|
||||
| `--redis-db` | — | `0` | Redis DB number |
|
||||
| `--model` | — | server default | Default model for new workstreams |
|
||||
| `--auto-approve` | — | `false` | Auto-approve ALL tool calls (skips approval buttons entirely) |
|
||||
| `--http-host` | — | `127.0.0.1` | HTTP server bind address for notify endpoint |
|
||||
| `--http-port` | `TURNSTONE_CHANNEL_PORT` | `8091` | HTTP server port |
|
||||
| `--auth-token` | `TURNSTONE_CHANNEL_AUTH_TOKEN` | — | Static auth token for `/v1/api/notify` (alternative to JWT) |
|
||||
| `--log-level` | `TURNSTONE_LOG_LEVEL` | `INFO` | Log level |
|
||||
| `--log-format` | `TURNSTONE_LOG_FORMAT` | `auto` | Log format (`auto`/`json`/`text`) |
|
||||
|
||||
---
|
||||
|
||||
## User Identity
|
||||
|
||||
- The `channel_users` table maps `(channel_type, channel_user_id)` to a
|
||||
turnstone `user_id`
|
||||
- Self-service linking via the `/link` slash command (modal input, not
|
||||
visible in Discord audit logs)
|
||||
- Admin can force-link or unlink via the console admin panel (Admin >
|
||||
Channels tab). Unlinking uses a styled confirmation modal.
|
||||
- Unlinked users' messages are silently dropped
|
||||
- A user can be linked across multiple platforms (e.g. Discord + Slack)
|
||||
|
||||
See [Security: Database Schema](security.md#database-schema) for the
|
||||
`channel_users` table definition.
|
||||
|
||||
---
|
||||
|
||||
## Workstream Lifecycle
|
||||
|
||||
1. **Creation** — @mention or `/ask` creates a Discord thread and a
|
||||
turnstone workstream. The `ChannelRouter` persists the mapping in the
|
||||
`channel_routes` table.
|
||||
2. **Active** — messages are routed bidirectionally. The bot streams
|
||||
responses via message edits (updated every ~1.5 seconds).
|
||||
3. **Eviction** — the server evicts an idle workstream for capacity. The
|
||||
route is preserved and the thread stays open.
|
||||
4. **Reactivation** — the next message in the thread detects the stale
|
||||
route (no MQ owner) and creates a new workstream with the old `ws_id`
|
||||
as `resume_ws` on the `CreateWorkstreamMessage`. The server resumes
|
||||
the workstream during creation (no separate command or reverse lookup
|
||||
needed). The bridge emits a `WorkstreamResumedEvent` to the channel, and
|
||||
the thread displays *"Resumed: {name} ({count} messages restored)"*.
|
||||
If the old workstream was pruned, a fresh one starts with no error.
|
||||
5. **Close** — `/close` command closes the workstream via MQ, deletes the
|
||||
route, unsubscribes from events, and archives the Discord thread.
|
||||
|
||||
---
|
||||
|
||||
## Notifications
|
||||
|
||||
> See also: [Notification Flow diagram](diagrams/png/17-notify-flow.png)
|
||||
|
||||
The `notify` tool allows the LLM to proactively send notifications to
|
||||
users or channels on external platforms. This is useful for alerting
|
||||
people about task completion, errors, or important updates without
|
||||
waiting for them to check in.
|
||||
|
||||
### Targeting
|
||||
|
||||
Two modes:
|
||||
|
||||
- **Username** — provide a turnstone `username`. The gateway resolves
|
||||
it via the `channel_users` table and sends to all linked channels
|
||||
(e.g. Discord + future Slack).
|
||||
- **Direct** — provide `channel_type` + `channel_id` to target a
|
||||
specific platform channel or user DM.
|
||||
|
||||
### Delivery Flow
|
||||
|
||||
Notifications bypass MQ for lower latency. The server calls the channel
|
||||
gateway directly over HTTP:
|
||||
|
||||
1. The LLM calls the `notify` tool with a message and target
|
||||
2. `_exec_notify()` queries the `services` table for healthy channel
|
||||
gateways (heartbeat within the last 120 seconds)
|
||||
3. The server mints a service JWT (`aud: turnstone-channel`) via
|
||||
`ServiceTokenManager` and POSTs to the first healthy gateway
|
||||
4. The gateway validates the JWT, resolves the target, and calls
|
||||
`adapter.send()` on the appropriate platform adapter
|
||||
5. On failure, the server tries the next gateway. If all fail, it
|
||||
retries up to 2 more times (delays: 1s, 3s), re-querying the
|
||||
service registry on each attempt
|
||||
|
||||
### Service Registry
|
||||
|
||||
The channel gateway registers itself in the `services` database table
|
||||
on startup and sends a heartbeat every 30 seconds. On shutdown it
|
||||
deregisters. Services are considered stale after 120 seconds (4 missed
|
||||
heartbeats) and are excluded from `list_services()` queries.
|
||||
|
||||
The `services` table schema:
|
||||
|
||||
| Column | Description |
|
||||
|--------|-------------|
|
||||
| `service_type` | Service category (e.g. `"channel"`) |
|
||||
| `service_id` | Unique instance ID (`channel-<hostname>-<random>`) |
|
||||
| `url` | HTTP base URL for the service |
|
||||
| `last_heartbeat` | ISO 8601 timestamp of last heartbeat |
|
||||
| `created` | ISO 8601 timestamp of initial registration |
|
||||
|
||||
### Security
|
||||
|
||||
- **Authentication** — the gateway's `POST /v1/api/notify` endpoint
|
||||
requires authentication. Configure either `TURNSTONE_JWT_SECRET`
|
||||
(the server mints JWTs with `aud: turnstone-channel` automatically)
|
||||
or a static token via `--auth-token`. If neither is set, the
|
||||
gateway fails closed and rejects all requests with 401. Server JWTs
|
||||
(`aud: turnstone-server`) are rejected.
|
||||
- **Rate limit** — maximum 5 notifications per turn. The counter only
|
||||
increments on successful delivery, so failures don't consume the
|
||||
budget.
|
||||
- **SSRF protection** — only `http://` and `https://` service URLs
|
||||
are allowed. Other schemes are silently skipped.
|
||||
- **Mention sanitization** — `discord.utils.escape_mentions()` is
|
||||
applied before sending, preventing `@everyone` / `@here` abuse.
|
||||
- **Error redaction** — generic error messages are returned to the
|
||||
LLM. Internal details (service IDs, URLs, exception messages) are
|
||||
logged server-side only.
|
||||
|
||||
---
|
||||
|
||||
## Adding New Adapters
|
||||
|
||||
The `ChannelAdapter` protocol defines the interface any platform adapter
|
||||
must implement:
|
||||
|
||||
```python
|
||||
class ChannelAdapter(Protocol):
|
||||
channel_type: str
|
||||
|
||||
async def start(self) -> None: ...
|
||||
async def stop(self) -> None: ...
|
||||
async def send(self, channel_id: str, content: str) -> str: ...
|
||||
async def edit_message(self, channel_id: str, message_id: str, content: str) -> None: ...
|
||||
async def send_approval_request(self, channel_id: str, ws_id: str, correlation_id: str, items: list[dict]) -> None: ...
|
||||
async def send_plan_review(self, channel_id: str, ws_id: str, correlation_id: str, content: str) -> None: ...
|
||||
async def create_thread(self, parent_channel_id: str, name: str, message_id: str = "") -> str: ...
|
||||
```
|
||||
|
||||
To add a new platform:
|
||||
|
||||
1. Create `turnstone/channels/<platform>/` package
|
||||
2. Implement the `ChannelAdapter` protocol
|
||||
3. Add a `--<platform>-token` flag and detection logic in
|
||||
`turnstone/channels/cli.py`
|
||||
4. Add the optional dependency in `pyproject.toml` (e.g.
|
||||
`turnstone[slack]`)
|
||||
|
||||
See `turnstone/channels/discord/` as a reference implementation.
|
||||
+329
-6
@@ -61,6 +61,8 @@ The collector (`turnstone/console/collector.py`) maintains an in-memory snapshot
|
||||
|
||||
3. **Poll loop** — fetches `GET /v1/api/dashboard` and `GET /health` from each known node every 10 seconds. Uses `ThreadPoolExecutor(max_workers=50)` for parallelism. Each poll replaces the node's workstream list with the authoritative server data.
|
||||
|
||||
A `get_snapshot()` method builds the full cluster state under a single lock acquisition — overview aggregates and per-node workstream lists in one atomic read. This is served both as a REST endpoint and as the initial SSE event on client connect.
|
||||
|
||||
### Thread Safety
|
||||
|
||||
All reads and writes to the node/workstream map are protected by a single `threading.Lock`. Query methods acquire the lock, copy data, and release before returning.
|
||||
@@ -146,9 +148,41 @@ Single node detail with all its workstreams.
|
||||
}
|
||||
```
|
||||
|
||||
### `GET /v1/api/cluster/snapshot`
|
||||
|
||||
Full cluster state in a single response — all nodes with their workstreams plus overview aggregates. Built under a single lock for internal consistency. Used by the browser on initial load and SSE reconnect.
|
||||
|
||||
```json
|
||||
{
|
||||
"nodes": [
|
||||
{
|
||||
"node_id": "db-west-04",
|
||||
"server_url": "http://10.0.3.4:8080",
|
||||
"max_ws": 10,
|
||||
"reachable": true,
|
||||
"version": "0.3.0",
|
||||
"health": {"status": "ok", "version": "0.3.0"},
|
||||
"aggregate": {"total_tokens": 48200, "total_tool_calls": 156},
|
||||
"workstreams": [
|
||||
{"id": "a1b2c3d4", "name": "perf-db-west", "state": "running", ...}
|
||||
]
|
||||
}
|
||||
],
|
||||
"overview": {
|
||||
"nodes": 847,
|
||||
"workstreams": 4219,
|
||||
"states": {"running": 1847, "thinking": 312, "attention": 89, "idle": 1940, "error": 31},
|
||||
"aggregate": {"total_tokens": 12400000, "total_tool_calls": 34200},
|
||||
"version_drift": false,
|
||||
"versions": ["0.3.0"]
|
||||
},
|
||||
"timestamp": 1709294400.0
|
||||
}
|
||||
```
|
||||
|
||||
### `POST /v1/api/cluster/workstreams/new`
|
||||
|
||||
Create a new workstream on a target node. Dispatches a `CreateWorkstreamMessage` through the Redis MQ pipeline — the bridge on the target node picks it up and creates the workstream on the server. Requires `"full"` auth role.
|
||||
Create a new workstream on a target node. Dispatches a `CreateWorkstreamMessage` through the Redis MQ pipeline — the bridge on the target node picks it up and creates the workstream on the server. Requires `write` scope.
|
||||
|
||||
Request:
|
||||
|
||||
@@ -182,7 +216,7 @@ Creation is asynchronous — the response confirms the MQ message was dispatched
|
||||
|
||||
### `GET /v1/api/cluster/events`
|
||||
|
||||
Server-Sent Events stream for real-time cluster updates.
|
||||
Server-Sent Events stream for real-time cluster updates. The first event is always a `snapshot` containing the full cluster state (same shape as `GET /v1/api/cluster/snapshot` with an added `type: "snapshot"` field), followed by incremental events:
|
||||
|
||||
```
|
||||
data: {"type":"cluster_state","ws_id":"a1b2","node_id":"db-west-04","state":"running"}
|
||||
@@ -207,6 +241,96 @@ Keepalive comments (`: keepalive\n\n`) are sent every 5 seconds. Clients should
|
||||
}
|
||||
```
|
||||
|
||||
### Admin API
|
||||
|
||||
User and token management endpoints. All admin endpoints require `approve` scope, except for the setup endpoint which is public.
|
||||
|
||||
#### `POST /v1/api/auth/setup`
|
||||
|
||||
Creates the first admin user when no users exist. Public endpoint (no auth required). Returns a JWT and sets a session cookie. Returns `409` if users already exist. See [Security: First-time setup](security.md#first-time-setup) for full details.
|
||||
|
||||
#### `POST /v1/api/admin/users`
|
||||
|
||||
Create a new user.
|
||||
|
||||
```json
|
||||
{
|
||||
"username": "alice",
|
||||
"password": "s3cret",
|
||||
"scopes": ["read", "write"]
|
||||
}
|
||||
```
|
||||
|
||||
#### `GET /v1/api/admin/users`
|
||||
|
||||
List all users.
|
||||
|
||||
```json
|
||||
{
|
||||
"users": [
|
||||
{"user_id": "u_abc123", "username": "alice", "scopes": ["read", "write"], "created": "2026-03-01T12:00:00Z"}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
#### `DELETE /v1/api/admin/users/{user_id}`
|
||||
|
||||
Delete a user and revoke all their tokens.
|
||||
|
||||
#### `POST /v1/api/admin/users/{user_id}/tokens`
|
||||
|
||||
Create an API token for the given user. Returns a `ts_`-prefixed token string that can be used for Bearer auth or passed to `client.login(token="ts_xxx")`.
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "CI pipeline",
|
||||
"scopes": ["read", "write"]
|
||||
}
|
||||
```
|
||||
|
||||
#### `GET /v1/api/admin/users/{user_id}/tokens`
|
||||
|
||||
List active tokens for a user (token strings are not returned, only metadata).
|
||||
|
||||
#### `DELETE /v1/api/admin/tokens/{token_id}`
|
||||
|
||||
Revoke a specific API token.
|
||||
|
||||
### Channel links
|
||||
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| GET | `/v1/api/admin/users/{user_id}/channels` | List channel links for a user |
|
||||
| POST | `/v1/api/admin/users/{user_id}/channels` | Link a channel account (channel_type, channel_user_id) |
|
||||
| DELETE | `/v1/api/admin/channels/{channel_type}/{channel_user_id}` | Unlink a channel account |
|
||||
|
||||
These endpoints manage the `channel_users` table mappings that connect external platform identities (e.g. Discord user IDs) to turnstone users. See [Channel Integrations](channels.md) for details on the linking flow.
|
||||
|
||||
#### `GET /v1/api/auth/status`
|
||||
|
||||
Public endpoint for login UI state detection. Returns auth configuration, not
|
||||
current-user identity.
|
||||
|
||||
```json
|
||||
{
|
||||
"auth_enabled": true,
|
||||
"has_users": true,
|
||||
"setup_required": false
|
||||
}
|
||||
```
|
||||
|
||||
### Auth Scopes
|
||||
|
||||
The auth system uses three scopes instead of the earlier read/full role model:
|
||||
|
||||
| Scope | Grants |
|
||||
|-------|--------|
|
||||
| `read` | Read-only access: dashboards, workstream lists, SSE streams, health |
|
||||
| `write` | Send messages, create/close workstreams, approve tool calls |
|
||||
| `approve` | Admin operations: manage users and API tokens |
|
||||
|
||||
Scopes are cumulative — a user with `approve` scope can also perform `write` and `read` operations.
|
||||
|
||||
---
|
||||
|
||||
## Reverse Proxy
|
||||
@@ -236,17 +360,17 @@ The server UI uses root-relative URLs (`/v1/api/send`, `/static/app.js`, `/share
|
||||
|
||||
### SSE Proxy
|
||||
|
||||
SSE streams (`/v1/api/events`, `/v1/api/events/global`) are proxied by creating a per-connection `httpx.AsyncClient(timeout=None)`, streaming the upstream response via `aiter_text()`, parsing SSE framing (`\n\n` delimiters), and re-emitting events through `EventSourceResponse`. Each proxied SSE stream requires its own httpx client since the shared client's 30-second timeout would kill long-lived connections.
|
||||
SSE streams (`/v1/api/events`, `/v1/api/events/global`) are proxied as raw byte passthrough — the console opens an `httpx.AsyncClient.stream()` to the upstream server (with `read=None` and `pool=None` timeouts since SSE connections are long-lived) and relays every byte via `StreamingResponse`. This preserves server-side ping comments, event framing, and keepalives verbatim without parsing or re-encoding.
|
||||
|
||||
### Authentication
|
||||
|
||||
The proxy forwards requests to server nodes using the console's `--auth-token`. The console's own auth middleware also checks proxy routes — `POST` requests to proxy write endpoints (`/v1/api/send`, `/v1/api/approve`, etc.) require the `"full"` auth role, preventing read-only tokens from escalating to write operations.
|
||||
The proxy forwards the user's JWT to upstream server nodes — it extracts the token from the incoming request's cookie (or `Authorization` header) and adds it as a `Bearer` header on the proxied request. Since all services share the same `TURNSTONE_JWT_SECRET`, the user's JWT is valid on every node without re-authentication. The console's own auth middleware also checks proxy routes — `POST` requests to proxy write endpoints (`/v1/api/send`, `/v1/api/approve`, etc.) require `write` scope, preventing read-only tokens from escalating via proxy. The static `--auth-token` / `proxy_auth_token` is used as a fallback when no user JWT is present.
|
||||
|
||||
---
|
||||
|
||||
## Browser Dashboard
|
||||
|
||||
The web UI has four views, toggled client-side:
|
||||
The web UI has five views, toggled client-side:
|
||||
|
||||
### 1. Cluster Overview (landing)
|
||||
|
||||
@@ -276,7 +400,206 @@ Triggered by the "+ new" header button. A modal dialog with:
|
||||
|
||||
On submit, `POST /v1/api/cluster/workstreams/new` dispatches the creation request. A toast confirms success; the SSE stream delivers the `ws_created` event to update the dashboard.
|
||||
|
||||
All four views receive live updates via SSE — state cards update counts, node rows update metrics, workstream rows update state indicators.
|
||||
All five views receive live updates via SSE — state cards update counts, node rows update metrics, workstream rows update state indicators.
|
||||
|
||||
The browser maintains a local `clusterState` object that mirrors the cluster snapshot. It is initialized from the SSE `snapshot` event on connect (or via `GET /v1/api/cluster/snapshot` on initial page load) and updated incrementally by SSE events. View navigation reads from local state — no API round-trips needed after the initial snapshot.
|
||||
|
||||
### 5. Admin Panel
|
||||
|
||||
Accessed via the "admin" button in the header (visible when authenticated
|
||||
with `approve` scope). Provides user, API token, and channel link management
|
||||
with three tabs:
|
||||
|
||||
**Users tab:**
|
||||
|
||||
- Grid table listing all users (username, display name, role, creation date)
|
||||
- "Create User" button opens a modal with fields for username, display name,
|
||||
and password (validated: username 1-64 ASCII, password min 8 characters)
|
||||
- Delete button on each row opens a styled confirmation modal before
|
||||
removing the user and cascading to revoke all their tokens
|
||||
|
||||
**Tokens tab:**
|
||||
|
||||
- User selector dropdown to pick which user's tokens to manage
|
||||
- Grid table listing tokens for the selected user (name, prefix, scopes,
|
||||
creation date)
|
||||
- Scope badges rendered as colored pills for visual clarity
|
||||
- "Create Token" button opens a modal with fields for token name and scope
|
||||
checkboxes
|
||||
- On creation, a "Token Created" modal displays the raw `ts_`-prefixed
|
||||
token with a copy button. The token is shown once and cannot be retrieved
|
||||
again.
|
||||
- Revoke button on each row opens a styled confirmation modal before
|
||||
deleting the token
|
||||
|
||||
**Channels tab:**
|
||||
|
||||
- User selector dropdown to pick which user's channel links to manage
|
||||
- Grid table listing linked channel accounts for the selected user
|
||||
(channel type, channel user ID, creation date)
|
||||
- "Link Channel" button opens a modal with fields for channel type
|
||||
(e.g. `discord`) and the platform user ID
|
||||
- Unlink button on each row opens a styled confirmation modal before
|
||||
removing the channel mapping
|
||||
- Admins can force-link users who have not self-linked via `/link` in
|
||||
Discord
|
||||
|
||||
**Accessibility:**
|
||||
|
||||
- Full keyboard navigation: focus traps in modals, Escape to close, arrow
|
||||
keys for tab switching
|
||||
- Responsive layout with column hiding at 700px breakpoint
|
||||
|
||||
**First-time setup:**
|
||||
|
||||
The console also exposes `POST /v1/api/auth/setup` for first-time
|
||||
bootstrap. When no users exist, the setup wizard calls this public endpoint
|
||||
to create the initial admin user and receive a JWT in one step. See
|
||||
[Security: First-time setup](security.md#first-time-setup) for details.
|
||||
|
||||
---
|
||||
|
||||
## Scheduled Tasks
|
||||
|
||||
The console includes a background **TaskScheduler** daemon that creates workstreams on a timed basis via the MQ broker. It supports cron-based recurring schedules and one-shot `at` schedules.
|
||||
|
||||
### Architecture
|
||||
|
||||
The scheduler runs as a daemon thread inside the console process. Every `check_interval` seconds (default 15) it:
|
||||
|
||||
1. Acquires a distributed lock via Redis `SET NX EX` (prevents duplicate dispatch in multi-console deployments)
|
||||
2. Queries the storage backend for tasks whose `next_run <= now` and `enabled = true`
|
||||
3. Dispatches each due task as one or more `CreateWorkstreamMessage` via MQ
|
||||
4. Updates `last_run` and computes the next `next_run` (or disables one-shot `at` tasks)
|
||||
5. Releases the lock via Lua script (safe conditional delete)
|
||||
|
||||
Run history is automatically pruned (runs older than 90 days) approximately once per hour.
|
||||
|
||||
### Schedule Types
|
||||
|
||||
| Type | Field | Behavior |
|
||||
|------|-------|----------|
|
||||
| `cron` | `cron_expr` | Recurring schedule using standard 5-field cron syntax. Requires `croniter`. |
|
||||
| `at` | `at_time` | One-shot: fires once at the given ISO 8601 timestamp (must include timezone), then auto-disables. |
|
||||
|
||||
### Target Modes
|
||||
|
||||
| Mode | Behavior |
|
||||
|------|----------|
|
||||
| `auto` | Picks the reachable node with the most available capacity |
|
||||
| `pool` | Pushes to the shared inbound queue (any bridge picks it up) |
|
||||
| `all` | Fan-out to all reachable nodes (capped at `max_fan_out`, default 20) |
|
||||
| `<node_id>` | Targets a specific node by ID |
|
||||
|
||||
### Configuration
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `check_interval` | `15.0` | Seconds between scheduler ticks |
|
||||
| `lock_ttl` | `60` | Distributed lock TTL in seconds |
|
||||
| `max_fan_out` | `20` | Maximum nodes for `all` target mode |
|
||||
|
||||
Dependency: `croniter` (installed with turnstone).
|
||||
|
||||
### Schedule API
|
||||
|
||||
All schedule endpoints require `approve` scope. Maximum 200 schedules.
|
||||
|
||||
#### `GET /v1/api/admin/schedules`
|
||||
|
||||
List all scheduled tasks.
|
||||
|
||||
```json
|
||||
{
|
||||
"schedules": [
|
||||
{
|
||||
"task_id": "a1b2c3d4",
|
||||
"name": "nightly-checks",
|
||||
"description": "Run nightly health checks",
|
||||
"schedule_type": "cron",
|
||||
"cron_expr": "0 2 * * *",
|
||||
"at_time": "",
|
||||
"target_mode": "auto",
|
||||
"model": "",
|
||||
"initial_message": "Run the nightly health check suite.",
|
||||
"auto_approve": false,
|
||||
"auto_approve_tools": [],
|
||||
"enabled": true,
|
||||
"created_by": "u_admin",
|
||||
"last_run": "2026-03-05T02:00:00Z",
|
||||
"next_run": "2026-03-06T02:00:00Z",
|
||||
"created": "2026-03-01T12:00:00Z",
|
||||
"updated": "2026-03-05T02:00:01Z"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
#### `POST /v1/api/admin/schedules`
|
||||
|
||||
Create a scheduled task.
|
||||
|
||||
Request:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "nightly-checks",
|
||||
"description": "Run nightly health checks",
|
||||
"schedule_type": "cron",
|
||||
"cron_expr": "0 2 * * *",
|
||||
"target_mode": "auto",
|
||||
"initial_message": "Run the nightly health check suite.",
|
||||
"auto_approve": false,
|
||||
"enabled": true
|
||||
}
|
||||
```
|
||||
|
||||
Required fields: `name`, `schedule_type`, `initial_message`. For `cron` schedules provide `cron_expr`; for `at` schedules provide `at_time` (ISO 8601 with timezone, must be in the future).
|
||||
|
||||
Response: `ScheduleInfo` (same shape as list items above). Returns `400` for invalid cron syntax, naive timestamps, or past `at_time`. Returns `409` if the 200-schedule cap is reached.
|
||||
|
||||
#### `GET /v1/api/admin/schedules/{task_id}`
|
||||
|
||||
Get a single scheduled task. Returns `ScheduleInfo` or `404`.
|
||||
|
||||
#### `PUT /v1/api/admin/schedules/{task_id}`
|
||||
|
||||
Partial update — only include fields to change. If `schedule_type`, `cron_expr`, or `at_time` change, `next_run` is recomputed automatically.
|
||||
|
||||
```json
|
||||
{
|
||||
"enabled": false
|
||||
}
|
||||
```
|
||||
|
||||
Response: updated `ScheduleInfo`. Returns `400` for validation errors, `404` if not found.
|
||||
|
||||
#### `DELETE /v1/api/admin/schedules/{task_id}`
|
||||
|
||||
Delete a scheduled task and all its run history. Returns `{"status": "ok"}` or `404`.
|
||||
|
||||
#### `GET /v1/api/admin/schedules/{task_id}/runs?limit=50`
|
||||
|
||||
List execution history for a task (most recent first). `limit` defaults to 50, max 200.
|
||||
|
||||
```json
|
||||
{
|
||||
"runs": [
|
||||
{
|
||||
"run_id": "r_abc123",
|
||||
"task_id": "a1b2c3d4",
|
||||
"node_id": "db-west-04",
|
||||
"ws_id": "ws_xyz",
|
||||
"correlation_id": "corr_789",
|
||||
"started": "2026-03-05T02:00:00Z",
|
||||
"status": "dispatched",
|
||||
"error": ""
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
Status is `dispatched` on success or `failed` with an `error` message (e.g. no reachable nodes). Failed runs do not advance `next_run`.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -1,3 +0,0 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:d5a2bd1c55ac8cf3b777a8decb6f3bb3d063c10c8f3a9e63457079830e48f456
|
||||
size 162310
|
||||
@@ -1,3 +0,0 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:09cee5819bb7820641a466ed29a53f1b479bc2ded83079fc798c7410bf741a62
|
||||
size 329703
|
||||
@@ -40,7 +40,8 @@ package "turnstone/core/" <<Rectangle>> {
|
||||
component [auth.py\nAuthentication] as auth <<core>>
|
||||
component [healthcheck.py\nBackendHealthMonitor] as healthcheck <<core>>
|
||||
component [ratelimit.py\nRateLimiter] as ratelimit <<core>>
|
||||
component [mcp_client.py\nMCPClientManager] as mcp <<core>>
|
||||
component [mcp_client.py\nMCPClientManager\n(push + periodic refresh)] as mcp <<core>>
|
||||
component [tool_search.py\nToolSearchManager, BM25] as toolsearch <<core>>
|
||||
component [model_registry.py\nModelRegistry] as registry <<core>>
|
||||
}
|
||||
|
||||
@@ -95,7 +96,7 @@ package "turnstone/sdk/" <<Rectangle>> {
|
||||
|
||||
' Tool schemas
|
||||
package "turnstone/tools/" <<Rectangle>> {
|
||||
component [*.json\n14 tool schemas] as schemas <<artifact>>
|
||||
component [*.json\n15 tool schemas] as schemas <<artifact>>
|
||||
}
|
||||
|
||||
' Entry point dependencies
|
||||
@@ -136,6 +137,7 @@ session --> edit
|
||||
session --> web
|
||||
session --> healthcheck
|
||||
session --> mcp : optional
|
||||
session --> toolsearch : optional
|
||||
session --> registry : optional
|
||||
registry --> providers
|
||||
healthcheck --> metrics
|
||||
|
||||
@@ -1,3 +0,0 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:b3f76042c8046560502fa3132351821d56e8526b13d38d7be8b839fcf3d5f648
|
||||
size 373463
|
||||
@@ -108,6 +108,8 @@ class "ModelCapabilities" as ModelCaps <<frozen>> {
|
||||
+ thinking_mode: str
|
||||
+ supports_effort: bool
|
||||
+ supports_web_search: bool
|
||||
+ supports_tool_search: bool
|
||||
+ supports_vision: bool
|
||||
}
|
||||
|
||||
' ChatSession
|
||||
@@ -118,8 +120,9 @@ class "ChatSession" as ChatSession {
|
||||
- ui: SessionUI
|
||||
- messages: list[dict]
|
||||
- _msg_tokens: list[int]
|
||||
- _session_id: str
|
||||
- _ws_id: str
|
||||
- _mcp_client: MCPClientManager | None
|
||||
- _tool_search: ToolSearchManager | None
|
||||
- _registry: ModelRegistry | None
|
||||
+ model_alias: str | None {property}
|
||||
- _tools: list[dict]
|
||||
@@ -130,7 +133,7 @@ class "ChatSession" as ChatSession {
|
||||
--
|
||||
+ send(user_input: str)
|
||||
+ handle_command(command: str)
|
||||
+ resume_session(session_id: str)
|
||||
+ resume(ws_id: str)
|
||||
- _save_config()
|
||||
- _stream_response(stream) → dict
|
||||
- _create_stream_with_retry(msgs) → Stream (+ fallback)
|
||||
@@ -139,6 +142,12 @@ class "ChatSession" as ChatSession {
|
||||
- _prepare_tool(tc) → item dict
|
||||
- _prepare_mcp_tool(call_id, name, args) → item dict
|
||||
- _exec_mcp_tool(item) → (call_id, output)
|
||||
- _get_active_tools() → list[dict]
|
||||
- _prepare_tool_search() → None
|
||||
- _exec_tool_search(item) → (call_id, output)
|
||||
- _on_mcp_tools_changed()
|
||||
- _rebuild_tool_search()
|
||||
+ close()
|
||||
- _run_agent(messages, tools, ...) → str
|
||||
- _compact_messages(auto: bool)
|
||||
- _full_messages() → list[dict]
|
||||
@@ -201,22 +210,48 @@ enum "WorkstreamState" as WsState {
|
||||
' MCPClientManager
|
||||
class "MCPClientManager" as MCPMgr {
|
||||
- _sessions: dict[str, ClientSession]
|
||||
- _per_server_tools: dict[str, list[dict]]
|
||||
- _tools: list[dict]
|
||||
- _tool_map: dict[str, tuple]
|
||||
- _supports_list_changed: dict[str, bool]
|
||||
- _listeners: list[Callable]
|
||||
--
|
||||
+ start()
|
||||
+ get_tools() → list[dict]
|
||||
+ is_mcp_tool(name) → bool
|
||||
+ call_tool_sync(name, args) → str
|
||||
+ refresh_sync(server?) → dict
|
||||
+ add_listener(callback)
|
||||
+ remove_listener(callback)
|
||||
+ server_names: list[str] {property}
|
||||
+ shutdown()
|
||||
--
|
||||
Background asyncio event loop
|
||||
bridges async MCP SDK to
|
||||
sync ChatSession dispatch.
|
||||
Push + periodic + manual refresh.
|
||||
--
|
||||
core/mcp_client.py
|
||||
}
|
||||
|
||||
' ToolSearchManager
|
||||
class "ToolSearchManager" as ToolSearchMgr {
|
||||
- _all_tools: list[dict]
|
||||
- _always_on: list[dict]
|
||||
- _deferred: list[dict]
|
||||
- _expanded: dict[str, None]
|
||||
- _index: BM25Index
|
||||
--
|
||||
+ should_activate() → bool
|
||||
+ get_visible_tools() → list[dict]
|
||||
+ get_deferred_tools() → list[dict]
|
||||
+ get_expanded_names() → list[str]
|
||||
+ search(query, k) → list[dict]
|
||||
+ expand_visible(names) → list[dict]
|
||||
+ get_search_tool_definition() → dict
|
||||
+ format_search_results(tools) → str
|
||||
}
|
||||
|
||||
' ModelRegistry
|
||||
class "ModelRegistry" as ModelReg {
|
||||
- _models: dict[str, ModelConfig]
|
||||
@@ -317,6 +352,7 @@ LLMProvider <|.. AnthropicProv
|
||||
ChatSession --> SessionUI : uses
|
||||
ChatSession --> LLMProvider : delegates LLM calls
|
||||
ChatSession --> MCPMgr : optional
|
||||
ChatSession --o ToolSearchMgr : _tool_search
|
||||
ChatSession --> ModelReg : optional
|
||||
ChatSession <|-- HeadlessSession
|
||||
|
||||
|
||||
@@ -1,3 +0,0 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:fca0957b54101ce5c2b2e06d639b04dc5fff733f2641af0d732c49ea86883882
|
||||
size 279397
|
||||
@@ -18,7 +18,7 @@ User -> CS : send(user_input)
|
||||
activate CS
|
||||
|
||||
CS -> CS : messages.append({role: "user", content: input})
|
||||
CS -> DB : save_message(session_id, "user", input)
|
||||
CS -> DB : save_message(ws_id, "user", input)
|
||||
|
||||
== LLM Call Loop ==
|
||||
|
||||
@@ -65,8 +65,8 @@ group loop [while tool_calls present]
|
||||
|
||||
CS -> CS : _update_token_table()\ncalibrate chars_per_token ratio
|
||||
CS -> CS : messages.append(assistant_msg)
|
||||
CS -> DB : save_message(session_id, "assistant", content)
|
||||
CS -> DB : save_message(session_id, "tool_call", ...) ×N
|
||||
CS -> DB : save_message(ws_id, "assistant", content)
|
||||
CS -> DB : save_message(ws_id, "tool_call", ...) ×N
|
||||
|
||||
== Tool Dispatch (if tool_calls) ==
|
||||
|
||||
@@ -112,7 +112,7 @@ group loop [while tool_calls present]
|
||||
note right of TP
|
||||
Parallel execution:
|
||||
bash → Popen + line-by-line streaming
|
||||
read_file → open().read()
|
||||
read_file → open().read() or base64 image
|
||||
search → grep subprocess
|
||||
edit_file → string replace
|
||||
task/plan → _run_agent() sub-loop
|
||||
@@ -136,7 +136,7 @@ group loop [while tool_calls present]
|
||||
|
||||
loop for each result
|
||||
CS -> CS : messages.append({role: "tool", ...})
|
||||
CS -> DB : save_message(session_id, "tool_result", ...)
|
||||
CS -> DB : save_message(ws_id, "tool_result", ...)
|
||||
end
|
||||
|
||||
opt user_feedback from approval
|
||||
|
||||
@@ -1,3 +0,0 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:06fb9faa29c11395c6fc54ebddc79994b000dee78d56e0c13cb689fd6a82e37a
|
||||
size 237255
|
||||
@@ -24,27 +24,29 @@ partition "Phase 1: Prepare" #E8F5E9 {
|
||||
:Dispatch to _prepare_{func_name}();
|
||||
|
||||
note right
|
||||
**Dispatch table (14 tools):**
|
||||
┌─────────────┬──────────────────┐
|
||||
│ Tool │ Needs Approval? │
|
||||
├─────────────┼──────────────────┤
|
||||
│ bash │ ✓ Yes │
|
||||
│ read_file │ ✗ Auto-approve │
|
||||
│ write_file │ ✓ Yes │
|
||||
│ edit_file │ ✓ Yes │
|
||||
│ search │ ✗ Auto-approve │
|
||||
│ math │ ✓ Yes │
|
||||
│ man │ ✗ Auto-approve │
|
||||
│ web_fetch │ ✓ Yes │
|
||||
│ web_search │ ✓ Yes │
|
||||
│ task │ ✓ Yes │
|
||||
│ plan │ ✓ Yes │
|
||||
│ remember │ ✗ Auto-approve │
|
||||
│ recall │ ✗ Auto-approve │
|
||||
│ forget │ ✗ Auto-approve │
|
||||
├─────────────┼──────────────────┤
|
||||
│ mcp__* │ ✓ Yes (external) │
|
||||
└─────────────┴──────────────────┘
|
||||
**Dispatch table (16 tools):**
|
||||
┌──────────────┬──────────────────┐
|
||||
│ Tool │ Needs Approval? │
|
||||
├──────────────┼──────────────────┤
|
||||
│ bash │ ✓ Yes │
|
||||
│ read_file │ ✗ Auto-approve │
|
||||
│ write_file │ ✓ Yes │
|
||||
│ edit_file │ ✓ Yes │
|
||||
│ search │ ✗ Auto-approve │
|
||||
│ math │ ✓ Yes │
|
||||
│ man │ ✗ Auto-approve │
|
||||
│ web_fetch │ ✓ Yes │
|
||||
│ web_search │ ✓ Yes │
|
||||
│ tool_search │ ✗ Auto-approve │
|
||||
│ task │ ✓ Yes │
|
||||
│ plan │ ✓ Yes │
|
||||
│ remember │ ✗ Auto-approve │
|
||||
│ recall │ ✗ Auto-approve │
|
||||
│ forget │ ✗ Auto-approve │
|
||||
│ notify │ ✗ Auto-approve │
|
||||
├──────────────┼──────────────────┤
|
||||
│ mcp__* │ ✓ Yes (external) │
|
||||
└──────────────┴──────────────────┘
|
||||
end note
|
||||
|
||||
:Build item dict:
|
||||
@@ -98,7 +100,7 @@ partition "Phase 3: Execute" #E3F2FD {
|
||||
if item.denied → return denial message
|
||||
else → item["execute"](item)
|
||||
├─ _exec_bash: subprocess.run(["bash", script.sh])
|
||||
├─ _exec_read_file: open().readlines()
|
||||
├─ _exec_read_file: open().readlines() or _exec_read_image (base64)
|
||||
├─ _exec_write_file: makedirs + write
|
||||
├─ _exec_edit_file: find_occurrences + replace
|
||||
├─ _exec_search: grep subprocess
|
||||
@@ -106,8 +108,10 @@ partition "Phase 3: Execute" #E3F2FD {
|
||||
├─ _exec_man: man/info subprocess
|
||||
├─ _exec_web_fetch: httpx.get + LLM summary
|
||||
├─ _exec_web_search: Tavily API POST (fallback for local models)
|
||||
├─ _exec_tool_search: BM25 search + expand_visible()
|
||||
├─ _exec_task: _run_agent(TASK_AGENT_TOOLS)
|
||||
├─ _exec_plan: _run_agent(AGENT_TOOLS, read-only)
|
||||
├─ _exec_notify: HTTP POST to channel gateway
|
||||
├─ _exec_remember: SQLite INSERT OR REPLACE
|
||||
├─ _exec_recall: SQLite FTS5/LIKE search
|
||||
├─ _exec_forget: SQLite DELETE
|
||||
|
||||
@@ -1,3 +0,0 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:13b2da77312e5abb44c81aabb7f0addccab31d9bc7e8ef2f1c3563ef985ed503
|
||||
size 186941
|
||||
@@ -78,7 +78,19 @@ activate NodeA
|
||||
NodeA --> CC : {status:"ok", version:"0.3.0",\nmodel:"...", workstreams:{...}}
|
||||
deactivate NodeA
|
||||
|
||||
CC -> CC : Diff old vs new workstream IDs
|
||||
CC -> CC : Replace NodeSnapshot["nodeA"]\n.workstreams, .health, .aggregate
|
||||
CC -> CC : _fanout(ws_created) for\nnewly appeared workstreams
|
||||
CC -> CC : _fanout(ws_closed) for\nremoved workstreams
|
||||
|
||||
note right of CC
|
||||
Poll-diff fanout ensures
|
||||
browser SSE clients learn
|
||||
about workstreams that
|
||||
appeared without a real-time
|
||||
cluster event (e.g. bridge
|
||||
startup recovery).
|
||||
end note
|
||||
|
||||
CC -x NodeB : (SKIPPED: sim:// URL)
|
||||
|
||||
@@ -89,10 +101,15 @@ deactivate CC
|
||||
Browser -> Server : GET /v1/api/cluster/events
|
||||
activate Server
|
||||
|
||||
Server -> CC : get_snapshot()
|
||||
CC --> Server : ClusterSnapshot\n(full current state)
|
||||
|
||||
Server -> CC : register_listener(queue)
|
||||
note right : Per-client queue.Queue(maxsize=500)\nSSE via EventSourceResponse + run_in_executor()
|
||||
|
||||
loop continuous
|
||||
Server -> Browser : data: {"type":"snapshot",...}\n(full state as first SSE event)
|
||||
|
||||
loop continuous (incremental updates)
|
||||
CC -> Server : event via listener queue\n(from any of the 3 threads)
|
||||
Server -> Browser : data: {"type":"cluster_state",...}\n\n
|
||||
end
|
||||
@@ -105,6 +122,13 @@ Browser -> Server : connection closed
|
||||
Server -> CC : unregister_listener(queue)
|
||||
deactivate Server
|
||||
|
||||
== Browser REST: Snapshot ==
|
||||
|
||||
Browser -> Server : GET /v1/api/cluster/snapshot
|
||||
Server -> CC : get_snapshot()
|
||||
CC --> Server : ClusterSnapshot\n(full current state)
|
||||
Server --> Browser : JSON response
|
||||
|
||||
== Browser REST Requests ==
|
||||
|
||||
Browser -> Server : GET /v1/api/cluster/overview
|
||||
|
||||
@@ -35,7 +35,7 @@ package "turnstone/sdk/ (Python)" {
|
||||
+ stream_events(ws_id)
|
||||
+ stream_global_events()
|
||||
+ send_and_wait()
|
||||
+ list_sessions()
|
||||
+ list_saved_workstreams()
|
||||
+ login() / logout()
|
||||
+ health()
|
||||
}
|
||||
@@ -45,6 +45,7 @@ package "turnstone/sdk/ (Python)" {
|
||||
+ nodes()
|
||||
+ workstreams()
|
||||
+ node_detail()
|
||||
+ snapshot()
|
||||
+ create_workstream()
|
||||
+ stream_cluster_events()
|
||||
+ login() / logout()
|
||||
@@ -129,6 +130,7 @@ package "sdk/typescript/ (TypeScript)" {
|
||||
class "TurnstoneConsole" as TSConsole <<ts>> {
|
||||
+ overview()
|
||||
+ nodes()
|
||||
+ snapshot()
|
||||
+ clusterEvents()
|
||||
...
|
||||
}
|
||||
|
||||
@@ -13,18 +13,19 @@ skinparam class {
|
||||
|
||||
' -- Protocol --
|
||||
interface "StorageBackend" as SB <<protocol>> {
|
||||
+register_session(session_id, title)
|
||||
+save_message(session_id, role, content, ...)
|
||||
+load_session_messages(session_id) → list[dict]
|
||||
+list_sessions(limit) → list
|
||||
+delete_session(session_id) → bool
|
||||
+prune_sessions(retention_days) → (int, int)
|
||||
+resolve_session(alias_or_id) → str | None
|
||||
+save_session_config(session_id, config)
|
||||
+load_session_config(session_id) → dict
|
||||
+set_session_alias(session_id, alias) → bool
|
||||
+get_session_name(session_id) → str | None
|
||||
+update_session_title(session_id, title)
|
||||
+save_message(ws_id, role, content, ...)
|
||||
+load_messages(ws_id) → list[dict]
|
||||
+register_workstream(ws_id, node_id, name, state)
|
||||
+update_workstream_state(ws_id, state)
|
||||
+update_workstream_name(ws_id, name)
|
||||
+set_workstream_alias(ws_id, alias) → bool
|
||||
+update_workstream_title(ws_id, title)
|
||||
+resolve_workstream(alias_or_id) → str | None
|
||||
+delete_workstream(ws_id) → bool
|
||||
+prune_workstreams(retention_days) → (int, int)
|
||||
+list_workstreams(node_id, limit) → list
|
||||
+save_workstream_config(ws_id, config)
|
||||
+load_workstream_config(ws_id) → dict
|
||||
+kv_get(key) → str | None
|
||||
+kv_set(key, value) → str | None
|
||||
+kv_delete(key) → bool
|
||||
@@ -32,6 +33,11 @@ interface "StorageBackend" as SB <<protocol>> {
|
||||
+kv_search(query) → list[(str, str)]
|
||||
+search_history(query, limit) → list
|
||||
+search_history_recent(limit) → list
|
||||
+create_user(user_id, username, display_name, pw_hash)
|
||||
+get_user(user_id) / get_user_by_username(username)
|
||||
+list_users() / delete_user(user_id)
|
||||
+create_api_token(...) / get_api_token_by_hash(hash)
|
||||
+list_api_tokens(user_id) / delete_api_token(id)
|
||||
+close()
|
||||
}
|
||||
|
||||
@@ -58,8 +64,11 @@ class "_schema.py" as Schema <<schema>> {
|
||||
+metadata: MetaData
|
||||
+memories: Table
|
||||
+conversations: Table
|
||||
+sessions: Table
|
||||
+session_config: Table
|
||||
+workstreams: Table (node_id, alias, title, state)
|
||||
+workstream_config: Table
|
||||
+users: Table (username, password_hash)
|
||||
+api_tokens: Table (token_hash, scopes)
|
||||
+channel_users: Table (channel_type)
|
||||
--
|
||||
SQLAlchemy Core
|
||||
Single source of truth
|
||||
@@ -76,6 +85,7 @@ class "_migrate.py" as Migrate <<migration>> {
|
||||
|
||||
class "migrations/" as Versions <<migration>> {
|
||||
001_initial_schema.py
|
||||
002_user_identity.py
|
||||
}
|
||||
|
||||
' -- Registry --
|
||||
@@ -91,12 +101,14 @@ class "_registry.py" as Registry {
|
||||
|
||||
' -- Facade --
|
||||
class "memory.py" as Facade <<facade>> {
|
||||
+register_session()
|
||||
+save_message()
|
||||
+load_session_messages()
|
||||
+load_messages()
|
||||
+register_workstream()
|
||||
+update_workstream_state()
|
||||
+save_workstream_config()
|
||||
+save_memory() / delete_memory()
|
||||
+search_memories()
|
||||
+... (all 18 functions)
|
||||
+... (all delegated functions)
|
||||
--
|
||||
Thin delegation to
|
||||
get_storage()
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
@startuml
|
||||
!theme plain
|
||||
title Turnstone — Authentication Architecture
|
||||
|
||||
skinparam class {
|
||||
BackgroundColor<<core>> #E8EAF6
|
||||
BackgroundColor<<jwt>> #C8E6C9
|
||||
BackgroundColor<<storage>> #B3E5FC
|
||||
BackgroundColor<<endpoint>> #FFE0B2
|
||||
BackgroundColor<<scope>> #F3E5F5
|
||||
}
|
||||
|
||||
' -- Core Auth --
|
||||
class "AuthConfig" as AC <<core>> {
|
||||
+enabled: bool
|
||||
+tokens: dict[str, str]
|
||||
+check(token) → role | None
|
||||
--
|
||||
Static config-file tokens
|
||||
hmac.compare_digest
|
||||
}
|
||||
|
||||
class "AuthResult" as AR <<core>> {
|
||||
+user_id: str
|
||||
+scopes: frozenset[str]
|
||||
+token_source: str
|
||||
+has_scope(scope) → bool
|
||||
}
|
||||
|
||||
class "check_request()" as CR <<core>> {
|
||||
auth_config, method, path,
|
||||
auth_header, cookie_header,
|
||||
jwt_secret, storage
|
||||
→ (allowed, status, msg, AuthResult)
|
||||
--
|
||||
1. Auth disabled → allow
|
||||
2. Public path → allow
|
||||
3. Extract Bearer / cookie
|
||||
4. Detect token type
|
||||
5. Validate → AuthResult
|
||||
6. Check scope vs path
|
||||
}
|
||||
|
||||
' -- Token Types --
|
||||
class "JWT (HS256)" as JWT <<jwt>> {
|
||||
sub: user_id
|
||||
scopes: "read,write,approve"
|
||||
src: "password" | "database"
|
||||
iat, exp (24h default)
|
||||
--
|
||||
Detected by: contains "."
|
||||
Validated locally
|
||||
No DB call
|
||||
}
|
||||
|
||||
class "API Token" as AT <<jwt>> {
|
||||
Format: ts_ + 64 hex
|
||||
Stored: SHA-256 hash
|
||||
--
|
||||
Detected by: starts with "ts_"
|
||||
Lookup by hash in DB
|
||||
Expiry check
|
||||
}
|
||||
|
||||
class "Config Token" as CT <<core>> {
|
||||
Raw value in memory
|
||||
Role: "read" | "full"
|
||||
--
|
||||
Detected by: fallback
|
||||
hmac.compare_digest
|
||||
No DB needed
|
||||
}
|
||||
|
||||
' -- Scopes --
|
||||
class "Scope Hierarchy" as SH <<scope>> {
|
||||
read: {read}
|
||||
write: {read, write}
|
||||
approve: {read, write, approve}
|
||||
--
|
||||
GET → read
|
||||
POST write paths → write
|
||||
POST /api/approve → approve
|
||||
/api/admin/* → approve
|
||||
}
|
||||
|
||||
' -- Storage --
|
||||
class "users" as UT <<storage>> {
|
||||
user_id (PK)
|
||||
username (unique)
|
||||
display_name
|
||||
password_hash (bcrypt)
|
||||
created
|
||||
}
|
||||
|
||||
class "api_tokens" as TT <<storage>> {
|
||||
token_id (PK)
|
||||
token_hash (SHA-256, unique)
|
||||
token_prefix
|
||||
user_id → users
|
||||
name, scopes
|
||||
created, expires
|
||||
}
|
||||
|
||||
' -- Endpoints --
|
||||
class "POST /api/auth/login" as Login <<endpoint>> {
|
||||
{username, password}
|
||||
OR {token: "ts_xxx"}
|
||||
→ {jwt, role, scopes, user_id}
|
||||
--
|
||||
Sets HttpOnly cookie
|
||||
}
|
||||
|
||||
class "GET /api/auth/status" as Status <<endpoint>> {
|
||||
→ {auth_enabled, has_users,
|
||||
setup_required}
|
||||
--
|
||||
Public (no auth)
|
||||
Drives UI setup wizard
|
||||
}
|
||||
|
||||
class "POST /api/auth/setup" as Setup <<endpoint>> {
|
||||
{username, display_name, password}
|
||||
→ {jwt, user_id, scopes}
|
||||
--
|
||||
Public (no auth)
|
||||
Only when zero users exist
|
||||
Returns 409 if already set up
|
||||
}
|
||||
|
||||
class "Admin API (Console)" as Admin <<endpoint>> {
|
||||
POST/GET/DELETE users
|
||||
POST/GET tokens
|
||||
DELETE tokens/{id}
|
||||
--
|
||||
Requires approve scope
|
||||
}
|
||||
|
||||
' -- Relationships --
|
||||
CR --> AC : config tokens
|
||||
CR --> JWT : validate
|
||||
CR --> AT : hash lookup
|
||||
CR --> CT : hmac check
|
||||
CR --> AR : returns
|
||||
CR --> SH : checks
|
||||
|
||||
Login --> JWT : issues
|
||||
Login --> UT : verify password
|
||||
Login --> TT : verify API token
|
||||
|
||||
Setup --> UT : create first user
|
||||
Setup --> JWT : issues
|
||||
|
||||
AT --> TT : lookup by hash
|
||||
Admin --> UT : CRUD
|
||||
Admin --> TT : CRUD
|
||||
|
||||
AR --> SH : scopes from
|
||||
|
||||
JWT ..> AR : produces
|
||||
AT ..> AR : produces
|
||||
CT ..> AR : produces
|
||||
|
||||
note right of CR
|
||||
**Middleware Flow**
|
||||
AuthMiddleware on every request:
|
||||
1. Extract token from header/cookie
|
||||
2. Detect type (JWT / ts_ / config)
|
||||
3. Validate → AuthResult
|
||||
4. Set ctx_user_id for logging
|
||||
5. Store auth_result in scope state
|
||||
end note
|
||||
|
||||
note bottom of SH
|
||||
**Console** owns admin endpoints
|
||||
**Server** validates JWT + config only
|
||||
Both share JWT signing secret
|
||||
end note
|
||||
|
||||
@enduml
|
||||
@@ -0,0 +1,250 @@
|
||||
@startuml
|
||||
!theme plain
|
||||
title Turnstone — Channel Integration Architecture
|
||||
|
||||
skinparam class {
|
||||
BackgroundColor<<platform>> #E1BEE7
|
||||
BackgroundColor<<service>> #E8EAF6
|
||||
BackgroundColor<<mq>> #FFCDD2
|
||||
BackgroundColor<<bridge>> #C8E6C9
|
||||
BackgroundColor<<server>> #FFE0B2
|
||||
BackgroundColor<<storage>> #B3E5FC
|
||||
}
|
||||
|
||||
' -- External Platforms --
|
||||
class "Discord" as Discord <<platform>> {
|
||||
Gateway WebSocket (v10)
|
||||
Message events
|
||||
Interaction callbacks (buttons)
|
||||
Thread-per-workstream
|
||||
--
|
||||
discord.py 2.x
|
||||
asyncio event loop
|
||||
}
|
||||
|
||||
class "Slack (future)" as Slack <<platform>> {
|
||||
Socket Mode / Events API
|
||||
Block Kit messages
|
||||
--
|
||||
Planned integration
|
||||
}
|
||||
|
||||
class "Teams (future)" as Teams <<platform>> {
|
||||
Bot Framework
|
||||
Adaptive Cards
|
||||
--
|
||||
Planned integration
|
||||
}
|
||||
|
||||
' -- Channel Service --
|
||||
class "turnstone-channel" as ChannelService <<service>> {
|
||||
entry point: turnstone-channel
|
||||
--
|
||||
One process per platform
|
||||
asyncio event loop
|
||||
Structured logging (structlog)
|
||||
--log-level, --log-format
|
||||
--
|
||||
POST /v1/api/notify (HTTP)
|
||||
GET /health
|
||||
}
|
||||
|
||||
class "DiscordBot" as Bot <<service>> {
|
||||
+on_message(msg)
|
||||
+on_interaction(interaction)
|
||||
+send(channel_id, content)
|
||||
+run(token)
|
||||
--
|
||||
discord.py Client
|
||||
Receives message events
|
||||
Sends replies + embeds
|
||||
Creates threads for workstreams
|
||||
Renders approval buttons
|
||||
escape_mentions() on send
|
||||
}
|
||||
|
||||
class "ChannelRouter" as Router <<service>> {
|
||||
+resolve_route(platform, channel_id)
|
||||
→ ws_id | None
|
||||
+register_route(channel_id, ws_id)
|
||||
+resolve_identity(platform, platform_user_id)
|
||||
→ user_id | None
|
||||
--
|
||||
Maps channels → workstreams
|
||||
Maps platform users → turnstone users
|
||||
Caches routes in memory
|
||||
}
|
||||
|
||||
class "AsyncRedisBroker" as Broker <<service>> {
|
||||
+push_inbound(msg)
|
||||
+subscribe(ws_id) → AsyncIterator
|
||||
+subscribe_global() → AsyncIterator
|
||||
+push_response(correlation_id, msg)
|
||||
--
|
||||
redis.asyncio client
|
||||
Pub/sub + queue operations
|
||||
}
|
||||
|
||||
' -- Redis MQ --
|
||||
class "Redis MQ" as Redis <<mq>> {
|
||||
turnstone:inbound (LIST)
|
||||
turnstone:events:{ws_id} (PUBSUB)
|
||||
turnstone:events:global (PUBSUB)
|
||||
turnstone:resp:{corr_id} (LIST)
|
||||
--
|
||||
Shared message bus
|
||||
Same queues as bridge protocol
|
||||
}
|
||||
|
||||
' -- Bridge + Server --
|
||||
class "turnstone-bridge" as Bridge <<bridge>> {
|
||||
BLPOP turnstone:inbound
|
||||
Drive server via HTTP
|
||||
Relay SSE → Redis pub/sub
|
||||
--
|
||||
Owns workstream lifecycle
|
||||
Auto-approve / manual approve
|
||||
}
|
||||
|
||||
class "turnstone-server" as Server <<server>> {
|
||||
POST /v1/api/send
|
||||
POST /v1/api/approve
|
||||
POST /v1/api/workstreams/new
|
||||
GET /v1/api/events?ws_id=
|
||||
--
|
||||
LLM execution + tool use
|
||||
SSE event stream
|
||||
--
|
||||
notify tool: _exec_notify()
|
||||
ServiceTokenManager (JWT)
|
||||
}
|
||||
|
||||
' -- Storage --
|
||||
class "channel_users" as CU <<storage>> {
|
||||
channel_user_id (PK)
|
||||
platform: "discord" | "slack"
|
||||
platform_user_id
|
||||
user_id → users
|
||||
linked_at
|
||||
--
|
||||
/link command creates row
|
||||
Resolved on each inbound message
|
||||
}
|
||||
|
||||
class "channel_routes" as CR <<storage>> {
|
||||
channel_type (PK)
|
||||
channel_id (PK)
|
||||
ws_id
|
||||
node_id
|
||||
created
|
||||
--
|
||||
Maps platform channels
|
||||
to turnstone workstreams
|
||||
}
|
||||
|
||||
class "services" as SVC <<storage>> {
|
||||
service_type (PK)
|
||||
service_id (PK)
|
||||
url
|
||||
last_heartbeat
|
||||
created
|
||||
--
|
||||
Heartbeat every 30s
|
||||
Stale after 120s
|
||||
ON CONFLICT DO UPDATE
|
||||
}
|
||||
|
||||
' -- Relationships --
|
||||
Discord --> Bot : gateway\nevents
|
||||
Bot --> Router : on_message\non_interaction
|
||||
Router --> Broker : SendMessage\nApproveMessage
|
||||
Router --> CU : resolve identity
|
||||
Router --> CR : resolve / register route
|
||||
Broker --> Redis : RPUSH inbound\nRPUSH resp:{id}
|
||||
|
||||
Redis --> Bridge : BLPOP inbound
|
||||
Bridge --> Server : HTTP API
|
||||
Server --> Bridge : SSE events
|
||||
Bridge --> Redis : PUBLISH events:{ws_id}\nPUBLISH events:global
|
||||
|
||||
Redis --> Broker : SUBSCRIBE events:{ws_id}
|
||||
Broker --> Bot : event stream
|
||||
Bot --> Discord : reply / embed\nbutton callback
|
||||
|
||||
Slack .[hidden]. Discord
|
||||
Teams .[hidden]. Slack
|
||||
|
||||
ChannelService --> Bot : creates + runs
|
||||
ChannelService --> Router : creates
|
||||
ChannelService --> Broker : creates
|
||||
ChannelService --> SVC : register / heartbeat /\nderegister
|
||||
|
||||
' -- Notification path (direct HTTP, bypasses MQ) --
|
||||
Server --> ChannelService : POST /v1/api/notify\n(JWT: aud=turnstone-channel)
|
||||
Server --> SVC : list_services("channel",\nmax_age_seconds=120)
|
||||
|
||||
' -- Notes --
|
||||
note right of Bot
|
||||
**Inbound Flow**
|
||||
1. Discord message arrives via gateway
|
||||
2. Bot.on_message() fires
|
||||
3. ChannelRouter resolves channel → ws_id
|
||||
(or creates new workstream)
|
||||
4. ChannelRouter resolves platform user → user_id
|
||||
via channel_users table
|
||||
5. Broker.push_inbound(SendMessage)
|
||||
6. Bridge pops from Redis, drives server
|
||||
|
||||
**Workstream Resume (evicted workstreams)**
|
||||
1. Stale route detected (no MQ owner)
|
||||
2. Existing ws_id reused directly from route
|
||||
3. CreateWorkstreamMessage sent with
|
||||
resume_ws=<ws_id>
|
||||
4. Server resumes atomically during creation
|
||||
5. Bridge emits WorkstreamResumedEvent → thread
|
||||
end note
|
||||
|
||||
note right of Broker
|
||||
**Outbound Flow**
|
||||
1. Server emits SSE events
|
||||
2. Bridge relays to Redis events:{ws_id}
|
||||
3. Broker.subscribe(ws_id) yields events
|
||||
4. Bot formats and sends to Discord thread
|
||||
end note
|
||||
|
||||
note bottom of CR
|
||||
**Approval Flow**
|
||||
1. ApprovalRequestEvent arrives via events:{ws_id}
|
||||
2. Bot renders Discord buttons (Approve / Deny)
|
||||
3. User clicks button → on_interaction()
|
||||
4. Router builds ApproveMessage
|
||||
5. Broker.push_response(correlation_id, msg)
|
||||
6. Bridge pops from resp:{id}, calls POST /api/approve
|
||||
end note
|
||||
|
||||
note bottom of CU
|
||||
**Identity Linking**
|
||||
1. User runs /link in Discord
|
||||
2. Bot opens modal requesting API token
|
||||
3. User submits ts_... API token
|
||||
4. Bot validates token against storage
|
||||
5. On success, inserts channel_users row
|
||||
6. Subsequent messages carry user_id
|
||||
7. AuthResult scopes applied by server
|
||||
end note
|
||||
|
||||
note bottom of SVC
|
||||
**Notification Flow** (direct HTTP, bypasses MQ)
|
||||
1. LLM calls notify tool → _prepare_notify()
|
||||
2. _exec_notify() checks rate limit (5/turn)
|
||||
3. Queries services table for healthy gateways
|
||||
4. Mints JWT (aud: turnstone-channel) via
|
||||
ServiceTokenManager
|
||||
5. POSTs to first healthy gateway
|
||||
6. Gateway validates JWT, resolves target
|
||||
7. adapter.send() → Discord API
|
||||
8. On failure: retry up to 3× (1s, 3s backoff)
|
||||
9. SSRF: only http(s) URLs allowed
|
||||
end note
|
||||
|
||||
@enduml
|
||||
@@ -0,0 +1,116 @@
|
||||
@startuml
|
||||
!theme plain
|
||||
title Turnstone — Notification Delivery Flow
|
||||
|
||||
skinparam participant {
|
||||
BackgroundColor<<server>> #FFE0B2
|
||||
BackgroundColor<<storage>> #B3E5FC
|
||||
BackgroundColor<<service>> #E8EAF6
|
||||
BackgroundColor<<platform>> #E1BEE7
|
||||
}
|
||||
|
||||
participant "ChatSession\n(turnstone-server)" as Session <<server>>
|
||||
participant "StorageBackend" as Storage <<storage>>
|
||||
participant "ServiceTokenManager" as STM <<server>>
|
||||
participant "Channel Gateway\n(_http.py)" as Gateway <<service>>
|
||||
participant "ChannelAdapter\n(Discord bot)" as Adapter <<service>>
|
||||
participant "Discord API" as Discord <<platform>>
|
||||
|
||||
== Prepare Phase ==
|
||||
|
||||
Session -> Session : _prepare_notify(call_id, args)
|
||||
note right
|
||||
Validates:
|
||||
- message (required, ≤2000 chars)
|
||||
- target: username OR channel_type+channel_id
|
||||
- no ambiguous targeting (both set)
|
||||
- partial targeting errors
|
||||
end note
|
||||
|
||||
== Execute Phase ==
|
||||
|
||||
Session -> Session : _exec_notify(item)
|
||||
Session -> Session : check rate limit\n(≥5 per turn?)
|
||||
|
||||
alt rate limit exceeded
|
||||
Session --> Session : "Error: rate limit exceeded"
|
||||
end
|
||||
|
||||
loop up to 3 attempts (retry delays: 1s, 3s)
|
||||
|
||||
Session -> Storage : list_services("channel",\nmax_age_seconds=120)
|
||||
Storage --> Session : services[] (sorted by\nlast_heartbeat DESC)
|
||||
|
||||
alt no healthy services
|
||||
Session -> Session : log.warning("notify.no_services")
|
||||
Session -> Session : sleep(delay)
|
||||
else services available
|
||||
|
||||
Session -> STM : bearer_header
|
||||
note right
|
||||
Lazy-init ServiceTokenManager
|
||||
aud: turnstone-channel
|
||||
scope: write
|
||||
Auto-rotates 1h JWTs
|
||||
end note
|
||||
STM --> Session : Authorization: Bearer <jwt>
|
||||
|
||||
loop for each gateway (first-healthy)
|
||||
Session -> Session : SSRF check:\nurl.startswith("http://"|"https://")
|
||||
|
||||
Session -> Gateway : POST /v1/api/notify\n+ Authorization header
|
||||
Gateway -> Gateway : _check_auth()\nvalidate JWT (aud=turnstone-channel)\nor static token
|
||||
|
||||
alt auth failed
|
||||
Gateway --> Session : 401 Unauthorized
|
||||
else auth ok
|
||||
|
||||
alt username target
|
||||
Gateway -> Storage : get_user_by_username()
|
||||
Storage --> Gateway : user
|
||||
Gateway -> Storage : list_channel_users_by_user()
|
||||
Storage --> Gateway : linked channels
|
||||
else direct target
|
||||
Gateway -> Gateway : use channel_type + channel_id
|
||||
end
|
||||
|
||||
Gateway -> Adapter : send(channel_id, content)
|
||||
note right
|
||||
escape_mentions() applied
|
||||
Chunked for 2000-char limit
|
||||
end note
|
||||
Adapter -> Discord : POST message
|
||||
Discord --> Adapter : message_id
|
||||
Adapter --> Gateway : message_id
|
||||
Gateway --> Session : 200 {results: [{status: "sent"}]}
|
||||
|
||||
Session -> Session : _notify_count += 1
|
||||
Session --> Session : "Notification sent successfully"
|
||||
note right : Return — no further\ngateways tried
|
||||
end
|
||||
end
|
||||
|
||||
alt all gateways failed
|
||||
Session -> Session : log.warning(\n"notify.all_gateways_failed")
|
||||
Session -> Session : sleep(delay)
|
||||
end
|
||||
|
||||
end
|
||||
end
|
||||
|
||||
alt all retries exhausted
|
||||
Session -> Session : log.warning("notify.delivery_failed")
|
||||
Session --> Session : "Error: notification delivery failed"
|
||||
end
|
||||
|
||||
== Service Registry (Background) ==
|
||||
|
||||
note over Gateway, Storage
|
||||
**Heartbeat Lifecycle**
|
||||
1. Gateway startup: register_service("channel", id, url)
|
||||
2. Every 30s: heartbeat_service("channel", id)
|
||||
3. Shutdown: deregister_service("channel", id)
|
||||
4. Stale after 120s (4 missed heartbeats)
|
||||
end note
|
||||
|
||||
@enduml
|
||||
@@ -0,0 +1,286 @@
|
||||
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 1200 540" font-family="-apple-system, BlinkMacSystemFont, 'Segoe UI', Helvetica, Arial, sans-serif">
|
||||
<defs>
|
||||
<!-- Arrowhead markers -->
|
||||
<marker id="arrow" viewBox="0 0 10 7" refX="10" refY="3.5" markerWidth="8" markerHeight="6" orient="auto-start-auto">
|
||||
<path d="M 0 0 L 10 3.5 L 0 7 z" fill="#484f58"/>
|
||||
</marker>
|
||||
<marker id="arrow-blue" viewBox="0 0 10 7" refX="10" refY="3.5" markerWidth="8" markerHeight="6" orient="auto-start-auto">
|
||||
<path d="M 0 0 L 10 3.5 L 0 7 z" fill="#58a6ff"/>
|
||||
</marker>
|
||||
<marker id="arrow-green" viewBox="0 0 10 7" refX="10" refY="3.5" markerWidth="8" markerHeight="6" orient="auto-start-auto">
|
||||
<path d="M 0 0 L 10 3.5 L 0 7 z" fill="#3fb950"/>
|
||||
</marker>
|
||||
<marker id="arrow-orange" viewBox="0 0 10 7" refX="10" refY="3.5" markerWidth="8" markerHeight="6" orient="auto-start-auto">
|
||||
<path d="M 0 0 L 10 3.5 L 0 7 z" fill="#f0883e"/>
|
||||
</marker>
|
||||
<marker id="arrow-coral" viewBox="0 0 10 7" refX="10" refY="3.5" markerWidth="8" markerHeight="6" orient="auto-start-auto">
|
||||
<path d="M 0 0 L 10 3.5 L 0 7 z" fill="#f47067"/>
|
||||
</marker>
|
||||
<marker id="arrow-muted" viewBox="0 0 10 7" refX="10" refY="3.5" markerWidth="8" markerHeight="6" orient="auto-start-auto">
|
||||
<path d="M 0 0 L 10 3.5 L 0 7 z" fill="#8b949e"/>
|
||||
</marker>
|
||||
|
||||
<!-- Card shadow filter -->
|
||||
<filter id="shadow" x="-4%" y="-4%" width="108%" height="112%">
|
||||
<feDropShadow dx="0" dy="1" stdDeviation="2" flood-color="#000" flood-opacity="0.4"/>
|
||||
</filter>
|
||||
</defs>
|
||||
|
||||
<!-- Background -->
|
||||
<rect width="1200" height="540" rx="8" fill="#0d1117"/>
|
||||
|
||||
<!-- Title -->
|
||||
<text x="600" y="36" text-anchor="middle" fill="#e6edf3" font-size="15" font-weight="700" letter-spacing="3">TURNSTONE</text>
|
||||
<text x="600" y="54" text-anchor="middle" fill="#8b949e" font-size="11" letter-spacing="1">SYSTEM ARCHITECTURE</text>
|
||||
|
||||
<!-- ==================== COLUMN HEADERS ==================== -->
|
||||
<text x="90" y="86" text-anchor="middle" fill="#58a6ff" font-size="9" font-weight="600" letter-spacing="2">CLIENTS</text>
|
||||
<text x="276" y="86" text-anchor="middle" fill="#3fb950" font-size="9" font-weight="600" letter-spacing="2">GATEWAYS</text>
|
||||
<text x="480" y="86" text-anchor="middle" fill="#f0883e" font-size="9" font-weight="600" letter-spacing="2">MESSAGE QUEUE</text>
|
||||
<text x="700" y="86" text-anchor="middle" fill="#f47067" font-size="9" font-weight="600" letter-spacing="2">CLUSTER NODES</text>
|
||||
<text x="940" y="86" text-anchor="middle" fill="#f778ba" font-size="9" font-weight="600" letter-spacing="2">LLM PROVIDERS</text>
|
||||
|
||||
<!-- ==================== CLIENT BOXES ==================== -->
|
||||
<!-- CLI -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="30" y="108" width="120" height="46" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="30" y="108" width="120" height="5" rx="5" fill="#58a6ff"/>
|
||||
<rect x="30" y="108" width="120" height="5" fill="#58a6ff"/>
|
||||
<rect x="30" y="111" width="120" height="2" fill="#161b22"/>
|
||||
<text x="90" y="130" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">CLI</text>
|
||||
<text x="90" y="145" text-anchor="middle" fill="#8b949e" font-size="9">terminal REPL</text>
|
||||
</g>
|
||||
|
||||
<!-- Browser UI -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="30" y="174" width="120" height="46" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="30" y="174" width="120" height="5" rx="5" fill="#58a6ff"/>
|
||||
<rect x="30" y="174" width="120" height="5" fill="#58a6ff"/>
|
||||
<rect x="30" y="177" width="120" height="2" fill="#161b22"/>
|
||||
<text x="90" y="196" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">Browser UI</text>
|
||||
<text x="90" y="211" text-anchor="middle" fill="#8b949e" font-size="9">HTTP + SSE</text>
|
||||
</g>
|
||||
|
||||
<!-- SDK / API -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="30" y="244" width="120" height="46" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="30" y="244" width="120" height="5" rx="5" fill="#58a6ff"/>
|
||||
<rect x="30" y="244" width="120" height="5" fill="#58a6ff"/>
|
||||
<rect x="30" y="247" width="120" height="2" fill="#161b22"/>
|
||||
<text x="90" y="266" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">SDK / API</text>
|
||||
<text x="90" y="281" text-anchor="middle" fill="#8b949e" font-size="9">programmatic</text>
|
||||
</g>
|
||||
|
||||
<!-- Discord / Slack -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="30" y="314" width="120" height="46" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="30" y="314" width="120" height="5" rx="5" fill="#58a6ff"/>
|
||||
<rect x="30" y="314" width="120" height="5" fill="#58a6ff"/>
|
||||
<rect x="30" y="317" width="120" height="2" fill="#161b22"/>
|
||||
<text x="90" y="336" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">Discord / Slack</text>
|
||||
<text x="90" y="351" text-anchor="middle" fill="#8b949e" font-size="9">chat platforms</text>
|
||||
</g>
|
||||
|
||||
<!-- ==================== GATEWAY BOXES ==================== -->
|
||||
<!-- Console -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="216" y="118" width="120" height="58" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="216" y="118" width="120" height="5" rx="5" fill="#3fb950"/>
|
||||
<rect x="216" y="118" width="120" height="5" fill="#3fb950"/>
|
||||
<rect x="216" y="121" width="120" height="2" fill="#161b22"/>
|
||||
<text x="276" y="142" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">Console</text>
|
||||
<text x="276" y="157" text-anchor="middle" fill="#8b949e" font-size="9">dashboard + proxy</text>
|
||||
<text x="276" y="169" text-anchor="middle" fill="#8b949e" font-size="9">cluster management</text>
|
||||
</g>
|
||||
|
||||
<!-- Channel Gateway -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="216" y="292" width="120" height="58" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="216" y="292" width="120" height="5" rx="5" fill="#3fb950"/>
|
||||
<rect x="216" y="292" width="120" height="5" fill="#3fb950"/>
|
||||
<rect x="216" y="295" width="120" height="2" fill="#161b22"/>
|
||||
<text x="276" y="316" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">Channel Gateway</text>
|
||||
<text x="276" y="331" text-anchor="middle" fill="#8b949e" font-size="9">platform adapter</text>
|
||||
<text x="276" y="343" text-anchor="middle" fill="#8b949e" font-size="9">Discord, Slack, ...</text>
|
||||
</g>
|
||||
|
||||
<!-- ==================== REDIS MQ ==================== -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="420" y="168" width="120" height="132" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="420" y="168" width="120" height="5" rx="5" fill="#f0883e"/>
|
||||
<rect x="420" y="168" width="120" height="5" fill="#f0883e"/>
|
||||
<rect x="420" y="171" width="120" height="2" fill="#161b22"/>
|
||||
<text x="480" y="198" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">Redis MQ</text>
|
||||
<line x1="438" y1="210" x2="522" y2="210" stroke="#30363d" stroke-width="1"/>
|
||||
<text x="480" y="228" text-anchor="middle" fill="#8b949e" font-size="9">inbound queues</text>
|
||||
<text x="480" y="243" text-anchor="middle" fill="#8b949e" font-size="9">event pub/sub</text>
|
||||
<text x="480" y="258" text-anchor="middle" fill="#8b949e" font-size="9">node heartbeats</text>
|
||||
<text x="480" y="273" text-anchor="middle" fill="#8b949e" font-size="9">workstream routing</text>
|
||||
<text x="480" y="288" text-anchor="middle" fill="#8b949e" font-size="9">cluster state</text>
|
||||
</g>
|
||||
|
||||
<!-- ==================== CLUSTER NODES ==================== -->
|
||||
<!-- Cluster outline -->
|
||||
<rect x="598" y="100" width="204" height="310" rx="8" fill="none" stroke="#30363d" stroke-width="1" stroke-dasharray="4,3"/>
|
||||
<text x="700" y="422" text-anchor="middle" fill="#30363d" font-size="9" letter-spacing="1">CLUSTER</text>
|
||||
|
||||
<!-- Node A -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="614" y="118" width="170" height="100" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="614" y="118" width="170" height="5" rx="5" fill="#f47067"/>
|
||||
<rect x="614" y="118" width="170" height="5" fill="#f47067"/>
|
||||
<rect x="614" y="121" width="170" height="2" fill="#161b22"/>
|
||||
<text x="699" y="142" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">Node A</text>
|
||||
<line x1="632" y1="152" x2="766" y2="152" stroke="#30363d" stroke-width="1"/>
|
||||
<!-- Bridge -->
|
||||
<rect x="626" y="162" width="70" height="28" rx="3" fill="#1c2128" stroke="#30363d" stroke-width="1"/>
|
||||
<text x="661" y="180" text-anchor="middle" fill="#8b949e" font-size="9">bridge</text>
|
||||
<!-- Server -->
|
||||
<rect x="704" y="162" width="70" height="28" rx="3" fill="#1c2128" stroke="#30363d" stroke-width="1"/>
|
||||
<text x="739" y="180" text-anchor="middle" fill="#8b949e" font-size="9">server</text>
|
||||
<!-- Arrow bridge to server -->
|
||||
<line x1="696" y1="176" x2="702" y2="176" stroke="#484f58" stroke-width="1" marker-end="url(#arrow)"/>
|
||||
<!-- Tools label -->
|
||||
<text x="699" y="206" text-anchor="middle" fill="#484f58" font-size="8">14 tools + MCP</text>
|
||||
</g>
|
||||
|
||||
<!-- Node B -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="614" y="238" width="170" height="100" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="614" y="238" width="170" height="5" rx="5" fill="#f47067"/>
|
||||
<rect x="614" y="238" width="170" height="5" fill="#f47067"/>
|
||||
<rect x="614" y="241" width="170" height="2" fill="#161b22"/>
|
||||
<text x="699" y="262" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">Node B</text>
|
||||
<line x1="632" y1="272" x2="766" y2="272" stroke="#30363d" stroke-width="1"/>
|
||||
<!-- Bridge -->
|
||||
<rect x="626" y="282" width="70" height="28" rx="3" fill="#1c2128" stroke="#30363d" stroke-width="1"/>
|
||||
<text x="661" y="300" text-anchor="middle" fill="#8b949e" font-size="9">bridge</text>
|
||||
<!-- Server -->
|
||||
<rect x="704" y="282" width="70" height="28" rx="3" fill="#1c2128" stroke="#30363d" stroke-width="1"/>
|
||||
<text x="739" y="300" text-anchor="middle" fill="#8b949e" font-size="9">server</text>
|
||||
<!-- Arrow bridge to server -->
|
||||
<line x1="696" y1="296" x2="702" y2="296" stroke="#484f58" stroke-width="1" marker-end="url(#arrow)"/>
|
||||
<!-- Tools label -->
|
||||
<text x="699" y="326" text-anchor="middle" fill="#484f58" font-size="8">14 tools + MCP</text>
|
||||
</g>
|
||||
|
||||
<!-- ==================== LLM PROVIDERS ==================== -->
|
||||
<!-- OpenAI -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="870" y="130" width="140" height="46" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="870" y="130" width="140" height="5" rx="5" fill="#f778ba"/>
|
||||
<rect x="870" y="130" width="140" height="5" fill="#f778ba"/>
|
||||
<rect x="870" y="133" width="140" height="2" fill="#161b22"/>
|
||||
<text x="940" y="153" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">OpenAI</text>
|
||||
<text x="940" y="167" text-anchor="middle" fill="#8b949e" font-size="9">GPT-5, o-series</text>
|
||||
</g>
|
||||
|
||||
<!-- Anthropic -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="870" y="196" width="140" height="46" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="870" y="196" width="140" height="5" rx="5" fill="#f778ba"/>
|
||||
<rect x="870" y="196" width="140" height="5" fill="#f778ba"/>
|
||||
<rect x="870" y="199" width="140" height="2" fill="#161b22"/>
|
||||
<text x="940" y="219" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">Anthropic</text>
|
||||
<text x="940" y="233" text-anchor="middle" fill="#8b949e" font-size="9">Claude 4.5 / 4.6</text>
|
||||
</g>
|
||||
|
||||
<!-- Local / vLLM -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="870" y="262" width="140" height="46" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="870" y="262" width="140" height="5" rx="5" fill="#f778ba"/>
|
||||
<rect x="870" y="262" width="140" height="5" fill="#f778ba"/>
|
||||
<rect x="870" y="265" width="140" height="2" fill="#161b22"/>
|
||||
<text x="940" y="285" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">Local / vLLM</text>
|
||||
<text x="940" y="299" text-anchor="middle" fill="#8b949e" font-size="9">llama.cpp, NIM</text>
|
||||
</g>
|
||||
|
||||
<!-- ==================== STORAGE ==================== -->
|
||||
<g filter="url(#shadow)">
|
||||
<rect x="614" y="450" width="170" height="52" rx="5" fill="#161b22" stroke="#30363d" stroke-width="1"/>
|
||||
<rect x="614" y="450" width="170" height="5" rx="5" fill="#bc8cff"/>
|
||||
<rect x="614" y="450" width="170" height="5" fill="#bc8cff"/>
|
||||
<rect x="614" y="453" width="170" height="2" fill="#161b22"/>
|
||||
<text x="699" y="476" text-anchor="middle" fill="#e6edf3" font-size="11" font-weight="600">PostgreSQL / SQLite</text>
|
||||
<text x="699" y="492" text-anchor="middle" fill="#8b949e" font-size="9">conversations, memory, auth</text>
|
||||
</g>
|
||||
<text x="699" y="444" text-anchor="middle" fill="#bc8cff" font-size="9" font-weight="600" letter-spacing="2">STORAGE</text>
|
||||
|
||||
<!-- ==================== CONNECTION LINES ==================== -->
|
||||
|
||||
<!-- CLIENT -> GATEWAY connections -->
|
||||
<!-- Browser -> Console -->
|
||||
<line x1="150" y1="197" x2="214" y2="155" stroke="#58a6ff" stroke-width="1.2" stroke-opacity="0.6" marker-end="url(#arrow-blue)"/>
|
||||
<!-- Discord -> Channel -->
|
||||
<line x1="150" y1="337" x2="214" y2="325" stroke="#58a6ff" stroke-width="1.2" stroke-opacity="0.6" marker-end="url(#arrow-blue)"/>
|
||||
|
||||
<!-- CLI -> direct to Node A server (top path, curved) -->
|
||||
<path d="M 150 131 C 200 131, 200 100, 400 100 L 400 100 C 500 100, 570 140, 612 168" stroke="#58a6ff" stroke-width="1.2" stroke-opacity="0.5" fill="none" stroke-dasharray="6,3" marker-end="url(#arrow-blue)"/>
|
||||
<text x="370" y="96" fill="#484f58" font-size="8" text-anchor="middle">direct</text>
|
||||
|
||||
<!-- SDK -> Redis (direct push) -->
|
||||
<line x1="150" y1="267" x2="418" y2="240" stroke="#58a6ff" stroke-width="1.2" stroke-opacity="0.6" marker-end="url(#arrow-blue)"/>
|
||||
|
||||
<!-- GATEWAY -> REDIS connections -->
|
||||
<!-- Console -> Redis -->
|
||||
<line x1="336" y1="160" x2="418" y2="200" stroke="#3fb950" stroke-width="1.2" stroke-opacity="0.6" marker-end="url(#arrow-green)"/>
|
||||
<!-- Channel -> Redis -->
|
||||
<line x1="336" y1="318" x2="418" y2="272" stroke="#3fb950" stroke-width="1.2" stroke-opacity="0.6" marker-end="url(#arrow-green)"/>
|
||||
|
||||
<!-- REDIS -> NODE connections -->
|
||||
<!-- Redis -> Node A bridge -->
|
||||
<line x1="540" y1="210" x2="624" y2="176" stroke="#f0883e" stroke-width="1.2" stroke-opacity="0.6" marker-end="url(#arrow-orange)"/>
|
||||
<!-- Redis -> Node B bridge -->
|
||||
<line x1="540" y1="260" x2="624" y2="296" stroke="#f0883e" stroke-width="1.2" stroke-opacity="0.6" marker-end="url(#arrow-orange)"/>
|
||||
|
||||
<!-- Console -> Node (proxy, dashed) -->
|
||||
<path d="M 336 147 C 380 130, 500 108, 612 145" stroke="#3fb950" stroke-width="1" stroke-opacity="0.4" fill="none" stroke-dasharray="4,3" marker-end="url(#arrow-green)"/>
|
||||
<text x="468" y="120" fill="#484f58" font-size="8" text-anchor="middle">proxy</text>
|
||||
|
||||
<!-- NODE -> LLM connections -->
|
||||
<!-- Node A -> LLM providers -->
|
||||
<line x1="784" y1="168" x2="868" y2="155" stroke="#f47067" stroke-width="1.2" stroke-opacity="0.5" marker-end="url(#arrow-coral)"/>
|
||||
<line x1="784" y1="176" x2="868" y2="219" stroke="#f47067" stroke-width="1.2" stroke-opacity="0.3"/>
|
||||
<line x1="784" y1="180" x2="868" y2="282" stroke="#f47067" stroke-width="1.2" stroke-opacity="0.2"/>
|
||||
|
||||
<!-- Node B -> LLM providers -->
|
||||
<line x1="784" y1="288" x2="868" y2="163" stroke="#f47067" stroke-width="1.2" stroke-opacity="0.2"/>
|
||||
<line x1="784" y1="296" x2="868" y2="222" stroke="#f47067" stroke-width="1.2" stroke-opacity="0.3"/>
|
||||
<line x1="784" y1="300" x2="868" y2="288" stroke="#f47067" stroke-width="1.2" stroke-opacity="0.5" marker-end="url(#arrow-coral)"/>
|
||||
|
||||
<!-- NODE -> STORAGE connections -->
|
||||
<line x1="680" y1="338" x2="680" y2="448" stroke="#bc8cff" stroke-width="1.2" stroke-opacity="0.4" stroke-dasharray="4,3" marker-end="url(#arrow-muted)"/>
|
||||
<line x1="718" y1="218" x2="718" y2="236" stroke="#484f58" stroke-width="1" stroke-opacity="0.3" stroke-dasharray="2,2"/>
|
||||
|
||||
<!-- Extensibility hint -->
|
||||
<text x="699" y="392" text-anchor="middle" fill="#30363d" font-size="10">...</text>
|
||||
|
||||
<!-- Event flow: Bridges -> Redis (dashed, bidirectional feel) -->
|
||||
<line x1="624" y1="186" x2="542" y2="220" stroke="#f0883e" stroke-width="1" stroke-opacity="0.3" stroke-dasharray="3,3"/>
|
||||
<line x1="624" y1="286" x2="542" y2="250" stroke="#f0883e" stroke-width="1" stroke-opacity="0.3" stroke-dasharray="3,3"/>
|
||||
<text x="574" y="242" fill="#484f58" font-size="7" text-anchor="middle">events</text>
|
||||
|
||||
<!-- ==================== FLOW LABELS ==================== -->
|
||||
<!-- Interactive flow label -->
|
||||
<rect x="30" y="395" width="10" height="10" rx="2" fill="none" stroke="#58a6ff" stroke-width="1.5" stroke-dasharray="3,2"/>
|
||||
<text x="46" y="404" fill="#8b949e" font-size="9">interactive (direct)</text>
|
||||
|
||||
<!-- Queue flow label -->
|
||||
<rect x="160" y="395" width="10" height="10" rx="2" fill="none" stroke="#f0883e" stroke-width="1.5"/>
|
||||
<text x="176" y="404" fill="#8b949e" font-size="9">queue-driven</text>
|
||||
|
||||
<!-- Proxy/event label -->
|
||||
<rect x="275" y="395" width="10" height="10" rx="2" fill="none" stroke="#3fb950" stroke-width="1.5" stroke-dasharray="3,2"/>
|
||||
<text x="291" y="404" fill="#8b949e" font-size="9">proxy / events</text>
|
||||
|
||||
<!-- ==================== BOTTOM DETAILS ==================== -->
|
||||
<line x1="30" y1="430" x2="1170" y2="430" stroke="#21262d" stroke-width="1"/>
|
||||
|
||||
<!-- Routing rules at bottom, left-aligned -->
|
||||
<text x="44" y="456" fill="#30363d" font-size="9" font-weight="600" letter-spacing="1">ROUTING</text>
|
||||
<circle cx="44" cy="474" r="3" fill="#f47067" opacity="0.6"/>
|
||||
<text x="54" y="477" fill="#484f58" font-size="9">target_node set → route to specific node queue</text>
|
||||
<circle cx="44" cy="494" r="3" fill="#f0883e" opacity="0.6"/>
|
||||
<text x="54" y="497" fill="#484f58" font-size="9">ws_id set → route to owning node</text>
|
||||
<circle cx="44" cy="514" r="3" fill="#58a6ff" opacity="0.6"/>
|
||||
<text x="54" y="517" fill="#484f58" font-size="9">neither → shared queue, any node picks up</text></svg>
|
||||
|
After Width: | Height: | Size: 18 KiB |
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:9a1b0361c466327d0011a488847ea3c0365983713537d4a7c27cd7f5538ba33c
|
||||
size 164829
|
||||
oid sha256:d8ce6d2a43a991655c3f64a20b6e810fdb2f78eb767acc3d3d1b8d2c9f443181
|
||||
size 165011
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7d75da92a657525bcbb7a425dc6c8d3cafe074c3ac3bff3cf4b1d44aea607b50
|
||||
size 330156
|
||||
oid sha256:0ee0a9391bd19d92e9271bf6bd531e9c2e18baf8c5a11ead49b3c10db4d8939b
|
||||
size 329625
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:637458e0d78df82752746e519cd7300a830c8ce211f21625694ad0c162ca316d
|
||||
size 481637
|
||||
oid sha256:c53ddce800c59f9432d7a016c7d66282a449555d452b7fe9dd393f4282f08c46
|
||||
size 554721
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:dc3b64c9e48153641af62ed43fbc1d89a31d1a8a61e7e71cfc550c805000310d
|
||||
size 288290
|
||||
oid sha256:e3044c738d6d6853aab5c4990e6c67bab0165eba991a4f5bebdfc4d4a0b305ee
|
||||
size 289165
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:b842683d238664a3e35d04358fecfc56cefd013f7dca5f13357b0376f881e1b3
|
||||
size 245043
|
||||
oid sha256:282820fe416961e735d050f86ecdc079e29824d2b3c4d5c8c174d0533d41f211
|
||||
size 258045
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:b22d5980fe5cc4b8466ba0797113dc8fa83fab8df24b5dacceaf97e62e2e25b0
|
||||
size 187649
|
||||
oid sha256:32a0665cceffcc0517265bde12cfb227688aa8585284b5e946ab23bcc52daee6
|
||||
size 187650
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:90e4f74be795b530e711faa87bc6eb2b3bf6abb68d8fac8ebff7aaf30c6fbe53
|
||||
oid sha256:09535722ba975e47cf0557a40b6c481f125ff2022c396f79715c3bba9f715871
|
||||
size 222032
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:d33b9b3affcdb07086b5aebca8a3b9c2b009cdfc6f360950a0e72e65fbcb8f17
|
||||
size 201602
|
||||
oid sha256:ed457b10b534b5fc2a5e190b281d7ded4dd1615da2229d67a373cf5dddccd059
|
||||
size 201601
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:adac93a0bb062d7199b819a600a0983ff011a75d16928fb80322cbb41f9284ea
|
||||
size 158866
|
||||
oid sha256:e0a3f48cca1b8408862dc4ba04fd340703346f44d84048c99e9900f48e9c7e22
|
||||
size 158867
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:69f201cff948cb0a19810b7c4ad26d346f869ee2dd3141eba4f353332efa2e21
|
||||
size 373649
|
||||
oid sha256:35cf3a6942f62dabcbbe012ac2f9e6f155332c894692981b076de5a25c1f3330
|
||||
size 374055
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:97e7210cd8f1ad195f4d5e25e778d82df3c08c5c6e0f09722e84a7a453714867
|
||||
size 411664
|
||||
oid sha256:a74b4b8b5dbfb1a51a01100b731477968942b01218bad9451a3d5a9cb3003294
|
||||
size 411665
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:4c3214ef416c1dfe4fa17834c2b6f4071a8093cfdb2b862848ca79938f726a13
|
||||
oid sha256:84524f4bc900708ac8adf081591d336f862830188eb8505e71a0f071b339d923
|
||||
size 252599
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:c9823a41e09611c5c0530d9fc12ad4139cfcc3ae238dc665b2888ec94d7d6781
|
||||
size 195708
|
||||
oid sha256:435a58aa09d0e6615e78c0be62e5fd9aa6d7329b1e96619744355c42ade649c9
|
||||
size 196502
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:0c615984373b4893b6cc5755604f137541e9d391862122746a4fcbae63543563
|
||||
size 201041
|
||||
oid sha256:5faa5335152685cf1c8bf77ed93847d751cde59e1afed651e5991113f2f0f31b
|
||||
size 242670
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:af5ab3126bf685afe68e24bc4b0ed97371d0ebdb77bf4d76c0331ab120580cc0
|
||||
size 248809
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:1380065cbb5f95b5ea7dc6b2a00986c455b82888af60784980dffbd936460dcf
|
||||
size 431129
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:f0f6097840fccdbfe16cd5e4c9f5d063b2c36942460a944df68b8ec947e63ea3
|
||||
size 221452
|
||||
+48
-6
@@ -27,6 +27,9 @@ Console dashboard: http://localhost:8090
|
||||
| `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
|
||||
@@ -37,6 +40,18 @@ Console dashboard: http://localhost:8090
|
||||
docker compose up
|
||||
```
|
||||
|
||||
**Production** — adds PostgreSQL and the channel gateway. Requires `POSTGRES_PASSWORD` and (for Discord) `TURNSTONE_DISCORD_TOKEN`:
|
||||
|
||||
```bash
|
||||
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`:
|
||||
|
||||
```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
|
||||
@@ -84,8 +99,35 @@ All configuration is via environment variables in `.env` (copy from `.env.exampl
|
||||
|
||||
| Variable | Default | Description |
|
||||
|----------|---------|-------------|
|
||||
| `TURNSTONE_AUTH_ENABLED` | — | Set to `1` to require Bearer token auth |
|
||||
| `TURNSTONE_AUTH_TOKEN` | — | Shared auth token for server/bridge/console |
|
||||
| `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_JWT_SECRET` | — | Secret key for signing JWTs (required when using user identity / JWT auth) |
|
||||
|
||||
### Database
|
||||
|
||||
| Variable | Default | Description |
|
||||
|----------|---------|-------------|
|
||||
| `TURNSTONE_DB_BACKEND` | `sqlite` | Storage backend: `sqlite` or `postgresql` |
|
||||
| `TURNSTONE_DB_URL` | — | Database URL (e.g. `postgresql://user:pass@db:5432/turnstone`). For SQLite, defaults to `/data/.turnstone.db` |
|
||||
|
||||
The database stores workstream history, user accounts, and API tokens. When using JWT auth, a database backend is required for user storage.
|
||||
|
||||
> **First-time setup:** After deploying with auth enabled, create an initial admin user by running `turnstone-admin create-user` inside the container:
|
||||
>
|
||||
> ```bash
|
||||
> docker compose exec server turnstone-admin create-user --username admin --name "Admin"
|
||||
> ```
|
||||
>
|
||||
> You will be prompted to set a password. Use it to log in via the UI or SDK, then create additional users through the admin API. Pass `--token --scopes read,write,approve` to also generate an initial API token.
|
||||
|
||||
### Channel Gateway
|
||||
|
||||
| Variable | Default | Description |
|
||||
|----------|---------|-------------|
|
||||
| `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
|
||||
|
||||
@@ -101,13 +143,13 @@ All configuration is via environment variables in `.env` (copy from `.env.exampl
|
||||
|
||||
## Scaling
|
||||
|
||||
Scale to multiple server/bridge pairs:
|
||||
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:
|
||||
|
||||
```bash
|
||||
docker compose up --scale server=3 --scale bridge=3
|
||||
POSTGRES_PASSWORD=secret docker compose --profile cluster up
|
||||
```
|
||||
|
||||
Each bridge auto-generates a unique node ID from its container hostname. When scaling `server`, remove the host port mapping (or use a reverse proxy) to avoid port conflicts.
|
||||
The default `server` and `bridge` also run alongside the cluster nodes (11 total). All nodes are accessible via the console dashboard at `:8090`.
|
||||
|
||||
## Volumes
|
||||
|
||||
@@ -128,7 +170,7 @@ docker compose build
|
||||
docker compose build --no-cache
|
||||
```
|
||||
|
||||
All five entry points are installed in a single image: `turnstone-server`, `turnstone-bridge`, `turnstone-console`, `turnstone-sim`, `turnstone-eval`.
|
||||
All entry points are installed in a single image: `turnstone-server`, `turnstone-bridge`, `turnstone-console`, `turnstone-channel`, `turnstone-admin`, `turnstone-sim`, `turnstone-eval`.
|
||||
|
||||
## Cleanup
|
||||
|
||||
|
||||
+75
-16
@@ -15,8 +15,10 @@ The Python SDK is included in the `turnstone` package — no extra install requi
|
||||
```python
|
||||
from turnstone.sdk import TurnstoneServer
|
||||
|
||||
# Synchronous client
|
||||
with TurnstoneServer("http://localhost:8080", token="tok_xxx") as client:
|
||||
# Synchronous client — login with username/password
|
||||
with TurnstoneServer("http://localhost:8080") as client:
|
||||
client.login(username="alice", password="s3cret")
|
||||
|
||||
# Create a workstream
|
||||
ws = client.create_workstream(name="Analysis")
|
||||
|
||||
@@ -33,6 +35,15 @@ with TurnstoneServer("http://localhost:8080", token="tok_xxx") as client:
|
||||
client.close_workstream(ws.ws_id)
|
||||
```
|
||||
|
||||
Alternatively, authenticate with an API token:
|
||||
|
||||
```python
|
||||
with TurnstoneServer("http://localhost:8080") as client:
|
||||
client.login(token="ts_abc123...")
|
||||
ws = client.create_workstream(name="CI run")
|
||||
result = client.send_and_wait("Run the test suite.", ws.ws_id)
|
||||
```
|
||||
|
||||
### Async Client
|
||||
|
||||
```python
|
||||
@@ -40,7 +51,8 @@ import asyncio
|
||||
from turnstone.sdk import AsyncTurnstoneServer
|
||||
|
||||
async def main():
|
||||
async with AsyncTurnstoneServer("http://localhost:8080", token="tok_xxx") as client:
|
||||
async with AsyncTurnstoneServer("http://localhost:8080") as client:
|
||||
await client.login(username="alice", password="s3cret")
|
||||
ws = await client.create_workstream(name="demo")
|
||||
async for event in client.stream_events(ws.ws_id):
|
||||
if event.type == "content":
|
||||
@@ -66,9 +78,11 @@ Both `TurnstoneServer` (sync) and `AsyncTurnstoneServer` (async) expose:
|
||||
| **Streaming** | `stream_events(ws_id)` | `Iterator[ServerEvent]` |
|
||||
| | `stream_global_events()` | `Iterator[ServerEvent]` |
|
||||
| **High-level** | `send_and_wait(message, ws_id, *, timeout, on_event)` | `TurnResult` |
|
||||
| **Sessions** | `list_sessions()` | `ListSessionsResponse` |
|
||||
| **Auth** | `login(token)` | `AuthLoginResponse` |
|
||||
| **Saved** | `list_saved_workstreams()` | `ListSavedWorkstreamsResponse` |
|
||||
| **Auth** | `login(username=..., password=...)` | `AuthLoginResponse` |
|
||||
| | `login(token="ts_xxx")` | `AuthLoginResponse` |
|
||||
| | `logout()` | `StatusResponse` |
|
||||
| | `auth_status()` | `AuthStatusResponse` |
|
||||
| **Health** | `health()` | `HealthResponse` |
|
||||
|
||||
### Console Client API
|
||||
@@ -81,9 +95,17 @@ Both `TurnstoneConsole` (sync) and `AsyncTurnstoneConsole` (async) expose:
|
||||
| | `nodes(*, sort, limit, offset)` | `ClusterNodesResponse` |
|
||||
| | `workstreams(*, state, node, search, sort, page, per_page)` | `ClusterWorkstreamsResponse` |
|
||||
| | `node_detail(node_id)` | `NodeDetailResponse` |
|
||||
| | `snapshot()` | `ClusterSnapshotResponse` |
|
||||
| | `create_workstream(*, node_id, name, model, initial_message)` | `ConsoleCreateWsResponse` |
|
||||
| **Schedules** | `list_schedules()` | `ListSchedulesResponse` |
|
||||
| | `create_schedule(*, name, schedule_type, initial_message, ...)` | `ScheduleInfo` |
|
||||
| | `get_schedule(task_id)` | `ScheduleInfo` |
|
||||
| | `update_schedule(task_id, *, name=..., enabled=..., ...)` | `ScheduleInfo` |
|
||||
| | `delete_schedule(task_id)` | `StatusResponse` |
|
||||
| | `list_schedule_runs(task_id, *, limit=50)` | `ListScheduleRunsResponse` |
|
||||
| **Streaming** | `stream_cluster_events()` | `Iterator[ClusterEvent]` |
|
||||
| **Auth** | `login(token)` / `logout()` | `AuthLoginResponse` / `StatusResponse` |
|
||||
| **Auth** | `login(username=..., password=...)` / `login(token="ts_xxx")` | `AuthLoginResponse` |
|
||||
| | `logout()` | `StatusResponse` |
|
||||
| **Health** | `health()` | `ConsoleHealthResponse` |
|
||||
|
||||
### Event Types
|
||||
@@ -125,6 +147,9 @@ SSE events are deserialized into typed dataclasses. Use `event.type` to discrimi
|
||||
| `node_lost` | `NodeLostEvent` | `node_id` |
|
||||
| `cluster_state` | `ClusterStateEvent` | `ws_id`, `node_id`, `state`, `tokens` |
|
||||
| `ws_created` | `ClusterWsCreatedEvent` | `ws_id`, `node_id`, `name` |
|
||||
| `ws_closed` | `ClusterWsClosedEvent` | `ws_id` |
|
||||
| `ws_rename` | `ClusterWsRenameEvent` | `ws_id`, `name` |
|
||||
| `snapshot` | `ClusterSnapshotEvent` | `nodes`, `overview`, `timestamp` |
|
||||
|
||||
### TurnResult
|
||||
|
||||
@@ -165,10 +190,11 @@ Located at `sdk/typescript/`. Zero runtime dependencies for browsers; uses nativ
|
||||
```typescript
|
||||
import { TurnstoneServer } from "@turnstone/sdk";
|
||||
|
||||
const client = new TurnstoneServer({
|
||||
baseUrl: "http://localhost:8080",
|
||||
token: "tok_xxx",
|
||||
});
|
||||
const client = new TurnstoneServer({ baseUrl: "http://localhost:8080" });
|
||||
|
||||
// Login with username/password or API token
|
||||
await client.login({ username: "alice", password: "s3cret" });
|
||||
// or: await client.login({ token: "ts_abc123..." });
|
||||
|
||||
// Create workstream and send message
|
||||
const ws = await client.createWorkstream({ name: "demo" });
|
||||
@@ -188,16 +214,14 @@ for await (const event of client.streamEvents(ws.ws_id)) {
|
||||
```typescript
|
||||
import { TurnstoneConsole } from "@turnstone/sdk";
|
||||
|
||||
const console = new TurnstoneConsole({
|
||||
baseUrl: "http://localhost:8081",
|
||||
token: "tok_xxx",
|
||||
});
|
||||
const client = new TurnstoneConsole({ baseUrl: "http://localhost:8090" });
|
||||
await client.login({ username: "alice", password: "s3cret" });
|
||||
|
||||
const overview = await console.overview();
|
||||
const overview = await client.overview();
|
||||
console.log(`Nodes: ${overview.nodes}, Workstreams: ${overview.workstreams}`);
|
||||
|
||||
// Stream cluster events
|
||||
for await (const event of console.clusterEvents()) {
|
||||
for await (const event of client.clusterEvents()) {
|
||||
console.log(event.type, event);
|
||||
}
|
||||
```
|
||||
@@ -256,3 +280,38 @@ sdk/typescript/ TypeScript SDK (npm package)
|
||||
The Python SDK reuses Pydantic models from `turnstone/api/` directly — no schema duplication. The TypeScript SDK has hand-written interfaces matching those models.
|
||||
|
||||
Both SDKs follow the same design: typed methods for REST endpoints, async iterators for SSE streams, and a high-level `send_and_wait` method for simple request-response patterns.
|
||||
|
||||
---
|
||||
|
||||
## Authentication
|
||||
|
||||
When auth is enabled on the server, the SDK handles JWT-based authentication automatically.
|
||||
|
||||
### Login Flow
|
||||
|
||||
There are two ways to authenticate:
|
||||
|
||||
1. **Username + password** — calls `POST /v1/api/auth/login` with credentials. The server validates against the user database and returns a JWT.
|
||||
|
||||
2. **API token** — calls `POST /v1/api/auth/login` with a `ts_`-prefixed token string. The server looks up the token, resolves the associated user, and returns a JWT.
|
||||
|
||||
In both cases the server returns the JWT in the response body and as a `Set-Cookie` header. The SDK extracts the JWT and includes it as a `Bearer` token in the `Authorization` header on all subsequent requests.
|
||||
|
||||
```python
|
||||
# Username + password
|
||||
client.login(username="alice", password="s3cret")
|
||||
|
||||
# API token (created via admin API or turnstone-admin CLI)
|
||||
client.login(token="ts_abc123...")
|
||||
```
|
||||
|
||||
### Token Lifecycle
|
||||
|
||||
- JWTs have a configurable expiry (default: 24 hours).
|
||||
- `client.auth_status()` returns the current user identity and scopes without refreshing the token.
|
||||
- `client.logout()` clears the stored JWT from the client.
|
||||
- If a request returns 401, the SDK raises `TurnstoneAPIError` — the caller is responsible for re-authenticating.
|
||||
|
||||
### Backward Compatibility
|
||||
|
||||
The config-file token (`TURNSTONE_AUTH_TOKEN`) still works as a simple Bearer token for environments that do not use the user/JWT system. When the server receives a non-JWT Bearer token, it falls back to the legacy token check.
|
||||
|
||||
@@ -0,0 +1,457 @@
|
||||
# Security and Authentication
|
||||
|
||||
Turnstone uses a layered authentication system with three token types,
|
||||
hierarchical scopes, and a split architecture where the console manages
|
||||
credentials while individual server nodes validate JWTs locally.
|
||||
|
||||
---
|
||||
|
||||
## Token Types
|
||||
|
||||
### Config-file tokens
|
||||
|
||||
Static tokens defined in `config.toml` or the `TURNSTONE_AUTH_TOKEN`
|
||||
environment variable. Validated in-memory using `hmac.compare_digest`
|
||||
(timing-safe). Each token maps to a role that determines its scopes.
|
||||
|
||||
```toml
|
||||
[[auth.tokens]]
|
||||
value = "tok_legacy"
|
||||
role = "full" # full → {read, write, approve}
|
||||
```
|
||||
|
||||
Role mappings: `"read"` → `{read}`, `"full"` → `{read, write, approve}`.
|
||||
|
||||
Config tokens are sent directly as `Authorization: Bearer tok_legacy`
|
||||
on every request. No JWT exchange is needed.
|
||||
|
||||
### API tokens
|
||||
|
||||
Database-backed tokens prefixed with `ts_`. Created via the admin CLI
|
||||
(`turnstone-admin create-token`) or the console admin API. Stored as
|
||||
SHA-256 hashes — the raw token is shown exactly once at creation and
|
||||
never persisted in plaintext.
|
||||
|
||||
```
|
||||
$ turnstone-admin create-token --user abc123 --scopes read,write --name "CI bot"
|
||||
Token created: ts_a1b2c3d4e5f6...
|
||||
(save this — it will not be shown again)
|
||||
```
|
||||
|
||||
API tokens can be used directly as `Bearer ts_xxx` headers or exchanged
|
||||
for a JWT via the login endpoint.
|
||||
|
||||
### JWTs
|
||||
|
||||
Short-lived session tokens (24 hours by default). Issued after
|
||||
authenticating with username/password or by exchanging an API token.
|
||||
HS256-signed with a shared secret. Validated locally on every service
|
||||
node — no database call per request.
|
||||
|
||||
Claims:
|
||||
|
||||
| Claim | Description |
|
||||
|-------|-------------|
|
||||
| `sub` | User ID |
|
||||
| `scopes` | Comma-separated scope list (`read,write,approve`) |
|
||||
| `src` | Token source (`password`, `api_token`, `config`) |
|
||||
| `iss` | Issuer — always `turnstone` |
|
||||
| `aud` | Audience — `turnstone-server` or `turnstone-console` |
|
||||
| `iat` | Issued-at timestamp |
|
||||
| `exp` | Expiry timestamp |
|
||||
|
||||
The `aud` claim prevents cross-service token reuse — a JWT issued for the
|
||||
console cannot be used to authenticate against a server node, and vice versa.
|
||||
Tokens without an `aud` claim are accepted during the rollout window when
|
||||
`audience` validation is not specified.
|
||||
|
||||
---
|
||||
|
||||
## Scope Model
|
||||
|
||||
Scopes are hierarchical — higher scopes imply all lower ones.
|
||||
|
||||
| Scope | Grants | Implies |
|
||||
|-------|--------|---------|
|
||||
| `read` | View workstreams, saved workstreams, history | — |
|
||||
| `write` | Send messages, create/close workstreams | `read` |
|
||||
| `approve` | Approve tool calls, admin endpoints | `read`, `write` |
|
||||
|
||||
### Path-to-scope mapping
|
||||
|
||||
| Method | Path pattern | Required scope |
|
||||
|--------|-------------|----------------|
|
||||
| GET | Any protected path | `read` |
|
||||
| POST | `/api/send`, `/api/plan`, `/api/command` | `write` |
|
||||
| POST | `/api/workstreams/new`, `/api/workstreams/close` | `write` |
|
||||
| POST | `/api/cluster/workstreams/new` | `write` |
|
||||
| POST | `/api/approve` | `approve` |
|
||||
| Any | `/api/admin/*` | `approve` |
|
||||
|
||||
Public paths bypass authentication entirely: `/`, `/health`, `/metrics`,
|
||||
`/static/*`, `/shared/*`, `/docs`, `/openapi.json`, `/api/auth/login`,
|
||||
`/api/auth/logout`, `/api/auth/status`, `/api/auth/setup`.
|
||||
|
||||
---
|
||||
|
||||
## Login Flows
|
||||
|
||||
### Username and password
|
||||
|
||||
```
|
||||
POST /v1/api/auth/login
|
||||
Content-Type: application/json
|
||||
|
||||
{"username": "admin", "password": "s3cret"}
|
||||
```
|
||||
|
||||
Returns a JWT in the response body and sets an `HttpOnly` session cookie.
|
||||
|
||||
### API token exchange
|
||||
|
||||
```
|
||||
POST /v1/api/auth/login
|
||||
Content-Type: application/json
|
||||
|
||||
{"token": "ts_a1b2c3d4e5f6..."}
|
||||
```
|
||||
|
||||
The API token is hashed, looked up in the database, and exchanged for a
|
||||
JWT with the token's scopes. This is the recommended flow for SDKs and
|
||||
automated clients that need cookie-based sessions.
|
||||
|
||||
### Config-file tokens (direct)
|
||||
|
||||
Config tokens are validated per-request via `hmac.compare_digest`. No
|
||||
login exchange is needed — include the token as a `Bearer` header:
|
||||
|
||||
```
|
||||
Authorization: Bearer tok_legacy
|
||||
```
|
||||
|
||||
### First-time setup
|
||||
|
||||
When no users exist in the database:
|
||||
|
||||
1. `GET /v1/api/auth/status` returns `{"setup_required": true}`
|
||||
2. The UI presents a setup wizard
|
||||
3. `POST /v1/api/auth/setup` creates the first admin user and returns a
|
||||
JWT in one atomic step (no auth required — this is a public endpoint)
|
||||
4. The endpoint returns `409 Conflict` if setup has already been completed
|
||||
(i.e. users already exist in the database)
|
||||
5. Subsequent admin requests require `approve` scope
|
||||
|
||||
The `/api/auth/setup` endpoint is available on both the server and
|
||||
console. It validates input before creating the user:
|
||||
|
||||
- **username**: 1-64 ASCII characters
|
||||
- **display_name**: required (non-empty)
|
||||
- **password**: minimum 8 characters
|
||||
|
||||
```
|
||||
POST /v1/api/auth/setup
|
||||
Content-Type: application/json
|
||||
|
||||
{"username": "admin", "display_name": "Admin", "password": "strongpass"}
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{
|
||||
"status": "ok",
|
||||
"user_id": "u_abc123",
|
||||
"username": "admin",
|
||||
"role": "full",
|
||||
"scopes": "approve,read,write",
|
||||
"jwt": "eyJhbGciOiJIUzI1NiIs..."
|
||||
}
|
||||
```
|
||||
|
||||
The response also sets an `HttpOnly` session cookie containing the JWT,
|
||||
so the browser is immediately authenticated after setup completes.
|
||||
|
||||
---
|
||||
|
||||
## Token Detection Order
|
||||
|
||||
The auth middleware inspects the `Authorization: Bearer <token>` header
|
||||
and classifies the token:
|
||||
|
||||
1. **Contains `.`** → JWT → validate HS256 signature and expiry
|
||||
2. **Starts with `ts_`** → API token → SHA-256 hash, database lookup
|
||||
3. **Otherwise** → config-file token → `hmac.compare_digest` against
|
||||
each configured token
|
||||
|
||||
If a session cookie is present and no `Authorization` header is sent,
|
||||
the cookie value is treated as a JWT (step 1).
|
||||
|
||||
---
|
||||
|
||||
## Password Storage
|
||||
|
||||
Passwords are hashed with **bcrypt** using a random salt per password.
|
||||
Plaintext passwords are only accepted over HTTPS in production
|
||||
deployments.
|
||||
|
||||
---
|
||||
|
||||
## Cookie Security
|
||||
|
||||
| Attribute | Value | Purpose |
|
||||
|-----------|-------|---------|
|
||||
| `HttpOnly` | `true` | Prevents JavaScript access |
|
||||
| `SameSite` | `Lax` | CSRF protection |
|
||||
| `Path` | `/` | Available to all routes |
|
||||
| `Max-Age` | 24 hours | Matches JWT expiry |
|
||||
| `Secure` | `true` (default) | Always set unless explicitly disabled for dev |
|
||||
|
||||
---
|
||||
|
||||
## JWT Configuration
|
||||
|
||||
| Setting | Config key | Env var | Default |
|
||||
|---------|-----------|---------|---------|
|
||||
| Signing secret | `[auth] jwt_secret` | `TURNSTONE_JWT_SECRET` | Auto-generated ephemeral (warning logged) |
|
||||
| Expiry | `[auth] jwt_expiry_hours` | — | 24 hours |
|
||||
| Algorithm | — | — | HS256 (not configurable) |
|
||||
| Minimum secret length | — | — | 32 characters (warning if shorter) |
|
||||
|
||||
All service nodes that need to validate JWTs must share the same signing
|
||||
secret. If no secret is configured, an ephemeral key is generated at
|
||||
startup and a warning is logged — JWTs will not survive restarts or work
|
||||
across nodes.
|
||||
|
||||
The bridge and console **require** `TURNSTONE_JWT_SECRET` when no
|
||||
`--auth-token` is provided. They exit with an error if the secret is
|
||||
missing, since ephemeral secrets would silently break inter-service
|
||||
communication.
|
||||
|
||||
---
|
||||
|
||||
## Admin API Endpoints
|
||||
|
||||
All admin endpoints require `approve` scope.
|
||||
|
||||
### Users
|
||||
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| POST | `/v1/api/admin/users` | Create user (username, display_name, password) |
|
||||
| GET | `/v1/api/admin/users` | List all users |
|
||||
| DELETE | `/v1/api/admin/users/{user_id}` | Delete user and cascade tokens |
|
||||
|
||||
### API tokens
|
||||
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| POST | `/v1/api/admin/users/{user_id}/tokens` | Create API token (returns raw value once) |
|
||||
| GET | `/v1/api/admin/users/{user_id}/tokens` | List tokens (prefix only, no hashes) |
|
||||
| DELETE | `/v1/api/admin/tokens/{token_id}` | Revoke token |
|
||||
|
||||
---
|
||||
|
||||
## CLI Administration
|
||||
|
||||
The `turnstone-admin` command provides offline user and token management:
|
||||
|
||||
```
|
||||
turnstone-admin create-user --username admin --name "Admin" [--password] [--token]
|
||||
turnstone-admin create-token --user <user_id> --scopes read,write --name "CI bot"
|
||||
turnstone-admin list-users
|
||||
turnstone-admin list-tokens
|
||||
turnstone-admin revoke-token <token_id>
|
||||
```
|
||||
|
||||
When `--password` is omitted, the CLI prompts interactively. When
|
||||
`--token` is passed to `create-user`, an API token is created alongside
|
||||
the user and printed to stdout.
|
||||
|
||||
---
|
||||
|
||||
## Database Schema
|
||||
|
||||
```sql
|
||||
CREATE TABLE users (
|
||||
user_id TEXT PRIMARY KEY,
|
||||
username TEXT NOT NULL UNIQUE,
|
||||
display_name TEXT NOT NULL,
|
||||
password_hash TEXT NOT NULL,
|
||||
created TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE api_tokens (
|
||||
token_id TEXT PRIMARY KEY,
|
||||
token_hash TEXT NOT NULL, -- SHA-256 of raw token
|
||||
token_prefix TEXT NOT NULL, -- first 8 chars for display
|
||||
user_id TEXT NOT NULL REFERENCES users(user_id),
|
||||
name TEXT NOT NULL,
|
||||
scopes TEXT NOT NULL, -- comma-separated
|
||||
created TEXT NOT NULL,
|
||||
expires TEXT -- nullable, ISO 8601
|
||||
);
|
||||
|
||||
CREATE UNIQUE INDEX ix_api_tokens_hash ON api_tokens(token_hash);
|
||||
|
||||
CREATE TABLE channel_users (
|
||||
channel_type TEXT NOT NULL,
|
||||
channel_user_id TEXT NOT NULL,
|
||||
user_id TEXT NOT NULL REFERENCES users(user_id),
|
||||
created TEXT NOT NULL,
|
||||
PRIMARY KEY (channel_type, channel_user_id)
|
||||
);
|
||||
```
|
||||
|
||||
The `sessions` and `workstreams` tables have a nullable `user_id`
|
||||
column for attribution when auth is enabled.
|
||||
|
||||
---
|
||||
|
||||
## Revocation
|
||||
|
||||
- **API tokens**: Deleting a token via the admin API or CLI prevents new
|
||||
JWTs from being issued with that token. Existing JWTs derived from the
|
||||
token remain valid until they expire (at most 24 hours).
|
||||
- **Config-file tokens**: Remove the token from `config.toml` and
|
||||
restart the service. No JWTs are involved, so revocation is immediate.
|
||||
- **JWTs**: Cannot be individually revoked. Rely on short expiry (24h)
|
||||
and revoke the underlying credential to prevent renewal.
|
||||
|
||||
---
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
Console (cluster-wide) Server (per-node)
|
||||
┌──────────────────────┐ ┌──────────────────────┐
|
||||
│ User/Token CRUD (DB) │ │ JWT validation only │
|
||||
│ Login: creds → JWT │ │ (shared signing key) │
|
||||
│ Admin API endpoints │ │ Config tokens: hmac │
|
||||
│ Storage: users, │ │ No auth DB needed │
|
||||
│ api_tokens tables │ │ │
|
||||
└──────────────────────┘ └──────────────────────┘
|
||||
```
|
||||
|
||||
The console owns the credential database and handles all user/token
|
||||
CRUD. Individual server nodes only need the JWT signing secret to
|
||||
validate session tokens. Config-file tokens are validated locally
|
||||
without any database.
|
||||
|
||||
### Proxy auth forwarding
|
||||
|
||||
When the console proxies requests to server nodes (via `/node/{id}/...`
|
||||
routes), it uses a dedicated **service proxy token** with
|
||||
`aud: turnstone-server` and `write` scope. The user's console JWT
|
||||
(which has `aud: turnstone-console`) is **not** forwarded — it would be
|
||||
rejected by the server's audience validation.
|
||||
|
||||
The proxy token is managed by a `ServiceTokenManager` that auto-rotates
|
||||
1-hour JWTs, refreshing at 80% of lifetime. If `--auth-token` is
|
||||
provided, that static token is used instead.
|
||||
|
||||
### Service-to-service authentication
|
||||
|
||||
The bridge and console collector use `ServiceTokenManager` for
|
||||
auto-rotating JWTs when communicating with server nodes:
|
||||
|
||||
| Service | Identity | Scope | Audience | Purpose |
|
||||
|---------|----------|-------|----------|---------|
|
||||
| Bridge | `bridge` | `approve` | `turnstone-server` | Tool approval proxy, message relay |
|
||||
| Console collector | `console-collector` | `read` | `turnstone-server` | Node health polling |
|
||||
| Console proxy | `console-proxy` | `write` | `turnstone-server` | Proxied API calls |
|
||||
| Channel notify | `system` | `write` | `turnstone-channel` | Notification delivery to channel gateway |
|
||||
|
||||
Service tokens use 1-hour expiry with automatic refresh via
|
||||
`ServiceTokenManager`. The bridge injects auth headers per-request via
|
||||
httpx event hooks to ensure rotated tokens are picked up on SSE
|
||||
reconnects.
|
||||
|
||||
Note that the channel gateway uses a distinct JWT audience
|
||||
(`turnstone-channel`) from the server (`turnstone-server`) and console
|
||||
(`turnstone-console`). A server-scoped JWT cannot authenticate to the
|
||||
channel gateway endpoint, and vice versa.
|
||||
|
||||
---
|
||||
|
||||
## Configuration Reference
|
||||
|
||||
### config.toml
|
||||
|
||||
```toml
|
||||
[auth]
|
||||
enabled = true
|
||||
jwt_secret = "your-secret-key-here"
|
||||
jwt_expiry_hours = 24
|
||||
|
||||
[[auth.tokens]]
|
||||
value = "tok_legacy"
|
||||
role = "full"
|
||||
```
|
||||
|
||||
### Environment variables
|
||||
|
||||
| Variable | Description |
|
||||
|----------|-------------|
|
||||
| `TURNSTONE_AUTH_ENABLED=1` | Enable authentication |
|
||||
| `TURNSTONE_AUTH_TOKEN=tok_xxx` | Register a config-file token with `full` access |
|
||||
| `TURNSTONE_JWT_SECRET=xxx` | JWT signing secret (must match across nodes) |
|
||||
| `TURNSTONE_CORS_ORIGINS=` | CORS allowed origins (comma-separated; empty = same-origin only) |
|
||||
|
||||
---
|
||||
|
||||
## Login Rate Limiting
|
||||
|
||||
The `/api/auth/login` endpoint is protected by a dedicated
|
||||
`LoginRateLimiter` (separate from the general API rate limiter).
|
||||
Limits are enforced per-IP and per-username with a sliding window:
|
||||
|
||||
- **5 attempts** per **5-minute window** per key
|
||||
- Failed logins record against both `ip:{client_ip}` and `user:{username}`
|
||||
- Returns `429 Too Many Requests` with `Retry-After` header when exceeded
|
||||
- Successful logins do not consume the budget
|
||||
|
||||
---
|
||||
|
||||
## CORS Policy
|
||||
|
||||
By default, no CORS headers are sent (same-origin only). To allow
|
||||
cross-origin requests, set `TURNSTONE_CORS_ORIGINS`:
|
||||
|
||||
```bash
|
||||
# Allow specific origins
|
||||
TURNSTONE_CORS_ORIGINS=https://app.example.com,https://admin.example.com
|
||||
|
||||
# Allow all origins (development only)
|
||||
TURNSTONE_CORS_ORIGINS=*
|
||||
```
|
||||
|
||||
When the variable is empty or unset, the CORS middleware is not added
|
||||
and browsers enforce same-origin policy.
|
||||
|
||||
---
|
||||
|
||||
## Security Properties
|
||||
|
||||
- **Timing-safe comparison** for config-file tokens via
|
||||
`hmac.compare_digest` — no timing side-channel.
|
||||
- **Hash-based lookup** for API tokens — the database stores only
|
||||
SHA-256 hashes, eliminating timing attacks on token comparison.
|
||||
- **Local JWT validation** — no network call or database query needed
|
||||
per request on server nodes.
|
||||
- **One-time display** of raw API tokens at creation. The plaintext is
|
||||
never stored; `token_hash` never appears in API responses or logs.
|
||||
- **Structured logging audit trail** — `ctx_user_id` is set on every
|
||||
authenticated request and injected into all log events.
|
||||
- **Scope enforcement** at the middleware layer before any handler
|
||||
executes. Path-to-scope mapping is defined statically.
|
||||
- **JWT audience isolation** — server and console JWTs have distinct
|
||||
`aud` claims, preventing cross-service token reuse.
|
||||
- **Login brute-force protection** — per-IP and per-username rate
|
||||
limiting on the login endpoint.
|
||||
- **Secure cookies by default** — `Secure` flag set unconditionally;
|
||||
24-hour max-age matches JWT expiry.
|
||||
- **CORS restriction** — no CORS headers by default (same-origin only).
|
||||
- **Service JWT auto-rotation** — 1-hour expiry with transparent
|
||||
refresh, eliminating long-lived static tokens for inter-service auth.
|
||||
- **Secret strength validation** — warning logged when JWT secret is
|
||||
shorter than 32 characters.
|
||||
+162
-9
@@ -1,6 +1,6 @@
|
||||
# Tools Reference
|
||||
|
||||
turnstone exposes 14 built-in tools plus any number of external MCP tools to the
|
||||
turnstone exposes 15 built-in tools plus any number of external MCP tools to the
|
||||
LLM via the OpenAI function-calling interface. Built-in tools are defined as JSON
|
||||
files under `turnstone/tools/` and loaded at startup by `turnstone/core/tools.py`.
|
||||
MCP tools are discovered from configured MCP servers at startup by
|
||||
@@ -46,11 +46,12 @@ schema plus turnstone-specific metadata keys:
|
||||
|
||||
| Name | Description |
|
||||
|---------------------|-------------|
|
||||
| `TOOLS` | All 14 tool definitions (sent to the model). |
|
||||
| `TOOLS` | All 15 tool definitions (sent to the model). |
|
||||
| `AGENT_TOOLS` | Tools with `agent: true` -- available to plan sub-agents. Read-only tools. |
|
||||
| `TASK_AGENT_TOOLS` | Tools with `task_agent: true` -- available to task sub-agents. Includes write operations. |
|
||||
| `AGENT_AUTO_TOOLS` | Set of tool names with `auto_approve: true` -- no user confirmation needed. |
|
||||
| `TASK_AUTO_TOOLS` | Same as `AGENT_AUTO_TOOLS` (identical filter). |
|
||||
| `BUILTIN_TOOL_NAMES`| Frozenset of all 15 built-in tool names. Used by tool search to distinguish always-on tools from deferrable MCP tools. |
|
||||
| `PRIMARY_KEY_MAP` | Dict mapping tool name to its `primary_key` parameter name. |
|
||||
|
||||
---
|
||||
@@ -68,6 +69,9 @@ Tool execution follows a three-phase pipeline inside `ChatSession._execute_tools
|
||||
- Parses the JSON arguments (with fallback for malformed JSON).
|
||||
- If JSON parsing fails entirely, uses `PRIMARY_KEY_MAP` to map a bare string
|
||||
to the correct parameter.
|
||||
- Dispatches to the matching `_prepare_{func_name}()` handler. There are 15
|
||||
built-in tools plus `tool_search` (synthetic, client-side BM25 fallback) and
|
||||
the generic `_prepare_mcp_tool()` handler for MCP tools.
|
||||
- Validates arguments and builds a preview dict containing:
|
||||
- `call_id`, `func_name`, `header`, `preview` (for display)
|
||||
- `needs_approval` (bool)
|
||||
@@ -113,6 +117,7 @@ Each item's `execute` callable is invoked:
|
||||
- `remember` -- writes to persistent memory database (lightweight, always auto-approved)
|
||||
- `recall` -- reads from persistent memory database
|
||||
- `forget` -- deletes from persistent memory database (lightweight, always auto-approved)
|
||||
- `notify` -- sends notifications to linked channels (time-sensitive, auto-approved for urgency)
|
||||
|
||||
**Requires user confirmation** (write operations, network access, side effects):
|
||||
- `bash` -- arbitrary command execution
|
||||
@@ -162,6 +167,7 @@ Every tool defines a `primary_key`. The mapping is:
|
||||
| `remember` | `key` |
|
||||
| `recall` | `query` |
|
||||
| `forget` | `key` |
|
||||
| `notify` | `message` |
|
||||
|
||||
---
|
||||
|
||||
@@ -183,15 +189,17 @@ Execute a bash command and return stdout + stderr.
|
||||
|
||||
### read_file
|
||||
|
||||
Read the contents of a file, returning numbered lines.
|
||||
Read the contents of a file, returning numbered lines for text files or
|
||||
base64-encoded image data for supported image formats.
|
||||
|
||||
| Parameter | Type | Required | Description |
|
||||
|-----------|---------|----------|-------------|
|
||||
| `path` | string | yes | Absolute or relative file path. |
|
||||
| `offset` | integer | no | Line number to start from (1-based, default: 1). |
|
||||
| `limit` | integer | no | Maximum number of lines to read. Omit for full file. |
|
||||
| `offset` | integer | no | Line number to start from (1-based, default: 1). Text files only. |
|
||||
| `limit` | integer | no | Maximum number of lines to read. Omit for full file. Text files only. |
|
||||
|
||||
- **What it does**: Reads the file and returns content with line numbers. Must be called before `edit_file` on the same path (the session tracks which files have been read).
|
||||
- **What it does**: For text files, reads and returns content with line numbers. For image files (PNG, JPEG, GIF, WebP, BMP, TIFF, ICO), returns image data as multi-part content when the model supports vision, or a text description when it does not. SVG files are read as text. Images larger than 4 MB are rejected. Must be called before `edit_file` on the same path (the session tracks which files have been read).
|
||||
- **Vision support**: Controlled by `ModelCapabilities.supports_vision`. All commercial OpenAI and Anthropic models have vision enabled. Local models (vLLM, llama.cpp, NIM) default to off — enable via `[models.*.capabilities] supports_vision = true` in config.toml.
|
||||
- **Auto-approve**: Yes.
|
||||
- **Agent availability**: `agent` and `task_agent`.
|
||||
|
||||
@@ -335,7 +343,7 @@ Plan before implementing -- an autonomous agent explores the codebase and writes
|
||||
|-----------|--------|----------|-------------|
|
||||
| `prompt` | string | yes | What to plan -- the goal, constraints, and scope. |
|
||||
|
||||
- **What it does**: Spawns a planning sub-agent with `AGENT_TOOLS` (read-only tools: `read_file`, `search`, `math`, `man`, `web_fetch`, `web_search`). The agent explores the codebase and writes a structured plan to `.plan-<session_id>.md` (unique per session, so concurrent workstreams never collide). If the `plan` tool has been called before in the same session, the prior plan is passed to the agent as context so it refines rather than restarts. After completion, the user is prompted to review and can accept, reject, or annotate the plan.
|
||||
- **What it does**: Spawns a planning sub-agent with `AGENT_TOOLS` (read-only tools: `read_file`, `search`, `math`, `man`, `web_fetch`, `web_search`). The agent explores the codebase and writes a structured plan to `.plan-<ws_id>.md` (unique per workstream, so concurrent workstreams never collide). If the `plan` tool has been called before in the same session, the prior plan is passed to the agent as context so it refines rather than restarts. After completion, the user is prompted to review and can accept, reject, or annotate the plan.
|
||||
- **Auto-approve**: No -- requires user confirmation, plus post-execution review gate.
|
||||
- **Agent availability**: Not available to sub-agents (top-level only).
|
||||
|
||||
@@ -387,6 +395,33 @@ Remove a persistent memory by key.
|
||||
|
||||
---
|
||||
|
||||
## Notifications
|
||||
|
||||
### notify
|
||||
|
||||
Send a notification to a user or channel on an external platform.
|
||||
|
||||
| Parameter | Type | Required | Description |
|
||||
|----------------|--------|----------|-------------|
|
||||
| `message` | string | yes | Notification content (plain text, max 2000 chars). |
|
||||
| `username` | string | no | Turnstone username — sends to all linked channels. |
|
||||
| `channel_type` | string | no | Platform for direct targeting (`discord`). |
|
||||
| `channel_id` | string | no | Platform-specific channel or user ID for direct targeting. |
|
||||
| `title` | string | no | Optional short title (rendered as bold prefix). |
|
||||
|
||||
Provide either `username` for user-based targeting or `channel_type` +
|
||||
`channel_id` for direct targeting. Do not combine both.
|
||||
|
||||
- **What it does**: Sends a notification via the channel gateway's HTTP endpoint (`POST /v1/api/notify`). The server queries the `services` table for healthy channel gateways, authenticates with a service JWT (`aud: turnstone-channel`), and delivers to the first healthy gateway. On failure, retries up to 2 additional times with backoff (1s, 3s). Rate-limited to 5 notifications per turn (counter only increments on success).
|
||||
- **Auto-approve**: Yes — notifications are time-sensitive and auto-approved so the model can alert users urgently.
|
||||
- **Agent availability**: `agent` and `task_agent`.
|
||||
|
||||
> See [Channel Integrations: Notifications](channels.md#notifications)
|
||||
> for the full delivery flow, service registry details, and security
|
||||
> measures.
|
||||
|
||||
---
|
||||
|
||||
## Summary Table
|
||||
|
||||
| Tool | Category | Auto-approve | agent | task_agent | primary_key |
|
||||
@@ -405,6 +440,78 @@ Remove a persistent memory by key.
|
||||
| `remember` | Memory | Yes | No | No | `key` |
|
||||
| `recall` | Memory | Yes | No | No | `query` |
|
||||
| `forget` | Memory | Yes | No | No | `key` |
|
||||
| `notify` | Notify | Yes | Yes | Yes | `message` |
|
||||
| `tool_search`| Search | Yes | No | No | `query` |
|
||||
|
||||
---
|
||||
|
||||
## Dynamic Tool Search
|
||||
|
||||
When many MCP tools are connected, the total tool count can grow large enough to
|
||||
consume significant context window tokens and reduce model accuracy. Dynamic tool
|
||||
search addresses this by deferring tools the model is unlikely to need on the
|
||||
current turn and letting it search for them on demand.
|
||||
|
||||
### Three-tier approach
|
||||
|
||||
Tool search uses the best available mechanism for each provider:
|
||||
|
||||
1. **Anthropic (native)** -- Models that support it receive `defer_loading: true`
|
||||
on deferred tool definitions plus the `tool_search_tool_bm25_20251119` server-side
|
||||
search tool. Anthropic's API handles search and expansion transparently.
|
||||
|
||||
2. **OpenAI GPT-5.4+ (native)** -- Models with hosted tool search receive
|
||||
`defer_loading: true` on deferred definitions. The API handles search internally.
|
||||
|
||||
3. **vLLM / llama.cpp / NIM (client-side BM25)** -- A synthetic `tool_search`
|
||||
function tool is injected into the tool list. When the model calls it,
|
||||
`_exec_tool_search()` runs a pure-Python BM25 index over tool names and
|
||||
descriptions, then expands the matched tools into the visible set.
|
||||
|
||||
### Configuration
|
||||
|
||||
Tool search is configured in `config.toml` under the `[tools]` section:
|
||||
|
||||
```toml
|
||||
[tools]
|
||||
search = "auto" # "auto", "on", or "off"
|
||||
search_threshold = 20 # minimum total tool count to activate
|
||||
search_max_results = 5 # max tools returned per search call
|
||||
```
|
||||
|
||||
CLI flags override the config file:
|
||||
|
||||
- `--tool-search {auto,on,off}` -- force tool search on or off, or let turnstone
|
||||
decide based on threshold (default: `auto`).
|
||||
- `--tool-search-threshold N` -- minimum tool count to activate (default: 20).
|
||||
- `--tool-search-max-results N` -- max results per search (default: 5).
|
||||
|
||||
### How it works
|
||||
|
||||
1. **Threshold check**: At session startup, `ToolSearchManager.should_activate()`
|
||||
counts total tools (built-in + MCP). If the count is below the threshold, tool
|
||||
search stays off and all tools are sent to the model directly.
|
||||
|
||||
2. **Partitioning**: When active, tools are split into two sets:
|
||||
- **Always-on** -- the 15 built-in tools (members of `BUILTIN_TOOL_NAMES`).
|
||||
These are always visible to the model.
|
||||
- **Deferred** -- all MCP tools. These are not sent in the tool list unless
|
||||
the model searches for them.
|
||||
|
||||
3. **Search and expand**: When the model calls `tool_search` (client-side) or the
|
||||
provider's native search returns results, the matched tools are added to the
|
||||
visible set via `expand_visible()`. Once expanded, a tool stays visible for
|
||||
the remainder of the session.
|
||||
|
||||
4. **Multi-turn persistence**: Expanded tools are never removed. This avoids
|
||||
confusing the model when it references a tool it discovered in an earlier turn.
|
||||
|
||||
### Agent exemption
|
||||
|
||||
Plan and task sub-agents do not use tool search. They operate on scoped tool
|
||||
sets (`AGENT_TOOLS` for plan agents, `TASK_AGENT_TOOLS` for task agents) with
|
||||
MCP tools merged in. Tool search is only active for the top-level session,
|
||||
where the model can interactively search for tools it needs.
|
||||
|
||||
---
|
||||
|
||||
@@ -421,13 +528,17 @@ MCP-compatible service.
|
||||
|
||||
2. **Discovery**: At startup, `MCPClientManager` connects to each configured server
|
||||
(via stdio subprocess or HTTP), performs the MCP `initialize` handshake, and calls
|
||||
`tools/list` to discover available tools.
|
||||
`tools/list` to discover available tools. During the handshake, the manager checks
|
||||
each server's capabilities for `tools.listChanged` support (push notifications).
|
||||
|
||||
3. **Schema conversion**: Each MCP tool's `inputSchema` is converted to OpenAI
|
||||
function-calling format. The tool name is prefixed: `mcp__{server}__{tool}`.
|
||||
|
||||
4. **Merging**: MCP tools are appended after the 14 built-in tools via
|
||||
4. **Merging**: MCP tools are appended after the 15 built-in tools via
|
||||
`merge_mcp_tools()`. Built-in tools appear first, giving them natural LLM priority.
|
||||
When dynamic tool search is active, MCP tools are deferred rather than directly
|
||||
visible -- the model discovers them via search as needed (see
|
||||
[Dynamic Tool Search](#dynamic-tool-search) above).
|
||||
|
||||
5. **Dispatch**: When the LLM calls an MCP tool, `_prepare_mcp_tool()` builds a
|
||||
generic approval preview and `_exec_mcp_tool()` calls `MCPClientManager.call_tool_sync()`,
|
||||
@@ -500,3 +611,45 @@ MCP tools (3):
|
||||
mcp__github__create_issue [MCP: github] Create a GitHub issue
|
||||
mcp__postgres__query [MCP: postgres] Run a SQL query
|
||||
```
|
||||
|
||||
### Dynamic tool refresh
|
||||
|
||||
MCP tool lists stay up-to-date without restart through three mechanisms:
|
||||
|
||||
1. **Push notifications** -- MCP servers that declare `tools.listChanged: true` in
|
||||
their capabilities send `notifications/tools/list_changed` when their tool list
|
||||
changes. `MCPClientManager` registers a `message_handler` on each `ClientSession`
|
||||
that triggers an immediate refresh for that server.
|
||||
|
||||
2. **Periodic timer** -- Servers that do *not* support push notifications are polled
|
||||
on a configurable interval (default 4 hours). The timer is staggered using a
|
||||
launch-time seed (`monotonic_ns ^ pid`) so cluster nodes don't all hit MCP
|
||||
servers simultaneously. Configure via `[mcp] refresh_interval` in `config.toml`
|
||||
or `--mcp-refresh-interval SECONDS` on the CLI. Set to `0` to disable.
|
||||
|
||||
3. **Manual** -- `/mcp refresh` re-fetches tools from all servers immediately.
|
||||
`/mcp refresh <server>` targets a single server. If a server has disconnected,
|
||||
manual refresh attempts reconnection.
|
||||
|
||||
When tools change, `MCPClientManager` rebuilds its merged tool list using copy-on-write
|
||||
(new list/dict objects assigned atomically) and notifies all active `ChatSession`
|
||||
instances via registered listener callbacks. Each session rebuilds its `_tools`,
|
||||
`_task_tools`, `_agent_tools`, and reconstructs its `ToolSearchManager` (if active),
|
||||
preserving the set of previously expanded (discovered) tools.
|
||||
|
||||
```toml
|
||||
[mcp]
|
||||
refresh_interval = 14400 # seconds (default 4h), 0 to disable
|
||||
```
|
||||
|
||||
```
|
||||
/mcp refresh
|
||||
MCP refresh complete:
|
||||
github: +1 added
|
||||
+ mcp__github__create_pr
|
||||
postgres: no changes
|
||||
|
||||
/mcp refresh github
|
||||
MCP refresh complete:
|
||||
github: no changes
|
||||
```
|
||||
|
||||
+31
-3
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "turnstone"
|
||||
version = "0.3.5"
|
||||
version = "0.5.2"
|
||||
description = "Multi-node AI orchestration platform with tool use, agent routing, and cluster simulation."
|
||||
readme = "README.md"
|
||||
license = "BUSL-1.1"
|
||||
@@ -32,6 +32,9 @@ dependencies = [
|
||||
"pydantic>=2.0",
|
||||
"sqlalchemy>=2.0",
|
||||
"alembic>=1.14",
|
||||
"structlog>=24.1",
|
||||
"PyJWT>=2.8",
|
||||
"bcrypt>=4.0",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
@@ -40,13 +43,14 @@ Repository = "https://github.com/turnstonelabs/turnstone"
|
||||
Issues = "https://github.com/turnstonelabs/turnstone/issues"
|
||||
|
||||
[project.optional-dependencies]
|
||||
test = ["pytest>=9.0", "pytest-cov>=6.0"]
|
||||
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"]
|
||||
console = ["redis>=7.2", "croniter>=3.0"]
|
||||
sim = ["redis>=7.2"]
|
||||
anthropic = ["anthropic>=0.39"]
|
||||
postgres = ["psycopg[binary]>=3.2"]
|
||||
discord = ["discord.py>=2.4", "redis>=7.2"]
|
||||
|
||||
|
||||
[project.scripts]
|
||||
@@ -56,6 +60,8 @@ 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"
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
include = [
|
||||
@@ -128,10 +134,32 @@ ignore_missing_imports = true
|
||||
module = ["sqlalchemy", "sqlalchemy.*", "alembic", "alembic.*"]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = ["structlog", "structlog.*"]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = ["jwt", "jwt.*"]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = ["anthropic", "anthropic.*"]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = ["discord", "discord.*"]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = ["croniter", "croniter.*"]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = ["turnstone.channels.discord.*"]
|
||||
disallow_subclassing_any = false
|
||||
disallow_untyped_decorators = false
|
||||
warn_unused_ignores = false
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = "tests.*"
|
||||
disallow_untyped_defs = false
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"openapi": "3.1.0",
|
||||
"info": {
|
||||
"title": "turnstone Server API",
|
||||
"version": "0.3.0",
|
||||
"version": "0.4.2",
|
||||
"description": "Single-node workstream management, chat interaction, and real-time streaming."
|
||||
},
|
||||
"paths": {
|
||||
@@ -365,12 +365,12 @@
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/api/sessions": {
|
||||
"/v1/api/workstreams/saved": {
|
||||
"get": {
|
||||
"summary": "List saved sessions",
|
||||
"operationId": "v1_api_sessions_get",
|
||||
"summary": "List saved workstreams",
|
||||
"operationId": "v1_api_workstreams_saved_get",
|
||||
"tags": [
|
||||
"Sessions"
|
||||
"Workstreams"
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
@@ -378,7 +378,7 @@
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ListSessionsResponse"
|
||||
"$ref": "#/components/schemas/ListSavedWorkstreamsResponse"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -427,6 +427,88 @@
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/api/auth/setup": {
|
||||
"post": {
|
||||
"summary": "Create first admin user",
|
||||
"operationId": "v1_api_auth_setup_post",
|
||||
"tags": [
|
||||
"Auth"
|
||||
],
|
||||
"requestBody": {
|
||||
"required": true,
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/AuthSetupRequest"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"responses": {
|
||||
"200": {
|
||||
"description": "Success",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/AuthSetupResponse"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"400": {
|
||||
"description": "Error 400",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ErrorResponse"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"409": {
|
||||
"description": "Error 409",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ErrorResponse"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"503": {
|
||||
"description": "Error 503",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ErrorResponse"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/api/auth/status": {
|
||||
"get": {
|
||||
"summary": "Return auth state",
|
||||
"operationId": "v1_api_auth_status_get",
|
||||
"tags": [
|
||||
"Auth"
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"description": "Success",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/AuthStatusResponse"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/api/auth/logout": {
|
||||
"post": {
|
||||
"summary": "Clear auth cookie",
|
||||
@@ -503,17 +585,27 @@
|
||||
"type": "object"
|
||||
},
|
||||
"AuthLoginRequest": {
|
||||
"description": "POST /v1/api/auth/login request body.",
|
||||
"description": "POST /v1/api/auth/login request body.\n\nEither username+password or token must be provided.",
|
||||
"properties": {
|
||||
"username": {
|
||||
"default": "",
|
||||
"description": "Login username",
|
||||
"title": "Username",
|
||||
"type": "string"
|
||||
},
|
||||
"password": {
|
||||
"default": "",
|
||||
"description": "Login password",
|
||||
"title": "Password",
|
||||
"type": "string"
|
||||
},
|
||||
"token": {
|
||||
"description": "Bearer token to authenticate",
|
||||
"default": "",
|
||||
"description": "Legacy: bearer token to authenticate",
|
||||
"title": "Token",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"token"
|
||||
],
|
||||
"title": "AuthLoginRequest",
|
||||
"type": "object"
|
||||
},
|
||||
@@ -525,14 +617,35 @@
|
||||
"title": "Status",
|
||||
"type": "string"
|
||||
},
|
||||
"user_id": {
|
||||
"default": "",
|
||||
"description": "Authenticated user ID",
|
||||
"title": "User Id",
|
||||
"type": "string"
|
||||
},
|
||||
"role": {
|
||||
"description": "Assigned role",
|
||||
"description": "Legacy role",
|
||||
"examples": [
|
||||
"full",
|
||||
"read"
|
||||
],
|
||||
"title": "Role",
|
||||
"type": "string"
|
||||
},
|
||||
"scopes": {
|
||||
"default": "",
|
||||
"description": "Comma-separated scopes",
|
||||
"examples": [
|
||||
"read,write,approve"
|
||||
],
|
||||
"title": "Scopes",
|
||||
"type": "string"
|
||||
},
|
||||
"jwt": {
|
||||
"default": "",
|
||||
"description": "JWT session token (if JWT auth is configured)",
|
||||
"title": "Jwt",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
@@ -541,6 +654,97 @@
|
||||
"title": "AuthLoginResponse",
|
||||
"type": "object"
|
||||
},
|
||||
"AuthSetupRequest": {
|
||||
"description": "POST /v1/api/auth/setup request body.",
|
||||
"properties": {
|
||||
"username": {
|
||||
"description": "Login username (1-64 ASCII characters)",
|
||||
"title": "Username",
|
||||
"type": "string"
|
||||
},
|
||||
"display_name": {
|
||||
"description": "Display name",
|
||||
"title": "Display Name",
|
||||
"type": "string"
|
||||
},
|
||||
"password": {
|
||||
"description": "Password (minimum 8 characters)",
|
||||
"title": "Password",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"username",
|
||||
"display_name",
|
||||
"password"
|
||||
],
|
||||
"title": "AuthSetupRequest",
|
||||
"type": "object"
|
||||
},
|
||||
"AuthSetupResponse": {
|
||||
"description": "POST /v1/api/auth/setup success response.",
|
||||
"properties": {
|
||||
"status": {
|
||||
"default": "ok",
|
||||
"title": "Status",
|
||||
"type": "string"
|
||||
},
|
||||
"user_id": {
|
||||
"title": "User Id",
|
||||
"type": "string"
|
||||
},
|
||||
"username": {
|
||||
"title": "Username",
|
||||
"type": "string"
|
||||
},
|
||||
"role": {
|
||||
"default": "full",
|
||||
"title": "Role",
|
||||
"type": "string"
|
||||
},
|
||||
"scopes": {
|
||||
"default": "approve,read,write",
|
||||
"title": "Scopes",
|
||||
"type": "string"
|
||||
},
|
||||
"jwt": {
|
||||
"default": "",
|
||||
"description": "JWT session token",
|
||||
"title": "Jwt",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"user_id",
|
||||
"username"
|
||||
],
|
||||
"title": "AuthSetupResponse",
|
||||
"type": "object"
|
||||
},
|
||||
"AuthStatusResponse": {
|
||||
"description": "GET /v1/api/auth/status response.",
|
||||
"properties": {
|
||||
"auth_enabled": {
|
||||
"title": "Auth Enabled",
|
||||
"type": "boolean"
|
||||
},
|
||||
"has_users": {
|
||||
"title": "Has Users",
|
||||
"type": "boolean"
|
||||
},
|
||||
"setup_required": {
|
||||
"title": "Setup Required",
|
||||
"type": "boolean"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"auth_enabled",
|
||||
"has_users",
|
||||
"setup_required"
|
||||
],
|
||||
"title": "AuthStatusResponse",
|
||||
"type": "object"
|
||||
},
|
||||
"SendRequest": {
|
||||
"properties": {
|
||||
"message": {
|
||||
@@ -677,6 +881,12 @@
|
||||
"description": "Auto-approve all tool calls",
|
||||
"title": "Auto Approve",
|
||||
"type": "boolean"
|
||||
},
|
||||
"resume_ws": {
|
||||
"default": "",
|
||||
"description": "Workstream ID to resume atomically during creation (empty = fresh start)",
|
||||
"title": "Resume Ws",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"title": "CreateWorkstreamRequest",
|
||||
@@ -693,6 +903,18 @@
|
||||
"description": "Assigned workstream name",
|
||||
"title": "Name",
|
||||
"type": "string"
|
||||
},
|
||||
"resumed": {
|
||||
"default": false,
|
||||
"description": "Whether a previous workstream was resumed",
|
||||
"title": "Resumed",
|
||||
"type": "boolean"
|
||||
},
|
||||
"message_count": {
|
||||
"default": 0,
|
||||
"description": "Number of messages in the resumed workstream",
|
||||
"title": "Message Count",
|
||||
"type": "integer"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
@@ -745,18 +967,6 @@
|
||||
"state": {
|
||||
"title": "State",
|
||||
"type": "string"
|
||||
},
|
||||
"session_id": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"default": null,
|
||||
"title": "Session Id"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
@@ -837,18 +1047,6 @@
|
||||
"title": "State",
|
||||
"type": "string"
|
||||
},
|
||||
"session_id": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"default": null,
|
||||
"title": "Session Id"
|
||||
},
|
||||
"title": {
|
||||
"default": "",
|
||||
"title": "Title",
|
||||
@@ -903,26 +1101,26 @@
|
||||
"title": "DashboardWorkstream",
|
||||
"type": "object"
|
||||
},
|
||||
"ListSessionsResponse": {
|
||||
"ListSavedWorkstreamsResponse": {
|
||||
"properties": {
|
||||
"sessions": {
|
||||
"workstreams": {
|
||||
"items": {
|
||||
"$ref": "#/components/schemas/SessionInfo"
|
||||
"$ref": "#/components/schemas/SavedWorkstreamInfo"
|
||||
},
|
||||
"title": "Sessions",
|
||||
"title": "Workstreams",
|
||||
"type": "array"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"sessions"
|
||||
"workstreams"
|
||||
],
|
||||
"title": "ListSessionsResponse",
|
||||
"title": "ListSavedWorkstreamsResponse",
|
||||
"type": "object"
|
||||
},
|
||||
"SessionInfo": {
|
||||
"SavedWorkstreamInfo": {
|
||||
"properties": {
|
||||
"session_id": {
|
||||
"title": "Session Id",
|
||||
"ws_id": {
|
||||
"title": "Ws Id",
|
||||
"type": "string"
|
||||
},
|
||||
"alias": {
|
||||
@@ -963,12 +1161,12 @@
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"session_id",
|
||||
"ws_id",
|
||||
"created",
|
||||
"updated",
|
||||
"message_count"
|
||||
],
|
||||
"title": "SessionInfo",
|
||||
"title": "SavedWorkstreamInfo",
|
||||
"type": "object"
|
||||
},
|
||||
"HealthResponse": {
|
||||
|
||||
@@ -2,15 +2,23 @@ import { BaseClient, type ClientOptions } from "./base.js";
|
||||
import type { ClusterEvent } from "./events.js";
|
||||
import type {
|
||||
AuthLoginResponse,
|
||||
AuthSetupResponse,
|
||||
AuthStatusResponse,
|
||||
ClusterNodesResponse,
|
||||
ClusterOverviewResponse,
|
||||
ClusterSnapshotResponse,
|
||||
ClusterWorkstreamsResponse,
|
||||
ConsoleCreateWsRequest,
|
||||
ConsoleCreateWsResponse,
|
||||
ConsoleHealthResponse,
|
||||
CreateScheduleRequest,
|
||||
ListScheduleRunsResponse,
|
||||
ListSchedulesResponse,
|
||||
NodeDetailResponse,
|
||||
NodesOptions,
|
||||
ScheduleInfo,
|
||||
StatusResponse,
|
||||
UpdateScheduleRequest,
|
||||
WorkstreamsOptions,
|
||||
} from "./types.js";
|
||||
|
||||
@@ -26,6 +34,10 @@ export class TurnstoneConsole extends BaseClient {
|
||||
return this.request("GET", "/v1/api/cluster/overview");
|
||||
}
|
||||
|
||||
async snapshot(): Promise<ClusterSnapshotResponse> {
|
||||
return this.request("GET", "/v1/api/cluster/snapshot");
|
||||
}
|
||||
|
||||
async nodes(opts?: NodesOptions): Promise<ClusterNodesResponse> {
|
||||
return this.request("GET", "/v1/api/cluster/nodes", {
|
||||
params: {
|
||||
@@ -70,9 +82,33 @@ export class TurnstoneConsole extends BaseClient {
|
||||
|
||||
// -- Auth -----------------------------------------------------------------
|
||||
|
||||
async login(token: string): Promise<AuthLoginResponse> {
|
||||
return this.request("POST", "/v1/api/auth/login", {
|
||||
json: { token },
|
||||
async login(opts: {
|
||||
token?: string;
|
||||
username?: string;
|
||||
password?: string;
|
||||
}): Promise<AuthLoginResponse> {
|
||||
const body =
|
||||
opts.username && opts.password
|
||||
? { username: opts.username, password: opts.password }
|
||||
: { token: opts.token ?? "" };
|
||||
return this.request("POST", "/v1/api/auth/login", { json: body });
|
||||
}
|
||||
|
||||
async authStatus(): Promise<AuthStatusResponse> {
|
||||
return this.request("GET", "/v1/api/auth/status");
|
||||
}
|
||||
|
||||
async setup(opts: {
|
||||
username: string;
|
||||
displayName: string;
|
||||
password: string;
|
||||
}): Promise<AuthSetupResponse> {
|
||||
return this.request("POST", "/v1/api/auth/setup", {
|
||||
json: {
|
||||
username: opts.username,
|
||||
display_name: opts.displayName,
|
||||
password: opts.password,
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
@@ -85,4 +121,40 @@ export class TurnstoneConsole extends BaseClient {
|
||||
async health(): Promise<ConsoleHealthResponse> {
|
||||
return this.request("GET", "/health");
|
||||
}
|
||||
|
||||
// -- Schedules ------------------------------------------------------------
|
||||
|
||||
async listSchedules(): Promise<ListSchedulesResponse> {
|
||||
return this.request("GET", "/v1/api/admin/schedules");
|
||||
}
|
||||
|
||||
async createSchedule(opts: CreateScheduleRequest): Promise<ScheduleInfo> {
|
||||
return this.request("POST", "/v1/api/admin/schedules", { json: opts });
|
||||
}
|
||||
|
||||
async getSchedule(taskId: string): Promise<ScheduleInfo> {
|
||||
return this.request("GET", `/v1/api/admin/schedules/${taskId}`);
|
||||
}
|
||||
|
||||
async updateSchedule(
|
||||
taskId: string,
|
||||
opts: UpdateScheduleRequest,
|
||||
): Promise<ScheduleInfo> {
|
||||
return this.request("PUT", `/v1/api/admin/schedules/${taskId}`, {
|
||||
json: opts,
|
||||
});
|
||||
}
|
||||
|
||||
async deleteSchedule(taskId: string): Promise<StatusResponse> {
|
||||
return this.request("DELETE", `/v1/api/admin/schedules/${taskId}`);
|
||||
}
|
||||
|
||||
async listScheduleRuns(
|
||||
taskId: string,
|
||||
opts?: { limit?: number },
|
||||
): Promise<ListScheduleRunsResponse> {
|
||||
return this.request("GET", `/v1/api/admin/schedules/${taskId}/runs`, {
|
||||
params: { limit: opts?.limit ?? 50 },
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import type { ClusterOverviewResponse, ClusterSnapshotNode } from "./types.js";
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Server SSE events
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -191,6 +193,13 @@ export interface ClusterWsRenameEvent {
|
||||
name: string;
|
||||
}
|
||||
|
||||
export interface ClusterSnapshotEvent {
|
||||
type: "snapshot";
|
||||
nodes: ClusterSnapshotNode[];
|
||||
overview: ClusterOverviewResponse;
|
||||
timestamp: number;
|
||||
}
|
||||
|
||||
/** Discriminated union of all console cluster SSE event types. */
|
||||
export type ClusterEvent =
|
||||
| NodeJoinedEvent
|
||||
@@ -198,7 +207,8 @@ export type ClusterEvent =
|
||||
| ClusterStateEvent
|
||||
| ClusterWsCreatedEvent
|
||||
| ClusterWsClosedEvent
|
||||
| ClusterWsRenameEvent;
|
||||
| ClusterWsRenameEvent
|
||||
| ClusterSnapshotEvent;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Type guards
|
||||
|
||||
@@ -55,6 +55,7 @@ export type {
|
||||
ClusterWsCreatedEvent,
|
||||
ClusterWsClosedEvent,
|
||||
ClusterWsRenameEvent,
|
||||
ClusterSnapshotEvent,
|
||||
} from "./events.js";
|
||||
|
||||
export {
|
||||
@@ -83,24 +84,34 @@ export type {
|
||||
DashboardWorkstream,
|
||||
DashboardAggregate,
|
||||
DashboardResponse,
|
||||
SessionInfo,
|
||||
ListSessionsResponse,
|
||||
SavedWorkstreamInfo,
|
||||
ListSavedWorkstreamsResponse,
|
||||
BackendStatus,
|
||||
WorkstreamCounts,
|
||||
HealthResponse,
|
||||
AuthLoginRequest,
|
||||
AuthLoginResponse,
|
||||
AuthStatusResponse,
|
||||
AuthSetupResponse,
|
||||
StatusResponse,
|
||||
ErrorResponse,
|
||||
ClusterOverviewResponse,
|
||||
ClusterNodeInfo,
|
||||
ClusterNodesResponse,
|
||||
ClusterSnapshotNode,
|
||||
ClusterSnapshotResponse,
|
||||
ClusterWorkstreamInfo,
|
||||
ClusterWorkstreamsResponse,
|
||||
NodeDetailResponse,
|
||||
ConsoleCreateWsRequest,
|
||||
ConsoleCreateWsResponse,
|
||||
ConsoleHealthResponse,
|
||||
CreateScheduleRequest,
|
||||
UpdateScheduleRequest,
|
||||
ScheduleInfo,
|
||||
ScheduleRunInfo,
|
||||
ListSchedulesResponse,
|
||||
ListScheduleRunsResponse,
|
||||
TurnResult,
|
||||
SendAndWaitOptions,
|
||||
NodesOptions,
|
||||
|
||||
@@ -2,11 +2,13 @@ import { BaseClient, type ClientOptions } from "./base.js";
|
||||
import type { ServerEvent } from "./events.js";
|
||||
import type {
|
||||
AuthLoginResponse,
|
||||
AuthSetupResponse,
|
||||
AuthStatusResponse,
|
||||
CreateWorkstreamRequest,
|
||||
CreateWorkstreamResponse,
|
||||
DashboardResponse,
|
||||
HealthResponse,
|
||||
ListSessionsResponse,
|
||||
ListSavedWorkstreamsResponse,
|
||||
ListWorkstreamsResponse,
|
||||
SendAndWaitOptions,
|
||||
SendResponse,
|
||||
@@ -176,17 +178,41 @@ export class TurnstoneServer extends BaseClient {
|
||||
return result;
|
||||
}
|
||||
|
||||
// -- Sessions -------------------------------------------------------------
|
||||
// -- Saved workstreams ----------------------------------------------------
|
||||
|
||||
async listSessions(): Promise<ListSessionsResponse> {
|
||||
return this.request("GET", "/v1/api/sessions");
|
||||
async listSavedWorkstreams(): Promise<ListSavedWorkstreamsResponse> {
|
||||
return this.request("GET", "/v1/api/workstreams/saved");
|
||||
}
|
||||
|
||||
// -- Auth -----------------------------------------------------------------
|
||||
|
||||
async login(token: string): Promise<AuthLoginResponse> {
|
||||
return this.request("POST", "/v1/api/auth/login", {
|
||||
json: { token },
|
||||
async login(opts: {
|
||||
token?: string;
|
||||
username?: string;
|
||||
password?: string;
|
||||
}): Promise<AuthLoginResponse> {
|
||||
const body =
|
||||
opts.username && opts.password
|
||||
? { username: opts.username, password: opts.password }
|
||||
: { token: opts.token ?? "" };
|
||||
return this.request("POST", "/v1/api/auth/login", { json: body });
|
||||
}
|
||||
|
||||
async authStatus(): Promise<AuthStatusResponse> {
|
||||
return this.request("GET", "/v1/api/auth/status");
|
||||
}
|
||||
|
||||
async setup(opts: {
|
||||
username: string;
|
||||
displayName: string;
|
||||
password: string;
|
||||
}): Promise<AuthSetupResponse> {
|
||||
return this.request("POST", "/v1/api/auth/setup", {
|
||||
json: {
|
||||
username: opts.username,
|
||||
display_name: opts.displayName,
|
||||
password: opts.password,
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
|
||||
+114
-7
@@ -17,6 +17,24 @@ export interface AuthLoginRequest {
|
||||
export interface AuthLoginResponse {
|
||||
status: string;
|
||||
role: string;
|
||||
scopes?: string;
|
||||
jwt?: string;
|
||||
user_id?: string;
|
||||
}
|
||||
|
||||
export interface AuthStatusResponse {
|
||||
auth_enabled: boolean;
|
||||
has_users: boolean;
|
||||
setup_required: boolean;
|
||||
}
|
||||
|
||||
export interface AuthSetupResponse {
|
||||
status: string;
|
||||
user_id: string;
|
||||
username: string;
|
||||
role: string;
|
||||
scopes: string;
|
||||
jwt?: string;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -53,11 +71,14 @@ export interface CreateWorkstreamRequest {
|
||||
name?: string;
|
||||
model?: string;
|
||||
auto_approve?: boolean;
|
||||
resume_ws?: string;
|
||||
}
|
||||
|
||||
export interface CreateWorkstreamResponse {
|
||||
ws_id: string;
|
||||
name: string;
|
||||
resumed?: boolean;
|
||||
message_count?: number;
|
||||
}
|
||||
|
||||
export interface CloseWorkstreamRequest {
|
||||
@@ -68,7 +89,6 @@ export interface WorkstreamInfo {
|
||||
id: string;
|
||||
name: string;
|
||||
state: string;
|
||||
session_id?: string | null;
|
||||
}
|
||||
|
||||
export interface ListWorkstreamsResponse {
|
||||
@@ -79,7 +99,6 @@ export interface DashboardWorkstream {
|
||||
id: string;
|
||||
name: string;
|
||||
state: string;
|
||||
session_id?: string | null;
|
||||
title?: string;
|
||||
tokens?: number;
|
||||
context_ratio?: number;
|
||||
@@ -106,11 +125,11 @@ export interface DashboardResponse {
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Server API — Sessions
|
||||
// Server API — Saved workstreams
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
export interface SessionInfo {
|
||||
session_id: string;
|
||||
export interface SavedWorkstreamInfo {
|
||||
ws_id: string;
|
||||
alias?: string | null;
|
||||
title?: string | null;
|
||||
created: string;
|
||||
@@ -118,8 +137,8 @@ export interface SessionInfo {
|
||||
message_count: number;
|
||||
}
|
||||
|
||||
export interface ListSessionsResponse {
|
||||
sessions: SessionInfo[];
|
||||
export interface ListSavedWorkstreamsResponse {
|
||||
workstreams: SavedWorkstreamInfo[];
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -225,6 +244,23 @@ export interface NodeDetailResponse {
|
||||
aggregate: ClusterAggregate;
|
||||
}
|
||||
|
||||
export interface ClusterSnapshotNode {
|
||||
node_id: string;
|
||||
server_url: string;
|
||||
max_ws: number;
|
||||
reachable: boolean;
|
||||
version: string;
|
||||
health: Record<string, string>;
|
||||
aggregate: Record<string, number>;
|
||||
workstreams: ClusterWorkstreamInfo[];
|
||||
}
|
||||
|
||||
export interface ClusterSnapshotResponse {
|
||||
nodes: ClusterSnapshotNode[];
|
||||
overview: ClusterOverviewResponse;
|
||||
timestamp: number;
|
||||
}
|
||||
|
||||
export interface ConsoleCreateWsRequest {
|
||||
node_id?: string;
|
||||
name?: string;
|
||||
@@ -247,6 +283,77 @@ export interface ConsoleHealthResponse {
|
||||
versions: string[];
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Console API — Schedules
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
export interface CreateScheduleRequest {
|
||||
name: string;
|
||||
schedule_type: string;
|
||||
initial_message: string;
|
||||
description?: string;
|
||||
cron_expr?: string;
|
||||
at_time?: string;
|
||||
target_mode?: string;
|
||||
model?: string;
|
||||
auto_approve?: boolean;
|
||||
auto_approve_tools?: string[];
|
||||
enabled?: boolean;
|
||||
}
|
||||
|
||||
export interface UpdateScheduleRequest {
|
||||
name?: string;
|
||||
description?: string;
|
||||
schedule_type?: string;
|
||||
cron_expr?: string;
|
||||
at_time?: string;
|
||||
target_mode?: string;
|
||||
model?: string;
|
||||
initial_message?: string;
|
||||
auto_approve?: boolean;
|
||||
auto_approve_tools?: string[];
|
||||
enabled?: boolean;
|
||||
}
|
||||
|
||||
export interface ScheduleInfo {
|
||||
task_id: string;
|
||||
name: string;
|
||||
description: string;
|
||||
schedule_type: string;
|
||||
cron_expr: string;
|
||||
at_time: string;
|
||||
target_mode: string;
|
||||
model: string;
|
||||
initial_message: string;
|
||||
auto_approve: boolean;
|
||||
auto_approve_tools: string[];
|
||||
enabled: boolean;
|
||||
created_by: string;
|
||||
last_run: string | null;
|
||||
next_run: string | null;
|
||||
created: string;
|
||||
updated: string;
|
||||
}
|
||||
|
||||
export interface ListSchedulesResponse {
|
||||
schedules: ScheduleInfo[];
|
||||
}
|
||||
|
||||
export interface ScheduleRunInfo {
|
||||
run_id: string;
|
||||
task_id: string;
|
||||
node_id: string;
|
||||
ws_id: string;
|
||||
correlation_id: string;
|
||||
started: string;
|
||||
status: string;
|
||||
error: string;
|
||||
}
|
||||
|
||||
export interface ListScheduleRunsResponse {
|
||||
runs: ScheduleRunInfo[];
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SDK-specific types
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
"""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
|
||||
+388
-79
@@ -1,7 +1,9 @@
|
||||
"""Tests for turnstone.core.auth — bearer token authentication and cookies."""
|
||||
|
||||
import os
|
||||
from unittest.mock import patch
|
||||
import queue
|
||||
import threading
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -15,7 +17,7 @@ from turnstone.core.auth import (
|
||||
load_auth_config,
|
||||
make_clear_cookie,
|
||||
make_set_cookie,
|
||||
required_role,
|
||||
required_scope,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -87,69 +89,59 @@ class TestIsPublicPath:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRequiredRole:
|
||||
class TestRequiredScope:
|
||||
def test_get_api_needs_read(self):
|
||||
assert required_role("GET", "/api/workstreams") == "read"
|
||||
assert required_scope("GET", "/api/workstreams") == "read"
|
||||
|
||||
def test_get_events_needs_read(self):
|
||||
assert required_role("GET", "/api/events") == "read"
|
||||
assert required_scope("GET", "/api/events") == "read"
|
||||
|
||||
def test_get_dashboard_needs_read(self):
|
||||
assert required_role("GET", "/api/dashboard") == "read"
|
||||
def test_post_send_needs_write(self):
|
||||
assert required_scope("POST", "/api/send") == "write"
|
||||
|
||||
def test_post_send_needs_full(self):
|
||||
assert required_role("POST", "/api/send") == "full"
|
||||
def test_post_approve_needs_approve(self):
|
||||
assert required_scope("POST", "/api/approve") == "approve"
|
||||
|
||||
def test_post_approve_needs_full(self):
|
||||
assert required_role("POST", "/api/approve") == "full"
|
||||
def test_post_plan_needs_write(self):
|
||||
assert required_scope("POST", "/api/plan") == "write"
|
||||
|
||||
def test_post_plan_needs_full(self):
|
||||
assert required_role("POST", "/api/plan") == "full"
|
||||
def test_post_command_needs_write(self):
|
||||
assert required_scope("POST", "/api/command") == "write"
|
||||
|
||||
def test_post_command_needs_full(self):
|
||||
assert required_role("POST", "/api/command") == "full"
|
||||
def test_post_workstreams_new_needs_write(self):
|
||||
assert required_scope("POST", "/api/workstreams/new") == "write"
|
||||
|
||||
def test_post_workstreams_new_needs_full(self):
|
||||
assert required_role("POST", "/api/workstreams/new") == "full"
|
||||
def test_post_workstreams_close_needs_write(self):
|
||||
assert required_scope("POST", "/api/workstreams/close") == "write"
|
||||
|
||||
def test_post_workstreams_close_needs_full(self):
|
||||
assert required_role("POST", "/api/workstreams/close") == "full"
|
||||
|
||||
def test_all_write_paths_need_full(self):
|
||||
def test_all_write_paths_need_write(self):
|
||||
for path in WRITE_PATHS:
|
||||
assert required_role("POST", path) == "full"
|
||||
scope = required_scope("POST", path)
|
||||
assert scope in ("write", "approve"), f"{path} should need write or approve"
|
||||
|
||||
def test_post_unknown_path_needs_read(self):
|
||||
assert required_role("POST", "/api/unknown") == "read"
|
||||
assert required_scope("POST", "/api/unknown") == "read"
|
||||
|
||||
def test_v1_post_send_needs_full(self):
|
||||
assert required_role("POST", "/v1/api/send") == "full"
|
||||
def test_v1_post_send_needs_write(self):
|
||||
assert required_scope("POST", "/v1/api/send") == "write"
|
||||
|
||||
def test_v1_post_approve_needs_full(self):
|
||||
assert required_role("POST", "/v1/api/approve") == "full"
|
||||
def test_v1_post_approve_needs_approve(self):
|
||||
assert required_scope("POST", "/v1/api/approve") == "approve"
|
||||
|
||||
def test_v1_get_workstreams_needs_read(self):
|
||||
assert required_role("GET", "/v1/api/workstreams") == "read"
|
||||
assert required_scope("GET", "/v1/api/workstreams") == "read"
|
||||
|
||||
def test_v1_post_cluster_ws_new_needs_full(self):
|
||||
assert required_role("POST", "/v1/api/cluster/workstreams/new") == "full"
|
||||
def test_v1_post_cluster_ws_new_needs_write(self):
|
||||
assert required_scope("POST", "/v1/api/cluster/workstreams/new") == "write"
|
||||
|
||||
def test_v1_all_write_paths_need_full(self):
|
||||
for path in WRITE_PATHS:
|
||||
v1_path = "/v1" + path
|
||||
assert required_role("POST", v1_path) == "full", f"{v1_path} should need full"
|
||||
def test_proxy_v1_send_needs_write(self):
|
||||
assert required_scope("POST", "/node/node-a/v1/api/send") == "write"
|
||||
|
||||
def test_proxy_v1_send_needs_full(self):
|
||||
assert required_role("POST", "/node/node-a/v1/api/send") == "full"
|
||||
|
||||
def test_proxy_v1_approve_needs_full(self):
|
||||
assert required_role("POST", "/node/node-a/v1/api/approve") == "full"
|
||||
|
||||
def test_proxy_v1_cluster_ws_new_needs_full(self):
|
||||
assert required_role("POST", "/node/node-a/v1/api/cluster/workstreams/new") == "full"
|
||||
def test_proxy_v1_approve_needs_approve(self):
|
||||
assert required_scope("POST", "/node/node-a/v1/api/approve") == "approve"
|
||||
|
||||
def test_proxy_v1_read_endpoint_needs_read(self):
|
||||
assert required_role("GET", "/node/node-a/v1/api/workstreams") == "read"
|
||||
assert required_scope("GET", "/node/node-a/v1/api/workstreams") == "read"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -268,12 +260,24 @@ class TestMakeSetCookie:
|
||||
|
||||
def test_max_age_default(self):
|
||||
val = make_set_cookie("tok_abc")
|
||||
assert "Max-Age=2592000" in val # 30 days
|
||||
assert "Max-Age=86400" in val # 24 hours (matches JWT expiry)
|
||||
|
||||
def test_max_age_custom(self):
|
||||
val = make_set_cookie("tok_abc", max_age=3600)
|
||||
assert "Max-Age=3600" in val
|
||||
|
||||
def test_secure_default(self):
|
||||
val = make_set_cookie("tok_abc")
|
||||
assert "; Secure" in val
|
||||
|
||||
def test_secure_false(self):
|
||||
val = make_set_cookie("tok_abc", secure=False)
|
||||
assert "; Secure" not in val
|
||||
|
||||
def test_secure_true(self):
|
||||
val = make_set_cookie("tok_abc", secure=True)
|
||||
assert "; Secure" in val
|
||||
|
||||
|
||||
class TestMakeClearCookie:
|
||||
def test_max_age_zero(self):
|
||||
@@ -306,68 +310,78 @@ class TestCheckRequest:
|
||||
)
|
||||
|
||||
def test_disabled_allows_all(self, disabled):
|
||||
allowed, status, msg = check_request(disabled, "POST", "/api/send", None)
|
||||
allowed, status, msg, _result = check_request(disabled, "POST", "/api/send", None)
|
||||
assert allowed is True
|
||||
assert status == 200
|
||||
|
||||
def test_disabled_allows_no_header(self, disabled):
|
||||
allowed, status, msg = check_request(disabled, "GET", "/api/workstreams", None)
|
||||
allowed, status, msg, _result = check_request(disabled, "GET", "/api/workstreams", None)
|
||||
assert allowed is True
|
||||
|
||||
def test_public_path_no_token_ok(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "GET", "/health", None)
|
||||
allowed, status, msg, _result = check_request(enabled, "GET", "/health", None)
|
||||
assert allowed is True
|
||||
assert status == 200
|
||||
|
||||
def test_public_root_no_token_ok(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "GET", "/", None)
|
||||
allowed, status, msg, _result = check_request(enabled, "GET", "/", None)
|
||||
assert allowed is True
|
||||
|
||||
def test_public_static_no_token_ok(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "GET", "/static/style.css", None)
|
||||
allowed, status, msg, _result = check_request(enabled, "GET", "/static/style.css", None)
|
||||
assert allowed is True
|
||||
|
||||
def test_api_no_token_401(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "GET", "/api/workstreams", None)
|
||||
allowed, status, msg, _result = check_request(enabled, "GET", "/api/workstreams", None)
|
||||
assert allowed is False
|
||||
assert status == 401
|
||||
assert "Unauthorized" in msg
|
||||
|
||||
def test_api_invalid_token_401(self, enabled):
|
||||
allowed, status, msg = check_request(
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "GET", "/api/workstreams", "Bearer wrong_token"
|
||||
)
|
||||
assert allowed is False
|
||||
assert status == 401
|
||||
|
||||
def test_api_read_token_ok(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "GET", "/api/workstreams", "Bearer tok_read")
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "GET", "/api/workstreams", "Bearer tok_read"
|
||||
)
|
||||
assert allowed is True
|
||||
assert status == 200
|
||||
|
||||
def test_api_full_token_ok(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "GET", "/api/workstreams", "Bearer tok_full")
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "GET", "/api/workstreams", "Bearer tok_full"
|
||||
)
|
||||
assert allowed is True
|
||||
|
||||
def test_write_read_token_403(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "POST", "/api/send", "Bearer tok_read")
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/api/send", "Bearer tok_read"
|
||||
)
|
||||
assert allowed is False
|
||||
assert status == 403
|
||||
assert "Forbidden" in msg
|
||||
|
||||
def test_write_full_token_ok(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "POST", "/api/send", "Bearer tok_full")
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/api/send", "Bearer tok_full"
|
||||
)
|
||||
assert allowed is True
|
||||
assert status == 200
|
||||
|
||||
def test_approve_read_token_403(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "POST", "/api/approve", "Bearer tok_read")
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/api/approve", "Bearer tok_read"
|
||||
)
|
||||
assert allowed is False
|
||||
assert status == 403
|
||||
|
||||
def test_proxy_write_read_token_403(self, enabled):
|
||||
"""Read tokens cannot escalate to write ops via proxy routes."""
|
||||
allowed, status, msg = check_request(
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/node/node-a/api/send", "Bearer tok_read"
|
||||
)
|
||||
assert allowed is False
|
||||
@@ -375,7 +389,7 @@ class TestCheckRequest:
|
||||
|
||||
def test_proxy_write_trailing_slash_read_token_403(self, enabled):
|
||||
"""Trailing slash must not bypass write-role check on proxy routes."""
|
||||
allowed, status, msg = check_request(
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/node/node-a/api/send/", "Bearer tok_read"
|
||||
)
|
||||
assert allowed is False
|
||||
@@ -383,20 +397,22 @@ class TestCheckRequest:
|
||||
|
||||
def test_direct_write_trailing_slash_read_token_403(self, enabled):
|
||||
"""Trailing slash must not bypass write-role check on direct routes."""
|
||||
allowed, status, msg = check_request(enabled, "POST", "/api/send/", "Bearer tok_read")
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/api/send/", "Bearer tok_read"
|
||||
)
|
||||
assert allowed is False
|
||||
assert status == 403
|
||||
|
||||
def test_proxy_write_full_token_ok(self, enabled):
|
||||
"""Full tokens pass through proxy write routes."""
|
||||
allowed, status, msg = check_request(
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/node/node-a/api/send", "Bearer tok_full"
|
||||
)
|
||||
assert allowed is True
|
||||
|
||||
def test_proxy_v1_write_read_token_403(self, enabled):
|
||||
"""Read tokens cannot escalate to write ops via v1 proxy routes."""
|
||||
allowed, status, msg = check_request(
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/node/node-a/v1/api/send", "Bearer tok_read"
|
||||
)
|
||||
assert allowed is False
|
||||
@@ -404,14 +420,14 @@ class TestCheckRequest:
|
||||
|
||||
def test_proxy_v1_write_full_token_ok(self, enabled):
|
||||
"""Full tokens pass through v1 proxy write routes."""
|
||||
allowed, status, msg = check_request(
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/node/node-a/v1/api/send", "Bearer tok_full"
|
||||
)
|
||||
assert allowed is True
|
||||
|
||||
def test_proxy_v1_cluster_ws_new_read_403(self, enabled):
|
||||
"""Read tokens cannot create workstreams via v1 proxy."""
|
||||
allowed, status, msg = check_request(
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled,
|
||||
"POST",
|
||||
"/node/node-a/v1/api/cluster/workstreams/new",
|
||||
@@ -422,25 +438,27 @@ class TestCheckRequest:
|
||||
|
||||
def test_proxy_read_endpoint_read_token_ok(self, enabled):
|
||||
"""Read tokens can access proxy read endpoints."""
|
||||
allowed, status, msg = check_request(
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "GET", "/node/node-a/api/workstreams", "Bearer tok_read"
|
||||
)
|
||||
assert allowed is True
|
||||
|
||||
def test_console_create_ws_read_token_403(self, enabled):
|
||||
"""Read tokens cannot create workstreams."""
|
||||
allowed, status, msg = check_request(
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/api/cluster/workstreams/new", "Bearer tok_read"
|
||||
)
|
||||
assert allowed is False
|
||||
assert status == 403
|
||||
|
||||
def test_approve_full_token_ok(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "POST", "/api/approve", "Bearer tok_full")
|
||||
allowed, status, msg, _result = check_request(
|
||||
enabled, "POST", "/api/approve", "Bearer tok_full"
|
||||
)
|
||||
assert allowed is True
|
||||
|
||||
def test_no_auth_header_string(self, enabled):
|
||||
allowed, status, msg = check_request(enabled, "GET", "/api/dashboard", "")
|
||||
allowed, status, msg, _result = check_request(enabled, "GET", "/api/dashboard", "")
|
||||
assert allowed is False
|
||||
assert status == 401
|
||||
|
||||
@@ -461,7 +479,7 @@ class TestCheckRequestWithCookie:
|
||||
)
|
||||
|
||||
def test_cookie_fallback_when_no_bearer(self, enabled):
|
||||
allowed, status, _ = check_request(
|
||||
allowed, status, _, _r = check_request(
|
||||
enabled,
|
||||
"GET",
|
||||
"/api/workstreams",
|
||||
@@ -473,7 +491,7 @@ class TestCheckRequestWithCookie:
|
||||
|
||||
def test_bearer_takes_precedence_over_cookie(self, enabled):
|
||||
# Bearer is full, cookie is read — Bearer should win
|
||||
allowed, status, _ = check_request(
|
||||
allowed, status, _, _r = check_request(
|
||||
enabled,
|
||||
"POST",
|
||||
"/api/send",
|
||||
@@ -483,7 +501,7 @@ class TestCheckRequestWithCookie:
|
||||
assert allowed is True
|
||||
|
||||
def test_invalid_cookie_401(self, enabled):
|
||||
allowed, status, _ = check_request(
|
||||
allowed, status, _, _r = check_request(
|
||||
enabled,
|
||||
"GET",
|
||||
"/api/workstreams",
|
||||
@@ -494,7 +512,7 @@ class TestCheckRequestWithCookie:
|
||||
assert status == 401
|
||||
|
||||
def test_cookie_read_on_write_403(self, enabled):
|
||||
allowed, status, _ = check_request(
|
||||
allowed, status, _, _r = check_request(
|
||||
enabled,
|
||||
"POST",
|
||||
"/api/send",
|
||||
@@ -505,7 +523,7 @@ class TestCheckRequestWithCookie:
|
||||
assert status == 403
|
||||
|
||||
def test_cookie_full_on_write_ok(self, enabled):
|
||||
allowed, status, _ = check_request(
|
||||
allowed, status, _, _r = check_request(
|
||||
enabled,
|
||||
"POST",
|
||||
"/api/send",
|
||||
@@ -515,7 +533,7 @@ class TestCheckRequestWithCookie:
|
||||
assert allowed is True
|
||||
|
||||
def test_no_cookie_no_bearer_401(self, enabled):
|
||||
allowed, status, _ = check_request(
|
||||
allowed, status, _, _r = check_request(
|
||||
enabled,
|
||||
"GET",
|
||||
"/api/workstreams",
|
||||
@@ -526,7 +544,7 @@ class TestCheckRequestWithCookie:
|
||||
assert status == 401
|
||||
|
||||
def test_login_path_public(self, enabled):
|
||||
allowed, status, _ = check_request(
|
||||
allowed, status, _, _r = check_request(
|
||||
enabled,
|
||||
"POST",
|
||||
"/api/auth/login",
|
||||
@@ -535,7 +553,7 @@ class TestCheckRequestWithCookie:
|
||||
assert allowed is True
|
||||
|
||||
def test_logout_path_public(self, enabled):
|
||||
allowed, status, _ = check_request(
|
||||
allowed, status, _, _r = check_request(
|
||||
enabled,
|
||||
"POST",
|
||||
"/api/auth/logout",
|
||||
@@ -552,12 +570,28 @@ class TestCheckRequestWithCookie:
|
||||
class TestLoadAuthConfig:
|
||||
"""Tests for load_auth_config with mocked config + env vars."""
|
||||
|
||||
def test_default_disabled(self):
|
||||
def test_default_enabled(self):
|
||||
with patch("turnstone.core.config.load_config", return_value={}):
|
||||
cfg = load_auth_config()
|
||||
assert cfg.enabled is False
|
||||
assert cfg.enabled is True
|
||||
assert cfg.tokens == {}
|
||||
|
||||
def test_explicit_disable(self):
|
||||
with (
|
||||
patch("turnstone.core.config.load_config", return_value={"enabled": False}),
|
||||
patch.dict(os.environ, {}, clear=True),
|
||||
):
|
||||
cfg = load_auth_config()
|
||||
assert cfg.enabled is False
|
||||
|
||||
def test_env_disable(self):
|
||||
with (
|
||||
patch("turnstone.core.config.load_config", return_value={}),
|
||||
patch.dict(os.environ, {"TURNSTONE_AUTH_ENABLED": "0"}, clear=True),
|
||||
):
|
||||
cfg = load_auth_config()
|
||||
assert cfg.enabled is False
|
||||
|
||||
def test_config_file_tokens(self):
|
||||
mock_cfg = {
|
||||
"enabled": True,
|
||||
@@ -685,7 +719,7 @@ class TestServerAuth:
|
||||
srv_mod._metrics.model = "test-model"
|
||||
|
||||
mock_session = MagicMock()
|
||||
mock_session.session_id = "test-session-id"
|
||||
mock_session.ws_id = "test-session-id"
|
||||
|
||||
mock_ws = MagicMock()
|
||||
mock_ws.id = "test-ws"
|
||||
@@ -705,6 +739,7 @@ class TestServerAuth:
|
||||
enabled=True,
|
||||
tokens={"tok_full": "full", "tok_read": "read"},
|
||||
),
|
||||
cors_origins=["*"],
|
||||
)
|
||||
cls.client = TestClient(app, raise_server_exceptions=False)
|
||||
|
||||
@@ -902,7 +937,7 @@ class TestServerLogin:
|
||||
srv_mod._metrics.model = "test-model"
|
||||
|
||||
mock_session = MagicMock()
|
||||
mock_session.session_id = "test-session-id"
|
||||
mock_session.ws_id = "test-session-id"
|
||||
|
||||
mock_ws = MagicMock()
|
||||
mock_ws.id = "test-ws"
|
||||
@@ -1040,3 +1075,277 @@ class TestConsoleLogin:
|
||||
self.test_client.post("/v1/api/auth/logout")
|
||||
resp = self.test_client.get("/v1/api/cluster/overview")
|
||||
assert resp.status_code == 401
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Security hardening tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLoginRateLimiter:
|
||||
def test_allows_under_limit(self):
|
||||
from turnstone.core.auth import LoginRateLimiter
|
||||
|
||||
limiter = LoginRateLimiter(max_attempts=3, window_seconds=60)
|
||||
for _ in range(3):
|
||||
ok, _ = limiter.check("ip:1.2.3.4")
|
||||
assert ok
|
||||
limiter.record("ip:1.2.3.4")
|
||||
# 4th should be blocked (3 recorded)
|
||||
ok, retry = limiter.check("ip:1.2.3.4")
|
||||
assert not ok
|
||||
assert retry > 0
|
||||
|
||||
def test_different_keys_independent(self):
|
||||
from turnstone.core.auth import LoginRateLimiter
|
||||
|
||||
limiter = LoginRateLimiter(max_attempts=2, window_seconds=60)
|
||||
limiter.record("ip:a")
|
||||
limiter.record("ip:a")
|
||||
ok_a, _ = limiter.check("ip:a")
|
||||
ok_b, _ = limiter.check("ip:b")
|
||||
assert not ok_a
|
||||
assert ok_b
|
||||
|
||||
def test_cleanup(self):
|
||||
from turnstone.core.auth import LoginRateLimiter
|
||||
|
||||
limiter = LoginRateLimiter(max_attempts=1, window_seconds=60)
|
||||
limiter.record("ip:old")
|
||||
removed = limiter.cleanup(max_age=0.0)
|
||||
assert removed == 1
|
||||
ok, _ = limiter.check("ip:old")
|
||||
assert ok
|
||||
|
||||
def test_max_keys_protection(self):
|
||||
from turnstone.core.auth import LoginRateLimiter
|
||||
|
||||
limiter = LoginRateLimiter(max_attempts=5, window_seconds=60)
|
||||
limiter.MAX_KEYS = 2
|
||||
limiter.record("a")
|
||||
limiter.record("b")
|
||||
limiter.record("c") # should be silently dropped (at capacity)
|
||||
assert "c" not in limiter._attempts
|
||||
|
||||
|
||||
class TestJWTAudienceIssuer:
|
||||
SECRET = "test-secret-that-is-at-least-32-chars"
|
||||
|
||||
def test_create_jwt_includes_iss(self):
|
||||
import jwt as pyjwt
|
||||
|
||||
from turnstone.core.auth import JWT_ISSUER, create_jwt
|
||||
|
||||
token = create_jwt("user1", frozenset({"read"}), "test", self.SECRET)
|
||||
payload = pyjwt.decode(
|
||||
token, self.SECRET, algorithms=["HS256"], options={"verify_aud": False}
|
||||
)
|
||||
assert payload["iss"] == JWT_ISSUER
|
||||
|
||||
def test_create_jwt_with_audience(self):
|
||||
import jwt as pyjwt
|
||||
|
||||
from turnstone.core.auth import JWT_AUD_SERVER, create_jwt
|
||||
|
||||
token = create_jwt(
|
||||
"user1", frozenset({"read"}), "test", self.SECRET, audience=JWT_AUD_SERVER
|
||||
)
|
||||
payload = pyjwt.decode(token, self.SECRET, algorithms=["HS256"], audience=JWT_AUD_SERVER)
|
||||
assert payload["aud"] == JWT_AUD_SERVER
|
||||
|
||||
def test_validate_jwt_wrong_audience_rejected(self):
|
||||
from turnstone.core.auth import JWT_AUD_CONSOLE, JWT_AUD_SERVER, create_jwt, validate_jwt
|
||||
|
||||
token = create_jwt(
|
||||
"user1", frozenset({"read"}), "test", self.SECRET, audience=JWT_AUD_SERVER
|
||||
)
|
||||
result = validate_jwt(token, self.SECRET, audience=JWT_AUD_CONSOLE)
|
||||
assert result is None
|
||||
|
||||
def test_validate_jwt_correct_audience_accepted(self):
|
||||
from turnstone.core.auth import JWT_AUD_SERVER, create_jwt, validate_jwt
|
||||
|
||||
token = create_jwt(
|
||||
"user1", frozenset({"read"}), "test", self.SECRET, audience=JWT_AUD_SERVER
|
||||
)
|
||||
result = validate_jwt(token, self.SECRET, audience=JWT_AUD_SERVER)
|
||||
assert result is not None
|
||||
assert result.user_id == "user1"
|
||||
|
||||
def test_validate_jwt_no_audience_backward_compat(self):
|
||||
from turnstone.core.auth import create_jwt, validate_jwt
|
||||
|
||||
# Token without aud claim should be accepted when audience="" (backward compat)
|
||||
token = create_jwt("user1", frozenset({"read"}), "test", self.SECRET)
|
||||
result = validate_jwt(token, self.SECRET, audience="")
|
||||
assert result is not None
|
||||
|
||||
|
||||
class TestServiceTokenManager:
|
||||
SECRET = "test-secret-that-is-at-least-32-chars"
|
||||
|
||||
def test_auto_mints_on_first_access(self):
|
||||
from turnstone.core.auth import ServiceTokenManager
|
||||
|
||||
mgr = ServiceTokenManager(
|
||||
user_id="svc",
|
||||
scopes=frozenset({"read"}),
|
||||
source="test",
|
||||
secret=self.SECRET,
|
||||
)
|
||||
token = mgr.token
|
||||
assert token # non-empty
|
||||
assert isinstance(token, str)
|
||||
|
||||
def test_bearer_header_format(self):
|
||||
from turnstone.core.auth import ServiceTokenManager
|
||||
|
||||
mgr = ServiceTokenManager(
|
||||
user_id="svc",
|
||||
scopes=frozenset({"read"}),
|
||||
source="test",
|
||||
secret=self.SECRET,
|
||||
)
|
||||
header = mgr.bearer_header
|
||||
assert "Authorization" in header
|
||||
assert header["Authorization"].startswith("Bearer ")
|
||||
|
||||
def test_token_stable_within_window(self):
|
||||
from turnstone.core.auth import ServiceTokenManager
|
||||
|
||||
mgr = ServiceTokenManager(
|
||||
user_id="svc",
|
||||
scopes=frozenset({"read"}),
|
||||
source="test",
|
||||
secret=self.SECRET,
|
||||
expiry_hours=1,
|
||||
)
|
||||
t1 = mgr.token
|
||||
t2 = mgr.token
|
||||
assert t1 == t2
|
||||
|
||||
def test_token_rotates_near_expiry(self):
|
||||
from turnstone.core.auth import ServiceTokenManager
|
||||
|
||||
mgr = ServiceTokenManager(
|
||||
user_id="svc",
|
||||
scopes=frozenset({"read"}),
|
||||
source="test",
|
||||
secret=self.SECRET,
|
||||
expiry_hours=1,
|
||||
)
|
||||
_ = mgr.token # initial mint
|
||||
# Simulate expiry by backdating _expires_at
|
||||
mgr._expires_at = 0.0
|
||||
t2 = mgr.token
|
||||
# Token was re-minted (even if payload matches within same second,
|
||||
# the internal state was refreshed)
|
||||
assert t2 # non-empty, valid token
|
||||
assert mgr._expires_at > 0.0 # was refreshed
|
||||
|
||||
def test_audience_included(self):
|
||||
import jwt as pyjwt
|
||||
|
||||
from turnstone.core.auth import JWT_AUD_SERVER, ServiceTokenManager
|
||||
|
||||
mgr = ServiceTokenManager(
|
||||
user_id="svc",
|
||||
scopes=frozenset({"read"}),
|
||||
source="test",
|
||||
secret=self.SECRET,
|
||||
audience=JWT_AUD_SERVER,
|
||||
)
|
||||
payload = pyjwt.decode(
|
||||
mgr.token, self.SECRET, algorithms=["HS256"], audience=JWT_AUD_SERVER
|
||||
)
|
||||
assert payload["aud"] == JWT_AUD_SERVER
|
||||
|
||||
|
||||
class TestIsSecureRequest:
|
||||
def test_https_scheme(self):
|
||||
from turnstone.core.auth import is_secure_request
|
||||
|
||||
assert is_secure_request({}, scheme="https") is True
|
||||
|
||||
def test_http_scheme(self):
|
||||
from turnstone.core.auth import is_secure_request
|
||||
|
||||
assert is_secure_request({}, scheme="http") is False
|
||||
|
||||
def test_x_forwarded_proto_https(self):
|
||||
from turnstone.core.auth import is_secure_request
|
||||
|
||||
assert is_secure_request({"x-forwarded-proto": "https"}, scheme="http") is True
|
||||
|
||||
def test_x_forwarded_proto_http(self):
|
||||
from turnstone.core.auth import is_secure_request
|
||||
|
||||
assert is_secure_request({"x-forwarded-proto": "http"}, scheme="http") is False
|
||||
|
||||
|
||||
class TestSecretStrength:
|
||||
def test_short_secret_warns(self, caplog):
|
||||
import logging
|
||||
|
||||
from turnstone.core.auth import _MIN_SECRET_LENGTH
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="turnstone.core.auth"):
|
||||
import turnstone.core.auth as auth_mod
|
||||
|
||||
old = os.environ.get("TURNSTONE_JWT_SECRET", "")
|
||||
os.environ["TURNSTONE_JWT_SECRET"] = "short"
|
||||
try:
|
||||
secret = auth_mod.load_jwt_secret()
|
||||
assert secret == "short"
|
||||
assert any(str(_MIN_SECRET_LENGTH) in r.message for r in caplog.records)
|
||||
finally:
|
||||
if old:
|
||||
os.environ["TURNSTONE_JWT_SECRET"] = old
|
||||
else:
|
||||
os.environ.pop("TURNSTONE_JWT_SECRET", None)
|
||||
|
||||
|
||||
class TestCorsConfigurable:
|
||||
"""Verify CORS middleware is only added when origins are configured."""
|
||||
|
||||
def test_no_cors_origins_no_cors_headers(self):
|
||||
"""Without cors_origins, no Access-Control headers."""
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
import turnstone.server as srv_mod
|
||||
|
||||
app = srv_mod.create_app(
|
||||
workstreams=MagicMock(),
|
||||
global_queue=queue.Queue(),
|
||||
global_listeners=[],
|
||||
global_listeners_lock=threading.Lock(),
|
||||
skip_permissions=False,
|
||||
auth_config=AuthConfig(enabled=False),
|
||||
)
|
||||
client = TestClient(app)
|
||||
resp = client.get("/health", headers={"Origin": "http://evil.com"})
|
||||
assert "Access-Control-Allow-Origin" not in resp.headers
|
||||
client.close()
|
||||
|
||||
def test_cors_origins_set(self):
|
||||
"""With cors_origins, CORS headers are present."""
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
import turnstone.server as srv_mod
|
||||
|
||||
app = srv_mod.create_app(
|
||||
workstreams=MagicMock(),
|
||||
global_queue=queue.Queue(),
|
||||
global_listeners=[],
|
||||
global_listeners_lock=threading.Lock(),
|
||||
skip_permissions=False,
|
||||
auth_config=AuthConfig(enabled=False),
|
||||
cors_origins=["http://example.com"],
|
||||
)
|
||||
client = TestClient(app)
|
||||
resp = client.get(
|
||||
"/health",
|
||||
headers={"Origin": "http://example.com"},
|
||||
)
|
||||
assert resp.headers.get("Access-Control-Allow-Origin") == "http://example.com"
|
||||
client.close()
|
||||
|
||||
@@ -0,0 +1,357 @@
|
||||
"""Tests for user identity, API tokens, JWT, and scoped auth."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.auth import (
|
||||
AuthConfig,
|
||||
AuthResult,
|
||||
_authenticate_token,
|
||||
check_request,
|
||||
create_jwt,
|
||||
generate_token,
|
||||
hash_password,
|
||||
hash_token,
|
||||
parse_scopes,
|
||||
required_scope,
|
||||
token_prefix,
|
||||
validate_jwt,
|
||||
verify_password,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AuthResult
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAuthResult:
|
||||
def test_frozen(self):
|
||||
r = AuthResult(user_id="u1", scopes=frozenset({"read"}), token_source="config")
|
||||
with pytest.raises(AttributeError):
|
||||
r.user_id = "u2" # type: ignore[misc]
|
||||
|
||||
def test_has_scope(self):
|
||||
r = AuthResult(user_id="", scopes=frozenset({"read", "write"}), token_source="config")
|
||||
assert r.has_scope("read")
|
||||
assert r.has_scope("write")
|
||||
assert not r.has_scope("approve")
|
||||
|
||||
def test_empty_scopes(self):
|
||||
r = AuthResult(user_id="", scopes=frozenset(), token_source="config")
|
||||
assert not r.has_scope("read")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Token generation and hashing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestTokenHelpers:
|
||||
def test_generate_token_format(self):
|
||||
tok = generate_token()
|
||||
assert tok.startswith("ts_")
|
||||
assert len(tok) == 3 + 64 # ts_ + 64 hex chars
|
||||
|
||||
def test_generate_token_unique(self):
|
||||
tokens = {generate_token() for _ in range(10)}
|
||||
assert len(tokens) == 10
|
||||
|
||||
def test_hash_token_deterministic(self):
|
||||
assert hash_token("ts_abc") == hash_token("ts_abc")
|
||||
|
||||
def test_hash_token_hex(self):
|
||||
h = hash_token("test")
|
||||
assert len(h) == 64 # SHA-256 hex
|
||||
int(h, 16) # valid hex
|
||||
|
||||
def test_token_prefix(self):
|
||||
assert token_prefix("ts_abcdefgh1234") == "ts_abcde"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Password hashing (bcrypt)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPasswordHashing:
|
||||
def test_hash_and_verify(self):
|
||||
pw = "hunter2"
|
||||
hashed = hash_password(pw)
|
||||
assert verify_password(pw, hashed)
|
||||
|
||||
def test_wrong_password(self):
|
||||
hashed = hash_password("correct")
|
||||
assert not verify_password("wrong", hashed)
|
||||
|
||||
def test_hash_is_different_each_time(self):
|
||||
h1 = hash_password("same")
|
||||
h2 = hash_password("same")
|
||||
assert h1 != h2 # different salts
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Scope parsing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestParseScopes:
|
||||
def test_single_scope(self):
|
||||
assert parse_scopes("read") == frozenset({"read"})
|
||||
|
||||
def test_hierarchy_write(self):
|
||||
assert parse_scopes("write") == frozenset({"read", "write"})
|
||||
|
||||
def test_hierarchy_approve(self):
|
||||
assert parse_scopes("approve") == frozenset({"read", "write", "approve"})
|
||||
|
||||
def test_comma_separated(self):
|
||||
assert parse_scopes("read,write") == frozenset({"read", "write"})
|
||||
|
||||
def test_redundant_scopes(self):
|
||||
# approve already includes read,write
|
||||
assert parse_scopes("read,approve") == frozenset({"read", "write", "approve"})
|
||||
|
||||
def test_empty_string(self):
|
||||
assert parse_scopes("") == frozenset()
|
||||
|
||||
def test_invalid_scope_filtered(self):
|
||||
assert parse_scopes("bogus") == frozenset()
|
||||
|
||||
def test_mixed_valid_invalid(self):
|
||||
assert parse_scopes("read,bogus,approve") == frozenset({"read", "write", "approve"})
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# JWT create / validate
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestJWT:
|
||||
SECRET = "test-secret-key-for-jwt"
|
||||
|
||||
def test_round_trip(self):
|
||||
scopes = frozenset({"read", "write"})
|
||||
token = create_jwt("user123", scopes, "database", self.SECRET, expiry_hours=1)
|
||||
result = validate_jwt(token, self.SECRET)
|
||||
assert result is not None
|
||||
assert result.user_id == "user123"
|
||||
assert result.scopes == frozenset({"read", "write"})
|
||||
|
||||
def test_expired_token(self):
|
||||
import jwt
|
||||
|
||||
payload = {
|
||||
"sub": "user1",
|
||||
"scopes": "read",
|
||||
"src": "database",
|
||||
"iat": int(time.time()) - 7200,
|
||||
"exp": int(time.time()) - 3600,
|
||||
}
|
||||
token = jwt.encode(payload, self.SECRET, algorithm="HS256")
|
||||
assert validate_jwt(token, self.SECRET) is None
|
||||
|
||||
def test_invalid_signature(self):
|
||||
token = create_jwt("user1", frozenset({"read"}), "db", self.SECRET)
|
||||
assert validate_jwt(token, "wrong-secret") is None
|
||||
|
||||
def test_malformed_token(self):
|
||||
assert validate_jwt("not.a.jwt", self.SECRET) is None
|
||||
|
||||
def test_contains_dots(self):
|
||||
"""JWTs contain dots, used for detection."""
|
||||
token = create_jwt("u1", frozenset({"read"}), "db", self.SECRET)
|
||||
assert "." in token
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# required_scope
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRequiredScope:
|
||||
def test_get_read(self):
|
||||
assert required_scope("GET", "/api/workstreams") == "read"
|
||||
|
||||
def test_post_write(self):
|
||||
assert required_scope("POST", "/api/send") == "write"
|
||||
|
||||
def test_post_approve(self):
|
||||
assert required_scope("POST", "/api/approve") == "approve"
|
||||
|
||||
def test_admin_prefix(self):
|
||||
assert required_scope("GET", "/api/admin/users") == "approve"
|
||||
assert required_scope("POST", "/api/admin/users") == "approve"
|
||||
assert required_scope("DELETE", "/api/admin/users/abc") == "approve"
|
||||
|
||||
def test_versioned_path(self):
|
||||
assert required_scope("POST", "/v1/api/send") == "write"
|
||||
assert required_scope("POST", "/v1/api/approve") == "approve"
|
||||
|
||||
def test_proxy_write(self):
|
||||
assert required_scope("POST", "/node/n1/api/send") == "write"
|
||||
|
||||
def test_proxy_approve(self):
|
||||
assert required_scope("POST", "/node/n1/api/approve") == "approve"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _authenticate_token
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAuthenticateToken:
|
||||
def test_config_token_read(self):
|
||||
cfg = AuthConfig(enabled=True, tokens={"tok_read": "read"})
|
||||
result = _authenticate_token("tok_read", cfg)
|
||||
assert result is not None
|
||||
assert result.scopes == frozenset({"read"})
|
||||
assert result.token_source == "config"
|
||||
|
||||
def test_config_token_full(self):
|
||||
cfg = AuthConfig(enabled=True, tokens={"tok_full": "full"})
|
||||
result = _authenticate_token("tok_full", cfg)
|
||||
assert result is not None
|
||||
assert result.scopes == frozenset({"read", "write", "approve"})
|
||||
|
||||
def test_jwt_token(self):
|
||||
secret = "test-secret"
|
||||
jwt_tok = create_jwt("user1", frozenset({"read", "write"}), "db", secret)
|
||||
cfg = AuthConfig(enabled=True)
|
||||
result = _authenticate_token(jwt_tok, cfg, jwt_secret=secret)
|
||||
assert result is not None
|
||||
assert result.user_id == "user1"
|
||||
assert result.token_source == "db"
|
||||
|
||||
def test_api_token_with_storage(self):
|
||||
"""API tokens are looked up by hash in storage."""
|
||||
raw = generate_token()
|
||||
|
||||
class MockStorage:
|
||||
def get_api_token_by_hash(self, token_hash):
|
||||
expected = hash_token(raw)
|
||||
if token_hash == expected:
|
||||
return {
|
||||
"token_id": "tid",
|
||||
"token_prefix": "ts_abcde",
|
||||
"user_id": "user1",
|
||||
"name": "test",
|
||||
"scopes": "read,write",
|
||||
"created": "2026-01-01T00:00:00",
|
||||
}
|
||||
return None
|
||||
|
||||
cfg = AuthConfig(enabled=True)
|
||||
result = _authenticate_token(raw, cfg, storage=MockStorage())
|
||||
assert result is not None
|
||||
assert result.user_id == "user1"
|
||||
assert result.has_scope("write")
|
||||
assert result.token_source == "database"
|
||||
|
||||
def test_api_token_expired(self):
|
||||
"""Expired API tokens are rejected."""
|
||||
raw = generate_token()
|
||||
|
||||
class MockStorage:
|
||||
def get_api_token_by_hash(self, token_hash):
|
||||
return {
|
||||
"token_id": "tid",
|
||||
"token_prefix": "ts_abcde",
|
||||
"user_id": "user1",
|
||||
"name": "test",
|
||||
"scopes": "read",
|
||||
"created": "2020-01-01T00:00:00",
|
||||
"expires": "2020-01-02T00:00:00",
|
||||
}
|
||||
|
||||
cfg = AuthConfig(enabled=True)
|
||||
result = _authenticate_token(raw, cfg, storage=MockStorage())
|
||||
assert result is None
|
||||
|
||||
def test_unknown_token(self):
|
||||
cfg = AuthConfig(enabled=True, tokens={"tok": "full"})
|
||||
result = _authenticate_token("unknown", cfg)
|
||||
assert result is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# check_request with scopes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCheckRequestScopes:
|
||||
def test_config_read_on_write_403(self):
|
||||
cfg = AuthConfig(enabled=True, tokens={"tok_read": "read"})
|
||||
allowed, status, msg, _ = check_request(cfg, "POST", "/api/send", "Bearer tok_read")
|
||||
assert not allowed
|
||||
assert status == 403
|
||||
assert "write" in msg
|
||||
|
||||
def test_config_read_on_approve_403(self):
|
||||
cfg = AuthConfig(enabled=True, tokens={"tok_read": "read"})
|
||||
allowed, status, msg, _ = check_request(cfg, "POST", "/api/approve", "Bearer tok_read")
|
||||
assert not allowed
|
||||
assert status == 403
|
||||
assert "approve" in msg
|
||||
|
||||
def test_config_full_on_approve_ok(self):
|
||||
cfg = AuthConfig(enabled=True, tokens={"tok_full": "full"})
|
||||
allowed, status, msg, result = check_request(cfg, "POST", "/api/approve", "Bearer tok_full")
|
||||
assert allowed
|
||||
assert result is not None
|
||||
assert result.has_scope("approve")
|
||||
|
||||
def test_jwt_with_scopes(self):
|
||||
secret = "test"
|
||||
jwt_tok = create_jwt("u1", frozenset({"read", "write"}), "db", secret)
|
||||
cfg = AuthConfig(enabled=True)
|
||||
allowed, status, msg, result = check_request(
|
||||
cfg,
|
||||
"POST",
|
||||
"/api/send",
|
||||
f"Bearer {jwt_tok}",
|
||||
jwt_secret=secret,
|
||||
)
|
||||
assert allowed
|
||||
assert result is not None
|
||||
assert result.user_id == "u1"
|
||||
|
||||
def test_jwt_insufficient_scope(self):
|
||||
secret = "test"
|
||||
jwt_tok = create_jwt("u1", frozenset({"read"}), "db", secret)
|
||||
cfg = AuthConfig(enabled=True)
|
||||
allowed, status, msg, _ = check_request(
|
||||
cfg,
|
||||
"POST",
|
||||
"/api/send",
|
||||
f"Bearer {jwt_tok}",
|
||||
jwt_secret=secret,
|
||||
)
|
||||
assert not allowed
|
||||
assert status == 403
|
||||
|
||||
def test_admin_path_requires_approve(self):
|
||||
cfg = AuthConfig(enabled=True, tokens={"tok_read": "read"})
|
||||
allowed, status, msg, _ = check_request(
|
||||
cfg,
|
||||
"GET",
|
||||
"/v1/api/admin/users",
|
||||
"Bearer tok_read",
|
||||
)
|
||||
assert not allowed
|
||||
assert status == 403
|
||||
|
||||
def test_backward_compat_role_full(self):
|
||||
"""Config tokens with role='full' get all scopes."""
|
||||
cfg = AuthConfig(enabled=True, tokens={"tok_full": "full"})
|
||||
allowed, _, _, result = check_request(
|
||||
cfg,
|
||||
"GET",
|
||||
"/v1/api/admin/users",
|
||||
"Bearer tok_full",
|
||||
)
|
||||
assert allowed
|
||||
assert result is not None
|
||||
assert result.has_scope("approve")
|
||||
@@ -0,0 +1,328 @@
|
||||
"""Tests for the Discord channel adapter (bot, cog, views, config, CLI)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
discord = pytest.importorskip("discord")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _run(coro):
|
||||
"""Run an async coroutine in a fresh event loop (no pytest-asyncio needed)."""
|
||||
return asyncio.run(coro)
|
||||
|
||||
|
||||
def _make_message(*, bot=False, guild=True, content="hello", channel=None):
|
||||
"""Build a mock ``discord.Message``."""
|
||||
msg = MagicMock(spec=discord.Message)
|
||||
msg.author = MagicMock()
|
||||
msg.author.bot = bot
|
||||
msg.author.id = 12345
|
||||
msg.content = content
|
||||
msg.guild = MagicMock() if guild else None
|
||||
msg.channel = channel or MagicMock()
|
||||
msg.mentions = []
|
||||
return msg
|
||||
|
||||
|
||||
def _make_interaction(*, footer_text=None, has_embeds=True):
|
||||
"""Build a mock ``discord.Interaction``."""
|
||||
interaction = MagicMock(spec=discord.Interaction)
|
||||
interaction.user = MagicMock()
|
||||
interaction.user.id = 67890
|
||||
interaction.response = MagicMock()
|
||||
interaction.response.send_message = AsyncMock()
|
||||
|
||||
if has_embeds and footer_text is not None:
|
||||
embed = MagicMock()
|
||||
embed.footer.text = footer_text
|
||||
interaction.message = MagicMock()
|
||||
interaction.message.embeds = [embed]
|
||||
elif not has_embeds:
|
||||
interaction.message = MagicMock()
|
||||
interaction.message.embeds = []
|
||||
else:
|
||||
interaction.message = None
|
||||
|
||||
return interaction
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DiscordConfig
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDiscordConfig:
|
||||
"""Tests for DiscordConfig default and custom values."""
|
||||
|
||||
def test_defaults(self):
|
||||
from turnstone.channels.discord.config import DiscordConfig
|
||||
|
||||
cfg = DiscordConfig()
|
||||
assert cfg.bot_token == ""
|
||||
assert cfg.guild_id == 0
|
||||
assert cfg.allowed_channels == []
|
||||
assert cfg.thread_auto_archive == 1440
|
||||
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.model == ""
|
||||
assert cfg.auto_approve is False
|
||||
|
||||
def test_custom_values(self):
|
||||
from turnstone.channels.discord.config import DiscordConfig
|
||||
|
||||
cfg = DiscordConfig(
|
||||
bot_token="tok_123",
|
||||
guild_id=999,
|
||||
allowed_channels=[1, 2, 3],
|
||||
thread_auto_archive=60,
|
||||
max_message_length=4000,
|
||||
streaming_edit_interval=0.5,
|
||||
model="gpt-5",
|
||||
auto_approve=True,
|
||||
)
|
||||
assert cfg.bot_token == "tok_123"
|
||||
assert cfg.guild_id == 999
|
||||
assert cfg.allowed_channels == [1, 2, 3]
|
||||
assert cfg.thread_auto_archive == 60
|
||||
assert cfg.max_message_length == 4000
|
||||
assert cfg.streaming_edit_interval == 0.5
|
||||
assert cfg.model == "gpt-5"
|
||||
assert cfg.auto_approve is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# StreamingMessage
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestStreamingMessage:
|
||||
"""Tests for the StreamingMessage helper in bot.py."""
|
||||
|
||||
def test_append_accumulates(self):
|
||||
from turnstone.channels.discord.bot import StreamingMessage
|
||||
|
||||
channel = MagicMock()
|
||||
channel.send = AsyncMock()
|
||||
sm = StreamingMessage(channel=channel, edit_interval=999.0)
|
||||
|
||||
_run(sm.append("hello "))
|
||||
_run(sm.append("world"))
|
||||
|
||||
assert "".join(sm._buffer) == "hello world"
|
||||
|
||||
def test_finalize_sends_when_no_prior_message(self):
|
||||
from turnstone.channels.discord.bot import StreamingMessage
|
||||
|
||||
channel = MagicMock()
|
||||
channel.send = AsyncMock()
|
||||
sm = StreamingMessage(channel=channel, edit_interval=999.0)
|
||||
|
||||
_run(sm.append("hello"))
|
||||
_run(sm.finalize())
|
||||
|
||||
channel.send.assert_awaited_once_with("hello")
|
||||
|
||||
def test_finalize_edits_existing_message(self):
|
||||
from turnstone.channels.discord.bot import StreamingMessage
|
||||
|
||||
channel = MagicMock()
|
||||
sent_msg = MagicMock()
|
||||
sent_msg.edit = AsyncMock()
|
||||
channel.send = AsyncMock(return_value=sent_msg)
|
||||
sm = StreamingMessage(channel=channel, edit_interval=0.0)
|
||||
|
||||
# First append triggers flush (interval=0) which creates the message.
|
||||
_run(sm.append("hi"))
|
||||
assert sm._message is sent_msg
|
||||
|
||||
_run(sm.append(" there"))
|
||||
_run(sm.finalize())
|
||||
|
||||
# finalize edits the existing message with full content.
|
||||
sent_msg.edit.assert_awaited_with(content="hi there")
|
||||
|
||||
def test_finalize_chunks_long_content(self):
|
||||
from turnstone.channels.discord.bot import StreamingMessage
|
||||
|
||||
channel = MagicMock()
|
||||
channel.send = AsyncMock()
|
||||
sm = StreamingMessage(channel=channel, max_length=10, edit_interval=999.0)
|
||||
|
||||
# Content longer than max_length should be chunked on finalize.
|
||||
_run(sm.append("a" * 25))
|
||||
_run(sm.finalize())
|
||||
|
||||
# Should have sent multiple chunks via channel.send.
|
||||
assert channel.send.await_count >= 2
|
||||
|
||||
def test_finalize_empty_is_noop(self):
|
||||
from turnstone.channels.discord.bot import StreamingMessage
|
||||
|
||||
channel = MagicMock()
|
||||
channel.send = AsyncMock()
|
||||
sm = StreamingMessage(channel=channel)
|
||||
|
||||
_run(sm.finalize())
|
||||
channel.send.assert_not_awaited()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# MessageCog._on_message
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestMessageCog:
|
||||
"""Tests for the MessageCog on_message filtering logic."""
|
||||
|
||||
def _make_cog(self):
|
||||
"""Build a MessageCog with a fully mocked bot and TurnstoneBot."""
|
||||
from turnstone.channels.discord.cog import MessageCog
|
||||
|
||||
bot = MagicMock()
|
||||
bot.user = MagicMock()
|
||||
bot.user.id = 99999
|
||||
bot.user.mentioned_in = MagicMock(return_value=False)
|
||||
|
||||
ts = MagicMock()
|
||||
ts._is_allowed_channel = MagicMock(return_value=True)
|
||||
ts.storage = MagicMock()
|
||||
ts.router = MagicMock()
|
||||
ts.router.resolve_user = AsyncMock(return_value="u_abc")
|
||||
ts.router.send_message = AsyncMock()
|
||||
ts.config = MagicMock()
|
||||
ts._ws_tasks = {}
|
||||
bot.turnstone = ts
|
||||
|
||||
cog = MessageCog(bot)
|
||||
return cog, ts, bot
|
||||
|
||||
def test_ignores_bot_messages(self):
|
||||
cog, ts, _bot = self._make_cog()
|
||||
msg = _make_message(bot=True)
|
||||
|
||||
_run(cog._on_message(msg))
|
||||
|
||||
# No router interaction means the message was ignored.
|
||||
ts.router.send_message.assert_not_awaited()
|
||||
|
||||
def test_ignores_own_messages(self):
|
||||
cog, ts, bot = self._make_cog()
|
||||
msg = _make_message(bot=False)
|
||||
msg.author = bot.user # message from ourselves
|
||||
|
||||
_run(cog._on_message(msg))
|
||||
|
||||
ts.router.send_message.assert_not_awaited()
|
||||
|
||||
def test_ignores_dms(self):
|
||||
cog, ts, _bot = self._make_cog()
|
||||
msg = _make_message(guild=False)
|
||||
|
||||
_run(cog._on_message(msg))
|
||||
|
||||
ts.router.send_message.assert_not_awaited()
|
||||
|
||||
def test_ignores_non_allowed_channels(self):
|
||||
cog, ts, _bot = self._make_cog()
|
||||
ts._is_allowed_channel = MagicMock(return_value=False)
|
||||
|
||||
thread = MagicMock(spec=discord.Thread)
|
||||
thread.id = 111
|
||||
thread.parent_id = 222
|
||||
msg = _make_message(channel=thread)
|
||||
|
||||
_run(cog._on_message(msg))
|
||||
|
||||
ts.router.send_message.assert_not_awaited()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _parse_footer (views.py)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestParseFooter:
|
||||
"""Tests for _parse_footer in views.py."""
|
||||
|
||||
def test_valid_footer(self):
|
||||
from turnstone.channels.discord.views import _parse_footer
|
||||
|
||||
interaction = _make_interaction(footer_text="ws_abc|corr_123")
|
||||
result = _parse_footer(interaction)
|
||||
assert result == ("ws_abc", "corr_123")
|
||||
|
||||
def test_footer_with_pipe_in_correlation(self):
|
||||
from turnstone.channels.discord.views import _parse_footer
|
||||
|
||||
interaction = _make_interaction(footer_text="ws_abc|corr|extra")
|
||||
result = _parse_footer(interaction)
|
||||
# split("|", 1) means the second part includes everything after first pipe.
|
||||
assert result == ("ws_abc", "corr|extra")
|
||||
|
||||
def test_no_message_returns_none(self):
|
||||
from turnstone.channels.discord.views import _parse_footer
|
||||
|
||||
interaction = MagicMock()
|
||||
interaction.message = None
|
||||
assert _parse_footer(interaction) is None
|
||||
|
||||
def test_no_embeds_returns_none(self):
|
||||
from turnstone.channels.discord.views import _parse_footer
|
||||
|
||||
interaction = _make_interaction(has_embeds=False)
|
||||
assert _parse_footer(interaction) is None
|
||||
|
||||
def test_empty_footer_returns_none(self):
|
||||
from turnstone.channels.discord.views import _parse_footer
|
||||
|
||||
# Build an interaction whose embed has footer.text = None.
|
||||
interaction = MagicMock(spec=discord.Interaction)
|
||||
embed = MagicMock()
|
||||
embed.footer.text = None
|
||||
interaction.message = MagicMock()
|
||||
interaction.message.embeds = [embed]
|
||||
assert _parse_footer(interaction) is None
|
||||
|
||||
def test_footer_without_pipe_returns_none(self):
|
||||
from turnstone.channels.discord.views import _parse_footer
|
||||
|
||||
interaction = _make_interaction(footer_text="no_pipe_here")
|
||||
# footer text has no "|" separator
|
||||
embed = MagicMock()
|
||||
embed.footer.text = "no_pipe_here"
|
||||
interaction.message.embeds = [embed]
|
||||
assert _parse_footer(interaction) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# CLI main() — no adapter configured
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestChannelCLI:
|
||||
"""Tests for the channel CLI entry point."""
|
||||
|
||||
def test_exits_without_adapter_token(self):
|
||||
from turnstone.channels.cli import main
|
||||
|
||||
with (
|
||||
patch.object(sys, "argv", ["turnstone-channel"]),
|
||||
patch.dict("os.environ", {}, clear=True),
|
||||
pytest.raises(SystemExit) as exc_info,
|
||||
):
|
||||
main()
|
||||
|
||||
assert exc_info.value.code == 1
|
||||
@@ -0,0 +1,203 @@
|
||||
"""Tests for turnstone.channels._protocol and turnstone.channels._formatter."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from turnstone.channels._formatter import (
|
||||
chunk_message,
|
||||
format_approval_request,
|
||||
format_plan_review,
|
||||
truncate,
|
||||
)
|
||||
from turnstone.channels._protocol import ChannelEvent
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ChannelEvent
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestChannelEvent:
|
||||
def test_construction(self) -> None:
|
||||
evt = ChannelEvent(
|
||||
channel_type="discord",
|
||||
channel_id="ch-1",
|
||||
channel_user_id="u-42",
|
||||
message="hello",
|
||||
parent_channel_id="parent",
|
||||
metadata={"key": "val"},
|
||||
)
|
||||
assert evt.channel_type == "discord"
|
||||
assert evt.channel_id == "ch-1"
|
||||
assert evt.channel_user_id == "u-42"
|
||||
assert evt.message == "hello"
|
||||
assert evt.parent_channel_id == "parent"
|
||||
assert evt.metadata == {"key": "val"}
|
||||
|
||||
def test_defaults(self) -> None:
|
||||
evt = ChannelEvent(
|
||||
channel_type="slack",
|
||||
channel_id="ch-2",
|
||||
channel_user_id="u-7",
|
||||
message="hi",
|
||||
)
|
||||
assert evt.parent_channel_id == ""
|
||||
assert evt.metadata == {}
|
||||
|
||||
def test_metadata_independence(self) -> None:
|
||||
"""Default metadata dicts are independent across instances."""
|
||||
a = ChannelEvent(channel_type="x", channel_id="1", channel_user_id="u", message="m")
|
||||
b = ChannelEvent(channel_type="x", channel_id="2", channel_user_id="u", message="m")
|
||||
a.metadata["key"] = "val"
|
||||
assert "key" not in b.metadata
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# chunk_message
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestChunkMessage:
|
||||
def test_empty_string(self) -> None:
|
||||
assert chunk_message("") == [""]
|
||||
|
||||
def test_under_limit(self) -> None:
|
||||
assert chunk_message("short text", max_length=100) == ["short text"]
|
||||
|
||||
def test_exactly_at_limit(self) -> None:
|
||||
text = "a" * 50
|
||||
assert chunk_message(text, max_length=50) == [text]
|
||||
|
||||
def test_splits_at_newline(self) -> None:
|
||||
text = "line one\nline two\nline three"
|
||||
chunks = chunk_message(text, max_length=18)
|
||||
assert len(chunks) >= 2
|
||||
# The split should happen at a newline boundary within the text.
|
||||
# Reassembled chunks (with newline separators) should cover all content.
|
||||
rejoined = "\n".join(chunks)
|
||||
assert "line one" in rejoined
|
||||
assert "line three" in rejoined
|
||||
|
||||
def test_splits_at_word_boundary(self) -> None:
|
||||
text = "word1 word2 word3 word4"
|
||||
chunks = chunk_message(text, max_length=12)
|
||||
assert len(chunks) >= 2
|
||||
# No chunk should start with a space (lstrip handles newlines).
|
||||
for chunk in chunks:
|
||||
assert not chunk.startswith("\n")
|
||||
|
||||
def test_hard_splits(self) -> None:
|
||||
text = "a" * 30
|
||||
chunks = chunk_message(text, max_length=10)
|
||||
assert len(chunks) == 3
|
||||
assert "".join(chunks) == text
|
||||
|
||||
def test_code_block_spanning_boundary(self) -> None:
|
||||
text = "before\n```\ncode line 1\ncode line 2\ncode line 3\n```\nafter"
|
||||
chunks = chunk_message(text, max_length=30)
|
||||
assert len(chunks) >= 2
|
||||
# If a chunk opens a code block without closing it, the chunker
|
||||
# should close it and reopen in the next chunk.
|
||||
for chunk in chunks:
|
||||
fence_count = chunk.count("```")
|
||||
assert fence_count % 2 == 0, f"Unmatched code fence in chunk: {chunk!r}"
|
||||
|
||||
def test_multiple_code_blocks(self) -> None:
|
||||
text = "```\nblock1\n```\ntext\n```\nblock2\n```"
|
||||
chunks = chunk_message(text, max_length=20)
|
||||
for chunk in chunks:
|
||||
fence_count = chunk.count("```")
|
||||
assert fence_count % 2 == 0, f"Unmatched code fence in chunk: {chunk!r}"
|
||||
|
||||
def test_custom_max_length(self) -> None:
|
||||
text = "hello world"
|
||||
chunks = chunk_message(text, max_length=5)
|
||||
assert len(chunks) >= 2
|
||||
assert chunks[0] == "hello"
|
||||
|
||||
def test_very_long_single_line(self) -> None:
|
||||
text = "x" * 5000
|
||||
chunks = chunk_message(text, max_length=2000)
|
||||
assert len(chunks) == 3
|
||||
total = "".join(chunks)
|
||||
assert total == text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# format_approval_request
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFormatApprovalRequest:
|
||||
def test_single_tool(self) -> None:
|
||||
items = [{"function": {"name": "read_file", "arguments": "/etc/hosts"}}]
|
||||
result = format_approval_request(items)
|
||||
assert "Tool approval required" in result
|
||||
assert "`read_file`" in result
|
||||
|
||||
def test_multiple_tools(self) -> None:
|
||||
items = [
|
||||
{"function": {"name": "tool_a", "arguments": "arg1"}},
|
||||
{"function": {"name": "tool_b", "arguments": "arg2"}},
|
||||
]
|
||||
result = format_approval_request(items)
|
||||
assert "`tool_a`" in result
|
||||
assert "`tool_b`" in result
|
||||
|
||||
def test_long_arguments_truncated(self) -> None:
|
||||
long_args = "x" * 500
|
||||
items = [{"function": {"name": "fn", "arguments": long_args}}]
|
||||
result = format_approval_request(items)
|
||||
# The result should be shorter than the original args.
|
||||
assert len(result) < 500
|
||||
|
||||
def test_server_sse_format(self) -> None:
|
||||
"""Items from the server SSE use func_name/preview, not function.name."""
|
||||
items = [
|
||||
{
|
||||
"call_id": "c1",
|
||||
"func_name": "bash",
|
||||
"preview": "ls -la",
|
||||
"header": "Execute: ls -la",
|
||||
"needs_approval": True,
|
||||
}
|
||||
]
|
||||
result = format_approval_request(items)
|
||||
assert "`bash`" in result
|
||||
assert "Execute: ls -la" in result
|
||||
|
||||
def test_server_sse_format_no_header(self) -> None:
|
||||
items = [{"func_name": "read_file", "preview": "/etc/hosts"}]
|
||||
result = format_approval_request(items)
|
||||
assert "`read_file`" in result
|
||||
assert "/etc/hosts" in result
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# format_plan_review
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFormatPlanReview:
|
||||
def test_format(self) -> None:
|
||||
result = format_plan_review("Step 1: do stuff")
|
||||
assert result.startswith("**Plan review requested:**")
|
||||
assert "Step 1: do stuff" in result
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# truncate
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestTruncate:
|
||||
def test_short_text_unchanged(self) -> None:
|
||||
assert truncate("hello", max_length=200) == "hello"
|
||||
|
||||
def test_long_text_truncated(self) -> None:
|
||||
text = "a" * 300
|
||||
result = truncate(text, max_length=200)
|
||||
assert len(result) == 200
|
||||
assert result.endswith("\u2026")
|
||||
|
||||
def test_exactly_at_limit(self) -> None:
|
||||
text = "b" * 200
|
||||
assert truncate(text, max_length=200) == text
|
||||
@@ -0,0 +1,107 @@
|
||||
"""Tests for turnstone.channels._routing.ChannelRouter."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
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."""
|
||||
storage = MagicMock()
|
||||
storage.get_channel_user = MagicMock(return_value=None)
|
||||
storage.get_channel_route = MagicMock(return_value=None)
|
||||
storage.get_channel_route_by_ws = MagicMock(return_value=None)
|
||||
storage.create_channel_route = MagicMock()
|
||||
storage.delete_channel_route = MagicMock(return_value=True)
|
||||
return storage
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def router(mock_broker: AsyncMock, mock_storage: MagicMock) -> ChannelRouter:
|
||||
return ChannelRouter(broker=mock_broker, storage=mock_storage)
|
||||
|
||||
|
||||
class TestResolveUser:
|
||||
@pytest.mark.anyio
|
||||
async def test_linked_user(self, router: ChannelRouter, mock_storage: MagicMock) -> None:
|
||||
mock_storage.get_channel_user.return_value = {"user_id": "usr-1", "channel_user_id": "d-42"}
|
||||
result = await router.resolve_user("discord", "d-42")
|
||||
assert result == "usr-1"
|
||||
mock_storage.get_channel_user.assert_called_once_with("discord", "d-42")
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_unlinked_user(self, router: ChannelRouter, mock_storage: MagicMock) -> None:
|
||||
mock_storage.get_channel_user.return_value = None
|
||||
result = await router.resolve_user("slack", "s-99")
|
||||
assert result is None
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
class TestSendApproval:
|
||||
@pytest.mark.anyio
|
||||
async def test_pushes_to_response_queue(
|
||||
self, router: ChannelRouter, mock_broker: AsyncMock
|
||||
) -> None:
|
||||
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"
|
||||
|
||||
|
||||
class TestSendPlanFeedback:
|
||||
@pytest.mark.anyio
|
||||
async def test_pushes_to_response_queue(
|
||||
self, router: ChannelRouter, mock_broker: AsyncMock
|
||||
) -> None:
|
||||
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"
|
||||
|
||||
|
||||
class TestDeleteRoute:
|
||||
@pytest.mark.anyio
|
||||
async def test_calls_storage_delete(
|
||||
self, router: ChannelRouter, mock_storage: MagicMock
|
||||
) -> None:
|
||||
await router.delete_route("discord", "ch-123")
|
||||
mock_storage.delete_channel_route.assert_called_once_with("discord", "ch-123")
|
||||
@@ -0,0 +1,136 @@
|
||||
"""Tests for channel_users and channel_routes storage CRUD."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def db(tmp_path):
|
||||
"""Fresh SQLite backend for each test."""
|
||||
backend = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
return backend
|
||||
|
||||
|
||||
class TestChannelUserCRUD:
|
||||
"""Tests for channel_users table operations."""
|
||||
|
||||
def test_create_and_get(self, db):
|
||||
db.create_channel_user("discord", "12345", "u_abc")
|
||||
result = db.get_channel_user("discord", "12345")
|
||||
assert result is not None
|
||||
assert result["channel_type"] == "discord"
|
||||
assert result["channel_user_id"] == "12345"
|
||||
assert result["user_id"] == "u_abc"
|
||||
assert "created" in result
|
||||
|
||||
def test_get_nonexistent(self, db):
|
||||
assert db.get_channel_user("discord", "99999") is None
|
||||
|
||||
def test_create_duplicate_noop(self, db):
|
||||
db.create_channel_user("discord", "12345", "u_abc")
|
||||
db.create_channel_user("discord", "12345", "u_different")
|
||||
result = db.get_channel_user("discord", "12345")
|
||||
assert result is not None
|
||||
assert result["user_id"] == "u_abc" # first write wins
|
||||
|
||||
def test_same_user_different_channels(self, db):
|
||||
db.create_channel_user("discord", "d_123", "u_abc")
|
||||
db.create_channel_user("slack", "s_456", "u_abc")
|
||||
d = db.get_channel_user("discord", "d_123")
|
||||
s = db.get_channel_user("slack", "s_456")
|
||||
assert d is not None and d["user_id"] == "u_abc"
|
||||
assert s is not None and s["user_id"] == "u_abc"
|
||||
|
||||
def test_list_by_user(self, db):
|
||||
db.create_channel_user("discord", "d_123", "u_abc")
|
||||
db.create_channel_user("slack", "s_456", "u_abc")
|
||||
db.create_channel_user("discord", "d_999", "u_other")
|
||||
results = db.list_channel_users_by_user("u_abc")
|
||||
assert len(results) == 2
|
||||
types = {r["channel_type"] for r in results}
|
||||
assert types == {"discord", "slack"}
|
||||
|
||||
def test_list_by_user_empty(self, db):
|
||||
assert db.list_channel_users_by_user("u_nobody") == []
|
||||
|
||||
def test_delete(self, db):
|
||||
db.create_channel_user("discord", "12345", "u_abc")
|
||||
assert db.delete_channel_user("discord", "12345") is True
|
||||
assert db.get_channel_user("discord", "12345") is None
|
||||
|
||||
def test_delete_nonexistent(self, db):
|
||||
assert db.delete_channel_user("discord", "99999") is False
|
||||
|
||||
def test_delete_user_cascades_channel_users(self, db):
|
||||
"""Deleting a turnstone user should cascade to channel_users."""
|
||||
db.create_user("u_abc", "admin", "Admin", "hash123")
|
||||
db.create_channel_user("discord", "12345", "u_abc")
|
||||
db.delete_user("u_abc")
|
||||
assert db.get_channel_user("discord", "12345") is None
|
||||
|
||||
|
||||
class TestChannelRouteCRUD:
|
||||
"""Tests for channel_routes table operations."""
|
||||
|
||||
def test_create_and_get(self, db):
|
||||
db.create_channel_route("discord", "thread_123", "ws_abc", "node_1")
|
||||
result = db.get_channel_route("discord", "thread_123")
|
||||
assert result is not None
|
||||
assert result["channel_type"] == "discord"
|
||||
assert result["channel_id"] == "thread_123"
|
||||
assert result["ws_id"] == "ws_abc"
|
||||
assert result["node_id"] == "node_1"
|
||||
assert "created" in result
|
||||
|
||||
def test_get_nonexistent(self, db):
|
||||
assert db.get_channel_route("discord", "thread_999") is None
|
||||
|
||||
def test_create_duplicate_noop(self, db):
|
||||
db.create_channel_route("discord", "thread_123", "ws_abc")
|
||||
db.create_channel_route("discord", "thread_123", "ws_different")
|
||||
result = db.get_channel_route("discord", "thread_123")
|
||||
assert result is not None
|
||||
assert result["ws_id"] == "ws_abc" # first write wins
|
||||
|
||||
def test_default_empty_node_id(self, db):
|
||||
db.create_channel_route("discord", "thread_123", "ws_abc")
|
||||
result = db.get_channel_route("discord", "thread_123")
|
||||
assert result is not None
|
||||
assert result["node_id"] == ""
|
||||
|
||||
def test_get_by_ws(self, db):
|
||||
db.create_channel_route("discord", "thread_123", "ws_abc", "node_1")
|
||||
result = db.get_channel_route_by_ws("ws_abc")
|
||||
assert result is not None
|
||||
assert result["channel_id"] == "thread_123"
|
||||
assert result["ws_id"] == "ws_abc"
|
||||
|
||||
def test_get_by_ws_nonexistent(self, db):
|
||||
assert db.get_channel_route_by_ws("ws_nobody") is None
|
||||
|
||||
def test_delete(self, db):
|
||||
db.create_channel_route("discord", "thread_123", "ws_abc")
|
||||
assert db.delete_channel_route("discord", "thread_123") is True
|
||||
assert db.get_channel_route("discord", "thread_123") is None
|
||||
|
||||
def test_delete_nonexistent(self, db):
|
||||
assert db.delete_channel_route("discord", "thread_999") is False
|
||||
|
||||
def test_multiple_channels_same_type(self, db):
|
||||
db.create_channel_route("discord", "thread_1", "ws_1")
|
||||
db.create_channel_route("discord", "thread_2", "ws_2")
|
||||
r1 = db.get_channel_route("discord", "thread_1")
|
||||
r2 = db.get_channel_route("discord", "thread_2")
|
||||
assert r1 is not None and r1["ws_id"] == "ws_1"
|
||||
assert r2 is not None and r2["ws_id"] == "ws_2"
|
||||
|
||||
def test_different_channel_types(self, db):
|
||||
db.create_channel_route("discord", "thread_1", "ws_1")
|
||||
db.create_channel_route("slack", "channel_1", "ws_2")
|
||||
d = db.get_channel_route("discord", "thread_1")
|
||||
s = db.get_channel_route("slack", "channel_1")
|
||||
assert d is not None and d["ws_id"] == "ws_1"
|
||||
assert s is not None and s["ws_id"] == "ws_2"
|
||||
@@ -1,5 +1,6 @@
|
||||
"""Tests for turnstone.console — collector and HTTP server."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import queue
|
||||
from unittest.mock import MagicMock
|
||||
@@ -201,6 +202,68 @@ class TestCollectorPolling:
|
||||
# Should not raise
|
||||
c._apply_poll("unknown", _dashboard_response(), {})
|
||||
|
||||
def test_apply_poll_emits_ws_created_for_new_workstream(self):
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
dashboard = _dashboard_response(
|
||||
workstreams=[{"id": "ws1", "name": "new-task", "state": "idle"}]
|
||||
)
|
||||
c._apply_poll("node-a", dashboard, {})
|
||||
|
||||
event = q.get_nowait()
|
||||
assert event["type"] == "ws_created"
|
||||
assert event["ws_id"] == "ws1"
|
||||
assert event["name"] == "new-task"
|
||||
assert event["node_id"] == "node-a"
|
||||
|
||||
def test_apply_poll_emits_ws_closed_for_removed_workstream(self):
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
node_id="node-a",
|
||||
server_url="http://a:8080",
|
||||
workstreams={"ws1": {"id": "ws1", "name": "old", "state": "idle"}},
|
||||
)
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
c._apply_poll("node-a", _dashboard_response(), {})
|
||||
|
||||
event = q.get_nowait()
|
||||
assert event["type"] == "ws_closed"
|
||||
assert event["ws_id"] == "ws1"
|
||||
|
||||
def test_apply_poll_no_events_when_unchanged(self):
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
node_id="node-a",
|
||||
server_url="http://a:8080",
|
||||
workstreams={"ws1": {"id": "ws1", "name": "same", "state": "idle"}},
|
||||
)
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
dashboard = _dashboard_response(
|
||||
workstreams=[{"id": "ws1", "name": "same", "state": "running"}]
|
||||
)
|
||||
c._apply_poll("node-a", dashboard, {})
|
||||
|
||||
assert q.empty()
|
||||
|
||||
def test_apply_poll_skips_empty_id_workstream(self):
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
dashboard = _dashboard_response(workstreams=[{"name": "no-id", "state": "idle"}])
|
||||
c._apply_poll("node-a", dashboard, {})
|
||||
|
||||
assert q.empty()
|
||||
assert len(c._nodes["node-a"].workstreams) == 0
|
||||
|
||||
|
||||
class TestCollectorEvents:
|
||||
"""Real-time event handling from cluster channel."""
|
||||
@@ -445,6 +508,44 @@ class TestCollectorQueries:
|
||||
def test_get_node_detail_not_found(self, populated_collector):
|
||||
assert populated_collector.get_node_detail("nonexistent") is None
|
||||
|
||||
def test_get_snapshot_empty(self):
|
||||
c = _make_collector()
|
||||
snap = c.get_snapshot()
|
||||
assert snap["nodes"] == []
|
||||
assert snap["overview"]["nodes"] == 0
|
||||
assert snap["overview"]["workstreams"] == 0
|
||||
assert snap["overview"]["states"]["running"] == 0
|
||||
assert "timestamp" in snap
|
||||
|
||||
def test_get_snapshot_with_nodes(self, populated_collector):
|
||||
snap = populated_collector.get_snapshot()
|
||||
assert len(snap["nodes"]) == 2
|
||||
assert snap["overview"]["nodes"] == 2
|
||||
assert snap["overview"]["workstreams"] == 3
|
||||
assert snap["overview"]["states"]["running"] == 1
|
||||
assert snap["overview"]["states"]["attention"] == 1
|
||||
assert snap["overview"]["states"]["idle"] == 1
|
||||
assert snap["overview"]["aggregate"]["total_tokens"] == 17000
|
||||
assert snap["timestamp"] > 0
|
||||
# Each node should embed its workstreams
|
||||
node_ids = {n["node_id"] for n in snap["nodes"]}
|
||||
assert node_ids == {"node-a", "node-b"}
|
||||
for n in snap["nodes"]:
|
||||
if n["node_id"] == "node-a":
|
||||
assert len(n["workstreams"]) == 2
|
||||
elif n["node_id"] == "node-b":
|
||||
assert len(n["workstreams"]) == 1
|
||||
|
||||
def test_get_snapshot_consistency(self, populated_collector):
|
||||
"""Snapshot overview should match get_overview()."""
|
||||
snap = populated_collector.get_snapshot()
|
||||
overview = populated_collector.get_overview()
|
||||
assert snap["overview"]["nodes"] == overview["nodes"]
|
||||
assert snap["overview"]["workstreams"] == overview["workstreams"]
|
||||
assert snap["overview"]["states"] == overview["states"]
|
||||
assert snap["overview"]["aggregate"] == overview["aggregate"]
|
||||
assert snap["overview"]["version_drift"] == overview["version_drift"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ClusterStateEvent protocol tests
|
||||
@@ -535,6 +636,31 @@ class TestConsoleHTTPEndpoints:
|
||||
"workstreams": [],
|
||||
"aggregate": {},
|
||||
}
|
||||
collector.get_snapshot.return_value = {
|
||||
"nodes": [
|
||||
{
|
||||
"node_id": "node-a",
|
||||
"server_url": "http://a:8080",
|
||||
"max_ws": 10,
|
||||
"reachable": True,
|
||||
"version": "0.5.0",
|
||||
"health": {},
|
||||
"aggregate": {"total_tokens": 50000, "total_tool_calls": 200},
|
||||
"workstreams": [
|
||||
{"id": "ws1", "name": "test", "state": "running", "node": "node-a"},
|
||||
],
|
||||
},
|
||||
],
|
||||
"overview": {
|
||||
"nodes": 3,
|
||||
"workstreams": 15,
|
||||
"states": {"running": 5, "thinking": 2, "attention": 1, "idle": 6, "error": 1},
|
||||
"aggregate": {"total_tokens": 50000, "total_tool_calls": 200},
|
||||
"version_drift": False,
|
||||
"versions": ["0.5.0"],
|
||||
},
|
||||
"timestamp": 1234567890.0,
|
||||
}
|
||||
return collector
|
||||
|
||||
@pytest.fixture()
|
||||
@@ -614,6 +740,16 @@ class TestConsoleHTTPEndpoints:
|
||||
assert status == 404
|
||||
assert "error" in data
|
||||
|
||||
def test_get_snapshot(self, client, mock_collector):
|
||||
status, data = self._get(client, "/v1/api/cluster/snapshot")
|
||||
assert status == 200
|
||||
assert len(data["nodes"]) == 1
|
||||
assert data["nodes"][0]["node_id"] == "node-a"
|
||||
assert data["overview"]["nodes"] == 3
|
||||
assert data["overview"]["workstreams"] == 15
|
||||
assert data["timestamp"] == 1234567890.0
|
||||
mock_collector.get_snapshot.assert_called_once()
|
||||
|
||||
def test_health_endpoint(self, client, mock_collector):
|
||||
status, data = self._get(client, "/health")
|
||||
assert status == 200
|
||||
@@ -1299,3 +1435,177 @@ class TestProxySharedStatic:
|
||||
resp = client.get("/node/unknown/shared/base.css")
|
||||
assert resp.status_code == 404
|
||||
client.close()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SSE proxy — raw byte passthrough
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSSEProxy:
|
||||
"""Verify _proxy_sse forwards raw bytes including ping comments."""
|
||||
|
||||
def test_proxy_sse_preserves_pings_and_events(self):
|
||||
"""SSE proxy should forward ping comments and events verbatim."""
|
||||
from turnstone.console.server import _proxy_sse
|
||||
|
||||
# Simulate an upstream SSE response with a ping comment and a real event
|
||||
sse_payload = b': ping - 2026-03-08T12:00:00Z\n\nevent: message\ndata: {"type": "test"}\n\n'
|
||||
|
||||
class FakeResponse:
|
||||
status_code = 200
|
||||
headers = {"content-type": "text/event-stream"}
|
||||
|
||||
async def aiter_bytes(self):
|
||||
yield sse_payload
|
||||
|
||||
async def aclose(self):
|
||||
pass
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
pass
|
||||
|
||||
class FakeClient:
|
||||
def stream(self, method, url, **kwargs):
|
||||
return FakeResponse()
|
||||
|
||||
class FakeRequest:
|
||||
class url: # noqa: N801
|
||||
query = "ws_id=test123"
|
||||
|
||||
class app: # noqa: N801
|
||||
class state: # noqa: N801
|
||||
proxy_sse_client = FakeClient()
|
||||
proxy_auth_token = ""
|
||||
|
||||
headers = {}
|
||||
|
||||
async def is_disconnected(self):
|
||||
return False
|
||||
|
||||
async def _run():
|
||||
response = await _proxy_sse(
|
||||
FakeRequest(), "http://fake:8080", "events", api_prefix="v1/api"
|
||||
)
|
||||
assert response.media_type == "text/event-stream"
|
||||
# Collect the streamed bytes
|
||||
chunks: list[bytes] = []
|
||||
async for chunk in response.body_iterator:
|
||||
chunks.append(chunk if isinstance(chunk, bytes) else chunk.encode())
|
||||
body = b"".join(chunks)
|
||||
# Ping comment must be preserved (not filtered)
|
||||
assert b": ping" in body
|
||||
# Real event must be preserved
|
||||
assert b"event: message" in body
|
||||
assert b'"type": "test"' in body
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
def test_proxy_sse_upstream_error_status(self):
|
||||
"""Non-200 upstream status should yield an error event."""
|
||||
|
||||
from turnstone.console.server import _proxy_sse
|
||||
|
||||
class FakeResponse:
|
||||
status_code = 502
|
||||
|
||||
async def aiter_bytes(self):
|
||||
return
|
||||
yield # make it an async generator
|
||||
|
||||
async def aclose(self):
|
||||
pass
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
pass
|
||||
|
||||
class FakeClient:
|
||||
def stream(self, method, url, **kwargs):
|
||||
return FakeResponse()
|
||||
|
||||
class FakeRequest:
|
||||
class url: # noqa: N801
|
||||
query = ""
|
||||
|
||||
class app: # noqa: N801
|
||||
class state: # noqa: N801
|
||||
proxy_sse_client = FakeClient()
|
||||
proxy_auth_token = ""
|
||||
|
||||
headers = {}
|
||||
|
||||
async def is_disconnected(self):
|
||||
return False
|
||||
|
||||
async def _run():
|
||||
response = await _proxy_sse(FakeRequest(), "http://fake:8080", "events")
|
||||
chunks: list[bytes] = []
|
||||
async for chunk in response.body_iterator:
|
||||
chunks.append(chunk if isinstance(chunk, bytes) else chunk.encode())
|
||||
body = b"".join(chunks)
|
||||
assert b"event: error" in body
|
||||
assert b"502" in body
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
def test_proxy_sse_disconnect_handling(self):
|
||||
"""Proxy should stop when browser disconnects."""
|
||||
|
||||
from turnstone.console.server import _proxy_sse
|
||||
|
||||
class FakeResponse:
|
||||
status_code = 200
|
||||
|
||||
async def aiter_bytes(self):
|
||||
yield b"data: chunk1\n\n"
|
||||
yield b"data: chunk2\n\n" # should not be reached
|
||||
yield b"data: chunk3\n\n"
|
||||
|
||||
async def aclose(self):
|
||||
pass
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
pass
|
||||
|
||||
class FakeClient:
|
||||
def stream(self, method, url, **kwargs):
|
||||
return FakeResponse()
|
||||
|
||||
call_count = 0
|
||||
|
||||
class FakeRequest:
|
||||
class url: # noqa: N801
|
||||
query = ""
|
||||
|
||||
class app: # noqa: N801
|
||||
class state: # noqa: N801
|
||||
proxy_sse_client = FakeClient()
|
||||
proxy_auth_token = ""
|
||||
|
||||
headers = {}
|
||||
|
||||
async def is_disconnected(self):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return call_count > 1 # disconnect after first chunk
|
||||
|
||||
async def _run():
|
||||
response = await _proxy_sse(FakeRequest(), "http://fake:8080", "events")
|
||||
chunks: list[bytes] = []
|
||||
async for chunk in response.body_iterator:
|
||||
chunks.append(chunk if isinstance(chunk, bytes) else chunk.encode())
|
||||
body = b"".join(chunks)
|
||||
assert b"chunk1" in body
|
||||
# Should have stopped before chunk3
|
||||
assert b"chunk3" not in body
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
"""Tests for turnstone.core.log — structured logging configuration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
|
||||
import structlog
|
||||
|
||||
from turnstone.core.log import (
|
||||
configure_logging,
|
||||
ctx_node_id,
|
||||
ctx_request_id,
|
||||
ctx_user_id,
|
||||
ctx_ws_id,
|
||||
get_logger,
|
||||
)
|
||||
|
||||
|
||||
class TestConfigureLogging:
|
||||
"""Test configure_logging() sets up handlers and formatters."""
|
||||
|
||||
def setup_method(self):
|
||||
# Reset structlog and stdlib between tests
|
||||
structlog.reset_defaults()
|
||||
root = logging.getLogger()
|
||||
root.handlers.clear()
|
||||
root.setLevel(logging.WARNING)
|
||||
# Reset context vars
|
||||
for var in (ctx_node_id, ctx_ws_id, ctx_user_id, ctx_request_id):
|
||||
var.set("")
|
||||
|
||||
def test_sets_root_handler(self):
|
||||
configure_logging(level="INFO", json_output=False, service="test")
|
||||
root = logging.getLogger()
|
||||
assert len(root.handlers) == 1
|
||||
assert root.level == logging.INFO
|
||||
|
||||
def test_level_debug(self):
|
||||
configure_logging(level="DEBUG", json_output=False)
|
||||
root = logging.getLogger()
|
||||
assert root.level == logging.DEBUG
|
||||
|
||||
def test_level_warning(self):
|
||||
configure_logging(level="WARNING", json_output=False)
|
||||
root = logging.getLogger()
|
||||
assert root.level == logging.WARNING
|
||||
|
||||
def test_json_output(self, capsys):
|
||||
configure_logging(level="INFO", json_output=True, service="test-svc")
|
||||
log = logging.getLogger("test.json_output")
|
||||
log.info("hello world")
|
||||
captured = capsys.readouterr()
|
||||
# JSON goes to stderr
|
||||
line = captured.err.strip()
|
||||
data = json.loads(line)
|
||||
assert data["event"] == "hello world"
|
||||
assert data["level"] == "info"
|
||||
assert data["service"] == "test-svc"
|
||||
assert "timestamp" in data
|
||||
|
||||
def test_console_output(self, capsys):
|
||||
configure_logging(level="INFO", json_output=False)
|
||||
log = logging.getLogger("test.console_output")
|
||||
log.info("console hello")
|
||||
captured = capsys.readouterr()
|
||||
assert "console hello" in captured.err
|
||||
|
||||
def test_quiet_third_party(self):
|
||||
configure_logging(level="DEBUG", json_output=False)
|
||||
for name in ("httpx", "httpcore", "openai", "anthropic", "uvicorn.access"):
|
||||
assert logging.getLogger(name).level == logging.WARNING
|
||||
|
||||
def test_replaces_existing_handlers(self):
|
||||
root = logging.getLogger()
|
||||
# Count existing handlers (pytest may add its own)
|
||||
before = len(root.handlers)
|
||||
root.addHandler(logging.StreamHandler())
|
||||
root.addHandler(logging.StreamHandler())
|
||||
assert len(root.handlers) == before + 2
|
||||
configure_logging(level="INFO", json_output=False)
|
||||
# configure_logging clears all and adds exactly 1
|
||||
assert len(root.handlers) == 1
|
||||
|
||||
def test_env_var_level_override(self, monkeypatch):
|
||||
monkeypatch.setenv("TURNSTONE_LOG_LEVEL", "ERROR")
|
||||
configure_logging(level="DEBUG", json_output=False)
|
||||
root = logging.getLogger()
|
||||
assert root.level == logging.ERROR
|
||||
|
||||
def test_env_var_format_json(self, monkeypatch, capsys):
|
||||
monkeypatch.setenv("TURNSTONE_LOG_FORMAT", "json")
|
||||
configure_logging(level="INFO", service="test")
|
||||
log = logging.getLogger("test.env_json")
|
||||
log.info("env json test")
|
||||
captured = capsys.readouterr()
|
||||
data = json.loads(captured.err.strip())
|
||||
assert data["event"] == "env json test"
|
||||
|
||||
def test_env_var_format_text(self, monkeypatch, capsys):
|
||||
monkeypatch.setenv("TURNSTONE_LOG_FORMAT", "text")
|
||||
configure_logging(level="INFO", json_output=True) # json_output overridden by env
|
||||
log = logging.getLogger("test.env_text")
|
||||
log.info("env text test")
|
||||
captured = capsys.readouterr()
|
||||
# Should NOT be JSON
|
||||
line = captured.err.strip()
|
||||
assert "env text test" in line
|
||||
# Verify it's not JSON
|
||||
try:
|
||||
json.loads(line)
|
||||
is_json = True
|
||||
except json.JSONDecodeError:
|
||||
is_json = False
|
||||
assert not is_json
|
||||
|
||||
|
||||
class TestContextInjection:
|
||||
"""Test that context variables appear in log output."""
|
||||
|
||||
def setup_method(self):
|
||||
structlog.reset_defaults()
|
||||
root = logging.getLogger()
|
||||
root.handlers.clear()
|
||||
root.setLevel(logging.WARNING)
|
||||
for var in (ctx_node_id, ctx_ws_id, ctx_user_id, ctx_request_id):
|
||||
var.set("")
|
||||
|
||||
def test_node_id_in_output(self, capsys):
|
||||
configure_logging(level="INFO", json_output=True)
|
||||
ctx_node_id.set("worker-01_a3f2")
|
||||
log = logging.getLogger("test.ctx")
|
||||
log.info("ctx test")
|
||||
data = json.loads(capsys.readouterr().err.strip())
|
||||
assert data["node_id"] == "worker-01_a3f2"
|
||||
|
||||
def test_ws_id_in_output(self, capsys):
|
||||
configure_logging(level="INFO", json_output=True)
|
||||
ctx_ws_id.set("abc123")
|
||||
log = logging.getLogger("test.ctx")
|
||||
log.info("ws test")
|
||||
data = json.loads(capsys.readouterr().err.strip())
|
||||
assert data["ws_id"] == "abc123"
|
||||
|
||||
def test_empty_context_omitted(self, capsys):
|
||||
configure_logging(level="INFO", json_output=True)
|
||||
# All context vars are empty string (default)
|
||||
log = logging.getLogger("test.ctx")
|
||||
log.info("empty ctx")
|
||||
data = json.loads(capsys.readouterr().err.strip())
|
||||
assert "node_id" not in data
|
||||
assert "ws_id" not in data
|
||||
assert "user_id" not in data
|
||||
assert "request_id" not in data
|
||||
|
||||
def test_multiple_context_vars(self, capsys):
|
||||
configure_logging(level="INFO", json_output=True)
|
||||
ctx_node_id.set("node-1")
|
||||
ctx_ws_id.set("ws-2")
|
||||
ctx_request_id.set("req-3")
|
||||
log = logging.getLogger("test.ctx")
|
||||
log.info("multi ctx")
|
||||
data = json.loads(capsys.readouterr().err.strip())
|
||||
assert data["node_id"] == "node-1"
|
||||
assert data["ws_id"] == "ws-2"
|
||||
assert data["request_id"] == "req-3"
|
||||
assert "user_id" not in data
|
||||
|
||||
|
||||
class TestGetLogger:
|
||||
"""Test get_logger() returns a usable bound logger."""
|
||||
|
||||
def setup_method(self):
|
||||
structlog.reset_defaults()
|
||||
root = logging.getLogger()
|
||||
root.handlers.clear()
|
||||
root.setLevel(logging.WARNING)
|
||||
|
||||
def test_get_logger_returns_bound_logger(self):
|
||||
configure_logging(level="INFO", json_output=False)
|
||||
log = get_logger("test.bound")
|
||||
assert log is not None
|
||||
|
||||
def test_get_logger_outputs(self, capsys):
|
||||
configure_logging(level="INFO", json_output=True)
|
||||
log = get_logger("test.bound")
|
||||
log.info("bound logger test", extra_key="extra_val")
|
||||
data = json.loads(capsys.readouterr().err.strip())
|
||||
assert data["event"] == "bound logger test"
|
||||
assert data["extra_key"] == "extra_val"
|
||||
|
||||
|
||||
class TestServiceField:
|
||||
"""Test that service name is injected when configured."""
|
||||
|
||||
def setup_method(self):
|
||||
structlog.reset_defaults()
|
||||
root = logging.getLogger()
|
||||
root.handlers.clear()
|
||||
root.setLevel(logging.WARNING)
|
||||
|
||||
def test_service_present(self, capsys):
|
||||
configure_logging(level="INFO", json_output=True, service="myservice")
|
||||
log = logging.getLogger("test.svc")
|
||||
log.info("svc test")
|
||||
data = json.loads(capsys.readouterr().err.strip())
|
||||
assert data["service"] == "myservice"
|
||||
|
||||
def test_no_service_when_empty(self, capsys):
|
||||
configure_logging(level="INFO", json_output=True)
|
||||
log = logging.getLogger("test.svc")
|
||||
log.info("no svc")
|
||||
data = json.loads(capsys.readouterr().err.strip())
|
||||
assert "service" not in data
|
||||
+299
-1
@@ -2,10 +2,11 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from contextlib import AsyncExitStack
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -411,3 +412,300 @@ class TestCreateMcpClient:
|
||||
|
||||
result = create_mcp_client()
|
||||
assert result is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tool refresh — _rebuild_tools, _refresh_server, listeners
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRebuildTools:
|
||||
def test_rebuild_from_per_server(self):
|
||||
mgr = MCPClientManager({})
|
||||
mgr._per_server_tools = {
|
||||
"github": [_fake_openai_tool("mcp__github__search")],
|
||||
"slack": [_fake_openai_tool("mcp__slack__send")],
|
||||
}
|
||||
mgr._rebuild_tools()
|
||||
assert len(mgr._tools) == 2
|
||||
names = {t["function"]["name"] for t in mgr._tools}
|
||||
assert names == {"mcp__github__search", "mcp__slack__send"}
|
||||
assert mgr._tool_map["mcp__github__search"] == ("github", "search")
|
||||
assert mgr._tool_map["mcp__slack__send"] == ("slack", "send")
|
||||
|
||||
def test_rebuild_copy_on_write(self):
|
||||
mgr = MCPClientManager({})
|
||||
mgr._per_server_tools = {"a": [_fake_openai_tool("mcp__a__x")]}
|
||||
mgr._rebuild_tools()
|
||||
old_tools = mgr._tools
|
||||
old_map = mgr._tool_map
|
||||
mgr._per_server_tools["b"] = [_fake_openai_tool("mcp__b__y")]
|
||||
mgr._rebuild_tools()
|
||||
assert mgr._tools is not old_tools
|
||||
assert mgr._tool_map is not old_map
|
||||
|
||||
def test_rebuild_empty(self):
|
||||
mgr = MCPClientManager({})
|
||||
mgr._per_server_tools = {}
|
||||
mgr._rebuild_tools()
|
||||
assert mgr._tools == []
|
||||
assert mgr._tool_map == {}
|
||||
|
||||
|
||||
class TestRefreshServer:
|
||||
def test_refresh_detects_added_tools(self):
|
||||
async def _run() -> None:
|
||||
mgr = MCPClientManager({})
|
||||
mock_session = MagicMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.tools = [
|
||||
_fake_mcp_tool("search"),
|
||||
_fake_mcp_tool("create"), # new tool
|
||||
]
|
||||
mock_session.list_tools = AsyncMock(return_value=mock_result)
|
||||
mgr._sessions["github"] = mock_session
|
||||
mgr._per_server_tools["github"] = [_fake_openai_tool("mcp__github__search")]
|
||||
mgr._rebuild_tools()
|
||||
|
||||
added, removed = await mgr._refresh_server("github")
|
||||
assert "mcp__github__create" in added
|
||||
assert removed == []
|
||||
assert len(mgr._tools) == 2
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
def test_refresh_detects_removed_tools(self):
|
||||
async def _run() -> None:
|
||||
mgr = MCPClientManager({})
|
||||
mock_session = MagicMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.tools = [] # all tools removed
|
||||
mock_session.list_tools = AsyncMock(return_value=mock_result)
|
||||
mgr._sessions["github"] = mock_session
|
||||
mgr._per_server_tools["github"] = [_fake_openai_tool("mcp__github__search")]
|
||||
mgr._rebuild_tools()
|
||||
|
||||
added, removed = await mgr._refresh_server("github")
|
||||
assert added == []
|
||||
assert "mcp__github__search" in removed
|
||||
assert mgr._tools == []
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
def test_refresh_no_changes(self):
|
||||
async def _run() -> None:
|
||||
mgr = MCPClientManager({})
|
||||
mock_session = MagicMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.tools = [_fake_mcp_tool("search")]
|
||||
mock_session.list_tools = AsyncMock(return_value=mock_result)
|
||||
mgr._sessions["github"] = mock_session
|
||||
mgr._per_server_tools["github"] = [_fake_openai_tool("mcp__github__search")]
|
||||
mgr._rebuild_tools()
|
||||
|
||||
added, removed = await mgr._refresh_server("github")
|
||||
assert added == []
|
||||
assert removed == []
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
def test_refresh_disconnected_raises(self):
|
||||
async def _run() -> None:
|
||||
mgr = MCPClientManager({})
|
||||
with pytest.raises(RuntimeError, match="not connected"):
|
||||
await mgr._refresh_server("ghost")
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
class TestListeners:
|
||||
def test_add_and_notify(self):
|
||||
mgr = MCPClientManager({})
|
||||
calls: list[int] = []
|
||||
mgr.add_listener(lambda: calls.append(1))
|
||||
mgr._per_server_tools = {"a": [_fake_openai_tool("mcp__a__x")]}
|
||||
mgr._rebuild_tools()
|
||||
assert len(calls) == 1
|
||||
|
||||
def test_remove_listener(self):
|
||||
mgr = MCPClientManager({})
|
||||
calls: list[int] = []
|
||||
cb = lambda: calls.append(1) # noqa: E731
|
||||
mgr.add_listener(cb)
|
||||
mgr.remove_listener(cb)
|
||||
mgr._rebuild_tools()
|
||||
assert calls == []
|
||||
|
||||
def test_remove_nonexistent_listener(self):
|
||||
mgr = MCPClientManager({})
|
||||
mgr.remove_listener(lambda: None) # should not raise
|
||||
|
||||
def test_listener_error_does_not_propagate(self):
|
||||
mgr = MCPClientManager({})
|
||||
mgr.add_listener(lambda: 1 / 0) # will raise ZeroDivisionError
|
||||
mgr._rebuild_tools() # should not raise
|
||||
|
||||
|
||||
class TestServerNames:
|
||||
def test_server_names_property(self):
|
||||
mgr = MCPClientManager({"github": {}, "slack": {}})
|
||||
assert sorted(mgr.server_names) == ["github", "slack"]
|
||||
|
||||
def test_server_names_empty(self):
|
||||
mgr = MCPClientManager({})
|
||||
assert mgr.server_names == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Session integration — tool refresh propagation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSessionRefresh:
|
||||
@pytest.fixture()
|
||||
def tmp_db(self, tmp_path):
|
||||
from turnstone.core.storage import init_storage, reset_storage
|
||||
|
||||
reset_storage()
|
||||
init_storage("sqlite", path=str(tmp_path / "test.db"), run_migrations=False)
|
||||
yield
|
||||
reset_storage()
|
||||
|
||||
def _make_session(self, mcp_client=None, **kwargs):
|
||||
from turnstone.core.session import ChatSession
|
||||
|
||||
defaults: dict[str, Any] = dict(
|
||||
client=MagicMock(),
|
||||
model="test-model",
|
||||
ui=MagicMock(),
|
||||
instructions=None,
|
||||
temperature=0.5,
|
||||
max_tokens=4096,
|
||||
tool_timeout=30,
|
||||
mcp_client=mcp_client,
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return ChatSession(**defaults)
|
||||
|
||||
def test_listener_registered_on_init(self, tmp_db):
|
||||
mock_mcp = MagicMock()
|
||||
mock_mcp.get_tools.return_value = []
|
||||
session = self._make_session(mcp_client=mock_mcp)
|
||||
mock_mcp.add_listener.assert_called_once()
|
||||
assert session._mcp_refresh_cb is not None
|
||||
|
||||
def test_no_listener_without_mcp(self, tmp_db):
|
||||
session = self._make_session(mcp_client=None)
|
||||
assert session._mcp_refresh_cb is None
|
||||
|
||||
def test_close_removes_listener(self, tmp_db):
|
||||
mock_mcp = MagicMock()
|
||||
mock_mcp.get_tools.return_value = []
|
||||
session = self._make_session(mcp_client=mock_mcp)
|
||||
session.close()
|
||||
mock_mcp.remove_listener.assert_called_once()
|
||||
assert session._mcp_refresh_cb is None
|
||||
|
||||
def test_close_idempotent(self, tmp_db):
|
||||
mock_mcp = MagicMock()
|
||||
mock_mcp.get_tools.return_value = []
|
||||
session = self._make_session(mcp_client=mock_mcp)
|
||||
session.close()
|
||||
session.close() # should not raise
|
||||
assert mock_mcp.remove_listener.call_count == 1
|
||||
|
||||
def test_on_mcp_tools_changed_rebuilds_tools(self, tmp_db):
|
||||
mock_mcp = MagicMock()
|
||||
mock_mcp.get_tools.return_value = [_fake_openai_tool("mcp__test__a")]
|
||||
session = self._make_session(mcp_client=mock_mcp)
|
||||
initial_count = len(session._tools)
|
||||
|
||||
# Simulate a tool refresh — MCP now has 2 tools
|
||||
mock_mcp.get_tools.return_value = [
|
||||
_fake_openai_tool("mcp__test__a"),
|
||||
_fake_openai_tool("mcp__test__b"),
|
||||
]
|
||||
session._on_mcp_tools_changed()
|
||||
assert len(session._tools) == initial_count + 1
|
||||
|
||||
def test_tool_search_preserved_across_refresh(self, tmp_db):
|
||||
# Create enough MCP tools to trigger tool search
|
||||
mcp_tools = [_fake_openai_tool(f"mcp__srv__tool{i}") for i in range(25)]
|
||||
mock_mcp = MagicMock()
|
||||
mock_mcp.get_tools.return_value = mcp_tools
|
||||
session = self._make_session(
|
||||
mcp_client=mock_mcp,
|
||||
tool_search="auto",
|
||||
tool_search_threshold=20,
|
||||
)
|
||||
assert session._tool_search is not None
|
||||
|
||||
# Expand a tool
|
||||
session._tool_search.expand_visible(["mcp__srv__tool0"])
|
||||
assert "mcp__srv__tool0" in session._tool_search.get_expanded_names()
|
||||
|
||||
# Refresh with same tools
|
||||
session._on_mcp_tools_changed()
|
||||
assert session._tool_search is not None
|
||||
assert "mcp__srv__tool0" in session._tool_search.get_expanded_names()
|
||||
|
||||
def test_tool_search_prunes_removed_from_expanded(self, tmp_db):
|
||||
mcp_tools = [_fake_openai_tool(f"mcp__srv__tool{i}") for i in range(25)]
|
||||
mock_mcp = MagicMock()
|
||||
mock_mcp.get_tools.return_value = mcp_tools
|
||||
session = self._make_session(
|
||||
mcp_client=mock_mcp,
|
||||
tool_search="auto",
|
||||
tool_search_threshold=20,
|
||||
)
|
||||
session._tool_search.expand_visible(["mcp__srv__tool0"])
|
||||
|
||||
# Refresh with tool0 removed
|
||||
new_tools = [_fake_openai_tool(f"mcp__srv__tool{i}") for i in range(1, 25)]
|
||||
mock_mcp.get_tools.return_value = new_tools
|
||||
session._on_mcp_tools_changed()
|
||||
# tool0 was removed, so it should no longer be in expanded
|
||||
expanded = session._tool_search.get_expanded_names()
|
||||
assert "mcp__srv__tool0" not in expanded
|
||||
|
||||
def test_mcp_refresh_command(self, tmp_db):
|
||||
mock_mcp = MagicMock()
|
||||
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
|
||||
mock_mcp.server_names = ["test"]
|
||||
mock_mcp.refresh_sync.return_value = {"test": (["mcp__test__new"], [])}
|
||||
session = self._make_session(mcp_client=mock_mcp)
|
||||
|
||||
session.handle_command("/mcp refresh")
|
||||
mock_mcp.refresh_sync.assert_called_once_with(None)
|
||||
session.ui.on_info.assert_called()
|
||||
|
||||
def test_mcp_refresh_specific_server(self, tmp_db):
|
||||
mock_mcp = MagicMock()
|
||||
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
|
||||
mock_mcp.server_names = ["github", "slack"]
|
||||
mock_mcp.refresh_sync.return_value = {"github": ([], [])}
|
||||
session = self._make_session(mcp_client=mock_mcp)
|
||||
|
||||
session.handle_command("/mcp refresh github")
|
||||
mock_mcp.refresh_sync.assert_called_once_with("github")
|
||||
|
||||
def test_mcp_refresh_unknown_server(self, tmp_db):
|
||||
mock_mcp = MagicMock()
|
||||
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
|
||||
mock_mcp.server_names = ["github"]
|
||||
session = self._make_session(mcp_client=mock_mcp)
|
||||
|
||||
session.handle_command("/mcp refresh nonexistent")
|
||||
session.ui.on_error.assert_called_once()
|
||||
assert "Unknown MCP server" in session.ui.on_error.call_args[0][0]
|
||||
|
||||
def test_mcp_refresh_error_handling(self, tmp_db):
|
||||
mock_mcp = MagicMock()
|
||||
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
|
||||
mock_mcp.server_names = ["test"]
|
||||
mock_mcp.refresh_sync.side_effect = TimeoutError("timed out")
|
||||
session = self._make_session(mcp_client=mock_mcp)
|
||||
|
||||
session.handle_command("/mcp refresh")
|
||||
session.ui.on_error.assert_called_once()
|
||||
assert "MCP refresh failed" in session.ui.on_error.call_args[0][0]
|
||||
|
||||
@@ -530,11 +530,11 @@ class TestWorkstreamModelParam:
|
||||
|
||||
captured_alias = None
|
||||
|
||||
def factory(ui: Any, model_alias: str | None = None) -> Any:
|
||||
def factory(ui: Any, model_alias: str | None = None, ws_id: str | None = None) -> Any:
|
||||
nonlocal captured_alias
|
||||
captured_alias = model_alias
|
||||
mock_session = MagicMock()
|
||||
mock_session.session_id = "test123"
|
||||
mock_session.ws_id = "test123"
|
||||
return mock_session
|
||||
|
||||
mgr = WorkstreamManager(factory)
|
||||
@@ -544,11 +544,11 @@ class TestWorkstreamModelParam:
|
||||
def test_create_without_model(self) -> None:
|
||||
captured_alias = None
|
||||
|
||||
def factory(ui: Any, model_alias: str | None = None) -> Any:
|
||||
def factory(ui: Any, model_alias: str | None = None, ws_id: str | None = None) -> Any:
|
||||
nonlocal captured_alias
|
||||
captured_alias = model_alias
|
||||
mock_session = MagicMock()
|
||||
mock_session.session_id = "test123"
|
||||
mock_session.ws_id = "test123"
|
||||
return mock_session
|
||||
|
||||
from turnstone.core.workstream import WorkstreamManager
|
||||
|
||||
@@ -0,0 +1,338 @@
|
||||
"""Tests for the channel gateway HTTP notify endpoint."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from turnstone.channels._http import create_channel_app
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path):
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_adapter():
|
||||
adapter = AsyncMock()
|
||||
adapter.channel_type = "discord"
|
||||
adapter.send = AsyncMock(return_value="msg_001")
|
||||
return adapter
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def no_auth_client(storage, mock_adapter):
|
||||
"""Client with no auth configured (for fail-closed tests)."""
|
||||
app = create_channel_app({"discord": mock_adapter}, storage)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(storage, mock_adapter):
|
||||
"""Default client with static auth token configured."""
|
||||
app = create_channel_app({"discord": mock_adapter}, storage, auth_token="test-secret-token")
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def authed_client(storage, mock_adapter):
|
||||
"""Alias — same as client, for auth-specific test clarity."""
|
||||
app = create_channel_app({"discord": mock_adapter}, storage, auth_token="test-secret-token")
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def jwt_client(storage, mock_adapter):
|
||||
"""Client with JWT auth configured."""
|
||||
app = create_channel_app({"discord": mock_adapter}, storage, jwt_secret="a" * 32)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
class TestNotifyEndpoint:
|
||||
def test_health(self, client):
|
||||
resp = client.get("/health")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["status"] == "ok"
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
return {"Authorization": "Bearer test-secret-token"}
|
||||
|
||||
def test_direct_discord_target(self, client, mock_adapter):
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123456"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers=self._headers(),
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
results = resp.json()["results"]
|
||||
assert len(results) == 1
|
||||
assert results[0]["status"] == "sent"
|
||||
assert results[0]["message_id"] == "msg_001"
|
||||
mock_adapter.send.assert_called_once_with("123456", "Hello!")
|
||||
|
||||
def test_with_title(self, client, mock_adapter):
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123456"},
|
||||
"message": "Hello!",
|
||||
"title": "Alert",
|
||||
},
|
||||
headers=self._headers(),
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
mock_adapter.send.assert_called_once_with("123456", "**Alert**\nHello!")
|
||||
|
||||
def test_username_resolution(self, client, storage, mock_adapter):
|
||||
# Create a user and link a channel
|
||||
storage.create_user("u1", "testuser", "Test User", "hash")
|
||||
storage.create_channel_user("discord", "disc_123", "u1")
|
||||
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"username": "testuser"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers=self._headers(),
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
results = resp.json()["results"]
|
||||
assert len(results) == 1
|
||||
assert results[0]["status"] == "sent"
|
||||
mock_adapter.send.assert_called_once_with("disc_123", "Hello!")
|
||||
|
||||
def test_unknown_username(self, client):
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"username": "nobody"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers=self._headers(),
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
error = resp.json()["error"]
|
||||
assert "nobody" not in error
|
||||
assert "not found or has no linked channels" in error
|
||||
|
||||
def test_user_no_channels(self, authed_client, storage):
|
||||
storage.create_user("u1", "testuser", "Test User", "hash")
|
||||
|
||||
resp = authed_client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"username": "testuser"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers={"Authorization": "Bearer test-secret-token"},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
# Generic message — must not differentiate "not found" vs "no channels"
|
||||
error = resp.json()["error"]
|
||||
assert "testuser" not in error
|
||||
assert "not found or has no linked channels" in error
|
||||
|
||||
def test_missing_fields(self, client):
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
json={"target": {"username": "x"}},
|
||||
headers=self._headers(),
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
|
||||
def test_missing_target(self, client):
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
json={"message": "Hello!"},
|
||||
headers=self._headers(),
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
|
||||
def test_invalid_target(self, client):
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"invalid": "field"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers=self._headers(),
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
|
||||
def test_no_adapter(self, client, storage):
|
||||
# App has discord adapter, try email target
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "email", "channel_id": "test@example.com"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers=self._headers(),
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
results = resp.json()["results"]
|
||||
assert results[0]["status"] == "no_adapter"
|
||||
|
||||
def test_adapter_failure(self, client, mock_adapter):
|
||||
mock_adapter.send.side_effect = RuntimeError("Discord API error")
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123456"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers=self._headers(),
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
results = resp.json()["results"]
|
||||
assert results[0]["status"] == "failed"
|
||||
|
||||
def test_invalid_json(self, client):
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
content=b"not json",
|
||||
headers={
|
||||
"content-type": "application/json",
|
||||
"Authorization": "Bearer test-secret-token",
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
|
||||
def test_whitespace_only_message(self, client):
|
||||
"""Whitespace-only messages should be rejected."""
|
||||
resp = client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123"},
|
||||
"message": " ",
|
||||
},
|
||||
headers=self._headers(),
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
|
||||
|
||||
class TestNotifyAuth:
|
||||
"""Tests for authentication on the /v1/api/notify endpoint."""
|
||||
|
||||
def test_reject_when_unconfigured(self, no_auth_client):
|
||||
"""Requests are rejected (fail closed) when no auth is configured."""
|
||||
resp = no_auth_client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_reject_without_token(self, authed_client):
|
||||
"""Requests without Authorization header are rejected when auth is configured."""
|
||||
resp = authed_client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_reject_wrong_token(self, authed_client):
|
||||
"""Requests with wrong token are rejected."""
|
||||
resp = authed_client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers={"Authorization": "Bearer wrong-token"},
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_accept_valid_static_token(self, authed_client, mock_adapter):
|
||||
"""Requests with correct static token are accepted."""
|
||||
resp = authed_client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers={"Authorization": "Bearer test-secret-token"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["results"][0]["status"] == "sent"
|
||||
|
||||
def test_accept_valid_jwt(self, jwt_client, mock_adapter):
|
||||
"""Requests with a valid JWT for the channel audience are accepted."""
|
||||
from turnstone.core.auth import JWT_AUD_CHANNEL, create_jwt
|
||||
|
||||
token = create_jwt(
|
||||
user_id="system",
|
||||
scopes=frozenset({"write"}),
|
||||
source="service",
|
||||
secret="a" * 32,
|
||||
audience=JWT_AUD_CHANNEL,
|
||||
)
|
||||
resp = jwt_client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_reject_jwt_wrong_audience(self, jwt_client):
|
||||
"""JWTs with wrong audience are rejected."""
|
||||
from turnstone.core.auth import create_jwt
|
||||
|
||||
token = create_jwt(
|
||||
user_id="system",
|
||||
scopes=frozenset({"write"}),
|
||||
source="service",
|
||||
secret="a" * 32,
|
||||
audience="turnstone-server", # wrong audience
|
||||
)
|
||||
resp = jwt_client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_reject_jwt_wrong_secret(self, jwt_client):
|
||||
"""JWTs signed with wrong secret are rejected."""
|
||||
from turnstone.core.auth import JWT_AUD_CHANNEL, create_jwt
|
||||
|
||||
token = create_jwt(
|
||||
user_id="system",
|
||||
scopes=frozenset({"write"}),
|
||||
source="service",
|
||||
secret="b" * 32, # wrong secret
|
||||
audience=JWT_AUD_CHANNEL,
|
||||
)
|
||||
resp = jwt_client.post(
|
||||
"/v1/api/notify",
|
||||
json={
|
||||
"target": {"channel_type": "discord", "channel_id": "123"},
|
||||
"message": "Hello!",
|
||||
},
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_health_bypasses_auth(self, authed_client):
|
||||
"""Health endpoint is always accessible regardless of auth config."""
|
||||
resp = authed_client.get("/health")
|
||||
assert resp.status_code == 200
|
||||
@@ -0,0 +1,618 @@
|
||||
"""Tests for the notify tool (prepare + execute) in ChatSession."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from turnstone.core.session import ChatSession
|
||||
|
||||
|
||||
def _make_session() -> ChatSession:
|
||||
"""Create a minimal ChatSession with mocked dependencies."""
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.memory.register_workstream"),
|
||||
patch("turnstone.core.session.save_message"),
|
||||
):
|
||||
from turnstone.core.session import ChatSession
|
||||
|
||||
ui = MagicMock()
|
||||
session = ChatSession(
|
||||
client=MagicMock(),
|
||||
model="test-model",
|
||||
ui=ui,
|
||||
instructions=None,
|
||||
temperature=0.7,
|
||||
max_tokens=1000,
|
||||
tool_timeout=30,
|
||||
)
|
||||
return session
|
||||
|
||||
|
||||
class TestPrepareNotify:
|
||||
def test_valid_username_target(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify(
|
||||
"call_1",
|
||||
{
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
},
|
||||
)
|
||||
assert "execute" in result
|
||||
assert result["func_name"] == "notify"
|
||||
assert result["needs_approval"] is False
|
||||
assert "@admin" in result["header"]
|
||||
assert result["username"] == "admin"
|
||||
assert result["message"] == "Hello!"
|
||||
|
||||
def test_valid_direct_target(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify(
|
||||
"call_1",
|
||||
{
|
||||
"message": "Hello!",
|
||||
"channel_type": "discord",
|
||||
"channel_id": "123456",
|
||||
},
|
||||
)
|
||||
assert "execute" in result
|
||||
assert result["channel_type"] == "discord"
|
||||
assert result["channel_id"] == "123456"
|
||||
assert "discord:123456" in result["header"]
|
||||
|
||||
def test_missing_message(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify("call_1", {"username": "admin"})
|
||||
assert "error" in result
|
||||
assert "message" in result["error"].lower()
|
||||
|
||||
def test_empty_message(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify(
|
||||
"call_1",
|
||||
{
|
||||
"message": "",
|
||||
"username": "admin",
|
||||
},
|
||||
)
|
||||
assert "error" in result
|
||||
|
||||
def test_message_too_long(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify(
|
||||
"call_1",
|
||||
{
|
||||
"message": "x" * 2001,
|
||||
"username": "admin",
|
||||
},
|
||||
)
|
||||
assert "error" in result
|
||||
assert "2000" in result["error"]
|
||||
|
||||
def test_both_username_and_direct(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify(
|
||||
"call_1",
|
||||
{
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "discord",
|
||||
"channel_id": "123",
|
||||
},
|
||||
)
|
||||
assert "error" in result
|
||||
assert "both" in result["error"].lower() or "ambiguous" in result["error"].lower()
|
||||
|
||||
def test_no_target(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify("call_1", {"message": "Hello!"})
|
||||
assert "error" in result
|
||||
|
||||
def test_channel_type_without_id(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify(
|
||||
"call_1",
|
||||
{
|
||||
"message": "Hello!",
|
||||
"channel_type": "discord",
|
||||
},
|
||||
)
|
||||
assert "error" in result
|
||||
assert "channel_id" in result["error"]
|
||||
|
||||
def test_channel_id_without_type(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify(
|
||||
"call_1",
|
||||
{
|
||||
"message": "Hello!",
|
||||
"channel_id": "123456",
|
||||
},
|
||||
)
|
||||
assert "error" in result
|
||||
assert "channel_type" in result["error"]
|
||||
|
||||
def test_preview_truncated(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify(
|
||||
"call_1",
|
||||
{
|
||||
"message": "a" * 200,
|
||||
"username": "admin",
|
||||
},
|
||||
)
|
||||
assert result["preview"].endswith("...")
|
||||
assert len(result["preview"]) <= 123 # 120 chars + "..."
|
||||
|
||||
def test_title_passed_through(self):
|
||||
session = _make_session()
|
||||
result = session._prepare_notify(
|
||||
"call_1",
|
||||
{
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"title": "Alert",
|
||||
},
|
||||
)
|
||||
assert result["title"] == "Alert"
|
||||
|
||||
|
||||
class TestExecNotify:
|
||||
def test_sends_http_to_channel_gateway(self, tmp_path):
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "Alert",
|
||||
}
|
||||
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.json.return_value = {
|
||||
"results": [{"channel_type": "discord", "channel_id": "123", "status": "sent"}]
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch("turnstone.core.session.httpx.post", return_value=mock_resp) as mock_post,
|
||||
patch.dict("os.environ", {}, clear=False),
|
||||
):
|
||||
call_id, msg = session._exec_notify(item)
|
||||
|
||||
assert call_id == "call_1"
|
||||
assert "sent successfully" in msg.lower()
|
||||
mock_post.assert_called_once()
|
||||
post_kwargs = mock_post.call_args
|
||||
assert post_kwargs.kwargs["json"]["target"] == {"username": "admin"}
|
||||
assert post_kwargs.kwargs["json"]["message"] == "Hello!"
|
||||
|
||||
def test_no_services_available(self, tmp_path):
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
# No services registered
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch("turnstone.core.session.time.sleep"),
|
||||
):
|
||||
call_id, msg = session._exec_notify(item)
|
||||
|
||||
assert "no channel gateway" in msg.lower()
|
||||
|
||||
def test_rate_limit(self, tmp_path):
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.json.return_value = {
|
||||
"results": [{"channel_type": "discord", "channel_id": "123", "status": "sent"}]
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch("turnstone.core.session.httpx.post", return_value=mock_resp),
|
||||
patch.dict("os.environ", {}, clear=False),
|
||||
):
|
||||
for _i in range(5):
|
||||
call_id, msg = session._exec_notify(item)
|
||||
assert "sent successfully" in msg.lower()
|
||||
|
||||
# 6th should fail
|
||||
call_id, msg = session._exec_notify(item)
|
||||
assert "rate limit" in msg.lower()
|
||||
|
||||
def test_rate_limit_not_consumed_on_failure(self, tmp_path):
|
||||
"""Failed delivery should not consume rate limit slots."""
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch(
|
||||
"turnstone.core.session.httpx.post",
|
||||
side_effect=ConnectionError("refused"),
|
||||
),
|
||||
patch.dict("os.environ", {}, clear=False),
|
||||
patch("turnstone.core.session.time.sleep"),
|
||||
):
|
||||
# All fail — counter should stay at 0
|
||||
for _i in range(3):
|
||||
session._exec_notify(item)
|
||||
assert session._notify_count == 0
|
||||
|
||||
def test_counter_on_init(self):
|
||||
session = _make_session()
|
||||
assert session._notify_count == 0
|
||||
|
||||
def test_http_failure_reported(self, tmp_path):
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "",
|
||||
"channel_type": "discord",
|
||||
"channel_id": "999",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch(
|
||||
"turnstone.core.session.httpx.post",
|
||||
side_effect=ConnectionError("refused"),
|
||||
),
|
||||
patch.dict("os.environ", {}, clear=False),
|
||||
patch("turnstone.core.session.time.sleep"),
|
||||
):
|
||||
call_id, msg = session._exec_notify(item)
|
||||
|
||||
# Error message should be generic (no internal details)
|
||||
assert "delivery failed" in msg.lower()
|
||||
assert "refused" not in msg
|
||||
assert "ch-1" not in msg
|
||||
|
||||
def test_first_healthy_only(self, tmp_path):
|
||||
"""Only the first healthy gateway should receive the request."""
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
storage.register_service("channel", "ch-2", "http://localhost:8092")
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.json.return_value = {
|
||||
"results": [{"channel_type": "discord", "channel_id": "123", "status": "sent"}]
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch("turnstone.core.session.httpx.post", return_value=mock_resp) as mock_post,
|
||||
patch.dict("os.environ", {}, clear=False),
|
||||
):
|
||||
session._exec_notify(item)
|
||||
|
||||
# Should only have been called once (first healthy)
|
||||
assert mock_post.call_count == 1
|
||||
|
||||
def test_ssrf_protection(self, tmp_path):
|
||||
"""URLs with non-http(s) schemes should be skipped."""
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
# Register a service with an invalid scheme
|
||||
storage.register_service("channel", "ch-bad", "ftp://evil.example.com")
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch("turnstone.core.session.httpx.post") as mock_post,
|
||||
patch.dict("os.environ", {}, clear=False),
|
||||
patch("turnstone.core.session.time.sleep"),
|
||||
):
|
||||
call_id, msg = session._exec_notify(item)
|
||||
|
||||
# httpx.post should never be called for ftp:// URL
|
||||
mock_post.assert_not_called()
|
||||
assert "delivery failed" in msg.lower()
|
||||
|
||||
def test_retry_on_no_services(self, tmp_path):
|
||||
"""Retries service lookup when no gateways are initially available."""
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
# First two calls return empty, third returns a service
|
||||
call_count = 0
|
||||
|
||||
def _list_services(stype: str, max_age_seconds: int = 120) -> list[dict[str, str]]:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count <= 2:
|
||||
return []
|
||||
return [
|
||||
{
|
||||
"service_type": "channel",
|
||||
"service_id": "ch-1",
|
||||
"url": "http://localhost:8091",
|
||||
"metadata": "{}",
|
||||
"last_heartbeat": "",
|
||||
"created": "",
|
||||
}
|
||||
]
|
||||
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.json.return_value = {
|
||||
"results": [{"channel_type": "discord", "channel_id": "123", "status": "sent"}]
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
mock_storage = MagicMock()
|
||||
mock_storage.list_services = _list_services
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=mock_storage),
|
||||
patch("turnstone.core.session.httpx.post", return_value=mock_resp),
|
||||
patch.dict("os.environ", {}, clear=False),
|
||||
patch("turnstone.core.session.time.sleep") as mock_sleep,
|
||||
):
|
||||
call_id, msg = session._exec_notify(item)
|
||||
|
||||
assert "sent successfully" in msg.lower()
|
||||
# Should have slept twice (retry delays)
|
||||
assert mock_sleep.call_count == 2
|
||||
|
||||
def test_retry_on_all_gateways_failed(self, tmp_path):
|
||||
"""Retries when all gateways fail on first attempt but succeed on retry."""
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
call_count = 0
|
||||
|
||||
def _post(*args, **kwargs):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count <= 1:
|
||||
raise ConnectionError("refused")
|
||||
resp = MagicMock()
|
||||
resp.status_code = 200
|
||||
resp.json.return_value = {
|
||||
"results": [{"channel_type": "discord", "channel_id": "123", "status": "sent"}]
|
||||
}
|
||||
return resp
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch("turnstone.core.session.httpx.post", side_effect=_post),
|
||||
patch.dict("os.environ", {}, clear=False),
|
||||
patch("turnstone.core.session.time.sleep") as mock_sleep,
|
||||
):
|
||||
call_id, msg = session._exec_notify(item)
|
||||
|
||||
assert "sent successfully" in msg.lower()
|
||||
assert mock_sleep.call_count == 1
|
||||
|
||||
def test_no_services_logs_warning(self, tmp_path):
|
||||
"""Server-side warning is logged when no services are available."""
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch("turnstone.core.session.time.sleep"),
|
||||
patch("turnstone.core.session.log") as mock_log,
|
||||
):
|
||||
session._exec_notify(item)
|
||||
|
||||
# Should have logged warnings for retries + final exhaustion
|
||||
warning_calls = [c for c in mock_log.warning.call_args_list]
|
||||
assert len(warning_calls) >= 3 # 2 retry warnings + 1 exhaustion
|
||||
events = [c.args[0] for c in warning_calls]
|
||||
assert "notify.no_services" in events
|
||||
assert "notify.no_services_exhausted" in events
|
||||
|
||||
def test_all_gateways_failed_logs_warning(self, tmp_path):
|
||||
"""Server-side warning is logged when all gateways fail."""
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch(
|
||||
"turnstone.core.session.httpx.post",
|
||||
side_effect=ConnectionError("refused"),
|
||||
),
|
||||
patch.dict("os.environ", {}, clear=False),
|
||||
patch("turnstone.core.session.time.sleep"),
|
||||
patch("turnstone.core.session.log") as mock_log,
|
||||
):
|
||||
session._exec_notify(item)
|
||||
|
||||
warning_calls = [c for c in mock_log.warning.call_args_list]
|
||||
events = [c.args[0] for c in warning_calls]
|
||||
# 2 retry warnings + 1 final failure
|
||||
assert "notify.all_gateways_failed" in events
|
||||
assert "notify.delivery_failed" in events
|
||||
|
||||
def test_gateway_200_but_no_delivery(self, tmp_path):
|
||||
"""HTTP 200 with all results failed should not count as success."""
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
|
||||
session = _make_session()
|
||||
item = {
|
||||
"call_id": "call_1",
|
||||
"func_name": "notify",
|
||||
"message": "Hello!",
|
||||
"username": "admin",
|
||||
"channel_type": "",
|
||||
"channel_id": "",
|
||||
"title": "",
|
||||
}
|
||||
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.json.return_value = {
|
||||
"results": [{"channel_type": "discord", "channel_id": "123", "status": "no_adapter"}]
|
||||
}
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
with (
|
||||
patch("turnstone.core.session.get_storage", return_value=storage),
|
||||
patch("turnstone.core.session.httpx.post", return_value=mock_resp),
|
||||
patch.dict("os.environ", {}, clear=False),
|
||||
patch("turnstone.core.session.time.sleep"),
|
||||
):
|
||||
call_id, msg = session._exec_notify(item)
|
||||
|
||||
assert "delivery failed" in msg.lower()
|
||||
assert session._notify_count == 0
|
||||
@@ -27,7 +27,7 @@ class TestServerSpec:
|
||||
expected = {
|
||||
"/v1/api/workstreams",
|
||||
"/v1/api/dashboard",
|
||||
"/v1/api/sessions",
|
||||
"/v1/api/workstreams/saved",
|
||||
"/v1/api/send",
|
||||
"/v1/api/approve",
|
||||
"/v1/api/plan",
|
||||
|
||||
@@ -1103,6 +1103,43 @@ class TestOpenAIParameterGating:
|
||||
assert "temperature" not in kwargs
|
||||
assert "reasoning_effort" not in kwargs
|
||||
|
||||
def test_gpt5_pro_unsupported_effort_falls_back(self) -> None:
|
||||
"""GPT-5 pro only supports 'high'; unsupported values fall back to default."""
|
||||
caps = self.provider.get_capabilities("gpt-5-pro")
|
||||
kwargs: dict[str, Any] = {}
|
||||
self.provider._apply_model_params(kwargs, caps, temperature=0.7, reasoning_effort="medium")
|
||||
assert "temperature" not in kwargs
|
||||
assert kwargs["reasoning_effort"] == "high" # fell back to default
|
||||
|
||||
def test_gpt5_pro_supported_effort_passes_through(self) -> None:
|
||||
"""GPT-5 pro accepts 'high' directly."""
|
||||
caps = self.provider.get_capabilities("gpt-5-pro")
|
||||
kwargs: dict[str, Any] = {}
|
||||
self.provider._apply_model_params(kwargs, caps, temperature=0.7, reasoning_effort="high")
|
||||
assert kwargs["reasoning_effort"] == "high"
|
||||
|
||||
def test_gpt54_1m_context_and_effort(self) -> None:
|
||||
"""GPT-5.4: 1M context, temperature when effort=none, xhigh supported."""
|
||||
caps = self.provider.get_capabilities("gpt-5.4")
|
||||
assert caps.context_window == 1050000
|
||||
kwargs: dict[str, Any] = {}
|
||||
self.provider._apply_model_params(kwargs, caps, temperature=0.7, reasoning_effort="none")
|
||||
assert kwargs["temperature"] == 0.7
|
||||
assert "reasoning_effort" not in kwargs
|
||||
kwargs2: dict[str, Any] = {}
|
||||
self.provider._apply_model_params(kwargs2, caps, temperature=0.7, reasoning_effort="xhigh")
|
||||
assert "temperature" not in kwargs2
|
||||
assert kwargs2["reasoning_effort"] == "xhigh"
|
||||
|
||||
def test_gpt54_pro_no_temperature_always_reasoning(self) -> None:
|
||||
"""GPT-5.4 pro: no temperature, medium/high/xhigh only."""
|
||||
caps = self.provider.get_capabilities("gpt-5.4-pro")
|
||||
assert caps.context_window == 1050000
|
||||
kwargs: dict[str, Any] = {}
|
||||
self.provider._apply_model_params(kwargs, caps, temperature=0.7, reasoning_effort="low")
|
||||
assert "temperature" not in kwargs
|
||||
assert kwargs["reasoning_effort"] == "medium" # fell back from unsupported "low"
|
||||
|
||||
|
||||
class TestAnthropicReasoningNone:
|
||||
"""Verify 'none' effort disables thinking for manual-thinking models."""
|
||||
@@ -1896,3 +1933,254 @@ class TestAnthropicProviderBlocks:
|
||||
assert blocks[1]["input"] == {"query": "test"} # parsed from accumulated JSON
|
||||
assert blocks[2]["type"] == "web_search_tool_result"
|
||||
assert blocks[2]["encrypted_content"] == "enc_data"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tool search tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAnthropicToolSearch:
|
||||
"""Test Anthropic provider tool search injection."""
|
||||
|
||||
@pytest.fixture()
|
||||
def provider(self):
|
||||
from turnstone.core.providers._anthropic import AnthropicProvider
|
||||
|
||||
return AnthropicProvider()
|
||||
|
||||
def test_tool_search_capability_flag(self, provider):
|
||||
caps = provider.get_capabilities("claude-opus-4-6-20260101")
|
||||
assert caps.supports_tool_search is True
|
||||
|
||||
def test_tool_search_not_supported_on_haiku(self, provider):
|
||||
caps = provider.get_capabilities("claude-haiku-4-5-20251001")
|
||||
assert caps.supports_tool_search is False
|
||||
|
||||
def test_inject_tool_search_marks_deferred(self, provider):
|
||||
caps = provider.get_capabilities("claude-opus-4-6-20260101")
|
||||
tools = [
|
||||
{"name": "bash", "description": "Run commands", "input_schema": {}},
|
||||
{
|
||||
"name": "mcp__github__create_issue",
|
||||
"description": "Create issue",
|
||||
"input_schema": {},
|
||||
},
|
||||
]
|
||||
deferred = frozenset(["mcp__github__create_issue"])
|
||||
result = provider._inject_tool_search(tools, caps, deferred)
|
||||
# bash should not be deferred
|
||||
assert result[0].get("defer_loading") is None or result[0].get("defer_loading") is False
|
||||
# MCP tool should be deferred
|
||||
assert result[1]["defer_loading"] is True
|
||||
# Search tool should be appended
|
||||
assert result[-1]["type"] == "tool_search_tool_bm25_20251119"
|
||||
assert result[-1]["name"] == "tool_search"
|
||||
|
||||
def test_inject_tool_search_no_op_without_deferred(self, provider):
|
||||
caps = provider.get_capabilities("claude-opus-4-6-20260101")
|
||||
tools = [{"name": "bash", "description": "Run commands", "input_schema": {}}]
|
||||
result = provider._inject_tool_search(tools, caps, None)
|
||||
assert result == tools
|
||||
|
||||
def test_inject_tool_search_no_op_on_unsupported_model(self, provider):
|
||||
caps = provider.get_capabilities("claude-haiku-4-5-20251001")
|
||||
tools = [{"name": "bash", "description": "Run commands", "input_schema": {}}]
|
||||
deferred = frozenset(["some_tool"])
|
||||
result = provider._inject_tool_search(tools, caps, deferred)
|
||||
assert result == tools
|
||||
|
||||
|
||||
class TestOpenAIToolSearch:
|
||||
"""Test OpenAI provider tool search injection."""
|
||||
|
||||
@pytest.fixture()
|
||||
def provider(self):
|
||||
return OpenAIProvider()
|
||||
|
||||
def test_tool_search_capability_on_gpt54(self, provider):
|
||||
caps = provider.get_capabilities("gpt-5.4")
|
||||
assert caps.supports_tool_search is True
|
||||
|
||||
def test_tool_search_not_supported_on_gpt5(self, provider):
|
||||
caps = provider.get_capabilities("gpt-5")
|
||||
assert caps.supports_tool_search is False
|
||||
|
||||
def test_apply_tool_search_marks_deferred(self, provider):
|
||||
caps = provider.get_capabilities("gpt-5.4")
|
||||
tools = [
|
||||
{"type": "function", "function": {"name": "bash", "description": "Run commands"}},
|
||||
{
|
||||
"type": "function",
|
||||
"function": {"name": "mcp__slack__send", "description": "Send message"},
|
||||
},
|
||||
]
|
||||
deferred = frozenset(["mcp__slack__send"])
|
||||
result = provider._apply_tool_search(caps, tools, deferred)
|
||||
assert result is not None
|
||||
# bash not deferred
|
||||
assert result[0].get("defer_loading") is None or result[0].get("defer_loading") is False
|
||||
# slack tool deferred
|
||||
assert result[1]["defer_loading"] is True
|
||||
|
||||
def test_apply_tool_search_no_op_without_deferred(self, provider):
|
||||
caps = provider.get_capabilities("gpt-5.4")
|
||||
tools = [
|
||||
{"type": "function", "function": {"name": "bash", "description": "Run commands"}},
|
||||
]
|
||||
result = provider._apply_tool_search(caps, tools, None)
|
||||
assert result == tools
|
||||
|
||||
def test_apply_tool_search_no_op_on_unsupported_model(self, provider):
|
||||
caps = provider.get_capabilities("gpt-5")
|
||||
tools = [
|
||||
{"type": "function", "function": {"name": "bash", "description": "Run commands"}},
|
||||
]
|
||||
deferred = frozenset(["some_tool"])
|
||||
result = provider._apply_tool_search(caps, tools, deferred)
|
||||
assert result == tools
|
||||
|
||||
|
||||
class TestModelCapabilitiesToolSearch:
|
||||
"""Test supports_tool_search defaults and values."""
|
||||
|
||||
def test_default_is_false(self):
|
||||
from turnstone.core.providers._protocol import ModelCapabilities
|
||||
|
||||
caps = ModelCapabilities()
|
||||
assert caps.supports_tool_search is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Vision support
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestVisionCapabilities:
|
||||
"""Test supports_vision flag across providers."""
|
||||
|
||||
def test_default_is_false(self) -> None:
|
||||
from turnstone.core.providers._protocol import ModelCapabilities
|
||||
|
||||
caps = ModelCapabilities()
|
||||
assert caps.supports_vision is False
|
||||
|
||||
def test_openai_commercial_supports_vision(self) -> None:
|
||||
provider = OpenAIProvider()
|
||||
for model in ("gpt-5", "gpt-5-mini", "gpt-5.4", "o3", "o4-mini"):
|
||||
caps = provider.get_capabilities(model)
|
||||
assert caps.supports_vision is True, f"{model} should support vision"
|
||||
|
||||
def test_openai_default_no_vision(self) -> None:
|
||||
"""Unknown models (local servers) default to no vision."""
|
||||
provider = OpenAIProvider()
|
||||
caps = provider.get_capabilities("some-local-model")
|
||||
assert caps.supports_vision is False
|
||||
|
||||
def test_anthropic_supports_vision(self) -> None:
|
||||
from turnstone.core.providers._anthropic import AnthropicProvider
|
||||
|
||||
provider = AnthropicProvider()
|
||||
for model in ("claude-opus-4-6", "claude-sonnet-4-6", "claude-haiku-4-5"):
|
||||
caps = provider.get_capabilities(model)
|
||||
assert caps.supports_vision is True, f"{model} should support vision"
|
||||
|
||||
def test_anthropic_default_supports_vision(self) -> None:
|
||||
"""Anthropic default (unknown Claude model) supports vision."""
|
||||
from turnstone.core.providers._anthropic import AnthropicProvider
|
||||
|
||||
provider = AnthropicProvider()
|
||||
caps = provider.get_capabilities("claude-unknown-9")
|
||||
assert caps.supports_vision is True
|
||||
|
||||
|
||||
class TestAnthropicVisionConversion:
|
||||
"""Test image content conversion in _convert_messages."""
|
||||
|
||||
def setup_method(self) -> None:
|
||||
from turnstone.core.providers._anthropic import AnthropicProvider
|
||||
|
||||
self.provider = AnthropicProvider()
|
||||
|
||||
def test_tool_result_with_image_content(self) -> None:
|
||||
"""Tool result with list content converts image_url to Anthropic image."""
|
||||
messages = [
|
||||
{"role": "user", "content": "Read this image"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"function": {"name": "read_file", "arguments": '{"path": "img.png"}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_1",
|
||||
"content": [
|
||||
{"type": "text", "text": "Image file: img.png (1024 bytes)"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,iVBORw0KGgo="},
|
||||
},
|
||||
],
|
||||
},
|
||||
]
|
||||
_, converted = self.provider._convert_messages(messages)
|
||||
# Tool result should be in a user message
|
||||
tool_user_msg = converted[2]
|
||||
assert tool_user_msg["role"] == "user"
|
||||
tool_result = tool_user_msg["content"][0]
|
||||
assert tool_result["type"] == "tool_result"
|
||||
assert tool_result["tool_use_id"] == "call_1"
|
||||
# Content should be a list with converted image block
|
||||
content = tool_result["content"]
|
||||
assert isinstance(content, list)
|
||||
assert content[0] == {"type": "text", "text": "Image file: img.png (1024 bytes)"}
|
||||
assert content[1]["type"] == "image"
|
||||
assert content[1]["source"]["type"] == "base64"
|
||||
assert content[1]["source"]["media_type"] == "image/png"
|
||||
assert content[1]["source"]["data"] == "iVBORw0KGgo="
|
||||
|
||||
def test_tool_result_with_string_content_unchanged(self) -> None:
|
||||
"""Tool result with plain string content is unchanged."""
|
||||
messages = [
|
||||
{"role": "user", "content": "Read file"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_2",
|
||||
"function": {"name": "read_file", "arguments": '{"path": "f.py"}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_2",
|
||||
"content": " 1\tprint('hello')",
|
||||
},
|
||||
]
|
||||
_, converted = self.provider._convert_messages(messages)
|
||||
tool_result = converted[2]["content"][0]
|
||||
assert tool_result["content"] == " 1\tprint('hello')"
|
||||
|
||||
def test_convert_content_parts_static_method(self) -> None:
|
||||
"""_convert_content_parts handles both image_url and text."""
|
||||
from turnstone.core.providers._anthropic import AnthropicProvider
|
||||
|
||||
parts = [
|
||||
{"type": "text", "text": "description"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/jpeg;base64,/9j/4AAQ"},
|
||||
},
|
||||
]
|
||||
result = AnthropicProvider._convert_content_parts(parts)
|
||||
assert result[0] == {"type": "text", "text": "description"}
|
||||
assert result[1]["type"] == "image"
|
||||
assert result[1]["source"]["media_type"] == "image/jpeg"
|
||||
assert result[1]["source"]["data"] == "/9j/4AAQ"
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Tests for the atomic workstream resumption flow.
|
||||
|
||||
Covers CreateWorkstreamMessage resume_ws field, WorkstreamResumedEvent,
|
||||
WorkstreamCreatedEvent resumed fields, and server endpoint handling.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
from turnstone.mq.protocol import (
|
||||
CreateWorkstreamMessage,
|
||||
WorkstreamCreatedEvent,
|
||||
WorkstreamResumedEvent,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Protocol tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCreateWorkstreamMessageResumeField:
|
||||
def test_resume_ws_defaults_empty(self) -> None:
|
||||
msg = CreateWorkstreamMessage(name="test")
|
||||
assert msg.resume_ws == ""
|
||||
|
||||
def test_resume_ws_set(self) -> None:
|
||||
msg = CreateWorkstreamMessage(name="test", resume_ws="ws-abc")
|
||||
assert msg.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_deserializes(self) -> None:
|
||||
msg = CreateWorkstreamMessage(resume_ws="ws-123")
|
||||
raw = msg.to_json()
|
||||
from turnstone.mq.protocol import InboundMessage
|
||||
|
||||
restored = InboundMessage.from_json(raw)
|
||||
assert getattr(restored, "resume_ws", "") == "ws-123"
|
||||
|
||||
|
||||
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
|
||||
@@ -0,0 +1,264 @@
|
||||
"""Tests for scheduled task admin API endpoints."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from starlette.applications import Starlette
|
||||
from starlette.routing import Mount, Route
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from turnstone.console.server import (
|
||||
admin_create_schedule,
|
||||
admin_delete_schedule,
|
||||
admin_get_schedule,
|
||||
admin_list_schedule_runs,
|
||||
admin_list_schedules,
|
||||
admin_update_schedule,
|
||||
)
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path):
|
||||
"""Fresh SQLite backend for each test."""
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(storage):
|
||||
"""TestClient with storage and auth bypassed."""
|
||||
app = Starlette(
|
||||
routes=[
|
||||
Mount(
|
||||
"/v1",
|
||||
routes=[
|
||||
Route("/api/admin/schedules", admin_list_schedules),
|
||||
Route("/api/admin/schedules", admin_create_schedule, methods=["POST"]),
|
||||
Route("/api/admin/schedules/{task_id}", admin_get_schedule),
|
||||
Route(
|
||||
"/api/admin/schedules/{task_id}",
|
||||
admin_update_schedule,
|
||||
methods=["PUT"],
|
||||
),
|
||||
Route(
|
||||
"/api/admin/schedules/{task_id}",
|
||||
admin_delete_schedule,
|
||||
methods=["DELETE"],
|
||||
),
|
||||
Route(
|
||||
"/api/admin/schedules/{task_id}/runs",
|
||||
admin_list_schedule_runs,
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def _cron_payload(**overrides):
|
||||
"""Build default cron schedule creation payload."""
|
||||
defaults = {
|
||||
"name": "Daily report",
|
||||
"description": "Generate the summary",
|
||||
"schedule_type": "cron",
|
||||
"cron_expr": "0 9 * * *",
|
||||
"target_mode": "auto",
|
||||
"model": "gpt-5",
|
||||
"initial_message": "Generate the daily report",
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return defaults
|
||||
|
||||
|
||||
def _at_payload(**overrides):
|
||||
"""Build default at-time schedule creation payload."""
|
||||
defaults = {
|
||||
"name": "One-shot task",
|
||||
"description": "Run once",
|
||||
"schedule_type": "at",
|
||||
"at_time": "2099-01-01T00:00:00+00:00",
|
||||
"target_mode": "auto",
|
||||
"model": "gpt-5",
|
||||
"initial_message": "Do the thing",
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return defaults
|
||||
|
||||
|
||||
class TestScheduleAPI:
|
||||
"""Tests for the 6 admin schedule endpoints."""
|
||||
|
||||
def test_list_empty(self, client):
|
||||
resp = client.get("/v1/api/admin/schedules")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["schedules"] == []
|
||||
|
||||
def test_create_cron(self, client):
|
||||
resp = client.post("/v1/api/admin/schedules", json=_cron_payload())
|
||||
assert resp.status_code == 200
|
||||
task = resp.json()
|
||||
assert task["name"] == "Daily report"
|
||||
assert task["schedule_type"] == "cron"
|
||||
assert task["cron_expr"] == "0 9 * * *"
|
||||
assert task["enabled"] is True
|
||||
assert "task_id" in task
|
||||
assert "created" in task
|
||||
assert "next_run" in task
|
||||
assert task["next_run"] != ""
|
||||
|
||||
def test_create_at(self, client):
|
||||
resp = client.post("/v1/api/admin/schedules", json=_at_payload())
|
||||
assert resp.status_code == 200
|
||||
task = resp.json()
|
||||
assert task["schedule_type"] == "at"
|
||||
assert task["at_time"] == "2099-01-01T00:00:00+00:00"
|
||||
assert task["next_run"] == "2099-01-01T00:00:00+00:00"
|
||||
|
||||
def test_create_missing_name(self, client):
|
||||
payload = _cron_payload()
|
||||
del payload["name"]
|
||||
resp = client.post("/v1/api/admin/schedules", json=payload)
|
||||
assert resp.status_code == 400
|
||||
assert "name" in resp.json()["error"].lower()
|
||||
|
||||
def test_create_invalid_cron(self, client):
|
||||
resp = client.post(
|
||||
"/v1/api/admin/schedules",
|
||||
json=_cron_payload(cron_expr="not a cron"),
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
assert "cron" in resp.json()["error"].lower()
|
||||
|
||||
def test_create_naive_at_time(self, client):
|
||||
"""Naive timestamps (no timezone) should be rejected."""
|
||||
resp = client.post(
|
||||
"/v1/api/admin/schedules",
|
||||
json=_at_payload(at_time="2099-01-01T00:00:00"),
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
assert "timezone" in resp.json()["error"].lower()
|
||||
|
||||
def test_create_past_at_time(self, client):
|
||||
resp = client.post(
|
||||
"/v1/api/admin/schedules",
|
||||
json=_at_payload(at_time="2000-01-01T00:00:00+00:00"),
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
assert "future" in resp.json()["error"].lower()
|
||||
|
||||
def test_get_schedule(self, client):
|
||||
create_resp = client.post("/v1/api/admin/schedules", json=_cron_payload())
|
||||
task_id = create_resp.json()["task_id"]
|
||||
|
||||
resp = client.get(f"/v1/api/admin/schedules/{task_id}")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["task_id"] == task_id
|
||||
assert resp.json()["name"] == "Daily report"
|
||||
|
||||
def test_get_nonexistent(self, client):
|
||||
resp = client.get("/v1/api/admin/schedules/nonexistent_id")
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_update_schedule(self, client):
|
||||
create_resp = client.post("/v1/api/admin/schedules", json=_cron_payload())
|
||||
task_id = create_resp.json()["task_id"]
|
||||
|
||||
resp = client.put(
|
||||
f"/v1/api/admin/schedules/{task_id}",
|
||||
json={"name": "Weekly report"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["name"] == "Weekly report"
|
||||
|
||||
# Verify via GET
|
||||
get_resp = client.get(f"/v1/api/admin/schedules/{task_id}")
|
||||
assert get_resp.json()["name"] == "Weekly report"
|
||||
|
||||
def test_update_nonexistent(self, client):
|
||||
resp = client.put(
|
||||
"/v1/api/admin/schedules/nonexistent_id",
|
||||
json={"name": "Nope"},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_delete_schedule(self, client):
|
||||
create_resp = client.post("/v1/api/admin/schedules", json=_cron_payload())
|
||||
task_id = create_resp.json()["task_id"]
|
||||
|
||||
resp = client.delete(f"/v1/api/admin/schedules/{task_id}")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["status"] == "ok"
|
||||
|
||||
# Verify gone
|
||||
get_resp = client.get(f"/v1/api/admin/schedules/{task_id}")
|
||||
assert get_resp.status_code == 404
|
||||
|
||||
def test_delete_nonexistent(self, client):
|
||||
resp = client.delete("/v1/api/admin/schedules/nonexistent_id")
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_list_runs_empty(self, client):
|
||||
create_resp = client.post("/v1/api/admin/schedules", json=_cron_payload())
|
||||
task_id = create_resp.json()["task_id"]
|
||||
|
||||
resp = client.get(f"/v1/api/admin/schedules/{task_id}/runs")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["runs"] == []
|
||||
|
||||
def test_list_runs_nonexistent(self, client):
|
||||
resp = client.get("/v1/api/admin/schedules/nonexistent_id/runs")
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_create_specific_node_target(self, client):
|
||||
payload = _cron_payload(target_mode="node-custom-001")
|
||||
resp = client.post("/v1/api/admin/schedules", json=payload)
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["target_mode"] == "node-custom-001"
|
||||
|
||||
def test_list_runs_with_data(self, client, storage):
|
||||
create_resp = client.post("/v1/api/admin/schedules", json=_cron_payload())
|
||||
task_id = create_resp.json()["task_id"]
|
||||
|
||||
# Record runs directly in storage
|
||||
storage.record_task_run(
|
||||
run_id="run_001",
|
||||
task_id=task_id,
|
||||
node_id="node-1",
|
||||
ws_id="ws_abc",
|
||||
correlation_id="corr_001",
|
||||
started="2025-06-01T09:00:00",
|
||||
status="dispatched",
|
||||
error="",
|
||||
)
|
||||
storage.record_task_run(
|
||||
run_id="run_002",
|
||||
task_id=task_id,
|
||||
node_id="node-2",
|
||||
ws_id="",
|
||||
correlation_id="corr_002",
|
||||
started="2025-06-01T09:01:00",
|
||||
status="failed",
|
||||
error="No reachable nodes",
|
||||
)
|
||||
|
||||
resp = client.get(f"/v1/api/admin/schedules/{task_id}/runs")
|
||||
assert resp.status_code == 200
|
||||
runs = resp.json()["runs"]
|
||||
assert len(runs) == 2
|
||||
# Most recent first
|
||||
assert runs[0]["run_id"] == "run_002"
|
||||
assert runs[0]["status"] == "failed"
|
||||
assert runs[1]["run_id"] == "run_001"
|
||||
|
||||
def test_list_runs_invalid_limit(self, client):
|
||||
create_resp = client.post("/v1/api/admin/schedules", json=_cron_payload())
|
||||
task_id = create_resp.json()["task_id"]
|
||||
|
||||
# Invalid limit should not crash — falls back to 50
|
||||
resp = client.get(f"/v1/api/admin/schedules/{task_id}/runs?limit=abc")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["runs"] == []
|
||||
@@ -0,0 +1,259 @@
|
||||
"""Tests for scheduled_tasks and scheduled_task_runs storage CRUD."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def db(tmp_path):
|
||||
"""Fresh SQLite backend for each test."""
|
||||
backend = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
return backend
|
||||
|
||||
|
||||
def _make_task_kwargs(**overrides):
|
||||
"""Build default kwargs for create_scheduled_task."""
|
||||
defaults = {
|
||||
"task_id": "task_001",
|
||||
"name": "Daily report",
|
||||
"description": "Generate the daily summary",
|
||||
"schedule_type": "cron",
|
||||
"cron_expr": "0 9 * * *",
|
||||
"at_time": "",
|
||||
"target_mode": "auto",
|
||||
"model": "gpt-5",
|
||||
"initial_message": "Generate the daily report",
|
||||
"auto_approve": False,
|
||||
"auto_approve_tools": [],
|
||||
"created_by": "u_admin",
|
||||
"next_run": "2099-01-01T09:00:00",
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return defaults
|
||||
|
||||
|
||||
class TestScheduledTaskCRUD:
|
||||
"""Tests for scheduled_tasks table operations."""
|
||||
|
||||
def test_create_and_get(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs())
|
||||
result = db.get_scheduled_task("task_001")
|
||||
assert result is not None
|
||||
assert result["task_id"] == "task_001"
|
||||
assert result["name"] == "Daily report"
|
||||
assert result["description"] == "Generate the daily summary"
|
||||
assert result["schedule_type"] == "cron"
|
||||
assert result["cron_expr"] == "0 9 * * *"
|
||||
assert result["at_time"] == ""
|
||||
assert result["target_mode"] == "auto"
|
||||
assert result["model"] == "gpt-5"
|
||||
assert result["initial_message"] == "Generate the daily report"
|
||||
assert result["auto_approve"] == 0
|
||||
assert result["auto_approve_tools"] == ""
|
||||
assert result["enabled"] == 1
|
||||
assert result["created_by"] == "u_admin"
|
||||
assert result["next_run"] == "2099-01-01T09:00:00"
|
||||
assert "created" in result
|
||||
assert "updated" in result
|
||||
|
||||
def test_get_nonexistent(self, db):
|
||||
assert db.get_scheduled_task("no_such_task") is None
|
||||
|
||||
def test_create_duplicate_noop(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs(name="First"))
|
||||
db.create_scheduled_task(**_make_task_kwargs(name="Second"))
|
||||
result = db.get_scheduled_task("task_001")
|
||||
assert result is not None
|
||||
assert result["name"] == "First" # first write wins
|
||||
|
||||
def test_list_tasks(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs(task_id="task_a", name="Alpha"))
|
||||
# Ensure different created timestamps (resolution is 1 second)
|
||||
time.sleep(1.1)
|
||||
db.create_scheduled_task(**_make_task_kwargs(task_id="task_b", name="Beta"))
|
||||
tasks = db.list_scheduled_tasks()
|
||||
assert len(tasks) == 2
|
||||
# Ordered by created DESC — most recent first
|
||||
assert tasks[0]["task_id"] == "task_b"
|
||||
assert tasks[1]["task_id"] == "task_a"
|
||||
|
||||
def test_update_task(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs())
|
||||
original = db.get_scheduled_task("task_001")
|
||||
assert original is not None
|
||||
original_updated = original["updated"]
|
||||
|
||||
time.sleep(0.05)
|
||||
result = db.update_scheduled_task("task_001", name="Weekly report")
|
||||
assert result is True
|
||||
|
||||
updated = db.get_scheduled_task("task_001")
|
||||
assert updated is not None
|
||||
assert updated["name"] == "Weekly report"
|
||||
assert updated["updated"] >= original_updated
|
||||
|
||||
def test_update_enable_disable(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs())
|
||||
task = db.get_scheduled_task("task_001")
|
||||
assert task is not None
|
||||
assert task["enabled"] == 1
|
||||
|
||||
db.update_scheduled_task("task_001", enabled=False)
|
||||
task = db.get_scheduled_task("task_001")
|
||||
assert task is not None
|
||||
assert task["enabled"] == 0
|
||||
|
||||
db.update_scheduled_task("task_001", enabled=True)
|
||||
task = db.get_scheduled_task("task_001")
|
||||
assert task is not None
|
||||
assert task["enabled"] == 1
|
||||
|
||||
def test_delete_task(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs())
|
||||
assert db.delete_scheduled_task("task_001") is True
|
||||
assert db.get_scheduled_task("task_001") is None
|
||||
# Deleting again returns False
|
||||
assert db.delete_scheduled_task("task_001") is False
|
||||
|
||||
def test_delete_cascades_runs(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs())
|
||||
db.record_task_run(
|
||||
run_id="run_001",
|
||||
task_id="task_001",
|
||||
node_id="node_1",
|
||||
ws_id="ws_abc",
|
||||
correlation_id="corr_001",
|
||||
started="2025-01-01T09:00:00",
|
||||
status="dispatched",
|
||||
error="",
|
||||
)
|
||||
assert len(db.list_task_runs("task_001")) == 1
|
||||
|
||||
db.delete_scheduled_task("task_001")
|
||||
assert db.list_task_runs("task_001") == []
|
||||
|
||||
def test_list_due_tasks(self, db):
|
||||
db.create_scheduled_task(
|
||||
**_make_task_kwargs(task_id="past", next_run="2020-01-01T00:00:00")
|
||||
)
|
||||
db.create_scheduled_task(
|
||||
**_make_task_kwargs(task_id="future", next_run="2099-12-31T23:59:59")
|
||||
)
|
||||
now = "2025-06-01T12:00:00"
|
||||
due = db.list_due_tasks(now)
|
||||
assert len(due) == 1
|
||||
assert due[0]["task_id"] == "past"
|
||||
|
||||
def test_list_due_tasks_skips_disabled(self, db):
|
||||
db.create_scheduled_task(
|
||||
**_make_task_kwargs(task_id="disabled_task", next_run="2020-01-01T00:00:00")
|
||||
)
|
||||
db.update_scheduled_task("disabled_task", enabled=False)
|
||||
due = db.list_due_tasks("2025-06-01T12:00:00")
|
||||
assert len(due) == 0
|
||||
|
||||
def test_list_due_tasks_empty_next_run(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs(task_id="empty_next", next_run=""))
|
||||
due = db.list_due_tasks("2099-12-31T23:59:59")
|
||||
assert len(due) == 0
|
||||
|
||||
def test_at_task_fields(self, db):
|
||||
db.create_scheduled_task(
|
||||
**_make_task_kwargs(
|
||||
task_id="at_task",
|
||||
schedule_type="at",
|
||||
cron_expr="",
|
||||
at_time="2099-06-15T14:00:00",
|
||||
next_run="2099-06-15T14:00:00",
|
||||
)
|
||||
)
|
||||
result = db.get_scheduled_task("at_task")
|
||||
assert result is not None
|
||||
assert result["schedule_type"] == "at"
|
||||
assert result["at_time"] == "2099-06-15T14:00:00"
|
||||
|
||||
|
||||
class TestScheduledTaskRuns:
|
||||
"""Tests for scheduled_task_runs table operations."""
|
||||
|
||||
def test_record_and_list(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs())
|
||||
db.record_task_run(
|
||||
run_id="run_a",
|
||||
task_id="task_001",
|
||||
node_id="node_1",
|
||||
ws_id="ws_1",
|
||||
correlation_id="corr_a",
|
||||
started="2025-01-01T09:00:00",
|
||||
status="dispatched",
|
||||
error="",
|
||||
)
|
||||
db.record_task_run(
|
||||
run_id="run_b",
|
||||
task_id="task_001",
|
||||
node_id="node_2",
|
||||
ws_id="ws_2",
|
||||
correlation_id="corr_b",
|
||||
started="2025-01-02T09:00:00",
|
||||
status="dispatched",
|
||||
error="",
|
||||
)
|
||||
runs = db.list_task_runs("task_001")
|
||||
assert len(runs) == 2
|
||||
# Ordered by started DESC — most recent first
|
||||
assert runs[0]["run_id"] == "run_b"
|
||||
assert runs[1]["run_id"] == "run_a"
|
||||
|
||||
def test_list_runs_respects_limit(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs())
|
||||
for i in range(3):
|
||||
db.record_task_run(
|
||||
run_id=f"run_{i}",
|
||||
task_id="task_001",
|
||||
node_id="node_1",
|
||||
ws_id="",
|
||||
correlation_id=f"corr_{i}",
|
||||
started=f"2025-01-0{i + 1}T09:00:00",
|
||||
status="dispatched",
|
||||
error="",
|
||||
)
|
||||
runs = db.list_task_runs("task_001", limit=2)
|
||||
assert len(runs) == 2
|
||||
|
||||
def test_list_runs_empty(self, db):
|
||||
assert db.list_task_runs("no_such_task") == []
|
||||
|
||||
def test_prune_task_runs(self, db):
|
||||
db.create_scheduled_task(**_make_task_kwargs())
|
||||
# Old run (should be pruned)
|
||||
db.record_task_run(
|
||||
run_id="old_run",
|
||||
task_id="task_001",
|
||||
node_id="node_1",
|
||||
ws_id="",
|
||||
correlation_id="c_old",
|
||||
started="2020-01-01T00:00:00",
|
||||
status="dispatched",
|
||||
error="",
|
||||
)
|
||||
# Recent run (should survive)
|
||||
db.record_task_run(
|
||||
run_id="new_run",
|
||||
task_id="task_001",
|
||||
node_id="node_1",
|
||||
ws_id="",
|
||||
correlation_id="c_new",
|
||||
started="2099-01-01T00:00:00",
|
||||
status="dispatched",
|
||||
error="",
|
||||
)
|
||||
pruned = db.prune_task_runs(retention_days=90)
|
||||
assert pruned == 1
|
||||
runs = db.list_task_runs("task_001")
|
||||
assert len(runs) == 1
|
||||
assert runs[0]["run_id"] == "new_run"
|
||||
@@ -0,0 +1,272 @@
|
||||
"""Tests for turnstone.console.scheduler — TaskScheduler tick and dispatch."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.console.scheduler import TaskScheduler
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mocks():
|
||||
"""Broker, collector, and storage mocks for scheduler tests."""
|
||||
broker = MagicMock()
|
||||
broker._redis = MagicMock()
|
||||
collector = MagicMock()
|
||||
storage = MagicMock()
|
||||
return broker, collector, storage
|
||||
|
||||
|
||||
def _make_task(**overrides):
|
||||
"""Build a minimal task dict matching storage row format."""
|
||||
defaults = {
|
||||
"task_id": "task_001",
|
||||
"name": "Test task",
|
||||
"description": "",
|
||||
"schedule_type": "cron",
|
||||
"cron_expr": "0 9 * * *",
|
||||
"at_time": "",
|
||||
"target_mode": "auto",
|
||||
"model": "gpt-5",
|
||||
"initial_message": "Run the tests",
|
||||
"auto_approve": 0,
|
||||
"auto_approve_tools": "",
|
||||
"enabled": 1,
|
||||
"created_by": "u_admin",
|
||||
"next_run": "2020-01-01T09:00:00",
|
||||
"last_run": "",
|
||||
"created": "2020-01-01T00:00:00",
|
||||
"updated": "2020-01-01T00:00:00",
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return defaults
|
||||
|
||||
|
||||
def _make_node(node_id="node-001", reachable=True, ws_total=2, max_ws=10):
|
||||
"""Build a minimal node dict matching collector output."""
|
||||
return {
|
||||
"node_id": node_id,
|
||||
"reachable": reachable,
|
||||
"ws_total": ws_total,
|
||||
"max_ws": max_ws,
|
||||
}
|
||||
|
||||
|
||||
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
|
||||
storage.list_due_tasks.return_value = []
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage)
|
||||
scheduler._tick()
|
||||
|
||||
broker._redis.set.assert_called_once()
|
||||
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
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage)
|
||||
scheduler._tick()
|
||||
|
||||
storage.list_due_tasks.assert_not_called()
|
||||
|
||||
def test_dispatch_auto_mode(self, mocks):
|
||||
broker, collector, storage = mocks
|
||||
broker._redis.set.return_value = True
|
||||
|
||||
task = _make_task(target_mode="auto")
|
||||
storage.list_due_tasks.return_value = [task]
|
||||
collector.get_nodes.return_value = ([_make_node()], 1)
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage)
|
||||
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"
|
||||
)
|
||||
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
|
||||
|
||||
task = _make_task(target_mode="pool")
|
||||
storage.list_due_tasks.return_value = [task]
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage)
|
||||
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
|
||||
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
|
||||
|
||||
task = _make_task(target_mode="all")
|
||||
storage.list_due_tasks.return_value = [task]
|
||||
collector.get_nodes.return_value = (
|
||||
[_make_node("node-001"), _make_node("node-002")],
|
||||
2,
|
||||
)
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage)
|
||||
scheduler._tick()
|
||||
|
||||
assert broker.push_inbound.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
|
||||
|
||||
task = _make_task(target_mode="node-001")
|
||||
storage.list_due_tasks.return_value = [task]
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage)
|
||||
scheduler._tick()
|
||||
|
||||
broker.push_inbound.assert_called_once()
|
||||
_, kwargs = broker.push_inbound.call_args
|
||||
assert kwargs["node_id"] == "node-001"
|
||||
|
||||
def test_at_task_disables_after_dispatch(self, mocks):
|
||||
broker, collector, storage = mocks
|
||||
broker._redis.set.return_value = True
|
||||
|
||||
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)
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage)
|
||||
scheduler._tick()
|
||||
|
||||
# At-task should be disabled after dispatch
|
||||
update_calls = storage.update_scheduled_task.call_args_list
|
||||
assert len(update_calls) == 1
|
||||
args, kwargs = update_calls[0]
|
||||
assert args[0] == "task_001"
|
||||
assert kwargs["enabled"] is False
|
||||
assert kwargs["next_run"] == ""
|
||||
|
||||
def test_cron_task_updates_next_run(self, mocks):
|
||||
broker, collector, storage = mocks
|
||||
broker._redis.set.return_value = True
|
||||
|
||||
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)
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage)
|
||||
scheduler._tick()
|
||||
|
||||
update_calls = storage.update_scheduled_task.call_args_list
|
||||
assert len(update_calls) == 1
|
||||
_, kwargs = update_calls[0]
|
||||
assert kwargs["next_run"] != ""
|
||||
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
|
||||
|
||||
task = _make_task(target_mode="auto")
|
||||
storage.list_due_tasks.return_value = [task]
|
||||
# No reachable nodes
|
||||
collector.get_nodes.return_value = (
|
||||
[_make_node("node-001", reachable=False)],
|
||||
1,
|
||||
)
|
||||
|
||||
scheduler = TaskScheduler(broker, 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"
|
||||
assert run_kwargs["error"] != ""
|
||||
|
||||
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
|
||||
|
||||
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._tick()
|
||||
|
||||
# update_scheduled_task should NOT be called (no last_run/next_run advance)
|
||||
storage.update_scheduled_task.assert_not_called()
|
||||
|
||||
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
|
||||
|
||||
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)
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage, max_fan_out=3)
|
||||
scheduler._tick()
|
||||
|
||||
assert broker.push_inbound.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
|
||||
|
||||
task = _make_task(target_mode="node-custom-123")
|
||||
storage.list_due_tasks.return_value = [task]
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage)
|
||||
scheduler._tick()
|
||||
|
||||
broker.push_inbound.assert_called_once()
|
||||
call_kwargs = broker.push_inbound.call_args
|
||||
assert call_kwargs[1]["node_id"] == "node-custom-123"
|
||||
|
||||
def test_user_id_in_dispatched_message(self, mocks):
|
||||
"""Dispatched message should include created_by as user_id."""
|
||||
import json
|
||||
|
||||
broker, collector, storage = mocks
|
||||
broker._redis.set.return_value = True
|
||||
|
||||
task = _make_task(target_mode="pool", created_by="u_scheduler_admin")
|
||||
storage.list_due_tasks.return_value = [task]
|
||||
|
||||
scheduler = TaskScheduler(broker, collector, storage)
|
||||
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"
|
||||
@@ -2,6 +2,8 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
@@ -240,3 +242,140 @@ async def test_query_params_passed():
|
||||
assert "state=running" in captured_url[0]
|
||||
assert "page=2" in captured_url[0]
|
||||
assert "per_page=25" in captured_url[0]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Schedules
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_SCHEDULE_FIXTURE = {
|
||||
"task_id": "t1",
|
||||
"name": "nightly",
|
||||
"description": "",
|
||||
"schedule_type": "cron",
|
||||
"cron_expr": "0 2 * * *",
|
||||
"at_time": "",
|
||||
"target_mode": "auto",
|
||||
"model": "",
|
||||
"initial_message": "Run nightly checks",
|
||||
"auto_approve": False,
|
||||
"auto_approve_tools": [],
|
||||
"enabled": True,
|
||||
"created_by": "u1",
|
||||
"last_run": None,
|
||||
"next_run": "2026-03-06T02:00:00Z",
|
||||
"created": "2026-03-05T12:00:00Z",
|
||||
"updated": "2026-03-05T12:00:00Z",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_list_schedules():
|
||||
transport = _mock_transport(
|
||||
{"GET /v1/api/admin/schedules": _json_response({"schedules": [_SCHEDULE_FIXTURE]})}
|
||||
)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as hc:
|
||||
client = AsyncTurnstoneConsole(httpx_client=hc)
|
||||
resp = await client.list_schedules()
|
||||
assert len(resp.schedules) == 1
|
||||
assert resp.schedules[0].task_id == "t1"
|
||||
assert resp.schedules[0].name == "nightly"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_create_schedule():
|
||||
captured_body: list[dict] = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured_body.append(json.loads(request.content))
|
||||
return _json_response(_SCHEDULE_FIXTURE)
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as hc:
|
||||
client = AsyncTurnstoneConsole(httpx_client=hc)
|
||||
resp = await client.create_schedule(
|
||||
name="nightly",
|
||||
schedule_type="cron",
|
||||
initial_message="Run nightly checks",
|
||||
cron_expr="0 2 * * *",
|
||||
)
|
||||
assert resp.task_id == "t1"
|
||||
body = captured_body[0]
|
||||
assert body["name"] == "nightly"
|
||||
assert body["schedule_type"] == "cron"
|
||||
assert body["cron_expr"] == "0 2 * * *"
|
||||
assert body["initial_message"] == "Run nightly checks"
|
||||
# Optional fields with defaults should not appear when not set
|
||||
assert "description" not in body
|
||||
assert "model" not in body
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_get_schedule():
|
||||
transport = _mock_transport(
|
||||
{"GET /v1/api/admin/schedules/t1": _json_response(_SCHEDULE_FIXTURE)}
|
||||
)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as hc:
|
||||
client = AsyncTurnstoneConsole(httpx_client=hc)
|
||||
resp = await client.get_schedule("t1")
|
||||
assert resp.task_id == "t1"
|
||||
assert resp.schedule_type == "cron"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_update_schedule_partial():
|
||||
"""Only explicitly-passed fields should appear in the request body."""
|
||||
captured_body: list[dict] = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured_body.append(json.loads(request.content))
|
||||
return _json_response({**_SCHEDULE_FIXTURE, "enabled": False})
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as hc:
|
||||
client = AsyncTurnstoneConsole(httpx_client=hc)
|
||||
resp = await client.update_schedule("t1", enabled=False)
|
||||
assert resp.enabled is False
|
||||
body = captured_body[0]
|
||||
assert body == {"enabled": False}
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_delete_schedule():
|
||||
transport = _mock_transport(
|
||||
{"DELETE /v1/api/admin/schedules/t1": _json_response({"status": "ok"})}
|
||||
)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as hc:
|
||||
client = AsyncTurnstoneConsole(httpx_client=hc)
|
||||
resp = await client.delete_schedule("t1")
|
||||
assert resp.status == "ok"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_list_schedule_runs():
|
||||
transport = _mock_transport(
|
||||
{
|
||||
"GET /v1/api/admin/schedules/t1/runs": _json_response(
|
||||
{
|
||||
"runs": [
|
||||
{
|
||||
"run_id": "r1",
|
||||
"task_id": "t1",
|
||||
"node_id": "n1",
|
||||
"ws_id": "ws1",
|
||||
"correlation_id": "c1",
|
||||
"started": "2026-03-05T02:00:00Z",
|
||||
"status": "dispatched",
|
||||
"error": "",
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
}
|
||||
)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as hc:
|
||||
client = AsyncTurnstoneConsole(httpx_client=hc)
|
||||
resp = await client.list_schedule_runs("t1", limit=10)
|
||||
assert len(resp.runs) == 1
|
||||
assert resp.runs[0].run_id == "r1"
|
||||
assert resp.runs[0].status == "dispatched"
|
||||
|
||||
@@ -148,19 +148,19 @@ async def test_command():
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Sessions
|
||||
# History
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_list_sessions():
|
||||
async def test_list_saved_workstreams():
|
||||
transport = _mock_transport(
|
||||
{
|
||||
"GET /v1/api/sessions": _json_response(
|
||||
"GET /v1/api/workstreams/saved": _json_response(
|
||||
{
|
||||
"sessions": [
|
||||
"workstreams": [
|
||||
{
|
||||
"session_id": "s1",
|
||||
"ws_id": "s1",
|
||||
"title": "test",
|
||||
"created": "2024-01-01",
|
||||
"updated": "2024-01-02",
|
||||
@@ -173,8 +173,8 @@ async def test_list_sessions():
|
||||
)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as hc:
|
||||
client = AsyncTurnstoneServer(httpx_client=hc)
|
||||
resp = await client.list_sessions()
|
||||
assert len(resp.sessions) == 1
|
||||
resp = await client.list_saved_workstreams()
|
||||
assert len(resp.workstreams) == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -49,6 +49,21 @@ def test_sync_runner_iter():
|
||||
runner.close()
|
||||
|
||||
|
||||
def test_sync_runner_iter_empty():
|
||||
"""_SyncRunner.run_iter handles empty async generator via sentinel."""
|
||||
runner = _SyncRunner()
|
||||
try:
|
||||
|
||||
async def _empty():
|
||||
return
|
||||
yield # pragma: no cover # makes this an async generator
|
||||
|
||||
items = list(runner.run_iter(_empty()))
|
||||
assert items == []
|
||||
finally:
|
||||
runner.close()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TurnstoneServer (sync)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -609,7 +609,7 @@ class TestServerHealthMetrics:
|
||||
mock_ui._ws_context_ratio = 0.0
|
||||
|
||||
mock_session = MagicMock()
|
||||
mock_session.session_id = "test-session-id"
|
||||
mock_session.ws_id = "test-session-id"
|
||||
|
||||
mock_ws = MagicMock()
|
||||
mock_ws.id = "test-ws"
|
||||
@@ -785,7 +785,7 @@ class TestServerRateLimiting:
|
||||
mock_ui._ws_context_ratio = 0.0
|
||||
|
||||
mock_session = MagicMock()
|
||||
mock_session.session_id = "test-session-id"
|
||||
mock_session.ws_id = "test-session-id"
|
||||
|
||||
mock_ws = MagicMock()
|
||||
mock_ws.id = "test-ws"
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
"""Tests for the services registry storage methods."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path):
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
class TestServiceRegistry:
|
||||
def test_register_and_list(self, storage):
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
services = storage.list_services("channel", max_age_seconds=120)
|
||||
assert len(services) == 1
|
||||
assert services[0]["service_type"] == "channel"
|
||||
assert services[0]["service_id"] == "ch-1"
|
||||
assert services[0]["url"] == "http://localhost:8091"
|
||||
|
||||
def test_register_upsert(self, storage):
|
||||
storage.register_service("channel", "ch-1", "http://old:8091")
|
||||
storage.register_service("channel", "ch-1", "http://new:8091")
|
||||
services = storage.list_services("channel", max_age_seconds=120)
|
||||
assert len(services) == 1
|
||||
assert services[0]["url"] == "http://new:8091"
|
||||
|
||||
def test_heartbeat(self, storage):
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
result = storage.heartbeat_service("channel", "ch-1")
|
||||
assert result is True
|
||||
|
||||
def test_heartbeat_nonexistent(self, storage):
|
||||
result = storage.heartbeat_service("channel", "nonexistent")
|
||||
assert result is False
|
||||
|
||||
def test_list_filters_stale(self, storage):
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
# Manually set heartbeat to the past so it's stale
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.storage._schema import services
|
||||
|
||||
old_time = (datetime.now(UTC) - timedelta(seconds=300)).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with storage._engine.connect() as conn:
|
||||
conn.execute(sa.update(services).values(last_heartbeat=old_time))
|
||||
conn.commit()
|
||||
|
||||
# Should be excluded with 120s max age
|
||||
result = storage.list_services("channel", max_age_seconds=120)
|
||||
assert len(result) == 0
|
||||
|
||||
def test_list_empty(self, storage):
|
||||
services = storage.list_services("channel", max_age_seconds=120)
|
||||
assert services == []
|
||||
|
||||
def test_list_filters_by_type(self, storage):
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
storage.register_service("bridge", "br-1", "http://localhost:8080")
|
||||
channels = storage.list_services("channel", max_age_seconds=120)
|
||||
bridges = storage.list_services("bridge", max_age_seconds=120)
|
||||
assert len(channels) == 1
|
||||
assert len(bridges) == 1
|
||||
|
||||
def test_deregister(self, storage):
|
||||
storage.register_service("channel", "ch-1", "http://localhost:8091")
|
||||
result = storage.deregister_service("channel", "ch-1")
|
||||
assert result is True
|
||||
services = storage.list_services("channel", max_age_seconds=120)
|
||||
assert services == []
|
||||
|
||||
def test_deregister_nonexistent(self, storage):
|
||||
result = storage.deregister_service("channel", "nonexistent")
|
||||
assert result is False
|
||||
|
||||
def test_metadata(self, storage):
|
||||
storage.register_service(
|
||||
"channel", "ch-1", "http://localhost:8091", metadata='{"adapter": "discord"}'
|
||||
)
|
||||
services = storage.list_services("channel", max_age_seconds=120)
|
||||
assert services[0]["metadata"] == '{"adapter": "discord"}'
|
||||
+158
-6
@@ -1,9 +1,10 @@
|
||||
"""Tests for turnstone.core.session — ChatSession construction."""
|
||||
|
||||
import base64
|
||||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from turnstone.core.session import ChatSession
|
||||
from turnstone.core.session import _IMAGE_EXTENSIONS, _IMAGE_SIZE_CAP, ChatSession
|
||||
|
||||
|
||||
class NullUI:
|
||||
@@ -161,12 +162,12 @@ class TestPlanExec:
|
||||
|
||||
return call_id, content, captured.get("messages", [])
|
||||
|
||||
def test_plan_file_uses_session_id(self, tmp_db, tmp_path, monkeypatch):
|
||||
"""Plan file is named .plan-<session_id>.md, not .plan.md."""
|
||||
def test_plan_file_uses_ws_id(self, tmp_db, tmp_path, monkeypatch):
|
||||
"""Plan file is named .plan-<ws_id>.md, not .plan.md."""
|
||||
monkeypatch.chdir(tmp_path)
|
||||
session = _make_session()
|
||||
self._run_plan(session, "add feature")
|
||||
expected = tmp_path / f".plan-{session._session_id}.md"
|
||||
expected = tmp_path / f".plan-{session._ws_id}.md"
|
||||
assert expected.exists(), f"Expected {expected} to be created"
|
||||
assert not (tmp_path / ".plan.md").exists()
|
||||
|
||||
@@ -176,7 +177,7 @@ class TestPlanExec:
|
||||
session = _make_session()
|
||||
plan_content = "## Goal\n\nAdd a new endpoint."
|
||||
self._run_plan(session, "add endpoint", agent_return=plan_content)
|
||||
plan_file = tmp_path / f".plan-{session._session_id}.md"
|
||||
plan_file = tmp_path / f".plan-{session._ws_id}.md"
|
||||
assert plan_file.read_text() == plan_content
|
||||
|
||||
def test_two_sessions_produce_different_files(self, tmp_db, tmp_path, monkeypatch):
|
||||
@@ -184,7 +185,7 @@ class TestPlanExec:
|
||||
monkeypatch.chdir(tmp_path)
|
||||
s1 = _make_session()
|
||||
s2 = _make_session()
|
||||
assert s1._session_id != s2._session_id
|
||||
assert s1._ws_id != s2._ws_id
|
||||
self._run_plan(s1, "feature A")
|
||||
self._run_plan(s2, "feature B")
|
||||
files = list(tmp_path.glob(".plan-*.md"))
|
||||
@@ -265,3 +266,154 @@ class TestPlanExec:
|
||||
call_id, content, _ = self._run_plan(session, "do stuff", agent_return=agent_output)
|
||||
assert call_id == "test-call-1"
|
||||
assert content == agent_output
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Vision / image support
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestImageExtensions:
|
||||
"""Test _IMAGE_EXTENSIONS constant and detection logic."""
|
||||
|
||||
def test_common_image_extensions(self):
|
||||
for ext in (".png", ".jpg", ".jpeg", ".gif", ".webp", ".bmp", ".tiff", ".tif", ".ico"):
|
||||
assert ext in _IMAGE_EXTENSIONS, f"{ext} should be in _IMAGE_EXTENSIONS"
|
||||
|
||||
def test_svg_excluded(self):
|
||||
assert ".svg" not in _IMAGE_EXTENSIONS
|
||||
|
||||
def test_text_extensions_excluded(self):
|
||||
for ext in (".py", ".txt", ".json", ".md", ".rs", ".go"):
|
||||
assert ext not in _IMAGE_EXTENSIONS
|
||||
|
||||
|
||||
class TestExecReadImage:
|
||||
"""Test _exec_read_image method."""
|
||||
|
||||
def _make_png(self, path: str, size: int = 100) -> None:
|
||||
"""Write a minimal valid-ish PNG header to a file."""
|
||||
# 8-byte PNG signature + enough bytes to reach target size
|
||||
header = b"\x89PNG\r\n\x1a\n"
|
||||
with open(path, "wb") as f:
|
||||
f.write(header + b"\x00" * max(0, size - len(header)))
|
||||
|
||||
def test_image_returns_content_parts(self, tmp_db, tmp_path):
|
||||
"""read_file on a PNG with vision support returns content parts."""
|
||||
img = tmp_path / "test.png"
|
||||
self._make_png(str(img))
|
||||
|
||||
session = _make_session()
|
||||
# Mock provider to report vision support
|
||||
mock_caps = MagicMock()
|
||||
mock_caps.supports_vision = True
|
||||
session._provider.get_capabilities = MagicMock(return_value=mock_caps)
|
||||
|
||||
item = {"call_id": "c1", "path": str(img), "offset": None, "limit": None}
|
||||
call_id, output = session._exec_read_file(item)
|
||||
|
||||
assert call_id == "c1"
|
||||
assert isinstance(output, list)
|
||||
assert len(output) == 2
|
||||
assert output[0]["type"] == "text"
|
||||
assert "test.png" in output[0]["text"]
|
||||
assert output[1]["type"] == "image_url"
|
||||
url = output[1]["image_url"]["url"]
|
||||
assert url.startswith("data:image/png;base64,")
|
||||
# Verify base64 round-trip
|
||||
b64part = url.split(",", 1)[1]
|
||||
decoded = base64.b64decode(b64part)
|
||||
assert decoded == img.read_bytes()
|
||||
|
||||
def test_no_vision_returns_text(self, tmp_db, tmp_path):
|
||||
"""read_file on image with non-vision model returns text description."""
|
||||
img = tmp_path / "photo.jpg"
|
||||
self._make_png(str(img), size=2048)
|
||||
|
||||
session = _make_session()
|
||||
mock_caps = MagicMock()
|
||||
mock_caps.supports_vision = False
|
||||
session._provider.get_capabilities = MagicMock(return_value=mock_caps)
|
||||
|
||||
item = {"call_id": "c2", "path": str(img), "offset": None, "limit": None}
|
||||
call_id, output = session._exec_read_file(item)
|
||||
|
||||
assert call_id == "c2"
|
||||
assert isinstance(output, str)
|
||||
assert "does not support vision" in output
|
||||
assert "photo.jpg" in output
|
||||
|
||||
def test_oversized_image_returns_error(self, tmp_db, tmp_path):
|
||||
"""Images exceeding _IMAGE_SIZE_CAP return an error string."""
|
||||
img = tmp_path / "huge.png"
|
||||
# Write slightly over the cap
|
||||
with open(img, "wb") as f:
|
||||
f.write(b"\x89PNG\r\n\x1a\n" + b"\x00" * _IMAGE_SIZE_CAP)
|
||||
|
||||
session = _make_session()
|
||||
mock_caps = MagicMock()
|
||||
mock_caps.supports_vision = True
|
||||
session._provider.get_capabilities = MagicMock(return_value=mock_caps)
|
||||
|
||||
item = {"call_id": "c3", "path": str(img), "offset": None, "limit": None}
|
||||
call_id, output = session._exec_read_file(item)
|
||||
|
||||
assert call_id == "c3"
|
||||
assert isinstance(output, str)
|
||||
assert "exceeds" in output
|
||||
|
||||
def test_missing_image_returns_error(self, tmp_db, tmp_path):
|
||||
"""read_file on non-existent image returns error."""
|
||||
session = _make_session()
|
||||
mock_caps = MagicMock()
|
||||
mock_caps.supports_vision = True
|
||||
session._provider.get_capabilities = MagicMock(return_value=mock_caps)
|
||||
|
||||
item = {"call_id": "c4", "path": str(tmp_path / "nope.png"), "offset": None, "limit": None}
|
||||
call_id, output = session._exec_read_file(item)
|
||||
assert isinstance(output, str)
|
||||
assert "not found" in output
|
||||
|
||||
def test_svg_read_as_text(self, tmp_db, tmp_path):
|
||||
"""SVG files are read as text, not as images."""
|
||||
svg = tmp_path / "icon.svg"
|
||||
svg.write_text('<svg xmlns="http://www.w3.org/2000/svg"><circle r="10"/></svg>')
|
||||
|
||||
session = _make_session()
|
||||
item = {"call_id": "c5", "path": str(svg), "offset": None, "limit": None}
|
||||
call_id, output = session._exec_read_file(item)
|
||||
assert isinstance(output, str)
|
||||
assert "<svg" in output # Read as text
|
||||
|
||||
|
||||
class TestGetCapabilitiesOverride:
|
||||
"""Test _get_capabilities with config.toml overrides."""
|
||||
|
||||
def test_config_override_applies(self, tmp_db):
|
||||
"""capabilities dict from ModelConfig is merged onto provider caps."""
|
||||
from turnstone.core.model_registry import ModelConfig, ModelRegistry
|
||||
from turnstone.core.providers._protocol import ModelCapabilities
|
||||
|
||||
cfg = ModelConfig(
|
||||
alias="qwen-vl",
|
||||
base_url="http://localhost:8000/v1",
|
||||
api_key="dummy",
|
||||
model="qwen-3.5-vl",
|
||||
capabilities={"supports_vision": True},
|
||||
)
|
||||
registry = ModelRegistry(
|
||||
models={"qwen-vl": cfg},
|
||||
default="qwen-vl",
|
||||
)
|
||||
session = _make_session(registry=registry, model_alias="qwen-vl")
|
||||
# Ensure provider returns a real ModelCapabilities (not MagicMock)
|
||||
session._provider.get_capabilities = MagicMock(return_value=ModelCapabilities())
|
||||
caps = session._get_capabilities()
|
||||
assert caps.supports_vision is True
|
||||
|
||||
def test_no_override_uses_provider_default(self, tmp_db):
|
||||
"""Without config override, provider defaults are used."""
|
||||
session = _make_session()
|
||||
caps = session._get_capabilities()
|
||||
# Default OpenAI provider for unknown model → no vision
|
||||
assert caps.supports_vision is False
|
||||
|
||||
+178
-179
@@ -1,155 +1,155 @@
|
||||
"""Tests for session persistence and resume functionality."""
|
||||
"""Tests for workstream persistence and resume functionality."""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.memory import (
|
||||
delete_session,
|
||||
list_sessions,
|
||||
load_session_config,
|
||||
load_session_messages,
|
||||
prune_sessions,
|
||||
register_session,
|
||||
resolve_session,
|
||||
delete_workstream,
|
||||
list_workstreams_with_history,
|
||||
load_messages,
|
||||
load_workstream_config,
|
||||
prune_workstreams,
|
||||
register_workstream,
|
||||
resolve_workstream,
|
||||
save_message,
|
||||
save_session_config,
|
||||
set_session_alias,
|
||||
update_session_title,
|
||||
save_workstream_config,
|
||||
set_workstream_alias,
|
||||
update_workstream_title,
|
||||
)
|
||||
from turnstone.core.session import ChatSession
|
||||
from turnstone.core.storage import get_storage
|
||||
|
||||
# ── Session registration ──────────────────────────────────────────────
|
||||
# ── Workstream registration ───────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRegisterSession:
|
||||
class TestRegisterWorkstream:
|
||||
def test_register_creates_row(self, tmp_db):
|
||||
register_session("abc123")
|
||||
# Session exists in DB (resolve works) even without messages
|
||||
assert resolve_session("abc123") == "abc123"
|
||||
register_workstream("abc123")
|
||||
# Workstream exists in DB (resolve works) even without messages
|
||||
assert resolve_workstream("abc123") == "abc123"
|
||||
|
||||
def test_register_with_title(self, tmp_db):
|
||||
register_session("abc123", title="My Session")
|
||||
register_workstream("abc123", name="My Workstream")
|
||||
save_message("abc123", "user", "hello")
|
||||
rows = list_sessions()
|
||||
assert rows[0][2] == "My Session" # title
|
||||
rows = list_workstreams_with_history()
|
||||
assert rows[0][2] is None # title column (name is separate)
|
||||
|
||||
def test_register_idempotent(self, tmp_db):
|
||||
register_session("abc123", title="First")
|
||||
register_session("abc123", title="Second") # should be ignored
|
||||
register_workstream("abc123")
|
||||
update_workstream_title("abc123", "First")
|
||||
register_workstream("abc123") # should be ignored
|
||||
update_workstream_title("abc123", "First") # title is set via update
|
||||
save_message("abc123", "user", "hello")
|
||||
rows = list_sessions()
|
||||
rows = list_workstreams_with_history()
|
||||
assert len(rows) == 1
|
||||
assert rows[0][2] == "First" # original title preserved
|
||||
assert rows[0][2] == "First" # title preserved
|
||||
|
||||
def test_update_title(self, tmp_db):
|
||||
register_session("abc123")
|
||||
update_session_title("abc123", "New Title")
|
||||
register_workstream("abc123")
|
||||
update_workstream_title("abc123", "New Title")
|
||||
save_message("abc123", "user", "hello")
|
||||
rows = list_sessions()
|
||||
rows = list_workstreams_with_history()
|
||||
assert rows[0][2] == "New Title"
|
||||
|
||||
|
||||
# ── Session alias ─────────────────────────────────────────────────────
|
||||
# ── Workstream alias ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSessionAlias:
|
||||
class TestWorkstreamAlias:
|
||||
def test_set_alias(self, tmp_db):
|
||||
register_session("abc123")
|
||||
assert set_session_alias("abc123", "my-session") is True
|
||||
register_workstream("abc123")
|
||||
assert set_workstream_alias("abc123", "my-session") is True
|
||||
save_message("abc123", "user", "hello")
|
||||
rows = list_sessions()
|
||||
rows = list_workstreams_with_history()
|
||||
assert rows[0][1] == "my-session" # alias
|
||||
|
||||
def test_alias_conflict(self, tmp_db):
|
||||
register_session("abc123")
|
||||
register_session("def456")
|
||||
set_session_alias("abc123", "taken")
|
||||
assert set_session_alias("def456", "taken") is False
|
||||
register_workstream("abc123")
|
||||
register_workstream("def456")
|
||||
set_workstream_alias("abc123", "taken")
|
||||
assert set_workstream_alias("def456", "taken") is False
|
||||
|
||||
def test_alias_same_session_ok(self, tmp_db):
|
||||
register_session("abc123")
|
||||
set_session_alias("abc123", "mine")
|
||||
assert set_session_alias("abc123", "mine") is True # no-op, same session
|
||||
def test_alias_same_workstream_ok(self, tmp_db):
|
||||
register_workstream("abc123")
|
||||
set_workstream_alias("abc123", "mine")
|
||||
assert set_workstream_alias("abc123", "mine") is True # no-op, same workstream
|
||||
|
||||
|
||||
# ── Session resolution ────────────────────────────────────────────────
|
||||
# ── Workstream resolution ─────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestResolveSession:
|
||||
class TestResolveWorkstream:
|
||||
def test_resolve_by_alias(self, tmp_db):
|
||||
register_session("abc123")
|
||||
set_session_alias("abc123", "my-alias")
|
||||
assert resolve_session("my-alias") == "abc123"
|
||||
register_workstream("abc123")
|
||||
set_workstream_alias("abc123", "my-alias")
|
||||
assert resolve_workstream("my-alias") == "abc123"
|
||||
|
||||
def test_resolve_by_exact_id(self, tmp_db):
|
||||
register_session("abc123def456")
|
||||
assert resolve_session("abc123def456") == "abc123def456"
|
||||
register_workstream("abc123def456")
|
||||
assert resolve_workstream("abc123def456") == "abc123def456"
|
||||
|
||||
def test_resolve_by_prefix(self, tmp_db):
|
||||
register_session("abc123def456")
|
||||
assert resolve_session("abc123") == "abc123def456"
|
||||
register_workstream("abc123def456")
|
||||
assert resolve_workstream("abc123") == "abc123def456"
|
||||
|
||||
def test_resolve_prefix_ambiguous(self, tmp_db):
|
||||
register_session("abc123aaaaaa")
|
||||
register_session("abc123bbbbbb")
|
||||
register_workstream("abc123aaaaaa")
|
||||
register_workstream("abc123bbbbbb")
|
||||
# Ambiguous prefix should return None
|
||||
assert resolve_session("abc123") is None
|
||||
assert resolve_workstream("abc123") is None
|
||||
|
||||
def test_resolve_not_found(self, tmp_db):
|
||||
assert resolve_session("nonexistent") is None
|
||||
|
||||
def test_resolve_legacy_session(self, tmp_db):
|
||||
"""Sessions that exist only in conversations (pre-migration) should auto-register."""
|
||||
save_message("legacy123456", "user", "old message")
|
||||
result = resolve_session("legacy123456")
|
||||
assert result == "legacy123456"
|
||||
# Should now appear in sessions list
|
||||
rows = list_sessions()
|
||||
assert any(r[0] == "legacy123456" for r in rows)
|
||||
assert resolve_workstream("nonexistent") is None
|
||||
|
||||
|
||||
# ── List sessions ─────────────────────────────────────────────────────
|
||||
# ── List workstreams with history ──────────────────────────────────────
|
||||
|
||||
|
||||
class TestListSessions:
|
||||
class TestListWorkstreamsWithHistory:
|
||||
def test_empty(self, tmp_db):
|
||||
assert list_sessions() == []
|
||||
assert list_workstreams_with_history() == []
|
||||
|
||||
def test_ordered_by_updated(self, tmp_db):
|
||||
register_session("first")
|
||||
register_workstream("first")
|
||||
save_message("first", "user", "hello")
|
||||
register_session("second")
|
||||
# Force an older timestamp so ordering is deterministic
|
||||
engine = get_storage()._engine # noqa: SLF001
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.text("UPDATE workstreams SET updated = '2020-01-01' WHERE ws_id = 'first'")
|
||||
)
|
||||
conn.commit()
|
||||
register_workstream("second")
|
||||
save_message("second", "user", "hello")
|
||||
# second is more recent
|
||||
rows = list_sessions()
|
||||
rows = list_workstreams_with_history()
|
||||
assert rows[0][0] == "second"
|
||||
assert rows[1][0] == "first"
|
||||
|
||||
def test_includes_message_count(self, tmp_db):
|
||||
register_session("sess1")
|
||||
register_workstream("sess1")
|
||||
save_message("sess1", "user", "hello")
|
||||
save_message("sess1", "assistant", "hi")
|
||||
rows = list_sessions()
|
||||
rows = list_workstreams_with_history()
|
||||
assert rows[0][5] == 2 # msg_count
|
||||
|
||||
def test_respects_limit(self, tmp_db):
|
||||
for i in range(5):
|
||||
register_session(f"sess{i}")
|
||||
register_workstream(f"sess{i}")
|
||||
save_message(f"sess{i}", "user", "hello")
|
||||
rows = list_sessions(limit=3)
|
||||
rows = list_workstreams_with_history(limit=3)
|
||||
assert len(rows) == 3
|
||||
|
||||
|
||||
# ── Load session messages ─────────────────────────────────────────────
|
||||
# ── Load messages ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestLoadSessionMessages:
|
||||
class TestLoadMessages:
|
||||
def test_simple_user_assistant(self, tmp_db):
|
||||
save_message("s1", "user", "hello")
|
||||
save_message("s1", "assistant", "hi there")
|
||||
msgs = load_session_messages("s1")
|
||||
msgs = load_messages("s1")
|
||||
assert len(msgs) == 2
|
||||
assert msgs[0] == {"role": "user", "content": "hello"}
|
||||
assert msgs[1] == {"role": "assistant", "content": "hi there"}
|
||||
@@ -159,7 +159,7 @@ class TestLoadSessionMessages:
|
||||
save_message("s1", "assistant", "Let me check.")
|
||||
save_message("s1", "tool_call", None, "bash", '{"command":"ls"}', tool_call_id="call_abc")
|
||||
save_message("s1", "tool_result", "file1.txt\nfile2.txt", "bash", tool_call_id="call_abc")
|
||||
msgs = load_session_messages("s1")
|
||||
msgs = load_messages("s1")
|
||||
assert len(msgs) == 3 # user, assistant+tool_calls, tool
|
||||
# Assistant should have content merged with tool_calls
|
||||
assert msgs[1]["role"] == "assistant"
|
||||
@@ -177,7 +177,7 @@ class TestLoadSessionMessages:
|
||||
save_message("s1", "user", "do stuff")
|
||||
save_message("s1", "tool_call", None, "bash", '{"command":"ls"}')
|
||||
save_message("s1", "tool_result", "output", "bash")
|
||||
msgs = load_session_messages("s1")
|
||||
msgs = load_messages("s1")
|
||||
assert len(msgs) == 3
|
||||
# Synthetic IDs should match
|
||||
tc_id = msgs[1]["tool_calls"][0]["id"]
|
||||
@@ -189,36 +189,36 @@ class TestLoadSessionMessages:
|
||||
save_message("s1", "tool_call", None, "search", '{"query":"b"}', tool_call_id="call_2")
|
||||
save_message("s1", "tool_result", "result a", "search", tool_call_id="call_1")
|
||||
save_message("s1", "tool_result", "result b", "search", tool_call_id="call_2")
|
||||
msgs = load_session_messages("s1")
|
||||
msgs = load_messages("s1")
|
||||
assert len(msgs) == 4 # user, assistant+2 tool_calls, 2 tool results
|
||||
assert len(msgs[1]["tool_calls"]) == 2
|
||||
assert msgs[2]["tool_call_id"] == "call_1"
|
||||
assert msgs[3]["tool_call_id"] == "call_2"
|
||||
|
||||
def test_empty_session(self, tmp_db):
|
||||
assert load_session_messages("nonexistent") == []
|
||||
def test_empty_workstream(self, tmp_db):
|
||||
assert load_messages("nonexistent") == []
|
||||
|
||||
def test_orphaned_tool_result_skipped(self, tmp_db):
|
||||
save_message("s1", "user", "hello")
|
||||
save_message("s1", "tool_result", "orphan", "bash")
|
||||
msgs = load_session_messages("s1")
|
||||
msgs = load_messages("s1")
|
||||
assert len(msgs) == 1 # only the user message
|
||||
|
||||
|
||||
# ── Delete session ────────────────────────────────────────────────────
|
||||
# ── Delete workstream ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestDeleteSession:
|
||||
def test_delete_removes_session_and_messages(self, tmp_db):
|
||||
register_session("abc123")
|
||||
class TestDeleteWorkstream:
|
||||
def test_delete_removes_workstream_and_messages(self, tmp_db):
|
||||
register_workstream("abc123")
|
||||
save_message("abc123", "user", "hello")
|
||||
save_message("abc123", "assistant", "hi")
|
||||
assert delete_session("abc123") is True
|
||||
assert list_sessions() == []
|
||||
assert load_session_messages("abc123") == []
|
||||
assert delete_workstream("abc123") is True
|
||||
assert list_workstreams_with_history() == []
|
||||
assert load_messages("abc123") == []
|
||||
|
||||
def test_delete_nonexistent(self, tmp_db):
|
||||
assert delete_session("nonexistent") is True # no-op, still returns True
|
||||
assert delete_workstream("nonexistent") is False
|
||||
|
||||
|
||||
# ── save_message with tool_call_id ────────────────────────────────────
|
||||
@@ -230,7 +230,7 @@ class TestSaveMessageToolCallId:
|
||||
engine = get_storage()._engine # noqa: SLF001
|
||||
with engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
sa.text("SELECT tool_call_id FROM conversations WHERE session_id = 's1'")
|
||||
sa.text("SELECT tool_call_id FROM conversations WHERE ws_id = 's1'")
|
||||
).fetchone()
|
||||
assert row[0] == "call_xyz"
|
||||
|
||||
@@ -239,20 +239,20 @@ class TestSaveMessageToolCallId:
|
||||
engine = get_storage()._engine # noqa: SLF001
|
||||
with engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
sa.text("SELECT tool_call_id FROM conversations WHERE session_id = 's1'")
|
||||
sa.text("SELECT tool_call_id FROM conversations WHERE ws_id = 's1'")
|
||||
).fetchone()
|
||||
assert row[0] is None
|
||||
|
||||
|
||||
# ── Sessions table creation ───────────────────────────────────────────
|
||||
# ── Workstreams table creation ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSessionsTable:
|
||||
def test_sessions_table_exists(self, tmp_db):
|
||||
class TestWorkstreamsTable:
|
||||
def test_workstreams_table_exists(self, tmp_db):
|
||||
engine = get_storage()._engine # noqa: SLF001
|
||||
with engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
sa.text("SELECT name FROM sqlite_master WHERE type='table' AND name='sessions'")
|
||||
sa.text("SELECT name FROM sqlite_master WHERE type='table' AND name='workstreams'")
|
||||
).fetchall()
|
||||
assert len(rows) == 1
|
||||
|
||||
@@ -263,15 +263,15 @@ class TestSessionsTable:
|
||||
conn.execute(sa.text("SELECT tool_call_id FROM conversations LIMIT 0"))
|
||||
|
||||
|
||||
# ── ChatSession.resume_session ────────────────────────────────────────
|
||||
# ── ChatSession.resume ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestResumeSession:
|
||||
class TestResumeWorkstream:
|
||||
def test_resume_loads_messages(self, tmp_db, mock_openai_client):
|
||||
# Set up a session with messages in DB
|
||||
register_session("old_sess_123")
|
||||
save_message("old_sess_123", "user", "hello world")
|
||||
save_message("old_sess_123", "assistant", "hi there")
|
||||
# Set up a workstream with messages in DB
|
||||
register_workstream("old_ws_123")
|
||||
save_message("old_ws_123", "user", "hello world")
|
||||
save_message("old_ws_123", "assistant", "hi there")
|
||||
|
||||
# Create a new session and resume
|
||||
session = ChatSession(
|
||||
@@ -283,12 +283,12 @@ class TestResumeSession:
|
||||
max_tokens=1000,
|
||||
tool_timeout=10,
|
||||
)
|
||||
original_id = session._session_id
|
||||
assert original_id != "old_sess_123"
|
||||
original_id = session._ws_id
|
||||
assert original_id != "old_ws_123"
|
||||
|
||||
result = session.resume_session("old_sess_123")
|
||||
result = session.resume("old_ws_123")
|
||||
assert result is True
|
||||
assert session._session_id == "old_sess_123"
|
||||
assert session._ws_id == "old_ws_123"
|
||||
assert len(session.messages) == 2
|
||||
assert session.messages[0]["content"] == "hello world"
|
||||
assert session._title_generated is True
|
||||
@@ -303,9 +303,9 @@ class TestResumeSession:
|
||||
max_tokens=1000,
|
||||
tool_timeout=10,
|
||||
)
|
||||
assert session.resume_session("nonexistent") is False
|
||||
assert session.resume("nonexistent") is False
|
||||
|
||||
def test_session_registered_on_init(self, tmp_db, mock_openai_client):
|
||||
def test_workstream_not_registered_until_message(self, tmp_db, mock_openai_client):
|
||||
session = ChatSession(
|
||||
client=mock_openai_client,
|
||||
model="test-model",
|
||||
@@ -315,20 +315,19 @@ class TestResumeSession:
|
||||
max_tokens=1000,
|
||||
tool_timeout=10,
|
||||
)
|
||||
# Session is registered in DB (resolvable) even before any messages
|
||||
assert resolve_session(session._session_id) == session._session_id
|
||||
# But does not appear in list_sessions until a message is saved
|
||||
assert not any(r[0] == session._session_id for r in list_sessions())
|
||||
# Workstream is not auto-registered on init — only on /new or server creation
|
||||
assert resolve_workstream(session._ws_id) is None
|
||||
assert not any(r[0] == session._ws_id for r in list_workstreams_with_history())
|
||||
|
||||
|
||||
# ── save_message updates sessions.updated ─────────────────────────────
|
||||
# ── save_message updates workstreams.updated ──────────────────────────
|
||||
|
||||
|
||||
class TestSaveMessageUpdatesSession:
|
||||
class TestSaveMessageUpdatesWorkstream:
|
||||
def test_updated_timestamp_bumped(self, tmp_db):
|
||||
register_session("s1")
|
||||
register_workstream("s1")
|
||||
save_message("s1", "user", "first")
|
||||
rows = list_sessions()
|
||||
rows = list_workstreams_with_history()
|
||||
_original_updated = rows[0][4]
|
||||
|
||||
import time
|
||||
@@ -336,18 +335,18 @@ class TestSaveMessageUpdatesSession:
|
||||
time.sleep(0.01) # ensure different timestamp
|
||||
save_message("s1", "user", "hello")
|
||||
|
||||
rows = list_sessions()
|
||||
rows = list_workstreams_with_history()
|
||||
new_updated = rows[0][4]
|
||||
# updated should be same or later (sqlite datetime resolution is seconds,
|
||||
# so they may be equal in fast tests — just verify no error)
|
||||
assert new_updated is not None
|
||||
|
||||
|
||||
# ── Interrupted session repair ───────────────────────────────────────
|
||||
# ── Interrupted workstream repair ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestInterruptedSessionRepair:
|
||||
"""load_session_messages() should strip trailing incomplete tool call turns."""
|
||||
class TestInterruptedWorkstreamRepair:
|
||||
"""load_messages() should strip trailing incomplete tool call turns."""
|
||||
|
||||
def test_complete_tool_turn_preserved(self, tmp_db):
|
||||
"""2 tool_calls + 2 tool_results = complete, no stripping."""
|
||||
@@ -356,7 +355,7 @@ class TestInterruptedSessionRepair:
|
||||
save_message("s1", "tool_call", None, "bash", '{"command":"pwd"}', "call_2")
|
||||
save_message("s1", "tool_result", "file.txt", tool_call_id="call_1")
|
||||
save_message("s1", "tool_result", "/home", tool_call_id="call_2")
|
||||
msgs = load_session_messages("s1")
|
||||
msgs = load_messages("s1")
|
||||
assert len(msgs) == 4 # user + assistant(2 calls) + 2 tool results
|
||||
|
||||
def test_partial_tool_results_stripped(self, tmp_db):
|
||||
@@ -365,7 +364,7 @@ class TestInterruptedSessionRepair:
|
||||
save_message("s1", "tool_call", None, "bash", '{"command":"ls"}', "call_1")
|
||||
save_message("s1", "tool_call", None, "bash", '{"command":"pwd"}', "call_2")
|
||||
save_message("s1", "tool_result", "file.txt", tool_call_id="call_1")
|
||||
msgs = load_session_messages("s1")
|
||||
msgs = load_messages("s1")
|
||||
assert len(msgs) == 1 # only user message remains
|
||||
assert msgs[0]["role"] == "user"
|
||||
|
||||
@@ -375,7 +374,7 @@ class TestInterruptedSessionRepair:
|
||||
save_message("s1", "assistant", "Let me check")
|
||||
save_message("s1", "tool_call", None, "bash", '{"command":"ls"}', "call_1")
|
||||
save_message("s1", "tool_call", None, "bash", '{"command":"pwd"}', "call_2")
|
||||
msgs = load_session_messages("s1")
|
||||
msgs = load_messages("s1")
|
||||
# assistant with content was merged into tool_call assistant, so stripped
|
||||
assert len(msgs) == 1
|
||||
assert msgs[0]["role"] == "user"
|
||||
@@ -386,42 +385,42 @@ class TestInterruptedSessionRepair:
|
||||
save_message("s1", "assistant", "response")
|
||||
save_message("s1", "user", "second")
|
||||
save_message("s1", "tool_call", None, "bash", '{"command":"ls"}', "call_1")
|
||||
msgs = load_session_messages("s1")
|
||||
msgs = load_messages("s1")
|
||||
assert len(msgs) == 3 # user + assistant + user (incomplete turn stripped)
|
||||
assert msgs[0]["role"] == "user"
|
||||
assert msgs[1]["role"] == "assistant"
|
||||
assert msgs[2]["role"] == "user"
|
||||
|
||||
|
||||
# ── Session config persistence ───────────────────────────────────────
|
||||
# ── Workstream config persistence ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestSessionConfig:
|
||||
class TestWorkstreamConfig:
|
||||
def test_save_load_roundtrip(self, tmp_db):
|
||||
config = {"temperature": "0.3", "reasoning_effort": "high", "creative_mode": "False"}
|
||||
save_session_config("s1", config)
|
||||
loaded = load_session_config("s1")
|
||||
save_workstream_config("s1", config)
|
||||
loaded = load_workstream_config("s1")
|
||||
assert loaded == config
|
||||
|
||||
def test_update_existing_key(self, tmp_db):
|
||||
save_session_config("s1", {"temperature": "0.3"})
|
||||
save_session_config("s1", {"temperature": "0.7"})
|
||||
loaded = load_session_config("s1")
|
||||
save_workstream_config("s1", {"temperature": "0.3"})
|
||||
save_workstream_config("s1", {"temperature": "0.7"})
|
||||
loaded = load_workstream_config("s1")
|
||||
assert loaded["temperature"] == "0.7"
|
||||
|
||||
def test_missing_session_returns_empty(self, tmp_db):
|
||||
loaded = load_session_config("nonexistent")
|
||||
def test_missing_workstream_returns_empty(self, tmp_db):
|
||||
loaded = load_workstream_config("nonexistent")
|
||||
assert loaded == {}
|
||||
|
||||
def test_delete_session_removes_config(self, tmp_db):
|
||||
register_session("s1")
|
||||
def test_delete_workstream_removes_config(self, tmp_db):
|
||||
register_workstream("s1")
|
||||
save_message("s1", "user", "hi")
|
||||
save_session_config("s1", {"temperature": "0.5"})
|
||||
delete_session("s1")
|
||||
assert load_session_config("s1") == {}
|
||||
save_workstream_config("s1", {"temperature": "0.5"})
|
||||
delete_workstream("s1")
|
||||
assert load_workstream_config("s1") == {}
|
||||
|
||||
def test_resume_restores_config(self, tmp_db):
|
||||
"""ChatSession.resume_session() should restore persisted config."""
|
||||
"""ChatSession.resume() should restore persisted config."""
|
||||
client = MagicMock()
|
||||
client.models.list.return_value.data = [MagicMock(id="test-model")]
|
||||
ui = MagicMock()
|
||||
@@ -430,11 +429,11 @@ class TestSessionConfig:
|
||||
ui.on_state_change = MagicMock()
|
||||
ui.on_rename = MagicMock()
|
||||
|
||||
# Create a session with specific config
|
||||
register_session("orig")
|
||||
# Create a workstream with specific config
|
||||
register_workstream("orig")
|
||||
save_message("orig", "user", "hello")
|
||||
save_message("orig", "assistant", "hi there")
|
||||
save_session_config(
|
||||
save_workstream_config(
|
||||
"orig",
|
||||
{
|
||||
"temperature": "0.3",
|
||||
@@ -456,7 +455,7 @@ class TestSessionConfig:
|
||||
tool_timeout=30,
|
||||
)
|
||||
assert session.temperature == 0.7 # default
|
||||
result = session.resume_session("orig")
|
||||
result = session.resume("orig")
|
||||
assert result is True
|
||||
assert session.temperature == 0.3
|
||||
assert session.reasoning_effort == "high"
|
||||
@@ -465,84 +464,84 @@ class TestSessionConfig:
|
||||
assert session.creative_mode is True
|
||||
|
||||
|
||||
# ── Prune sessions ───────────────────────────────────────────────────
|
||||
# ── Prune workstreams ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPruneSessions:
|
||||
class TestPruneWorkstreams:
|
||||
def test_orphan_removed(self, tmp_db):
|
||||
"""Session registered with no messages should be pruned."""
|
||||
register_session("orphan")
|
||||
orphans, stale = prune_sessions()
|
||||
"""Workstream registered with no messages should be pruned."""
|
||||
register_workstream("orphan")
|
||||
orphans, stale = prune_workstreams()
|
||||
assert orphans == 1
|
||||
assert list_sessions() == []
|
||||
assert list_workstreams_with_history() == []
|
||||
|
||||
def test_session_with_messages_kept(self, tmp_db):
|
||||
"""Session with messages should not be pruned."""
|
||||
register_session("active")
|
||||
def test_workstream_with_messages_kept(self, tmp_db):
|
||||
"""Workstream with messages should not be pruned."""
|
||||
register_workstream("active")
|
||||
save_message("active", "user", "hello")
|
||||
orphans, _stale = prune_sessions()
|
||||
orphans, _stale = prune_workstreams()
|
||||
assert orphans == 0
|
||||
assert len(list_sessions()) == 1
|
||||
assert len(list_workstreams_with_history()) == 1
|
||||
|
||||
def test_stale_unnamed_removed(self, tmp_db):
|
||||
"""Old unnamed session should be pruned by retention policy."""
|
||||
register_session("old1")
|
||||
"""Old unnamed workstream should be pruned by retention policy."""
|
||||
register_workstream("old1")
|
||||
save_message("old1", "user", "ancient message")
|
||||
# Force the updated timestamp to the past so it looks stale
|
||||
engine = get_storage()._engine # noqa: SLF001
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.text("UPDATE sessions SET updated = '2020-01-01' WHERE session_id = 'old1'")
|
||||
sa.text("UPDATE workstreams SET updated = '2020-01-01' WHERE ws_id = 'old1'")
|
||||
)
|
||||
conn.commit()
|
||||
_orphans, stale = prune_sessions(retention_days=30)
|
||||
_orphans, stale = prune_workstreams(retention_days=30)
|
||||
assert stale == 1
|
||||
|
||||
def test_named_session_preserved(self, tmp_db):
|
||||
"""Session with alias should be kept regardless of age."""
|
||||
register_session("old2")
|
||||
set_session_alias("old2", "important")
|
||||
def test_named_workstream_preserved(self, tmp_db):
|
||||
"""Workstream with alias should be kept regardless of age."""
|
||||
register_workstream("old2")
|
||||
set_workstream_alias("old2", "important")
|
||||
save_message("old2", "user", "old but named")
|
||||
# Force old timestamp
|
||||
engine = get_storage()._engine # noqa: SLF001
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.text("UPDATE sessions SET updated = '2020-01-01' WHERE session_id = 'old2'")
|
||||
sa.text("UPDATE workstreams SET updated = '2020-01-01' WHERE ws_id = 'old2'")
|
||||
)
|
||||
conn.commit()
|
||||
_orphans, stale = prune_sessions(retention_days=30)
|
||||
_orphans, stale = prune_workstreams(retention_days=30)
|
||||
assert stale == 0
|
||||
assert len(list_sessions()) == 1
|
||||
assert len(list_workstreams_with_history()) == 1
|
||||
|
||||
def test_fresh_unnamed_preserved(self, tmp_db):
|
||||
"""Recent unnamed session should not be pruned."""
|
||||
register_session("fresh")
|
||||
"""Recent unnamed workstream should not be pruned."""
|
||||
register_workstream("fresh")
|
||||
save_message("fresh", "user", "just now")
|
||||
_orphans, stale = prune_sessions(retention_days=30)
|
||||
_orphans, stale = prune_workstreams(retention_days=30)
|
||||
assert stale == 0
|
||||
assert len(list_sessions()) == 1
|
||||
assert len(list_workstreams_with_history()) == 1
|
||||
|
||||
def test_prune_removes_session_config(self, tmp_db):
|
||||
"""Pruning orphan/stale sessions should also remove their config rows."""
|
||||
register_session("orphan_cfg")
|
||||
save_session_config("orphan_cfg", {"temperature": "0.5"})
|
||||
def test_prune_removes_workstream_config(self, tmp_db):
|
||||
"""Pruning orphan/stale workstreams should also remove their config rows."""
|
||||
register_workstream("orphan_cfg")
|
||||
save_workstream_config("orphan_cfg", {"temperature": "0.5"})
|
||||
|
||||
register_session("stale_cfg")
|
||||
register_workstream("stale_cfg")
|
||||
save_message("stale_cfg", "user", "old")
|
||||
save_session_config("stale_cfg", {"temperature": "0.9"})
|
||||
save_workstream_config("stale_cfg", {"temperature": "0.9"})
|
||||
engine = get_storage()._engine # noqa: SLF001
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.text("UPDATE sessions SET updated = '2020-01-01' WHERE session_id = 'stale_cfg'")
|
||||
sa.text("UPDATE workstreams SET updated = '2020-01-01' WHERE ws_id = 'stale_cfg'")
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
# Both should have config before prune
|
||||
assert load_session_config("orphan_cfg") == {"temperature": "0.5"}
|
||||
assert load_session_config("stale_cfg") == {"temperature": "0.9"}
|
||||
assert load_workstream_config("orphan_cfg") == {"temperature": "0.5"}
|
||||
assert load_workstream_config("stale_cfg") == {"temperature": "0.9"}
|
||||
|
||||
prune_sessions(retention_days=30)
|
||||
prune_workstreams(retention_days=30)
|
||||
|
||||
# Config rows should be cleaned up
|
||||
assert load_session_config("orphan_cfg") == {}
|
||||
assert load_session_config("stale_cfg") == {}
|
||||
assert load_workstream_config("orphan_cfg") == {}
|
||||
assert load_workstream_config("stale_cfg") == {}
|
||||
|
||||
+125
-71
@@ -14,28 +14,28 @@ def backend(tmp_path):
|
||||
reset_storage()
|
||||
|
||||
|
||||
# -- Session operations --------------------------------------------------------
|
||||
# -- Workstream registration ---------------------------------------------------
|
||||
|
||||
|
||||
class TestRegisterSession:
|
||||
def test_register_creates_session(self, backend):
|
||||
backend.register_session("s1", title="Test")
|
||||
name = backend.get_session_name("s1")
|
||||
class TestRegisterWorkstream:
|
||||
def test_register_creates_workstream(self, backend):
|
||||
backend.register_workstream("s1", title="Test")
|
||||
name = backend.get_workstream_display_name("s1")
|
||||
assert name == "Test"
|
||||
|
||||
def test_register_idempotent(self, backend):
|
||||
backend.register_session("s1", title="First")
|
||||
backend.register_session("s1", title="Second")
|
||||
name = backend.get_session_name("s1")
|
||||
backend.register_workstream("s1", title="First")
|
||||
backend.register_workstream("s1", title="Second")
|
||||
name = backend.get_workstream_display_name("s1")
|
||||
assert name == "First" # INSERT OR IGNORE preserves first
|
||||
|
||||
|
||||
class TestSaveAndLoadMessages:
|
||||
def test_roundtrip(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message("s1", "user", "hello")
|
||||
backend.save_message("s1", "assistant", "world")
|
||||
msgs = backend.load_session_messages("s1")
|
||||
msgs = backend.load_messages("s1")
|
||||
assert len(msgs) == 2
|
||||
assert msgs[0]["role"] == "user"
|
||||
assert msgs[0]["content"] == "hello"
|
||||
@@ -43,12 +43,12 @@ class TestSaveAndLoadMessages:
|
||||
assert msgs[1]["content"] == "world"
|
||||
|
||||
def test_tool_call_grouping(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message("s1", "user", "do something")
|
||||
backend.save_message("s1", "tool_call", None, "bash", '{"cmd":"ls"}', tool_call_id="c1")
|
||||
backend.save_message("s1", "tool_result", "file.txt", tool_call_id="c1")
|
||||
backend.save_message("s1", "assistant", "done")
|
||||
msgs = backend.load_session_messages("s1")
|
||||
msgs = backend.load_messages("s1")
|
||||
assert len(msgs) == 4
|
||||
assert msgs[1]["role"] == "assistant"
|
||||
assert len(msgs[1]["tool_calls"]) == 1
|
||||
@@ -57,136 +57,136 @@ class TestSaveAndLoadMessages:
|
||||
assert msgs[2]["content"] == "file.txt"
|
||||
|
||||
def test_incomplete_turn_repair(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message("s1", "user", "do something")
|
||||
backend.save_message("s1", "tool_call", None, "bash", '{"cmd":"ls"}', tool_call_id="c1")
|
||||
backend.save_message("s1", "tool_call", None, "read", '{"path":"a"}', tool_call_id="c2")
|
||||
# Only 1 result for 2 calls — incomplete turn
|
||||
backend.save_message("s1", "tool_result", "ok", tool_call_id="c1")
|
||||
msgs = backend.load_session_messages("s1")
|
||||
msgs = backend.load_messages("s1")
|
||||
# Incomplete turn should be stripped
|
||||
assert len(msgs) == 1 # only the user message remains
|
||||
|
||||
def test_provider_data_preserved(self, backend):
|
||||
import json
|
||||
|
||||
backend.register_session("s1")
|
||||
backend.register_workstream("s1")
|
||||
pd = json.dumps({"encrypted": True})
|
||||
backend.save_message("s1", "assistant", "hi", provider_data=pd)
|
||||
msgs = backend.load_session_messages("s1")
|
||||
msgs = backend.load_messages("s1")
|
||||
assert msgs[0].get("_provider_content") == {"encrypted": True}
|
||||
|
||||
def test_empty_session_returns_empty(self, backend):
|
||||
assert backend.load_session_messages("nonexistent") == []
|
||||
def test_empty_workstream_returns_empty(self, backend):
|
||||
assert backend.load_messages("nonexistent") == []
|
||||
|
||||
|
||||
class TestListSessions:
|
||||
def test_lists_sessions_with_messages(self, backend):
|
||||
backend.register_session("s1")
|
||||
class TestListWorkstreamsWithHistory:
|
||||
def test_lists_workstreams_with_messages(self, backend):
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message("s1", "user", "hi")
|
||||
backend.register_session("s2") # no messages
|
||||
rows = backend.list_sessions()
|
||||
backend.register_workstream("s2") # no messages
|
||||
rows = backend.list_workstreams_with_history()
|
||||
assert len(rows) == 1
|
||||
assert rows[0][0] == "s1"
|
||||
|
||||
def test_respects_limit(self, backend):
|
||||
for i in range(5):
|
||||
sid = f"s{i}"
|
||||
backend.register_session(sid)
|
||||
backend.register_workstream(sid)
|
||||
backend.save_message(sid, "user", f"msg {i}")
|
||||
rows = backend.list_sessions(limit=3)
|
||||
rows = backend.list_workstreams_with_history(limit=3)
|
||||
assert len(rows) == 3
|
||||
|
||||
|
||||
class TestDeleteSession:
|
||||
class TestDeleteWorkstream:
|
||||
def test_deletes_all_data(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message("s1", "user", "hi")
|
||||
backend.save_session_config("s1", {"temp": "0.5"})
|
||||
assert backend.delete_session("s1")
|
||||
assert backend.load_session_messages("s1") == []
|
||||
assert backend.load_session_config("s1") == {}
|
||||
assert backend.get_session_name("s1") is None
|
||||
backend.save_workstream_config("s1", {"temp": "0.5"})
|
||||
assert backend.delete_workstream("s1")
|
||||
assert backend.load_messages("s1") == []
|
||||
assert backend.load_workstream_config("s1") == {}
|
||||
assert backend.get_workstream_display_name("s1") is None
|
||||
|
||||
|
||||
class TestPruneSessions:
|
||||
class TestPruneWorkstreams:
|
||||
def test_orphan_removed(self, backend):
|
||||
backend.register_session("orphan")
|
||||
orphans, stale = backend.prune_sessions()
|
||||
backend.register_workstream("orphan")
|
||||
orphans, stale = backend.prune_workstreams()
|
||||
assert orphans == 1
|
||||
|
||||
def test_stale_removed(self, backend):
|
||||
import sqlalchemy as sa
|
||||
|
||||
backend.register_session("old")
|
||||
backend.register_workstream("old")
|
||||
backend.save_message("old", "user", "hi")
|
||||
# Force old timestamp
|
||||
with backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.text("UPDATE sessions SET updated = '2020-01-01' WHERE session_id = 'old'")
|
||||
sa.text("UPDATE workstreams SET updated = '2020-01-01' WHERE ws_id = 'old'")
|
||||
)
|
||||
conn.commit()
|
||||
_, stale = backend.prune_sessions(retention_days=30)
|
||||
_, stale = backend.prune_workstreams(retention_days=30)
|
||||
assert stale == 1
|
||||
|
||||
|
||||
class TestResolveSession:
|
||||
class TestResolveWorkstream:
|
||||
def test_exact_alias(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.set_session_alias("s1", "myalias")
|
||||
assert backend.resolve_session("myalias") == "s1"
|
||||
backend.register_workstream("s1")
|
||||
backend.set_workstream_alias("s1", "myalias")
|
||||
assert backend.resolve_workstream("myalias") == "s1"
|
||||
|
||||
def test_exact_id(self, backend):
|
||||
backend.register_session("abc-123-def")
|
||||
assert backend.resolve_session("abc-123-def") == "abc-123-def"
|
||||
backend.register_workstream("abc-123-def")
|
||||
assert backend.resolve_workstream("abc-123-def") == "abc-123-def"
|
||||
|
||||
def test_prefix_match(self, backend):
|
||||
backend.register_session("abc-123-def")
|
||||
assert backend.resolve_session("abc") == "abc-123-def"
|
||||
backend.register_workstream("abc-123-def")
|
||||
assert backend.resolve_workstream("abc") == "abc-123-def"
|
||||
|
||||
def test_not_found(self, backend):
|
||||
assert backend.resolve_session("nonexistent") is None
|
||||
assert backend.resolve_workstream("nonexistent") is None
|
||||
|
||||
|
||||
# -- Session config ------------------------------------------------------------
|
||||
# -- Workstream config ---------------------------------------------------------
|
||||
|
||||
|
||||
class TestSessionConfig:
|
||||
class TestWorkstreamConfig:
|
||||
def test_roundtrip(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.save_session_config("s1", {"temperature": "0.7", "effort": "high"})
|
||||
cfg = backend.load_session_config("s1")
|
||||
backend.register_workstream("s1")
|
||||
backend.save_workstream_config("s1", {"temperature": "0.7", "effort": "high"})
|
||||
cfg = backend.load_workstream_config("s1")
|
||||
assert cfg == {"temperature": "0.7", "effort": "high"}
|
||||
|
||||
def test_empty_config(self, backend):
|
||||
assert backend.load_session_config("nonexistent") == {}
|
||||
assert backend.load_workstream_config("nonexistent") == {}
|
||||
|
||||
|
||||
# -- Session metadata ----------------------------------------------------------
|
||||
# -- Workstream metadata ------------------------------------------------------
|
||||
|
||||
|
||||
class TestSessionMetadata:
|
||||
class TestWorkstreamMetadata:
|
||||
def test_alias(self, backend):
|
||||
backend.register_session("s1")
|
||||
assert backend.set_session_alias("s1", "my-session")
|
||||
assert backend.get_session_name("s1") == "my-session"
|
||||
backend.register_workstream("s1")
|
||||
assert backend.set_workstream_alias("s1", "my-session")
|
||||
assert backend.get_workstream_display_name("s1") == "my-session"
|
||||
|
||||
def test_alias_conflict(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.register_session("s2")
|
||||
backend.set_session_alias("s1", "taken")
|
||||
assert not backend.set_session_alias("s2", "taken")
|
||||
backend.register_workstream("s1")
|
||||
backend.register_workstream("s2")
|
||||
backend.set_workstream_alias("s1", "taken")
|
||||
assert not backend.set_workstream_alias("s2", "taken")
|
||||
|
||||
def test_title(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.update_session_title("s1", "My Title")
|
||||
assert backend.get_session_name("s1") == "My Title"
|
||||
backend.register_workstream("s1")
|
||||
backend.update_workstream_title("s1", "My Title")
|
||||
assert backend.get_workstream_display_name("s1") == "My Title"
|
||||
|
||||
def test_alias_preferred_over_title(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.update_session_title("s1", "Title")
|
||||
backend.set_session_alias("s1", "Alias")
|
||||
assert backend.get_session_name("s1") == "Alias"
|
||||
backend.register_workstream("s1")
|
||||
backend.update_workstream_title("s1", "Title")
|
||||
backend.set_workstream_alias("s1", "Alias")
|
||||
assert backend.get_workstream_display_name("s1") == "Alias"
|
||||
|
||||
|
||||
# -- Key-value store -----------------------------------------------------------
|
||||
@@ -234,7 +234,7 @@ class TestKVStore:
|
||||
|
||||
class TestSearch:
|
||||
def test_search_history(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message("s1", "user", "hello world")
|
||||
backend.save_message("s1", "user", "goodbye world")
|
||||
results = backend.search_history("hello")
|
||||
@@ -242,13 +242,67 @@ class TestSearch:
|
||||
assert any("hello" in str(r[3]) for r in results)
|
||||
|
||||
def test_search_recent(self, backend):
|
||||
backend.register_session("s1")
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message("s1", "user", "msg1")
|
||||
backend.save_message("s1", "user", "msg2")
|
||||
results = backend.search_history_recent(limit=1)
|
||||
assert len(results) == 1
|
||||
|
||||
|
||||
# -- Workstream operations -----------------------------------------------------
|
||||
|
||||
|
||||
class TestWorkstreams:
|
||||
def test_register_and_list(self, backend):
|
||||
backend.register_workstream("ws1", node_id="node-a", name="first")
|
||||
backend.register_workstream("ws2", node_id="node-a", name="second")
|
||||
rows = backend.list_workstreams()
|
||||
assert len(rows) == 2
|
||||
ws_ids = {r[0] for r in rows}
|
||||
assert ws_ids == {"ws1", "ws2"}
|
||||
|
||||
def test_register_idempotent(self, backend):
|
||||
backend.register_workstream("ws1", name="first")
|
||||
backend.register_workstream("ws1", name="overwrite")
|
||||
rows = backend.list_workstreams()
|
||||
assert len(rows) == 1
|
||||
assert rows[0][2] == "first" # name preserved from first insert
|
||||
|
||||
def test_update_state(self, backend):
|
||||
backend.register_workstream("ws1")
|
||||
backend.update_workstream_state("ws1", "running")
|
||||
rows = backend.list_workstreams()
|
||||
assert rows[0][3] == "running"
|
||||
|
||||
def test_update_name(self, backend):
|
||||
backend.register_workstream("ws1", name="old")
|
||||
backend.update_workstream_name("ws1", "new")
|
||||
rows = backend.list_workstreams()
|
||||
assert rows[0][2] == "new"
|
||||
|
||||
def test_delete(self, backend):
|
||||
backend.register_workstream("ws1")
|
||||
assert backend.delete_workstream("ws1") is True
|
||||
assert backend.list_workstreams() == []
|
||||
assert backend.delete_workstream("ws1") is False
|
||||
|
||||
def test_list_by_node(self, backend):
|
||||
backend.register_workstream("ws1", node_id="node-a")
|
||||
backend.register_workstream("ws2", node_id="node-b")
|
||||
rows = backend.list_workstreams(node_id="node-a")
|
||||
assert len(rows) == 1
|
||||
assert rows[0][0] == "ws1"
|
||||
|
||||
def test_workstream_with_messages_in_history(self, backend):
|
||||
backend.register_workstream("ws1", node_id="node-a")
|
||||
backend.save_message("ws1", "user", "hello")
|
||||
rows = backend.list_workstreams_with_history()
|
||||
assert len(rows) == 1
|
||||
# Columns: ws_id, alias, title, created, updated, count, node_id
|
||||
assert rows[0][0] == "ws1"
|
||||
assert rows[0][6] == "node-a"
|
||||
|
||||
|
||||
# -- Lifecycle -----------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,254 @@
|
||||
"""Tests for turnstone.core.tool_search — BM25 index and tool search manager."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.tool_search import (
|
||||
BM25Index,
|
||||
ToolSearchManager,
|
||||
_mcp_server_summary,
|
||||
_tokenize,
|
||||
_tool_name,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_tool(name: str, description: str = "") -> dict:
|
||||
"""Create a minimal OpenAI-format tool dict for testing."""
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": name,
|
||||
"description": description or f"Tool {name}",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BM25Index tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestTokenize:
|
||||
def test_basic_split(self):
|
||||
assert _tokenize("hello world") == ["hello", "world"]
|
||||
|
||||
def test_underscore_split(self):
|
||||
assert _tokenize("create_issue") == ["create", "issue"]
|
||||
|
||||
def test_mixed_delimiters(self):
|
||||
assert _tokenize("mcp__github__create-issue") == ["mcp", "github", "create", "issue"]
|
||||
|
||||
def test_empty_string(self):
|
||||
assert _tokenize("") == []
|
||||
|
||||
def test_lowercased(self):
|
||||
assert _tokenize("GitHub Create") == ["github", "create"]
|
||||
|
||||
|
||||
class TestBM25Index:
|
||||
def test_empty_corpus(self):
|
||||
idx = BM25Index([])
|
||||
assert idx.search("test") == []
|
||||
|
||||
def test_empty_query(self):
|
||||
idx = BM25Index(["hello world", "foo bar"])
|
||||
assert idx.search("") == []
|
||||
|
||||
def test_single_document(self):
|
||||
idx = BM25Index(["create github issue"])
|
||||
assert idx.search("github") == [0]
|
||||
|
||||
def test_ranking_order(self):
|
||||
docs = [
|
||||
"list_repos List all repositories",
|
||||
"create_issue Create a new GitHub issue",
|
||||
"get_issue Get details of a GitHub issue",
|
||||
]
|
||||
idx = BM25Index(docs)
|
||||
results = idx.search("github issue")
|
||||
# Both issue-related docs should rank above list_repos
|
||||
assert 1 in results[:2]
|
||||
assert 2 in results[:2]
|
||||
|
||||
def test_top_k_limit(self):
|
||||
docs = [f"tool_{i} description {i}" for i in range(20)]
|
||||
idx = BM25Index(docs)
|
||||
results = idx.search("tool description", k=3)
|
||||
assert len(results) <= 3
|
||||
|
||||
def test_no_match(self):
|
||||
idx = BM25Index(["alpha beta gamma"])
|
||||
assert idx.search("zzzzz") == []
|
||||
|
||||
def test_exact_name_match_ranks_high(self):
|
||||
docs = [
|
||||
"send_email Send an email message",
|
||||
"send_slack Send a Slack message",
|
||||
"read_email Read email inbox",
|
||||
]
|
||||
idx = BM25Index(docs)
|
||||
results = idx.search("send email")
|
||||
assert results[0] == 0 # send_email should rank first
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ToolSearchManager tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestToolSearchManager:
|
||||
@pytest.fixture()
|
||||
def builtin_tools(self):
|
||||
return [
|
||||
_make_tool("bash", "Execute shell commands"),
|
||||
_make_tool("read_file", "Read a file"),
|
||||
_make_tool("edit_file", "Edit a file"),
|
||||
]
|
||||
|
||||
@pytest.fixture()
|
||||
def mcp_tools(self):
|
||||
return [
|
||||
_make_tool("mcp__github__create_issue", "Create a new GitHub issue"),
|
||||
_make_tool("mcp__github__list_issues", "List GitHub issues"),
|
||||
_make_tool("mcp__github__get_repo", "Get repository details"),
|
||||
_make_tool("mcp__slack__send_message", "Send a Slack message"),
|
||||
_make_tool("mcp__slack__list_channels", "List Slack channels"),
|
||||
_make_tool("mcp__jira__create_ticket", "Create a Jira ticket"),
|
||||
]
|
||||
|
||||
@pytest.fixture()
|
||||
def manager(self, builtin_tools, mcp_tools):
|
||||
all_tools = builtin_tools + mcp_tools
|
||||
return ToolSearchManager(
|
||||
all_tools,
|
||||
always_on_names={"bash", "read_file", "edit_file"},
|
||||
threshold=5,
|
||||
max_results=3,
|
||||
)
|
||||
|
||||
def test_should_activate_above_threshold(self, manager):
|
||||
assert manager.should_activate()
|
||||
|
||||
def test_should_not_activate_below_threshold(self, builtin_tools):
|
||||
mgr = ToolSearchManager(builtin_tools, always_on_names={"bash", "read_file", "edit_file"})
|
||||
assert not mgr.should_activate()
|
||||
|
||||
def test_visible_tools_initially_builtin_only(self, manager):
|
||||
visible = manager.get_visible_tools()
|
||||
names = {_tool_name(t) for t in visible}
|
||||
assert names == {"bash", "read_file", "edit_file"}
|
||||
|
||||
def test_deferred_tools_excludes_builtin(self, manager):
|
||||
deferred = manager.get_deferred_tools()
|
||||
names = {_tool_name(t) for t in deferred}
|
||||
assert "bash" not in names
|
||||
assert "mcp__github__create_issue" in names
|
||||
|
||||
def test_search_returns_relevant_tools(self, manager):
|
||||
results = manager.search("github issue")
|
||||
names = {_tool_name(t) for t in results}
|
||||
assert "mcp__github__create_issue" in names or "mcp__github__list_issues" in names
|
||||
|
||||
def test_search_respects_max_results(self, manager):
|
||||
results = manager.search("tool")
|
||||
assert len(results) <= 3
|
||||
|
||||
def test_search_excludes_already_expanded(self, manager):
|
||||
# Expand a github tool, then search for github — expanded tool should not appear
|
||||
manager.expand_visible(["mcp__github__create_issue"])
|
||||
results = manager.search("github issue")
|
||||
names = {_tool_name(t) for t in results}
|
||||
assert "mcp__github__create_issue" not in names
|
||||
|
||||
def test_expand_visible_adds_tools(self, manager):
|
||||
manager.expand_visible(["mcp__github__create_issue"])
|
||||
visible = manager.get_visible_tools()
|
||||
names = {_tool_name(t) for t in visible}
|
||||
assert "mcp__github__create_issue" in names
|
||||
|
||||
def test_expand_visible_returns_newly_added(self, manager):
|
||||
added = manager.expand_visible(["mcp__github__create_issue", "mcp__slack__send_message"])
|
||||
assert len(added) == 2
|
||||
names = {_tool_name(t) for t in added}
|
||||
assert names == {"mcp__github__create_issue", "mcp__slack__send_message"}
|
||||
|
||||
def test_expand_visible_idempotent(self, manager):
|
||||
manager.expand_visible(["mcp__github__create_issue"])
|
||||
added = manager.expand_visible(["mcp__github__create_issue"])
|
||||
assert added == []
|
||||
|
||||
def test_expand_visible_ignores_unknown(self, manager):
|
||||
added = manager.expand_visible(["nonexistent_tool"])
|
||||
assert added == []
|
||||
|
||||
def test_get_expanded_names_empty(self, manager):
|
||||
assert manager.get_expanded_names() == []
|
||||
|
||||
def test_get_expanded_names_after_expand(self, manager):
|
||||
manager.expand_visible(["mcp__github__create_issue", "mcp__slack__send_message"])
|
||||
names = manager.get_expanded_names()
|
||||
assert names == ["mcp__github__create_issue", "mcp__slack__send_message"]
|
||||
|
||||
def test_deferred_excludes_expanded(self, manager):
|
||||
manager.expand_visible(["mcp__github__create_issue"])
|
||||
deferred = manager.get_deferred_tools()
|
||||
names = {_tool_name(t) for t in deferred}
|
||||
assert "mcp__github__create_issue" not in names
|
||||
|
||||
def test_get_all_tools_returns_everything(self, manager, builtin_tools, mcp_tools):
|
||||
assert len(manager.get_all_tools()) == len(builtin_tools) + len(mcp_tools)
|
||||
|
||||
def test_search_tool_definition_format(self, manager):
|
||||
defn = manager.get_search_tool_definition()
|
||||
assert defn["type"] == "function"
|
||||
fn = defn["function"]
|
||||
assert fn["name"] == "tool_search"
|
||||
assert "query" in fn["parameters"]["properties"]
|
||||
assert "query" in fn["parameters"]["required"]
|
||||
|
||||
def test_search_tool_description_has_server_hint(self, manager):
|
||||
defn = manager.get_search_tool_definition()
|
||||
desc = defn["function"]["description"]
|
||||
assert "github" in desc
|
||||
assert "slack" in desc
|
||||
assert "jira" in desc
|
||||
|
||||
def test_format_search_results_empty(self, manager):
|
||||
text = manager.format_search_results([])
|
||||
assert "No matching tools found" in text
|
||||
|
||||
def test_format_search_results_with_tools(self, manager, mcp_tools):
|
||||
text = manager.format_search_results(mcp_tools[:2])
|
||||
assert "Found 2" in text
|
||||
assert "mcp__github__create_issue" in text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helper function tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestMCPServerSummary:
|
||||
def test_groups_by_server(self):
|
||||
tools = [
|
||||
_make_tool("mcp__github__a"),
|
||||
_make_tool("mcp__github__b"),
|
||||
_make_tool("mcp__slack__c"),
|
||||
]
|
||||
summary = _mcp_server_summary(tools)
|
||||
assert "github (2 tools)" in summary
|
||||
assert "slack (1 tool)" in summary
|
||||
|
||||
def test_non_mcp_tools_counted_as_other(self):
|
||||
tools = [_make_tool("custom_tool")]
|
||||
summary = _mcp_server_summary(tools)
|
||||
assert "other (1 tool)" in summary
|
||||
|
||||
def test_empty_list(self):
|
||||
assert _mcp_server_summary([]) == ""
|
||||
@@ -72,16 +72,16 @@ class TestToolsMetadata:
|
||||
"""Validate the metadata extracted from JSON files."""
|
||||
|
||||
def test_tool_count(self):
|
||||
assert len(TOOLS) == 14
|
||||
assert len(TOOLS) == 15
|
||||
|
||||
def test_agent_tools_count(self):
|
||||
assert len(AGENT_TOOLS) == 6
|
||||
assert len(AGENT_TOOLS) == 7
|
||||
|
||||
def test_task_agent_tools_count(self):
|
||||
assert len(TASK_AGENT_TOOLS) == 9
|
||||
assert len(TASK_AGENT_TOOLS) == 10
|
||||
|
||||
def test_auto_approve_sets_match(self):
|
||||
expected = {"read_file", "search", "math", "man", "web_fetch", "web_search"}
|
||||
expected = {"read_file", "search", "math", "man", "web_fetch", "web_search", "notify"}
|
||||
assert expected == AGENT_AUTO_TOOLS
|
||||
assert expected == TASK_AUTO_TOOLS
|
||||
|
||||
@@ -101,6 +101,7 @@ class TestToolsMetadata:
|
||||
"remember": "key",
|
||||
"recall": "query",
|
||||
"forget": "key",
|
||||
"notify": "message",
|
||||
}
|
||||
assert expected == PRIMARY_KEY_MAP
|
||||
|
||||
|
||||
@@ -0,0 +1,155 @@
|
||||
"""Tests for user identity storage operations (SQLite backend)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db(tmp_path):
|
||||
"""Create a fresh SQLite backend for each test."""
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
class TestUserCRUD:
|
||||
def test_create_and_get(self, db):
|
||||
db.create_user("u1", "admin", "Admin User", "$2b$hash")
|
||||
user = db.get_user("u1")
|
||||
assert user is not None
|
||||
assert user["user_id"] == "u1"
|
||||
assert user["username"] == "admin"
|
||||
assert user["display_name"] == "Admin User"
|
||||
assert user["password_hash"] == "$2b$hash"
|
||||
|
||||
def test_get_nonexistent(self, db):
|
||||
assert db.get_user("missing") is None
|
||||
|
||||
def test_get_by_username(self, db):
|
||||
db.create_user("u1", "admin", "Admin", "$2b$hash")
|
||||
user = db.get_user_by_username("admin")
|
||||
assert user is not None
|
||||
assert user["user_id"] == "u1"
|
||||
|
||||
def test_get_by_username_nonexistent(self, db):
|
||||
assert db.get_user_by_username("nope") is None
|
||||
|
||||
def test_create_duplicate_noop(self, db):
|
||||
db.create_user("u1", "admin", "First", "$2b$hash1")
|
||||
db.create_user("u1", "admin2", "Second", "$2b$hash2")
|
||||
user = db.get_user("u1")
|
||||
assert user is not None
|
||||
assert user["display_name"] == "First"
|
||||
|
||||
def test_list_users(self, db):
|
||||
db.create_user("u1", "admin", "Admin", "$2b$h1")
|
||||
db.create_user("u2", "reader", "Reader", "$2b$h2")
|
||||
users = db.list_users()
|
||||
assert len(users) == 2
|
||||
assert "password_hash" not in users[0]
|
||||
|
||||
def test_delete_user(self, db):
|
||||
db.create_user("u1", "admin", "Admin", "$2b$hash")
|
||||
assert db.delete_user("u1")
|
||||
assert db.get_user("u1") is None
|
||||
|
||||
def test_delete_nonexistent(self, db):
|
||||
assert not db.delete_user("missing")
|
||||
|
||||
def test_delete_cascades_tokens(self, db):
|
||||
db.create_user("u1", "admin", "Admin", "$2b$hash")
|
||||
db.create_api_token("t1", "hash1", "ts_abcde", "u1", "tok1", "read,write")
|
||||
db.create_api_token("t2", "hash2", "ts_fghij", "u1", "tok2", "read")
|
||||
assert len(db.list_api_tokens("u1")) == 2
|
||||
db.delete_user("u1")
|
||||
assert len(db.list_api_tokens("u1")) == 0
|
||||
|
||||
|
||||
class TestApiTokenCRUD:
|
||||
def test_create_and_lookup_by_hash(self, db):
|
||||
db.create_user("u1", "admin", "Admin", "$2b$hash")
|
||||
db.create_api_token("t1", "tokenhash123", "ts_abcde", "u1", "My Token", "read,write")
|
||||
tok = db.get_api_token_by_hash("tokenhash123")
|
||||
assert tok is not None
|
||||
assert tok["token_id"] == "t1"
|
||||
assert tok["user_id"] == "u1"
|
||||
assert tok["scopes"] == "read,write"
|
||||
|
||||
def test_lookup_missing_hash(self, db):
|
||||
assert db.get_api_token_by_hash("nonexistent") is None
|
||||
|
||||
def test_list_tokens_excludes_hash(self, db):
|
||||
db.create_user("u1", "admin", "Admin", "$2b$hash")
|
||||
db.create_api_token("t1", "secret_hash", "ts_abcde", "u1", "tok1", "read")
|
||||
tokens = db.list_api_tokens("u1")
|
||||
assert len(tokens) == 1
|
||||
assert "token_hash" not in tokens[0]
|
||||
assert tokens[0]["token_prefix"] == "ts_abcde"
|
||||
|
||||
def test_list_tokens_by_user(self, db):
|
||||
db.create_user("u1", "admin", "Admin", "$2b$hash")
|
||||
db.create_user("u2", "reader", "Reader", "$2b$hash")
|
||||
db.create_api_token("t1", "h1", "ts_a", "u1", "tok1", "read")
|
||||
db.create_api_token("t2", "h2", "ts_b", "u2", "tok2", "read")
|
||||
assert len(db.list_api_tokens("u1")) == 1
|
||||
assert len(db.list_api_tokens("u2")) == 1
|
||||
|
||||
def test_delete_token(self, db):
|
||||
db.create_user("u1", "admin", "Admin", "$2b$hash")
|
||||
db.create_api_token("t1", "h1", "ts_a", "u1", "tok1", "read")
|
||||
assert db.delete_api_token("t1")
|
||||
assert db.get_api_token_by_hash("h1") is None
|
||||
|
||||
def test_delete_nonexistent_token(self, db):
|
||||
assert not db.delete_api_token("missing")
|
||||
|
||||
def test_token_with_expiry(self, db):
|
||||
db.create_user("u1", "admin", "Admin", "$2b$hash")
|
||||
db.create_api_token(
|
||||
"t1",
|
||||
"h1",
|
||||
"ts_a",
|
||||
"u1",
|
||||
"tok1",
|
||||
"read",
|
||||
expires="2030-01-01T00:00:00",
|
||||
)
|
||||
tok = db.get_api_token_by_hash("h1")
|
||||
assert tok is not None
|
||||
assert tok["expires"] == "2030-01-01T00:00:00"
|
||||
|
||||
def test_token_without_expiry(self, db):
|
||||
db.create_user("u1", "admin", "Admin", "$2b$hash")
|
||||
db.create_api_token("t1", "h1", "ts_a", "u1", "tok1", "read")
|
||||
tok = db.get_api_token_by_hash("h1")
|
||||
assert tok is not None
|
||||
assert "expires" not in tok
|
||||
|
||||
|
||||
class TestWorkstreamUserId:
|
||||
def test_register_workstream_with_user_id(self, db):
|
||||
db.register_workstream("ws1", user_id="u1")
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.storage._schema import workstreams
|
||||
|
||||
with db._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
sa.select(workstreams.c.user_id).where(workstreams.c.ws_id == "ws1")
|
||||
).fetchone()
|
||||
assert row is not None
|
||||
assert row[0] == "u1"
|
||||
|
||||
def test_register_workstream_without_user_id(self, db):
|
||||
db.register_workstream("ws1")
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.storage._schema import workstreams
|
||||
|
||||
with db._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
sa.select(workstreams.c.user_id).where(workstreams.c.ws_id == "ws1")
|
||||
).fetchone()
|
||||
assert row is not None
|
||||
assert row[0] is None
|
||||
@@ -20,7 +20,7 @@ class FakeSession:
|
||||
self.messages = []
|
||||
|
||||
|
||||
def _fake_factory(ui, model_alias=None):
|
||||
def _fake_factory(ui, model_alias=None, ws_id=None):
|
||||
return FakeSession()
|
||||
|
||||
|
||||
@@ -440,8 +440,10 @@ class TestManagerThreadSafety:
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
assert mgr.count == 5
|
||||
assert len(errors) == 5
|
||||
# All threads resolved (created or rejected)
|
||||
assert len(created) + len(errors) == 10
|
||||
# Never exceeded capacity
|
||||
assert mgr.count <= 5
|
||||
|
||||
def test_concurrent_switch(self):
|
||||
"""Concurrent switches should not corrupt state."""
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
"""turnstone - Multi-node AI orchestration platform with tool use, agent routing, and cluster simulation."""
|
||||
|
||||
__version__ = "0.3.5"
|
||||
__version__ = "0.5.2"
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
"""CLI admin commands for user and token management.
|
||||
|
||||
Entry point: turnstone-admin
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
|
||||
def _get_storage() -> Any:
|
||||
"""Initialize and return the storage backend."""
|
||||
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", "")
|
||||
return init_storage(db_backend, path=db_path, url=db_url)
|
||||
|
||||
|
||||
def _cmd_create_user(args: argparse.Namespace) -> None:
|
||||
import getpass
|
||||
|
||||
from turnstone.core.auth import (
|
||||
generate_token,
|
||||
hash_password,
|
||||
hash_token,
|
||||
is_valid_username,
|
||||
token_prefix,
|
||||
)
|
||||
|
||||
if not is_valid_username(args.username):
|
||||
print("Error: invalid username (1-64 chars: letters, digits, . _ -)", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
storage = _get_storage()
|
||||
user_id = uuid.uuid4().hex
|
||||
|
||||
# Prompt for password
|
||||
password = args.password
|
||||
if not password:
|
||||
password = getpass.getpass("Password: ")
|
||||
confirm = getpass.getpass("Confirm password: ")
|
||||
if password != confirm:
|
||||
print("Error: passwords do not match", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
pw_hash = hash_password(password)
|
||||
storage.create_user(user_id, args.username, args.name, pw_hash)
|
||||
print(f"Created user: {user_id}")
|
||||
print(f" Username: {args.username}")
|
||||
print(f" Name: {args.name}")
|
||||
|
||||
if args.token:
|
||||
scopes = args.scopes or "read,write,approve"
|
||||
raw = generate_token()
|
||||
tid = uuid.uuid4().hex
|
||||
storage.create_api_token(
|
||||
token_id=tid,
|
||||
token_hash=hash_token(raw),
|
||||
token_prefix=token_prefix(raw),
|
||||
user_id=user_id,
|
||||
name="initial",
|
||||
scopes=scopes,
|
||||
)
|
||||
print(f"\n Token: {raw}")
|
||||
print(f" Token ID: {tid}")
|
||||
print(f" Scopes: {scopes}")
|
||||
print(" (Save this token now — it cannot be retrieved again)")
|
||||
|
||||
|
||||
def _cmd_create_token(args: argparse.Namespace) -> None:
|
||||
from turnstone.core.auth import generate_token, hash_token, token_prefix
|
||||
|
||||
storage = _get_storage()
|
||||
|
||||
if storage.get_user(args.user) is None:
|
||||
print(f"Error: user {args.user} not found", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
expires = None
|
||||
if args.expires_days:
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
expires = (datetime.now(UTC) + timedelta(days=args.expires_days)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
)
|
||||
|
||||
raw = generate_token()
|
||||
tid = uuid.uuid4().hex
|
||||
storage.create_api_token(
|
||||
token_id=tid,
|
||||
token_hash=hash_token(raw),
|
||||
token_prefix=token_prefix(raw),
|
||||
user_id=args.user,
|
||||
name=args.name or "",
|
||||
scopes=args.scopes,
|
||||
expires=expires,
|
||||
)
|
||||
print(f"Token: {raw}")
|
||||
print(f" ID: {tid}")
|
||||
print(f" Scopes: {args.scopes}")
|
||||
if expires:
|
||||
print(f" Expires: {expires}")
|
||||
print(" (Save this token now — it cannot be retrieved again)")
|
||||
|
||||
|
||||
def _cmd_list_users(args: argparse.Namespace) -> None:
|
||||
storage = _get_storage()
|
||||
users = storage.list_users()
|
||||
if not users:
|
||||
print("No users found.")
|
||||
return
|
||||
for u in users:
|
||||
print(f" {u['user_id'][:12]}.. {u['display_name']} ({u['created']})")
|
||||
|
||||
|
||||
def _cmd_list_tokens(args: argparse.Namespace) -> None:
|
||||
storage = _get_storage()
|
||||
tokens = storage.list_api_tokens(args.user)
|
||||
if not tokens:
|
||||
print(f"No tokens found for user {args.user}.")
|
||||
return
|
||||
for t in tokens:
|
||||
exp = f" expires={t['expires']}" if t.get("expires") else ""
|
||||
print(
|
||||
f" {t['token_id'][:12]}.. {t['token_prefix']}.. scopes={t['scopes']}"
|
||||
f" name={t['name']}{exp}"
|
||||
)
|
||||
|
||||
|
||||
def _cmd_revoke_token(args: argparse.Namespace) -> None:
|
||||
storage = _get_storage()
|
||||
if storage.delete_api_token(args.token_id):
|
||||
print(f"Revoked token {args.token_id}")
|
||||
else:
|
||||
print("Token not found", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Entry point for turnstone-admin CLI."""
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="turnstone-admin",
|
||||
description="Turnstone user and token administration",
|
||||
)
|
||||
sub = parser.add_subparsers(dest="command")
|
||||
|
||||
p_cu = sub.add_parser("create-user", help="Create a new user")
|
||||
p_cu.add_argument("--username", required=True, help="Login username")
|
||||
p_cu.add_argument("--name", required=True, help="Display name")
|
||||
p_cu.add_argument("--password", default="", help="Password (prompted if not provided)")
|
||||
p_cu.add_argument("--token", action="store_true", help="Also create an initial API token")
|
||||
p_cu.add_argument("--scopes", default="read,write,approve", help="Scopes for initial token")
|
||||
|
||||
p_ct = sub.add_parser("create-token", help="Create an API token for a user")
|
||||
p_ct.add_argument("--user", required=True, help="User ID")
|
||||
p_ct.add_argument("--name", default="", help="Human label for the token")
|
||||
p_ct.add_argument("--scopes", default="read,write", help="Comma-separated scopes")
|
||||
p_ct.add_argument("--expires-days", type=int, default=None, help="Days until expiry")
|
||||
|
||||
sub.add_parser("list-users", help="List all users")
|
||||
|
||||
p_lt = sub.add_parser("list-tokens", help="List tokens for a user")
|
||||
p_lt.add_argument("--user", required=True, help="User ID")
|
||||
|
||||
p_rt = sub.add_parser("revoke-token", help="Revoke an API token")
|
||||
p_rt.add_argument("--token-id", required=True, help="Token ID to revoke")
|
||||
|
||||
args = parser.parse_args()
|
||||
if not args.command:
|
||||
parser.print_help()
|
||||
sys.exit(1)
|
||||
|
||||
dispatch = {
|
||||
"create-user": _cmd_create_user,
|
||||
"create-token": _cmd_create_token,
|
||||
"list-users": _cmd_list_users,
|
||||
"list-tokens": _cmd_list_tokens,
|
||||
"revoke-token": _cmd_revoke_token,
|
||||
}
|
||||
dispatch[args.command](args)
|
||||
@@ -2,6 +2,8 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -48,7 +50,7 @@ class ClusterNodeInfo(BaseModel):
|
||||
total_tokens: int = 0
|
||||
started: float = 0.0
|
||||
reachable: bool = True
|
||||
health: dict[str, str] = Field(default_factory=dict)
|
||||
health: dict[str, Any] = Field(default_factory=dict)
|
||||
version: str = ""
|
||||
|
||||
|
||||
@@ -91,12 +93,34 @@ class ClusterWorkstreamsResponse(BaseModel):
|
||||
class NodeDetailResponse(BaseModel):
|
||||
node_id: str
|
||||
server_url: str = ""
|
||||
health: dict[str, str] = Field(default_factory=dict)
|
||||
health: dict[str, Any] = Field(default_factory=dict)
|
||||
workstreams: list[ClusterWorkstreamInfo] = []
|
||||
aggregate: dict[str, int] = Field(default_factory=dict)
|
||||
reachable: bool = True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cluster snapshot
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ClusterSnapshotNode(BaseModel):
|
||||
node_id: str
|
||||
server_url: str = ""
|
||||
max_ws: int = 10
|
||||
reachable: bool = True
|
||||
version: str = ""
|
||||
health: dict[str, Any] = Field(default_factory=dict)
|
||||
aggregate: dict[str, int] = Field(default_factory=dict)
|
||||
workstreams: list[ClusterWorkstreamInfo] = []
|
||||
|
||||
|
||||
class ClusterSnapshotResponse(BaseModel):
|
||||
nodes: list[ClusterSnapshotNode]
|
||||
overview: ClusterOverviewResponse
|
||||
timestamp: float = 0.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Workstream creation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -10,6 +10,7 @@ if TYPE_CHECKING:
|
||||
from turnstone.api.console_schemas import (
|
||||
ClusterNodesResponse,
|
||||
ClusterOverviewResponse,
|
||||
ClusterSnapshotResponse,
|
||||
ClusterWorkstreamsResponse,
|
||||
ConsoleCreateWsRequest,
|
||||
ConsoleCreateWsResponse,
|
||||
@@ -20,8 +21,22 @@ from turnstone.api.openapi import EndpointSpec, QueryParam, build_openapi
|
||||
from turnstone.api.schemas import (
|
||||
AuthLoginRequest,
|
||||
AuthLoginResponse,
|
||||
AuthSetupRequest,
|
||||
AuthSetupResponse,
|
||||
AuthStatusResponse,
|
||||
CreateScheduleRequest,
|
||||
CreateTokenRequest,
|
||||
CreateTokenResponse,
|
||||
CreateUserRequest,
|
||||
ErrorResponse,
|
||||
ListScheduleRunsResponse,
|
||||
ListSchedulesResponse,
|
||||
ListTokensResponse,
|
||||
ListUsersResponse,
|
||||
ScheduleInfo,
|
||||
StatusResponse,
|
||||
UpdateScheduleRequest,
|
||||
UserInfo,
|
||||
)
|
||||
|
||||
CONSOLE_ENDPOINTS: list[EndpointSpec] = [
|
||||
@@ -83,14 +98,23 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [
|
||||
error_codes=[400, 404, 503],
|
||||
tags=["Cluster"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/cluster/snapshot",
|
||||
"GET",
|
||||
"Full cluster state snapshot",
|
||||
description="Returns the complete cluster state: all nodes with their workstreams "
|
||||
"and overview aggregates. Used for initial load and reconnection.",
|
||||
response_model=ClusterSnapshotResponse,
|
||||
tags=["Cluster"],
|
||||
),
|
||||
# --- Streaming ---
|
||||
EndpointSpec(
|
||||
"/v1/api/cluster/events",
|
||||
"GET",
|
||||
"Cluster SSE event stream",
|
||||
description="Server-Sent Events stream for real-time cluster updates. "
|
||||
"Returns text/event-stream with node_joined, node_lost, cluster_state, "
|
||||
"ws_created, ws_closed, ws_rename events.",
|
||||
"First event is a 'snapshot' with full cluster state, followed by "
|
||||
"node_joined, node_lost, cluster_state, ws_created, ws_closed, ws_rename events.",
|
||||
tags=["Streaming"],
|
||||
),
|
||||
# --- Auth ---
|
||||
@@ -103,6 +127,22 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [
|
||||
error_codes=[401],
|
||||
tags=["Auth"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/auth/setup",
|
||||
"POST",
|
||||
"Create first admin user",
|
||||
request_model=AuthSetupRequest,
|
||||
response_model=AuthSetupResponse,
|
||||
error_codes=[400, 409, 503],
|
||||
tags=["Auth"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/auth/status",
|
||||
"GET",
|
||||
"Return auth state",
|
||||
response_model=AuthStatusResponse,
|
||||
tags=["Auth"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/auth/logout",
|
||||
"POST",
|
||||
@@ -110,6 +150,108 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [
|
||||
response_model=StatusResponse,
|
||||
tags=["Auth"],
|
||||
),
|
||||
# --- Admin ---
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/users",
|
||||
"GET",
|
||||
"List all users",
|
||||
response_model=ListUsersResponse,
|
||||
tags=["Admin"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/users",
|
||||
"POST",
|
||||
"Create a user",
|
||||
request_model=CreateUserRequest,
|
||||
response_model=UserInfo,
|
||||
tags=["Admin"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/users/{user_id}",
|
||||
"DELETE",
|
||||
"Delete a user and their tokens",
|
||||
response_model=StatusResponse,
|
||||
error_codes=[404],
|
||||
tags=["Admin"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/users/{user_id}/tokens",
|
||||
"GET",
|
||||
"List tokens for a user",
|
||||
response_model=ListTokensResponse,
|
||||
tags=["Admin"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/users/{user_id}/tokens",
|
||||
"POST",
|
||||
"Create an API token (raw token shown once)",
|
||||
request_model=CreateTokenRequest,
|
||||
response_model=CreateTokenResponse,
|
||||
tags=["Admin"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/tokens/{token_id}",
|
||||
"DELETE",
|
||||
"Revoke an API token",
|
||||
response_model=StatusResponse,
|
||||
error_codes=[404],
|
||||
tags=["Admin"],
|
||||
),
|
||||
# --- Schedules ---
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/schedules",
|
||||
"GET",
|
||||
"List all scheduled tasks",
|
||||
response_model=ListSchedulesResponse,
|
||||
tags=["Schedules"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/schedules",
|
||||
"POST",
|
||||
"Create a scheduled task",
|
||||
request_model=CreateScheduleRequest,
|
||||
response_model=ScheduleInfo,
|
||||
error_codes=[400],
|
||||
tags=["Schedules"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/schedules/{task_id}",
|
||||
"GET",
|
||||
"Get a scheduled task",
|
||||
response_model=ScheduleInfo,
|
||||
error_codes=[404],
|
||||
tags=["Schedules"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/schedules/{task_id}",
|
||||
"PUT",
|
||||
"Update a scheduled task",
|
||||
request_model=UpdateScheduleRequest,
|
||||
response_model=ScheduleInfo,
|
||||
error_codes=[400, 404],
|
||||
tags=["Schedules"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/schedules/{task_id}",
|
||||
"DELETE",
|
||||
"Delete a scheduled task",
|
||||
response_model=StatusResponse,
|
||||
error_codes=[404],
|
||||
tags=["Schedules"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/schedules/{task_id}/runs",
|
||||
"GET",
|
||||
"List run history for a scheduled task",
|
||||
response_model=ListScheduleRunsResponse,
|
||||
query_params=[
|
||||
QueryParam(
|
||||
"limit", "Max results (default 50, max 200)", schema_type="integer", default=50
|
||||
),
|
||||
],
|
||||
error_codes=[404],
|
||||
tags=["Schedules"],
|
||||
),
|
||||
# --- Observability ---
|
||||
EndpointSpec(
|
||||
"/health",
|
||||
@@ -125,13 +267,28 @@ _ALL_MODELS: list[type[BaseModel]] = [
|
||||
StatusResponse,
|
||||
AuthLoginRequest,
|
||||
AuthLoginResponse,
|
||||
AuthSetupRequest,
|
||||
AuthSetupResponse,
|
||||
AuthStatusResponse,
|
||||
CreateUserRequest,
|
||||
UserInfo,
|
||||
ListUsersResponse,
|
||||
CreateTokenRequest,
|
||||
CreateTokenResponse,
|
||||
ListTokensResponse,
|
||||
ClusterOverviewResponse,
|
||||
ClusterNodesResponse,
|
||||
ClusterWorkstreamsResponse,
|
||||
NodeDetailResponse,
|
||||
ClusterSnapshotResponse,
|
||||
ConsoleCreateWsRequest,
|
||||
ConsoleCreateWsResponse,
|
||||
ConsoleHealthResponse,
|
||||
CreateScheduleRequest,
|
||||
UpdateScheduleRequest,
|
||||
ScheduleInfo,
|
||||
ListSchedulesResponse,
|
||||
ListScheduleRunsResponse,
|
||||
]
|
||||
|
||||
|
||||
|
||||
+197
-3
@@ -35,13 +35,207 @@ class StatusResponse(BaseModel):
|
||||
|
||||
|
||||
class AuthLoginRequest(BaseModel):
|
||||
"""POST /v1/api/auth/login request body."""
|
||||
"""POST /v1/api/auth/login request body.
|
||||
|
||||
token: str = Field(description="Bearer token to authenticate")
|
||||
Either username+password or token must be provided.
|
||||
"""
|
||||
|
||||
username: str = Field(default="", description="Login username")
|
||||
password: str = Field(default="", description="Login password")
|
||||
token: str = Field(default="", description="Legacy: bearer token to authenticate")
|
||||
|
||||
|
||||
class AuthLoginResponse(BaseModel):
|
||||
"""POST /v1/api/auth/login success response."""
|
||||
|
||||
status: str = Field(default="ok")
|
||||
role: str = Field(description="Assigned role", examples=["full", "read"])
|
||||
user_id: str = Field(default="", description="Authenticated user ID")
|
||||
role: str = Field(description="Legacy role", examples=["full", "read"])
|
||||
scopes: str = Field(
|
||||
default="", description="Comma-separated scopes", examples=["read,write,approve"]
|
||||
)
|
||||
jwt: str = Field(default="", description="JWT session token (if JWT auth is configured)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Admin — User identity + API tokens
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class CreateUserRequest(BaseModel):
|
||||
"""POST /v1/api/admin/users request body."""
|
||||
|
||||
username: str = Field(description="Login username (unique)")
|
||||
display_name: str = Field(description="Human-readable display name")
|
||||
password: str = Field(description="Initial password")
|
||||
|
||||
|
||||
class UserInfo(BaseModel):
|
||||
"""User record (no password_hash)."""
|
||||
|
||||
user_id: str
|
||||
username: str
|
||||
display_name: str
|
||||
created: str
|
||||
|
||||
|
||||
class ListUsersResponse(BaseModel):
|
||||
"""GET /v1/api/admin/users response."""
|
||||
|
||||
users: list[UserInfo]
|
||||
|
||||
|
||||
class CreateTokenRequest(BaseModel):
|
||||
"""POST /v1/api/admin/users/{user_id}/tokens request body."""
|
||||
|
||||
name: str = Field(default="", description="Human label for the token")
|
||||
scopes: str = Field(
|
||||
default="read,write,approve",
|
||||
description="Comma-separated scopes: read, write, approve",
|
||||
)
|
||||
expires_days: int | None = Field(
|
||||
default=None,
|
||||
description="Days until expiry (null = no expiry)",
|
||||
)
|
||||
|
||||
|
||||
class TokenInfo(BaseModel):
|
||||
"""Token metadata (never includes the hash or raw token)."""
|
||||
|
||||
token_id: str
|
||||
token_prefix: str
|
||||
name: str
|
||||
scopes: str
|
||||
created: str
|
||||
expires: str | None = None
|
||||
|
||||
|
||||
class CreateTokenResponse(BaseModel):
|
||||
"""POST /v1/api/admin/users/{user_id}/tokens response (raw token shown once)."""
|
||||
|
||||
token: str = Field(description="Raw API token — save this, it cannot be retrieved again")
|
||||
token_id: str
|
||||
token_prefix: str
|
||||
scopes: str
|
||||
|
||||
|
||||
class ListTokensResponse(BaseModel):
|
||||
"""GET /v1/api/admin/users/{user_id}/tokens response."""
|
||||
|
||||
tokens: list[TokenInfo]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth — Setup + status
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class AuthSetupRequest(BaseModel):
|
||||
"""POST /v1/api/auth/setup request body."""
|
||||
|
||||
username: str = Field(description="Login username (1-64 ASCII characters)")
|
||||
display_name: str = Field(description="Display name")
|
||||
password: str = Field(description="Password (minimum 8 characters)")
|
||||
|
||||
|
||||
class AuthSetupResponse(BaseModel):
|
||||
"""POST /v1/api/auth/setup success response."""
|
||||
|
||||
status: str = Field(default="ok")
|
||||
user_id: str
|
||||
username: str
|
||||
role: str = Field(default="full")
|
||||
scopes: str = Field(default="approve,read,write")
|
||||
jwt: str = Field(default="", description="JWT session token")
|
||||
|
||||
|
||||
class AuthStatusResponse(BaseModel):
|
||||
"""GET /v1/api/auth/status response."""
|
||||
|
||||
auth_enabled: bool
|
||||
has_users: bool
|
||||
setup_required: bool
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Schedules
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class CreateScheduleRequest(BaseModel):
|
||||
"""POST /v1/api/admin/schedules request body."""
|
||||
|
||||
name: str = Field(description="Human-readable schedule name")
|
||||
description: str = Field(default="", description="Optional description")
|
||||
schedule_type: str = Field(description="'cron' or 'at'")
|
||||
cron_expr: str = Field(default="", description="Cron expression (when schedule_type='cron')")
|
||||
at_time: str = Field(default="", description="ISO8601 timestamp (when schedule_type='at')")
|
||||
target_mode: str = Field(default="auto", description="auto, pool, all, or specific node_id")
|
||||
model: str = Field(default="", description="Model alias for the workstream")
|
||||
initial_message: str = Field(description="Message sent to the new workstream")
|
||||
auto_approve: bool = Field(default=False)
|
||||
auto_approve_tools: list[str] = Field(default_factory=list)
|
||||
enabled: bool = Field(default=True)
|
||||
|
||||
|
||||
class UpdateScheduleRequest(BaseModel):
|
||||
"""PUT /v1/api/admin/schedules/{task_id} request body (partial update)."""
|
||||
|
||||
name: str | None = None
|
||||
description: str | None = None
|
||||
schedule_type: str | None = None
|
||||
cron_expr: str | None = None
|
||||
at_time: str | None = None
|
||||
target_mode: str | None = None
|
||||
model: str | None = None
|
||||
initial_message: str | None = None
|
||||
auto_approve: bool | None = None
|
||||
auto_approve_tools: list[str] | None = None
|
||||
enabled: bool | None = None
|
||||
|
||||
|
||||
class ScheduleInfo(BaseModel):
|
||||
"""Scheduled task details."""
|
||||
|
||||
task_id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
schedule_type: str
|
||||
cron_expr: str = ""
|
||||
at_time: str = ""
|
||||
target_mode: str = "auto"
|
||||
model: str = ""
|
||||
initial_message: str
|
||||
auto_approve: bool = False
|
||||
auto_approve_tools: list[str] = Field(default_factory=list)
|
||||
enabled: bool = True
|
||||
created_by: str = ""
|
||||
last_run: str | None = None
|
||||
next_run: str | None = None
|
||||
created: str = ""
|
||||
updated: str = ""
|
||||
|
||||
|
||||
class ListSchedulesResponse(BaseModel):
|
||||
"""GET /v1/api/admin/schedules response."""
|
||||
|
||||
schedules: list[ScheduleInfo]
|
||||
|
||||
|
||||
class ScheduleRunInfo(BaseModel):
|
||||
"""Single execution record for a scheduled task."""
|
||||
|
||||
run_id: str
|
||||
task_id: str
|
||||
node_id: str = ""
|
||||
ws_id: str = ""
|
||||
correlation_id: str = ""
|
||||
started: str
|
||||
status: str = "dispatched"
|
||||
error: str = ""
|
||||
|
||||
|
||||
class ListScheduleRunsResponse(BaseModel):
|
||||
"""GET /v1/api/admin/schedules/{task_id}/runs response."""
|
||||
|
||||
runs: list[ScheduleRunInfo]
|
||||
|
||||
@@ -39,11 +39,19 @@ class CreateWorkstreamRequest(BaseModel):
|
||||
name: str = Field(default="", description="Workstream display name (auto-generated if empty)")
|
||||
model: str = Field(default="", description="Model alias from registry")
|
||||
auto_approve: bool = Field(default=False, description="Auto-approve all tool calls")
|
||||
resume_ws: str = Field(
|
||||
default="",
|
||||
description="Workstream ID to resume atomically during creation (empty = fresh start)",
|
||||
)
|
||||
|
||||
|
||||
class CreateWorkstreamResponse(BaseModel):
|
||||
ws_id: str = Field(description="Unique ID of the new workstream")
|
||||
name: str = Field(description="Assigned workstream name")
|
||||
resumed: bool = Field(default=False, description="Whether a previous workstream was resumed")
|
||||
message_count: int = Field(
|
||||
default=0, description="Number of messages in the resumed workstream"
|
||||
)
|
||||
|
||||
|
||||
class CloseWorkstreamRequest(BaseModel):
|
||||
@@ -59,7 +67,6 @@ class WorkstreamInfo(BaseModel):
|
||||
id: str
|
||||
name: str
|
||||
state: str
|
||||
session_id: str | None = None
|
||||
|
||||
|
||||
class ListWorkstreamsResponse(BaseModel):
|
||||
@@ -70,7 +77,6 @@ class DashboardWorkstream(BaseModel):
|
||||
id: str
|
||||
name: str
|
||||
state: str
|
||||
session_id: str | None = None
|
||||
title: str = ""
|
||||
tokens: int = 0
|
||||
context_ratio: float = 0.0
|
||||
@@ -97,12 +103,12 @@ class DashboardResponse(BaseModel):
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Sessions
|
||||
# Saved workstreams
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SessionInfo(BaseModel):
|
||||
session_id: str
|
||||
class SavedWorkstreamInfo(BaseModel):
|
||||
ws_id: str
|
||||
alias: str | None = None
|
||||
title: str | None = None
|
||||
created: str
|
||||
@@ -110,8 +116,8 @@ class SessionInfo(BaseModel):
|
||||
message_count: int
|
||||
|
||||
|
||||
class ListSessionsResponse(BaseModel):
|
||||
sessions: list[SessionInfo]
|
||||
class ListSavedWorkstreamsResponse(BaseModel):
|
||||
workstreams: list[SavedWorkstreamInfo]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -11,6 +11,9 @@ if TYPE_CHECKING:
|
||||
from turnstone.api.schemas import (
|
||||
AuthLoginRequest,
|
||||
AuthLoginResponse,
|
||||
AuthSetupRequest,
|
||||
AuthSetupResponse,
|
||||
AuthStatusResponse,
|
||||
ErrorResponse,
|
||||
StatusResponse,
|
||||
)
|
||||
@@ -22,7 +25,7 @@ from turnstone.api.server_schemas import (
|
||||
CreateWorkstreamResponse,
|
||||
DashboardResponse,
|
||||
HealthResponse,
|
||||
ListSessionsResponse,
|
||||
ListSavedWorkstreamsResponse,
|
||||
ListWorkstreamsResponse,
|
||||
PlanFeedbackRequest,
|
||||
SendRequest,
|
||||
@@ -119,13 +122,13 @@ SERVER_ENDPOINTS: list[EndpointSpec] = [
|
||||
"across all workstreams. Returns text/event-stream.",
|
||||
tags=["Streaming"],
|
||||
),
|
||||
# --- Sessions ---
|
||||
# --- Saved workstreams ---
|
||||
EndpointSpec(
|
||||
"/v1/api/sessions",
|
||||
"/v1/api/workstreams/saved",
|
||||
"GET",
|
||||
"List saved sessions",
|
||||
response_model=ListSessionsResponse,
|
||||
tags=["Sessions"],
|
||||
"List saved workstreams",
|
||||
response_model=ListSavedWorkstreamsResponse,
|
||||
tags=["Workstreams"],
|
||||
),
|
||||
# --- Auth ---
|
||||
EndpointSpec(
|
||||
@@ -137,6 +140,22 @@ SERVER_ENDPOINTS: list[EndpointSpec] = [
|
||||
error_codes=[401],
|
||||
tags=["Auth"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/auth/setup",
|
||||
"POST",
|
||||
"Create first admin user",
|
||||
request_model=AuthSetupRequest,
|
||||
response_model=AuthSetupResponse,
|
||||
error_codes=[400, 409, 503],
|
||||
tags=["Auth"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/auth/status",
|
||||
"GET",
|
||||
"Return auth state",
|
||||
response_model=AuthStatusResponse,
|
||||
tags=["Auth"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/auth/logout",
|
||||
"POST",
|
||||
@@ -159,6 +178,9 @@ _ALL_MODELS: list[type[BaseModel]] = [
|
||||
StatusResponse,
|
||||
AuthLoginRequest,
|
||||
AuthLoginResponse,
|
||||
AuthSetupRequest,
|
||||
AuthSetupResponse,
|
||||
AuthStatusResponse,
|
||||
SendRequest,
|
||||
SendResponse,
|
||||
ApproveRequest,
|
||||
@@ -169,7 +191,7 @@ _ALL_MODELS: list[type[BaseModel]] = [
|
||||
CloseWorkstreamRequest,
|
||||
ListWorkstreamsResponse,
|
||||
DashboardResponse,
|
||||
ListSessionsResponse,
|
||||
ListSavedWorkstreamsResponse,
|
||||
HealthResponse,
|
||||
]
|
||||
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
"""Shared channel infrastructure for turnstone communication integrations.
|
||||
|
||||
Provides the :class:`ChannelAdapter` protocol, the :class:`ChannelEvent`
|
||||
normalized event type, the :class:`ChannelRouter` for workstream mapping,
|
||||
and shared formatting / configuration utilities.
|
||||
"""
|
||||
|
||||
from turnstone.channels._protocol import ChannelAdapter, ChannelEvent
|
||||
from turnstone.channels._routing import ChannelRouter
|
||||
|
||||
__all__ = [
|
||||
"ChannelAdapter",
|
||||
"ChannelEvent",
|
||||
"ChannelRouter",
|
||||
]
|
||||
@@ -0,0 +1,23 @@
|
||||
"""Base configuration shared by all channel adapters."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChannelConfig:
|
||||
"""Base configuration shared by all channel adapters.
|
||||
|
||||
Individual adapters extend this with platform-specific fields (tokens,
|
||||
guild IDs, etc.).
|
||||
"""
|
||||
|
||||
redis_host: str = "localhost"
|
||||
redis_port: int = 6379
|
||||
redis_db: int = 0
|
||||
redis_password: str | None = None
|
||||
prefix: str = "turnstone"
|
||||
model: str = ""
|
||||
auto_approve: bool = False
|
||||
auto_approve_tools: list[str] = field(default_factory=list)
|
||||
@@ -0,0 +1,116 @@
|
||||
"""Message formatting utilities for channel adapters.
|
||||
|
||||
Handles chunking long messages for platforms with character limits, formatting
|
||||
tool-approval requests, and plan-review prompts.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
|
||||
def chunk_message(text: str, max_length: int = 2000) -> list[str]:
|
||||
"""Split *text* into chunks that fit within *max_length*.
|
||||
|
||||
Respects code-block boundaries: if a fenced code block (````` ```)
|
||||
spans a chunk boundary the current chunk is closed with ````` ``` ``
|
||||
and the next chunk reopens it. Prefers splitting at newline
|
||||
boundaries, then word boundaries, then hard-splits.
|
||||
"""
|
||||
if len(text) <= max_length:
|
||||
return [text]
|
||||
|
||||
chunks: list[str] = []
|
||||
remaining = text
|
||||
in_code_block = False
|
||||
|
||||
while remaining:
|
||||
if len(remaining) <= max_length:
|
||||
chunks.append(remaining)
|
||||
break
|
||||
|
||||
# Reserve space for a closing ``` if we're inside a code block.
|
||||
limit = max_length - 4 if in_code_block else max_length
|
||||
limit = max(limit, 1)
|
||||
|
||||
candidate = remaining[:limit]
|
||||
|
||||
# Prefer a newline boundary.
|
||||
split_idx = candidate.rfind("\n")
|
||||
if split_idx <= 0:
|
||||
# Fall back to a word boundary.
|
||||
split_idx = candidate.rfind(" ")
|
||||
if split_idx <= 0:
|
||||
# Hard split.
|
||||
split_idx = limit
|
||||
|
||||
chunk = remaining[:split_idx]
|
||||
remaining = remaining[split_idx:].lstrip("\n")
|
||||
|
||||
# Track code-block fences in this chunk.
|
||||
fence_count = chunk.count("```")
|
||||
block_open = in_code_block
|
||||
|
||||
if fence_count % 2 != 0:
|
||||
in_code_block = not in_code_block
|
||||
|
||||
# If we end inside a code block, close it in this chunk and
|
||||
# reopen in the next.
|
||||
if in_code_block:
|
||||
chunk += "\n```"
|
||||
remaining = "```\n" + remaining
|
||||
in_code_block = False
|
||||
elif block_open and fence_count % 2 != 0:
|
||||
# We were inside a code block and the chunk closed it
|
||||
# properly -- nothing extra needed.
|
||||
pass
|
||||
|
||||
chunks.append(chunk)
|
||||
|
||||
return chunks
|
||||
|
||||
|
||||
def format_approval_request(items: list[dict[str, Any]]) -> str:
|
||||
"""Format tool-approval *items* into a human-readable message.
|
||||
|
||||
Items use the server's SSE format: ``func_name``, ``preview``,
|
||||
``approval_label``, ``header``. Falls back to the nested
|
||||
``function.name`` format for compatibility.
|
||||
"""
|
||||
lines: list[str] = ["**Tool approval required:**"]
|
||||
for item in items:
|
||||
# Server SSE format: top-level func_name / preview
|
||||
name = item.get("func_name") or item.get("approval_label", "")
|
||||
if not name:
|
||||
# Fallback: nested function.name (SDK / older format)
|
||||
func = item.get("function", {})
|
||||
name = func.get("name", "unknown")
|
||||
preview = item.get("preview", "")
|
||||
if not preview:
|
||||
args = item.get("function", {}).get("arguments", "")
|
||||
if isinstance(args, dict):
|
||||
import json
|
||||
|
||||
args = json.dumps(args, ensure_ascii=False)
|
||||
preview = str(args)
|
||||
preview = truncate(preview)
|
||||
header = item.get("header", "")
|
||||
if header:
|
||||
lines.append(f"\u2022 `{name}`: {header}")
|
||||
elif preview:
|
||||
lines.append(f"\u2022 `{name}`: {preview}")
|
||||
else:
|
||||
lines.append(f"\u2022 `{name}`")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def format_plan_review(content: str) -> str:
|
||||
"""Format a plan-review prompt with a header."""
|
||||
return f"**Plan review requested:**\n\n{content}"
|
||||
|
||||
|
||||
def truncate(text: str, max_length: int = 200) -> str:
|
||||
"""Truncate *text* to *max_length*, appending an ellipsis if trimmed."""
|
||||
if len(text) <= max_length:
|
||||
return text
|
||||
return text[: max_length - 1] + "\u2026"
|
||||
@@ -0,0 +1,195 @@
|
||||
"""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``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import socket
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from starlette.applications import Starlette
|
||||
from starlette.responses import JSONResponse
|
||||
from starlette.routing import Mount, Route
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from starlette.requests import Request
|
||||
|
||||
from turnstone.channels._protocol import ChannelAdapter
|
||||
from turnstone.core.storage._protocol import StorageBackend
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
|
||||
async def _handle_health(request: Request) -> JSONResponse:
|
||||
return JSONResponse({"status": "ok", "service": "channel"})
|
||||
|
||||
|
||||
def _check_auth(request: Request) -> JSONResponse | None:
|
||||
"""Validate the request's Authorization header. Returns an error response or None."""
|
||||
auth_token: str = getattr(request.app.state, "auth_token", "")
|
||||
jwt_secret: str = getattr(request.app.state, "jwt_secret", "")
|
||||
|
||||
if not auth_token and not jwt_secret:
|
||||
log.warning("notify.auth_not_configured")
|
||||
return JSONResponse({"error": "authentication not configured"}, status_code=401)
|
||||
|
||||
header = request.headers.get("Authorization", "")
|
||||
if not header.startswith("Bearer "):
|
||||
return JSONResponse({"error": "Unauthorized"}, status_code=401)
|
||||
|
||||
token = header[7:]
|
||||
|
||||
# Static token check
|
||||
if auth_token:
|
||||
import hmac
|
||||
|
||||
if hmac.compare_digest(token, auth_token):
|
||||
return None
|
||||
|
||||
# JWT check
|
||||
if jwt_secret and "." in token:
|
||||
from turnstone.core.auth import JWT_AUD_CHANNEL, validate_jwt
|
||||
|
||||
result = validate_jwt(token, jwt_secret, audience=JWT_AUD_CHANNEL)
|
||||
if result is not None:
|
||||
return None
|
||||
|
||||
return JSONResponse({"error": "Unauthorized"}, status_code=401)
|
||||
|
||||
|
||||
async def _handle_notify(request: Request) -> JSONResponse:
|
||||
"""Deliver a notification to one or more channel adapters."""
|
||||
auth_err = _check_auth(request)
|
||||
if auth_err is not None:
|
||||
return auth_err
|
||||
|
||||
adapters: dict[str, ChannelAdapter] = request.app.state.adapters
|
||||
storage: StorageBackend = request.app.state.storage
|
||||
|
||||
try:
|
||||
body: dict[str, Any] = await request.json()
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return JSONResponse({"error": "invalid JSON"}, status_code=400)
|
||||
|
||||
target = body.get("target")
|
||||
message = body.get("message", "").strip() if isinstance(body.get("message"), str) else ""
|
||||
title = body.get("title", "").strip() if isinstance(body.get("title"), str) else ""
|
||||
|
||||
if not target or not message:
|
||||
return JSONResponse({"error": "target and message are required"}, status_code=400)
|
||||
|
||||
content = f"**{title}**\n{message}" if title else message
|
||||
|
||||
# Resolve targets
|
||||
targets: list[tuple[str, str]] = []
|
||||
if "username" in target:
|
||||
user = await asyncio.to_thread(storage.get_user_by_username, target["username"])
|
||||
if user is None:
|
||||
log.warning("notify.user_not_found", username=target["username"])
|
||||
return JSONResponse(
|
||||
{"error": "target not found or has no linked channels"},
|
||||
status_code=404,
|
||||
)
|
||||
links = await asyncio.to_thread(storage.list_channel_users_by_user, user["user_id"])
|
||||
for link in links:
|
||||
targets.append((link["channel_type"], link["channel_user_id"]))
|
||||
if not targets:
|
||||
log.warning("notify.user_no_linked_channels", username=target["username"])
|
||||
return JSONResponse(
|
||||
{"error": "target not found or has no linked channels"},
|
||||
status_code=404,
|
||||
)
|
||||
elif "channel_type" in target and "channel_id" in target:
|
||||
targets.append((target["channel_type"], target["channel_id"]))
|
||||
else:
|
||||
return JSONResponse(
|
||||
{"error": "target must have username or channel_type+channel_id"},
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
results: list[dict[str, str]] = []
|
||||
for channel_type, channel_id in targets:
|
||||
adapter = adapters.get(channel_type)
|
||||
if adapter is None:
|
||||
results.append(
|
||||
{
|
||||
"channel_type": channel_type,
|
||||
"channel_id": channel_id,
|
||||
"status": "no_adapter",
|
||||
}
|
||||
)
|
||||
log.warning(
|
||||
"notify.no_adapter",
|
||||
channel_type=channel_type,
|
||||
channel_id=channel_id,
|
||||
)
|
||||
continue
|
||||
try:
|
||||
msg_id = await adapter.send(channel_id, content)
|
||||
results.append(
|
||||
{
|
||||
"channel_type": channel_type,
|
||||
"channel_id": channel_id,
|
||||
"status": "sent",
|
||||
"message_id": msg_id,
|
||||
}
|
||||
)
|
||||
log.info(
|
||||
"notify.delivered",
|
||||
channel_type=channel_type,
|
||||
channel_id=channel_id,
|
||||
message_id=msg_id,
|
||||
)
|
||||
except Exception:
|
||||
log.exception(
|
||||
"notify.delivery_failed",
|
||||
channel_type=channel_type,
|
||||
channel_id=channel_id,
|
||||
)
|
||||
results.append(
|
||||
{
|
||||
"channel_type": channel_type,
|
||||
"channel_id": channel_id,
|
||||
"status": "failed",
|
||||
}
|
||||
)
|
||||
|
||||
return JSONResponse({"results": results})
|
||||
|
||||
|
||||
def create_channel_app(
|
||||
adapters: dict[str, ChannelAdapter],
|
||||
storage: StorageBackend,
|
||||
*,
|
||||
auth_token: str = "",
|
||||
jwt_secret: str = "",
|
||||
) -> Starlette:
|
||||
"""Create the channel gateway HTTP application."""
|
||||
app = Starlette(
|
||||
routes=[
|
||||
Route("/health", _handle_health),
|
||||
Mount(
|
||||
"/v1",
|
||||
routes=[
|
||||
Route("/api/notify", _handle_notify, methods=["POST"]),
|
||||
],
|
||||
),
|
||||
],
|
||||
)
|
||||
app.state.adapters = adapters
|
||||
app.state.storage = storage
|
||||
app.state.auth_token = auth_token
|
||||
app.state.jwt_secret = jwt_secret
|
||||
return app
|
||||
|
||||
|
||||
def _get_service_id() -> str:
|
||||
"""Generate a unique service ID from hostname + random suffix."""
|
||||
return f"channel-{socket.gethostname()}-{uuid.uuid4().hex[:8]}"
|
||||
@@ -0,0 +1,70 @@
|
||||
"""Channel adapter protocol and normalized event type.
|
||||
|
||||
Defines the :class:`ChannelEvent` data class for inbound events and the
|
||||
:class:`ChannelAdapter` structural protocol that all bidirectional channel
|
||||
adapters (Discord, Slack, etc.) must satisfy.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChannelEvent:
|
||||
"""Normalized inbound event from any channel."""
|
||||
|
||||
channel_type: str # "discord", "slack"
|
||||
channel_id: str # thread/channel ID
|
||||
channel_user_id: str # platform user ID
|
||||
message: str
|
||||
parent_channel_id: str = "" # main channel (for thread creation)
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ChannelAdapter(Protocol):
|
||||
"""Protocol for bidirectional channel adapters."""
|
||||
|
||||
channel_type: str
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Connect to the platform and begin listening for events."""
|
||||
...
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Disconnect and release resources."""
|
||||
...
|
||||
|
||||
async def send(self, channel_id: str, content: str) -> str:
|
||||
"""Send a message to a channel. Returns the platform message ID."""
|
||||
...
|
||||
|
||||
async def edit_message(self, channel_id: str, message_id: str, content: str) -> None:
|
||||
"""Edit an existing message in a channel."""
|
||||
...
|
||||
|
||||
async def send_approval_request(
|
||||
self,
|
||||
channel_id: str,
|
||||
ws_id: str,
|
||||
correlation_id: str,
|
||||
items: list[dict[str, Any]],
|
||||
) -> None:
|
||||
"""Send an interactive tool-approval prompt to a channel."""
|
||||
...
|
||||
|
||||
async def send_plan_review(
|
||||
self,
|
||||
channel_id: str,
|
||||
ws_id: str,
|
||||
correlation_id: str,
|
||||
content: str,
|
||||
) -> None:
|
||||
"""Send a plan-review prompt to a channel."""
|
||||
...
|
||||
|
||||
async def create_thread(self, parent_channel_id: str, name: str, message_id: str = "") -> str:
|
||||
"""Create a thread under a parent channel. Returns the new thread ID."""
|
||||
...
|
||||
@@ -0,0 +1,297 @@
|
||||
"""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.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
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__)
|
||||
|
||||
_WS_CREATE_TIMEOUT = 30.0 # seconds
|
||||
|
||||
|
||||
class ChannelRouter:
|
||||
"""Manage channel-to-workstream routing and MQ message dispatch.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
broker:
|
||||
An :class:`AsyncRedisBroker` used for pub/sub and queue operations.
|
||||
storage:
|
||||
A :class:`StorageBackend` instance for persistent route lookups.
|
||||
All storage calls are synchronous and will be wrapped in
|
||||
:func:`asyncio.to_thread`.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
broker: AsyncRedisBroker,
|
||||
storage: StorageBackend,
|
||||
*,
|
||||
auto_approve: bool = False,
|
||||
auto_approve_tools: list[str] | None = None,
|
||||
) -> None:
|
||||
self._broker = broker
|
||||
self._storage = storage
|
||||
self._auto_approve = auto_approve
|
||||
self._auto_approve_tools: list[str] = auto_approve_tools or []
|
||||
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] = {}
|
||||
|
||||
# -- 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 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")
|
||||
|
||||
# -- event handler -------------------------------------------------------
|
||||
|
||||
async def _on_global_event(self, raw: str) -> None:
|
||||
"""Handle events on the global pub/sub channel.
|
||||
|
||||
Exceptions are caught so the broker listener task stays alive.
|
||||
"""
|
||||
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"],
|
||||
)
|
||||
except Exception:
|
||||
log.exception("channel_router.global_event_error")
|
||||
|
||||
# -- workstream management -----------------------------------------------
|
||||
|
||||
async def get_or_create_workstream(
|
||||
self,
|
||||
channel_type: str,
|
||||
channel_id: str,
|
||||
name: str = "",
|
||||
model: str = "",
|
||||
initial_message: str = "",
|
||||
) -> tuple[str, bool]:
|
||||
"""Look up or create a workstream for a channel.
|
||||
|
||||
Returns ``(ws_id, is_new)`` where *is_new* is ``True`` when a new
|
||||
workstream was created.
|
||||
|
||||
A per-channel lock prevents duplicate workstreams when concurrent
|
||||
messages arrive for the same channel before the first creation
|
||||
completes.
|
||||
"""
|
||||
key = f"{channel_type}:{channel_id}"
|
||||
lock = self._create_locks.setdefault(key, asyncio.Lock())
|
||||
|
||||
old_ws_id: str | None = None
|
||||
|
||||
async with lock:
|
||||
# 1. Check for existing route.
|
||||
old_ws_id = ""
|
||||
route = await asyncio.to_thread(
|
||||
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:
|
||||
return route["ws_id"], False
|
||||
# Workstream was evicted/closed — capture old ws_id for
|
||||
# resume, then remove the stale route.
|
||||
old_ws_id = route["ws_id"]
|
||||
await asyncio.to_thread(
|
||||
self._storage.delete_channel_route, channel_type, channel_id
|
||||
)
|
||||
log.info(
|
||||
"channel_router.stale_route_cleared",
|
||||
ws_id=old_ws_id,
|
||||
channel_type=channel_type,
|
||||
channel_id=channel_id,
|
||||
)
|
||||
|
||||
# 2. Create via MQ with atomic resume (reuse old ws_id directly).
|
||||
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),
|
||||
)
|
||||
cid = msg.correlation_id
|
||||
waiter = asyncio.Event()
|
||||
self._pending[cid] = waiter
|
||||
|
||||
await self._broker.push_inbound(msg.to_json())
|
||||
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)
|
||||
|
||||
if not ws_id:
|
||||
msg_err = "workstream creation returned empty ws_id"
|
||||
raise RuntimeError(msg_err)
|
||||
|
||||
# 4. Persist the route.
|
||||
await asyncio.to_thread(
|
||||
self._storage.create_channel_route, channel_type, channel_id, ws_id
|
||||
)
|
||||
log.info(
|
||||
"channel_router.route_created",
|
||||
ws_id=ws_id,
|
||||
channel_type=channel_type,
|
||||
channel_id=channel_id,
|
||||
)
|
||||
|
||||
return ws_id, True
|
||||
|
||||
# -- user resolution -----------------------------------------------------
|
||||
|
||||
async def resolve_user(self, channel_type: str, channel_user_id: str) -> str | None:
|
||||
"""Resolve an external platform user to a turnstone ``user_id``.
|
||||
|
||||
Returns ``None`` if no mapping exists.
|
||||
"""
|
||||
result = await asyncio.to_thread(
|
||||
self._storage.get_channel_user, channel_type, channel_user_id
|
||||
)
|
||||
if result is None:
|
||||
return None
|
||||
return result.get("user_id")
|
||||
|
||||
# -- 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_approval(
|
||||
self,
|
||||
ws_id: str,
|
||||
correlation_id: str,
|
||||
approved: bool,
|
||||
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())
|
||||
log.debug(
|
||||
"channel_router.send_approval",
|
||||
ws_id=ws_id,
|
||||
correlation_id=correlation_id,
|
||||
approved=approved,
|
||||
)
|
||||
|
||||
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())
|
||||
log.debug(
|
||||
"channel_router.send_plan_feedback",
|
||||
ws_id=ws_id,
|
||||
correlation_id=correlation_id,
|
||||
)
|
||||
|
||||
# -- route management ----------------------------------------------------
|
||||
|
||||
async def delete_route(self, channel_type: str, channel_id: str) -> None:
|
||||
"""Remove a channel-to-workstream mapping."""
|
||||
deleted = await asyncio.to_thread(
|
||||
self._storage.delete_channel_route, channel_type, channel_id
|
||||
)
|
||||
log.info(
|
||||
"channel_router.delete_route",
|
||||
channel_type=channel_type,
|
||||
channel_id=channel_id,
|
||||
deleted=deleted,
|
||||
)
|
||||
@@ -0,0 +1,238 @@
|
||||
"""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.
|
||||
|
||||
Run as: ``turnstone-channel --discord-token $TURNSTONE_DISCORD_TOKEN``
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Parse arguments, initialize storage and broker, 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)
|
||||
|
||||
# -- Discord -------------------------------------------------------------
|
||||
parser.add_argument(
|
||||
"--discord-token",
|
||||
default=os.environ.get("TURNSTONE_DISCORD_TOKEN", ""),
|
||||
help="Discord bot token (default: $TURNSTONE_DISCORD_TOKEN)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--discord-guild",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Restrict to a single Discord guild (0 = all, default: %(default)s)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--discord-channels",
|
||||
default="",
|
||||
help="Comma-separated list of allowed Discord channel IDs (default: all)",
|
||||
)
|
||||
|
||||
# -- HTTP server ---------------------------------------------------------
|
||||
parser.add_argument(
|
||||
"--http-host",
|
||||
default="127.0.0.1",
|
||||
help="HTTP server bind address (default: %(default)s)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--http-port",
|
||||
type=int,
|
||||
default=int(os.environ.get("TURNSTONE_CHANNEL_PORT", "8091")),
|
||||
help="HTTP server port (default: $TURNSTONE_CHANNEL_PORT or 8091)",
|
||||
)
|
||||
|
||||
# -- Auth ----------------------------------------------------------------
|
||||
parser.add_argument(
|
||||
"--auth-token",
|
||||
default=os.environ.get("TURNSTONE_CHANNEL_AUTH_TOKEN", ""),
|
||||
help="Static auth token for /v1/api/notify (default: $TURNSTONE_CHANNEL_AUTH_TOKEN)",
|
||||
)
|
||||
|
||||
# -- Workstream defaults -------------------------------------------------
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
default="",
|
||||
help="Default model for new workstreams (default: server default)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto-approve",
|
||||
action="store_true",
|
||||
help="Auto-approve all tool calls",
|
||||
)
|
||||
|
||||
# -- Logging -------------------------------------------------------------
|
||||
from turnstone.core.log import add_log_args
|
||||
|
||||
add_log_args(parser)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# -- Logging setup -------------------------------------------------------
|
||||
from turnstone.core.log import configure_logging_from_args
|
||||
|
||||
configure_logging_from_args(args, "channel")
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
# -- Storage -------------------------------------------------------------
|
||||
from turnstone.core.storage._registry 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", "")
|
||||
|
||||
init_storage(
|
||||
backend=db_backend,
|
||||
url=db_url,
|
||||
path=db_path,
|
||||
)
|
||||
|
||||
# -- Auth config ---------------------------------------------------------
|
||||
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)
|
||||
|
||||
# -- Adapter selection ---------------------------------------------------
|
||||
adapters_configured = False
|
||||
|
||||
if args.discord_token:
|
||||
adapters_configured = True
|
||||
|
||||
if not adapters_configured:
|
||||
print(
|
||||
"Error: no channel adapters configured. "
|
||||
"Set --discord-token or $TURNSTONE_DISCORD_TOKEN.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
# -- Run -----------------------------------------------------------------
|
||||
if args.discord_token:
|
||||
import asyncio
|
||||
|
||||
from turnstone.channels._http import _get_service_id, create_channel_app
|
||||
from turnstone.channels.discord.bot import TurnstoneBot
|
||||
from turnstone.channels.discord.config import DiscordConfig
|
||||
from turnstone.core.storage._registry import get_storage
|
||||
|
||||
allowed_channels: list[int] = []
|
||||
if args.discord_channels:
|
||||
allowed_channels = [
|
||||
int(c.strip()) for c in args.discord_channels.split(",") if c.strip()
|
||||
]
|
||||
|
||||
config = DiscordConfig(
|
||||
redis_host=args.redis_host,
|
||||
redis_port=args.redis_port,
|
||||
redis_db=args.redis_db,
|
||||
redis_password=args.redis_password,
|
||||
model=args.model,
|
||||
auto_approve=args.auto_approve,
|
||||
bot_token=args.discord_token,
|
||||
guild_id=args.discord_guild,
|
||||
allowed_channels=allowed_channels,
|
||||
)
|
||||
|
||||
storage = get_storage()
|
||||
bot = TurnstoneBot(config, broker, storage)
|
||||
adapters = {"discord": bot}
|
||||
|
||||
# Create HTTP app for notification delivery
|
||||
channel_app = create_channel_app(
|
||||
adapters, # type: ignore[arg-type]
|
||||
storage,
|
||||
auth_token=auth_token,
|
||||
jwt_secret=jwt_secret,
|
||||
)
|
||||
|
||||
log.info(
|
||||
"channel.starting",
|
||||
adapter="discord",
|
||||
guild_id=config.guild_id,
|
||||
http_port=args.http_port,
|
||||
)
|
||||
|
||||
async def _run_all() -> None:
|
||||
"""Run Discord bot + HTTP server + service heartbeat concurrently."""
|
||||
import uvicorn
|
||||
|
||||
service_id = _get_service_id()
|
||||
|
||||
# Resolve advertise URL — env override for Docker/K8s,
|
||||
# otherwise derive from bind address.
|
||||
advertise_url = os.environ.get("TURNSTONE_CHANNEL_ADVERTISE_URL", "").strip()
|
||||
if not advertise_url:
|
||||
if args.http_host in ("0.0.0.0", "::"):
|
||||
advertise_host = socket.gethostname()
|
||||
else:
|
||||
advertise_host = args.http_host
|
||||
advertise_url = f"http://{advertise_host}:{args.http_port}"
|
||||
service_url = advertise_url
|
||||
|
||||
# Register in service registry
|
||||
storage.register_service("channel", service_id, service_url)
|
||||
log.info(
|
||||
"channel.service_registered",
|
||||
service_id=service_id,
|
||||
url=service_url,
|
||||
)
|
||||
|
||||
async def _heartbeat_loop() -> None:
|
||||
"""Periodically update service heartbeat."""
|
||||
while True:
|
||||
await asyncio.sleep(30)
|
||||
try:
|
||||
await asyncio.to_thread(storage.heartbeat_service, "channel", service_id)
|
||||
except Exception:
|
||||
log.exception("channel.heartbeat_failed")
|
||||
|
||||
uv_config = uvicorn.Config(
|
||||
channel_app,
|
||||
host=args.http_host,
|
||||
port=args.http_port,
|
||||
log_level="warning",
|
||||
)
|
||||
server = uvicorn.Server(uv_config)
|
||||
|
||||
heartbeat_task = asyncio.create_task(_heartbeat_loop())
|
||||
try:
|
||||
await asyncio.gather(
|
||||
bot.start(),
|
||||
server.serve(),
|
||||
)
|
||||
finally:
|
||||
heartbeat_task.cancel()
|
||||
await asyncio.to_thread(storage.deregister_service, "channel", service_id)
|
||||
log.info("channel.service_deregistered", service_id=service_id)
|
||||
|
||||
import contextlib
|
||||
|
||||
with contextlib.suppress(KeyboardInterrupt):
|
||||
asyncio.run(_run_all())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user