mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-13 07:22:24 -06:00
Compare commits
54 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| f3d33bf44a | |||
| fdb1a189e8 | |||
| 58e2d9348f | |||
| 771d03b8e6 | |||
| d147aaea36 | |||
| 8b747178e0 | |||
| d57280d807 | |||
| b9870f279c | |||
| 71ee340bc6 | |||
| c5d5d0b7cd | |||
| 1f9d03c3e0 | |||
| 24e082df05 | |||
| 4a78d20eea | |||
| e0d17e0f99 | |||
| cd6c49dd01 | |||
| 8454e961ba | |||
| 4ae38bc2ae | |||
| ab1a71c86c | |||
| ce57df6888 | |||
| 275f40eebb | |||
| 2c510f8617 | |||
| 2afb9c7f72 | |||
| 5b8ab94446 | |||
| 30828e9f9c | |||
| 6d0dc6df94 | |||
| 7e680ee883 | |||
| e950219246 | |||
| b3764a8035 | |||
| 756c4d8929 | |||
| 04c50568e9 | |||
| 3bf220c503 | |||
| 4b853e329e | |||
| 29ffdc36d0 | |||
| ada8b80509 | |||
| 1f47ca62de | |||
| 83d9233304 | |||
| bf06102d37 | |||
| e015b4512d | |||
| 0c1afff7fc | |||
| a94051a995 | |||
| 2f906ea1f9 | |||
| b61bfd1aa6 | |||
| 19c3a48b10 | |||
| 414eb52d67 | |||
| 86b404177b | |||
| 10165bb8a1 | |||
| 5d478573cc | |||
| 1b24e4717f | |||
| 9a2db63c07 | |||
| e159837b74 | |||
| ec3454ee2e | |||
| 5cbb832162 | |||
| 4e4ae2a91d | |||
| c7d0bac638 |
@@ -47,11 +47,37 @@ jobs:
|
||||
name: coverage-${{ matrix.python-version }}
|
||||
path: coverage.xml
|
||||
|
||||
test-postgres:
|
||||
runs-on: ubuntu-latest
|
||||
services:
|
||||
postgres:
|
||||
image: postgres:18
|
||||
env:
|
||||
POSTGRES_USER: postgres
|
||||
POSTGRES_PASSWORD: postgres
|
||||
POSTGRES_DB: turnstone_test
|
||||
ports:
|
||||
- 5432:5432
|
||||
options: >-
|
||||
--health-cmd="pg_isready -U postgres"
|
||||
--health-interval=10s
|
||||
--health-timeout=5s
|
||||
--health-retries=5
|
||||
steps:
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
|
||||
- uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6
|
||||
with:
|
||||
python-version: "3.12"
|
||||
- run: pip install -e ".[test,mq,postgres]"
|
||||
- run: pytest tests/ -m "not live" --storage-backend=postgresql -q
|
||||
env:
|
||||
TURNSTONE_TEST_PG_URL: postgresql+psycopg://postgres:postgres@localhost:5432/turnstone_test
|
||||
|
||||
lock-check:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
|
||||
- uses: astral-sh/setup-uv@e06108dd0aef18192324c70427afc47652e63a82 # v7
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
with:
|
||||
uv-version: "0.9.18"
|
||||
- run: uv lock --check
|
||||
@@ -60,7 +86,7 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
|
||||
- uses: astral-sh/setup-uv@e06108dd0aef18192324c70427afc47652e63a82 # v7
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
with:
|
||||
uv-version: "0.9.18"
|
||||
- uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6
|
||||
|
||||
+1
-1
@@ -8,7 +8,7 @@ FROM python:3.14-slim
|
||||
LABEL org.opencontainers.image.title="turnstone" \
|
||||
org.opencontainers.image.description="Multi-node AI orchestration platform"
|
||||
|
||||
COPY --from=ghcr.io/astral-sh/uv:0.10.10 /uv /usr/local/bin/uv
|
||||
COPY --from=ghcr.io/astral-sh/uv:0.10.12 /uv /usr/local/bin/uv
|
||||
|
||||
# System dependencies for psycopg (PostgreSQL client library)
|
||||
RUN apt-get update && apt-get upgrade -y && apt-get install -y --no-install-recommends libpq5 \
|
||||
|
||||
@@ -5,9 +5,11 @@
|
||||
[](https://pypi.org/project/turnstone/)
|
||||
[](LICENSE)
|
||||
|
||||
Multi-node AI orchestration platform. Deploy tool-using AI agents across a cluster of servers, driven by message queues or interactive interfaces.
|
||||
Experimental multi-node AI orchestration platform. Deploy tool-using AI agents across a cluster of servers, driven by message queues or interactive interfaces.
|
||||
|
||||
Named after the [Ruddy Turnstone](https://en.wikipedia.org/wiki/Ruddy_turnstone) — a bird that flips rocks to expose what's hiding underneath.
|
||||
> **Beta — Use at your own risk.** Turnstone is under active development and has not reached a stable release. APIs, configuration formats, and database schemas may change between versions without migration paths. We make no guarantees of determinism, reliability, or backward compatibility. Evaluate thoroughly before deploying to any environment where these properties matter.
|
||||
|
||||
Named after the [Ruddy Turnstone](https://en.wikipedia.org/wiki/Ruddy_turnstone) (*Arenaria interpres*) — a shorebird that flips stones to discover what's hiding underneath.
|
||||
|
||||
## What it does
|
||||
|
||||
@@ -308,7 +310,7 @@ search_max_results = 5 # max tools returned per search query
|
||||
[server]
|
||||
host = "0.0.0.0"
|
||||
port = 8080
|
||||
max_workstreams = 10 # auto-evicts oldest idle when full
|
||||
max_workstreams = 50 # auto-evicts oldest idle when full
|
||||
|
||||
[redis]
|
||||
host = "localhost"
|
||||
@@ -340,7 +342,7 @@ burst = 20
|
||||
backend = "sqlite" # "sqlite" (default) or "postgresql"
|
||||
path = ".turnstone.db" # SQLite file path (relative to working directory)
|
||||
# url = "postgresql+psycopg://user:pass@host:5432/turnstone" # PostgreSQL
|
||||
# pool_size = 5 # PostgreSQL connection pool size
|
||||
# pool_size = 2 # PostgreSQL connection pool size (per process)
|
||||
|
||||
[judge]
|
||||
enabled = true # intent validation for tool approvals (--no-judge to disable)
|
||||
@@ -394,7 +396,7 @@ Idle workstreams are automatically cleaned up after 2 hours (configurable). In m
|
||||
- `turnstone_judge_llm_latency_seconds` — LLM judge evaluation latency histogram
|
||||
- `turnstone_judge_enabled` — whether the intent validation judge is active (0/1)
|
||||
|
||||
Per-workstream metrics are labeled by `ws_id` (bounded to 10 max workstreams).
|
||||
Per-workstream metrics are labeled by `ws_id` (bounded by `[server].max_workstreams`).
|
||||
|
||||
### Health & Rate Limiting
|
||||
|
||||
@@ -404,7 +406,7 @@ Per-workstream metrics are labeled by `ws_id` (bounded to 10 max workstreams).
|
||||
|
||||
**Per-IP rate limiting.** When `[ratelimit].enabled` is true, each client IP is tracked with a token-bucket limiter (`requests_per_second` / `burst`). Rate limiting is applied in `do_GET`/`do_POST` after authentication but before route dispatch. `/health` and `/metrics` are exempt. Requests that exceed the limit receive HTTP 429 with a `Retry-After` header.
|
||||
|
||||
**Workstream eviction.** When `WorkstreamManager.create()` would exceed `max_workstreams`, the oldest IDLE workstream is automatically evicted and the `turnstone_workstreams_evicted_total` counter is incremented. Configure via `[server].max_workstreams` (default 10).
|
||||
**Workstream eviction.** When `WorkstreamManager.create()` would exceed `max_workstreams`, the oldest IDLE workstream is automatically evicted and the `turnstone_workstreams_evicted_total` counter is incremented. Configure via `[server].max_workstreams` (default 50).
|
||||
|
||||
## Requirements
|
||||
|
||||
|
||||
+2354
-5
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,49 @@
|
||||
# OpenShell inference routing for Turnstone.
|
||||
#
|
||||
# When using inference routing, the sandbox process connects to
|
||||
# https://inference.local instead of the real LLM API. The OpenShell
|
||||
# proxy intercepts, rewrites credentials, and forwards to the backend.
|
||||
#
|
||||
# This keeps real API keys out of the sandbox entirely — the process
|
||||
# only sees opaque placeholder tokens in its environment.
|
||||
#
|
||||
# Usage:
|
||||
# openshell sandbox run \
|
||||
# --inference-routes deploy/openshell/routes.yaml \
|
||||
# ...
|
||||
#
|
||||
# Then start turnstone with:
|
||||
# python3 -m turnstone.server --base-url https://inference.local
|
||||
#
|
||||
# CUSTOMIZE: uncomment one of the provider blocks below.
|
||||
|
||||
routes:
|
||||
|
||||
# --- OpenAI ---
|
||||
# - name: inference.local
|
||||
# endpoint: https://api.openai.com/v1
|
||||
# model: gpt-5
|
||||
# provider_type: openai
|
||||
# protocols:
|
||||
# - openai_chat_completions
|
||||
# - model_discovery
|
||||
# api_key_env: OPENAI_API_KEY
|
||||
|
||||
# --- Anthropic ---
|
||||
# - name: inference.local
|
||||
# endpoint: https://api.anthropic.com
|
||||
# model: claude-sonnet-4-6
|
||||
# provider_type: anthropic
|
||||
# protocols:
|
||||
# - anthropic_messages
|
||||
# api_key_env: ANTHROPIC_API_KEY
|
||||
|
||||
# --- Local model server (vLLM / llama.cpp) ---
|
||||
# No secret resolution needed — local servers typically have no auth.
|
||||
# Omit both api_key and api_key_env to skip credential injection.
|
||||
# - name: inference.local
|
||||
# endpoint: http://localhost:8000/v1
|
||||
# model: meta-llama/Llama-3.1-70B-Instruct
|
||||
# protocols:
|
||||
# - openai_chat_completions
|
||||
# - model_discovery
|
||||
@@ -0,0 +1,333 @@
|
||||
# OpenShell sandbox policy for Turnstone AI orchestration platform.
|
||||
#
|
||||
# This policy wraps a turnstone-server process (the primary sandbox target).
|
||||
# The bridge, console, and channel gateway are separate processes that would
|
||||
# each need their own sandbox with a tailored policy variant.
|
||||
#
|
||||
# Usage:
|
||||
# openshell sandbox run \
|
||||
# --policy deploy/openshell/turnstone-policy.yaml \
|
||||
# --workdir /project \
|
||||
# -- python3 -m turnstone.server --host 0.0.0.0 --port 8080
|
||||
#
|
||||
# For inference routing (keeps real API keys out of the sandbox):
|
||||
# openshell sandbox run \
|
||||
# --policy deploy/openshell/turnstone-policy.yaml \
|
||||
# --inference-routes deploy/openshell/routes.yaml \
|
||||
# --workdir /project \
|
||||
# -- python3 -m turnstone.server --host 0.0.0.0 --port 8080 \
|
||||
# --base-url https://inference.local
|
||||
#
|
||||
# Note: inference.local is intercepted by the OpenShell proxy before
|
||||
# network policy evaluation — no network_policies entry is needed for it.
|
||||
#
|
||||
# Customization points (search for "CUSTOMIZE"):
|
||||
# - OIDC issuer endpoint
|
||||
# - Redis host/port (if not localhost)
|
||||
# - MCP HTTP server endpoints
|
||||
# - Additional tool binaries
|
||||
# - web_fetch domain allowlist
|
||||
|
||||
version: 1
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Filesystem: Landlock kernel enforcement
|
||||
# ---------------------------------------------------------------------------
|
||||
# Static — cannot be changed after sandbox creation.
|
||||
# include_workdir adds the --workdir path to read_write automatically.
|
||||
|
||||
filesystem_policy:
|
||||
include_workdir: true
|
||||
|
||||
read_only:
|
||||
# Python runtime + installed packages (includes turnstone package)
|
||||
- /usr
|
||||
- /lib
|
||||
- /lib64
|
||||
# System essentials
|
||||
- /etc
|
||||
- /proc
|
||||
- /dev/urandom
|
||||
# Turnstone config (read-only — writes go to database)
|
||||
# CUSTOMIZE: adjust if config lives elsewhere
|
||||
- /home/sandbox/.config/turnstone
|
||||
|
||||
read_write:
|
||||
# Working directory is added via include_workdir
|
||||
# Temp files (bash tool scripts, eval workdirs)
|
||||
- /tmp
|
||||
# Shell redirections (2>/dev/null)
|
||||
- /dev/null
|
||||
# SQLite database (default location is workdir, covered by include_workdir)
|
||||
# Logs
|
||||
- /var/log
|
||||
|
||||
landlock:
|
||||
# best_effort: degrade gracefully on kernels without Landlock (< 5.13)
|
||||
# Change to hard_requirement for production hardened deployments
|
||||
compatibility: best_effort
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Process: privilege separation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
process:
|
||||
run_as_user: sandbox
|
||||
run_as_group: sandbox
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Network: per-endpoint, per-binary allowlisting
|
||||
# ---------------------------------------------------------------------------
|
||||
# Default-deny. Only listed host:port pairs are reachable.
|
||||
# Child processes (MCP servers, bash subcommands) inherit the network
|
||||
# namespace — they cannot bypass the proxy.
|
||||
|
||||
network_policies:
|
||||
|
||||
# --- LLM API providers ---
|
||||
|
||||
openai_api:
|
||||
name: openai-api
|
||||
endpoints:
|
||||
- host: api.openai.com
|
||||
port: 443
|
||||
binaries:
|
||||
- path: /usr/bin/python3*
|
||||
- path: /usr/local/bin/python3*
|
||||
|
||||
anthropic_api:
|
||||
name: anthropic-api
|
||||
endpoints:
|
||||
- host: api.anthropic.com
|
||||
port: 443
|
||||
binaries:
|
||||
- path: /usr/bin/python3*
|
||||
- path: /usr/local/bin/python3*
|
||||
|
||||
# --- Web search fallback (Tavily) ---
|
||||
|
||||
tavily_api:
|
||||
name: tavily-search
|
||||
endpoints:
|
||||
- host: api.tavily.com
|
||||
port: 443
|
||||
binaries:
|
||||
- path: /usr/bin/python3*
|
||||
- path: /usr/local/bin/python3*
|
||||
|
||||
# --- Skill discovery ---
|
||||
|
||||
skills_registry:
|
||||
name: skills-registry
|
||||
endpoints:
|
||||
- host: skills.sh
|
||||
port: 443
|
||||
binaries:
|
||||
- path: /usr/bin/python3*
|
||||
- path: /usr/local/bin/python3*
|
||||
|
||||
github_api:
|
||||
name: github-api
|
||||
endpoints:
|
||||
- host: api.github.com
|
||||
port: 443
|
||||
protocol: rest
|
||||
tls: terminate
|
||||
enforcement: enforce
|
||||
access: read-only
|
||||
- host: raw.githubusercontent.com
|
||||
port: 443
|
||||
binaries:
|
||||
- path: /usr/bin/python3*
|
||||
- path: /usr/local/bin/python3*
|
||||
|
||||
mcp_registry:
|
||||
name: mcp-registry
|
||||
endpoints:
|
||||
- host: registry.modelcontextprotocol.io
|
||||
port: 443
|
||||
protocol: rest
|
||||
tls: terminate
|
||||
enforcement: enforce
|
||||
access: read-only
|
||||
binaries:
|
||||
- path: /usr/bin/python3*
|
||||
- path: /usr/local/bin/python3*
|
||||
|
||||
# --- OIDC SSO ---
|
||||
# CUSTOMIZE: replace with your identity provider's hostname
|
||||
|
||||
# oidc_provider:
|
||||
# name: oidc-provider
|
||||
# endpoints:
|
||||
# - host: login.example.com
|
||||
# port: 443
|
||||
# binaries:
|
||||
# - path: /usr/bin/python3*
|
||||
# - path: /usr/local/bin/python3*
|
||||
|
||||
# --- Redis (MQ) ---
|
||||
# CUSTOMIZE: if Redis is not on localhost, add host + allowed_ips.
|
||||
# localhost is blocked by default SSRF protection, so we need allowed_ips.
|
||||
|
||||
redis:
|
||||
name: redis-mq
|
||||
endpoints:
|
||||
- port: 6379
|
||||
allowed_ips:
|
||||
- "127.0.0.1"
|
||||
binaries:
|
||||
- path: /usr/bin/python3*
|
||||
- path: /usr/local/bin/python3*
|
||||
|
||||
# --- Discord (channel integration) ---
|
||||
# Uncomment if using turnstone-channel with Discord adapter.
|
||||
|
||||
# discord:
|
||||
# name: discord
|
||||
# endpoints:
|
||||
# - host: discord.com
|
||||
# port: 443
|
||||
# - host: gateway.discord.gg
|
||||
# port: 443
|
||||
# - host: cdn.discordapp.com
|
||||
# port: 443
|
||||
# binaries:
|
||||
# - path: /usr/bin/python3*
|
||||
# - path: /usr/local/bin/python3*
|
||||
|
||||
# --- web_fetch tool: curated domain allowlist ---
|
||||
#
|
||||
# This is the hard tradeoff. Turnstone's web_fetch tool lets the LLM
|
||||
# fetch arbitrary public URLs. OpenShell cannot allow "all HTTPS" —
|
||||
# every domain must be enumerated.
|
||||
#
|
||||
# Strategy: allowlist the domains your workloads actually need.
|
||||
# The web_fetch tool will return a connection error for unlisted domains,
|
||||
# which the LLM handles gracefully (it tells the user it can't reach
|
||||
# that site).
|
||||
#
|
||||
# CUSTOMIZE: add domains your workstreams need to fetch from.
|
||||
|
||||
web_fetch_common:
|
||||
name: web-fetch-common
|
||||
endpoints:
|
||||
# Documentation sites
|
||||
- host: "**.readthedocs.io"
|
||||
port: 443
|
||||
- host: docs.python.org
|
||||
port: 443
|
||||
- host: "**.github.io"
|
||||
port: 443
|
||||
# Package registries (metadata lookups)
|
||||
- host: pypi.org
|
||||
port: 443
|
||||
- host: www.npmjs.com
|
||||
port: 443
|
||||
# Stack Overflow / reference
|
||||
- host: stackoverflow.com
|
||||
port: 443
|
||||
- host: "**.stackexchange.com"
|
||||
port: 443
|
||||
# Wikipedia
|
||||
- host: "**.wikipedia.org"
|
||||
port: 443
|
||||
binaries:
|
||||
- path: /usr/bin/python3*
|
||||
- path: /usr/local/bin/python3*
|
||||
|
||||
# --- MCP HTTP servers ---
|
||||
# CUSTOMIZE: add endpoints for any MCP servers using streamable-http
|
||||
# transport. stdio-transport MCP servers need no network entry (they
|
||||
# communicate via stdin/stdout pipes within the sandbox).
|
||||
|
||||
# mcp_http_servers:
|
||||
# name: mcp-http
|
||||
# endpoints:
|
||||
# - host: mcp.internal.example.com
|
||||
# port: 443
|
||||
# binaries:
|
||||
# - path: /usr/bin/python3*
|
||||
# - path: /usr/local/bin/python3*
|
||||
|
||||
# --- Bash tool: curl/wget ---
|
||||
# The bash tool can run curl/wget. These inherit the network namespace
|
||||
# so they can only reach allowed endpoints. But they need binary entries
|
||||
# to pass the proxy's identity check.
|
||||
|
||||
bash_network_tools:
|
||||
name: bash-network-tools
|
||||
endpoints:
|
||||
# Mirrors web_fetch_common — curl/wget should have the same reach.
|
||||
- host: "**.readthedocs.io"
|
||||
port: 443
|
||||
- host: docs.python.org
|
||||
port: 443
|
||||
- host: "**.github.io"
|
||||
port: 443
|
||||
- host: pypi.org
|
||||
port: 443
|
||||
- host: www.npmjs.com
|
||||
port: 443
|
||||
- host: stackoverflow.com
|
||||
port: 443
|
||||
- host: "**.stackexchange.com"
|
||||
port: 443
|
||||
- host: "**.wikipedia.org"
|
||||
port: 443
|
||||
binaries:
|
||||
- path: /usr/bin/curl
|
||||
- path: /usr/bin/wget
|
||||
|
||||
# --- Package installation ---
|
||||
# pip install / uv add from the bash tool.
|
||||
|
||||
package_registries:
|
||||
name: package-install
|
||||
endpoints:
|
||||
- host: pypi.org
|
||||
port: 443
|
||||
- host: files.pythonhosted.org
|
||||
port: 443
|
||||
- host: "**.pypi.org"
|
||||
port: 443
|
||||
binaries:
|
||||
- path: /usr/bin/pip*
|
||||
- path: /usr/local/bin/pip*
|
||||
- path: /usr/bin/uv
|
||||
- path: /usr/local/bin/uv
|
||||
- path: /usr/bin/python3*
|
||||
- path: /usr/local/bin/python3*
|
||||
|
||||
# --- Git operations ---
|
||||
# read-only: clone, fetch, pull. No push (L7 enforcement).
|
||||
|
||||
git_operations:
|
||||
name: git-read-only
|
||||
endpoints:
|
||||
- host: github.com
|
||||
port: 443
|
||||
protocol: rest
|
||||
tls: terminate
|
||||
enforcement: enforce
|
||||
rules:
|
||||
- allow:
|
||||
method: GET
|
||||
path: "/**/info/refs*"
|
||||
- allow:
|
||||
method: POST
|
||||
path: "/**/git-upload-pack"
|
||||
- host: gitlab.com
|
||||
port: 443
|
||||
protocol: rest
|
||||
tls: terminate
|
||||
enforcement: enforce
|
||||
rules:
|
||||
- allow:
|
||||
method: GET
|
||||
path: "/**/info/refs*"
|
||||
- allow:
|
||||
method: POST
|
||||
path: "/**/git-upload-pack"
|
||||
binaries:
|
||||
- path: /usr/bin/git
|
||||
@@ -554,7 +554,7 @@ Possible `state` values:
|
||||
| `error` | An error occurred |
|
||||
|
||||
**Fan-out pattern:** Each connected client receives its own bounded queue
|
||||
(`maxsize=500`). A dedicated fan-out thread reads from the shared global queue
|
||||
(`maxsize=1000`). A dedicated fan-out thread reads from the shared global queue
|
||||
and copies each event to every client queue. If a client queue is full, the
|
||||
event is silently dropped for that client.
|
||||
|
||||
|
||||
+25
-9
@@ -91,7 +91,7 @@ turnstone/
|
||||
_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)
|
||||
katex-0.16.38/ Vendored KaTeX math rendering library (MIT, woff2 fonts)
|
||||
katex-0.16.40/ Vendored KaTeX math rendering library (MIT, woff2 fonts)
|
||||
ui/
|
||||
colors.py ANSI color constants with NO_COLOR support
|
||||
markdown.py Streaming terminal markdown renderer (line-buffered)
|
||||
@@ -100,7 +100,7 @@ turnstone/
|
||||
index.html Single-page app shell (links to CSS and JS)
|
||||
style.css Page-specific UI styles (dashboard, markdown elements, approval blocks)
|
||||
renderer.js Markdown + LaTeX renderer (tables, nested lists, blockquotes, KaTeX math)
|
||||
app.js Page-specific client-side JavaScript (SSE, workstreams, tool approval)
|
||||
app.js Split-pane UI (Pane class, binary layout tree, SSE, tool approval)
|
||||
tools/
|
||||
*.json 15 tool schemas (OpenAI function-calling format + turnstone metadata)
|
||||
```
|
||||
@@ -353,7 +353,7 @@ remove the tab immediately. Controlled by `--workstream-idle-timeout` (default:
|
||||
|
||||
**Workstream eviction at capacity:** When `WorkstreamManager.create()` would
|
||||
exceed `max_workstreams` (configurable via `[server].max_workstreams`, default
|
||||
10), the oldest IDLE workstream is automatically evicted to make room. The
|
||||
50), the oldest IDLE workstream is automatically evicted to make room. The
|
||||
`turnstone_workstreams_evicted_total` counter is incremented on each eviction.
|
||||
If no IDLE workstream is available the create request fails as before.
|
||||
|
||||
@@ -375,12 +375,19 @@ non-idle background workstreams above the input prompt.
|
||||
### Web Workstreams
|
||||
|
||||
- **Tab bar**: Each workstream renders as a tab with a colored state indicator
|
||||
(CSS `@keyframes pulse` animation per state).
|
||||
- **Per-tab SSE**: `connectContentSSE(wsId)` opens
|
||||
`/v1/api/events?ws_id=<id>` for the active tab's event stream.
|
||||
(CSS `@keyframes pulse` animation per state). Clicking a tab switches the
|
||||
focused pane's workstream (or focuses an existing pane showing that ws).
|
||||
- **Split panes**: The UI supports tiling multiple workstreams side-by-side or
|
||||
stacked via a binary layout tree. Each `Pane` instance encapsulates its own
|
||||
SSE connection, message area, input, and state (busy, approval, streaming).
|
||||
Split via right-click context menu, pane header buttons, or keyboard
|
||||
(`Ctrl+\`, `Ctrl+Shift+\`). Max 6 panes; no duplicate workstreams across panes.
|
||||
Layout persisted to `localStorage`.
|
||||
- **Per-pane SSE**: `Pane.connectSSE(wsId)` opens
|
||||
`/v1/api/events?ws_id=<id>` for each pane's event stream independently.
|
||||
- **Global SSE**: `connectGlobalSSE()` opens `/v1/api/events/global` which
|
||||
receives `ws_state` broadcasts from all workstreams, used to update tab
|
||||
indicators without switching.
|
||||
indicators and pane headers without switching.
|
||||
- **New tab / close**: POST `/v1/api/workstreams/new`, POST `/v1/api/workstreams/close`.
|
||||
|
||||
### Thread Safety
|
||||
@@ -828,10 +835,16 @@ and are the single source of truth for both backends and Alembic migrations.
|
||||
backend = "sqlite" # "sqlite" | "postgresql"
|
||||
path = ".turnstone.db" # SQLite file path
|
||||
url = "" # PostgreSQL connection URL
|
||||
pool_size = 5 # PostgreSQL connection pool size
|
||||
pool_size = 2 # PostgreSQL connection pool size (per process)
|
||||
```
|
||||
|
||||
Environment variables: `TURNSTONE_DB_BACKEND`, `TURNSTONE_DB_URL`, `TURNSTONE_DB_PATH`.
|
||||
Environment variables: `TURNSTONE_DB_BACKEND`, `TURNSTONE_DB_URL`, `TURNSTONE_DB_PATH`,
|
||||
`TURNSTONE_DB_POOL_SIZE`.
|
||||
|
||||
The default pool is intentionally small (2 base + 3 overflow = 5 per process)
|
||||
because all database operations are short-burst queries that hold connections for
|
||||
milliseconds. For clusters with many nodes sharing a PostgreSQL instance, use
|
||||
[PgBouncer](pgbouncer.md) in transaction pooling mode.
|
||||
|
||||
### Persistence and Resume
|
||||
|
||||
@@ -953,6 +966,9 @@ warns if the summary was truncated.
|
||||
seeded with `history.replaceState({turnstone: 'dashboard'})` on load. The
|
||||
`popstate` listener restores the correct tab or shows the dashboard,
|
||||
guarded by `_historyNavigation = true` to prevent re-entrant pushState.
|
||||
- **Pane focus**: `mousedown` and `focusin` events on pane containers update
|
||||
`focusedPaneId`. Approval shortcuts (y/n/a) apply to the focused pane.
|
||||
`Ctrl+Alt+Arrow` cycles focus between panes.
|
||||
|
||||
### Eval Resilience
|
||||
|
||||
|
||||
+4
-3
@@ -69,10 +69,11 @@ All reads and writes to the node/workstream map are protected by a single `threa
|
||||
|
||||
### Scale Considerations
|
||||
|
||||
- **10,000 workstreams** at ~500 bytes each = ~5 MB in memory
|
||||
- **1,000 nodes** polled in parallel with 50 threads at ~100ms each = ~2 second poll cycle
|
||||
- **50,000 workstreams** (1,000 nodes × 50 per node) at ~500 bytes each = ~25 MB in memory
|
||||
- **1,000 nodes** polled in parallel — fan-out concurrency is configurable via `cluster.node_fan_out_limit` (default 200), yielding 5 batches at ~100ms each = ~0.5 second poll cycle
|
||||
- **Filtering and pagination** run in-memory on the full workstream list — sub-millisecond at this scale
|
||||
- **SSE fan-out** uses the same per-client queue pattern as the per-node server — backed-up clients get events dropped, not blocking
|
||||
- **SSE fan-out** uses per-client queues (2,000 events) — backed-up clients get events dropped, not blocking
|
||||
- **Database** — for clusters sharing PostgreSQL, use [PgBouncer](pgbouncer.md) in transaction pooling mode
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -105,7 +105,7 @@ 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()
|
||||
note right : Per-client queue.Queue(maxsize=2000)\nSSE via EventSourceResponse + run_in_executor()
|
||||
|
||||
Server -> Browser : data: {"type":"snapshot",...}\n(full state as first SSE event)
|
||||
|
||||
|
||||
@@ -56,6 +56,26 @@ node "Docker Host" as host {
|
||||
end note
|
||||
}
|
||||
|
||||
node "postgres (profile: production)" <<pgautoupgrade>> as pg_node {
|
||||
component [PostgreSQL\nport 5432] as postgres
|
||||
note bottom of postgres
|
||||
Healthcheck: pg_isready
|
||||
Volume: postgres-data
|
||||
Required for cluster
|
||||
and production profiles
|
||||
end note
|
||||
}
|
||||
|
||||
node "pgbouncer (optional)" <<bitnami/pgbouncer>> as pgb_node {
|
||||
component [PgBouncer\nport 6432] as pgbouncer
|
||||
note bottom of pgbouncer
|
||||
pool_mode: transaction
|
||||
Recommended for clusters
|
||||
> 50 nodes
|
||||
See docs/pgbouncer.md
|
||||
end note
|
||||
}
|
||||
|
||||
node "sim (profile: sim)" <<turnstone image>> as sim_node {
|
||||
component [turnstone-sim] as sim
|
||||
note bottom of sim
|
||||
@@ -92,6 +112,11 @@ console --> server : HTTP polling + proxy\n(GET /v1/api/dashboard,\nproxy /node/
|
||||
|
||||
sim --> redis : Redis protocol\n(queues + pubsub + keys)
|
||||
|
||||
' Database connections (production/cluster profiles)
|
||||
server ..> pgbouncer : PostgreSQL\n(pool_size=2)
|
||||
console ..> pgbouncer : PostgreSQL\n(auth/admin)
|
||||
pgbouncer --> postgres : transaction\npooling
|
||||
|
||||
' Environment variables
|
||||
note right of host
|
||||
**Environment Variables:**
|
||||
@@ -99,13 +124,17 @@ note right of host
|
||||
• OPENAI_API_KEY — API key
|
||||
• REDIS_PASSWORD — Redis auth
|
||||
• TURNSTONE_AUTH_TOKEN — API auth
|
||||
• TURNSTONE_DB_URL — PostgreSQL URL
|
||||
• POSTGRES_PASSWORD — DB password
|
||||
end note
|
||||
|
||||
' Volumes
|
||||
database "redis-data" as rv
|
||||
database "turnstone-data" as tv
|
||||
database "postgres-data" as pv
|
||||
|
||||
redis_node --> rv
|
||||
server_node --> tv
|
||||
pg_node --> pv
|
||||
|
||||
@enduml
|
||||
|
||||
@@ -53,10 +53,10 @@ class "SQLiteBackend" as SQLite <<sqlite>> {
|
||||
|
||||
class "PostgreSQLBackend" as PG <<postgres>> {
|
||||
-_engine: sa.Engine
|
||||
+__init__(url: str, pool_size: int)
|
||||
+__init__(url: str, pool_size: int = 2,\n max_overflow: int = 3)
|
||||
--
|
||||
tsvector + ILIKE search
|
||||
Connection pooling
|
||||
Connection pooling (5 max per process)
|
||||
}
|
||||
|
||||
' -- Schema --
|
||||
@@ -151,7 +151,7 @@ note right of Registry
|
||||
backend = "sqlite" | "postgresql"
|
||||
url = "postgresql+psycopg://..."
|
||||
path = ".turnstone.db"
|
||||
pool_size = 5
|
||||
pool_size = 2 (+ 3 overflow)
|
||||
end note
|
||||
|
||||
note bottom of SQLite
|
||||
@@ -162,8 +162,9 @@ end note
|
||||
|
||||
note bottom of PG
|
||||
Production backend.
|
||||
Multi-node / Docker
|
||||
default.
|
||||
Multi-node / Docker default.
|
||||
Use PgBouncer (transaction mode)
|
||||
for clusters > 50 nodes.
|
||||
end note
|
||||
|
||||
@enduml
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:a74b4b8b5dbfb1a51a01100b731477968942b01218bad9451a3d5a9cb3003294
|
||||
size 411665
|
||||
oid sha256:e3f1ad0fcd55eaca3b8ad9c5abc07432803641c54ede9fc93c79df144cf77d1c
|
||||
size 407761
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:84524f4bc900708ac8adf081591d336f862830188eb8505e71a0f071b339d923
|
||||
size 252599
|
||||
oid sha256:09065fef028d05e6df425fd8abefaf5a2ca04b66802f2e3975f289fa597f63ed
|
||||
size 309656
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:c94556889abb382cd5b818639fc0a4706beef3d9c7a0b4cbedc763943d657dd0
|
||||
size 244998
|
||||
oid sha256:b047cdc318c505f0f0895a65e14c5cc7552716053055cca57fa0a77db150e618
|
||||
size 255458
|
||||
|
||||
@@ -109,9 +109,12 @@ All configuration is via environment variables in `.env` (copy from `.env.exampl
|
||||
|----------|---------|-------------|
|
||||
| `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` |
|
||||
| `TURNSTONE_DB_POOL_SIZE` | `2` | PostgreSQL connection pool size per process (default: 2 base + 3 overflow = 5 max) |
|
||||
|
||||
The database stores workstream history, user accounts, and API tokens. When using JWT auth, a database backend is required for user storage.
|
||||
|
||||
> **Large clusters:** Each turnstone process maintains a small connection pool (5 max). At hundreds of nodes this adds up — use [PgBouncer](pgbouncer.md) in transaction pooling mode between turnstone and PostgreSQL.
|
||||
|
||||
> **First-time setup:** After deploying with auth enabled, create an initial admin user by running `turnstone-admin create-user` inside the container:
|
||||
>
|
||||
> ```bash
|
||||
@@ -151,6 +154,8 @@ POSTGRES_PASSWORD=secret docker compose --profile cluster up
|
||||
|
||||
The default `server` and `bridge` also run alongside the cluster nodes (11 total). All nodes are accessible via the console dashboard at `:8090`.
|
||||
|
||||
For production clusters beyond ~50 nodes, add PgBouncer between turnstone services and PostgreSQL. See [PgBouncer Connection Pooling](pgbouncer.md) for Docker Compose and Helm configuration.
|
||||
|
||||
## Volumes
|
||||
|
||||
| Volume | Mount | Purpose |
|
||||
|
||||
@@ -0,0 +1,288 @@
|
||||
# OpenShell Sandbox Integration
|
||||
|
||||
Turnstone can run inside an [OpenShell](https://github.com/NVIDIA/OpenShell)
|
||||
sandbox for kernel-enforced security boundaries around tool execution. OpenShell
|
||||
provides four layers of defense that Turnstone's application-level safety model
|
||||
does not cover:
|
||||
|
||||
| Layer | Mechanism | What it prevents |
|
||||
|-------|-----------|------------------|
|
||||
| Filesystem | Landlock | Writes to `/etc`, `~/.ssh`, system paths |
|
||||
| Network | Network namespace + seccomp + HTTP CONNECT proxy | Connections to unlisted hosts |
|
||||
| Process | `setuid` drop + verification | Privilege escalation to root |
|
||||
| Credentials | Proxy-level secret resolution | API keys in sandbox memory |
|
||||
|
||||
Turnstone's own safety layers (human approval, intent judge, tool policies,
|
||||
output guard) remain active inside the sandbox and handle threats at the semantic
|
||||
level -- what the LLM *means* to do with its legitimate access.
|
||||
|
||||
> See also: [Security and Authentication](security.md),
|
||||
> [Intent Validation](judge.md), [Governance](governance.md)
|
||||
|
||||
---
|
||||
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
# Run turnstone-server in an OpenShell sandbox
|
||||
openshell sandbox run \
|
||||
--policy deploy/openshell/turnstone-policy.yaml \
|
||||
--workdir /path/to/project \
|
||||
-- python3 -m turnstone.server --host 0.0.0.0 --port 8080
|
||||
```
|
||||
|
||||
With inference routing (API keys never enter the sandbox):
|
||||
|
||||
```bash
|
||||
openshell sandbox run \
|
||||
--policy deploy/openshell/turnstone-policy.yaml \
|
||||
--inference-routes deploy/openshell/routes.yaml \
|
||||
--workdir /path/to/project \
|
||||
-- python3 -m turnstone.server --host 0.0.0.0 --port 8080 \
|
||||
--base-url https://inference.local
|
||||
```
|
||||
|
||||
The `inference.local` hostname is intercepted by the OpenShell proxy before
|
||||
network policy evaluation -- no network policy entry is needed for it.
|
||||
|
||||
---
|
||||
|
||||
## Policy Files
|
||||
|
||||
### `deploy/openshell/turnstone-policy.yaml`
|
||||
|
||||
The main sandbox policy. Covers filesystem, process, and network rules.
|
||||
|
||||
### `deploy/openshell/routes.yaml`
|
||||
|
||||
Inference routing configuration. Maps `inference.local` to real LLM API
|
||||
backends. Uncomment and configure the provider(s) you use.
|
||||
|
||||
---
|
||||
|
||||
## Filesystem Policy
|
||||
|
||||
The policy uses Landlock (Linux 5.13+) for kernel-enforced filesystem access
|
||||
control. Paths are locked at sandbox creation and cannot be changed at runtime.
|
||||
|
||||
| Path | Access | Purpose |
|
||||
|------|--------|---------|
|
||||
| `--workdir` | read-write | Project files (auto-added via `include_workdir`) |
|
||||
| `/tmp` | read-write | Bash tool temp scripts, eval workdirs |
|
||||
| `/dev/null` | read-write | Shell redirections (`2>/dev/null`) |
|
||||
| `/var/log` | read-write | Log files |
|
||||
| `/usr`, `/lib`, `/lib64` | read-only | Python runtime, installed packages |
|
||||
| `/etc` | read-only | System config, SSL certificates |
|
||||
| `/proc`, `/dev/urandom` | read-only | Process info, entropy |
|
||||
| `~/.config/turnstone` | read-only | Config file (writes go to database) |
|
||||
|
||||
Landlock runs in `best_effort` mode by default -- degrades gracefully on kernels
|
||||
without Landlock support. Set `compatibility: hard_requirement` for production
|
||||
hardened deployments.
|
||||
|
||||
---
|
||||
|
||||
## Network Policy
|
||||
|
||||
Default-deny. Only explicitly listed host:port pairs are reachable. All child
|
||||
processes (MCP servers, bash commands, grep) inherit the network namespace and
|
||||
cannot bypass the proxy.
|
||||
|
||||
### Included endpoints
|
||||
|
||||
| Policy | Hosts | Purpose |
|
||||
|--------|-------|---------|
|
||||
| `openai_api` | `api.openai.com` | OpenAI LLM API |
|
||||
| `anthropic_api` | `api.anthropic.com` | Anthropic LLM API |
|
||||
| `tavily_api` | `api.tavily.com` | Web search fallback |
|
||||
| `skills_registry` | `skills.sh` | Skill discovery |
|
||||
| `github_api` | `api.github.com` (read-only L7), `raw.githubusercontent.com` | Skill fetch, GitHub API |
|
||||
| `mcp_registry` | `registry.modelcontextprotocol.io` (read-only L7) | MCP server discovery |
|
||||
| `redis` | `127.0.0.1:6379` | Message queue |
|
||||
| `web_fetch_common` | readthedocs, python docs, GitHub Pages, PyPI, npm, Stack Overflow, Wikipedia | Curated web_fetch domains |
|
||||
| `bash_network_tools` | Same as `web_fetch_common` | curl/wget from bash tool |
|
||||
| `package_registries` | `pypi.org`, `files.pythonhosted.org` | pip/uv package installs |
|
||||
| `git_operations` | `github.com`, `gitlab.com` (L7: clone/fetch only, no push) | Git read-only operations |
|
||||
|
||||
### L7 enforcement
|
||||
|
||||
Endpoints marked with `protocol: rest` and `tls: terminate` get HTTP-level
|
||||
inspection. The proxy TLS-terminates using an ephemeral per-sandbox CA, parses
|
||||
each request, and evaluates method + path against the rules.
|
||||
|
||||
The `github_api`, `mcp_registry`, and `git_operations` policies use L7
|
||||
enforcement:
|
||||
|
||||
- **GitHub API / MCP Registry**: `access: read-only` -- only GET, HEAD, OPTIONS
|
||||
allowed
|
||||
- **Git operations**: explicit rules allowing only `info/refs` (GET) and
|
||||
`git-upload-pack` (POST) -- clone and fetch work, push is blocked
|
||||
|
||||
### Commented-out sections
|
||||
|
||||
The policy includes commented blocks for optional integrations. Uncomment and
|
||||
configure as needed:
|
||||
|
||||
- **OIDC** -- add your identity provider's hostname
|
||||
- **Discord** -- `discord.com`, `gateway.discord.gg`, `cdn.discordapp.com`
|
||||
- **MCP HTTP servers** -- any MCP servers using streamable-http transport
|
||||
|
||||
---
|
||||
|
||||
## Customizing the Domain Allowlist
|
||||
|
||||
The `web_fetch` tool lets the LLM fetch arbitrary public URLs, but OpenShell
|
||||
cannot allow "all HTTPS" -- bare wildcard hosts are rejected by policy
|
||||
validation. Instead, the policy ships with a curated set of common reference
|
||||
domains.
|
||||
|
||||
To add domains your workloads need:
|
||||
|
||||
```yaml
|
||||
# In turnstone-policy.yaml, under web_fetch_common.endpoints:
|
||||
- host: docs.example.com
|
||||
port: 443
|
||||
|
||||
# Also add to bash_network_tools.endpoints if curl/wget should reach it:
|
||||
- host: docs.example.com
|
||||
port: 443
|
||||
```
|
||||
|
||||
Wildcard patterns are supported:
|
||||
|
||||
- `*.example.com` -- matches one subdomain level (e.g. `api.example.com`)
|
||||
- `**.example.com` -- matches any depth (e.g. `deep.sub.example.com`)
|
||||
|
||||
Unlisted domains return connection errors, which the LLM handles gracefully by
|
||||
telling the user it cannot reach that site.
|
||||
|
||||
---
|
||||
|
||||
## Inference Routing
|
||||
|
||||
Inference routing keeps real API keys completely outside the sandbox. The
|
||||
sandbox process only sees opaque placeholder tokens in its environment
|
||||
(`openshell:resolve:env:ANTHROPIC_API_KEY`). The proxy rewrites these to real
|
||||
credentials on the wire before forwarding to the upstream API.
|
||||
|
||||
### Setup
|
||||
|
||||
1. Edit `deploy/openshell/routes.yaml` -- uncomment your provider:
|
||||
|
||||
```yaml
|
||||
routes:
|
||||
# OpenAI
|
||||
- name: inference.local
|
||||
endpoint: https://api.openai.com/v1
|
||||
model: gpt-5
|
||||
provider_type: openai
|
||||
protocols:
|
||||
- openai_chat_completions
|
||||
- model_discovery
|
||||
api_key_env: OPENAI_API_KEY
|
||||
|
||||
# Or Anthropic
|
||||
- name: inference.local
|
||||
endpoint: https://api.anthropic.com
|
||||
model: claude-sonnet-4-6
|
||||
provider_type: anthropic
|
||||
protocols:
|
||||
- anthropic_messages
|
||||
api_key_env: ANTHROPIC_API_KEY
|
||||
```
|
||||
|
||||
2. Start with `--inference-routes` and point turnstone at `inference.local`:
|
||||
|
||||
```bash
|
||||
openshell sandbox run \
|
||||
--inference-routes deploy/openshell/routes.yaml \
|
||||
--base-url https://inference.local \
|
||||
...
|
||||
```
|
||||
|
||||
3. When inference routing is active, the `openai_api` and `anthropic_api`
|
||||
network policies can be removed from the sandbox policy -- the proxy handles
|
||||
LLM traffic on a separate code path that bypasses OPA entirely.
|
||||
|
||||
### Local model servers
|
||||
|
||||
For local servers (vLLM, llama.cpp) with no authentication, omit both
|
||||
`api_key` and `api_key_env` from the route config. No credential resolution
|
||||
is needed.
|
||||
|
||||
---
|
||||
|
||||
## MCP Server Subprocesses
|
||||
|
||||
MCP servers using stdio transport are spawned as child processes of turnstone.
|
||||
They automatically inherit all sandbox constraints:
|
||||
|
||||
- **Network namespace** -- kernel-level, cannot be bypassed
|
||||
- **Landlock filesystem** -- kernel-level, cannot be relaxed
|
||||
- **Seccomp socket filter** -- kernel-level, inherited on fork
|
||||
|
||||
No per-subprocess policy entries are needed for these constraints. However, if
|
||||
an MCP server makes outbound network requests (through the proxy), its binary
|
||||
must appear in a `binaries[]` entry for the relevant network policy. The proxy
|
||||
identifies the requesting process via `/proc/<pid>/exe` (not `argv[0]`, which
|
||||
is spoofable).
|
||||
|
||||
Example for a Python-based MCP server that calls an external API:
|
||||
|
||||
```yaml
|
||||
mcp_external_api:
|
||||
name: mcp-external
|
||||
endpoints:
|
||||
- host: api.example.com
|
||||
port: 443
|
||||
binaries:
|
||||
- path: /usr/bin/python3*
|
||||
- path: /usr/local/bin/python3*
|
||||
```
|
||||
|
||||
MCP servers using streamable-http transport are remote -- they need a network
|
||||
policy entry for their host:port but no binary entry (the Python process making
|
||||
the HTTP call is already covered by the standard `python3*` binary entries).
|
||||
|
||||
---
|
||||
|
||||
## Security Model: Which Layer Enforces What
|
||||
|
||||
```
|
||||
OpenShell (infrastructure) Turnstone (application)
|
||||
───────────────────────────── ──────────────────────────────
|
||||
Filesystem access Landlock kernel enforcement (no enforcement)
|
||||
Network egress Netns + seccomp + proxy + OPA SSRF check on web_fetch
|
||||
Credentials Placeholder injection + proxy Output guard redaction
|
||||
Privilege level setuid drop + verification (no enforcement)
|
||||
Tool semantics (no visibility) Heuristic + LLM judge
|
||||
Tool policies (no visibility) fnmatch admin policies
|
||||
Prompt injection (no visibility) Output guard detection
|
||||
Human approval (no visibility) Approval gate + "always"
|
||||
```
|
||||
|
||||
OpenShell constrains what the process can physically reach. Turnstone constrains
|
||||
what the LLM does with its legitimate access. Neither layer is sufficient alone:
|
||||
|
||||
- Without OpenShell: a bash command can `curl` secrets to any endpoint, write to
|
||||
`/etc/crontab`, or read `~/.ssh/id_rsa` -- all gated only by human approval
|
||||
- Without Turnstone: the LLM can `rm -rf` the entire workdir, run destructive
|
||||
commands, or consume prompt injection payloads -- all within the sandbox's
|
||||
allowed scope
|
||||
|
||||
---
|
||||
|
||||
## Hardening Checklist
|
||||
|
||||
For production deployments:
|
||||
|
||||
- [ ] Set `landlock.compatibility: hard_requirement`
|
||||
- [ ] Enable inference routing (removes API keys from sandbox)
|
||||
- [ ] Remove `openai_api`/`anthropic_api` network policies when using inference
|
||||
routing (traffic goes through the router, not direct)
|
||||
- [ ] Review and trim `web_fetch_common` domains to your actual needs
|
||||
- [ ] Remove `package_registries` policy if pip/uv installs are not needed
|
||||
- [ ] Add your OIDC provider endpoint if using SSO
|
||||
- [ ] Set Redis `allowed_ips` to your actual Redis host if not localhost
|
||||
- [ ] Consider removing `bash_network_tools` entirely if bash should not have
|
||||
network access
|
||||
@@ -0,0 +1,200 @@
|
||||
# PgBouncer Connection Pooling
|
||||
|
||||
Turnstone cluster deployments share a single PostgreSQL instance across
|
||||
all server nodes, bridge processes, and the console. Each process
|
||||
maintains a small connection pool (2 base + 3 overflow = 5 max). At
|
||||
scale this adds up — a 100-node cluster opens up to 500 connections,
|
||||
and a 1000-node cluster up to 5,000.
|
||||
|
||||
PostgreSQL's default `max_connections` is 100, and each real connection
|
||||
allocates ~5–10 MB of backend memory. PgBouncer sits between turnstone
|
||||
and PostgreSQL, multiplexing thousands of lightweight client connections
|
||||
down to a small number of real database connections.
|
||||
|
||||
---
|
||||
|
||||
## Why PgBouncer works well with turnstone
|
||||
|
||||
All turnstone database operations are short-burst queries: acquire a
|
||||
connection, execute 1–3 statements, commit, release. No operation holds
|
||||
a connection for more than a few milliseconds. This makes **transaction
|
||||
pooling mode** ideal — PgBouncer assigns a real connection only for the
|
||||
duration of each transaction, then returns it to the pool.
|
||||
|
||||
| Cluster size | Client connections (max) | PgBouncer server connections needed |
|
||||
|--------------|------------------------|-------------------------------------|
|
||||
| 10 nodes | 50 | 10–20 |
|
||||
| 100 nodes | 500 | 20–40 |
|
||||
| 500 nodes | 2,500 | 30–60 |
|
||||
| 1,000 nodes | 5,000 | 40–80 |
|
||||
|
||||
The server connection count stays low because most client connections
|
||||
are idle at any given moment.
|
||||
|
||||
---
|
||||
|
||||
## Docker Compose
|
||||
|
||||
Add PgBouncer between turnstone services and PostgreSQL:
|
||||
|
||||
```yaml
|
||||
services:
|
||||
pgbouncer:
|
||||
image: bitnami/pgbouncer:latest
|
||||
environment:
|
||||
POSTGRESQL_HOST: postgres
|
||||
POSTGRESQL_PORT: "5432"
|
||||
POSTGRESQL_DATABASE: turnstone
|
||||
POSTGRESQL_USERNAME: ${POSTGRES_USER:-turnstone}
|
||||
POSTGRESQL_PASSWORD: ${POSTGRES_PASSWORD:?}
|
||||
PGBOUNCER_POOL_MODE: transaction
|
||||
PGBOUNCER_DEFAULT_POOL_SIZE: "40"
|
||||
PGBOUNCER_MAX_CLIENT_CONN: "5000"
|
||||
PGBOUNCER_MAX_DB_CONNECTIONS: "80"
|
||||
PGBOUNCER_SERVER_IDLE_TIMEOUT: "300"
|
||||
ports:
|
||||
- "6432:6432"
|
||||
networks:
|
||||
- turnstone-net
|
||||
depends_on:
|
||||
postgres:
|
||||
condition: service_healthy
|
||||
healthcheck:
|
||||
test: ["CMD", "pg_isready", "-h", "127.0.0.1", "-p", "6432"]
|
||||
interval: 5s
|
||||
timeout: 3s
|
||||
retries: 5
|
||||
```
|
||||
|
||||
Then point turnstone services at PgBouncer instead of PostgreSQL
|
||||
directly by changing the `DATABASE_URL` (or `TURNSTONE_DB_URL`):
|
||||
|
||||
```bash
|
||||
# Before (direct)
|
||||
TURNSTONE_DB_URL=postgresql://turnstone:secret@postgres:5432/turnstone
|
||||
|
||||
# After (via PgBouncer)
|
||||
TURNSTONE_DB_URL=postgresql://turnstone:secret@pgbouncer:6432/turnstone
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Helm / Kubernetes
|
||||
|
||||
Add a PgBouncer deployment or use a Helm chart like
|
||||
[bitnami/pgbouncer](https://github.com/bitnami/charts/tree/main/bitnami/pgbouncer).
|
||||
|
||||
In `values.yaml`, point the database at PgBouncer:
|
||||
|
||||
```yaml
|
||||
database:
|
||||
backend: postgresql
|
||||
external:
|
||||
host: pgbouncer
|
||||
port: 6432
|
||||
database: turnstone
|
||||
username: turnstone
|
||||
existingSecret: turnstone-db-secret
|
||||
```
|
||||
|
||||
PgBouncer configuration:
|
||||
|
||||
```yaml
|
||||
pgbouncer:
|
||||
poolMode: transaction
|
||||
defaultPoolSize: 40
|
||||
maxClientConn: 5000
|
||||
maxDbConnections: 80
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Configuration reference
|
||||
|
||||
| PgBouncer setting | Recommended | Notes |
|
||||
|-------------------|-------------|-------|
|
||||
| `pool_mode` | `transaction` | Required — turnstone uses short-burst queries with no session state |
|
||||
| `default_pool_size` | 40 | Real PostgreSQL connections per database. Start here, increase if you see `no more connections allowed` |
|
||||
| `max_client_conn` | 5000 | Upper bound on client connections. Set to `cluster_nodes × 5` |
|
||||
| `max_db_connections` | 80 | Hard cap on real connections to PostgreSQL. Keep below PG `max_connections` minus headroom for admin/monitoring |
|
||||
| `server_idle_timeout` | 300 | Close idle server connections after 5 minutes |
|
||||
| `server_lifetime` | 3600 | Recycle server connections after 1 hour |
|
||||
|
||||
On the PostgreSQL side:
|
||||
|
||||
| PostgreSQL setting | Recommended | Notes |
|
||||
|--------------------|-------------|-------|
|
||||
| `max_connections` | 100 | Default is fine — PgBouncer is the only client. Set higher than `max_db_connections` to leave room for admin connections |
|
||||
| `shared_buffers` | 25% of RAM | Standard PostgreSQL tuning |
|
||||
|
||||
---
|
||||
|
||||
## Turnstone pool settings
|
||||
|
||||
Each turnstone process maintains its own SQLAlchemy connection pool to
|
||||
PgBouncer (which then multiplexes to PostgreSQL):
|
||||
|
||||
| Environment variable | Default | Description |
|
||||
|---------------------|---------|-------------|
|
||||
| `TURNSTONE_DB_POOL_SIZE` | 2 | Base pool size per process |
|
||||
| `TURNSTONE_DB_BACKEND` | sqlite | Set to `postgresql` for cluster deployments |
|
||||
| `TURNSTONE_DB_URL` | — | Connection URL (point at PgBouncer, not PostgreSQL directly) |
|
||||
|
||||
The default pool of 2 + 3 overflow = 5 connections per process is
|
||||
intentionally small to support large clusters. You should not need to
|
||||
increase this — turnstone's database operations are all short-burst
|
||||
context-managed queries that hold connections for milliseconds.
|
||||
|
||||
SQLAlchemy `pool_pre_ping` is enabled, so stale connections (e.g. after
|
||||
PgBouncer restarts) are automatically detected and replaced.
|
||||
|
||||
---
|
||||
|
||||
## Monitoring
|
||||
|
||||
PgBouncer exposes stats via its admin console (connect to
|
||||
PgBouncer port with user `pgbouncer`):
|
||||
|
||||
```sql
|
||||
-- Active and waiting clients
|
||||
SHOW POOLS;
|
||||
|
||||
-- Per-database stats
|
||||
SHOW STATS;
|
||||
|
||||
-- Current client connections
|
||||
SHOW CLIENTS;
|
||||
```
|
||||
|
||||
Key metrics to watch:
|
||||
|
||||
- **`cl_active`** — clients with a server connection assigned. Should be
|
||||
well below `max_db_connections`.
|
||||
- **`cl_waiting`** — clients waiting for a server connection. Sustained
|
||||
non-zero values mean you need more `default_pool_size`.
|
||||
- **`sv_active`** — active server (PostgreSQL) connections. Should stay
|
||||
below PostgreSQL `max_connections`.
|
||||
|
||||
---
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
**"no more connections allowed (max_client_conn)"** — PgBouncer is
|
||||
rejecting new client connections. Increase `max_client_conn` to match
|
||||
your cluster size × 5.
|
||||
|
||||
**"no more connections allowed (max_db_connections)"** — PgBouncer
|
||||
cannot open more connections to PostgreSQL. Increase
|
||||
`max_db_connections` and ensure PostgreSQL `max_connections` is higher.
|
||||
|
||||
**Connections timing out on startup** — If all nodes start
|
||||
simultaneously, the burst of initial connections (migrations, health
|
||||
checks) can temporarily exceed the pool. PgBouncer queues excess
|
||||
clients by default — this resolves itself within seconds.
|
||||
|
||||
**Prepared statements not supported** — PgBouncer in `transaction` mode
|
||||
does not support prepared statements. Turnstone's SQLAlchemy layer does
|
||||
not use server-side prepared statements by default, so this is not an
|
||||
issue.
|
||||
|
||||
See also: [Docker deployment](docker.md) · [Security](security.md)
|
||||
+5
-3
@@ -51,7 +51,7 @@ connection, Redis, auth secrets, server bind address). These stay in
|
||||
| Bridge identity | `[bridge]` | config.toml / env |
|
||||
| Console bind | `[console]` | config.toml / env |
|
||||
|
||||
**ConfigStore settings** (~40 settings) are loaded from the database after
|
||||
**ConfigStore settings** (48 settings) are loaded from the database after
|
||||
storage initialization:
|
||||
|
||||
| Section | Settings |
|
||||
@@ -60,10 +60,12 @@ storage initialization:
|
||||
| `session` | instructions, retention_days, compact_max_tokens, auto_compact_pct |
|
||||
| `tools` | timeout, truncation, agent_max_turns, skip_permissions, search, search_threshold, search_max_results |
|
||||
| `server` | workstream_idle_timeout, max_workstreams |
|
||||
| `cluster` | node_fan_out_limit, mcp_max_servers |
|
||||
| `mcp` | config_path, refresh_interval, registry_url |
|
||||
| `ratelimit` | enabled, requests_per_second, burst |
|
||||
| `ratelimit` | enabled, requests_per_second, burst, trusted_proxies |
|
||||
| `health` | backend_probe_interval, backend_probe_timeout, circuit_breaker_threshold, circuit_breaker_cooldown |
|
||||
| `judge` | enabled, model, provider, base_url, api_key, confidence_threshold, max_context_ratio, timeout, read_only_tools |
|
||||
| `judge` | enabled, model, provider, base_url, api_key, confidence_threshold, max_context_ratio, timeout, read_only_tools, output_guard, redact_secrets |
|
||||
| `skills` | discovery_url |
|
||||
| `memory` | relevance_k, fetch_limit, max_content, nudge_cooldown, nudges |
|
||||
|
||||
Settings are addressed by dotted key (e.g. `memory.relevance_k`). Each has a
|
||||
|
||||
+4
-3
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "turnstone"
|
||||
version = "0.8.3"
|
||||
version = "0.8.5"
|
||||
description = "Multi-node AI orchestration platform with tool use, agent routing, and cluster simulation."
|
||||
readme = "README.md"
|
||||
license = "BUSL-1.1"
|
||||
@@ -51,8 +51,9 @@ console = ["redis>=7.2", "croniter>=3.0"]
|
||||
sim = ["redis>=7.2"]
|
||||
anthropic = ["anthropic>=0.39"]
|
||||
postgres = ["psycopg[binary]>=3.2"]
|
||||
ddg = ["duckduckgo-search>=8.0"]
|
||||
discord = ["discord.py>=2.4", "redis>=7.2"]
|
||||
all = ["turnstone[mq,console,sim,anthropic,postgres,discord]"]
|
||||
all = ["turnstone[mq,console,sim,anthropic,postgres,discord,ddg]"]
|
||||
|
||||
[project.scripts]
|
||||
turnstone = "turnstone.cli:main"
|
||||
@@ -77,7 +78,7 @@ include = [
|
||||
"turnstone/console/static/*.js",
|
||||
"turnstone/shared_static/*.css",
|
||||
"turnstone/shared_static/*.js",
|
||||
"turnstone/shared_static/katex-0.16.38/**/*",
|
||||
"turnstone/shared_static/katex-0.16.40/**/*",
|
||||
"turnstone/shared_static/hljs-11.11.1/**/*",
|
||||
"turnstone/shared_static/mermaid-11.13.0/**/*",
|
||||
"turnstone/sdk/py.typed",
|
||||
|
||||
@@ -28,9 +28,19 @@ usage() {
|
||||
LIB="$1"
|
||||
VERSION="$2"
|
||||
|
||||
# Detect current version from pyproject.toml
|
||||
# Detect current version from the filesystem (not pyproject.toml, which
|
||||
# Renovate may have already updated). Falls back to pyproject.toml if
|
||||
# no directory is found.
|
||||
detect_old_version() {
|
||||
local pattern="$1"
|
||||
# Look for existing directory: e.g. turnstone/shared_static/katex-0.16.38
|
||||
local dir
|
||||
dir=$(find "${STATIC_DIR}" -maxdepth 1 -type d -name "${pattern}-*" | head -1)
|
||||
if [[ -n "$dir" ]]; then
|
||||
basename "$dir" | sed "s/${pattern}-//"
|
||||
return
|
||||
fi
|
||||
# Fallback to pyproject.toml
|
||||
grep -oE "${pattern}-[0-9.]+" pyproject.toml | head -1 | sed "s/${pattern}-//"
|
||||
}
|
||||
|
||||
@@ -51,9 +61,19 @@ update_refs() {
|
||||
done
|
||||
}
|
||||
|
||||
check_same_version() {
|
||||
if [[ "$1" == "$2" ]]; then
|
||||
echo "ERROR: Old version ($1) == new version ($2). Nothing to update."
|
||||
echo "If the old directory was already removed, re-download with:"
|
||||
echo " rm -rf ${STATIC_DIR}/${3}-${1} && $0 $3 $2"
|
||||
exit 1
|
||||
fi
|
||||
}
|
||||
|
||||
case "$LIB" in
|
||||
katex)
|
||||
OLD_VERSION=$(detect_old_version "katex")
|
||||
check_same_version "$OLD_VERSION" "$VERSION" "katex"
|
||||
OLD_DIR="${STATIC_DIR}/katex-${OLD_VERSION}"
|
||||
NEW_DIR="${STATIC_DIR}/katex-${VERSION}"
|
||||
|
||||
@@ -87,6 +107,7 @@ case "$LIB" in
|
||||
|
||||
hljs)
|
||||
OLD_VERSION=$(detect_old_version "hljs")
|
||||
check_same_version "$OLD_VERSION" "$VERSION" "hljs"
|
||||
OLD_DIR="${STATIC_DIR}/hljs-${OLD_VERSION}"
|
||||
NEW_DIR="${STATIC_DIR}/hljs-${VERSION}"
|
||||
|
||||
@@ -107,6 +128,7 @@ case "$LIB" in
|
||||
|
||||
mermaid)
|
||||
OLD_VERSION=$(detect_old_version "mermaid")
|
||||
check_same_version "$OLD_VERSION" "$VERSION" "mermaid"
|
||||
OLD_DIR="${STATIC_DIR}/mermaid-${OLD_VERSION}"
|
||||
NEW_DIR="${STATIC_DIR}/mermaid-${VERSION}"
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"openapi": "3.1.0",
|
||||
"info": {
|
||||
"title": "turnstone Console API",
|
||||
"version": "0.8.2",
|
||||
"version": "0.8.4",
|
||||
"description": "Cluster-wide visibility and control across all turnstone nodes."
|
||||
},
|
||||
"paths": {
|
||||
@@ -4378,6 +4378,12 @@
|
||||
"description": "Skill name (replaces default skills)",
|
||||
"title": "Skill",
|
||||
"type": "string"
|
||||
},
|
||||
"resume_ws": {
|
||||
"default": "",
|
||||
"description": "Workstream ID to resume (loads previous conversation)",
|
||||
"title": "Resume Ws",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"title": "ConsoleCreateWsRequest",
|
||||
@@ -6878,6 +6884,11 @@
|
||||
"title": "Enabled",
|
||||
"type": "boolean"
|
||||
},
|
||||
"priority": {
|
||||
"default": 0,
|
||||
"title": "Priority",
|
||||
"type": "integer"
|
||||
},
|
||||
"allowed_tools": {
|
||||
"default": "[]",
|
||||
"title": "Allowed Tools",
|
||||
@@ -7113,6 +7124,11 @@
|
||||
"title": "Enabled",
|
||||
"type": "boolean"
|
||||
},
|
||||
"priority": {
|
||||
"default": 0,
|
||||
"title": "Priority",
|
||||
"type": "integer"
|
||||
},
|
||||
"allowed_tools": {
|
||||
"default": "[]",
|
||||
"title": "Allowed Tools",
|
||||
@@ -7366,6 +7382,18 @@
|
||||
"default": null,
|
||||
"title": "Enabled"
|
||||
},
|
||||
"priority": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "integer"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"default": null,
|
||||
"title": "Priority"
|
||||
},
|
||||
"allowed_tools": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"openapi": "3.1.0",
|
||||
"info": {
|
||||
"title": "turnstone Server API",
|
||||
"version": "0.8.2",
|
||||
"version": "0.8.4",
|
||||
"description": "Single-node workstream management, chat interaction, and real-time streaming."
|
||||
},
|
||||
"paths": {
|
||||
@@ -1559,6 +1559,11 @@
|
||||
"title": "Version",
|
||||
"type": "string"
|
||||
},
|
||||
"node_id": {
|
||||
"default": "",
|
||||
"title": "Node Id",
|
||||
"type": "string"
|
||||
},
|
||||
"uptime_seconds": {
|
||||
"default": 0.0,
|
||||
"title": "Uptime Seconds",
|
||||
@@ -1569,6 +1574,12 @@
|
||||
"title": "Model",
|
||||
"type": "string"
|
||||
},
|
||||
"max_ws": {
|
||||
"default": 10,
|
||||
"description": "Maximum concurrent workstreams",
|
||||
"title": "Max Ws",
|
||||
"type": "integer"
|
||||
},
|
||||
"workstreams": {
|
||||
"$ref": "#/components/schemas/WorkstreamCounts",
|
||||
"default": {
|
||||
|
||||
Generated
+131
-142
@@ -9,14 +9,14 @@
|
||||
"version": "0.3.0",
|
||||
"license": "BUSL-1.1",
|
||||
"devDependencies": {
|
||||
"typescript": "^5.4",
|
||||
"typescript": "^6.0.0",
|
||||
"vitest": "^4.1"
|
||||
}
|
||||
},
|
||||
"node_modules/@emnapi/core": {
|
||||
"version": "1.9.0",
|
||||
"resolved": "https://registry.npmjs.org/@emnapi/core/-/core-1.9.0.tgz",
|
||||
"integrity": "sha512-0DQ98G9ZQZOxfUcQn1waV2yS8aWdZ6kJMbYCJB3oUBecjWYO1fqJ+a1DRfPF3O5JEkwqwP1A9QEN/9mYm2Yd0w==",
|
||||
"version": "1.9.1",
|
||||
"resolved": "https://registry.npmjs.org/@emnapi/core/-/core-1.9.1.tgz",
|
||||
"integrity": "sha512-mukuNALVsoix/w1BJwFzwXBN/dHeejQtuVzcDsfOEsdpCumXb/E9j8w11h5S54tT1xhifGfbbSm/ICrObRb3KA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"optional": true,
|
||||
@@ -26,9 +26,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@emnapi/runtime": {
|
||||
"version": "1.9.0",
|
||||
"resolved": "https://registry.npmjs.org/@emnapi/runtime/-/runtime-1.9.0.tgz",
|
||||
"integrity": "sha512-QN75eB0IH2ywSpRpNddCRfQIhmJYBCJ1x5Lb3IscKAL8bMnVAKnRg8dCoXbHzVLLH7P38N2Z3mtulB7W0J0FKw==",
|
||||
"version": "1.9.1",
|
||||
"resolved": "https://registry.npmjs.org/@emnapi/runtime/-/runtime-1.9.1.tgz",
|
||||
"integrity": "sha512-VYi5+ZVLhpgK4hQ0TAjiQiZ6ol0oe4mBx7mVv7IflsiEp0OWoVsp/+f9Vc1hOhE0TtkORVrI1GvzyreqpgWtkA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"optional": true,
|
||||
@@ -71,20 +71,10 @@
|
||||
"url": "https://github.com/sponsors/Brooooooklyn"
|
||||
}
|
||||
},
|
||||
"node_modules/@oxc-project/runtime": {
|
||||
"version": "0.115.0",
|
||||
"resolved": "https://registry.npmjs.org/@oxc-project/runtime/-/runtime-0.115.0.tgz",
|
||||
"integrity": "sha512-Rg8Wlt5dCbXhQnsXPrkOjL1DTSvXLgb2R/KYfnf1/K+R0k6UMLEmbQXPM+kwrWqSmWA2t0B1EtHy2/3zikQpvQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": "^20.19.0 || >=22.12.0"
|
||||
}
|
||||
},
|
||||
"node_modules/@oxc-project/types": {
|
||||
"version": "0.115.0",
|
||||
"resolved": "https://registry.npmjs.org/@oxc-project/types/-/types-0.115.0.tgz",
|
||||
"integrity": "sha512-4n91DKnebUS4yjUHl2g3/b2T+IUdCfmoZGhmwsovZCDaJSs+QkVAM+0AqqTxHSsHfeiMuueT75cZaZcT/m0pSw==",
|
||||
"version": "0.120.0",
|
||||
"resolved": "https://registry.npmjs.org/@oxc-project/types/-/types-0.120.0.tgz",
|
||||
"integrity": "sha512-k1YNu55DuvAip/MGE1FTsIuU3FUCn6v/ujG9V7Nq5Df/kX2CWb13hhwD0lmJGMGqE+bE1MXvv9SZVnMzEXlWcg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"funding": {
|
||||
@@ -92,9 +82,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-android-arm64": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-android-arm64/-/binding-android-arm64-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-lcJL0bN5hpgJfSIz/8PIf02irmyL43P+j1pTCfbD1DbLkmGRuFIA4DD3B3ZOvGqG0XiVvRznbKtN0COQVaKUTg==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-android-arm64/-/binding-android-arm64-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-jOHxwXhxmFKuXztiu1ORieJeTbx5vrTkcOkkkn2d35726+iwhrY1w/+nYY/AGgF12thg33qC3R1LMBF5tHTZHg==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -109,9 +99,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-darwin-arm64": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-darwin-arm64/-/binding-darwin-arm64-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-J7Zk3kLYFsLtuH6U+F4pS2sYVzac0qkjcO5QxHS7OS7yZu2LRs+IXo+uvJ/mvpyUljDJ3LROZPoQfgBIpCMhdQ==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-darwin-arm64/-/binding-darwin-arm64-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-gED05Teg/vtTZbIJBc4VNMAxAFDUPkuO/rAIyyxZjTj1a1/s6z5TII/5yMGZ0uLRCifEtwUQn8OlYzuYc0m70w==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -126,9 +116,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-darwin-x64": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-darwin-x64/-/binding-darwin-x64-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-iwtmmghy8nhfRGeNAIltcNXzD0QMNaaA5U/NyZc1Ia4bxrzFByNMDoppoC+hl7cDiUq5/1CnFthpT9n+UtfFyg==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-darwin-x64/-/binding-darwin-x64-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-rI15NcM1mA48lqrIxVkHfAqcyFLcQwyXWThy+BQ5+mkKKPvSO26ir+ZDp36AgYoYVkqvMcdS8zOE6SeBsR9e8A==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
@@ -143,9 +133,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-freebsd-x64": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-freebsd-x64/-/binding-freebsd-x64-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-DLFYI78SCiZr5VvdEplsVC2Vx53lnA4/Ga5C65iyldMVaErr86aiqCoNBLl92PXPfDtUYjUh+xFFor40ueNs4Q==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-freebsd-x64/-/binding-freebsd-x64-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-XZRXHdTa+4ME1MuDVp021+doQ+z6Ei4CCFmNc5/sKbqb8YmkiJdj8QKlV3rCI0AJtAeSB5n0WGPuJWNL9p/L2w==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
@@ -160,9 +150,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-arm-gnueabihf": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm-gnueabihf/-/binding-linux-arm-gnueabihf-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-CsjTmTwd0Hri6iTw/DRMK7kOZ7FwAkrO4h8YWKoX/kcj833e4coqo2wzIFywtch/8Eb5enQ/lwLM7w6JX1W5RQ==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm-gnueabihf/-/binding-linux-arm-gnueabihf-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-R0SQMRluISSLzFE20sPWYHVmJdDQnRyc/FzSCN72BqQmh2SOZUFG+N3/vBZpR4C6WpEUVYJLrYUXaj43sJsNLA==",
|
||||
"cpu": [
|
||||
"arm"
|
||||
],
|
||||
@@ -177,9 +167,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-arm64-gnu": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm64-gnu/-/binding-linux-arm64-gnu-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-2x9O2JbSPxpxMDhP9Z74mahAStibTlrBMW0520+epJH5sac7/LwZW5Bmg/E6CXuEF53JJFW509uP+lSedaUNxg==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm64-gnu/-/binding-linux-arm64-gnu-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-Y1reMrV/o+cwpduYhJuOE3OMKx32RMYCidf14y+HssARRmhDuWXJ4yVguDg2R/8SyyGNo+auzz64LnPK9Hq6jg==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -194,9 +184,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-arm64-musl": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm64-musl/-/binding-linux-arm64-musl-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-JA1QRW31ogheAIRhIg9tjMfsYbglXXYGNPLdPEYrwFxdbkQCAzvpSCSHCDWNl4hTtrol8WeboCSEpjdZK8qrCg==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm64-musl/-/binding-linux-arm64-musl-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-vELN+HNb2IzuzSBUOD4NHmP9yrGwl1DVM29wlQvx1OLSclL0NgVWnVDKl/8tEks79EFek/kebQKnNJkIAA4W2g==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -211,9 +201,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-ppc64-gnu": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-ppc64-gnu/-/binding-linux-ppc64-gnu-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-aOKU9dJheda8Kj8Y3w9gnt9QFOO+qKPAl8SWd7JPHP+Cu0EuDAE5wokQubLzIDQWg2myXq2XhTpOVS07qqvT+w==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-ppc64-gnu/-/binding-linux-ppc64-gnu-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-ZqrufYTgzxbHwpqOjzSsb0UV/aV2TFIY5rP8HdsiPTv/CuAgCRjM6s9cYFwQ4CNH+hf9Y4erHW1GjZuZ7WoI7w==",
|
||||
"cpu": [
|
||||
"ppc64"
|
||||
],
|
||||
@@ -228,9 +218,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-s390x-gnu": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-s390x-gnu/-/binding-linux-s390x-gnu-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-OalO94fqj7IWRn3VdXWty75jC5dk4C197AWEuMhIpvVv2lw9fiPhud0+bW2ctCxb3YoBZor71QHbY+9/WToadA==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-s390x-gnu/-/binding-linux-s390x-gnu-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-gSlmVS1FZJSRicA6IyjoRoKAFK7IIHBs7xJuHRSmjImqk3mPPWbR7RhbnfH2G6bcmMEllCt2vQ/7u9e6bBnByg==",
|
||||
"cpu": [
|
||||
"s390x"
|
||||
],
|
||||
@@ -245,9 +235,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-x64-gnu": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-x64-gnu/-/binding-linux-x64-gnu-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-cVEl1vZtBsBZna3YMjGXNvnYYrOJ7RzuWvZU0ffvJUexWkukMaDuGhUXn0rjnV0ptzGVkvc+vW9Yqy6h8YX4pg==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-x64-gnu/-/binding-linux-x64-gnu-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-eOCKUpluKgfObT2pHjztnaWEIbUabWzk3qPZ5PuacuPmr4+JtQG4k2vGTY0H15edaTnicgU428XW/IH6AimcQw==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
@@ -262,9 +252,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-x64-musl": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-x64-musl/-/binding-linux-x64-musl-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-UzYnKCIIc4heAKgI4PZ3dfBGUZefGCJ1TPDuLHoCzgrMYPb5Rv6TLFuYtyM4rWyHM7hymNdsg5ik2C+UD9VDbA==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-x64-musl/-/binding-linux-x64-musl-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-Xdf2jQbfQowJnLcgYfD/m0Uu0Qj5OdxKallD78/IPPfzaiaI4KRAwZzHcKQ4ig1gtg1SuzC7jovNiM2TzQsBXA==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
@@ -279,9 +269,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-openharmony-arm64": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-openharmony-arm64/-/binding-openharmony-arm64-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-+6zoiF+RRyf5cdlFQP7nm58mq7+/2PFaY2DNQeD4B87N36JzfF/l9mdBkkmTvSYcYPE8tMh/o3cRlsx1ldLfog==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-openharmony-arm64/-/binding-openharmony-arm64-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-o1hYe8hLi1EY6jgPFyxQgQ1wcycX+qz8eEbVmot2hFkgUzPxy9+kF0u0NIQBeDq+Mko47AkaFFaChcvZa9UX9Q==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -296,9 +286,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-wasm32-wasi": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-wasm32-wasi/-/binding-wasm32-wasi-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-rgFN6sA/dyebil3YTlL2evvi/M+ivhfnyxec7AccTpRPccno/rPoNlqybEZQBkcbZu8Hy+eqNJCqfBR8P7Pg8g==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-wasm32-wasi/-/binding-wasm32-wasi-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-Ugv9o7qYJudqQO5Y5y2N2SOo6S4WiqiNOpuQyoPInnhVzCY+wi/GHltcLHypG9DEUYMB0iTB/huJrpadiAcNcA==",
|
||||
"cpu": [
|
||||
"wasm32"
|
||||
],
|
||||
@@ -313,9 +303,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-win32-arm64-msvc": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-win32-arm64-msvc/-/binding-win32-arm64-msvc-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-lHVNUG/8nlF1IQk1C0Ci574qKYyty2goMiPlRqkC5R+3LkXDkL5Dhx8ytbxq35m+pkHVIvIxviD+TWLdfeuadA==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-win32-arm64-msvc/-/binding-win32-arm64-msvc-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-7UODQb4fQUNT/vmgDZBl3XOBAIOutP5R3O/rkxg0aLfEGQ4opbCgU5vOw/scPe4xOqBwL9fw7/RP1vAMZ6QlAQ==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -330,9 +320,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-win32-x64-msvc": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-win32-x64-msvc/-/binding-win32-x64-msvc-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-G0oA4+w1iY5AGi5HcDTxWsoxF509hrFIPB2rduV5aDqS9FtDg1CAfa7V34qImbjfhIcA8C+RekocJZA96EarwQ==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-win32-x64-msvc/-/binding-win32-x64-msvc-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-PYxKHMVHOb5NJuDL53vBUl1VwUjymDcYI6rzpIni0C9+9mTiJedvUxSk7/RPp7OOAm3v+EjgMu9bIy3N6b408w==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
@@ -347,9 +337,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/pluginutils": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/pluginutils/-/pluginutils-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-w6oiRWgEBl04QkFZgmW+jnU1EC9b57Oihi2ot3HNWIQRqgHp5PnYDia5iZ5FF7rpa4EQdiqMDXjlqKGXBhsoXw==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/pluginutils/-/pluginutils-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-UkVDEFk1w3mveXeKgaTuYfKWtPbvgck1dT8TUG3bnccrH0XtLTuAyfCoks4Q/M5ZGToSVJTIQYCzy2g/atAOeg==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
@@ -397,16 +387,16 @@
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/@vitest/expect": {
|
||||
"version": "4.1.0",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/expect/-/expect-4.1.0.tgz",
|
||||
"integrity": "sha512-EIxG7k4wlWweuCLG9Y5InKFwpMEOyrMb6ZJ1ihYu02LVj/bzUwn2VMU+13PinsjRW75XnITeFrQBMH5+dLvCDA==",
|
||||
"version": "4.1.1",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/expect/-/expect-4.1.1.tgz",
|
||||
"integrity": "sha512-xAV0fqBTk44Rn6SjJReEQkHP3RrqbJo6JQ4zZ7/uVOiJZRarBtblzrOfFIZeYUrukp2YD6snZG6IBqhOoHTm+A==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@standard-schema/spec": "^1.1.0",
|
||||
"@types/chai": "^5.2.2",
|
||||
"@vitest/spy": "4.1.0",
|
||||
"@vitest/utils": "4.1.0",
|
||||
"@vitest/spy": "4.1.1",
|
||||
"@vitest/utils": "4.1.1",
|
||||
"chai": "^6.2.2",
|
||||
"tinyrainbow": "^3.0.3"
|
||||
},
|
||||
@@ -415,13 +405,13 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/mocker": {
|
||||
"version": "4.1.0",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/mocker/-/mocker-4.1.0.tgz",
|
||||
"integrity": "sha512-evxREh+Hork43+Y4IOhTo+h5lGmVRyjqI739Rz4RlUPqwrkFFDF6EMvOOYjTx4E8Tl6gyCLRL8Mu7Ry12a13Tw==",
|
||||
"version": "4.1.1",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/mocker/-/mocker-4.1.1.tgz",
|
||||
"integrity": "sha512-h3BOylsfsCLPeceuCPAAJ+BvNwSENgJa4hXoXu4im0bs9Lyp4URc4JYK4pWLZ4pG/UQn7AT92K6IByi6rE6g3A==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@vitest/spy": "4.1.0",
|
||||
"@vitest/spy": "4.1.1",
|
||||
"estree-walker": "^3.0.3",
|
||||
"magic-string": "^0.30.21"
|
||||
},
|
||||
@@ -430,7 +420,7 @@
|
||||
},
|
||||
"peerDependencies": {
|
||||
"msw": "^2.4.9",
|
||||
"vite": "^6.0.0 || ^7.0.0 || ^8.0.0-0"
|
||||
"vite": "^6.0.0 || ^7.0.0 || ^8.0.0"
|
||||
},
|
||||
"peerDependenciesMeta": {
|
||||
"msw": {
|
||||
@@ -442,9 +432,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/pretty-format": {
|
||||
"version": "4.1.0",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/pretty-format/-/pretty-format-4.1.0.tgz",
|
||||
"integrity": "sha512-3RZLZlh88Ib0J7NQTRATfc/3ZPOnSUn2uDBUoGNn5T36+bALixmzphN26OUD3LRXWkJu4H0s5vvUeqBiw+kS0A==",
|
||||
"version": "4.1.1",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/pretty-format/-/pretty-format-4.1.1.tgz",
|
||||
"integrity": "sha512-GM+TEQN5WhOygr1lp7skeVjdLPqqWMHsfzXrcHAqZJi/lIVh63H0kaRCY8MDhNWikx19zBUK8ceaLB7X5AH9NQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
@@ -455,13 +445,13 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/runner": {
|
||||
"version": "4.1.0",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/runner/-/runner-4.1.0.tgz",
|
||||
"integrity": "sha512-Duvx2OzQ7d6OjchL+trw+aSrb9idh7pnNfxrklo14p3zmNL4qPCDeIJAK+eBKYjkIwG96Bc6vYuxhqDXQOWpoQ==",
|
||||
"version": "4.1.1",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/runner/-/runner-4.1.1.tgz",
|
||||
"integrity": "sha512-f7+FPy75vN91QGWsITueq0gedwUZy1fLtHOCMeQpjs8jTekAHeKP80zfDEnhrleviLHzVSDXIWuCIOFn3D3f8A==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@vitest/utils": "4.1.0",
|
||||
"@vitest/utils": "4.1.1",
|
||||
"pathe": "^2.0.3"
|
||||
},
|
||||
"funding": {
|
||||
@@ -469,14 +459,14 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/snapshot": {
|
||||
"version": "4.1.0",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/snapshot/-/snapshot-4.1.0.tgz",
|
||||
"integrity": "sha512-0Vy9euT1kgsnj1CHttwi9i9o+4rRLEaPRSOJ5gyv579GJkNpgJK+B4HSv/rAWixx2wdAFci1X4CEPjiu2bXIMg==",
|
||||
"version": "4.1.1",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/snapshot/-/snapshot-4.1.1.tgz",
|
||||
"integrity": "sha512-kMVSgcegWV2FibXEx9p9WIKgje58lcTbXgnJixfcg15iK8nzCXhmalL0ZLtTWLW9PH1+1NEDShiFFedB3tEgWg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@vitest/pretty-format": "4.1.0",
|
||||
"@vitest/utils": "4.1.0",
|
||||
"@vitest/pretty-format": "4.1.1",
|
||||
"@vitest/utils": "4.1.1",
|
||||
"magic-string": "^0.30.21",
|
||||
"pathe": "^2.0.3"
|
||||
},
|
||||
@@ -485,9 +475,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/spy": {
|
||||
"version": "4.1.0",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/spy/-/spy-4.1.0.tgz",
|
||||
"integrity": "sha512-pz77k+PgNpyMDv2FV6qmk5ZVau6c3R8HC8v342T2xlFxQKTrSeYw9waIJG8KgV9fFwAtTu4ceRzMivPTH6wSxw==",
|
||||
"version": "4.1.1",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/spy/-/spy-4.1.1.tgz",
|
||||
"integrity": "sha512-6Ti/KT5OVaiupdIZEuZN7l3CZcR0cxnxt70Z0//3CtwgObwA6jZhmVBA3yrXSVN3gmwjgd7oDNLlsXz526gpRA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"funding": {
|
||||
@@ -495,13 +485,13 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/utils": {
|
||||
"version": "4.1.0",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/utils/-/utils-4.1.0.tgz",
|
||||
"integrity": "sha512-XfPXT6a8TZY3dcGY8EdwsBulFCIw+BeeX0RZn2x/BtiY/75YGh8FeWGG8QISN/WhaqSrE2OrlDgtF8q5uhOTmw==",
|
||||
"version": "4.1.1",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/utils/-/utils-4.1.1.tgz",
|
||||
"integrity": "sha512-cNxAlaB3sHoCdL6pj6yyUXv9Gry1NHNg0kFTXdvSIZXLHsqKH7chiWOkwJ5s5+d/oMwcoG9T0bKU38JZWKusrQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@vitest/pretty-format": "4.1.0",
|
||||
"@vitest/pretty-format": "4.1.1",
|
||||
"convert-source-map": "^2.0.0",
|
||||
"tinyrainbow": "^3.0.3"
|
||||
},
|
||||
@@ -964,14 +954,14 @@
|
||||
}
|
||||
},
|
||||
"node_modules/rolldown": {
|
||||
"version": "1.0.0-rc.9",
|
||||
"resolved": "https://registry.npmjs.org/rolldown/-/rolldown-1.0.0-rc.9.tgz",
|
||||
"integrity": "sha512-9EbgWge7ZH+yqb4d2EnELAntgPTWbfL8ajiTW+SyhJEC4qhBbkCKbqFV4Ge4zmu5ziQuVbWxb/XwLZ+RIO7E8Q==",
|
||||
"version": "1.0.0-rc.10",
|
||||
"resolved": "https://registry.npmjs.org/rolldown/-/rolldown-1.0.0-rc.10.tgz",
|
||||
"integrity": "sha512-q7j6vvarRFmKpgJUT8HCAUljkgzEp4LAhPlJUvQhA5LA1SUL36s5QCysMutErzL3EbNOZOkoziSx9iZC4FddKA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@oxc-project/types": "=0.115.0",
|
||||
"@rolldown/pluginutils": "1.0.0-rc.9"
|
||||
"@oxc-project/types": "=0.120.0",
|
||||
"@rolldown/pluginutils": "1.0.0-rc.10"
|
||||
},
|
||||
"bin": {
|
||||
"rolldown": "bin/cli.mjs"
|
||||
@@ -980,21 +970,21 @@
|
||||
"node": "^20.19.0 || >=22.12.0"
|
||||
},
|
||||
"optionalDependencies": {
|
||||
"@rolldown/binding-android-arm64": "1.0.0-rc.9",
|
||||
"@rolldown/binding-darwin-arm64": "1.0.0-rc.9",
|
||||
"@rolldown/binding-darwin-x64": "1.0.0-rc.9",
|
||||
"@rolldown/binding-freebsd-x64": "1.0.0-rc.9",
|
||||
"@rolldown/binding-linux-arm-gnueabihf": "1.0.0-rc.9",
|
||||
"@rolldown/binding-linux-arm64-gnu": "1.0.0-rc.9",
|
||||
"@rolldown/binding-linux-arm64-musl": "1.0.0-rc.9",
|
||||
"@rolldown/binding-linux-ppc64-gnu": "1.0.0-rc.9",
|
||||
"@rolldown/binding-linux-s390x-gnu": "1.0.0-rc.9",
|
||||
"@rolldown/binding-linux-x64-gnu": "1.0.0-rc.9",
|
||||
"@rolldown/binding-linux-x64-musl": "1.0.0-rc.9",
|
||||
"@rolldown/binding-openharmony-arm64": "1.0.0-rc.9",
|
||||
"@rolldown/binding-wasm32-wasi": "1.0.0-rc.9",
|
||||
"@rolldown/binding-win32-arm64-msvc": "1.0.0-rc.9",
|
||||
"@rolldown/binding-win32-x64-msvc": "1.0.0-rc.9"
|
||||
"@rolldown/binding-android-arm64": "1.0.0-rc.10",
|
||||
"@rolldown/binding-darwin-arm64": "1.0.0-rc.10",
|
||||
"@rolldown/binding-darwin-x64": "1.0.0-rc.10",
|
||||
"@rolldown/binding-freebsd-x64": "1.0.0-rc.10",
|
||||
"@rolldown/binding-linux-arm-gnueabihf": "1.0.0-rc.10",
|
||||
"@rolldown/binding-linux-arm64-gnu": "1.0.0-rc.10",
|
||||
"@rolldown/binding-linux-arm64-musl": "1.0.0-rc.10",
|
||||
"@rolldown/binding-linux-ppc64-gnu": "1.0.0-rc.10",
|
||||
"@rolldown/binding-linux-s390x-gnu": "1.0.0-rc.10",
|
||||
"@rolldown/binding-linux-x64-gnu": "1.0.0-rc.10",
|
||||
"@rolldown/binding-linux-x64-musl": "1.0.0-rc.10",
|
||||
"@rolldown/binding-openharmony-arm64": "1.0.0-rc.10",
|
||||
"@rolldown/binding-wasm32-wasi": "1.0.0-rc.10",
|
||||
"@rolldown/binding-win32-arm64-msvc": "1.0.0-rc.10",
|
||||
"@rolldown/binding-win32-x64-msvc": "1.0.0-rc.10"
|
||||
}
|
||||
},
|
||||
"node_modules/siginfo": {
|
||||
@@ -1081,9 +1071,9 @@
|
||||
"optional": true
|
||||
},
|
||||
"node_modules/typescript": {
|
||||
"version": "5.9.3",
|
||||
"resolved": "https://registry.npmjs.org/typescript/-/typescript-5.9.3.tgz",
|
||||
"integrity": "sha512-jl1vZzPDinLr9eUt3J/t7V6FgNEw9QjvBPdysz9KfQDD41fQrC2Y4vKQdiaUpFT4bXlb1RHhLpp8wtm6M5TgSw==",
|
||||
"version": "6.0.2",
|
||||
"resolved": "https://registry.npmjs.org/typescript/-/typescript-6.0.2.tgz",
|
||||
"integrity": "sha512-bGdAIrZ0wiGDo5l8c++HWtbaNCWTS4UTv7RaTH/ThVIgjkveJt83m74bBHMJkuCbslY8ixgLBVZJIOiQlQTjfQ==",
|
||||
"dev": true,
|
||||
"license": "Apache-2.0",
|
||||
"bin": {
|
||||
@@ -1095,17 +1085,16 @@
|
||||
}
|
||||
},
|
||||
"node_modules/vite": {
|
||||
"version": "8.0.0",
|
||||
"resolved": "https://registry.npmjs.org/vite/-/vite-8.0.0.tgz",
|
||||
"integrity": "sha512-fPGaRNj9Zytaf8LEiBhY7Z6ijnFKdzU/+mL8EFBaKr7Vw1/FWcTBAMW0wLPJAGMPX38ZPVCVgLceWiEqeoqL2Q==",
|
||||
"version": "8.0.1",
|
||||
"resolved": "https://registry.npmjs.org/vite/-/vite-8.0.1.tgz",
|
||||
"integrity": "sha512-wt+Z2qIhfFt85uiyRt5LPU4oVEJBXj8hZNWKeqFG4gRG/0RaRGJ7njQCwzFVjO+v4+Ipmf5CY7VdmZRAYYBPHw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@oxc-project/runtime": "0.115.0",
|
||||
"lightningcss": "^1.32.0",
|
||||
"picomatch": "^4.0.3",
|
||||
"postcss": "^8.5.8",
|
||||
"rolldown": "1.0.0-rc.9",
|
||||
"rolldown": "1.0.0-rc.10",
|
||||
"tinyglobby": "^0.2.15"
|
||||
},
|
||||
"bin": {
|
||||
@@ -1122,7 +1111,7 @@
|
||||
},
|
||||
"peerDependencies": {
|
||||
"@types/node": "^20.19.0 || >=22.12.0",
|
||||
"@vitejs/devtools": "^0.0.0-alpha.31",
|
||||
"@vitejs/devtools": "^0.1.0",
|
||||
"esbuild": "^0.27.0",
|
||||
"jiti": ">=1.21.0",
|
||||
"less": "^4.0.0",
|
||||
@@ -1174,19 +1163,19 @@
|
||||
}
|
||||
},
|
||||
"node_modules/vitest": {
|
||||
"version": "4.1.0",
|
||||
"resolved": "https://registry.npmjs.org/vitest/-/vitest-4.1.0.tgz",
|
||||
"integrity": "sha512-YbDrMF9jM2Lqc++2530UourxZHmkKLxrs4+mYhEwqWS97WJ7wOYEkcr+QfRgJ3PW9wz3odRijLZjHEaRLTNbqw==",
|
||||
"version": "4.1.1",
|
||||
"resolved": "https://registry.npmjs.org/vitest/-/vitest-4.1.1.tgz",
|
||||
"integrity": "sha512-yF+o4POL41rpAzj5KVILUxm1GCjKnELvaqmU9TLLUbMfDzuN0UpUR9uaDs+mCtjPe+uYPksXDRLQGGPvj1cTmA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@vitest/expect": "4.1.0",
|
||||
"@vitest/mocker": "4.1.0",
|
||||
"@vitest/pretty-format": "4.1.0",
|
||||
"@vitest/runner": "4.1.0",
|
||||
"@vitest/snapshot": "4.1.0",
|
||||
"@vitest/spy": "4.1.0",
|
||||
"@vitest/utils": "4.1.0",
|
||||
"@vitest/expect": "4.1.1",
|
||||
"@vitest/mocker": "4.1.1",
|
||||
"@vitest/pretty-format": "4.1.1",
|
||||
"@vitest/runner": "4.1.1",
|
||||
"@vitest/snapshot": "4.1.1",
|
||||
"@vitest/spy": "4.1.1",
|
||||
"@vitest/utils": "4.1.1",
|
||||
"es-module-lexer": "^2.0.0",
|
||||
"expect-type": "^1.3.0",
|
||||
"magic-string": "^0.30.21",
|
||||
@@ -1198,7 +1187,7 @@
|
||||
"tinyexec": "^1.0.2",
|
||||
"tinyglobby": "^0.2.15",
|
||||
"tinyrainbow": "^3.0.3",
|
||||
"vite": "^6.0.0 || ^7.0.0 || ^8.0.0-0",
|
||||
"vite": "^6.0.0 || ^7.0.0 || ^8.0.0",
|
||||
"why-is-node-running": "^2.3.0"
|
||||
},
|
||||
"bin": {
|
||||
@@ -1214,13 +1203,13 @@
|
||||
"@edge-runtime/vm": "*",
|
||||
"@opentelemetry/api": "^1.9.0",
|
||||
"@types/node": "^20.0.0 || ^22.0.0 || >=24.0.0",
|
||||
"@vitest/browser-playwright": "4.1.0",
|
||||
"@vitest/browser-preview": "4.1.0",
|
||||
"@vitest/browser-webdriverio": "4.1.0",
|
||||
"@vitest/ui": "4.1.0",
|
||||
"@vitest/browser-playwright": "4.1.1",
|
||||
"@vitest/browser-preview": "4.1.1",
|
||||
"@vitest/browser-webdriverio": "4.1.1",
|
||||
"@vitest/ui": "4.1.1",
|
||||
"happy-dom": "*",
|
||||
"jsdom": "*",
|
||||
"vite": "^6.0.0 || ^7.0.0 || ^8.0.0-0"
|
||||
"vite": "^6.0.0 || ^7.0.0 || ^8.0.0"
|
||||
},
|
||||
"peerDependenciesMeta": {
|
||||
"@edge-runtime/vm": {
|
||||
|
||||
@@ -32,7 +32,7 @@
|
||||
],
|
||||
"license": "BUSL-1.1",
|
||||
"devDependencies": {
|
||||
"typescript": "^5.4",
|
||||
"typescript": "^6.0.0",
|
||||
"vitest": "^4.1"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -186,6 +186,7 @@ export interface SkillInfo {
|
||||
agent_max_turns: number | null;
|
||||
notify_on_complete: string;
|
||||
enabled: boolean;
|
||||
priority: number;
|
||||
allowed_tools: string;
|
||||
license: string;
|
||||
compatibility: string;
|
||||
@@ -215,6 +216,7 @@ export interface CreateSkillRequest {
|
||||
agent_max_turns?: number | null;
|
||||
notify_on_complete?: string;
|
||||
enabled?: boolean;
|
||||
priority?: number;
|
||||
allowed_tools?: string;
|
||||
license?: string;
|
||||
compatibility?: string;
|
||||
@@ -240,6 +242,7 @@ export interface UpdateSkillRequest {
|
||||
agent_max_turns?: number | null;
|
||||
notify_on_complete?: string;
|
||||
enabled?: boolean;
|
||||
priority?: number;
|
||||
allowed_tools?: string;
|
||||
license?: string;
|
||||
compatibility?: string;
|
||||
@@ -402,6 +405,7 @@ export interface ConsoleCreateWsRequest {
|
||||
model?: string;
|
||||
initial_message?: string;
|
||||
skill?: string;
|
||||
resume_ws?: string;
|
||||
}
|
||||
|
||||
export interface ConsoleCreateWsResponse {
|
||||
|
||||
+75
-1
@@ -1,11 +1,23 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def pytest_addoption(parser: pytest.Parser) -> None:
|
||||
parser.addoption(
|
||||
"--storage-backend",
|
||||
default="sqlite",
|
||||
choices=["sqlite", "postgresql"],
|
||||
help="Storage backend for integration tests (default: sqlite)",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tmp_db(tmp_path):
|
||||
"""Provide a temporary SQLite storage backend."""
|
||||
"""Provide a temporary SQLite storage backend (singleton registry)."""
|
||||
from turnstone.core.storage import init_storage, reset_storage
|
||||
|
||||
db_path = str(tmp_path / "test.db")
|
||||
@@ -15,6 +27,68 @@ def tmp_db(tmp_path):
|
||||
reset_storage()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage_backend(request, tmp_path):
|
||||
"""Shared storage backend fixture — respects --storage-backend flag.
|
||||
|
||||
Returns a StorageBackend instance (SQLite or PostgreSQL).
|
||||
Tests that use this fixture run against whichever backend CI selects.
|
||||
"""
|
||||
from turnstone.core.storage import init_storage, reset_storage
|
||||
|
||||
backend_type = request.config.getoption("--storage-backend")
|
||||
reset_storage()
|
||||
|
||||
if backend_type == "postgresql":
|
||||
pg_url = os.environ.get(
|
||||
"TURNSTONE_TEST_PG_URL",
|
||||
"postgresql+psycopg://postgres:postgres@localhost:5432/turnstone_test",
|
||||
)
|
||||
backend = init_storage("postgresql", url=pg_url, run_migrations=False)
|
||||
yield backend
|
||||
# Truncate all tables between tests — faster than DELETE and resets
|
||||
# autoincrement sequences. CASCADE handles any future FK constraints.
|
||||
# NOTE: accesses backend._engine (SQLAlchemy internal) — both SQLite
|
||||
# and PostgreSQL backends expose this. If a non-SQLAlchemy backend is
|
||||
# ever added, this cleanup will need a protocol-level hook.
|
||||
try:
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.storage._schema import metadata as db_metadata
|
||||
|
||||
with backend._engine.connect() as conn:
|
||||
table_names = ", ".join(t.name for t in reversed(db_metadata.sorted_tables))
|
||||
conn.execute(sa.text(f"TRUNCATE {table_names} RESTART IDENTITY CASCADE"))
|
||||
conn.commit()
|
||||
except Exception:
|
||||
pass # best-effort cleanup; reset_storage disposes engine
|
||||
finally:
|
||||
reset_storage()
|
||||
else:
|
||||
db_path = str(tmp_path / "test.db")
|
||||
backend = init_storage("sqlite", path=db_path, run_migrations=False)
|
||||
yield backend
|
||||
reset_storage()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def backend(storage_backend):
|
||||
"""Alias for storage_backend — used by test_storage_sqlite.py etc."""
|
||||
return storage_backend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def db(storage_backend):
|
||||
"""Alias for storage_backend — used by domain-specific storage tests."""
|
||||
return storage_backend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(storage_backend):
|
||||
"""Alias for storage_backend — used by services/skill resource tests."""
|
||||
return storage_backend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_openai_client():
|
||||
"""Return a minimal mock OpenAI client."""
|
||||
|
||||
@@ -19,6 +19,7 @@ class TestServerVersioning:
|
||||
|
||||
mock_mgr = MagicMock()
|
||||
mock_mgr.list_all.return_value = []
|
||||
mock_mgr.max_workstreams = 10
|
||||
app = create_app(
|
||||
workstreams=mock_mgr,
|
||||
global_queue=queue.Queue(),
|
||||
|
||||
+10
-2
@@ -783,6 +783,7 @@ class TestServerAuth:
|
||||
mock_ws.session = mock_session
|
||||
mock_mgr = MagicMock()
|
||||
mock_mgr.list_all.return_value = [mock_ws]
|
||||
mock_mgr.max_workstreams = 10
|
||||
|
||||
app = srv_mod.create_app(
|
||||
workstreams=mock_mgr,
|
||||
@@ -1001,6 +1002,7 @@ class TestServerLogin:
|
||||
mock_ws.session = mock_session
|
||||
mock_mgr = MagicMock()
|
||||
mock_mgr.list_all.return_value = [mock_ws]
|
||||
mock_mgr.max_workstreams = 10
|
||||
|
||||
app = srv_mod.create_app(
|
||||
workstreams=mock_mgr,
|
||||
@@ -1369,8 +1371,11 @@ class TestCorsConfigurable:
|
||||
|
||||
import turnstone.server as srv_mod
|
||||
|
||||
mgr = MagicMock()
|
||||
mgr.list_all.return_value = []
|
||||
mgr.max_workstreams = 10
|
||||
app = srv_mod.create_app(
|
||||
workstreams=MagicMock(),
|
||||
workstreams=mgr,
|
||||
global_queue=queue.Queue(),
|
||||
global_listeners=[],
|
||||
global_listeners_lock=threading.Lock(),
|
||||
@@ -1388,8 +1393,11 @@ class TestCorsConfigurable:
|
||||
|
||||
import turnstone.server as srv_mod
|
||||
|
||||
mgr = MagicMock()
|
||||
mgr.list_all.return_value = []
|
||||
mgr.max_workstreams = 10
|
||||
app = srv_mod.create_app(
|
||||
workstreams=MagicMock(),
|
||||
workstreams=mgr,
|
||||
global_queue=queue.Queue(),
|
||||
global_listeners=[],
|
||||
global_listeners_lock=threading.Lock(),
|
||||
|
||||
@@ -0,0 +1,351 @@
|
||||
"""Stress tests for bridge.py threading — race conditions in approval,
|
||||
plan review, and workstream lifecycle.
|
||||
|
||||
Each scenario is run many times (ITERATIONS) with threading.Barrier to
|
||||
maximize timing overlap. Uses mock broker (no Redis) and no HTTP calls.
|
||||
|
||||
Races tested:
|
||||
1. Duplicate approval on SSE reconnect (TOCTOU in _pending_approvals)
|
||||
2. Duplicate plan review on SSE reconnect (TOCTOU in _pending_plan_reviews)
|
||||
3. approve_set stale reference escape during concurrent update
|
||||
4. _running flag visibility across threads on shutdown
|
||||
5. Approval thread exits within bounded time after timeout
|
||||
6. Concurrent approval + workstream close leaves no orphaned state
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from collections import Counter
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from turnstone.mq.bridge import Bridge
|
||||
|
||||
ITERATIONS = 100
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_bridge(**overrides) -> Bridge:
|
||||
"""Create a Bridge with a mock broker (no Redis or HTTP)."""
|
||||
broker = MagicMock()
|
||||
defaults = dict(
|
||||
server_url="http://localhost:8080",
|
||||
broker=broker,
|
||||
node_id="test-node",
|
||||
approval_timeout=1,
|
||||
)
|
||||
defaults.update(overrides)
|
||||
return Bridge(**defaults)
|
||||
|
||||
|
||||
def _approval_items(tool_name: str = "bash") -> list[dict]:
|
||||
return [{"func_name": tool_name, "needs_approval": True, "approval_label": tool_name}]
|
||||
|
||||
|
||||
def _wait_pending_resolved(bridge: Bridge, key: str, attr: str, deadline_s: float = 3.0) -> bool:
|
||||
"""Poll until the pending entry is resolved (tombstone) or absent."""
|
||||
deadline = time.monotonic() + deadline_s
|
||||
while time.monotonic() < deadline:
|
||||
with bridge._lock:
|
||||
entries = getattr(bridge, attr)
|
||||
if key not in entries:
|
||||
return True
|
||||
_, resolved_at = entries[key]
|
||||
if resolved_at > 0:
|
||||
return True
|
||||
time.sleep(0.01)
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Race 1: Duplicate approval on SSE reconnect
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDuplicateApproval:
|
||||
"""Two threads call _handle_approval for the same ws_id simultaneously.
|
||||
Only one should create a pending entry; the other should be skipped."""
|
||||
|
||||
def test_no_duplicate_approvals(self):
|
||||
sent_count = Counter()
|
||||
|
||||
for _ in range(ITERATIONS):
|
||||
bridge = _make_bridge()
|
||||
bridge._broker.pop_response.return_value = '{"type": "approve", "approved": true}'
|
||||
barrier = threading.Barrier(2, timeout=5)
|
||||
|
||||
def _call_approval(bridge=bridge, barrier=barrier):
|
||||
barrier.wait()
|
||||
bridge._handle_approval("ws-1", {"items": _approval_items()})
|
||||
|
||||
t1 = threading.Thread(target=_call_approval)
|
||||
t2 = threading.Thread(target=_call_approval)
|
||||
with (
|
||||
patch.object(bridge, "_api_approve") as mock_approve,
|
||||
patch.object(bridge, "_publish_ws"),
|
||||
):
|
||||
t1.start()
|
||||
t2.start()
|
||||
t1.join(timeout=5)
|
||||
t2.join(timeout=5)
|
||||
assert not t1.is_alive(), "Thread 1 hung"
|
||||
assert not t2.is_alive(), "Thread 2 hung"
|
||||
|
||||
# Wait for spawned _wait_approval threads to resolve
|
||||
_wait_pending_resolved(bridge, "ws-1", "_pending_approvals")
|
||||
|
||||
sent_count[mock_approve.call_count] += 1
|
||||
|
||||
# At most 1 approval should be forwarded per iteration
|
||||
assert sent_count.get(2, 0) == 0, (
|
||||
f"Duplicate approvals sent in {sent_count[2]}/{ITERATIONS} iterations"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Race 2: Duplicate plan review on SSE reconnect
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDuplicatePlanReview:
|
||||
"""Two threads call _handle_plan_review simultaneously.
|
||||
Only one should create a pending entry."""
|
||||
|
||||
def test_no_duplicate_plan_reviews(self):
|
||||
sent_count = Counter()
|
||||
|
||||
for _ in range(ITERATIONS):
|
||||
bridge = _make_bridge()
|
||||
bridge._broker.pop_response.return_value = (
|
||||
'{"type": "plan_feedback", "feedback": "looks good"}'
|
||||
)
|
||||
barrier = threading.Barrier(2, timeout=5)
|
||||
|
||||
def _call_plan(bridge=bridge, barrier=barrier):
|
||||
barrier.wait()
|
||||
bridge._handle_plan_review("ws-1", {"content": "plan text"})
|
||||
|
||||
t1 = threading.Thread(target=_call_plan)
|
||||
t2 = threading.Thread(target=_call_plan)
|
||||
with patch.object(bridge, "_publish_ws"), patch.object(bridge._http, "post"):
|
||||
t1.start()
|
||||
t2.start()
|
||||
t1.join(timeout=5)
|
||||
t2.join(timeout=5)
|
||||
assert not t1.is_alive(), "Thread 1 hung"
|
||||
assert not t2.is_alive(), "Thread 2 hung"
|
||||
|
||||
# Wait for spawned _wait_plan threads to resolve
|
||||
_wait_pending_resolved(bridge, "ws-1", "_pending_plan_reviews")
|
||||
|
||||
sent_count[bridge._http.post.call_count] += 1
|
||||
|
||||
assert sent_count.get(2, 0) == 0, (
|
||||
f"Duplicate plan reviews sent in {sent_count[2]}/{ITERATIONS} iterations"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Race 3: approve_set stale reference during concurrent update
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestApproveSetConsistency:
|
||||
"""One thread reads approve_set for auto-approve check while another
|
||||
updates it via _wait_approval 'always' path. The auto-approve
|
||||
decision should be consistent (either all-approved or not)."""
|
||||
|
||||
def test_approve_set_never_partially_visible(self):
|
||||
for _ in range(ITERATIONS):
|
||||
bridge = _make_bridge()
|
||||
with bridge._lock:
|
||||
bridge._ws_approve_tools["ws-1"] = {"read_file", "search"}
|
||||
|
||||
barrier = threading.Barrier(2, timeout=5)
|
||||
results = []
|
||||
|
||||
def _reader(bridge=bridge, barrier=barrier, results=results):
|
||||
barrier.wait()
|
||||
with bridge._lock:
|
||||
snap = bridge._ws_approve_tools.get("ws-1", set()).copy()
|
||||
results.append(snap)
|
||||
|
||||
def _writer(bridge=bridge, barrier=barrier):
|
||||
barrier.wait()
|
||||
with bridge._lock:
|
||||
existing = bridge._ws_approve_tools.get("ws-1", set())
|
||||
bridge._ws_approve_tools["ws-1"] = existing | {"bash", "write_file"}
|
||||
|
||||
t1 = threading.Thread(target=_reader)
|
||||
t2 = threading.Thread(target=_writer)
|
||||
t1.start()
|
||||
t2.start()
|
||||
t1.join(timeout=5)
|
||||
t2.join(timeout=5)
|
||||
assert not t1.is_alive(), "Reader hung"
|
||||
assert not t2.is_alive(), "Writer hung"
|
||||
|
||||
snap = results[0]
|
||||
assert snap in (
|
||||
{"read_file", "search"},
|
||||
{"read_file", "search", "bash", "write_file"},
|
||||
), f"Partial set observed: {snap}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Race 4: _running flag visibility across threads
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRunningFlagVisibility:
|
||||
"""All threads reading _running should see False within a bounded time
|
||||
after the main thread sets it."""
|
||||
|
||||
def test_all_threads_observe_shutdown(self):
|
||||
bridge = _make_bridge()
|
||||
observed_false = threading.Event()
|
||||
threads_running = []
|
||||
|
||||
def _spin_checker():
|
||||
while bridge._running:
|
||||
time.sleep(0.001)
|
||||
observed_false.set()
|
||||
|
||||
for _ in range(5):
|
||||
t = threading.Thread(target=_spin_checker, daemon=True)
|
||||
threads_running.append(t)
|
||||
t.start()
|
||||
|
||||
time.sleep(0.01)
|
||||
bridge._running = False
|
||||
|
||||
for t in threads_running:
|
||||
t.join(timeout=1)
|
||||
assert not t.is_alive(), "Thread did not observe _running=False"
|
||||
|
||||
assert observed_false.is_set()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Race 5: Approval thread exits within bounded time
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestApprovalThreadTimeout:
|
||||
"""An approval thread blocked on pop_response should exit within the
|
||||
configured approval_timeout, not hang indefinitely."""
|
||||
|
||||
def test_approval_thread_exits_within_timeout(self):
|
||||
for _ in range(10):
|
||||
bridge = _make_bridge(approval_timeout=0.5)
|
||||
|
||||
def _slow_pop(queue_name, timeout=300):
|
||||
time.sleep(min(timeout, 0.5))
|
||||
return None
|
||||
|
||||
bridge._broker.pop_response.side_effect = _slow_pop
|
||||
|
||||
with patch.object(bridge, "_publish_ws"), patch.object(bridge, "_api_approve"):
|
||||
bridge._handle_approval("ws-1", {"items": _approval_items()})
|
||||
|
||||
# The pending entry should be resolved within the timeout
|
||||
resolved = _wait_pending_resolved(bridge, "ws-1", "_pending_approvals", deadline_s=3.0)
|
||||
assert resolved, "Approval thread did not exit within expected timeout"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Race 6: Concurrent approval + workstream close
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestApprovalDuringClose:
|
||||
"""An approval arriving at the exact same time as a ws_closed event
|
||||
should not leave orphaned state."""
|
||||
|
||||
def test_no_orphaned_pending_after_close(self):
|
||||
for _ in range(ITERATIONS):
|
||||
bridge = _make_bridge(approval_timeout=0.1)
|
||||
bridge._broker.pop_response.return_value = None # timeout
|
||||
|
||||
barrier = threading.Barrier(2, timeout=5)
|
||||
|
||||
def _send_approval(bridge=bridge, barrier=barrier):
|
||||
barrier.wait()
|
||||
with patch.object(bridge, "_publish_ws"), patch.object(bridge, "_api_approve"):
|
||||
bridge._handle_approval("ws-1", {"items": _approval_items()})
|
||||
|
||||
def _close_ws(bridge=bridge, barrier=barrier):
|
||||
barrier.wait()
|
||||
with (
|
||||
patch.object(bridge, "_publish_global"),
|
||||
patch.object(bridge, "_publish_cluster"),
|
||||
):
|
||||
bridge._handle_global_event({"type": "ws_closed", "ws_id": "ws-1"})
|
||||
|
||||
t1 = threading.Thread(target=_send_approval)
|
||||
t2 = threading.Thread(target=_close_ws)
|
||||
t1.start()
|
||||
t2.start()
|
||||
t1.join(timeout=5)
|
||||
t2.join(timeout=5)
|
||||
assert not t1.is_alive(), "Approval thread hung"
|
||||
assert not t2.is_alive(), "Close thread hung"
|
||||
|
||||
# Wait for spawned _wait_approval thread to resolve (if close
|
||||
# didn't remove the entry first)
|
||||
resolved = _wait_pending_resolved(bridge, "ws-1", "_pending_approvals")
|
||||
assert resolved, "Orphaned pending approval"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Race 7: Plan review refinement loop (tombstone → cleanup → re-entry)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPlanReviewRefinementLoop:
|
||||
"""After a plan review is resolved, a ws_state event should clean up the
|
||||
tombstone so the refinement-loop plan_review event is handled correctly."""
|
||||
|
||||
def test_refinement_loop_allows_reentry(self):
|
||||
for _ in range(ITERATIONS):
|
||||
bridge = _make_bridge()
|
||||
bridge._broker.pop_response.return_value = (
|
||||
'{"type": "plan_feedback", "feedback": "refine this"}'
|
||||
)
|
||||
|
||||
# Step 1: first plan review — creates pending entry, resolves it
|
||||
with patch.object(bridge, "_publish_ws"), patch.object(bridge._http, "post"):
|
||||
bridge._handle_plan_review("ws-1", {"content": "plan v1"})
|
||||
|
||||
_wait_pending_resolved(bridge, "ws-1", "_pending_plan_reviews")
|
||||
|
||||
# Verify tombstone is present (resolved_at > 0)
|
||||
with bridge._lock:
|
||||
assert "ws-1" in bridge._pending_plan_reviews
|
||||
assert bridge._pending_plan_reviews["ws-1"][1] > 0
|
||||
|
||||
# Step 2: ws_state event cleans up the resolved tombstone
|
||||
with (
|
||||
patch.object(bridge, "_publish_ws"),
|
||||
patch.object(bridge, "_publish_global"),
|
||||
patch.object(bridge, "_publish_cluster"),
|
||||
):
|
||||
bridge._handle_global_event(
|
||||
{"type": "ws_state", "ws_id": "ws-1", "state": "working"}
|
||||
)
|
||||
|
||||
with bridge._lock:
|
||||
assert "ws-1" not in bridge._pending_plan_reviews
|
||||
|
||||
# Step 3: refinement plan_review arrives — should create new entry
|
||||
with patch.object(bridge, "_publish_ws"), patch.object(bridge._http, "post"):
|
||||
bridge._handle_plan_review("ws-1", {"content": "plan v2"})
|
||||
|
||||
_wait_pending_resolved(bridge, "ws-1", "_pending_plan_reviews")
|
||||
|
||||
with bridge._lock:
|
||||
assert "ws-1" in bridge._pending_plan_reviews
|
||||
@@ -2,17 +2,6 @@
|
||||
|
||||
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."""
|
||||
|
||||
+60
-27
@@ -3,54 +3,55 @@
|
||||
import argparse
|
||||
|
||||
import turnstone.core.config as config_mod
|
||||
from turnstone.core.config import apply_config, load_config
|
||||
from turnstone.core.config import apply_config, load_config, set_config_path
|
||||
|
||||
|
||||
def _reset_cache():
|
||||
"""Clear the module-level config cache between tests."""
|
||||
config_mod._cache = None
|
||||
config_mod._config_path = None
|
||||
|
||||
|
||||
def test_load_config_missing_file(tmp_path, monkeypatch):
|
||||
def test_load_config_missing_file(tmp_path):
|
||||
_reset_cache()
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", tmp_path / "nope.toml")
|
||||
set_config_path(str(tmp_path / "nope.toml"))
|
||||
assert load_config() == {}
|
||||
|
||||
|
||||
def test_load_config_valid_toml(tmp_path, monkeypatch):
|
||||
def test_load_config_valid_toml(tmp_path):
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text('[redis]\nhost = "10.0.0.1"\nport = 6380\npassword = "secret"\n')
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
result = load_config()
|
||||
assert result["redis"]["host"] == "10.0.0.1"
|
||||
assert result["redis"]["port"] == 6380
|
||||
assert result["redis"]["password"] == "secret"
|
||||
|
||||
|
||||
def test_load_config_section(tmp_path, monkeypatch):
|
||||
def test_load_config_section(tmp_path):
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text('[api]\nbase_url = "http://x:8000/v1"\n[redis]\nhost = "y"\n')
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
assert load_config("redis") == {"host": "y"}
|
||||
assert load_config("api") == {"base_url": "http://x:8000/v1"}
|
||||
assert load_config("nonexistent") == {}
|
||||
|
||||
|
||||
def test_load_config_invalid_toml(tmp_path, monkeypatch):
|
||||
def test_load_config_invalid_toml(tmp_path):
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text("this is not valid toml [[[")
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
assert load_config() == {}
|
||||
|
||||
|
||||
def test_load_config_caches(tmp_path, monkeypatch):
|
||||
def test_load_config_caches(tmp_path):
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text('[api]\nbase_url = "http://first"\n')
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
first = load_config()
|
||||
assert first["api"]["base_url"] == "http://first"
|
||||
|
||||
@@ -60,14 +61,14 @@ def test_load_config_caches(tmp_path, monkeypatch):
|
||||
assert second["api"]["base_url"] == "http://first"
|
||||
|
||||
|
||||
def test_apply_config_sets_defaults(tmp_path, monkeypatch):
|
||||
def test_apply_config_sets_defaults(tmp_path):
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text(
|
||||
'[redis]\nhost = "redis.local"\nport = 7777\npassword = "pw"\n'
|
||||
'[bridge]\nserver_url = "http://bridge:9090"\n'
|
||||
)
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--redis-host", default="localhost")
|
||||
@@ -84,11 +85,11 @@ def test_apply_config_sets_defaults(tmp_path, monkeypatch):
|
||||
assert args.server_url == "http://bridge:9090"
|
||||
|
||||
|
||||
def test_apply_config_cli_overrides(tmp_path, monkeypatch):
|
||||
def test_apply_config_cli_overrides(tmp_path):
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text('[redis]\nhost = "config-host"\nport = 7777\n')
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--redis-host", default="localhost")
|
||||
@@ -102,11 +103,11 @@ def test_apply_config_cli_overrides(tmp_path, monkeypatch):
|
||||
assert args.redis_port == 7777 # config wins (no CLI override)
|
||||
|
||||
|
||||
def test_apply_config_missing_keys_keep_defaults(tmp_path, monkeypatch):
|
||||
def test_apply_config_missing_keys_keep_defaults(tmp_path):
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text('[redis]\nhost = "only-host"\n') # no port, no password
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--redis-host", default="localhost")
|
||||
@@ -121,9 +122,9 @@ def test_apply_config_missing_keys_keep_defaults(tmp_path, monkeypatch):
|
||||
assert args.redis_password is None # original default kept
|
||||
|
||||
|
||||
def test_apply_config_no_file(tmp_path, monkeypatch):
|
||||
def test_apply_config_no_file(tmp_path):
|
||||
_reset_cache()
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", tmp_path / "nope.toml")
|
||||
set_config_path(str(tmp_path / "nope.toml"))
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--redis-host", default="localhost")
|
||||
@@ -133,11 +134,11 @@ def test_apply_config_no_file(tmp_path, monkeypatch):
|
||||
assert args.redis_host == "localhost"
|
||||
|
||||
|
||||
def test_apply_config_model_section(tmp_path, monkeypatch):
|
||||
def test_apply_config_model_section(tmp_path):
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text('[model]\nname = "qwen-72b"\ntemperature = 0.3\n')
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--model", default=None)
|
||||
@@ -158,7 +159,7 @@ def test_tavily_key_from_config(tmp_path, monkeypatch):
|
||||
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text('[api]\ntavily_key = "tvly-from-config"\n')
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
monkeypatch.delenv("TAVILY_API_KEY", raising=False)
|
||||
|
||||
key = config_mod.get_tavily_key()
|
||||
@@ -174,14 +175,14 @@ def test_tavily_key_fallback_to_env(tmp_path, monkeypatch):
|
||||
# Config exists but no tavily_key in it
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text("[api]\n")
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
monkeypatch.setenv("TAVILY_API_KEY", "tvly-from-env")
|
||||
|
||||
key = config_mod.get_tavily_key()
|
||||
assert key == "tvly-from-env"
|
||||
|
||||
|
||||
def test_apply_config_judge_section(tmp_path, monkeypatch):
|
||||
def test_apply_config_judge_section(tmp_path):
|
||||
"""apply_config() loads [judge] section and maps to argparse dests."""
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
@@ -193,7 +194,7 @@ def test_apply_config_judge_section(tmp_path, monkeypatch):
|
||||
"timeout = 30.0\n"
|
||||
"read_only_tools = false\n"
|
||||
)
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--judge", dest="judge_enabled", action="store_true", default=False)
|
||||
@@ -212,12 +213,12 @@ def test_apply_config_judge_section(tmp_path, monkeypatch):
|
||||
assert args.judge_read_only_tools is False
|
||||
|
||||
|
||||
def test_apply_config_judge_cli_overrides(tmp_path, monkeypatch):
|
||||
def test_apply_config_judge_cli_overrides(tmp_path):
|
||||
"""CLI flags override config.toml [judge] values."""
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text("[judge]\nenabled = true\nconfidence_threshold = 0.85\n")
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", cfg)
|
||||
set_config_path(str(cfg))
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--judge", dest="judge_enabled", action="store_true", default=False)
|
||||
@@ -229,3 +230,35 @@ def test_apply_config_judge_cli_overrides(tmp_path, monkeypatch):
|
||||
|
||||
assert args.judge_enabled is False # CLI wins
|
||||
assert args.judge_confidence == 0.85 # config wins (no CLI override)
|
||||
|
||||
|
||||
def test_set_config_path_overrides_default(tmp_path):
|
||||
"""set_config_path() overrides the default config location."""
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "custom.toml"
|
||||
cfg.write_text('[api]\nbase_url = "http://custom:9999"\n')
|
||||
set_config_path(str(cfg))
|
||||
assert load_config("api") == {"base_url": "http://custom:9999"}
|
||||
|
||||
|
||||
def test_env_var_overrides_default(tmp_path, monkeypatch):
|
||||
"""$TURNSTONE_CONFIG env var overrides the default config location."""
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "env.toml"
|
||||
cfg.write_text('[api]\nbase_url = "http://env:7777"\n')
|
||||
monkeypatch.setenv("TURNSTONE_CONFIG", str(cfg))
|
||||
assert load_config("api") == {"base_url": "http://env:7777"}
|
||||
|
||||
|
||||
def test_set_config_path_overrides_env_var(tmp_path, monkeypatch):
|
||||
"""set_config_path() takes precedence over $TURNSTONE_CONFIG."""
|
||||
_reset_cache()
|
||||
env_cfg = tmp_path / "env.toml"
|
||||
env_cfg.write_text('[api]\nbase_url = "http://env"\n')
|
||||
monkeypatch.setenv("TURNSTONE_CONFIG", str(env_cfg))
|
||||
|
||||
explicit_cfg = tmp_path / "explicit.toml"
|
||||
explicit_cfg.write_text('[api]\nbase_url = "http://explicit"\n')
|
||||
set_config_path(str(explicit_cfg))
|
||||
|
||||
assert load_config("api") == {"base_url": "http://explicit"}
|
||||
|
||||
+105
-17
@@ -3,7 +3,7 @@
|
||||
import asyncio
|
||||
import json
|
||||
import queue
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -47,8 +47,8 @@ class MockBroker:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_collector(broker=None, poll_interval=999, discovery_interval=999):
|
||||
"""Create a collector with long intervals so threads don't auto-fire."""
|
||||
def _make_collector(broker=None, poll_interval=0, discovery_interval=999):
|
||||
"""Create a collector with zero poll interval (no jitter delay in tests)."""
|
||||
b = broker or MockBroker()
|
||||
return ClusterCollector(
|
||||
broker=b,
|
||||
@@ -264,6 +264,57 @@ class TestCollectorPolling:
|
||||
assert q.empty()
|
||||
assert len(c._nodes["node-a"].workstreams) == 0
|
||||
|
||||
def test_poll_401_preserves_workstreams_and_marks_unreachable(self):
|
||||
"""A 401 from the server must NOT wipe workstream data."""
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
node_id="node-a",
|
||||
server_url="http://a:8080",
|
||||
reachable=True,
|
||||
workstreams={"ws1": {"id": "ws1", "name": "existing", "state": "idle"}},
|
||||
)
|
||||
|
||||
# Mock httpx to return 401
|
||||
import httpx as _httpx
|
||||
|
||||
mock_response = _httpx.Response(
|
||||
401,
|
||||
json={"error": "Unauthorized"},
|
||||
request=_httpx.Request("GET", "http://a:8080/v1/api/dashboard"),
|
||||
)
|
||||
|
||||
with patch.object(c._http_client, "get", return_value=mock_response):
|
||||
c._poll_all_nodes()
|
||||
|
||||
# Workstream data must be preserved, node marked unreachable
|
||||
assert c._nodes["node-a"].reachable is False
|
||||
assert "ws1" in c._nodes["node-a"].workstreams
|
||||
assert c._nodes["node-a"].workstreams["ws1"]["name"] == "existing"
|
||||
|
||||
def test_poll_403_preserves_workstreams(self):
|
||||
"""A 403 should also preserve state and mark unreachable."""
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
node_id="node-a",
|
||||
server_url="http://a:8080",
|
||||
reachable=True,
|
||||
workstreams={"ws1": {"id": "ws1", "name": "keep-me", "state": "running"}},
|
||||
)
|
||||
|
||||
import httpx as _httpx
|
||||
|
||||
mock_response = _httpx.Response(
|
||||
403,
|
||||
json={"error": "Forbidden"},
|
||||
request=_httpx.Request("GET", "http://a:8080/v1/api/dashboard"),
|
||||
)
|
||||
|
||||
with patch.object(c._http_client, "get", return_value=mock_response):
|
||||
c._poll_all_nodes()
|
||||
|
||||
assert c._nodes["node-a"].reachable is False
|
||||
assert "ws1" in c._nodes["node-a"].workstreams
|
||||
|
||||
|
||||
class TestCollectorEvents:
|
||||
"""Real-time event handling from cluster channel."""
|
||||
@@ -902,6 +953,8 @@ class TestConsoleWorkstreamCreation:
|
||||
],
|
||||
2,
|
||||
)
|
||||
# get_all_nodes delegates to get_nodes (mirrors real implementation)
|
||||
collector.get_all_nodes.side_effect = lambda: collector.get_nodes.return_value[0]
|
||||
return collector
|
||||
|
||||
@pytest.fixture()
|
||||
@@ -1050,6 +1103,39 @@ class TestConsoleWorkstreamCreation:
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["target_node"] == "pool"
|
||||
|
||||
def test_create_with_resume_ws_directed(self, client_and_broker, mock_collector):
|
||||
"""resume_ws is forwarded in directed dispatch."""
|
||||
client, broker = client_and_broker
|
||||
resp = client.post(
|
||||
"/v1/api/cluster/workstreams/new",
|
||||
json={"node_id": "node-a", "resume_ws": "old-ws-id-123"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
msg = json.loads(broker.push_inbound.call_args[0][0])
|
||||
assert msg["resume_ws"] == "old-ws-id-123"
|
||||
|
||||
def test_create_with_resume_ws_pool(self, client_and_broker, mock_collector):
|
||||
"""resume_ws is forwarded in pool dispatch."""
|
||||
client, broker = client_and_broker
|
||||
resp = client.post(
|
||||
"/v1/api/cluster/workstreams/new",
|
||||
json={"node_id": "pool", "resume_ws": "old-ws-id-456"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
msg = json.loads(broker.push_inbound.call_args[0][0])
|
||||
assert msg["resume_ws"] == "old-ws-id-456"
|
||||
|
||||
def test_create_with_resume_ws_auto(self, client_and_broker, mock_collector):
|
||||
"""resume_ws is forwarded in auto-select dispatch."""
|
||||
client, broker = client_and_broker
|
||||
resp = client.post(
|
||||
"/v1/api/cluster/workstreams/new",
|
||||
json={"resume_ws": "old-ws-id-789"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
msg = json.loads(broker.push_inbound.call_args[0][0])
|
||||
assert msg["resume_ws"] == "old-ws-id-789"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Proxy tests
|
||||
@@ -1182,47 +1268,49 @@ class TestProxyRewriting:
|
||||
class TestPickBestNode:
|
||||
"""Test the _pick_best_node helper."""
|
||||
|
||||
@staticmethod
|
||||
def _mock_collector(nodes: list) -> MagicMock:
|
||||
collector = MagicMock(spec=ClusterCollector)
|
||||
collector.get_nodes.return_value = (nodes, len(nodes))
|
||||
collector.get_all_nodes.side_effect = lambda: collector.get_nodes.return_value[0]
|
||||
return collector
|
||||
|
||||
def test_picks_node_with_most_headroom(self):
|
||||
from turnstone.console.server import _pick_best_node
|
||||
|
||||
collector = MagicMock(spec=ClusterCollector)
|
||||
collector.get_nodes.return_value = (
|
||||
collector = self._mock_collector(
|
||||
[
|
||||
{"node_id": "busy", "reachable": True, "max_ws": 10, "ws_total": 9},
|
||||
{"node_id": "free", "reachable": True, "max_ws": 10, "ws_total": 2},
|
||||
{"node_id": "mid", "reachable": True, "max_ws": 10, "ws_total": 5},
|
||||
],
|
||||
3,
|
||||
]
|
||||
)
|
||||
assert _pick_best_node(collector) == "free"
|
||||
|
||||
def test_skips_unreachable_nodes(self):
|
||||
from turnstone.console.server import _pick_best_node
|
||||
|
||||
collector = MagicMock(spec=ClusterCollector)
|
||||
collector.get_nodes.return_value = (
|
||||
collector = self._mock_collector(
|
||||
[
|
||||
{"node_id": "down", "reachable": False, "max_ws": 10, "ws_total": 0},
|
||||
{"node_id": "up", "reachable": True, "max_ws": 10, "ws_total": 5},
|
||||
],
|
||||
2,
|
||||
]
|
||||
)
|
||||
assert _pick_best_node(collector) == "up"
|
||||
|
||||
def test_returns_empty_when_no_nodes(self):
|
||||
from turnstone.console.server import _pick_best_node
|
||||
|
||||
collector = MagicMock(spec=ClusterCollector)
|
||||
collector.get_nodes.return_value = ([], 0)
|
||||
collector = self._mock_collector([])
|
||||
assert _pick_best_node(collector) == ""
|
||||
|
||||
def test_returns_empty_when_all_unreachable(self):
|
||||
from turnstone.console.server import _pick_best_node
|
||||
|
||||
collector = MagicMock(spec=ClusterCollector)
|
||||
collector.get_nodes.return_value = (
|
||||
[{"node_id": "down", "reachable": False, "max_ws": 10, "ws_total": 0}],
|
||||
1,
|
||||
collector = self._mock_collector(
|
||||
[
|
||||
{"node_id": "down", "reachable": False, "max_ws": 10, "ws_total": 0},
|
||||
]
|
||||
)
|
||||
assert _pick_best_node(collector) == ""
|
||||
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
"""Tests for turnstone.core.env — subprocess environment scrubbing."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from unittest.mock import patch
|
||||
|
||||
from turnstone.core.env import _is_safe, _is_secret, scrubbed_env
|
||||
|
||||
|
||||
class TestIsSecret:
|
||||
def test_explicit_scrub_list(self):
|
||||
assert _is_secret("OPENAI_API_KEY") is True
|
||||
assert _is_secret("ANTHROPIC_API_KEY") is True
|
||||
assert _is_secret("TURNSTONE_JWT_SECRET") is True
|
||||
assert _is_secret("AWS_SECRET_ACCESS_KEY") is True
|
||||
|
||||
def test_suffix_matching(self):
|
||||
assert _is_secret("MY_CUSTOM_API_KEY") is True
|
||||
assert _is_secret("DB_PASSWORD") is True
|
||||
assert _is_secret("AUTH_TOKEN") is True
|
||||
assert _is_secret("SERVICE_CREDENTIAL") is True
|
||||
assert _is_secret("GCP_CREDENTIALS") is True
|
||||
|
||||
def test_safe_vars_not_secret(self):
|
||||
assert _is_secret("PATH") is False
|
||||
assert _is_secret("HOME") is False
|
||||
assert _is_secret("LANG") is False
|
||||
|
||||
def test_no_false_positives_on_substring(self):
|
||||
"""Suffix matching avoids false positives like MONKEYTYPE."""
|
||||
assert _is_secret("MONKEYTYPE") is False
|
||||
assert _is_secret("KEYBOARD_LAYOUT") is False
|
||||
assert _is_secret("PYTHONPATH") is False
|
||||
assert _is_secret("EDITOR") is False
|
||||
assert _is_secret("GOPATH") is False
|
||||
|
||||
|
||||
class TestIsSafe:
|
||||
def test_safe_names(self):
|
||||
assert _is_safe("PATH") is True
|
||||
assert _is_safe("HOME") is True
|
||||
assert _is_safe("TERM") is True
|
||||
assert _is_safe("MANWIDTH") is True
|
||||
|
||||
def test_safe_prefixes(self):
|
||||
assert _is_safe("LC_ALL") is True
|
||||
assert _is_safe("LC_CTYPE") is True
|
||||
assert _is_safe("XDG_RUNTIME_DIR") is True
|
||||
|
||||
def test_non_safe_names(self):
|
||||
assert _is_safe("OPENAI_API_KEY") is False
|
||||
assert _is_safe("CUSTOM_VAR") is False
|
||||
|
||||
|
||||
class TestScrubbedEnv:
|
||||
def test_strips_api_keys(self):
|
||||
fake_env = {
|
||||
"PATH": "/usr/bin",
|
||||
"HOME": "/home/user",
|
||||
"OPENAI_API_KEY": "sk-secret",
|
||||
"ANTHROPIC_API_KEY": "ant-secret",
|
||||
"CUSTOM_VAR": "safe_value",
|
||||
}
|
||||
with patch.dict(os.environ, fake_env, clear=True):
|
||||
result = scrubbed_env()
|
||||
|
||||
assert result["PATH"] == "/usr/bin"
|
||||
assert result["HOME"] == "/home/user"
|
||||
assert result["CUSTOM_VAR"] == "safe_value"
|
||||
assert "OPENAI_API_KEY" not in result
|
||||
assert "ANTHROPIC_API_KEY" not in result
|
||||
|
||||
def test_strips_pattern_matched_secrets(self):
|
||||
fake_env = {
|
||||
"PATH": "/usr/bin",
|
||||
"MY_SERVICE_TOKEN": "tok-123",
|
||||
"DB_PASSWORD": "pass123",
|
||||
}
|
||||
with patch.dict(os.environ, fake_env, clear=True):
|
||||
result = scrubbed_env()
|
||||
|
||||
assert "MY_SERVICE_TOKEN" not in result
|
||||
assert "DB_PASSWORD" not in result
|
||||
|
||||
def test_extra_vars_merged(self):
|
||||
fake_env = {"PATH": "/usr/bin"}
|
||||
with patch.dict(os.environ, fake_env, clear=True):
|
||||
result = scrubbed_env(extra={"MANWIDTH": "80"})
|
||||
|
||||
assert result["MANWIDTH"] == "80"
|
||||
assert result["PATH"] == "/usr/bin"
|
||||
|
||||
def test_passthrough_overrides_scrub(self):
|
||||
fake_env = {
|
||||
"PATH": "/usr/bin",
|
||||
"OPENAI_API_KEY": "sk-needed",
|
||||
}
|
||||
with patch.dict(os.environ, fake_env, clear=True):
|
||||
result = scrubbed_env(passthrough=["OPENAI_API_KEY"])
|
||||
|
||||
assert result["OPENAI_API_KEY"] == "sk-needed"
|
||||
|
||||
def test_preserves_locale_vars(self):
|
||||
fake_env = {
|
||||
"PATH": "/usr/bin",
|
||||
"LC_ALL": "en_US.UTF-8",
|
||||
"LC_CTYPE": "en_US.UTF-8",
|
||||
}
|
||||
with patch.dict(os.environ, fake_env, clear=True):
|
||||
result = scrubbed_env()
|
||||
|
||||
assert result["LC_ALL"] == "en_US.UTF-8"
|
||||
assert result["LC_CTYPE"] == "en_US.UTF-8"
|
||||
|
||||
def test_preserves_unknown_non_secret_vars(self):
|
||||
fake_env = {
|
||||
"PATH": "/usr/bin",
|
||||
"PYTHONPATH": "/opt/lib",
|
||||
"GOPATH": "/home/user/go",
|
||||
}
|
||||
with patch.dict(os.environ, fake_env, clear=True):
|
||||
result = scrubbed_env()
|
||||
|
||||
assert result["PYTHONPATH"] == "/opt/lib"
|
||||
assert result["GOPATH"] == "/home/user/go"
|
||||
|
||||
def test_extra_can_reintroduce_scrubbed_var(self):
|
||||
"""extra= intentionally overrides scrubbing (operator-controlled)."""
|
||||
fake_env = {"PATH": "/usr/bin", "OPENAI_API_KEY": "sk-original"}
|
||||
with patch.dict(os.environ, fake_env, clear=True):
|
||||
result = scrubbed_env(extra={"OPENAI_API_KEY": "sk-injected"})
|
||||
|
||||
assert result["OPENAI_API_KEY"] == "sk-injected"
|
||||
|
||||
def test_less_prefix_does_not_leak_secrets(self):
|
||||
"""LESS pager vars are safe but LESS_SECRET_TOKEN is not."""
|
||||
fake_env = {
|
||||
"PATH": "/usr/bin",
|
||||
"LESS": "-R",
|
||||
"LESSOPEN": "| lesspipe %s",
|
||||
"LESS_SECRET_TOKEN": "tok-secret",
|
||||
}
|
||||
with patch.dict(os.environ, fake_env, clear=True):
|
||||
result = scrubbed_env()
|
||||
|
||||
assert result["LESS"] == "-R"
|
||||
assert result["LESSOPEN"] == "| lesspipe %s"
|
||||
assert "LESS_SECRET_TOKEN" not in result
|
||||
@@ -8,18 +8,8 @@ from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime
|
||||
|
||||
import pytest
|
||||
import sqlalchemy as sa
|
||||
|
||||
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"))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Roles
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -221,6 +221,63 @@ class TestBackendHealthMonitor:
|
||||
mon.record_failure()
|
||||
assert mon.circuit_state == CircuitState.OPEN # type: ignore[comparison-overlap]
|
||||
|
||||
def test_probe_loop_autonomous_recovery(
|
||||
self, mock_client: MagicMock, mock_metrics: MagicMock
|
||||
) -> None:
|
||||
"""_probe_loop transitions OPEN → HALF_OPEN → CLOSED without user requests."""
|
||||
# Use very short intervals so the test is fast
|
||||
mon = BackendHealthMonitor(
|
||||
client=mock_client,
|
||||
probe_interval=0.05,
|
||||
probe_timeout=1.0,
|
||||
failure_threshold=1,
|
||||
cooldown=0.1,
|
||||
)
|
||||
# Trip the circuit
|
||||
mon.record_failure()
|
||||
assert mon.circuit_state == CircuitState.OPEN
|
||||
|
||||
# Backend is healthy — probe_once will succeed
|
||||
mock_client.with_options.return_value.models.list.return_value = MagicMock()
|
||||
|
||||
# Start the probe loop and wait for autonomous recovery
|
||||
mon.start()
|
||||
try:
|
||||
import time
|
||||
|
||||
deadline = time.monotonic() + 5.0
|
||||
while mon.circuit_state != CircuitState.CLOSED and time.monotonic() < deadline:
|
||||
time.sleep(0.05)
|
||||
assert mon.circuit_state == CircuitState.CLOSED
|
||||
# User requests should flow again without anyone calling acquire_request_permit
|
||||
assert mon.acquire_request_permit() is True
|
||||
finally:
|
||||
mon.stop()
|
||||
if mon._thread:
|
||||
mon._thread.join(timeout=2.0)
|
||||
|
||||
def test_probe_loop_no_user_permit_during_probe(
|
||||
self, mock_client: MagicMock, mock_metrics: MagicMock
|
||||
) -> None:
|
||||
"""While background probe is in HALF_OPEN, user requests are blocked."""
|
||||
mon = BackendHealthMonitor(
|
||||
client=mock_client,
|
||||
probe_interval=0.05,
|
||||
probe_timeout=1.0,
|
||||
failure_threshold=1,
|
||||
cooldown=0.1,
|
||||
)
|
||||
mon.record_failure()
|
||||
assert mon.circuit_state == CircuitState.OPEN
|
||||
|
||||
# Force into HALF_OPEN as the probe loop would
|
||||
with mon._lock:
|
||||
mon._state = CircuitState.HALF_OPEN
|
||||
mon._half_open_permit = False # probe consumes it
|
||||
|
||||
# User requests should be blocked — only the probe gets through
|
||||
assert mon.acquire_request_permit() is False
|
||||
|
||||
def test_stop_thread(self, mock_client: MagicMock) -> None:
|
||||
"""stop() signals the probe loop to exit."""
|
||||
mon = _make_monitor(mock_client)
|
||||
|
||||
@@ -4,16 +4,6 @@ from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db(tmp_path):
|
||||
"""Fresh SQLite backend for each test."""
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
def _make_verdict_kwargs(**overrides):
|
||||
"""Build default kwargs for create_intent_verdict."""
|
||||
|
||||
@@ -26,6 +26,7 @@ from turnstone.console.server import (
|
||||
admin_get_mcp_server,
|
||||
admin_import_mcp_config,
|
||||
admin_list_mcp_servers,
|
||||
admin_mcp_reload,
|
||||
admin_update_mcp_server,
|
||||
)
|
||||
from turnstone.core.auth import AuthResult
|
||||
@@ -96,6 +97,11 @@ _ROUTES = [
|
||||
admin_import_mcp_config,
|
||||
methods=["POST"],
|
||||
),
|
||||
Route(
|
||||
"/api/admin/mcp-servers/reload",
|
||||
admin_mcp_reload,
|
||||
methods=["POST"],
|
||||
),
|
||||
Route(
|
||||
"/api/admin/mcp-servers/{server_id}",
|
||||
admin_get_mcp_server,
|
||||
@@ -115,6 +121,21 @@ _ROUTES = [
|
||||
]
|
||||
|
||||
|
||||
def _routes_with_internal() -> list[Mount]:
|
||||
"""Routes including the node-side internal endpoint (lazy-imported)."""
|
||||
from turnstone.server import internal_mcp_reload
|
||||
|
||||
return [
|
||||
Mount(
|
||||
"/v1",
|
||||
routes=[
|
||||
*_ROUTES[0].routes, # type: ignore[union-attr]
|
||||
Route("/api/_internal/mcp-reload", internal_mcp_reload, methods=["POST"]),
|
||||
],
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path):
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
@@ -550,6 +571,7 @@ def _fake_request(*nodes: dict[str, Any], proxy_client: Any = None) -> MagicMock
|
||||
"""Build a minimal mock request with collector and proxy_client."""
|
||||
collector = MagicMock()
|
||||
collector.get_nodes.return_value = (list(nodes), len(nodes))
|
||||
collector.get_all_nodes.side_effect = lambda: collector.get_nodes.return_value[0]
|
||||
req = MagicMock()
|
||||
req.app.state.collector = collector
|
||||
req.app.state.proxy_client = proxy_client or AsyncMock()
|
||||
@@ -693,3 +715,175 @@ class TestNotifyNodesMcpReload:
|
||||
result = await _notify_nodes_mcp_reload(req)
|
||||
assert result["n1"] == {"reloaded": 2}
|
||||
assert "error" in result["n2"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Console reload endpoint: POST /v1/api/admin/mcp-servers/reload
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAdminMcpReloadEndpoint:
|
||||
"""HTTP-level tests for the console reload endpoint."""
|
||||
|
||||
def test_reload_success(self, client: TestClient) -> None:
|
||||
"""Reload endpoint returns status ok and fan-out results."""
|
||||
with patch(
|
||||
"turnstone.console.server._notify_nodes_mcp_reload",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"n1": {"reloaded": 3}},
|
||||
):
|
||||
r = client.post("/v1/api/admin/mcp-servers/reload")
|
||||
assert r.status_code == 200
|
||||
data = r.json()
|
||||
assert data["status"] == "ok"
|
||||
assert data["results"] == {"n1": {"reloaded": 3}}
|
||||
|
||||
def test_reload_empty_cluster(self, client: TestClient) -> None:
|
||||
"""Reload with no nodes returns empty results."""
|
||||
with patch(
|
||||
"turnstone.console.server._notify_nodes_mcp_reload",
|
||||
new_callable=AsyncMock,
|
||||
return_value={},
|
||||
):
|
||||
r = client.post("/v1/api/admin/mcp-servers/reload")
|
||||
assert r.status_code == 200
|
||||
data = r.json()
|
||||
assert data["status"] == "ok"
|
||||
assert data["results"] == {}
|
||||
|
||||
def test_reload_permission_denied(self, client_no_perm: TestClient) -> None:
|
||||
"""Reload without admin.mcp permission is rejected."""
|
||||
r = client_no_perm.post("/v1/api/admin/mcp-servers/reload")
|
||||
assert r.status_code == 403
|
||||
assert "admin.mcp" in r.json()["error"]
|
||||
|
||||
def test_reload_no_storage(self) -> None:
|
||||
"""Reload returns 503 when auth_storage is not available."""
|
||||
app = Starlette(
|
||||
routes=_ROUTES,
|
||||
middleware=[Middleware(_InjectAuthMiddleware)],
|
||||
)
|
||||
# Deliberately omit app.state.auth_storage
|
||||
no_storage_client = TestClient(app, raise_server_exceptions=False)
|
||||
r = no_storage_client.post("/v1/api/admin/mcp-servers/reload")
|
||||
assert r.status_code == 503
|
||||
|
||||
def test_reload_mixed_node_results(self, client: TestClient) -> None:
|
||||
"""Reload propagates per-node errors in results."""
|
||||
with patch(
|
||||
"turnstone.console.server._notify_nodes_mcp_reload",
|
||||
new_callable=AsyncMock,
|
||||
return_value={
|
||||
"n1": {"reloaded": 2},
|
||||
"n2": {"error": "Connection refused"},
|
||||
},
|
||||
):
|
||||
r = client.post("/v1/api/admin/mcp-servers/reload")
|
||||
assert r.status_code == 200
|
||||
data = r.json()
|
||||
assert data["results"]["n1"] == {"reloaded": 2}
|
||||
assert "error" in data["results"]["n2"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Node reload endpoint: POST /v1/api/_internal/mcp-reload
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestInternalMcpReloadEndpoint:
|
||||
"""HTTP-level tests for the node-side MCP reload endpoint."""
|
||||
|
||||
@pytest.fixture()
|
||||
def node_client(self, storage: SQLiteBackend) -> TestClient:
|
||||
"""TestClient with an MCP client manager on app.state."""
|
||||
app = Starlette(
|
||||
routes=_routes_with_internal(),
|
||||
middleware=[Middleware(_InjectAuthMiddleware)],
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
mgr = MagicMock()
|
||||
mgr.reconcile_sync.return_value = {
|
||||
"added": ["new-srv"],
|
||||
"removed": [],
|
||||
"updated": [],
|
||||
}
|
||||
app.state.mcp_client = mgr
|
||||
return TestClient(app, raise_server_exceptions=False)
|
||||
|
||||
def test_reload_calls_reconcile(self, node_client: TestClient, storage: SQLiteBackend) -> None:
|
||||
"""Reload endpoint calls reconcile_sync and returns its result."""
|
||||
with patch("turnstone.core.storage._registry.get_storage", return_value=storage):
|
||||
r = node_client.post("/v1/api/_internal/mcp-reload")
|
||||
assert r.status_code == 200
|
||||
data = r.json()
|
||||
assert data["status"] == "ok"
|
||||
assert data["added"] == ["new-srv"]
|
||||
assert data["removed"] == []
|
||||
assert data["updated"] == []
|
||||
|
||||
def test_reload_passes_storage_to_reconcile(
|
||||
self,
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""Verify reconcile_sync receives the storage backend."""
|
||||
app = Starlette(
|
||||
routes=_routes_with_internal(),
|
||||
middleware=[Middleware(_InjectAuthMiddleware)],
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
mgr = MagicMock()
|
||||
mgr.reconcile_sync.return_value = {"added": [], "removed": [], "updated": []}
|
||||
app.state.mcp_client = mgr
|
||||
c = TestClient(app, raise_server_exceptions=False)
|
||||
with patch("turnstone.core.storage._registry.get_storage", return_value=storage):
|
||||
r = c.post("/v1/api/_internal/mcp-reload")
|
||||
assert r.status_code == 200
|
||||
mgr.reconcile_sync.assert_called_once_with(storage)
|
||||
|
||||
def test_reload_creates_manager_when_missing(self, storage: SQLiteBackend) -> None:
|
||||
"""When mcp_client is absent, a new MCPClientManager is created."""
|
||||
app = Starlette(
|
||||
routes=_routes_with_internal(),
|
||||
middleware=[Middleware(_InjectAuthMiddleware)],
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
# No mcp_client on app.state
|
||||
c = TestClient(app, raise_server_exceptions=False)
|
||||
with (
|
||||
patch("turnstone.core.storage._registry.get_storage", return_value=storage),
|
||||
patch("turnstone.core.mcp_client.MCPClientManager") as mock_cls,
|
||||
):
|
||||
mock_mgr = MagicMock()
|
||||
mock_mgr.reconcile_sync.return_value = {
|
||||
"added": [],
|
||||
"removed": [],
|
||||
"updated": [],
|
||||
}
|
||||
mock_cls.return_value = mock_mgr
|
||||
r = c.post("/v1/api/_internal/mcp-reload")
|
||||
assert r.status_code == 200
|
||||
mock_cls.assert_called_once_with({})
|
||||
mock_mgr.start.assert_called_once()
|
||||
mock_mgr.reconcile_sync.assert_called_once_with(storage)
|
||||
|
||||
def test_reload_reconcile_result_in_response(self, storage: SQLiteBackend) -> None:
|
||||
"""Full reconcile result fields (added/removed/updated) appear in JSON."""
|
||||
app = Starlette(
|
||||
routes=_routes_with_internal(),
|
||||
middleware=[Middleware(_InjectAuthMiddleware)],
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
mgr = MagicMock()
|
||||
mgr.reconcile_sync.return_value = {
|
||||
"added": ["a"],
|
||||
"removed": ["b"],
|
||||
"updated": ["c"],
|
||||
}
|
||||
app.state.mcp_client = mgr
|
||||
c = TestClient(app, raise_server_exceptions=False)
|
||||
with patch("turnstone.core.storage._registry.get_storage", return_value=storage):
|
||||
r = c.post("/v1/api/_internal/mcp-reload")
|
||||
data = r.json()
|
||||
assert data["added"] == ["a"]
|
||||
assert data["removed"] == ["b"]
|
||||
assert data["updated"] == ["c"]
|
||||
|
||||
@@ -398,6 +398,72 @@ class TestResolveInstallConfig:
|
||||
config = resolve_install_config(server, "remote", 0)
|
||||
assert config["url"] == "https://us-east.example.com/mcp"
|
||||
|
||||
def test_remote_variable_substitution_invalid_scheme(self) -> None:
|
||||
server = RegistryServer(
|
||||
name="io.example/test",
|
||||
version="1.0.0",
|
||||
remotes=[
|
||||
RegistryRemote(
|
||||
type="streamable-http",
|
||||
url="{scheme}://evil.example.com/mcp",
|
||||
variables={
|
||||
"scheme": RegistryRemoteVariable(is_required=True),
|
||||
},
|
||||
)
|
||||
],
|
||||
)
|
||||
with pytest.raises(MCPRegistryError, match="Invalid URL scheme"):
|
||||
resolve_install_config(server, "remote", 0, variables={"scheme": "file"})
|
||||
|
||||
def test_remote_variable_substitution_preserves_valid_scheme(self) -> None:
|
||||
server = RegistryServer(
|
||||
name="io.example/test",
|
||||
version="1.0.0",
|
||||
remotes=[
|
||||
RegistryRemote(
|
||||
type="streamable-http",
|
||||
url="https://{host}.example.com/mcp",
|
||||
variables={
|
||||
"host": RegistryRemoteVariable(is_required=True),
|
||||
},
|
||||
)
|
||||
],
|
||||
)
|
||||
config = resolve_install_config(server, "remote", 0, variables={"host": "api"})
|
||||
assert config["url"] == "https://api.example.com/mcp"
|
||||
|
||||
def test_remote_variable_substitution_missing_hostname(self) -> None:
|
||||
"""URL like https:///mcp has valid scheme but no hostname."""
|
||||
server = RegistryServer(
|
||||
name="io.example/test",
|
||||
version="1.0.0",
|
||||
remotes=[
|
||||
RegistryRemote(
|
||||
type="streamable-http",
|
||||
url="https:///mcp",
|
||||
)
|
||||
],
|
||||
)
|
||||
with pytest.raises(MCPRegistryError, match="hostname is missing"):
|
||||
resolve_install_config(server, "remote", 0)
|
||||
|
||||
def test_remote_variable_substitution_embedded_credentials(self) -> None:
|
||||
server = RegistryServer(
|
||||
name="io.example/test",
|
||||
version="1.0.0",
|
||||
remotes=[
|
||||
RegistryRemote(
|
||||
type="streamable-http",
|
||||
url="https://{creds}@example.com/mcp",
|
||||
variables={
|
||||
"creds": RegistryRemoteVariable(is_required=True),
|
||||
},
|
||||
)
|
||||
],
|
||||
)
|
||||
with pytest.raises(MCPRegistryError, match="embedded credentials"):
|
||||
resolve_install_config(server, "remote", 0, variables={"creds": "user:pass"})
|
||||
|
||||
def test_remote_no_remotes(self) -> None:
|
||||
server = RegistryServer(name="io.example/test", version="1.0.0")
|
||||
with pytest.raises(MCPRegistryError, match="no remote"):
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from unittest.mock import AsyncMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from starlette.applications import Starlette
|
||||
@@ -18,11 +18,13 @@ if TYPE_CHECKING:
|
||||
from starlette.responses import Response
|
||||
|
||||
from turnstone.console.server import (
|
||||
_get_registry_url,
|
||||
admin_registry_install,
|
||||
admin_registry_search,
|
||||
)
|
||||
from turnstone.core.auth import AuthResult
|
||||
from turnstone.core.mcp_registry import (
|
||||
DEFAULT_REGISTRY_URL,
|
||||
MCPRegistryError,
|
||||
RegistryPackage,
|
||||
RegistryRemote,
|
||||
@@ -362,7 +364,7 @@ class TestRegistryInstall:
|
||||
def test_install_max_servers(self, client: TestClient, storage: SQLiteBackend) -> None:
|
||||
import uuid
|
||||
|
||||
for i in range(50):
|
||||
for i in range(200):
|
||||
storage.create_mcp_server(
|
||||
server_id=uuid.uuid4().hex,
|
||||
name=f"server-{i}",
|
||||
@@ -549,3 +551,73 @@ class TestRegistryInstall:
|
||||
|
||||
assert resp.status_code == 409
|
||||
assert "custom 'name'" in resp.json()["error"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _get_registry_url fallback chain tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _mock_request(storage: Any = None, config_store: Any = None) -> MagicMock:
|
||||
"""Build a mock Request with app.state.auth_storage and app.state.config_store."""
|
||||
request = MagicMock()
|
||||
request.app.state.auth_storage = storage
|
||||
request.app.state.config_store = config_store
|
||||
return request
|
||||
|
||||
|
||||
class TestGetRegistryUrl:
|
||||
"""Verify three-tier URL resolution: DB setting -> config.toml -> default."""
|
||||
|
||||
def test_returns_db_setting_when_available(self) -> None:
|
||||
config_store = MagicMock()
|
||||
config_store.get.return_value = "https://custom.registry.example.com"
|
||||
request = _mock_request(config_store=config_store)
|
||||
|
||||
with patch("turnstone.core.config.load_config", return_value={}):
|
||||
url = _get_registry_url(request)
|
||||
|
||||
assert url == "https://custom.registry.example.com"
|
||||
config_store.get.assert_called_once_with("mcp.registry_url")
|
||||
|
||||
def test_falls_back_to_config_when_config_store_returns_empty(self) -> None:
|
||||
config_store = MagicMock()
|
||||
config_store.get.return_value = ""
|
||||
request = _mock_request(config_store=config_store)
|
||||
|
||||
with patch(
|
||||
"turnstone.core.config.load_config",
|
||||
return_value={"registry_url": "https://config.registry.example.com"},
|
||||
):
|
||||
url = _get_registry_url(request)
|
||||
|
||||
assert url == "https://config.registry.example.com"
|
||||
|
||||
def test_falls_back_to_config_when_no_config_store(self) -> None:
|
||||
request = _mock_request()
|
||||
|
||||
with patch(
|
||||
"turnstone.core.config.load_config",
|
||||
return_value={"registry_url": "https://config.registry.example.com"},
|
||||
):
|
||||
url = _get_registry_url(request)
|
||||
|
||||
assert url == "https://config.registry.example.com"
|
||||
|
||||
def test_falls_back_to_default_when_both_unavailable(self) -> None:
|
||||
config_store = MagicMock()
|
||||
config_store.get.return_value = ""
|
||||
request = _mock_request(config_store=config_store)
|
||||
|
||||
with patch("turnstone.core.config.load_config", return_value={}):
|
||||
url = _get_registry_url(request)
|
||||
|
||||
assert url == DEFAULT_REGISTRY_URL
|
||||
|
||||
def test_falls_back_to_default_when_no_config_store_or_config(self) -> None:
|
||||
request = _mock_request()
|
||||
|
||||
with patch("turnstone.core.config.load_config", return_value={}):
|
||||
url = _get_registry_url(request)
|
||||
|
||||
assert url == DEFAULT_REGISTRY_URL
|
||||
|
||||
@@ -3,16 +3,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def db(tmp_path):
|
||||
"""Fresh SQLite backend for each test."""
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
if TYPE_CHECKING:
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
def _make_id() -> str:
|
||||
|
||||
@@ -10,6 +10,7 @@ import pytest
|
||||
from turnstone.core.model_registry import (
|
||||
ModelConfig,
|
||||
ModelRegistry,
|
||||
detect_model,
|
||||
load_model_registry,
|
||||
)
|
||||
|
||||
@@ -590,3 +591,36 @@ class TestProtocolModel:
|
||||
assert isinstance(restored, CreateWorkstreamMessage)
|
||||
assert restored.model == "local"
|
||||
assert restored.name == "ws1"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# detect_model — startup timeout
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDetectModelTimeout:
|
||||
def test_uses_short_timeout_and_no_retries(self) -> None:
|
||||
"""detect_model() uses with_options(timeout=10, max_retries=0)."""
|
||||
mock_model = MagicMock()
|
||||
mock_model.id = "test-model"
|
||||
mock_model.owned_by = "test"
|
||||
|
||||
fast_client = MagicMock()
|
||||
fast_client.models.list.return_value = MagicMock(data=[mock_model])
|
||||
|
||||
client = MagicMock()
|
||||
client.with_options.return_value = fast_client
|
||||
|
||||
result = detect_model(client, provider="openai")
|
||||
client.with_options.assert_called_once_with(timeout=10.0, max_retries=0)
|
||||
fast_client.models.list.assert_called_once()
|
||||
assert result[0] == "test-model"
|
||||
|
||||
def test_connection_error_non_fatal(self) -> None:
|
||||
"""detect_model(fatal=False) returns (None, None) on connection error."""
|
||||
client = MagicMock()
|
||||
client.with_options.return_value = client
|
||||
client.models.list.side_effect = OSError("Connection refused")
|
||||
|
||||
result = detect_model(client, provider="openai", fatal=False)
|
||||
assert result == (None, None)
|
||||
|
||||
+172
-3
@@ -22,6 +22,7 @@ from turnstone.core.oidc import (
|
||||
load_oidc_config,
|
||||
provision_oidc_user,
|
||||
validate_id_token,
|
||||
validate_issuer_url,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -293,6 +294,162 @@ class TestLoadOIDCConfig:
|
||||
assert cfg.redirect_base == "http://localhost:8000"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SSRF Validation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestValidateIssuerURL:
|
||||
"""Tests for ``validate_issuer_url`` SSRF protection."""
|
||||
|
||||
def test_valid_https_url(self):
|
||||
"""Public HTTPS issuer URL passes validation."""
|
||||
# Should not raise -- mock DNS to return a public IP.
|
||||
with patch(
|
||||
"socket.getaddrinfo",
|
||||
return_value=[
|
||||
(2, 1, 6, "", ("93.184.216.34", 0)),
|
||||
],
|
||||
):
|
||||
validate_issuer_url("https://idp.example.com")
|
||||
|
||||
def test_rejects_http_non_localhost(self):
|
||||
"""HTTP is rejected for non-localhost hosts."""
|
||||
with pytest.raises(OIDCError, match="must use HTTPS"):
|
||||
validate_issuer_url("http://idp.example.com")
|
||||
|
||||
def test_allows_http_localhost(self):
|
||||
"""HTTP is allowed for localhost (development)."""
|
||||
with patch(
|
||||
"socket.getaddrinfo",
|
||||
return_value=[
|
||||
(2, 1, 6, "", ("127.0.0.1", 0)),
|
||||
],
|
||||
):
|
||||
validate_issuer_url("http://localhost:8080")
|
||||
|
||||
def test_allows_http_localhost_subdomain(self):
|
||||
"""HTTP is allowed for *.localhost subdomains."""
|
||||
with patch(
|
||||
"socket.getaddrinfo",
|
||||
return_value=[
|
||||
(2, 1, 6, "", ("127.0.0.1", 0)),
|
||||
],
|
||||
):
|
||||
validate_issuer_url("http://keycloak.localhost:8080")
|
||||
|
||||
def test_rejects_embedded_credentials(self):
|
||||
"""URLs with userinfo (user:pass@host) are rejected."""
|
||||
with pytest.raises(OIDCError, match="embedded credentials"):
|
||||
validate_issuer_url("https://admin:secret@idp.example.com")
|
||||
|
||||
def test_rejects_username_only(self):
|
||||
"""URLs with just a username are rejected."""
|
||||
with pytest.raises(OIDCError, match="embedded credentials"):
|
||||
validate_issuer_url("https://admin@idp.example.com")
|
||||
|
||||
def test_rejects_no_hostname(self):
|
||||
"""URLs without a hostname are rejected."""
|
||||
with pytest.raises(OIDCError, match="no hostname"):
|
||||
validate_issuer_url("https://")
|
||||
|
||||
def test_rejects_private_10_range(self):
|
||||
"""Hostnames resolving to 10.x.x.x are rejected."""
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("10.0.0.1", 0))]),
|
||||
pytest.raises(OIDCError, match="non-public address.*10.0.0.1"),
|
||||
):
|
||||
validate_issuer_url("https://internal.corp.example.com")
|
||||
|
||||
def test_rejects_private_172_range(self):
|
||||
"""Hostnames resolving to 172.16-31.x.x are rejected."""
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("172.16.0.1", 0))]),
|
||||
pytest.raises(OIDCError, match="non-public address.*172.16.0.1"),
|
||||
):
|
||||
validate_issuer_url("https://internal.corp.example.com")
|
||||
|
||||
def test_rejects_private_192_168_range(self):
|
||||
"""Hostnames resolving to 192.168.x.x are rejected."""
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("192.168.1.1", 0))]),
|
||||
pytest.raises(OIDCError, match="non-public address.*192.168.1.1"),
|
||||
):
|
||||
validate_issuer_url("https://internal.corp.example.com")
|
||||
|
||||
def test_rejects_loopback_127(self):
|
||||
"""Hostnames resolving to 127.x.x.x are rejected (non-localhost host)."""
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("127.0.0.1", 0))]),
|
||||
pytest.raises(OIDCError, match="non-public address.*127.0.0.1"),
|
||||
):
|
||||
validate_issuer_url("https://evil.example.com")
|
||||
|
||||
def test_rejects_ipv6_loopback(self):
|
||||
"""Hostnames resolving to ::1 are rejected (non-localhost host)."""
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=[(10, 1, 6, "", ("::1", 0, 0, 0))]),
|
||||
pytest.raises(OIDCError, match="non-public address.*::1"),
|
||||
):
|
||||
validate_issuer_url("https://evil.example.com")
|
||||
|
||||
def test_rejects_ipv6_private(self):
|
||||
"""Hostnames resolving to fc00::/7 are rejected."""
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=[(10, 1, 6, "", ("fd00::1", 0, 0, 0))]),
|
||||
pytest.raises(OIDCError, match="non-public address.*fd00::1"),
|
||||
):
|
||||
validate_issuer_url("https://evil.example.com")
|
||||
|
||||
def test_rejects_link_local(self):
|
||||
"""Hostnames resolving to link-local addresses are rejected."""
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("169.254.169.254", 0))]),
|
||||
pytest.raises(OIDCError, match="non-public address.*169.254.169.254"),
|
||||
):
|
||||
validate_issuer_url("https://metadata.internal")
|
||||
|
||||
def test_rejects_unresolvable_hostname(self):
|
||||
"""DNS resolution failure is rejected."""
|
||||
import socket as _socket
|
||||
|
||||
with (
|
||||
patch("socket.getaddrinfo", side_effect=_socket.gaierror("not found")),
|
||||
pytest.raises(OIDCError, match="cannot be resolved"),
|
||||
):
|
||||
validate_issuer_url("https://nonexistent.invalid")
|
||||
|
||||
def test_rejects_mixed_addresses(self):
|
||||
"""If any resolved address is private, the URL is rejected."""
|
||||
with (
|
||||
patch(
|
||||
"socket.getaddrinfo",
|
||||
return_value=[
|
||||
(2, 1, 6, "", ("93.184.216.34", 0)),
|
||||
(2, 1, 6, "", ("10.0.0.1", 0)),
|
||||
],
|
||||
),
|
||||
pytest.raises(OIDCError, match="non-public address.*10.0.0.1"),
|
||||
):
|
||||
validate_issuer_url("https://dual-homed.example.com")
|
||||
|
||||
def test_discover_rejects_ssrf(self):
|
||||
"""discover_oidc returns enabled=False when issuer URL fails SSRF check."""
|
||||
config = _make_config(
|
||||
issuer="http://10.0.0.1:8080",
|
||||
authorization_endpoint="",
|
||||
token_endpoint="",
|
||||
userinfo_endpoint="",
|
||||
jwks_uri="",
|
||||
)
|
||||
|
||||
async def _run():
|
||||
result = await discover_oidc(config)
|
||||
assert result.enabled is False
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Redirect URI Builder
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -869,6 +1026,9 @@ class TestApplyRoleMapping:
|
||||
|
||||
|
||||
class TestDiscoverOIDC:
|
||||
# Mock DNS result for a public IP — reused across discovery tests.
|
||||
_PUBLIC_ADDR = [(2, 1, 6, "", ("93.184.216.34", 0))]
|
||||
|
||||
def test_discover_oidc_success(self):
|
||||
"""Mock httpx response, verify endpoints populated."""
|
||||
config = _make_config(
|
||||
@@ -891,7 +1051,10 @@ class TestDiscoverOIDC:
|
||||
|
||||
async def _run():
|
||||
client = _mock_async_client(lambda url: _async_return(mock_response))
|
||||
with patch("httpx.AsyncClient", return_value=client):
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
||||
patch("httpx.AsyncClient", return_value=client),
|
||||
):
|
||||
result = await discover_oidc(config)
|
||||
|
||||
assert result.authorization_endpoint == "https://idp.example.com/authorize"
|
||||
@@ -916,7 +1079,10 @@ class TestDiscoverOIDC:
|
||||
|
||||
async def _run():
|
||||
client = _mock_async_client(_failing_get)
|
||||
with patch("httpx.AsyncClient", return_value=client):
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
||||
patch("httpx.AsyncClient", return_value=client),
|
||||
):
|
||||
result = await discover_oidc(config)
|
||||
|
||||
assert result.enabled is False
|
||||
@@ -954,7 +1120,10 @@ class TestDiscoverOIDC:
|
||||
|
||||
async def _run():
|
||||
client = _mock_async_client(lambda url: _async_return(mock_response))
|
||||
with patch("httpx.AsyncClient", return_value=client):
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
||||
patch("httpx.AsyncClient", return_value=client),
|
||||
):
|
||||
result = await discover_oidc(config)
|
||||
|
||||
assert result.enabled is False
|
||||
|
||||
@@ -6,15 +6,6 @@ import time
|
||||
|
||||
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"))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OIDC Identity CRUD
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -4,16 +4,6 @@ from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db(tmp_path):
|
||||
"""Fresh SQLite backend for each test."""
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
def _make_assessment_kwargs(**overrides):
|
||||
"""Build default kwargs for record_output_assessment."""
|
||||
|
||||
@@ -4,17 +4,6 @@ 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."""
|
||||
|
||||
@@ -0,0 +1,440 @@
|
||||
"""Integration tests for SDK governance methods against a real Starlette app.
|
||||
|
||||
Verifies round-trip serialization: SDK -> HTTP -> Starlette handler -> storage
|
||||
-> JSON response -> Pydantic model validation in the SDK client.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from starlette.applications import Starlette
|
||||
from starlette.middleware import Middleware
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.routing import Mount, Route
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
from turnstone.api.console_schemas import (
|
||||
ListOrgsResponse,
|
||||
ListRolesResponse,
|
||||
ListToolPoliciesResponse,
|
||||
OrgInfo,
|
||||
RoleInfo,
|
||||
ToolPolicyInfo,
|
||||
)
|
||||
from turnstone.api.schemas import StatusResponse
|
||||
from turnstone.console.server import (
|
||||
admin_create_policy,
|
||||
admin_create_role,
|
||||
admin_delete_policy,
|
||||
admin_delete_role,
|
||||
admin_get_org,
|
||||
admin_list_orgs,
|
||||
admin_list_policies,
|
||||
admin_list_roles,
|
||||
admin_update_org,
|
||||
admin_update_policy,
|
||||
admin_update_role,
|
||||
)
|
||||
from turnstone.core.auth import AuthResult
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
from turnstone.sdk.console import AsyncTurnstoneConsole
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth bypass middleware — injects a full-access AuthResult on every request.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _InjectAuthMiddleware(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next: Any) -> Response:
|
||||
request.state.auth_result = AuthResult(
|
||||
user_id="test-admin",
|
||||
scopes=frozenset({"approve"}),
|
||||
token_source="config",
|
||||
permissions=frozenset(
|
||||
{
|
||||
"read",
|
||||
"write",
|
||||
"approve",
|
||||
"admin.roles",
|
||||
"admin.orgs",
|
||||
"admin.policies",
|
||||
}
|
||||
),
|
||||
)
|
||||
resp: Response = await call_next(request)
|
||||
return resp
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_app() -> Starlette:
|
||||
return Starlette(
|
||||
routes=[
|
||||
Mount(
|
||||
"/v1",
|
||||
routes=[
|
||||
# Roles
|
||||
Route("/api/admin/roles", admin_list_roles),
|
||||
Route("/api/admin/roles", admin_create_role, methods=["POST"]),
|
||||
Route("/api/admin/roles/{role_id}", admin_update_role, methods=["PUT"]),
|
||||
Route("/api/admin/roles/{role_id}", admin_delete_role, methods=["DELETE"]),
|
||||
# Orgs
|
||||
Route("/api/admin/orgs", admin_list_orgs),
|
||||
Route("/api/admin/orgs/{org_id}", admin_get_org),
|
||||
Route("/api/admin/orgs/{org_id}", admin_update_org, methods=["PUT"]),
|
||||
# Policies
|
||||
Route("/api/admin/policies", admin_list_policies),
|
||||
Route("/api/admin/policies", admin_create_policy, methods=["POST"]),
|
||||
Route(
|
||||
"/api/admin/policies/{policy_id}",
|
||||
admin_update_policy,
|
||||
methods=["PUT"],
|
||||
),
|
||||
Route(
|
||||
"/api/admin/policies/{policy_id}",
|
||||
admin_delete_policy,
|
||||
methods=["DELETE"],
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
middleware=[Middleware(_InjectAuthMiddleware)],
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path: Any) -> SQLiteBackend:
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def sdk_client(storage: SQLiteBackend):
|
||||
"""SDK client wired to a real Starlette app via ASGITransport."""
|
||||
app = _make_app()
|
||||
app.state.auth_storage = storage
|
||||
transport = httpx.ASGITransport(app=app)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as hc:
|
||||
yield AsyncTurnstoneConsole(httpx_client=hc)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests — Roles round-trip
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRolesRoundTrip:
|
||||
@pytest.mark.anyio
|
||||
async def test_list_roles_empty(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
resp = await sdk_client.list_roles()
|
||||
assert isinstance(resp, ListRolesResponse)
|
||||
assert resp.roles == []
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_create_and_list_role(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
role = await sdk_client.create_role(
|
||||
"analyst", display_name="Data Analyst", permissions="read,write"
|
||||
)
|
||||
assert isinstance(role, RoleInfo)
|
||||
assert role.name == "analyst"
|
||||
assert role.display_name == "Data Analyst"
|
||||
assert role.permissions == "read,write"
|
||||
assert role.builtin is False
|
||||
assert role.role_id # non-empty
|
||||
|
||||
# List should now contain the new role
|
||||
resp = await sdk_client.list_roles()
|
||||
assert len(resp.roles) == 1
|
||||
assert resp.roles[0].role_id == role.role_id
|
||||
assert resp.roles[0].name == "analyst"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_create_update_role(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
role = await sdk_client.create_role("ops", permissions="read")
|
||||
assert role.permissions == "read"
|
||||
|
||||
updated = await sdk_client.update_role(
|
||||
role.role_id, display_name="Operations", permissions="read,write,approve"
|
||||
)
|
||||
assert isinstance(updated, RoleInfo)
|
||||
assert updated.display_name == "Operations"
|
||||
assert updated.permissions == "read,write,approve"
|
||||
assert updated.role_id == role.role_id
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_create_delete_role(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
role = await sdk_client.create_role("temp-role", permissions="read")
|
||||
|
||||
result = await sdk_client.delete_role(role.role_id)
|
||||
assert isinstance(result, StatusResponse)
|
||||
assert result.status == "ok"
|
||||
|
||||
# Verify gone
|
||||
resp = await sdk_client.list_roles()
|
||||
assert resp.roles == []
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_full_lifecycle(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
"""Create -> list -> update -> list -> delete -> list."""
|
||||
# Create
|
||||
role = await sdk_client.create_role(
|
||||
"lifecycle", display_name="Lifecycle", permissions="read"
|
||||
)
|
||||
role_id = role.role_id
|
||||
|
||||
# List confirms creation
|
||||
roles = (await sdk_client.list_roles()).roles
|
||||
assert len(roles) == 1
|
||||
assert roles[0].role_id == role_id
|
||||
|
||||
# Update
|
||||
updated = await sdk_client.update_role(role_id, permissions="read,write")
|
||||
assert updated.permissions == "read,write"
|
||||
|
||||
# List still has one
|
||||
roles = (await sdk_client.list_roles()).roles
|
||||
assert len(roles) == 1
|
||||
assert roles[0].permissions == "read,write"
|
||||
|
||||
# Delete
|
||||
await sdk_client.delete_role(role_id)
|
||||
|
||||
# List is empty
|
||||
roles = (await sdk_client.list_roles()).roles
|
||||
assert roles == []
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_delete_nonexistent_role_raises(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
from turnstone.sdk._types import TurnstoneAPIError
|
||||
|
||||
with pytest.raises(TurnstoneAPIError) as exc_info:
|
||||
await sdk_client.delete_role("nonexistent")
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_update_nonexistent_role_raises(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
from turnstone.sdk._types import TurnstoneAPIError
|
||||
|
||||
with pytest.raises(TurnstoneAPIError) as exc_info:
|
||||
await sdk_client.update_role("nonexistent", display_name="Nope")
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests — Policies round-trip
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPoliciesRoundTrip:
|
||||
@pytest.mark.anyio
|
||||
async def test_list_policies_empty(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
resp = await sdk_client.list_policies()
|
||||
assert isinstance(resp, ListToolPoliciesResponse)
|
||||
assert resp.policies == []
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_create_and_list_policy(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
policy = await sdk_client.create_policy("Allow bash", "bash_*", "allow", priority=10)
|
||||
assert isinstance(policy, ToolPolicyInfo)
|
||||
assert policy.name == "Allow bash"
|
||||
assert policy.tool_pattern == "bash_*"
|
||||
assert policy.action == "allow"
|
||||
assert policy.priority == 10
|
||||
assert policy.enabled is True
|
||||
assert policy.policy_id # non-empty
|
||||
|
||||
resp = await sdk_client.list_policies()
|
||||
assert len(resp.policies) == 1
|
||||
assert resp.policies[0].policy_id == policy.policy_id
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_create_update_policy(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
policy = await sdk_client.create_policy("Deny write", "write_*", "deny", priority=5)
|
||||
|
||||
updated = await sdk_client.update_policy(
|
||||
policy.policy_id, name="Allow write", action="allow", priority=20
|
||||
)
|
||||
assert isinstance(updated, ToolPolicyInfo)
|
||||
assert updated.name == "Allow write"
|
||||
assert updated.action == "allow"
|
||||
assert updated.priority == 20
|
||||
assert updated.policy_id == policy.policy_id
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_create_delete_policy(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
policy = await sdk_client.create_policy("Temp policy", "temp_*", "ask")
|
||||
|
||||
result = await sdk_client.delete_policy(policy.policy_id)
|
||||
assert isinstance(result, StatusResponse)
|
||||
assert result.status == "ok"
|
||||
|
||||
resp = await sdk_client.list_policies()
|
||||
assert resp.policies == []
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_full_lifecycle(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
"""Create -> list -> update -> list -> delete -> list."""
|
||||
policy = await sdk_client.create_policy("Lifecycle", "test_*", "deny", priority=1)
|
||||
pid = policy.policy_id
|
||||
|
||||
policies = (await sdk_client.list_policies()).policies
|
||||
assert len(policies) == 1
|
||||
|
||||
await sdk_client.update_policy(pid, action="allow", priority=99)
|
||||
policies = (await sdk_client.list_policies()).policies
|
||||
assert policies[0].action == "allow"
|
||||
assert policies[0].priority == 99
|
||||
|
||||
await sdk_client.delete_policy(pid)
|
||||
policies = (await sdk_client.list_policies()).policies
|
||||
assert policies == []
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_delete_nonexistent_policy_raises(
|
||||
self, sdk_client: AsyncTurnstoneConsole
|
||||
) -> None:
|
||||
from turnstone.sdk._types import TurnstoneAPIError
|
||||
|
||||
with pytest.raises(TurnstoneAPIError) as exc_info:
|
||||
await sdk_client.delete_policy("nonexistent")
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_update_nonexistent_policy_raises(
|
||||
self, sdk_client: AsyncTurnstoneConsole
|
||||
) -> None:
|
||||
from turnstone.sdk._types import TurnstoneAPIError
|
||||
|
||||
with pytest.raises(TurnstoneAPIError) as exc_info:
|
||||
await sdk_client.update_policy("nonexistent", name="Nope")
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_create_policy_invalid_action_raises(
|
||||
self, sdk_client: AsyncTurnstoneConsole
|
||||
) -> None:
|
||||
from turnstone.sdk._types import TurnstoneAPIError
|
||||
|
||||
with pytest.raises(TurnstoneAPIError) as exc_info:
|
||||
await sdk_client.create_policy("Bad", "tool_*", "yolo")
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests — Orgs round-trip
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOrgsRoundTrip:
|
||||
@pytest.mark.anyio
|
||||
async def test_list_orgs_empty(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
resp = await sdk_client.list_orgs()
|
||||
assert isinstance(resp, ListOrgsResponse)
|
||||
assert resp.orgs == []
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_get_org(self, sdk_client: AsyncTurnstoneConsole, storage: SQLiteBackend) -> None:
|
||||
storage.create_org(
|
||||
org_id="org-1", name="acme", display_name="Acme Corp", settings='{"k": "v"}'
|
||||
)
|
||||
org = await sdk_client.get_org("org-1")
|
||||
assert isinstance(org, OrgInfo)
|
||||
assert org.org_id == "org-1"
|
||||
assert org.name == "acme"
|
||||
assert org.display_name == "Acme Corp"
|
||||
assert org.settings == '{"k": "v"}'
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_list_orgs_after_seed(
|
||||
self, sdk_client: AsyncTurnstoneConsole, storage: SQLiteBackend
|
||||
) -> None:
|
||||
storage.create_org(org_id="org-a", name="alpha", display_name="Alpha")
|
||||
storage.create_org(org_id="org-b", name="beta", display_name="Beta")
|
||||
|
||||
resp = await sdk_client.list_orgs()
|
||||
assert len(resp.orgs) == 2
|
||||
names = {o.name for o in resp.orgs}
|
||||
assert names == {"alpha", "beta"}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_update_org(
|
||||
self, sdk_client: AsyncTurnstoneConsole, storage: SQLiteBackend
|
||||
) -> None:
|
||||
storage.create_org(org_id="org-1", name="acme", display_name="Acme Corp")
|
||||
|
||||
updated = await sdk_client.update_org("org-1", display_name="Acme Inc.")
|
||||
assert isinstance(updated, OrgInfo)
|
||||
assert updated.display_name == "Acme Inc."
|
||||
assert updated.org_id == "org-1"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_get_nonexistent_org_raises(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
from turnstone.sdk._types import TurnstoneAPIError
|
||||
|
||||
with pytest.raises(TurnstoneAPIError) as exc_info:
|
||||
await sdk_client.get_org("nonexistent")
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_update_nonexistent_org_raises(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
from turnstone.sdk._types import TurnstoneAPIError
|
||||
|
||||
with pytest.raises(TurnstoneAPIError) as exc_info:
|
||||
await sdk_client.update_org("nonexistent", display_name="Nope")
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests — Pydantic model field validation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestModelValidation:
|
||||
"""Verify that all expected fields are populated and correctly typed."""
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_role_info_fields(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
role = await sdk_client.create_role("reviewer", permissions="read")
|
||||
assert isinstance(role.role_id, str)
|
||||
assert isinstance(role.name, str)
|
||||
assert isinstance(role.display_name, str)
|
||||
assert isinstance(role.permissions, str)
|
||||
assert isinstance(role.builtin, bool)
|
||||
assert isinstance(role.org_id, str)
|
||||
assert isinstance(role.created, str)
|
||||
assert isinstance(role.updated, str)
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_policy_info_fields(self, sdk_client: AsyncTurnstoneConsole) -> None:
|
||||
policy = await sdk_client.create_policy("Test", "read_*", "allow", priority=5)
|
||||
assert isinstance(policy.policy_id, str)
|
||||
assert isinstance(policy.name, str)
|
||||
assert isinstance(policy.tool_pattern, str)
|
||||
assert isinstance(policy.action, str)
|
||||
assert isinstance(policy.priority, int)
|
||||
assert isinstance(policy.org_id, str)
|
||||
assert isinstance(policy.enabled, bool)
|
||||
assert isinstance(policy.created_by, str)
|
||||
assert isinstance(policy.created, str)
|
||||
assert isinstance(policy.updated, str)
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_org_info_fields(
|
||||
self, sdk_client: AsyncTurnstoneConsole, storage: SQLiteBackend
|
||||
) -> None:
|
||||
storage.create_org(org_id="org-v", name="validate", display_name="Validate")
|
||||
org = await sdk_client.get_org("org-v")
|
||||
assert isinstance(org.org_id, str)
|
||||
assert isinstance(org.name, str)
|
||||
assert isinstance(org.display_name, str)
|
||||
assert isinstance(org.settings, str)
|
||||
assert isinstance(org.created, str)
|
||||
assert isinstance(org.updated, str)
|
||||
@@ -623,6 +623,7 @@ class TestServerHealthMetrics:
|
||||
|
||||
mock_mgr = MagicMock()
|
||||
mock_mgr.list_all.return_value = [mock_ws]
|
||||
mock_mgr.max_workstreams = 10
|
||||
|
||||
app = srv_mod.create_app(
|
||||
workstreams=mock_mgr,
|
||||
@@ -799,6 +800,7 @@ class TestServerRateLimiting:
|
||||
|
||||
mock_mgr = MagicMock()
|
||||
mock_mgr.list_all.return_value = [mock_ws]
|
||||
mock_mgr.max_workstreams = 10
|
||||
|
||||
app = srv_mod.create_app(
|
||||
workstreams=mock_mgr,
|
||||
|
||||
@@ -2,15 +2,6 @@
|
||||
|
||||
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):
|
||||
|
||||
@@ -753,3 +753,268 @@ class TestGetCapabilitiesOverride:
|
||||
caps = session._get_capabilities()
|
||||
# Default OpenAI provider for unknown model → no vision
|
||||
assert caps.supports_vision is False
|
||||
|
||||
|
||||
class TestTitleRetry:
|
||||
"""_generate_title resets _title_generated on failure."""
|
||||
|
||||
def test_title_generated_reset_on_failure(self, tmp_db):
|
||||
session = _make_session()
|
||||
session._title_generated = True
|
||||
session.messages = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "Hi there"},
|
||||
]
|
||||
# Mock provider to raise
|
||||
session._provider = MagicMock()
|
||||
session._provider.create_completion.side_effect = RuntimeError("API error")
|
||||
|
||||
session._generate_title()
|
||||
|
||||
assert session._title_generated is False
|
||||
|
||||
def test_title_generated_stays_true_on_success(self, tmp_db):
|
||||
session = _make_session()
|
||||
session._title_generated = True
|
||||
session.messages = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "Hi there"},
|
||||
]
|
||||
result = MagicMock()
|
||||
result.content = "Test Title"
|
||||
session._provider = MagicMock()
|
||||
session._provider.create_completion.return_value = result
|
||||
|
||||
with patch("turnstone.core.session.update_workstream_title"):
|
||||
session._generate_title()
|
||||
|
||||
# Flag stays True after successful generation
|
||||
assert session._title_generated is True
|
||||
|
||||
def test_title_skipped_after_resume_changes_ws_id(self, tmp_db):
|
||||
"""If ws_id changes (via resume) during title generation, discard the result."""
|
||||
session = _make_session()
|
||||
session._title_generated = True
|
||||
session.messages = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "Hi there"},
|
||||
]
|
||||
original_ws_id = session._ws_id
|
||||
result = MagicMock()
|
||||
result.content = "Test Title"
|
||||
session._provider = MagicMock()
|
||||
session._provider.create_completion.return_value = result
|
||||
|
||||
# Simulate resume() changing ws_id while title generation is in flight
|
||||
def _change_ws_id(*args, **kwargs):
|
||||
session._ws_id = "different-ws-id"
|
||||
return result
|
||||
|
||||
session._provider.create_completion.side_effect = _change_ws_id
|
||||
|
||||
with patch("turnstone.core.session.update_workstream_title") as mock_update:
|
||||
session._generate_title()
|
||||
|
||||
# Title should NOT be applied to the new workstream
|
||||
mock_update.assert_not_called()
|
||||
# Restore for cleanup
|
||||
session._ws_id = original_ws_id
|
||||
|
||||
|
||||
class TestLiveConfigUpdate:
|
||||
"""ConfigStore-backed sessions pick up settings changes at point-of-use."""
|
||||
|
||||
def test_memory_config_reads_from_config_store(self, tmp_db):
|
||||
"""_mem_cfg returns live values from ConfigStore when present."""
|
||||
from turnstone.core.config_store import ConfigStore
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_db), create_tables=True)
|
||||
cs = ConfigStore(storage)
|
||||
session = _make_session(config_store=cs)
|
||||
|
||||
# Default: relevance_k=5
|
||||
assert session._mem_cfg.relevance_k == 5
|
||||
|
||||
# Admin changes the setting
|
||||
cs.set("memory.relevance_k", 10, changed_by="test")
|
||||
assert session._mem_cfg.relevance_k == 10
|
||||
|
||||
def test_judge_config_reads_from_config_store(self, tmp_db):
|
||||
"""_judge_cfg returns live behavioral flags from ConfigStore."""
|
||||
from turnstone.core.config_store import ConfigStore
|
||||
from turnstone.core.judge import JudgeConfig
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_db), create_tables=True)
|
||||
cs = ConfigStore(storage)
|
||||
session = _make_session(
|
||||
judge_config=JudgeConfig(),
|
||||
config_store=cs,
|
||||
)
|
||||
|
||||
# Default: enabled=True
|
||||
assert session._judge_cfg.enabled is True
|
||||
|
||||
# Admin disables the judge
|
||||
cs.set("judge.enabled", False, changed_by="test")
|
||||
assert session._judge_cfg.enabled is False
|
||||
|
||||
def test_judge_client_config_stays_frozen(self, tmp_db):
|
||||
"""LLM client fields (model, provider) are frozen from creation time."""
|
||||
from turnstone.core.config_store import ConfigStore
|
||||
from turnstone.core.judge import JudgeConfig
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_db), create_tables=True)
|
||||
cs = ConfigStore(storage)
|
||||
session = _make_session(
|
||||
judge_config=JudgeConfig(model="original-model"),
|
||||
config_store=cs,
|
||||
)
|
||||
|
||||
# Change the model in ConfigStore — should NOT affect the session
|
||||
cs.set("judge.model", "new-model", changed_by="test")
|
||||
assert session._judge_cfg.model == "original-model"
|
||||
|
||||
def test_judge_disable_after_init_stops_future_use(self, tmp_db):
|
||||
"""Disabling judge.enabled after IntentJudge is created returns None."""
|
||||
from turnstone.core.config_store import ConfigStore
|
||||
from turnstone.core.judge import JudgeConfig
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
storage = SQLiteBackend(str(tmp_db), create_tables=True)
|
||||
cs = ConfigStore(storage)
|
||||
session = _make_session(
|
||||
judge_config=JudgeConfig(),
|
||||
config_store=cs,
|
||||
)
|
||||
|
||||
# Force judge initialization by setting a mock
|
||||
session._judge = MagicMock()
|
||||
assert session._ensure_judge() is not None
|
||||
|
||||
# Admin disables the judge — cached instance should NOT be returned
|
||||
cs.set("judge.enabled", False, changed_by="test")
|
||||
assert session._ensure_judge() is None
|
||||
|
||||
def test_fallback_to_frozen_without_config_store(self, tmp_db):
|
||||
"""Without ConfigStore (CLI mode), frozen config is used."""
|
||||
from turnstone.core.memory_relevance import MemoryConfig
|
||||
|
||||
session = _make_session(memory_config=MemoryConfig(relevance_k=3))
|
||||
assert session._mem_cfg.relevance_k == 3
|
||||
|
||||
|
||||
class TestAgentOutputGuard:
|
||||
"""Output guard should evaluate tool results in _run_agent, not just the main loop."""
|
||||
|
||||
def test_agent_loop_calls_evaluate_output(self):
|
||||
"""_run_agent passes tool output through _evaluate_output when output_guard is enabled."""
|
||||
from turnstone.core.judge import JudgeConfig
|
||||
|
||||
session = _make_session(judge_config=JudgeConfig(output_guard=True))
|
||||
|
||||
with patch.object(session, "_evaluate_output", wraps=lambda cid, o, fn: o) as mock_eval:
|
||||
# Simulate _run_agent getting a tool call response then a text response
|
||||
call_count = [0]
|
||||
|
||||
def fake_create(**kwargs):
|
||||
call_count[0] += 1
|
||||
resp = MagicMock()
|
||||
if call_count[0] == 1:
|
||||
# First call: model returns a tool call
|
||||
choice = MagicMock()
|
||||
choice.finish_reason = "tool_calls"
|
||||
tc = MagicMock()
|
||||
tc.id = "call_1"
|
||||
tc.function.name = "read_file"
|
||||
tc.function.arguments = '{"path": "/tmp/test"}'
|
||||
choice.message.tool_calls = [tc]
|
||||
choice.message.content = None
|
||||
resp.choices = [choice]
|
||||
resp.usage = MagicMock(prompt_tokens=10, completion_tokens=5)
|
||||
else:
|
||||
# Second call: model returns text (done)
|
||||
choice = MagicMock()
|
||||
choice.finish_reason = "stop"
|
||||
choice.message.tool_calls = None
|
||||
choice.message.content = "Done"
|
||||
resp.choices = [choice]
|
||||
resp.usage = MagicMock(prompt_tokens=10, completion_tokens=5)
|
||||
return resp
|
||||
|
||||
session.client.chat.completions.create = fake_create
|
||||
|
||||
# Mock tool preparation to return a simple output
|
||||
def fake_prepare(tc_dict, **kwargs):
|
||||
return {
|
||||
"call_id": tc_dict["id"],
|
||||
"func_name": "read_file",
|
||||
"needs_approval": False,
|
||||
"execute": lambda p: ("call_1", "file contents with sk-proj-SECRET123"),
|
||||
}
|
||||
|
||||
with patch.object(session, "_prepare_tool", side_effect=fake_prepare):
|
||||
session._run_agent(
|
||||
[{"role": "user", "content": "test"}],
|
||||
tools=[{"type": "function", "function": {"name": "read_file"}}],
|
||||
label="test",
|
||||
)
|
||||
|
||||
mock_eval.assert_called_once()
|
||||
args = mock_eval.call_args[0]
|
||||
assert args[0] == "call_1" # call_id
|
||||
assert "sk-proj-SECRET123" in args[1] # output
|
||||
assert args[2] == "read_file" # func_name
|
||||
|
||||
def test_agent_loop_skips_guard_when_disabled(self):
|
||||
"""_run_agent does not call _evaluate_output when output_guard is disabled."""
|
||||
from turnstone.core.judge import JudgeConfig
|
||||
|
||||
session = _make_session(judge_config=JudgeConfig(output_guard=False))
|
||||
|
||||
with patch.object(session, "_evaluate_output") as mock_eval:
|
||||
call_count = [0]
|
||||
|
||||
def fake_create(**kwargs):
|
||||
call_count[0] += 1
|
||||
resp = MagicMock()
|
||||
if call_count[0] == 1:
|
||||
choice = MagicMock()
|
||||
choice.finish_reason = "tool_calls"
|
||||
tc = MagicMock()
|
||||
tc.id = "call_1"
|
||||
tc.function.name = "read_file"
|
||||
tc.function.arguments = '{"path": "/tmp/test"}'
|
||||
choice.message.tool_calls = [tc]
|
||||
choice.message.content = None
|
||||
resp.choices = [choice]
|
||||
resp.usage = MagicMock(prompt_tokens=10, completion_tokens=5)
|
||||
else:
|
||||
choice = MagicMock()
|
||||
choice.finish_reason = "stop"
|
||||
choice.message.tool_calls = None
|
||||
choice.message.content = "Done"
|
||||
resp.choices = [choice]
|
||||
resp.usage = MagicMock(prompt_tokens=10, completion_tokens=5)
|
||||
return resp
|
||||
|
||||
session.client.chat.completions.create = fake_create
|
||||
|
||||
def fake_prepare(tc_dict, **kwargs):
|
||||
return {
|
||||
"call_id": tc_dict["id"],
|
||||
"func_name": "read_file",
|
||||
"needs_approval": False,
|
||||
"execute": lambda p: ("call_1", "safe output"),
|
||||
}
|
||||
|
||||
with patch.object(session, "_prepare_tool", side_effect=fake_prepare):
|
||||
session._run_agent(
|
||||
[{"role": "user", "content": "test"}],
|
||||
tools=[{"type": "function", "function": {"name": "read_file"}}],
|
||||
label="test",
|
||||
)
|
||||
|
||||
mock_eval.assert_not_called()
|
||||
|
||||
@@ -179,7 +179,10 @@ class TestDeleteSetting:
|
||||
# Delete it
|
||||
r = client.delete("/v1/api/admin/settings/tools.timeout")
|
||||
assert r.status_code == 200
|
||||
assert r.json()["status"] == "ok"
|
||||
body = r.json()
|
||||
assert body["status"] == "ok"
|
||||
assert body["key"] == "tools.timeout"
|
||||
assert body["default"] == 120 # registry default for tools.timeout
|
||||
|
||||
def test_delete_then_list_shows_default(self, client):
|
||||
client.put(
|
||||
|
||||
@@ -4,16 +4,6 @@ from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
|
||||
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"))
|
||||
|
||||
|
||||
class TestDeleteSkillResourceByPath:
|
||||
def test_delete_existing(self, storage):
|
||||
|
||||
+366
-1
@@ -132,6 +132,7 @@ def _create_template(db, template_id, name, content, **kwargs):
|
||||
notify_on_complete=kwargs.get("notify_on_complete", "{}"),
|
||||
enabled=kwargs.get("enabled", True),
|
||||
allowed_tools=kwargs.get("allowed_tools", "[]"),
|
||||
priority=kwargs.get("priority", 0),
|
||||
)
|
||||
|
||||
|
||||
@@ -293,7 +294,7 @@ class TestSkillStorage:
|
||||
assert result == []
|
||||
|
||||
def test_list_skills_by_activation_ordered_by_name(self, db):
|
||||
"""Results are ordered by name ascending."""
|
||||
"""Results are ordered by name ascending when priority is equal."""
|
||||
_create_template(db, "s2", "beta-search", "B", activation="search")
|
||||
_create_template(db, "s1", "alpha-search", "A", activation="search")
|
||||
results = db.list_skills_by_activation("search")
|
||||
@@ -301,6 +302,51 @@ class TestSkillStorage:
|
||||
assert results[0]["name"] == "alpha-search"
|
||||
assert results[1]["name"] == "beta-search"
|
||||
|
||||
def test_list_skills_by_activation_ordered_by_priority(self, db):
|
||||
"""Results are ordered by priority ascending, then name."""
|
||||
_create_template(db, "s1", "style", "S", activation="default", priority=20)
|
||||
_create_template(db, "s2", "safety", "F", activation="default", priority=10)
|
||||
_create_template(db, "s3", "tone", "T", activation="default", priority=10)
|
||||
results = db.list_skills_by_activation("default")
|
||||
assert len(results) == 3
|
||||
assert results[0]["name"] == "safety"
|
||||
assert results[1]["name"] == "tone"
|
||||
assert results[2]["name"] == "style"
|
||||
|
||||
def test_priority_default_is_zero(self, db):
|
||||
"""Priority defaults to 0 when not specified."""
|
||||
_create_template(db, "s1", "skill", "content")
|
||||
tpl = db.get_prompt_template("s1")
|
||||
assert tpl is not None
|
||||
assert tpl["priority"] == 0
|
||||
|
||||
def test_priority_roundtrip(self, db):
|
||||
"""Priority can be set on create and retrieved."""
|
||||
_create_template(db, "s1", "skill", "content", priority=42)
|
||||
tpl = db.get_prompt_template("s1")
|
||||
assert tpl is not None
|
||||
assert tpl["priority"] == 42
|
||||
|
||||
def test_priority_update(self, db):
|
||||
"""Priority can be updated."""
|
||||
_create_template(db, "s1", "skill", "content", priority=10)
|
||||
db.update_prompt_template("s1", priority=99)
|
||||
tpl = db.get_prompt_template("s1")
|
||||
assert tpl is not None
|
||||
assert tpl["priority"] == 99
|
||||
|
||||
def test_list_default_templates_ordered_by_priority(self, db):
|
||||
"""list_default_templates() respects priority ordering."""
|
||||
_create_template(db, "s1", "beta", "b", activation="default", priority=10)
|
||||
_create_template(db, "s2", "alpha", "a", activation="default", priority=5)
|
||||
_create_template(db, "s3", "gamma", "g", activation="default", priority=1)
|
||||
|
||||
results = db.list_default_templates()
|
||||
assert len(results) == 3
|
||||
assert results[0]["name"] == "gamma"
|
||||
assert results[1]["name"] == "alpha"
|
||||
assert results[2]["name"] == "beta"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1b. Skill resource storage tests
|
||||
@@ -1555,3 +1601,322 @@ class TestSkillAdminEndpoints:
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
assert "integer" in resp.json()["error"].lower()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 9. Skill session config applied to workstream via server handler
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSkillConfigAppliedToWorkstream:
|
||||
"""Verify that skill session config fields are applied to the ChatSession
|
||||
when a workstream is created via the server ``create_workstream`` handler.
|
||||
"""
|
||||
|
||||
@pytest.fixture()
|
||||
def _ws_app(self, tmp_path):
|
||||
"""Build a minimal Starlette app with the real ``create_workstream``
|
||||
handler, a real ``WorkstreamManager``, and a temp SQLite storage
|
||||
backend. Returns ``(TestClient, WorkstreamManager, storage)``.
|
||||
"""
|
||||
import queue
|
||||
import threading
|
||||
|
||||
import turnstone.core.storage._registry as _reg
|
||||
from turnstone.core.workstream import WorkstreamManager
|
||||
from turnstone.server import create_workstream
|
||||
|
||||
storage = SQLiteBackend(str(tmp_path / "ws_test.db"))
|
||||
|
||||
# Inject the test storage as the global singleton so that
|
||||
# get_storage() / get_skill_by_name() resolve against it.
|
||||
old_storage = _reg._storage
|
||||
_reg._storage = storage
|
||||
|
||||
def _session_factory(
|
||||
ui: Any, model_alias: Any = None, ws_id: Any = None, **kwargs: Any
|
||||
) -> ChatSession:
|
||||
return ChatSession(
|
||||
client=MagicMock(),
|
||||
model=model_alias or "test-model",
|
||||
ui=ui,
|
||||
instructions=None,
|
||||
temperature=0.5,
|
||||
max_tokens=4096,
|
||||
tool_timeout=30,
|
||||
ws_id=ws_id,
|
||||
skill=kwargs.get("skill"),
|
||||
)
|
||||
|
||||
mgr = WorkstreamManager(_session_factory)
|
||||
|
||||
routes = [
|
||||
Mount(
|
||||
"/v1",
|
||||
routes=[
|
||||
Route(
|
||||
"/api/workstreams/new",
|
||||
create_workstream,
|
||||
methods=["POST"],
|
||||
),
|
||||
],
|
||||
),
|
||||
]
|
||||
app = Starlette(
|
||||
routes=routes,
|
||||
middleware=[Middleware(_InjectAuthMiddleware)],
|
||||
)
|
||||
app.state.workstreams = mgr
|
||||
app.state.skip_permissions = True
|
||||
app.state.global_queue = queue.Queue()
|
||||
app.state.global_listeners = []
|
||||
app.state.global_listeners_lock = threading.Lock()
|
||||
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
|
||||
yield client, mgr, storage
|
||||
|
||||
# Restore original storage singleton.
|
||||
_reg._storage = old_storage
|
||||
|
||||
def test_session_receives_temperature(self, _ws_app):
|
||||
"""Skill temperature overrides the session default."""
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(storage, "s1", "warm-skill", "Be warm.", temperature=0.9, enabled=True)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "warm-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws_id = resp.json()["ws_id"]
|
||||
ws = mgr.get(ws_id)
|
||||
assert ws is not None and ws.session is not None
|
||||
assert ws.session.temperature == 0.9
|
||||
|
||||
def test_session_receives_max_tokens(self, _ws_app):
|
||||
"""Skill max_tokens overrides the session default."""
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(storage, "s1", "token-skill", "Be concise.", max_tokens=1024, enabled=True)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "token-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws = mgr.get(resp.json()["ws_id"])
|
||||
assert ws is not None and ws.session is not None
|
||||
assert ws.session.max_tokens == 1024
|
||||
|
||||
def test_session_receives_token_budget(self, _ws_app):
|
||||
"""Skill token_budget is applied to the session."""
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(
|
||||
storage, "s1", "budget-skill", "Stay on budget.", token_budget=50000, enabled=True
|
||||
)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "budget-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws = mgr.get(resp.json()["ws_id"])
|
||||
assert ws is not None and ws.session is not None
|
||||
assert ws.session._token_budget == 50000
|
||||
|
||||
def test_session_receives_reasoning_effort(self, _ws_app):
|
||||
"""Skill reasoning_effort is applied to the session."""
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(
|
||||
storage, "s1", "effort-skill", "Think hard.", reasoning_effort="high", enabled=True
|
||||
)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "effort-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws = mgr.get(resp.json()["ws_id"])
|
||||
assert ws is not None and ws.session is not None
|
||||
assert ws.session.reasoning_effort == "high"
|
||||
|
||||
def test_session_receives_agent_max_turns(self, _ws_app):
|
||||
"""Skill agent_max_turns is applied to the session."""
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(
|
||||
storage, "s1", "turns-skill", "Few turns.", agent_max_turns=3, enabled=True
|
||||
)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "turns-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws = mgr.get(resp.json()["ws_id"])
|
||||
assert ws is not None and ws.session is not None
|
||||
assert ws.session.agent_max_turns == 3
|
||||
|
||||
def test_auto_approve_set_on_ui(self, _ws_app):
|
||||
"""Skill auto_approve=True propagates to the WebUI."""
|
||||
from turnstone.server import WebUI
|
||||
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(
|
||||
storage, "s1", "approve-skill", "Auto approve.", auto_approve=True, enabled=True
|
||||
)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "approve-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws = mgr.get(resp.json()["ws_id"])
|
||||
assert ws is not None
|
||||
assert isinstance(ws.ui, WebUI)
|
||||
assert ws.ui.auto_approve is True
|
||||
|
||||
def test_allowed_tools_set_on_ui(self, _ws_app):
|
||||
"""Skill allowed_tools are parsed and set as auto_approve_tools on the UI."""
|
||||
from turnstone.server import WebUI
|
||||
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(
|
||||
storage,
|
||||
"s1",
|
||||
"tools-skill",
|
||||
"Restricted tools.",
|
||||
allowed_tools='["bash", "read_file"]',
|
||||
enabled=True,
|
||||
)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "tools-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws = mgr.get(resp.json()["ws_id"])
|
||||
assert ws is not None
|
||||
assert isinstance(ws.ui, WebUI)
|
||||
assert ws.ui.auto_approve_tools == {"bash", "read_file"}
|
||||
|
||||
def test_skill_model_overrides_resolved_model(self, _ws_app):
|
||||
"""Skill model field overrides the default session model."""
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(
|
||||
storage, "s1", "model-skill", "Use specific model.", model="gpt-5", enabled=True
|
||||
)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "model-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws = mgr.get(resp.json()["ws_id"])
|
||||
assert ws is not None and ws.session is not None
|
||||
assert ws.session.model == "gpt-5"
|
||||
|
||||
def test_all_session_config_fields_applied(self, _ws_app):
|
||||
"""All session config fields from a skill are applied together."""
|
||||
from turnstone.server import WebUI
|
||||
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(
|
||||
storage,
|
||||
"s1",
|
||||
"full-skill",
|
||||
"Full config skill.",
|
||||
model="gpt-5",
|
||||
temperature=0.8,
|
||||
reasoning_effort="high",
|
||||
max_tokens=2048,
|
||||
token_budget=100000,
|
||||
agent_max_turns=10,
|
||||
auto_approve=True,
|
||||
allowed_tools='["bash", "write_file", "read_file"]',
|
||||
notify_on_complete='{"channel": "discord"}',
|
||||
enabled=True,
|
||||
)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "full-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws = mgr.get(resp.json()["ws_id"])
|
||||
assert ws is not None and ws.session is not None
|
||||
sess = ws.session
|
||||
assert sess.model == "gpt-5"
|
||||
assert sess.temperature == 0.8
|
||||
assert sess.reasoning_effort == "high"
|
||||
assert sess.max_tokens == 2048
|
||||
assert sess._token_budget == 100000
|
||||
assert sess.agent_max_turns == 10
|
||||
assert sess._notify_on_complete == '{"channel": "discord"}'
|
||||
assert sess._applied_skill_id == "s1"
|
||||
assert sess._applied_skill_content == "Full config skill."
|
||||
assert isinstance(ws.ui, WebUI)
|
||||
assert ws.ui.auto_approve is True
|
||||
assert ws.ui.auto_approve_tools == {"bash", "write_file", "read_file"}
|
||||
|
||||
def test_disabled_skill_returns_400(self, _ws_app):
|
||||
"""Creating a workstream with a disabled skill returns 400."""
|
||||
client, _mgr, storage = _ws_app
|
||||
_create_template(storage, "s1", "disabled-skill", "Disabled.", enabled=False)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "disabled-skill"})
|
||||
assert resp.status_code == 400
|
||||
assert "disabled" in resp.json()["error"].lower()
|
||||
|
||||
def test_unknown_skill_returns_400(self, _ws_app):
|
||||
"""Creating a workstream with a nonexistent skill returns 400."""
|
||||
client, _mgr, _storage = _ws_app
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "no-such-skill"})
|
||||
assert resp.status_code == 400
|
||||
assert "not found" in resp.json()["error"].lower()
|
||||
|
||||
def test_zero_token_budget_is_noop(self, _ws_app):
|
||||
"""Skill with token_budget=0 — handler skips budget application (> 0 guard)."""
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(
|
||||
storage, "s1", "no-budget-skill", "No budget.", token_budget=0, enabled=True
|
||||
)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "no-budget-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws = mgr.get(resp.json()["ws_id"])
|
||||
assert ws is not None and ws.session is not None
|
||||
# Budget stays at default (0) — the handler's > 0 guard prevents application
|
||||
assert ws.session._token_budget == 0
|
||||
|
||||
def test_empty_allowed_tools_is_noop(self, _ws_app):
|
||||
"""Skill with allowed_tools='[]' — handler skips (empty check)."""
|
||||
client, mgr, storage = _ws_app
|
||||
_create_template(
|
||||
storage, "s1", "no-tools-skill", "No tools.", allowed_tools="[]", enabled=True
|
||||
)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "no-tools-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws = mgr.get(resp.json()["ws_id"])
|
||||
assert ws is not None
|
||||
# auto_approve_tools stays at default (empty set)
|
||||
assert ws.ui.auto_approve_tools == set()
|
||||
|
||||
def test_skill_lineage_in_workstreams_table(self, _ws_app):
|
||||
"""skill_id and skill_version columns are populated in the workstreams table."""
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.storage._schema import workstreams
|
||||
|
||||
client, _mgr, storage = _ws_app
|
||||
_create_template(storage, "s1", "lineage-skill", "Track me.", enabled=True)
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={"skill": "lineage-skill"})
|
||||
assert resp.status_code == 200
|
||||
ws_id = resp.json()["ws_id"]
|
||||
|
||||
with storage._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
sa.select(workstreams.c.skill_id, workstreams.c.skill_version).where(
|
||||
workstreams.c.ws_id == ws_id
|
||||
)
|
||||
).fetchone()
|
||||
assert row is not None
|
||||
assert row[0] == "s1"
|
||||
assert row[1] == 1
|
||||
|
||||
def test_no_skill_lineage_when_no_skill(self, _ws_app):
|
||||
"""Workstream without a skill has empty skill_id and zero skill_version."""
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.storage._schema import workstreams
|
||||
|
||||
client, _mgr, storage = _ws_app
|
||||
|
||||
resp = client.post("/v1/api/workstreams/new", json={})
|
||||
assert resp.status_code == 200
|
||||
ws_id = resp.json()["ws_id"]
|
||||
|
||||
with storage._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
sa.select(workstreams.c.skill_id, workstreams.c.skill_version).where(
|
||||
workstreams.c.ws_id == ws_id
|
||||
)
|
||||
).fetchone()
|
||||
assert row is not None
|
||||
assert row[0] == ""
|
||||
assert row[1] == 0
|
||||
|
||||
@@ -1,18 +1,8 @@
|
||||
"""Tests for the SQLite storage backend."""
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.storage import init_storage, reset_storage
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def backend(tmp_path):
|
||||
"""Create a fresh SQLiteBackend for each test."""
|
||||
reset_storage()
|
||||
b = init_storage("sqlite", path=str(tmp_path / "test.db"), run_migrations=False)
|
||||
yield b
|
||||
reset_storage()
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
# -- Workstream registration ---------------------------------------------------
|
||||
|
||||
@@ -289,6 +279,73 @@ class TestWorkstreams:
|
||||
assert rows[0][6] == "node-a"
|
||||
|
||||
|
||||
# -- Structured memory touch ---------------------------------------------------
|
||||
|
||||
|
||||
class TestTouchStructuredMemory:
|
||||
@staticmethod
|
||||
def _create_memory(
|
||||
backend: Any, name: str = "m1", scope: str = "global", scope_id: str = ""
|
||||
) -> None:
|
||||
import uuid
|
||||
|
||||
backend.create_structured_memory(
|
||||
memory_id=str(uuid.uuid4()),
|
||||
name=name,
|
||||
description="test desc",
|
||||
mem_type="project",
|
||||
scope=scope,
|
||||
scope_id=scope_id,
|
||||
content="test content",
|
||||
)
|
||||
|
||||
def test_batch_touch_multiple(self, backend):
|
||||
self._create_memory(backend, name="a")
|
||||
self._create_memory(backend, name="b")
|
||||
self._create_memory(backend, name="c")
|
||||
|
||||
count = backend.touch_structured_memories(
|
||||
[
|
||||
("a", "global", ""),
|
||||
("b", "global", ""),
|
||||
("c", "global", ""),
|
||||
]
|
||||
)
|
||||
assert count == 3
|
||||
|
||||
for name in ("a", "b", "c"):
|
||||
mem = backend.get_structured_memory_by_name(name, "global", "")
|
||||
assert int(mem["access_count"]) == 1
|
||||
|
||||
def test_batch_touch_empty_list(self, backend):
|
||||
assert backend.touch_structured_memories([]) == 0
|
||||
|
||||
def test_batch_touch_partial_match(self, backend):
|
||||
self._create_memory(backend, name="exists")
|
||||
|
||||
count = backend.touch_structured_memories(
|
||||
[
|
||||
("exists", "global", ""),
|
||||
("missing", "global", ""),
|
||||
]
|
||||
)
|
||||
assert count == 1
|
||||
|
||||
mem = backend.get_structured_memory_by_name("exists", "global", "")
|
||||
assert int(mem["access_count"]) == 1
|
||||
|
||||
def test_batch_touch_with_duplicates(self, backend):
|
||||
"""Duplicate keys in batch should each increment access_count once."""
|
||||
self._create_memory(backend, name="dup")
|
||||
|
||||
# Two identical keys — storage gets called twice for the same row
|
||||
count = backend.touch_structured_memories([("dup", "global", ""), ("dup", "global", "")])
|
||||
assert count == 2
|
||||
|
||||
mem = backend.get_structured_memory_by_name("dup", "global", "")
|
||||
assert int(mem["access_count"]) == 2
|
||||
|
||||
|
||||
# -- Lifecycle -----------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
@@ -1,14 +1,5 @@
|
||||
"""Tests for structured memory storage backend operations."""
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def backend(tmp_path):
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
class TestCreateAndGet:
|
||||
def test_create_and_get_by_id(self, backend):
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
"""Tests for tool policy enforcement across CLI, bridge, and channel entry points."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from turnstone.cli import TerminalUI
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# CLI
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCLIPolicyEnforcement:
|
||||
"""Tool policies should be enforced in CLI approve_tools()."""
|
||||
|
||||
def _make_items(self, *tool_names: str) -> list[dict]:
|
||||
return [
|
||||
{
|
||||
"call_id": f"call_{i}",
|
||||
"header": f"Tool: {name}",
|
||||
"preview": "",
|
||||
"func_name": name,
|
||||
"approval_label": name,
|
||||
"needs_approval": True,
|
||||
}
|
||||
for i, name in enumerate(tool_names)
|
||||
]
|
||||
|
||||
def test_deny_policy_blocks_tool(self):
|
||||
"""A 'deny' policy verdict should block the tool without prompting."""
|
||||
ui = TerminalUI()
|
||||
items = self._make_items("bash")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.policy.evaluate_tool_policies_batch",
|
||||
return_value={"bash": "deny"},
|
||||
),
|
||||
patch(
|
||||
"turnstone.core.storage._registry.get_storage",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
):
|
||||
approved, _ = ui.approve_tools(items)
|
||||
|
||||
assert items[0].get("denied") is True
|
||||
assert items[0].get("error")
|
||||
assert "policy" in items[0]["error"].lower()
|
||||
|
||||
def test_allow_policy_auto_approves(self):
|
||||
"""An 'allow' policy verdict should auto-approve without prompting."""
|
||||
ui = TerminalUI()
|
||||
items = self._make_items("read_file")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.policy.evaluate_tool_policies_batch",
|
||||
return_value={"read_file": "allow"},
|
||||
),
|
||||
patch(
|
||||
"turnstone.core.storage._registry.get_storage",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
):
|
||||
approved, _ = ui.approve_tools(items)
|
||||
|
||||
assert approved is True
|
||||
|
||||
def test_no_storage_skips_policies(self):
|
||||
"""When storage is unavailable, policies are skipped (best-effort)."""
|
||||
ui = TerminalUI()
|
||||
items = self._make_items("bash")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.storage._registry.get_storage",
|
||||
return_value=None,
|
||||
),
|
||||
patch("builtins.input", return_value="y"),
|
||||
):
|
||||
approved, _ = ui.approve_tools(items)
|
||||
|
||||
# Should fall through to normal prompt (which we answered 'y')
|
||||
assert approved is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Bridge
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBridgePolicyEnforcement:
|
||||
"""Tool policies should be enforced in bridge _handle_approval()."""
|
||||
|
||||
def _make_bridge(self):
|
||||
from turnstone.mq.bridge import Bridge
|
||||
|
||||
broker = MagicMock()
|
||||
return Bridge(
|
||||
server_url="http://localhost:8080",
|
||||
broker=broker,
|
||||
node_id="test-node",
|
||||
approval_timeout=1,
|
||||
)
|
||||
|
||||
def _approval_items(self, *tool_names: str) -> list[dict]:
|
||||
return [
|
||||
{"func_name": name, "needs_approval": True, "approval_label": name}
|
||||
for name in tool_names
|
||||
]
|
||||
|
||||
def test_deny_policy_rejects_approval(self):
|
||||
"""A 'deny' policy should reject the approval."""
|
||||
bridge = self._make_bridge()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.policy.evaluate_tool_policies_batch",
|
||||
return_value={"bash": "deny"},
|
||||
),
|
||||
patch(
|
||||
"turnstone.core.storage._registry._storage",
|
||||
new=MagicMock(),
|
||||
),
|
||||
patch.object(bridge, "_api_approve") as mock_approve,
|
||||
patch.object(bridge, "_publish_ws"),
|
||||
):
|
||||
bridge._handle_approval("ws-1", {"items": self._approval_items("bash")})
|
||||
|
||||
mock_approve.assert_called_once()
|
||||
assert mock_approve.call_args.kwargs.get("approved") is False
|
||||
|
||||
def test_allow_policy_approves(self):
|
||||
"""An 'allow' policy should auto-approve."""
|
||||
bridge = self._make_bridge()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.policy.evaluate_tool_policies_batch",
|
||||
return_value={"read_file": "allow"},
|
||||
),
|
||||
patch(
|
||||
"turnstone.core.storage._registry.get_storage",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch.object(bridge, "_api_approve") as mock_approve,
|
||||
patch.object(bridge, "_publish_ws"),
|
||||
):
|
||||
bridge._handle_approval("ws-1", {"items": self._approval_items("read_file")})
|
||||
|
||||
mock_approve.assert_called_once()
|
||||
assert mock_approve.call_args.kwargs.get("approved") is True
|
||||
|
||||
def test_mixed_deny_rejects_batch(self):
|
||||
"""If any tool is denied, the whole batch is rejected."""
|
||||
bridge = self._make_bridge()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.policy.evaluate_tool_policies_batch",
|
||||
return_value={"bash": "deny", "read_file": "allow"},
|
||||
),
|
||||
patch(
|
||||
"turnstone.core.storage._registry._storage",
|
||||
new=MagicMock(),
|
||||
),
|
||||
patch.object(bridge, "_api_approve") as mock_approve,
|
||||
patch.object(bridge, "_publish_ws"),
|
||||
):
|
||||
bridge._handle_approval("ws-1", {"items": self._approval_items("bash", "read_file")})
|
||||
|
||||
mock_approve.assert_called_once()
|
||||
assert mock_approve.call_args.kwargs.get("approved") is False
|
||||
@@ -2,16 +2,6 @@
|
||||
|
||||
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):
|
||||
|
||||
@@ -2,16 +2,6 @@
|
||||
|
||||
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."""
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
def _make_watch_kwargs(**overrides):
|
||||
"""Build default kwargs for create_watch."""
|
||||
|
||||
@@ -0,0 +1,175 @@
|
||||
"""Tests for turnstone.core.web_search — pluggable web search backends."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from turnstone.core.web_search import (
|
||||
DuckDuckGoClient,
|
||||
MCPSearchClient,
|
||||
TavilyClient,
|
||||
_format_ddg,
|
||||
_format_tavily,
|
||||
resolve_web_search_client,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Formatters
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFormatTavily:
|
||||
def test_formats_answer_and_results(self):
|
||||
data = {
|
||||
"answer": "Python is great",
|
||||
"results": [
|
||||
{"title": "Python.org", "url": "https://python.org", "content": "Official site"},
|
||||
{"title": "PyPI", "url": "https://pypi.org", "content": "Package index"},
|
||||
],
|
||||
}
|
||||
out = _format_tavily(data, "python")
|
||||
assert "Answer: Python is great" in out
|
||||
assert "[Python.org](https://python.org)" in out
|
||||
assert "[PyPI](https://pypi.org)" in out
|
||||
|
||||
def test_no_results(self):
|
||||
out = _format_tavily({"results": []}, "nothing")
|
||||
assert "No results for 'nothing'" in out
|
||||
|
||||
def test_no_answer(self):
|
||||
data = {
|
||||
"results": [{"title": "T", "url": "http://t", "content": "C"}],
|
||||
}
|
||||
out = _format_tavily(data, "q")
|
||||
assert "Answer:" not in out
|
||||
assert "[T](http://t)" in out
|
||||
|
||||
|
||||
class TestFormatDDG:
|
||||
def test_formats_results(self):
|
||||
results = [
|
||||
{"title": "DDG Result", "href": "https://ddg.example.com", "body": "Search body"},
|
||||
]
|
||||
out = _format_ddg(results, "test")
|
||||
assert "[DDG Result](https://ddg.example.com)" in out
|
||||
assert "Search body" in out
|
||||
|
||||
def test_no_results(self):
|
||||
out = _format_ddg([], "nothing")
|
||||
assert "No results for 'nothing'" in out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Client tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestTavilyClient:
|
||||
def test_search_calls_api(self):
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.json.return_value = {
|
||||
"answer": "42",
|
||||
"results": [{"title": "T", "url": "http://t", "content": "C"}],
|
||||
}
|
||||
with patch("turnstone.core.web_search.httpx.post", return_value=mock_resp) as mock_post:
|
||||
client = TavilyClient("test-key", timeout=10)
|
||||
result = client.search("meaning of life", max_results=3)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
call_kwargs = mock_post.call_args
|
||||
assert call_kwargs.kwargs["json"]["query"] == "meaning of life"
|
||||
assert call_kwargs.kwargs["json"]["max_results"] == 3
|
||||
assert "Answer: 42" in result
|
||||
|
||||
|
||||
class TestDuckDuckGoClient:
|
||||
def test_integration_via_mock_ddgs(self):
|
||||
"""Patch the duckduckgo_search import inside DuckDuckGoClient.search."""
|
||||
mock_ddgs = MagicMock()
|
||||
mock_ddgs.__enter__ = MagicMock(return_value=mock_ddgs)
|
||||
mock_ddgs.__exit__ = MagicMock(return_value=False)
|
||||
mock_ddgs.text.return_value = [
|
||||
{"title": "DDG Result", "href": "https://ddg.co", "body": "Found it"},
|
||||
]
|
||||
mock_module = MagicMock()
|
||||
mock_module.DDGS.return_value = mock_ddgs
|
||||
with patch.dict("sys.modules", {"duckduckgo_search": mock_module}):
|
||||
client = DuckDuckGoClient(timeout=10)
|
||||
result = client.search("test query", max_results=3)
|
||||
mock_ddgs.text.assert_called_once_with("test query", max_results=3)
|
||||
assert "[DDG Result](https://ddg.co)" in result
|
||||
assert "Found it" in result
|
||||
|
||||
|
||||
class TestMCPSearchClient:
|
||||
def test_delegates_to_mcp(self):
|
||||
mcp = MagicMock()
|
||||
mcp.call_tool_sync.return_value = "MCP search results"
|
||||
client = MCPSearchClient(mcp, "mcp__ddg__search", timeout=30)
|
||||
result = client.search("test", max_results=3, topic="news")
|
||||
mcp.call_tool_sync.assert_called_once_with(
|
||||
"mcp__ddg__search",
|
||||
{"query": "test", "max_results": 3, "topic": "news"},
|
||||
timeout=30,
|
||||
)
|
||||
assert result == "MCP search results"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Resolver
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestResolveClient:
|
||||
def test_auto_tavily_when_key_present(self):
|
||||
client = resolve_web_search_client("", tavily_key="key")
|
||||
assert isinstance(client, TavilyClient)
|
||||
|
||||
def test_auto_ddg_when_no_tavily(self):
|
||||
with patch("turnstone.core.web_search._ddg_available", return_value=True):
|
||||
client = resolve_web_search_client("", tavily_key=None)
|
||||
assert isinstance(client, DuckDuckGoClient)
|
||||
|
||||
def test_auto_none_when_nothing_available(self):
|
||||
with patch("turnstone.core.web_search._ddg_available", return_value=False):
|
||||
client = resolve_web_search_client("", tavily_key=None)
|
||||
assert client is None
|
||||
|
||||
def test_explicit_tavily(self):
|
||||
client = resolve_web_search_client("tavily", tavily_key="key")
|
||||
assert isinstance(client, TavilyClient)
|
||||
|
||||
def test_explicit_tavily_no_key(self):
|
||||
client = resolve_web_search_client("tavily", tavily_key=None)
|
||||
assert client is None
|
||||
|
||||
def test_explicit_ddg(self):
|
||||
with patch("turnstone.core.web_search._ddg_available", return_value=True):
|
||||
client = resolve_web_search_client("ddg", tavily_key=None)
|
||||
assert isinstance(client, DuckDuckGoClient)
|
||||
|
||||
def test_explicit_ddg_not_installed(self):
|
||||
with patch("turnstone.core.web_search._ddg_available", return_value=False):
|
||||
client = resolve_web_search_client("ddg", tavily_key=None)
|
||||
assert client is None
|
||||
|
||||
def test_mcp_backend(self):
|
||||
mcp = MagicMock()
|
||||
mcp.is_mcp_tool.return_value = True
|
||||
client = resolve_web_search_client("mcp:ddg:search", tavily_key=None, mcp_client=mcp)
|
||||
assert isinstance(client, MCPSearchClient)
|
||||
mcp.is_mcp_tool.assert_called_with("mcp__ddg__search")
|
||||
|
||||
def test_mcp_backend_not_connected(self):
|
||||
mcp = MagicMock()
|
||||
mcp.is_mcp_tool.return_value = False
|
||||
client = resolve_web_search_client("mcp:ddg:search", tavily_key=None, mcp_client=mcp)
|
||||
assert client is None
|
||||
|
||||
def test_mcp_backend_no_client(self):
|
||||
client = resolve_web_search_client("mcp:ddg:search", tavily_key=None, mcp_client=None)
|
||||
assert client is None
|
||||
|
||||
def test_unknown_backend_returns_none(self):
|
||||
client = resolve_web_search_client("typo_backend", tavily_key="key")
|
||||
assert client is None
|
||||
@@ -1,3 +1,3 @@
|
||||
"""turnstone - Multi-node AI orchestration platform with tool use, agent routing, and cluster simulation."""
|
||||
|
||||
__version__ = "0.8.3"
|
||||
__version__ = "0.8.5"
|
||||
|
||||
@@ -137,6 +137,9 @@ class ConsoleCreateWsRequest(BaseModel):
|
||||
default="", description="Optional first message sent after creation"
|
||||
)
|
||||
skill: str = Field(default="", description="Skill name (replaces default skills)")
|
||||
resume_ws: str = Field(
|
||||
default="", description="Workstream ID to resume (loads previous conversation)"
|
||||
)
|
||||
|
||||
|
||||
class ConsoleCreateWsResponse(BaseModel):
|
||||
@@ -306,6 +309,7 @@ class SkillInfo(BaseModel):
|
||||
agent_max_turns: int | None = None
|
||||
notify_on_complete: str = "{}"
|
||||
enabled: bool = True
|
||||
priority: int = 0
|
||||
allowed_tools: str = "[]"
|
||||
license: str = ""
|
||||
compatibility: str = ""
|
||||
@@ -338,6 +342,7 @@ class CreateSkillRequest(BaseModel):
|
||||
agent_max_turns: int | None = None
|
||||
notify_on_complete: str = "{}"
|
||||
enabled: bool = True
|
||||
priority: int = 0
|
||||
allowed_tools: str = "[]"
|
||||
license: str = ""
|
||||
compatibility: str = ""
|
||||
@@ -363,6 +368,7 @@ class UpdateSkillRequest(BaseModel):
|
||||
agent_max_turns: int | None = None
|
||||
notify_on_complete: str | None = None
|
||||
enabled: bool | None = None
|
||||
priority: int | None = None
|
||||
allowed_tools: str | None = None
|
||||
license: str | None = None
|
||||
compatibility: str | None = None
|
||||
|
||||
@@ -82,6 +82,7 @@ from turnstone.api.schemas import (
|
||||
CreateTokenRequest,
|
||||
CreateTokenResponse,
|
||||
CreateUserRequest,
|
||||
DeleteSettingResponse,
|
||||
ErrorResponse,
|
||||
ListScheduleRunsResponse,
|
||||
ListSchedulesResponse,
|
||||
@@ -751,7 +752,7 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [
|
||||
"/v1/api/admin/settings/{key}",
|
||||
"DELETE",
|
||||
"Reset a setting to its default value",
|
||||
response_model=StatusResponse,
|
||||
response_model=DeleteSettingResponse,
|
||||
query_params=[
|
||||
QueryParam("node_id", "Node ID for node-scoped settings"),
|
||||
],
|
||||
@@ -855,6 +856,7 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [
|
||||
_ALL_MODELS: list[type[BaseModel]] = [
|
||||
ErrorResponse,
|
||||
StatusResponse,
|
||||
DeleteSettingResponse,
|
||||
AuthLoginRequest,
|
||||
AuthLoginResponse,
|
||||
AuthSetupRequest,
|
||||
|
||||
@@ -8,6 +8,7 @@ as the single source of truth for the generated OpenAPI spec.
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import StrEnum
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
@@ -34,6 +35,14 @@ class StatusResponse(BaseModel):
|
||||
status: str = Field(default="ok", examples=["ok"])
|
||||
|
||||
|
||||
class DeleteSettingResponse(BaseModel):
|
||||
"""DELETE /v1/api/admin/settings/{key} response."""
|
||||
|
||||
status: str = Field(default="ok", examples=["ok"])
|
||||
key: str = Field(description="Dotted setting key that was reset")
|
||||
default: Any = Field(description="Registry default value the setting reverted to")
|
||||
|
||||
|
||||
class AuthLoginRequest(BaseModel):
|
||||
"""POST /v1/api/auth/login request body.
|
||||
|
||||
|
||||
@@ -157,8 +157,10 @@ class McpStatus(BaseModel):
|
||||
class HealthResponse(BaseModel):
|
||||
status: str = Field(examples=["ok", "degraded"])
|
||||
version: str = ""
|
||||
node_id: str = ""
|
||||
uptime_seconds: float = 0.0
|
||||
model: str = ""
|
||||
max_ws: int = Field(default=10, description="Maximum concurrent workstreams")
|
||||
workstreams: WorkstreamCounts = WorkstreamCounts()
|
||||
backend: BackendStatus | None = None
|
||||
mcp: McpStatus | None = None
|
||||
|
||||
@@ -306,7 +306,49 @@ class TurnstoneBot:
|
||||
await sm.append(event.text)
|
||||
|
||||
elif isinstance(event, ApprovalRequestEvent):
|
||||
if self.config.auto_approve or self._should_auto_approve(event):
|
||||
# Evaluate admin tool policies before auto-approve.
|
||||
_policy_handled = False
|
||||
if self.storage is not None:
|
||||
try:
|
||||
from turnstone.core.policy import evaluate_tool_policies_batch
|
||||
|
||||
_tool_names = [
|
||||
it.get("approval_label", "") or it.get("func_name", "")
|
||||
for it in event.items
|
||||
if it.get("needs_approval") and it.get("func_name") and not it.get("error")
|
||||
]
|
||||
_tool_names = [n for n in _tool_names if n]
|
||||
if _tool_names:
|
||||
verdicts = await asyncio.to_thread(
|
||||
evaluate_tool_policies_batch,
|
||||
self.storage,
|
||||
_tool_names,
|
||||
)
|
||||
if any(v == "deny" for v in verdicts.values()):
|
||||
denied = [n for n, v in verdicts.items() if v == "deny"]
|
||||
await self.router.send_approval(
|
||||
ws_id,
|
||||
event.correlation_id,
|
||||
approved=False,
|
||||
feedback=f"Blocked by tool policy: {', '.join(denied)}",
|
||||
)
|
||||
await thread.send(
|
||||
f"*Tool blocked by admin policy: {', '.join(denied)}*"
|
||||
)
|
||||
_policy_handled = True
|
||||
elif all(verdicts.get(n) == "allow" for n in _tool_names):
|
||||
await self.router.send_approval(
|
||||
ws_id,
|
||||
event.correlation_id,
|
||||
approved=True,
|
||||
)
|
||||
await thread.send("*Tool approved by policy.*")
|
||||
_policy_handled = True
|
||||
except Exception:
|
||||
log.debug("Tool policy evaluation failed for ws %s", ws_id, exc_info=True)
|
||||
if not _policy_handled and (
|
||||
self.config.auto_approve or self._should_auto_approve(event)
|
||||
):
|
||||
await self.router.send_approval(ws_id, event.correlation_id, approved=True)
|
||||
await thread.send("*Tool auto-approved.*")
|
||||
else:
|
||||
|
||||
@@ -1,42 +0,0 @@
|
||||
"""chat.py — Backward-compatibility shim.
|
||||
|
||||
All functionality has been moved to submodules:
|
||||
- turnstone.core.session: ChatSession, SessionUI
|
||||
- turnstone.core.tools: TOOLS, AGENT_TOOLS, TASK_AGENT_TOOLS
|
||||
- turnstone.core.edit: find_occurrences, pick_nearest
|
||||
- turnstone.core.sandbox: validate_math_code, execute_math_sandboxed
|
||||
- turnstone.core.safety: is_command_blocked, sanitize_command
|
||||
- turnstone.core.web: strip_html, check_ssrf
|
||||
- turnstone.core.memory: save_message, structured memory facade, etc.
|
||||
- turnstone.ui.colors: ANSI constants and helpers
|
||||
- turnstone.ui.markdown: MarkdownRenderer
|
||||
- turnstone.ui.spinner: Spinner
|
||||
- turnstone.cli: TerminalUI, main, detect_model
|
||||
"""
|
||||
|
||||
# Re-export public API for backward compatibility
|
||||
from turnstone.cli import detect_model, main # noqa: F401
|
||||
from turnstone.core.session import ChatSession, SessionUI # noqa: F401
|
||||
from turnstone.core.tools import AGENT_TOOLS, TASK_AGENT_TOOLS, TOOLS # noqa: F401
|
||||
from turnstone.core.web import strip_html as _strip_html # noqa: F401
|
||||
from turnstone.ui.colors import ( # noqa: F401
|
||||
BLUE,
|
||||
BOLD,
|
||||
CYAN,
|
||||
DIM,
|
||||
GRAY,
|
||||
GREEN,
|
||||
ITALIC,
|
||||
MAGENTA,
|
||||
RED,
|
||||
RESET,
|
||||
YELLOW,
|
||||
bold,
|
||||
cyan,
|
||||
dim,
|
||||
green,
|
||||
red,
|
||||
yellow,
|
||||
)
|
||||
from turnstone.ui.markdown import MarkdownRenderer # noqa: F401
|
||||
from turnstone.ui.spinner import Spinner # noqa: F401
|
||||
+74
-6
@@ -14,6 +14,7 @@ import textwrap
|
||||
import threading
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from turnstone.core.judge import JudgeConfig
|
||||
from turnstone.core.session import ChatSession, SessionUI
|
||||
from turnstone.core.workstream import Workstream, WorkstreamManager, WorkstreamState
|
||||
from turnstone.ui.colors import (
|
||||
@@ -134,11 +135,42 @@ class TerminalUI(SessionUI):
|
||||
"""
|
||||
pending = [it for it in items if it.get("needs_approval") and not it.get("error")]
|
||||
|
||||
# Evaluate admin tool policies (deny/allow/ask) before prompting.
|
||||
if pending:
|
||||
try:
|
||||
from turnstone.core.policy import evaluate_tool_policies_batch
|
||||
from turnstone.core.storage._registry import get_storage
|
||||
|
||||
storage = get_storage()
|
||||
if storage is not None:
|
||||
_policy_names = [
|
||||
it.get("approval_label", "") or it.get("func_name", "")
|
||||
for it in pending
|
||||
if it.get("func_name")
|
||||
]
|
||||
if _policy_names:
|
||||
verdicts = evaluate_tool_policies_batch(storage, _policy_names)
|
||||
for it in pending:
|
||||
policy_name = it.get("approval_label", "") or it.get("func_name", "")
|
||||
verdict = verdicts.get(policy_name)
|
||||
if verdict == "deny":
|
||||
it["denied"] = True
|
||||
it["error"] = f"Blocked by tool policy ('{policy_name}')"
|
||||
it["needs_approval"] = False
|
||||
elif verdict == "allow":
|
||||
it["needs_approval"] = False
|
||||
pending = [
|
||||
it for it in items if it.get("needs_approval") and not it.get("error")
|
||||
]
|
||||
except Exception:
|
||||
pass # Best-effort — no policy enforcement on error
|
||||
|
||||
with self._print_lock:
|
||||
# Print all headers, previews, and heuristic verdicts
|
||||
for item in items:
|
||||
if item.get("error"):
|
||||
sys.stdout.write(f" {red(item['header'])}\n")
|
||||
sys.stdout.write(f" {red(item['error'])}\n")
|
||||
else:
|
||||
sys.stdout.write(f" {yellow(item['header'])}\n")
|
||||
if item.get("preview"):
|
||||
@@ -162,7 +194,11 @@ class TerminalUI(SessionUI):
|
||||
|
||||
# Per-tool auto-approve check
|
||||
if self.auto_approve_tools:
|
||||
pending_names = {it.get("func_name", "") for it in pending if it.get("func_name")}
|
||||
pending_names = {
|
||||
it.get("approval_label", "") or it.get("func_name", "")
|
||||
for it in pending
|
||||
if it.get("func_name")
|
||||
}
|
||||
if pending_names and pending_names.issubset(self.auto_approve_tools):
|
||||
return True, None
|
||||
|
||||
@@ -195,7 +231,11 @@ class TerminalUI(SessionUI):
|
||||
break
|
||||
|
||||
if decision in ("a", "always"):
|
||||
tool_names = {it.get("func_name", "") for it in pending if it.get("func_name")}
|
||||
tool_names = {
|
||||
it.get("approval_label", "") or it.get("func_name", "")
|
||||
for it in pending
|
||||
if it.get("func_name") and not it.get("error")
|
||||
}
|
||||
tool_names.discard("")
|
||||
tool_names.discard("__budget_override__")
|
||||
self.auto_approve_tools.update(tool_names)
|
||||
@@ -752,10 +792,15 @@ def _handle_cluster_command(cmd_line: str, console_url: str | None, auth_token:
|
||||
|
||||
|
||||
def detect_model(client: Any, provider: str = "openai") -> tuple[str, int | None]:
|
||||
"""Auto-detect model — delegates to :func:`turnstone.core.model_registry.detect_model`."""
|
||||
"""Auto-detect model — delegates to :func:`turnstone.core.model_registry.detect_model`.
|
||||
|
||||
CLI always uses fatal=True, so model is never None.
|
||||
"""
|
||||
from turnstone.core.model_registry import detect_model as _detect
|
||||
|
||||
return _detect(client, provider=provider)
|
||||
model, ctx = _detect(client, provider=provider)
|
||||
assert model is not None # fatal=True guarantees non-None or SystemExit
|
||||
return model, ctx
|
||||
|
||||
|
||||
# ─── Main ──────────────────────────────────────────────────────────────────
|
||||
@@ -870,6 +915,12 @@ def main() -> None:
|
||||
default=5,
|
||||
help="Max tools returned per tool search query (default: 5)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--web-search-backend",
|
||||
default="",
|
||||
metavar="BACKEND",
|
||||
help="Web search backend: '' (auto), 'tavily', 'ddg', or 'mcp:server:tool'",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--resume",
|
||||
default=None,
|
||||
@@ -959,8 +1010,9 @@ def main() -> None:
|
||||
default=0.7,
|
||||
help="Confidence threshold for judge (default: 0.7)",
|
||||
)
|
||||
from turnstone.core.config import apply_config
|
||||
from turnstone.core.config import add_config_arg, apply_config
|
||||
|
||||
add_config_arg(parser)
|
||||
apply_config(
|
||||
parser,
|
||||
["api", "model", "session", "tools", "console", "auth", "mcp", "database", "judge"],
|
||||
@@ -980,7 +1032,7 @@ def main() -> None:
|
||||
db_url = getattr(args, "db_url", None) or os.environ.get("TURNSTONE_DB_URL", "")
|
||||
db_path = getattr(args, "db_path", None) or os.environ.get("TURNSTONE_DB_PATH", "")
|
||||
db_pool_size = int(
|
||||
getattr(args, "db_pool_size", None) or os.environ.get("TURNSTONE_DB_POOL_SIZE", "5")
|
||||
getattr(args, "db_pool_size", None) or os.environ.get("TURNSTONE_DB_POOL_SIZE", "2")
|
||||
)
|
||||
init_storage(db_backend, path=db_path, url=db_url, pool_size=db_pool_size)
|
||||
|
||||
@@ -1040,6 +1092,20 @@ def main() -> None:
|
||||
storage=_get_storage(),
|
||||
)
|
||||
|
||||
# apply_config() merges [judge] config.toml values into args as
|
||||
# judge_base_url, judge_api_key, etc. Output_guard and redact_secrets
|
||||
# default to True, enabling the heuristic guard even when the LLM judge
|
||||
# is disabled via --no-judge.
|
||||
judge_config = JudgeConfig(
|
||||
enabled=args.judge_enabled,
|
||||
model=args.judge_model,
|
||||
provider=args.judge_provider,
|
||||
base_url=getattr(args, "judge_base_url", ""),
|
||||
api_key=getattr(args, "judge_api_key", ""),
|
||||
confidence_threshold=args.judge_confidence,
|
||||
timeout=args.judge_timeout,
|
||||
)
|
||||
|
||||
# ChatSession factory — captures shared config for creating workstreams
|
||||
def session_factory(
|
||||
ui: SessionUI | None,
|
||||
@@ -1070,7 +1136,9 @@ def main() -> None:
|
||||
tool_search=args.tool_search,
|
||||
tool_search_threshold=args.tool_search_threshold,
|
||||
tool_search_max_results=args.tool_search_max_results,
|
||||
web_search_backend=args.web_search_backend,
|
||||
skill=skill or args.skill or None,
|
||||
judge_config=judge_config,
|
||||
)
|
||||
|
||||
# Create workstream manager and initial workstream
|
||||
|
||||
@@ -20,6 +20,7 @@ from typing import TYPE_CHECKING, Any
|
||||
import httpx
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from turnstone.core.auth import ServiceTokenManager
|
||||
from turnstone.mq.broker import RedisBroker
|
||||
|
||||
log = logging.getLogger("turnstone.console.collector")
|
||||
@@ -53,11 +54,12 @@ class ClusterCollector:
|
||||
self,
|
||||
broker: RedisBroker,
|
||||
prefix: str = "turnstone",
|
||||
poll_interval: float = 10.0,
|
||||
poll_interval: float = 15.0,
|
||||
discovery_interval: float = 15.0,
|
||||
max_poll_workers: int = 50,
|
||||
http_timeout: float = 5.0,
|
||||
max_poll_workers: int = 200,
|
||||
http_timeout: float = 30.0,
|
||||
auth_token: str = "",
|
||||
token_manager: ServiceTokenManager | None = None,
|
||||
):
|
||||
self._broker = broker
|
||||
self._prefix = prefix
|
||||
@@ -65,16 +67,26 @@ class ClusterCollector:
|
||||
self._discovery_interval = discovery_interval
|
||||
self._max_poll_workers = max_poll_workers
|
||||
self._http_timeout = http_timeout
|
||||
self._token_manager = token_manager
|
||||
# Static auth header — only used when no token_manager is present.
|
||||
# When a token_manager exists, auth is injected per-request via
|
||||
# extra_headers in _poll_all_nodes to avoid stale JWT expiry.
|
||||
self._static_auth: dict[str, str] | None = None
|
||||
if auth_token and token_manager is None:
|
||||
self._static_auth = {"Authorization": f"Bearer {auth_token}"}
|
||||
|
||||
self._lock = threading.Lock()
|
||||
self._nodes: dict[str, NodeSnapshot] = {}
|
||||
self._running = False
|
||||
self._threads: list[threading.Thread] = []
|
||||
self._poll_pool = ThreadPoolExecutor(max_workers=max_poll_workers)
|
||||
headers = {}
|
||||
if auth_token:
|
||||
headers["Authorization"] = f"Bearer {auth_token}"
|
||||
self._http_client = httpx.Client(timeout=http_timeout, headers=headers)
|
||||
self._http_client = httpx.Client(
|
||||
timeout=httpx.Timeout(connect=10, read=http_timeout, write=5, pool=http_timeout),
|
||||
limits=httpx.Limits(
|
||||
max_connections=max_poll_workers + 10,
|
||||
max_keepalive_connections=min(max_poll_workers, 200),
|
||||
),
|
||||
)
|
||||
|
||||
# SSE fan-out to browser clients
|
||||
self._listeners: list[queue.Queue[dict[str, Any]]] = []
|
||||
@@ -212,6 +224,7 @@ class ClusterCollector:
|
||||
self._nodes[nid].server_url = meta.get(
|
||||
"server_url", self._nodes[nid].server_url
|
||||
)
|
||||
self._nodes[nid].max_ws = meta.get("max_ws", self._nodes[nid].max_ws)
|
||||
|
||||
# Remove nodes whose heartbeats expired
|
||||
lost = [nid for nid in self._nodes if nid not in active_ids]
|
||||
@@ -233,8 +246,35 @@ class ClusterCollector:
|
||||
log.exception("Poll loop error")
|
||||
time.sleep(self._poll_interval)
|
||||
|
||||
@staticmethod
|
||||
def _node_jitter(node_id: str, window: float) -> float:
|
||||
"""Deterministic per-node delay within a sliding window.
|
||||
|
||||
Uses a Mersenne prime (2^31 - 1) to hash the node_id into a
|
||||
stable offset so each node is polled at a different point in
|
||||
the cycle. The offset is consistent across restarts for the
|
||||
same node_id, giving an even spread without randomness.
|
||||
"""
|
||||
h = hash(node_id) & 0x7FFFFFFF # positive 31-bit
|
||||
return (h % 2147483647) / 2147483647 * window # M31 = 2^31 - 1
|
||||
|
||||
def _poll_all_nodes(self) -> None:
|
||||
"""Fetch dashboard data from all known nodes in parallel."""
|
||||
"""Fetch dashboard data from all known nodes in parallel.
|
||||
|
||||
Submissions are throttled by the thread pool size to avoid a
|
||||
thundering herd — at most ``max_poll_workers`` concurrent HTTP
|
||||
requests are in flight at any time. Each worker sleeps a
|
||||
deterministic per-node jitter (derived from its node_id) to
|
||||
spread requests across the first half of the poll interval.
|
||||
"""
|
||||
# Snapshot current auth header for this poll cycle. Per-request
|
||||
# headers avoid mutating shared client state (thread-safe).
|
||||
if self._token_manager is not None:
|
||||
poll_headers: dict[str, str] | None = {
|
||||
"Authorization": f"Bearer {self._token_manager.token}"
|
||||
}
|
||||
else:
|
||||
poll_headers = self._static_auth
|
||||
with self._lock:
|
||||
targets = [
|
||||
(n.node_id, n.server_url)
|
||||
@@ -245,27 +285,57 @@ class ClusterCollector:
|
||||
if not targets:
|
||||
return
|
||||
|
||||
futures = {self._poll_pool.submit(self._fetch_node, nid, url): nid for nid, url in targets}
|
||||
jitter_window = self._poll_interval / 2
|
||||
|
||||
def _jittered_fetch(
|
||||
nid: str, url: str, headers: dict[str, str] | None
|
||||
) -> tuple[dict[str, Any], dict[str, Any]]:
|
||||
delay = self._node_jitter(nid, jitter_window)
|
||||
if delay > 0.1:
|
||||
time.sleep(delay)
|
||||
return self._fetch_node(nid, url, headers)
|
||||
|
||||
futures = {
|
||||
self._poll_pool.submit(_jittered_fetch, nid, url, poll_headers): nid
|
||||
for nid, url in targets
|
||||
}
|
||||
for future in as_completed(futures):
|
||||
nid = futures[future]
|
||||
try:
|
||||
dashboard, health = future.result()
|
||||
self._apply_poll(nid, dashboard, health)
|
||||
except httpx.HTTPStatusError as exc:
|
||||
if exc.response.status_code in (401, 403):
|
||||
log.warning(
|
||||
"Auth failure polling node %s: HTTP %d", nid, exc.response.status_code
|
||||
)
|
||||
else:
|
||||
log.debug("Failed to poll node %s: HTTP %d", nid, exc.response.status_code)
|
||||
with self._lock:
|
||||
if nid in self._nodes:
|
||||
self._nodes[nid].reachable = False
|
||||
except Exception:
|
||||
log.debug("Failed to poll node %s", nid)
|
||||
log.warning("Failed to poll node %s", nid, exc_info=True)
|
||||
with self._lock:
|
||||
if nid in self._nodes:
|
||||
self._nodes[nid].reachable = False
|
||||
|
||||
def _fetch_node(self, node_id: str, server_url: str) -> tuple[dict[str, Any], dict[str, Any]]:
|
||||
def _fetch_node(
|
||||
self,
|
||||
node_id: str,
|
||||
server_url: str,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[dict[str, Any], dict[str, Any]]:
|
||||
"""Fetch /v1/api/dashboard and /health from a single node."""
|
||||
base = server_url.rstrip("/")
|
||||
dash_resp = self._http_client.get(f"{base}/v1/api/dashboard")
|
||||
dash_resp = self._http_client.get(f"{base}/v1/api/dashboard", headers=extra_headers)
|
||||
dash_resp.raise_for_status()
|
||||
dash_data: dict[str, Any] = dash_resp.json()
|
||||
try:
|
||||
health_resp = self._http_client.get(f"{base}/health")
|
||||
health_resp = self._http_client.get(f"{base}/health", headers=extra_headers)
|
||||
health_data: dict[str, Any] = health_resp.json()
|
||||
except Exception:
|
||||
log.debug("Failed to fetch health from %s", node_id, exc_info=True)
|
||||
health_data = {}
|
||||
return dash_data, health_data
|
||||
|
||||
@@ -373,9 +443,12 @@ class ClusterCollector:
|
||||
}
|
||||
|
||||
def get_nodes(
|
||||
self, sort_by: str = "activity", limit: int = 100, offset: int = 0
|
||||
self, sort_by: str = "activity", limit: int | None = 100, offset: int = 0
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
"""Return sorted, paginated node list with per-node counts."""
|
||||
"""Return sorted, paginated node list with per-node counts.
|
||||
|
||||
Pass ``limit=None`` to return all nodes (no pagination).
|
||||
"""
|
||||
with self._lock:
|
||||
items = []
|
||||
for node in self._nodes.values():
|
||||
@@ -423,8 +496,15 @@ class ClusterCollector:
|
||||
elif sort_by == "name":
|
||||
items.sort(key=lambda n: n["node_id"])
|
||||
|
||||
if limit is None:
|
||||
return items[offset:], total
|
||||
return items[offset : offset + limit], total
|
||||
|
||||
def get_all_nodes(self) -> list[dict[str, Any]]:
|
||||
"""Return all nodes without pagination (for fan-out operations)."""
|
||||
nodes, _ = self.get_nodes(sort_by="activity", limit=None)
|
||||
return nodes
|
||||
|
||||
def get_workstreams(
|
||||
self,
|
||||
state: str | None = None,
|
||||
|
||||
+142
-79
@@ -168,7 +168,7 @@ def _get_server_url(request: Request, node_id: str) -> str | None:
|
||||
|
||||
def _pick_best_node(collector: ClusterCollector) -> str:
|
||||
"""Select the reachable node with the most available capacity."""
|
||||
nodes, _ = collector.get_nodes(sort_by="activity", limit=1000, offset=0)
|
||||
nodes = collector.get_all_nodes()
|
||||
best_id = ""
|
||||
best_headroom = -1
|
||||
for n in nodes:
|
||||
@@ -252,7 +252,7 @@ async def cluster_snapshot(request: Request) -> JSONResponse:
|
||||
|
||||
async def cluster_events_sse(request: Request) -> Response:
|
||||
collector: ClusterCollector = request.app.state.collector
|
||||
client_queue: queue.Queue[dict[str, Any]] = queue.Queue(maxsize=500)
|
||||
client_queue: queue.Queue[dict[str, Any]] = queue.Queue(maxsize=2000)
|
||||
|
||||
async def event_generator() -> AsyncGenerator[dict[str, str], None]:
|
||||
loop = asyncio.get_running_loop()
|
||||
@@ -371,6 +371,7 @@ async def create_workstream(request: Request) -> JSONResponse:
|
||||
raw_model = body.get("model", "")
|
||||
raw_initial_message = body.get("initial_message", "")
|
||||
raw_skill = body.get("skill", "")
|
||||
raw_resume_ws = body.get("resume_ws", "")
|
||||
if not isinstance(raw_node_id, str):
|
||||
raw_node_id = "" if raw_node_id is None else None
|
||||
if not isinstance(raw_name, str):
|
||||
@@ -381,15 +382,20 @@ async def create_workstream(request: Request) -> JSONResponse:
|
||||
raw_initial_message = "" if raw_initial_message is None else None
|
||||
if not isinstance(raw_skill, str):
|
||||
raw_skill = "" if raw_skill is None else None
|
||||
if not isinstance(raw_resume_ws, str):
|
||||
raw_resume_ws = "" if raw_resume_ws is None else None
|
||||
if (
|
||||
raw_node_id is None
|
||||
or raw_name is None
|
||||
or raw_model is None
|
||||
or raw_initial_message is None
|
||||
or raw_skill is None
|
||||
or raw_resume_ws is None
|
||||
):
|
||||
return JSONResponse(
|
||||
{"error": "node_id, name, model, initial_message, and skill must be strings"},
|
||||
{
|
||||
"error": "node_id, name, model, initial_message, skill, and resume_ws must be strings"
|
||||
},
|
||||
status_code=400,
|
||||
)
|
||||
node_id = raw_node_id
|
||||
@@ -397,6 +403,7 @@ async def create_workstream(request: Request) -> JSONResponse:
|
||||
model = raw_model[:128]
|
||||
initial_message = raw_initial_message[:4096]
|
||||
skill = raw_skill[:256]
|
||||
resume_ws = raw_resume_ws[:64]
|
||||
|
||||
from turnstone.mq.protocol import CreateWorkstreamMessage
|
||||
|
||||
@@ -407,6 +414,7 @@ async def create_workstream(request: Request) -> JSONResponse:
|
||||
model=model,
|
||||
initial_message=initial_message,
|
||||
skill=skill,
|
||||
resume_ws=resume_ws,
|
||||
)
|
||||
broker.push_inbound(msg.to_json())
|
||||
log.debug("Pool dispatch: correlation_id=%s name=%r", msg.correlation_id, name)
|
||||
@@ -435,6 +443,7 @@ async def create_workstream(request: Request) -> JSONResponse:
|
||||
target_node=node_id,
|
||||
initial_message=initial_message,
|
||||
skill=skill,
|
||||
resume_ws=resume_ws,
|
||||
)
|
||||
broker.push_inbound(msg.to_json(), node_id=node_id)
|
||||
|
||||
@@ -670,17 +679,40 @@ async def _proxy_sse(
|
||||
|
||||
@asynccontextmanager
|
||||
async def _lifespan(app: Starlette) -> AsyncGenerator[None, None]:
|
||||
# Create async HTTP client for proxy routes
|
||||
headers: dict[str, str] = {}
|
||||
token = app.state.proxy_auth_token
|
||||
if token:
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
app.state.proxy_client = httpx.AsyncClient(timeout=30, headers=headers)
|
||||
# Separate client for SSE streams — longer read timeout, shared connection pool
|
||||
# Create async HTTP clients for proxy routes. Auth headers are NOT baked
|
||||
# in — _proxy_auth_headers() injects a fresh token per-request so JWTs
|
||||
# auto-rotate via ServiceTokenManager instead of expiring after 1 hour.
|
||||
# Size the pool above the fan-out limit to leave headroom for non-fan-out
|
||||
# proxy traffic (UI proxying, SSE streams, etc.).
|
||||
#
|
||||
# Build a ConfigStore so console settings reads get type validation and
|
||||
# caching instead of raw storage.get_system_setting() calls.
|
||||
storage = getattr(app.state, "auth_storage", None)
|
||||
config_store = None
|
||||
if storage:
|
||||
try:
|
||||
from turnstone.core.config_store import ConfigStore
|
||||
|
||||
config_store = ConfigStore(storage)
|
||||
except Exception:
|
||||
log.warning("Failed to initialise ConfigStore", exc_info=True)
|
||||
app.state.config_store = config_store
|
||||
fan_out = (
|
||||
config_store.get("cluster.node_fan_out_limit") if config_store else _NODE_FAN_OUT_LIMIT
|
||||
)
|
||||
app.state.fan_out_limit = fan_out
|
||||
app.state.proxy_client = httpx.AsyncClient(
|
||||
timeout=30,
|
||||
limits=httpx.Limits(
|
||||
max_connections=fan_out + 50,
|
||||
max_keepalive_connections=min(fan_out // 4, 100),
|
||||
),
|
||||
)
|
||||
app.state.proxy_sse_client = httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(connect=5, read=30, write=5, pool=5),
|
||||
limits=httpx.Limits(keepalive_expiry=30),
|
||||
headers=headers,
|
||||
limits=httpx.Limits(
|
||||
max_connections=1100, max_keepalive_connections=100, keepalive_expiry=30
|
||||
),
|
||||
)
|
||||
# Start scheduler if configured
|
||||
scheduler = getattr(app.state, "scheduler", None)
|
||||
@@ -1467,10 +1499,10 @@ async def admin_list_watches(request: Request) -> JSONResponse:
|
||||
if err:
|
||||
return err
|
||||
collector: ClusterCollector = request.app.state.collector
|
||||
nodes, _ = collector.get_nodes(limit=500)
|
||||
nodes = collector.get_all_nodes()
|
||||
client: httpx.AsyncClient = request.app.state.proxy_client
|
||||
headers = _proxy_auth_headers(request)
|
||||
sem = asyncio.Semaphore(_NODE_FAN_OUT_LIMIT)
|
||||
sem = asyncio.Semaphore(_get_fan_out_limit(request))
|
||||
|
||||
async def _fetch_node(node: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
server_url = (node.get("server_url") or "").rstrip("/")
|
||||
@@ -1509,9 +1541,14 @@ async def admin_list_watches(request: Request) -> JSONResponse:
|
||||
_VALID_WATCH_ID = re.compile(r"^[a-fA-F0-9]+$")
|
||||
|
||||
# Max concurrent outbound requests when fanning out to cluster nodes.
|
||||
# Sized below the default httpx pool limit (100) to leave headroom for
|
||||
# other proxy traffic (UI proxying, SSE streams, etc.).
|
||||
_NODE_FAN_OUT_LIMIT = 50
|
||||
# Must stay below the httpx pool limit (set in _lifespan) to leave
|
||||
# headroom for non-fan-out proxy traffic (UI proxying, SSE streams).
|
||||
_NODE_FAN_OUT_LIMIT = 200 # fallback; prefer cluster.node_fan_out_limit from storage
|
||||
|
||||
|
||||
def _get_fan_out_limit(request: Request) -> int:
|
||||
"""Return the fan-out limit cached at startup on app.state."""
|
||||
return int(getattr(request.app.state, "fan_out_limit", _NODE_FAN_OUT_LIMIT))
|
||||
|
||||
|
||||
async def admin_cancel_watch(request: Request) -> Response:
|
||||
@@ -2149,6 +2186,7 @@ _SKILL_RUNTIME_CONFIG_FIELDS = frozenset(
|
||||
"allowed_tools",
|
||||
"enabled",
|
||||
"notify_on_complete",
|
||||
"priority",
|
||||
}
|
||||
)
|
||||
|
||||
@@ -2303,6 +2341,7 @@ def _skill_to_response(r: dict[str, Any], resource_count: int = 0) -> dict[str,
|
||||
"agent_max_turns": r.get("agent_max_turns"),
|
||||
"notify_on_complete": r.get("notify_on_complete", "{}"),
|
||||
"enabled": r.get("enabled", True),
|
||||
"priority": r.get("priority", 0),
|
||||
"allowed_tools": r.get("allowed_tools", "[]"),
|
||||
"license": r.get("license", ""),
|
||||
"compatibility": r.get("compatibility", ""),
|
||||
@@ -2417,6 +2456,11 @@ async def admin_create_skill(request: Request) -> JSONResponse:
|
||||
if activation == "default":
|
||||
is_default = True
|
||||
|
||||
try:
|
||||
priority = max(-1000, min(1000, int(body.get("priority", 0) or 0)))
|
||||
except (ValueError, TypeError):
|
||||
priority = 0
|
||||
|
||||
if not name:
|
||||
return JSONResponse({"error": "name is required"}, status_code=400)
|
||||
if not content:
|
||||
@@ -2444,6 +2488,7 @@ async def admin_create_skill(request: Request) -> JSONResponse:
|
||||
compatibility=compatibility,
|
||||
activation=activation,
|
||||
token_estimate=token_estimate,
|
||||
priority=priority,
|
||||
**session_fields,
|
||||
)
|
||||
|
||||
@@ -2535,6 +2580,11 @@ async def admin_update_skill(request: Request) -> JSONResponse:
|
||||
except (ValueError, TypeError):
|
||||
tag_str = "[]"
|
||||
updates["tags"] = tag_str
|
||||
if "priority" in body:
|
||||
try:
|
||||
updates["priority"] = max(-1000, min(1000, int(body["priority"] or 0)))
|
||||
except (ValueError, TypeError):
|
||||
updates["priority"] = 0
|
||||
|
||||
# Installed (readonly) skills: restrict updates to runtime config only.
|
||||
# Spec/content fields are locked to preserve external-source fidelity.
|
||||
@@ -3058,20 +3108,18 @@ async def admin_delete_skill_resource(request: Request) -> JSONResponse:
|
||||
|
||||
|
||||
def _get_discovery_url(request: Request) -> str:
|
||||
"""Get skills discovery URL from DB settings, config.toml, or default."""
|
||||
"""Get skills discovery URL via ConfigStore, config.toml, or default."""
|
||||
from turnstone.core.config import load_config
|
||||
from turnstone.core.skill_sources import DEFAULT_DISCOVERY_URL
|
||||
|
||||
storage = getattr(request.app.state, "auth_storage", None)
|
||||
if storage:
|
||||
try:
|
||||
row = storage.get_system_setting("skills.discovery_url")
|
||||
if row:
|
||||
val = json.loads(row["value"])
|
||||
if val:
|
||||
return str(val)
|
||||
except (KeyError, json.JSONDecodeError, TypeError, AttributeError):
|
||||
pass
|
||||
# ConfigStore: validated + cached
|
||||
config_store = getattr(request.app.state, "config_store", None)
|
||||
if config_store:
|
||||
val = config_store.get("skills.discovery_url")
|
||||
if val:
|
||||
return str(val)
|
||||
|
||||
# Fall back to config.toml [skills] section
|
||||
skills_cfg = load_config("skills")
|
||||
url = skills_cfg.get("discovery_url", "")
|
||||
if url:
|
||||
@@ -3430,31 +3478,40 @@ async def admin_delete_memory(request: Request) -> JSONResponse:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _publish_config_change(request: Request, *, key: str, node_id: str, action: str) -> None:
|
||||
"""Fan out config-reload to all known server nodes (best-effort).
|
||||
async def _publish_config_change(request: Request) -> None:
|
||||
"""Fan out config-reload to all known server nodes (best-effort, async).
|
||||
|
||||
Uses the collector's node registry and the existing proxy auth
|
||||
mechanism — no MQ dependency.
|
||||
Uses the collector's node registry, the shared async proxy client,
|
||||
and bounded concurrency via the fan-out semaphore.
|
||||
"""
|
||||
import contextlib
|
||||
|
||||
import httpx
|
||||
# Reload the console's own ConfigStore so cached values stay fresh
|
||||
# (must happen even when collector is absent — e.g. standalone console)
|
||||
config_store = getattr(request.app.state, "config_store", None)
|
||||
if config_store:
|
||||
config_store.reload()
|
||||
|
||||
collector = getattr(request.app.state, "collector", None)
|
||||
if not collector:
|
||||
return
|
||||
client: httpx.AsyncClient = request.app.state.proxy_client
|
||||
headers = _proxy_auth_headers(request)
|
||||
with contextlib.suppress(Exception):
|
||||
nodes = collector.get_nodes()
|
||||
for node in nodes.get("nodes", []):
|
||||
url = node.get("url", "")
|
||||
if url:
|
||||
with contextlib.suppress(Exception):
|
||||
httpx.post(
|
||||
f"{url}/v1/api/_internal/config-reload",
|
||||
headers=headers,
|
||||
timeout=5.0,
|
||||
)
|
||||
sem = asyncio.Semaphore(_get_fan_out_limit(request))
|
||||
|
||||
async def _notify(url: str) -> None:
|
||||
async with sem:
|
||||
try:
|
||||
await client.post(
|
||||
f"{url.rstrip('/')}/v1/api/_internal/config-reload",
|
||||
headers=headers,
|
||||
timeout=5.0,
|
||||
)
|
||||
except Exception:
|
||||
log.warning("Config reload failed for %s", url, exc_info=True)
|
||||
|
||||
nodes = collector.get_all_nodes()
|
||||
tasks = [_notify(n["server_url"]) for n in nodes if n.get("server_url")]
|
||||
if tasks:
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
|
||||
async def admin_list_settings(request: Request) -> JSONResponse:
|
||||
@@ -3610,7 +3667,7 @@ async def admin_update_setting(request: Request) -> JSONResponse:
|
||||
ip,
|
||||
)
|
||||
|
||||
_publish_config_change(request, key=key, node_id=node_id, action="set")
|
||||
await _publish_config_change(request)
|
||||
|
||||
return JSONResponse(
|
||||
{
|
||||
@@ -3645,7 +3702,7 @@ async def admin_delete_setting(request: Request) -> JSONResponse:
|
||||
|
||||
key = request.path_params["key"]
|
||||
try:
|
||||
validate_key(key)
|
||||
defn = validate_key(key)
|
||||
except ValueError:
|
||||
return JSONResponse({"error": f"Unknown setting: {key}"}, status_code=400)
|
||||
|
||||
@@ -3665,9 +3722,9 @@ async def admin_delete_setting(request: Request) -> JSONResponse:
|
||||
ip,
|
||||
)
|
||||
|
||||
_publish_config_change(request, key=key, node_id=node_id, action="delete")
|
||||
await _publish_config_change(request)
|
||||
|
||||
return JSONResponse({"status": "ok", "key": key})
|
||||
return JSONResponse({"status": "ok", "key": key, "default": defn.default})
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -3676,21 +3733,16 @@ async def admin_delete_setting(request: Request) -> JSONResponse:
|
||||
|
||||
|
||||
def _get_registry_url(request: Request) -> str:
|
||||
"""Get the MCP Registry URL from DB settings, config.toml, or default."""
|
||||
"""Get the MCP Registry URL via ConfigStore, config.toml, or default."""
|
||||
from turnstone.core.config import load_config
|
||||
from turnstone.core.mcp_registry import DEFAULT_REGISTRY_URL
|
||||
|
||||
# Check database settings first
|
||||
storage = getattr(request.app.state, "auth_storage", None)
|
||||
if storage:
|
||||
try:
|
||||
row = storage.get_system_setting("mcp.registry_url")
|
||||
if row:
|
||||
val = json.loads(row["value"])
|
||||
if val:
|
||||
return str(val)
|
||||
except (KeyError, json.JSONDecodeError, TypeError, AttributeError):
|
||||
pass
|
||||
# ConfigStore: validated + cached
|
||||
config_store = getattr(request.app.state, "config_store", None)
|
||||
if config_store:
|
||||
val = config_store.get("mcp.registry_url")
|
||||
if val:
|
||||
return str(val)
|
||||
|
||||
# Fall back to config.toml [mcp] section
|
||||
mcp_cfg = load_config("mcp")
|
||||
@@ -3841,8 +3893,9 @@ async def admin_registry_install(request: Request) -> JSONResponse:
|
||||
|
||||
# Check max servers
|
||||
current = storage.list_mcp_servers()
|
||||
if len(current) >= _MCP_MAX_SERVERS:
|
||||
return JSONResponse({"error": f"Maximum {_MCP_MAX_SERVERS} servers"}, status_code=400)
|
||||
max_servers = _get_mcp_max_servers(request)
|
||||
if len(current) >= max_servers:
|
||||
return JSONResponse({"error": f"Maximum {max_servers} servers"}, status_code=400)
|
||||
|
||||
# Fetch the specific server from the registry
|
||||
registry_url = _get_registry_url(request)
|
||||
@@ -3938,7 +3991,15 @@ async def admin_registry_install(request: Request) -> JSONResponse:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_MCP_NAME_RE = re.compile(r"^[a-zA-Z0-9._-]+$")
|
||||
_MCP_MAX_SERVERS = 50
|
||||
_MCP_MAX_SERVERS = 200 # fallback; prefer cluster.mcp_max_servers from storage
|
||||
|
||||
|
||||
def _get_mcp_max_servers(request: Request) -> int:
|
||||
"""Read cluster.mcp_max_servers via ConfigStore (validated + cached)."""
|
||||
config_store = getattr(request.app.state, "config_store", None)
|
||||
if config_store:
|
||||
return int(config_store.get("cluster.mcp_max_servers"))
|
||||
return _MCP_MAX_SERVERS
|
||||
|
||||
|
||||
def _mask_mcp_secrets(server: dict[str, Any], reveal: bool = False) -> dict[str, Any]:
|
||||
@@ -3976,10 +4037,10 @@ async def _collect_mcp_status(
|
||||
) -> dict[str, dict[str, dict[str, Any]]]:
|
||||
"""Query all nodes for MCP status. Returns {node_id: {server_name: status}}."""
|
||||
collector: ClusterCollector = request.app.state.collector
|
||||
nodes, _ = collector.get_nodes(sort_by="activity", limit=1000, offset=0)
|
||||
nodes = collector.get_all_nodes()
|
||||
client: httpx.AsyncClient = request.app.state.proxy_client
|
||||
headers = _proxy_auth_headers(request)
|
||||
sem = asyncio.Semaphore(_NODE_FAN_OUT_LIMIT)
|
||||
sem = asyncio.Semaphore(_get_fan_out_limit(request))
|
||||
|
||||
async def _fetch(node: dict[str, Any]) -> tuple[str, dict[str, dict[str, Any]] | None]:
|
||||
node_id = node.get("node_id", "")
|
||||
@@ -4123,9 +4184,10 @@ async def admin_create_mcp_server(request: Request) -> JSONResponse:
|
||||
|
||||
# Check max servers
|
||||
existing = storage.list_mcp_servers()
|
||||
if len(existing) >= _MCP_MAX_SERVERS:
|
||||
max_servers = _get_mcp_max_servers(request)
|
||||
if len(existing) >= max_servers:
|
||||
return JSONResponse(
|
||||
{"error": f"Maximum {_MCP_MAX_SERVERS} servers"},
|
||||
{"error": f"Maximum {max_servers} servers"},
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
@@ -4327,10 +4389,10 @@ async def admin_delete_mcp_server(request: Request) -> JSONResponse:
|
||||
async def _notify_nodes_mcp_reload(request: Request) -> dict[str, Any]:
|
||||
"""Tell all nodes to re-read the mcp_servers DB table and reconcile."""
|
||||
collector: ClusterCollector = request.app.state.collector
|
||||
nodes, _ = collector.get_nodes(sort_by="activity", limit=1000, offset=0)
|
||||
nodes = collector.get_all_nodes()
|
||||
client: httpx.AsyncClient = request.app.state.proxy_client
|
||||
headers = _proxy_auth_headers(request)
|
||||
sem = asyncio.Semaphore(_NODE_FAN_OUT_LIMIT)
|
||||
sem = asyncio.Semaphore(_get_fan_out_limit(request))
|
||||
|
||||
async def _notify(node: dict[str, Any]) -> tuple[str, Any]:
|
||||
node_id = node.get("node_id", "")
|
||||
@@ -4406,6 +4468,7 @@ async def admin_import_mcp_config(request: Request) -> JSONResponse:
|
||||
errors: list[str] = []
|
||||
audit_uid, ip = _audit_context(request)
|
||||
current_count = len(storage.list_mcp_servers())
|
||||
max_servers = _get_mcp_max_servers(request)
|
||||
|
||||
for srv_name, cfg in servers.items():
|
||||
srv_name = str(srv_name).strip()[:64]
|
||||
@@ -4415,7 +4478,7 @@ async def admin_import_mcp_config(request: Request) -> JSONResponse:
|
||||
if storage.get_mcp_server_by_name(srv_name):
|
||||
skipped.append(srv_name)
|
||||
continue
|
||||
if current_count >= _MCP_MAX_SERVERS:
|
||||
if current_count >= max_servers:
|
||||
errors.append(f"{srv_name}: max servers reached")
|
||||
break
|
||||
|
||||
@@ -4817,8 +4880,9 @@ def main() -> None:
|
||||
help="Bearer token for polling turnstone-server nodes (default: $TURNSTONE_AUTH_TOKEN)",
|
||||
)
|
||||
|
||||
from turnstone.core.config import apply_config
|
||||
from turnstone.core.config import add_config_arg, apply_config
|
||||
|
||||
add_config_arg(parser)
|
||||
apply_config(parser, ["console", "redis", "auth"])
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -4854,13 +4918,13 @@ def main() -> None:
|
||||
audience=JWT_AUD_SERVER,
|
||||
expiry_hours=1,
|
||||
)
|
||||
collector_token = collector_token_mgr.token
|
||||
log.info("console.collector_jwt_minted")
|
||||
log.info("console.collector_token_manager_created")
|
||||
|
||||
collector = ClusterCollector(
|
||||
broker=broker,
|
||||
poll_interval=args.poll_interval,
|
||||
auth_token=collector_token,
|
||||
auth_token=collector_token if collector_token_mgr is None else "",
|
||||
token_manager=collector_token_mgr,
|
||||
)
|
||||
collector.start()
|
||||
|
||||
@@ -4898,8 +4962,7 @@ def main() -> None:
|
||||
audience=JWT_AUD_SERVER,
|
||||
expiry_hours=1,
|
||||
)
|
||||
proxy_token = proxy_token_mgr.token
|
||||
log.info("console.proxy_jwt_minted")
|
||||
log.info("console.proxy_token_manager_created")
|
||||
|
||||
from turnstone.core.web_helpers import parse_cors_origins
|
||||
|
||||
@@ -4911,7 +4974,7 @@ def main() -> None:
|
||||
auth_config=auth_config,
|
||||
jwt_secret=jwt_secret,
|
||||
auth_storage=auth_storage,
|
||||
proxy_auth_token=proxy_token,
|
||||
proxy_auth_token=proxy_token if proxy_token_mgr is None else "",
|
||||
proxy_token_mgr=proxy_token_mgr,
|
||||
cors_origins=cors_origins,
|
||||
)
|
||||
|
||||
@@ -2057,10 +2057,12 @@ var _settingsSectionOrder = [
|
||||
"session",
|
||||
"tools",
|
||||
"server",
|
||||
"cluster",
|
||||
"mcp",
|
||||
"ratelimit",
|
||||
"health",
|
||||
"judge",
|
||||
"skills",
|
||||
"memory",
|
||||
];
|
||||
|
||||
@@ -2070,10 +2072,12 @@ function _settingsSectionLabel(section) {
|
||||
session: "Session",
|
||||
tools: "Tools",
|
||||
server: "Server",
|
||||
cluster: "Cluster",
|
||||
mcp: "MCP",
|
||||
ratelimit: "Rate Limiting",
|
||||
health: "Health",
|
||||
judge: "Judge",
|
||||
skills: "Skills",
|
||||
memory: "Memory",
|
||||
};
|
||||
return labels[section] || section;
|
||||
|
||||
@@ -7,14 +7,15 @@ call after mutations to create a persistent audit trail.
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from turnstone.core.storage._protocol import StorageBackend
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
|
||||
def record_audit(
|
||||
|
||||
+13
-7
@@ -18,11 +18,9 @@ always accessible without authentication.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
@@ -40,7 +38,9 @@ if TYPE_CHECKING:
|
||||
|
||||
from turnstone.core.oidc import OIDCConfig
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constants
|
||||
@@ -1010,7 +1010,7 @@ async def handle_auth_status(request: Request) -> Response:
|
||||
users = storage.list_users()
|
||||
has_users = len(users) > 0
|
||||
except Exception:
|
||||
pass
|
||||
log.warning("Failed to check user existence for auth status", exc_info=True)
|
||||
|
||||
# OIDC configuration
|
||||
oidc_config = getattr(request.app.state, "oidc_config", None)
|
||||
@@ -1079,8 +1079,10 @@ async def handle_auth_setup(request: Request, audience: str) -> Response:
|
||||
except Exception:
|
||||
log.error("Failed to assign admin role to first user %s — aborting setup", user_id)
|
||||
# Roll back the user creation so setup can be retried
|
||||
with contextlib.suppress(Exception):
|
||||
try:
|
||||
storage.delete_user(user_id)
|
||||
except Exception:
|
||||
log.error("Failed to roll back user %s during setup abort", user_id, exc_info=True)
|
||||
return JSONResponse(
|
||||
{"error": "Failed to assign admin role. Ensure migrations have run."},
|
||||
status_code=503,
|
||||
@@ -1092,8 +1094,10 @@ async def handle_auth_setup(request: Request, audience: str) -> Response:
|
||||
log.error(
|
||||
"First user %s has no permissions after role assignment — aborting setup", user_id
|
||||
)
|
||||
with contextlib.suppress(Exception):
|
||||
try:
|
||||
storage.delete_user(user_id)
|
||||
except Exception:
|
||||
log.error("Failed to roll back user %s during setup abort", user_id, exc_info=True)
|
||||
return JSONResponse(
|
||||
{"error": "Failed to load permissions. Ensure migrations have run."},
|
||||
status_code=503,
|
||||
@@ -1229,8 +1233,10 @@ async def handle_oidc_callback(request: Request, audience: str) -> Response:
|
||||
return RedirectResponse("/?oidc_error=Too+many+login+attempts", status_code=302)
|
||||
|
||||
# Lazy cleanup of expired pending states
|
||||
with contextlib.suppress(Exception):
|
||||
try:
|
||||
storage.cleanup_expired_oidc_states(300)
|
||||
except Exception:
|
||||
log.debug("OIDC state cleanup failed", exc_info=True)
|
||||
|
||||
def _record_oidc_failure() -> None:
|
||||
if login_limiter is not None:
|
||||
|
||||
+64
-10
@@ -1,28 +1,59 @@
|
||||
"""Unified configuration for turnstone.
|
||||
|
||||
Loads ``~/.config/turnstone/config.toml`` and applies values as argparse defaults.
|
||||
Precedence: CLI args > env vars > config file > hardcoded defaults.
|
||||
Loads config.toml and applies values as argparse defaults.
|
||||
Precedence: CLI args > config file > hardcoded defaults.
|
||||
|
||||
Config file resolution:
|
||||
1. ``--config PATH`` CLI flag (via ``add_config_arg`` pre-parser)
|
||||
2. ``$TURNSTONE_CONFIG`` environment variable
|
||||
3. ``~/.config/turnstone/config.toml`` (default)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import tomllib
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import argparse
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
CONFIG_DIR = Path("~/.config/turnstone").expanduser()
|
||||
CONFIG_PATH = CONFIG_DIR / "config.toml"
|
||||
_DEFAULT_CONFIG_PATH = CONFIG_DIR / "config.toml"
|
||||
|
||||
# Resolved config path — set by set_config_path() or $TURNSTONE_CONFIG
|
||||
_config_path: Path | None = None
|
||||
|
||||
# Cache: None = not loaded yet, {} = loaded but empty/missing
|
||||
_cache: dict[str, Any] | None = None
|
||||
|
||||
|
||||
def _resolve_config_path() -> Path:
|
||||
"""Return the effective config file path."""
|
||||
if _config_path is not None:
|
||||
return _config_path
|
||||
env = os.environ.get("TURNSTONE_CONFIG", "").strip()
|
||||
if env:
|
||||
return Path(env).expanduser()
|
||||
return _DEFAULT_CONFIG_PATH
|
||||
|
||||
|
||||
def set_config_path(path: str) -> None:
|
||||
"""Override the config file path.
|
||||
|
||||
Invalidates the cache so subsequent ``load_config()`` calls re-read
|
||||
from the new path. Typically called from ``add_config_arg()``.
|
||||
"""
|
||||
global _config_path, _cache
|
||||
_config_path = Path(path).expanduser()
|
||||
_cache = None # invalidate cache so next load_config() re-reads
|
||||
|
||||
|
||||
def load_config(section: str | None = None) -> dict[str, Any]:
|
||||
"""Load config.toml and return the full dict or a specific section.
|
||||
|
||||
@@ -32,11 +63,12 @@ def load_config(section: str | None = None) -> dict[str, Any]:
|
||||
global _cache
|
||||
if _cache is None:
|
||||
_cache = {}
|
||||
if CONFIG_PATH.is_file():
|
||||
cfg_path = _resolve_config_path()
|
||||
if cfg_path.is_file():
|
||||
try:
|
||||
_cache = tomllib.loads(CONFIG_PATH.read_text(encoding="utf-8"))
|
||||
_cache = tomllib.loads(cfg_path.read_text(encoding="utf-8"))
|
||||
except Exception as exc:
|
||||
log.warning("Failed to parse %s: %s", CONFIG_PATH, exc)
|
||||
log.warning("Failed to parse %s: %s", cfg_path, exc)
|
||||
if section:
|
||||
result = _cache.get(section, {})
|
||||
return result if isinstance(result, dict) else {}
|
||||
@@ -73,6 +105,7 @@ _CONFIG_MAP: dict[str, dict[str, str]] = {
|
||||
"search": "tool_search",
|
||||
"search_threshold": "tool_search_threshold",
|
||||
"search_max_results": "tool_search_max_results",
|
||||
"web_search_backend": "web_search_backend",
|
||||
},
|
||||
"server": {
|
||||
"host": "host",
|
||||
@@ -156,8 +189,6 @@ def get_tavily_key() -> str | None:
|
||||
|
||||
Precedence: config.toml [api] tavily_key -> $TAVILY_API_KEY
|
||||
"""
|
||||
import os
|
||||
|
||||
global _tavily_key, _tavily_key_loaded
|
||||
if _tavily_key_loaded:
|
||||
return _tavily_key
|
||||
@@ -202,6 +233,29 @@ def apply_config(parser: argparse.ArgumentParser, sections: list[str]) -> None:
|
||||
parser.set_defaults(**defaults)
|
||||
|
||||
|
||||
def add_config_arg(parser: argparse.ArgumentParser) -> None:
|
||||
"""Add ``--config`` to *parser* and resolve the path before returning.
|
||||
|
||||
Uses a separate pre-parser (``add_help=False``) so ``--help`` on the
|
||||
main parser still works and shows config-derived defaults.
|
||||
"""
|
||||
import argparse as _ap
|
||||
import sys
|
||||
|
||||
parser.add_argument(
|
||||
"--config",
|
||||
default=None,
|
||||
metavar="PATH",
|
||||
help="Path to config.toml (default: $TURNSTONE_CONFIG or ~/.config/turnstone/config.toml)",
|
||||
)
|
||||
# Pre-parse only --config without intercepting --help
|
||||
pre = _ap.ArgumentParser(add_help=False)
|
||||
pre.add_argument("--config", default=None)
|
||||
pre_args, _ = pre.parse_known_args(sys.argv[1:])
|
||||
if pre_args.config:
|
||||
set_config_path(pre_args.config)
|
||||
|
||||
|
||||
def warn_migrated_settings() -> None:
|
||||
"""Log warnings for config.toml keys that are now managed by ConfigStore.
|
||||
|
||||
|
||||
@@ -18,10 +18,10 @@ ConfigStore) — it is a standalone tool, not a cluster node.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.settings_registry import (
|
||||
SETTINGS,
|
||||
deserialize_value,
|
||||
@@ -33,7 +33,7 @@ from turnstone.core.settings_registry import (
|
||||
if TYPE_CHECKING:
|
||||
from turnstone.core.storage._protocol import StorageBackend
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
_UNSET: Any = object()
|
||||
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
"""Subprocess environment scrubbing.
|
||||
|
||||
Builds a sanitized copy of ``os.environ`` that strips secrets
|
||||
(API keys, tokens, passwords) while preserving variables needed
|
||||
for normal tool operation (PATH, HOME, locale, etc.).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
# Env var names that are always preserved regardless of pattern matching.
|
||||
_SAFE_NAMES: frozenset[str] = frozenset(
|
||||
{
|
||||
"PATH",
|
||||
"HOME",
|
||||
"USER",
|
||||
"SHELL",
|
||||
"LANG",
|
||||
"TERM",
|
||||
"TMPDIR",
|
||||
"TMP",
|
||||
"TEMP",
|
||||
"EDITOR",
|
||||
"VISUAL",
|
||||
"COLORTERM",
|
||||
"COLUMNS",
|
||||
"LINES",
|
||||
"PWD",
|
||||
"OLDPWD",
|
||||
"HOSTNAME",
|
||||
"LOGNAME",
|
||||
"DISPLAY",
|
||||
"WAYLAND_DISPLAY",
|
||||
"SSH_AUTH_SOCK",
|
||||
"GPG_AGENT_INFO",
|
||||
"SHLVL",
|
||||
"MANWIDTH",
|
||||
"MAN_KEEP_FORMATTING",
|
||||
"LESS",
|
||||
"LESSOPEN",
|
||||
"LESSCLOSE",
|
||||
"LESSPIPE",
|
||||
"LESSCHARSET",
|
||||
}
|
||||
)
|
||||
|
||||
# Prefixes that are always preserved (locale, XDG, etc.).
|
||||
_SAFE_PREFIXES: tuple[str, ...] = ("LC_", "XDG_")
|
||||
|
||||
# Suffixes that cause a variable to be scrubbed (e.g. *_KEY, *_TOKEN).
|
||||
# Suffix matching avoids false positives on MONKEYTYPE, KEYBOARD_LAYOUT, etc.
|
||||
_SECRET_SUFFIXES: tuple[str, ...] = (
|
||||
"_KEY",
|
||||
"_SECRET",
|
||||
"_TOKEN",
|
||||
"_PASSWORD",
|
||||
"_CREDENTIAL",
|
||||
"_CREDENTIALS",
|
||||
)
|
||||
|
||||
# Exact names that are always scrubbed (even if they don't match patterns).
|
||||
_EXPLICIT_SCRUB: frozenset[str] = frozenset(
|
||||
{
|
||||
"OPENAI_API_KEY",
|
||||
"ANTHROPIC_API_KEY",
|
||||
"TAVILY_API_KEY",
|
||||
"TURNSTONE_JWT_SECRET",
|
||||
"TURNSTONE_AUTH_TOKEN",
|
||||
"TURNSTONE_DISCORD_TOKEN",
|
||||
"TURNSTONE_GITHUB_TOKEN",
|
||||
"TURNSTONE_OIDC_CLIENT_SECRET",
|
||||
"AWS_SECRET_ACCESS_KEY",
|
||||
"AZURE_CLIENT_SECRET",
|
||||
"GCP_SERVICE_ACCOUNT_KEY",
|
||||
"GOOGLE_APPLICATION_CREDENTIALS",
|
||||
"DATABASE_URL",
|
||||
"TURNSTONE_DB_URL",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _is_secret(name: str) -> bool:
|
||||
"""Return True if *name* looks like a secret variable."""
|
||||
if name in _EXPLICIT_SCRUB:
|
||||
return True
|
||||
upper = name.upper()
|
||||
return any(upper.endswith(sfx) for sfx in _SECRET_SUFFIXES)
|
||||
|
||||
|
||||
def _is_safe(name: str) -> bool:
|
||||
"""Return True if *name* should always be preserved."""
|
||||
if name in _SAFE_NAMES:
|
||||
return True
|
||||
return any(name.startswith(pfx) for pfx in _SAFE_PREFIXES)
|
||||
|
||||
|
||||
def scrubbed_env(
|
||||
extra: dict[str, str] | None = None,
|
||||
passthrough: list[str] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""Return a copy of ``os.environ`` with secrets removed.
|
||||
|
||||
Args:
|
||||
extra: Additional variables to merge on top (e.g. ``MANWIDTH``).
|
||||
passthrough: Explicit variable names to preserve even if they
|
||||
match secret patterns (operator override).
|
||||
"""
|
||||
passthrough_set = frozenset(passthrough) if passthrough else frozenset()
|
||||
env: dict[str, str] = {}
|
||||
for name, value in os.environ.items():
|
||||
if name in passthrough_set or _is_safe(name):
|
||||
env[name] = value
|
||||
elif _is_secret(name):
|
||||
continue
|
||||
else:
|
||||
env[name] = value
|
||||
if extra:
|
||||
env.update(extra)
|
||||
return env
|
||||
@@ -3,15 +3,16 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from openai import OpenAI
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
|
||||
class CircuitState(enum.Enum):
|
||||
@@ -149,11 +150,44 @@ class BackendHealthMonitor:
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _probe_loop(self) -> None:
|
||||
"""Background: probe backend every interval."""
|
||||
"""Background: probe backend every interval.
|
||||
|
||||
An initial jitter (derived from the PID) staggers probes across
|
||||
cluster nodes so they don't all hit the LLM backend at once.
|
||||
"""
|
||||
import os
|
||||
|
||||
# Deterministic per-process jitter: spread across half the interval
|
||||
jitter = ((os.getpid() * 2654435761) & 0x7FFFFFFF) / 0x7FFFFFFF * (self._probe_interval / 2)
|
||||
self._stop_event.wait(jitter)
|
||||
while not self._stop_event.is_set():
|
||||
self._stop_event.wait(self._probe_interval)
|
||||
if self._stop_event.is_set():
|
||||
break
|
||||
# When circuit is OPEN, only probe after cooldown expires.
|
||||
with self._lock:
|
||||
if self._state == CircuitState.OPEN:
|
||||
elapsed = time.monotonic() - self._last_state_change
|
||||
remaining = self._cooldown - elapsed
|
||||
if remaining > 0:
|
||||
# Wait precisely for cooldown rather than skipping
|
||||
# a full probe_interval (which could overshoot).
|
||||
self._lock.release()
|
||||
try:
|
||||
self._stop_event.wait(remaining)
|
||||
finally:
|
||||
self._lock.acquire()
|
||||
if self._stop_event.is_set():
|
||||
break
|
||||
# Transition to HALF_OPEN for the probe. The background
|
||||
# probe itself is the single HALF_OPEN request — keep
|
||||
# _half_open_permit False so concurrent user requests
|
||||
# are blocked until the probe completes.
|
||||
self._state = CircuitState.HALF_OPEN
|
||||
self._half_open_permit = False
|
||||
self._last_state_change = time.monotonic()
|
||||
log.info("Circuit breaker HALF_OPEN: cooldown elapsed, probing")
|
||||
self._update_metrics()
|
||||
success = self._probe_once()
|
||||
if success:
|
||||
self.record_success()
|
||||
|
||||
@@ -9,7 +9,6 @@ from __future__ import annotations
|
||||
|
||||
import fnmatch
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
@@ -20,12 +19,14 @@ from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from turnstone.core.providers._protocol import LLMProvider
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Data structures
|
||||
|
||||
@@ -149,6 +149,24 @@ def configure_logging(
|
||||
logging.getLogger(name).setLevel(logging.WARNING)
|
||||
|
||||
|
||||
def _ensure_stdlib_factory() -> None:
|
||||
"""Ensure structlog routes through stdlib even before configure_logging().
|
||||
|
||||
Without this, ``structlog.get_logger()`` defaults to ``PrintLogger``
|
||||
which bypasses stdlib handlers (and pytest caplog). Calling
|
||||
``configure_logging()`` later overwrites this minimal config.
|
||||
"""
|
||||
cfg = structlog.get_config()
|
||||
if not isinstance(cfg.get("logger_factory"), structlog.stdlib.LoggerFactory):
|
||||
structlog.configure(
|
||||
logger_factory=structlog.stdlib.LoggerFactory(),
|
||||
wrapper_class=structlog.stdlib.BoundLogger,
|
||||
)
|
||||
|
||||
|
||||
_ensure_stdlib_factory()
|
||||
|
||||
|
||||
def get_logger(name: str) -> structlog.stdlib.BoundLogger:
|
||||
"""Return a structlog bound logger backed by the stdlib."""
|
||||
result: structlog.stdlib.BoundLogger = structlog.get_logger(name)
|
||||
|
||||
@@ -22,7 +22,6 @@ import asyncio
|
||||
import concurrent.futures
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
import threading
|
||||
@@ -41,8 +40,9 @@ from mcp.client.stdio import stdio_client
|
||||
from mcp.client.streamable_http import streamablehttp_client
|
||||
|
||||
from turnstone.core.config import load_config
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
log = logging.getLogger("turnstone.mcp")
|
||||
log = get_logger("turnstone.mcp")
|
||||
|
||||
_DEFAULT_REFRESH_INTERVAL: float = 14400 # 4 hours
|
||||
|
||||
@@ -162,9 +162,11 @@ class MCPClientManager:
|
||||
future = asyncio.run_coroutine_threadsafe(self._connect_all(), self._loop)
|
||||
self._connected.wait(timeout=30)
|
||||
# Surface any exception from _connect_all (unlikely — per-server errors are caught)
|
||||
if future.done() and future.exception():
|
||||
self._error = str(future.exception())
|
||||
log.error("MCP initialization error: %s", self._error)
|
||||
if future.done() and not future.cancelled():
|
||||
exc = future.exception()
|
||||
if exc:
|
||||
self._error = str(exc)
|
||||
log.error("MCP initialization error: %s", self._error)
|
||||
|
||||
async def _connect_all(self) -> None:
|
||||
"""Connect to every configured server (runs on the background loop)."""
|
||||
@@ -224,7 +226,9 @@ class MCPClientManager:
|
||||
log.warning("MCP server '%s' has no command configured", name)
|
||||
await stack.aclose()
|
||||
return
|
||||
env = {**os.environ, **cfg.get("env", {})}
|
||||
from turnstone.core.env import scrubbed_env
|
||||
|
||||
env = scrubbed_env(extra=cfg.get("env", {}))
|
||||
params = StdioServerParameters(
|
||||
command=command,
|
||||
args=cfg.get("args", []),
|
||||
@@ -1415,7 +1419,7 @@ def create_mcp_client(
|
||||
if rows:
|
||||
db_names = {r["name"] for r in rows}
|
||||
except Exception:
|
||||
pass
|
||||
log.warning("Failed to load DB-managed MCP servers", exc_info=True)
|
||||
|
||||
servers = load_mcp_config(config_path, storage=storage)
|
||||
if not servers:
|
||||
|
||||
@@ -11,6 +11,7 @@ from __future__ import annotations
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -400,6 +401,17 @@ def resolve_install_config(
|
||||
raise MCPRegistryError(f"Required URL variable '{var_name}' not provided")
|
||||
url = url.replace(placeholder, value)
|
||||
|
||||
# Validate URL after substitution to prevent SSRF-style redirection
|
||||
parsed = urlparse(url)
|
||||
if parsed.scheme not in ("http", "https"):
|
||||
raise MCPRegistryError(
|
||||
f"Invalid URL scheme '{parsed.scheme}' after variable substitution"
|
||||
)
|
||||
if not parsed.hostname:
|
||||
raise MCPRegistryError("Invalid URL (hostname is missing) after variable substitution")
|
||||
if parsed.username is not None or parsed.password is not None:
|
||||
raise MCPRegistryError("URLs with embedded credentials are not allowed in MCP remotes")
|
||||
|
||||
# Build headers dict (required keys only — values provided by user at install time)
|
||||
headers: dict[str, str] = {}
|
||||
for h in remote.headers:
|
||||
|
||||
@@ -3,20 +3,26 @@
|
||||
All functions maintain their existing signatures for consumers (session.py,
|
||||
server.py, cli.py). The actual storage implementation lives in
|
||||
``turnstone.core.storage``.
|
||||
|
||||
The no-raise contract is preserved — callers never see exceptions from this
|
||||
module. All failures are logged so storage issues are visible in logs
|
||||
rather than silently swallowed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.storage import get_storage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
|
||||
def normalize_key(key: str) -> str:
|
||||
"""Normalize a memory key for consistent lookup."""
|
||||
@@ -37,7 +43,7 @@ def save_message(
|
||||
tool_calls: str | None = None,
|
||||
) -> None:
|
||||
"""Log a message to the conversations table."""
|
||||
with contextlib.suppress(Exception):
|
||||
try:
|
||||
get_storage().save_message(
|
||||
ws_id,
|
||||
role,
|
||||
@@ -48,6 +54,8 @@ def save_message(
|
||||
provider_data,
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
except Exception:
|
||||
log.warning("Failed to save message for ws=%s role=%s", ws_id, role, exc_info=True)
|
||||
|
||||
|
||||
def load_messages(ws_id: str) -> list[dict[str, Any]]:
|
||||
@@ -55,6 +63,7 @@ def load_messages(ws_id: str) -> list[dict[str, Any]]:
|
||||
try:
|
||||
return get_storage().load_messages(ws_id)
|
||||
except Exception:
|
||||
log.warning("Failed to load messages for ws=%s", ws_id, exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
@@ -70,22 +79,28 @@ def register_workstream(
|
||||
skill_version: int = 0,
|
||||
) -> None:
|
||||
"""Persist a new workstream (no-op if already exists)."""
|
||||
with contextlib.suppress(Exception):
|
||||
try:
|
||||
get_storage().register_workstream(
|
||||
ws_id, node_id, name, state, skill_id=skill_id, skill_version=skill_version
|
||||
)
|
||||
except Exception:
|
||||
log.warning("Failed to register workstream ws=%s", ws_id, exc_info=True)
|
||||
|
||||
|
||||
def update_workstream_state(ws_id: str, state: str) -> None:
|
||||
"""Update a workstream's state."""
|
||||
with contextlib.suppress(Exception):
|
||||
try:
|
||||
get_storage().update_workstream_state(ws_id, state)
|
||||
except Exception:
|
||||
log.warning("Failed to update workstream state ws=%s state=%s", ws_id, state, exc_info=True)
|
||||
|
||||
|
||||
def update_workstream_name(ws_id: str, name: str) -> None:
|
||||
"""Update a workstream's display name."""
|
||||
with contextlib.suppress(Exception):
|
||||
try:
|
||||
get_storage().update_workstream_name(ws_id, name)
|
||||
except Exception:
|
||||
log.warning("Failed to update workstream name ws=%s", ws_id, exc_info=True)
|
||||
|
||||
|
||||
def list_workstreams(node_id: str | None = None, limit: int = 100) -> list[Any]:
|
||||
@@ -93,6 +108,7 @@ def list_workstreams(node_id: str | None = None, limit: int = 100) -> list[Any]:
|
||||
try:
|
||||
return get_storage().list_workstreams(node_id, limit)
|
||||
except Exception:
|
||||
log.warning("Failed to list workstreams", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
@@ -101,6 +117,7 @@ def list_workstreams_with_history(limit: int = 20) -> list[Any]:
|
||||
try:
|
||||
return get_storage().list_workstreams_with_history(limit)
|
||||
except Exception:
|
||||
log.warning("Failed to list workstreams with history", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
@@ -109,6 +126,7 @@ def delete_workstream(ws_id: str) -> bool:
|
||||
try:
|
||||
return get_storage().delete_workstream(ws_id)
|
||||
except Exception:
|
||||
log.warning("Failed to delete workstream ws=%s", ws_id, exc_info=True)
|
||||
return False
|
||||
|
||||
|
||||
@@ -120,6 +138,7 @@ def prune_workstreams(
|
||||
try:
|
||||
orphans, stale = get_storage().prune_workstreams(retention_days)
|
||||
except Exception:
|
||||
log.warning("Failed to prune workstreams", exc_info=True)
|
||||
return (0, 0)
|
||||
|
||||
if log_fn and (orphans or stale):
|
||||
@@ -140,6 +159,7 @@ def resolve_workstream(alias_or_id: str) -> str | None:
|
||||
try:
|
||||
return get_storage().resolve_workstream(alias_or_id)
|
||||
except Exception:
|
||||
log.warning("Failed to resolve workstream alias=%s", alias_or_id, exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
@@ -148,8 +168,10 @@ def resolve_workstream(alias_or_id: str) -> str | None:
|
||||
|
||||
def save_workstream_config(ws_id: str, config: dict[str, str]) -> None:
|
||||
"""Persist workstream configuration key/value pairs."""
|
||||
with contextlib.suppress(Exception):
|
||||
try:
|
||||
get_storage().save_workstream_config(ws_id, config)
|
||||
except Exception:
|
||||
log.warning("Failed to save workstream config ws=%s", ws_id, exc_info=True)
|
||||
|
||||
|
||||
def load_workstream_config(ws_id: str) -> dict[str, str]:
|
||||
@@ -157,6 +179,7 @@ def load_workstream_config(ws_id: str) -> dict[str, str]:
|
||||
try:
|
||||
return get_storage().load_workstream_config(ws_id)
|
||||
except Exception:
|
||||
log.warning("Failed to load workstream config ws=%s", ws_id, exc_info=True)
|
||||
return {}
|
||||
|
||||
|
||||
@@ -168,6 +191,7 @@ def get_skill_by_name(name: str) -> dict[str, Any] | None:
|
||||
try:
|
||||
return get_storage().get_prompt_template_by_name(name)
|
||||
except Exception:
|
||||
log.warning("Failed to get skill name=%s", name, exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
@@ -176,6 +200,7 @@ def list_default_skills(org_id: str = "") -> list[dict[str, Any]]:
|
||||
try:
|
||||
return get_storage().list_default_templates(org_id)
|
||||
except Exception:
|
||||
log.warning("Failed to list default skills", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
@@ -191,6 +216,7 @@ def list_skills_by_activation(
|
||||
activation, enabled_only=enabled_only, limit=limit
|
||||
)
|
||||
except Exception:
|
||||
log.warning("Failed to list skills by activation=%s", activation, exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
@@ -202,6 +228,7 @@ def set_workstream_alias(ws_id: str, alias: str) -> bool:
|
||||
try:
|
||||
return get_storage().set_workstream_alias(ws_id, alias)
|
||||
except Exception:
|
||||
log.warning("Failed to set alias ws=%s alias=%s", ws_id, alias, exc_info=True)
|
||||
return False
|
||||
|
||||
|
||||
@@ -210,13 +237,16 @@ def get_workstream_display_name(ws_id: str) -> str | None:
|
||||
try:
|
||||
return get_storage().get_workstream_display_name(ws_id)
|
||||
except Exception:
|
||||
log.warning("Failed to get display name ws=%s", ws_id, exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
def update_workstream_title(ws_id: str, title: str) -> None:
|
||||
"""Set or update the auto-generated title for a workstream."""
|
||||
with contextlib.suppress(Exception):
|
||||
try:
|
||||
get_storage().update_workstream_title(ws_id, title)
|
||||
except Exception:
|
||||
log.warning("Failed to update title ws=%s", ws_id, exc_info=True)
|
||||
|
||||
|
||||
# -- Conversation search -------------------------------------------------------
|
||||
@@ -227,6 +257,7 @@ def search_history(query: str, limit: int = 20) -> list[Any]:
|
||||
try:
|
||||
return get_storage().search_history(query, limit)
|
||||
except Exception:
|
||||
log.warning("Failed to search history", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
@@ -235,6 +266,7 @@ def search_history_recent(limit: int = 20) -> list[Any]:
|
||||
try:
|
||||
return get_storage().search_history_recent(limit)
|
||||
except Exception:
|
||||
log.warning("Failed to search recent history", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
@@ -280,6 +312,7 @@ def save_structured_memory(
|
||||
return existing["memory_id"], old_content
|
||||
return "", None
|
||||
except Exception:
|
||||
log.warning("Failed to save structured memory name=%s", name, exc_info=True)
|
||||
return "", None
|
||||
|
||||
|
||||
@@ -289,6 +322,7 @@ def delete_structured_memory(name: str, scope: str = "global", scope_id: str = "
|
||||
try:
|
||||
return get_storage().delete_structured_memory(name, scope, scope_id)
|
||||
except Exception:
|
||||
log.warning("Failed to delete structured memory name=%s", name, exc_info=True)
|
||||
return False
|
||||
|
||||
|
||||
@@ -297,6 +331,7 @@ def delete_structured_memory_by_id(memory_id: str) -> bool:
|
||||
try:
|
||||
return get_storage().delete_structured_memory_by_id(memory_id)
|
||||
except Exception:
|
||||
log.warning("Failed to delete structured memory id=%s", memory_id, exc_info=True)
|
||||
return False
|
||||
|
||||
|
||||
@@ -312,6 +347,7 @@ def list_structured_memories(
|
||||
mem_type=mem_type, scope=scope, scope_id=scope_id, limit=limit
|
||||
)
|
||||
except Exception:
|
||||
log.warning("Failed to list structured memories", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
@@ -328,9 +364,31 @@ def search_structured_memories(
|
||||
query, mem_type=mem_type, scope=scope, scope_id=scope_id, limit=limit
|
||||
)
|
||||
except Exception:
|
||||
log.warning("Failed to search structured memories", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
def touch_structured_memories(keys: list[tuple[str, str, str]]) -> int:
|
||||
"""Batch-touch memories (bump last_accessed, increment access_count).
|
||||
|
||||
Each key is ``(name, scope, scope_id)``. Duplicates are removed so each
|
||||
distinct memory is touched at most once. Returns count of rows updated.
|
||||
"""
|
||||
if not keys:
|
||||
return 0
|
||||
seen: set[tuple[str, str, str]] = set()
|
||||
unique: list[tuple[str, str, str]] = []
|
||||
for k in keys:
|
||||
if k not in seen:
|
||||
seen.add(k)
|
||||
unique.append(k)
|
||||
try:
|
||||
return get_storage().touch_structured_memories(unique)
|
||||
except Exception:
|
||||
log.warning("Failed to touch structured memories", exc_info=True)
|
||||
return 0
|
||||
|
||||
|
||||
def count_structured_memories(mem_type: str = "", scope: str = "", scope_id: str = "") -> int:
|
||||
"""Count structured memories with optional type/scope filter."""
|
||||
try:
|
||||
@@ -338,4 +396,5 @@ def count_structured_memories(mem_type: str = "", scope: str = "", scope_id: str
|
||||
mem_type=mem_type, scope=scope, scope_id=scope_id
|
||||
)
|
||||
except Exception:
|
||||
log.warning("Failed to count structured memories", exc_info=True)
|
||||
return 0
|
||||
|
||||
@@ -7,15 +7,15 @@ resilience when the primary model is unreachable.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from turnstone.core.config import load_config
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.providers import LLMProvider, create_client, create_provider
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -278,7 +278,9 @@ def detect_model(
|
||||
client: Any,
|
||||
log_fn: Any = print,
|
||||
provider: str = "openai",
|
||||
) -> tuple[str, int | None]:
|
||||
*,
|
||||
fatal: bool = True,
|
||||
) -> tuple[str | None, int | None]:
|
||||
"""Auto-detect the model and context window from the API's models endpoint.
|
||||
|
||||
Returns ``(model_id, context_window)`` where *context_window* is
|
||||
@@ -289,13 +291,25 @@ def detect_model(
|
||||
For local single-model servers (vLLM, llama.cpp), uses the first model.
|
||||
|
||||
Calls ``log_fn`` for informational messages (defaults to ``print``).
|
||||
Raises ``SystemExit`` on failure.
|
||||
|
||||
When *fatal* is ``True`` (default), raises ``SystemExit`` on failure.
|
||||
When ``False``, returns ``(None, None)`` so the server can start in
|
||||
degraded mode (useful for cluster deployments where the LLM backend
|
||||
may not be available at startup).
|
||||
"""
|
||||
try:
|
||||
models = client.models.list()
|
||||
# Use a short timeout for startup detection — the default OpenAI client
|
||||
# timeout is 600s read which blocks the main thread for minutes when the
|
||||
# backend is unreachable (TCP SYN dropped → kernel retransmit timeout).
|
||||
# Disable retries (default 2) to avoid compounding the delay.
|
||||
fast_client = client.with_options(timeout=10.0, max_retries=0)
|
||||
models = fast_client.models.list()
|
||||
if not models.data:
|
||||
log_fn("Error: No models found at server. Use --model to specify.")
|
||||
raise SystemExit(1)
|
||||
if fatal:
|
||||
log_fn("Error: No models found at server. Use --model to specify.")
|
||||
raise SystemExit(1)
|
||||
log_fn("Warning: No models found at server — starting in degraded mode.")
|
||||
return None, None
|
||||
|
||||
all_ids = [x.id for x in models.data]
|
||||
selected_id = _select_best_model(all_ids, provider)
|
||||
@@ -321,6 +335,10 @@ def detect_model(
|
||||
except SystemExit:
|
||||
raise
|
||||
except Exception as e:
|
||||
log_fn(f"Error: Could not connect to server: {e}")
|
||||
log_fn("Is the model server running? Start it or use --base-url to point elsewhere.")
|
||||
raise SystemExit(1) from e
|
||||
if fatal:
|
||||
log_fn(f"Error: Could not connect to server: {e}")
|
||||
log_fn("Is the model server running? Start it or use --base-url to point elsewhere.")
|
||||
raise SystemExit(1) from e
|
||||
log_fn(f"Warning: Could not connect to LLM backend: {e}")
|
||||
log_fn("Starting in degraded mode — requests will fail until backend is reachable.")
|
||||
return None, None
|
||||
|
||||
+66
-2
@@ -10,10 +10,11 @@ from __future__ import annotations
|
||||
import base64
|
||||
import dataclasses
|
||||
import hashlib
|
||||
import logging
|
||||
import ipaddress
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import socket
|
||||
import urllib.parse
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
@@ -21,7 +22,9 @@ from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
# Sentinel password hash for OIDC-provisioned users.
|
||||
# Not a valid bcrypt hash -- verify_password() always rejects it.
|
||||
@@ -212,6 +215,61 @@ def load_oidc_config() -> OIDCConfig:
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SSRF validation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _is_localhost(hostname: str) -> bool:
|
||||
"""Return True if *hostname* refers to the loopback interface."""
|
||||
return hostname in ("localhost", "127.0.0.1", "::1") or hostname.endswith(".localhost")
|
||||
|
||||
|
||||
def validate_issuer_url(url: str) -> None:
|
||||
"""Validate an OIDC issuer URL to prevent SSRF.
|
||||
|
||||
Rejects:
|
||||
- Non-HTTPS URLs (except localhost for development)
|
||||
- URLs with embedded credentials (userinfo)
|
||||
- Hostnames that resolve to private/internal/loopback IP addresses
|
||||
|
||||
Raises :class:`OIDCError` on validation failure.
|
||||
"""
|
||||
parsed = urllib.parse.urlparse(url)
|
||||
|
||||
# Require a hostname.
|
||||
hostname = parsed.hostname
|
||||
if not hostname:
|
||||
raise OIDCError(f"OIDC issuer URL has no hostname: {url}")
|
||||
|
||||
# Reject embedded credentials — redact userinfo from error message.
|
||||
if parsed.username or parsed.password:
|
||||
raise OIDCError("OIDC issuer URL must not contain embedded credentials (userinfo)")
|
||||
|
||||
# Require HTTPS (allow HTTP only for localhost development).
|
||||
if parsed.scheme != "https":
|
||||
if parsed.scheme == "http" and _is_localhost(hostname):
|
||||
pass # Allow http://localhost for dev
|
||||
else:
|
||||
raise OIDCError(f"OIDC issuer URL must use HTTPS (got {parsed.scheme}://): {url}")
|
||||
|
||||
# Resolve hostname and reject non-globally-routable addresses.
|
||||
try:
|
||||
addr_infos = socket.getaddrinfo(hostname, None, proto=socket.IPPROTO_TCP)
|
||||
except socket.gaierror as exc:
|
||||
raise OIDCError(f"OIDC issuer hostname cannot be resolved: {hostname}") from exc
|
||||
|
||||
for _family, _type, _proto, _canonname, sockaddr in addr_infos:
|
||||
try:
|
||||
addr = ipaddress.ip_address(sockaddr[0])
|
||||
except ValueError as exc:
|
||||
raise OIDCError(
|
||||
f"OIDC issuer hostname resolved to invalid IP {sockaddr[0]!r}: {hostname}"
|
||||
) from exc
|
||||
if not addr.is_global and not _is_localhost(hostname):
|
||||
raise OIDCError(f"OIDC issuer URL resolves to non-public address ({addr}): {url}")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Discovery
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -225,6 +283,12 @@ async def discover_oidc(config: OIDCConfig) -> OIDCConfig:
|
||||
if not config.issuer:
|
||||
return dataclasses.replace(config, enabled=False)
|
||||
|
||||
try:
|
||||
validate_issuer_url(config.issuer)
|
||||
except OIDCError as exc:
|
||||
log.warning("OIDC issuer URL rejected: %s", exc)
|
||||
return dataclasses.replace(config, enabled=False)
|
||||
|
||||
url = config.issuer.rstrip("/") + "/.well-known/openid-configuration"
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
|
||||
@@ -7,13 +7,14 @@ a tool should be auto-allowed, denied, or require human approval.
|
||||
from __future__ import annotations
|
||||
|
||||
import fnmatch
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from turnstone.core.storage._protocol import StorageBackend
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
|
||||
def evaluate_tool_policy(
|
||||
|
||||
@@ -7,11 +7,12 @@ Thread-safe. Zero external dependencies.
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
_NetworkType = ipaddress.IPv4Network | ipaddress.IPv6Network
|
||||
|
||||
|
||||
+177
-90
@@ -90,6 +90,7 @@ log = get_logger(__name__)
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Iterator
|
||||
|
||||
from turnstone.core.config_store import ConfigStore
|
||||
from turnstone.core.healthcheck import BackendHealthMonitor
|
||||
from turnstone.core.judge import IntentJudge, JudgeConfig
|
||||
from turnstone.core.mcp_client import MCPClientManager
|
||||
@@ -100,6 +101,7 @@ if TYPE_CHECKING:
|
||||
ModelCapabilities,
|
||||
StreamChunk,
|
||||
)
|
||||
from turnstone.core.web_search import WebSearchClient
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cancellation support
|
||||
@@ -247,6 +249,8 @@ class ChatSession:
|
||||
judge_config: JudgeConfig | None = None,
|
||||
user_id: str = "",
|
||||
memory_config: MemoryConfig | None = None,
|
||||
config_store: ConfigStore | None = None,
|
||||
web_search_backend: str = "",
|
||||
):
|
||||
self.client = client
|
||||
self.model = model
|
||||
@@ -259,6 +263,7 @@ class ChatSession:
|
||||
if registry and model_alias
|
||||
else create_provider("openai")
|
||||
)
|
||||
self._cached_capabilities: ModelCapabilities | None = None
|
||||
self.ui = ui
|
||||
self.instructions = instructions
|
||||
self.temperature = temperature
|
||||
@@ -281,6 +286,7 @@ class ChatSession:
|
||||
self.auto_approve = False
|
||||
self._node_id = node_id
|
||||
self._user_id = user_id
|
||||
self._config_store = config_store
|
||||
self._memory_config = memory_config or MemoryConfig()
|
||||
self._ws_id = ws_id or uuid.uuid4().hex
|
||||
self._title_generated = False
|
||||
@@ -302,7 +308,7 @@ class ChatSession:
|
||||
self._notify_count = 0
|
||||
# Watch support: server-level runner injected via set_watch_runner()
|
||||
self._watch_runner: Any = None # WatchRunner | None
|
||||
self._watch_pending: queue.Queue[dict[str, Any]] = queue.Queue()
|
||||
self._watch_pending: queue.Queue[dict[str, Any]] = queue.Queue(maxsize=20)
|
||||
self._watch_dispatch_depth = 0
|
||||
# Metacognitive nudges: ephemeral prompts for proactive memory use
|
||||
self._metacog_state: dict[str, float] = {}
|
||||
@@ -336,6 +342,8 @@ class ChatSession:
|
||||
self._tools = TOOLS
|
||||
self._task_tools = TASK_AGENT_TOOLS
|
||||
self._agent_tools = AGENT_TOOLS
|
||||
# Web search backend (pluggable: auto/tavily/ddg/mcp:server:tool)
|
||||
self._web_search_backend = web_search_backend
|
||||
# Dynamic tool search: defer MCP tools when tool count is high
|
||||
self._tool_search_setting = tool_search
|
||||
self._tool_search_threshold = tool_search_threshold
|
||||
@@ -365,6 +373,70 @@ class ChatSession:
|
||||
def model_alias(self) -> str | None:
|
||||
return self._model_alias
|
||||
|
||||
@property
|
||||
def _mem_cfg(self) -> MemoryConfig:
|
||||
"""Live memory config — reads from ConfigStore when available."""
|
||||
cs = getattr(self, "_config_store", None)
|
||||
if cs is None:
|
||||
return self._memory_config
|
||||
return MemoryConfig(
|
||||
relevance_k=cs.get("memory.relevance_k"),
|
||||
fetch_limit=cs.get("memory.fetch_limit"),
|
||||
max_content=cs.get("memory.max_content"),
|
||||
nudge_cooldown=cs.get("memory.nudge_cooldown"),
|
||||
nudges=cs.get("memory.nudges"),
|
||||
)
|
||||
|
||||
@property
|
||||
def _judge_cfg(self) -> JudgeConfig | None:
|
||||
"""Live judge behavioral config — reads from ConfigStore when available.
|
||||
|
||||
LLM client fields (model, provider, base_url, api_key) stay frozen
|
||||
from session creation time since changing them would require tearing
|
||||
down and rebuilding the IntentJudge instance.
|
||||
"""
|
||||
jc = self._judge_config
|
||||
if jc is None:
|
||||
return None
|
||||
cs = getattr(self, "_config_store", None)
|
||||
if cs is None:
|
||||
return jc
|
||||
from turnstone.core.judge import JudgeConfig
|
||||
|
||||
return JudgeConfig(
|
||||
enabled=cs.get("judge.enabled"),
|
||||
model=jc.model,
|
||||
provider=jc.provider,
|
||||
base_url=jc.base_url,
|
||||
api_key=jc.api_key,
|
||||
confidence_threshold=cs.get("judge.confidence_threshold"),
|
||||
max_context_ratio=cs.get("judge.max_context_ratio"),
|
||||
timeout=cs.get("judge.timeout"),
|
||||
read_only_tools=cs.get("judge.read_only_tools"),
|
||||
output_guard=cs.get("judge.output_guard"),
|
||||
redact_secrets=cs.get("judge.redact_secrets"),
|
||||
)
|
||||
|
||||
def _get_web_search_backend(self) -> str:
|
||||
"""Effective web search backend — reads from ConfigStore when available."""
|
||||
cs = getattr(self, "_config_store", None)
|
||||
if cs is not None:
|
||||
val = cs.get("tools.web_search_backend")
|
||||
if val:
|
||||
return str(val)
|
||||
return self._web_search_backend
|
||||
|
||||
def _resolve_search_client(self) -> WebSearchClient | None:
|
||||
"""Return a web search client for the configured backend, or None."""
|
||||
from turnstone.core.web_search import resolve_web_search_client
|
||||
|
||||
return resolve_web_search_client(
|
||||
backend=self._get_web_search_backend(),
|
||||
tavily_key=get_tavily_key(),
|
||||
mcp_client=self._mcp_client,
|
||||
timeout=self.tool_timeout,
|
||||
)
|
||||
|
||||
def _resolve_capabilities(
|
||||
self,
|
||||
provider: LLMProvider,
|
||||
@@ -382,9 +454,16 @@ class ChatSession:
|
||||
caps = dataclasses.replace(caps, **overrides)
|
||||
return caps
|
||||
|
||||
def _get_capabilities(self) -> ModelCapabilities:
|
||||
"""Get capabilities for the current model."""
|
||||
return self._resolve_capabilities(self._provider, self.model, self._model_alias)
|
||||
def _get_capabilities(self, provider: Any = None, model: str = "") -> ModelCapabilities:
|
||||
"""Get capabilities for a model. Cached for the primary session model."""
|
||||
p = provider or self._provider
|
||||
m = model or self.model
|
||||
# Only use cache for the primary session model — fallback models bypass.
|
||||
if p is self._provider and m == self.model:
|
||||
if self._cached_capabilities is None:
|
||||
self._cached_capabilities = self._resolve_capabilities(p, m, self._model_alias)
|
||||
return self._cached_capabilities
|
||||
return self._resolve_capabilities(p, m, "")
|
||||
|
||||
def _save_config(self) -> None:
|
||||
"""Persist LLM-affecting config so resumed workstreams behave identically."""
|
||||
@@ -549,7 +628,12 @@ class ChatSession:
|
||||
pending = self._watch_pending
|
||||
|
||||
def _enqueue(msg: str) -> None:
|
||||
pending.put({"message": msg})
|
||||
try:
|
||||
pending.put_nowait({"message": msg})
|
||||
except queue.Full:
|
||||
log.warning(
|
||||
"Watch pending queue full, dropping result for ws_id=%s", self._ws_id
|
||||
)
|
||||
|
||||
runner.set_dispatch_fn(self._ws_id, _enqueue)
|
||||
|
||||
@@ -623,15 +707,20 @@ class ChatSession:
|
||||
|
||||
def _generate_title(self) -> None:
|
||||
"""Generate a short title for this session via a background LLM call."""
|
||||
ws_id = self._ws_id # Capture before async work
|
||||
try:
|
||||
# Gather first user message and first assistant reply
|
||||
user_msg = ""
|
||||
asst_msg = ""
|
||||
for m in self.messages:
|
||||
content = m.get("content") or ""
|
||||
# Handle multi-part content (vision messages)
|
||||
if isinstance(content, list):
|
||||
content = " ".join(p.get("text", "") for p in content if isinstance(p, dict))
|
||||
if m["role"] == "user" and not user_msg:
|
||||
user_msg = (m.get("content") or "")[:300]
|
||||
user_msg = content[:300]
|
||||
elif m["role"] == "assistant" and not asst_msg:
|
||||
asst_msg = (m.get("content") or "")[:200]
|
||||
asst_msg = content[:200]
|
||||
if user_msg and asst_msg:
|
||||
break
|
||||
if not user_msg:
|
||||
@@ -665,10 +754,15 @@ class ChatSession:
|
||||
raw = (result.content or "").strip()
|
||||
# Take first line, strip quotes
|
||||
title = raw.split("\n")[0].strip().strip('"').strip("'")
|
||||
if title:
|
||||
update_workstream_title(self._ws_id, title[:80])
|
||||
if title and self._ws_id == ws_id:
|
||||
update_workstream_title(ws_id, title[:80])
|
||||
self.ui.on_rename(title[:80])
|
||||
except Exception:
|
||||
pass # Title generation is non-critical
|
||||
# Only reset if ws_id hasn't changed (e.g., via /resume) to
|
||||
# avoid re-enabling titling for a different workstream.
|
||||
if self._ws_id == ws_id:
|
||||
self._title_generated = False
|
||||
log.debug("Title generation failed for ws=%s", ws_id, exc_info=True)
|
||||
|
||||
def resume(self, ws_id: str) -> bool:
|
||||
"""Load messages from a previous workstream and resume it.
|
||||
@@ -719,12 +813,12 @@ class ChatSession:
|
||||
self._skill_name = None
|
||||
if "notify_on_complete" in config:
|
||||
self._notify_on_complete = config["notify_on_complete"]
|
||||
if self._memory_config.nudges and should_nudge(
|
||||
if self._mem_cfg.nudges and should_nudge(
|
||||
"resume",
|
||||
self._metacog_state,
|
||||
message_count=len(self.messages),
|
||||
memory_count=self._visible_memory_count(),
|
||||
cooldown_secs=self._memory_config.nudge_cooldown,
|
||||
cooldown_secs=self._mem_cfg.nudge_cooldown,
|
||||
):
|
||||
self._pending_nudge.append(format_nudge("resume"))
|
||||
self._init_system_messages()
|
||||
@@ -894,10 +988,10 @@ class ChatSession:
|
||||
if self.instructions:
|
||||
dev_parts.append("")
|
||||
dev_parts.append(self.instructions)
|
||||
visible_mems = self._get_visible_memories(limit=self._memory_config.fetch_limit)
|
||||
visible_mems = self._get_visible_memories(limit=self._mem_cfg.fetch_limit)
|
||||
if visible_mems:
|
||||
context = extract_recent_context(self.messages)
|
||||
relevant = score_memories(visible_mems, context, k=self._memory_config.relevance_k)
|
||||
relevant = score_memories(visible_mems, context, k=self._mem_cfg.relevance_k)
|
||||
if relevant:
|
||||
dev_parts.append("")
|
||||
dev_parts.append(build_memory_context(relevant))
|
||||
@@ -952,7 +1046,8 @@ class ChatSession:
|
||||
Without tool search: return self._tools unchanged.
|
||||
|
||||
Web search gating: ``web_search`` is removed when the model has
|
||||
no native search support and no Tavily API key is configured.
|
||||
no native search support and no search backend is available
|
||||
(Tavily, DDG, or MCP — see ``_resolve_search_client``).
|
||||
"""
|
||||
if self.creative_mode:
|
||||
return None
|
||||
@@ -969,7 +1064,7 @@ class ChatSession:
|
||||
tools = visible + [self._tool_search.get_search_tool_definition()]
|
||||
|
||||
# Gate web_search: only include when a backend exists
|
||||
if not caps.supports_web_search and not get_tavily_key():
|
||||
if not caps.supports_web_search and not self._resolve_search_client():
|
||||
tools = _without_tool(tools, "web_search")
|
||||
|
||||
return tools
|
||||
@@ -1196,11 +1291,7 @@ class ChatSession:
|
||||
_tc_names = {c["id"]: c.get("function", {}).get("name", "") for c in tool_calls}
|
||||
for tc_id, output in results:
|
||||
# Output guard: evaluate tool result before it enters context
|
||||
if (
|
||||
self._judge_config
|
||||
and self._judge_config.enabled
|
||||
and self._judge_config.output_guard
|
||||
):
|
||||
if self._judge_cfg and self._judge_cfg.output_guard:
|
||||
if isinstance(output, str):
|
||||
output = self._evaluate_output(tc_id, output, _tc_names.get(tc_id, ""))
|
||||
elif isinstance(output, list):
|
||||
@@ -1259,7 +1350,7 @@ class ChatSession:
|
||||
)
|
||||
# Metacognitive nudge: check memories on tool error
|
||||
if (
|
||||
self._memory_config.nudges
|
||||
self._mem_cfg.nudges
|
||||
and any(
|
||||
isinstance(out, str)
|
||||
and (
|
||||
@@ -1275,7 +1366,7 @@ class ChatSession:
|
||||
self._metacog_state,
|
||||
message_count=len(self.messages),
|
||||
memory_count=self._visible_memory_count(),
|
||||
cooldown_secs=self._memory_config.nudge_cooldown,
|
||||
cooldown_secs=self._mem_cfg.nudge_cooldown,
|
||||
)
|
||||
):
|
||||
self._pending_nudge.append(format_nudge("tool_error"))
|
||||
@@ -1858,7 +1949,6 @@ class ChatSession:
|
||||
|
||||
self.ui.on_thinking_start()
|
||||
try:
|
||||
_last_err: Exception | None = None
|
||||
result: CompletionResult | None = None
|
||||
for attempt in range(self._MAX_RETRIES + 1):
|
||||
try:
|
||||
@@ -1879,7 +1969,6 @@ class ChatSession:
|
||||
or attempt == self._MAX_RETRIES
|
||||
):
|
||||
raise
|
||||
_last_err = e
|
||||
delay = self._RETRY_BASE_DELAY * (2**attempt)
|
||||
self.ui.on_info(f"[Compact retrying in {delay:.0f}s: {ename}]")
|
||||
time.sleep(delay)
|
||||
@@ -1931,10 +2020,20 @@ class ChatSession:
|
||||
# -- Intent validation --------------------------------------------------------
|
||||
|
||||
def _ensure_judge(self) -> IntentJudge | None:
|
||||
"""Lazily initialize the intent judge if configured."""
|
||||
"""Lazily initialize the intent judge if configured.
|
||||
|
||||
Re-checks the live ``enabled`` flag every call so disabling the
|
||||
judge via admin settings takes immediate effect on existing sessions.
|
||||
"""
|
||||
if not self._judge_cfg or not self._judge_cfg.enabled:
|
||||
return None
|
||||
if self._judge is not None:
|
||||
return self._judge
|
||||
if not self._judge_config or not self._judge_config.enabled:
|
||||
return None
|
||||
# Frozen config required for IntentJudge init (LLM client fields).
|
||||
# _judge_cfg already returns None when _judge_config is None, but
|
||||
# this guard makes the dependency explicit for type narrowing.
|
||||
if self._judge_config is None:
|
||||
return None
|
||||
try:
|
||||
from turnstone.core.judge import IntentJudge
|
||||
@@ -2041,11 +2140,7 @@ class ChatSession:
|
||||
except Exception:
|
||||
log.debug("output_guard.callback_failed", exc_info=True)
|
||||
|
||||
if (
|
||||
assessment.sanitized is not None
|
||||
and self._judge_config
|
||||
and self._judge_config.redact_secrets
|
||||
):
|
||||
if assessment.sanitized is not None and self._judge_cfg and self._judge_cfg.redact_secrets:
|
||||
return assessment.sanitized
|
||||
return output
|
||||
|
||||
@@ -2082,12 +2177,12 @@ class ChatSession:
|
||||
f"Denied by user: {user_feedback}" if user_feedback else "Denied by user"
|
||||
)
|
||||
user_feedback = None # feedback is in the denial_msg
|
||||
if self._memory_config.nudges and should_nudge(
|
||||
if self._mem_cfg.nudges and should_nudge(
|
||||
"denial",
|
||||
self._metacog_state,
|
||||
message_count=len(self.messages),
|
||||
memory_count=self._visible_memory_count(),
|
||||
cooldown_secs=self._memory_config.nudge_cooldown,
|
||||
cooldown_secs=self._mem_cfg.nudge_cooldown,
|
||||
):
|
||||
self._pending_nudge.append(format_nudge("denial"))
|
||||
self._init_system_messages()
|
||||
@@ -2699,17 +2794,17 @@ class ChatSession:
|
||||
"needs_approval": False,
|
||||
"error": "Error: no query provided",
|
||||
}
|
||||
if not get_tavily_key():
|
||||
if not self._resolve_search_client():
|
||||
return {
|
||||
"call_id": call_id,
|
||||
"func_name": "web_search",
|
||||
"header": "\u2717 web_search: no API key",
|
||||
"header": "\u2717 web_search: no backend available",
|
||||
"preview": "",
|
||||
"needs_approval": False,
|
||||
"error": (
|
||||
"Error: Tavily API key not configured. "
|
||||
"Set it in ~/.config/turnstone/tavily_key or $TAVILY_API_KEY. "
|
||||
"Use web_fetch with a direct URL as an alternative."
|
||||
"Error: No web search backend available. "
|
||||
"Install duckduckgo-search (`pip install duckduckgo-search`), "
|
||||
"configure a Tavily API key, or set tools.web_search_backend."
|
||||
),
|
||||
}
|
||||
try:
|
||||
@@ -2866,11 +2961,11 @@ class ChatSession:
|
||||
|
||||
def _check_metacognitive_nudge(self, user_message: str) -> str | None:
|
||||
"""Check if a metacognitive nudge should be injected."""
|
||||
if not self._memory_config.nudges:
|
||||
if not self._mem_cfg.nudges:
|
||||
return None
|
||||
mem_count = self._visible_memory_count()
|
||||
msg_count = len(self.messages)
|
||||
cd = self._memory_config.nudge_cooldown
|
||||
cd = self._mem_cfg.nudge_cooldown
|
||||
|
||||
if should_nudge(
|
||||
"start",
|
||||
@@ -2918,14 +3013,14 @@ class ChatSession:
|
||||
"needs_approval": False,
|
||||
"error": "Error: both 'name' and 'content' are required for save",
|
||||
}
|
||||
if len(content) > self._memory_config.max_content:
|
||||
if len(content) > self._mem_cfg.max_content:
|
||||
return {
|
||||
"call_id": call_id,
|
||||
"func_name": "memory",
|
||||
"header": "\u2717 memory save: content too large",
|
||||
"preview": "",
|
||||
"needs_approval": False,
|
||||
"error": f"Error: content exceeds {self._memory_config.max_content} character limit",
|
||||
"error": f"Error: content exceeds {self._mem_cfg.max_content} character limit",
|
||||
}
|
||||
description = (args.get("description") or "").strip()
|
||||
mem_type = (args.get("type") or "project").strip().lower()
|
||||
@@ -3453,12 +3548,15 @@ class ChatSession:
|
||||
f.write(command)
|
||||
script_path = f.name
|
||||
try:
|
||||
from turnstone.core.env import scrubbed_env
|
||||
|
||||
proc = subprocess.Popen(
|
||||
["bash", script_path],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
start_new_session=True,
|
||||
env=scrubbed_env(),
|
||||
)
|
||||
# Drain stderr in background thread to avoid pipe deadlock
|
||||
stderr_lines: list[str] = []
|
||||
@@ -3492,8 +3590,10 @@ class ChatSession:
|
||||
assert proc.stdout is not None
|
||||
for line in proc.stdout:
|
||||
stdout_parts.append(line)
|
||||
with contextlib.suppress(Exception):
|
||||
try:
|
||||
self.ui.on_tool_output_chunk(call_id, line)
|
||||
except Exception:
|
||||
log.debug("UI callback error during tool output", exc_info=True)
|
||||
# Check cancellation during long-running commands
|
||||
if self._cancel_event.is_set():
|
||||
with contextlib.suppress(OSError, ProcessLookupError):
|
||||
@@ -3506,8 +3606,11 @@ class ChatSession:
|
||||
finally:
|
||||
timer.cancel()
|
||||
|
||||
proc.wait()
|
||||
stderr_thread.join()
|
||||
try:
|
||||
proc.wait(timeout=10)
|
||||
except subprocess.TimeoutExpired:
|
||||
log.warning("Process did not exit after SIGKILL, pid=%d", proc.pid)
|
||||
stderr_thread.join(timeout=5)
|
||||
finally:
|
||||
os.unlink(script_path)
|
||||
|
||||
@@ -3642,6 +3745,8 @@ class ChatSession:
|
||||
call_id = item["call_id"]
|
||||
pattern, path = item["pattern"], item["path"]
|
||||
try:
|
||||
from turnstone.core.env import scrubbed_env
|
||||
|
||||
result = subprocess.run(
|
||||
[
|
||||
"grep",
|
||||
@@ -3658,6 +3763,7 @@ class ChatSession:
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=self.tool_timeout,
|
||||
env=scrubbed_env(),
|
||||
)
|
||||
output = result.stdout.strip()
|
||||
if result.returncode == 1:
|
||||
@@ -3687,10 +3793,6 @@ class ChatSession:
|
||||
self.ui.on_error(msg)
|
||||
return call_id, msg
|
||||
|
||||
# Tools the agent can auto-execute without user approval (read-only).
|
||||
_AGENT_AUTO_TOOLS = AGENT_AUTO_TOOLS
|
||||
_TASK_AUTO_TOOLS = TASK_AUTO_TOOLS
|
||||
|
||||
def _run_agent(
|
||||
self,
|
||||
agent_messages: list[dict[str, Any]],
|
||||
@@ -3705,7 +3807,7 @@ class ChatSession:
|
||||
agent_messages: Pre-built message list (system + developer + user).
|
||||
label: Display prefix for progress lines ("agent" or "plan").
|
||||
tools: Tool definitions to send to the API. Defaults to AGENT_TOOLS (read-only).
|
||||
auto_tools: Set of tool names the agent may execute. Defaults to _AGENT_AUTO_TOOLS.
|
||||
auto_tools: Set of tool names the agent may execute. Defaults to AGENT_AUTO_TOOLS.
|
||||
reasoning_effort: Override reasoning effort for this agent.
|
||||
|
||||
Returns:
|
||||
@@ -3714,7 +3816,7 @@ class ChatSession:
|
||||
if tools is None:
|
||||
tools = self._agent_tools
|
||||
if auto_tools is None:
|
||||
auto_tools = self._AGENT_AUTO_TOOLS
|
||||
auto_tools = AGENT_AUTO_TOOLS
|
||||
max_tool_turns = self.agent_max_turns
|
||||
|
||||
# Resolve agent model and provider: use registry.agent_model if configured
|
||||
@@ -3728,7 +3830,7 @@ class ChatSession:
|
||||
# Gate web_search: remove when no backend exists for the agent model
|
||||
agent_alias = self._registry.agent_model if self._registry else None
|
||||
agent_caps = self._resolve_capabilities(agent_provider, agent_model, agent_alias)
|
||||
if not agent_caps.supports_web_search and not get_tavily_key():
|
||||
if not agent_caps.supports_web_search and not self._resolve_search_client():
|
||||
tools = _without_tool(tools, "web_search")
|
||||
|
||||
# Build extra params for agent calls
|
||||
@@ -3851,6 +3953,12 @@ class ChatSession:
|
||||
else:
|
||||
output = f"Unknown tool: {tool_name}"
|
||||
|
||||
# Output guard: evaluate before truncation so the guard
|
||||
# sees full output (credentials split by truncation would
|
||||
# evade detection). Agent outputs are always str.
|
||||
if self._judge_cfg and self._judge_cfg.output_guard and isinstance(output, str):
|
||||
output = self._evaluate_output(tc_dict["id"], output, tool_name)
|
||||
|
||||
# Truncate large tool outputs to avoid blowing context limits.
|
||||
# Agents operate autonomously; they can refine their queries
|
||||
# if truncation loses important detail.
|
||||
@@ -3919,7 +4027,7 @@ class ChatSession:
|
||||
agent_messages,
|
||||
label="task",
|
||||
tools=self._task_tools,
|
||||
auto_tools=self._TASK_AUTO_TOOLS,
|
||||
auto_tools=TASK_AUTO_TOOLS,
|
||||
)
|
||||
except (KeyboardInterrupt, GenerationCancelled):
|
||||
return call_id, "(task interrupted by user)"
|
||||
@@ -4812,12 +4920,14 @@ class ChatSession:
|
||||
|
||||
text = ""
|
||||
try:
|
||||
from turnstone.core.env import scrubbed_env
|
||||
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=10,
|
||||
env={**os.environ, "MANWIDTH": "80", "MAN_KEEP_FORMATTING": "0"},
|
||||
env=scrubbed_env(extra={"MANWIDTH": "80", "MAN_KEEP_FORMATTING": "0"}),
|
||||
)
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
# Strip formatting: backspace overstrikes and ANSI escapes
|
||||
@@ -4830,6 +4940,7 @@ class ChatSession:
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=10,
|
||||
env=scrubbed_env(),
|
||||
)
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
text = result.stdout
|
||||
@@ -4936,51 +5047,26 @@ class ChatSession:
|
||||
return call_id, answer
|
||||
|
||||
def _exec_web_search(self, item: dict[str, Any]) -> tuple[str, str]:
|
||||
"""Search the web via Tavily API."""
|
||||
"""Search the web via the configured backend (Tavily, DDG, or MCP)."""
|
||||
call_id = item["call_id"]
|
||||
query = item["query"]
|
||||
max_results = item.get("max_results", 5)
|
||||
topic = item.get("topic", "general")
|
||||
api_key = get_tavily_key()
|
||||
|
||||
try:
|
||||
resp = httpx.post(
|
||||
"https://api.tavily.com/search",
|
||||
json={
|
||||
"query": query,
|
||||
"max_results": max_results,
|
||||
"topic": topic,
|
||||
"include_answer": True,
|
||||
},
|
||||
headers={"Authorization": f"Bearer {api_key}"},
|
||||
timeout=self.tool_timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
except Exception as e:
|
||||
msg = f"Tavily search failed: {e}"
|
||||
client = self._resolve_search_client()
|
||||
if not client:
|
||||
msg = "Web search backend not available"
|
||||
self.ui.on_error(msg)
|
||||
return call_id, msg
|
||||
|
||||
parts: list[str] = []
|
||||
answer = (data.get("answer") or "").strip()
|
||||
if answer:
|
||||
parts.append(f"Answer: {answer}")
|
||||
|
||||
results = data.get("results") or []
|
||||
if results:
|
||||
lines = []
|
||||
for i, r in enumerate(results, 1):
|
||||
title = r.get("title", "")
|
||||
url = r.get("url", "")
|
||||
content = (r.get("content") or "")[:500]
|
||||
lines.append(f"{i}. [{title}]({url})\n {content}")
|
||||
parts.append("\n".join(lines))
|
||||
|
||||
output = "\n\n".join(parts) if parts else f"No results for '{query}'."
|
||||
try:
|
||||
output = client.search(query, max_results=max_results, topic=topic)
|
||||
except Exception as e:
|
||||
msg = f"Web search failed: {e}"
|
||||
self.ui.on_error(msg)
|
||||
return call_id, msg
|
||||
|
||||
self.ui.on_tool_result(call_id, "web_search", output)
|
||||
|
||||
return call_id, output
|
||||
|
||||
def handle_command(self, cmd_line: str) -> bool:
|
||||
@@ -5148,6 +5234,7 @@ class ChatSession:
|
||||
self.model = model_name
|
||||
self._model_alias = arg
|
||||
self._provider = self._registry.get_provider(arg)
|
||||
self._cached_capabilities = None
|
||||
self.context_window = cfg.context_window
|
||||
if not self._manual_tool_truncation:
|
||||
self.tool_truncation = int(cfg.context_window * self._chars_per_token * 0.5)
|
||||
|
||||
@@ -196,6 +196,17 @@ def _build_registry() -> dict[str, SettingDef]:
|
||||
min_value=1,
|
||||
max_value=50,
|
||||
),
|
||||
SettingDef(
|
||||
"tools.web_search_backend",
|
||||
"str",
|
||||
"",
|
||||
"Web search backend: '' (auto), 'tavily', 'ddg', or 'mcp:server:tool'",
|
||||
"tools",
|
||||
help="Controls which service handles web_search calls when the model lacks native "
|
||||
"search support. Empty string auto-detects (Tavily if key present, else DuckDuckGo "
|
||||
"if installed). 'ddg' uses DuckDuckGo (free, no API key). 'tavily' forces Tavily. "
|
||||
"'mcp:server:tool' routes to an MCP server (e.g. 'mcp:ddg:search').",
|
||||
),
|
||||
# -- server ---------------------------------------------------------
|
||||
SettingDef(
|
||||
"server.workstream_idle_timeout",
|
||||
@@ -212,13 +223,42 @@ def _build_registry() -> dict[str, SettingDef]:
|
||||
SettingDef(
|
||||
"server.max_workstreams",
|
||||
"int",
|
||||
10,
|
||||
50,
|
||||
"Max concurrent workstreams",
|
||||
"server",
|
||||
min_value=1,
|
||||
restart_required=True,
|
||||
help="Maximum number of active conversation threads on this server node. "
|
||||
"When the limit is reached, the oldest idle workstream is evicted to make room.",
|
||||
"When the limit is reached, the oldest idle workstream is evicted to make room. "
|
||||
"Each workstream uses memory proportional to its conversation history.",
|
||||
),
|
||||
# -- cluster --------------------------------------------------------
|
||||
SettingDef(
|
||||
"cluster.node_fan_out_limit",
|
||||
"int",
|
||||
200,
|
||||
"Max concurrent outbound requests during cluster-wide operations",
|
||||
"cluster",
|
||||
min_value=10,
|
||||
max_value=1000,
|
||||
restart_required=True,
|
||||
help="Controls how many nodes the console queries in parallel during "
|
||||
"fan-out operations (watch listing, MCP status, reload notifications). "
|
||||
"Higher values speed up large-cluster admin operations at the cost of "
|
||||
"more concurrent connections. The httpx proxy pool is sized to match "
|
||||
"this value (requires console restart to take effect).",
|
||||
),
|
||||
SettingDef(
|
||||
"cluster.mcp_max_servers",
|
||||
"int",
|
||||
200,
|
||||
"Max MCP server definitions in the cluster",
|
||||
"cluster",
|
||||
min_value=1,
|
||||
max_value=2000,
|
||||
help="Hard cap on the total number of MCP server definitions stored in the "
|
||||
"database. Each node only connects to the servers it needs, so this "
|
||||
"limit is on definitions, not active connections.",
|
||||
),
|
||||
# -- mcp ------------------------------------------------------------
|
||||
SettingDef(
|
||||
|
||||
@@ -7,7 +7,6 @@ and :func:`fetch_skill_from_github` for fetching SKILL.md from GitHub repos.
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
@@ -16,9 +15,10 @@ from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.skill_parser import ParsedSkill, parse_skill_md
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
DEFAULT_DISCOVERY_URL = "https://skills.sh"
|
||||
|
||||
@@ -176,7 +176,7 @@ def _check_rate_limit(resp: httpx.Response) -> None:
|
||||
)
|
||||
remaining = resp.headers.get("x-ratelimit-remaining", "")
|
||||
if remaining and remaining.isdigit() and int(remaining) < 10:
|
||||
logger.warning("GitHub API rate limit low: %s remaining", remaining)
|
||||
log.warning("GitHub API rate limit low: %s remaining", remaining)
|
||||
|
||||
|
||||
_FETCH_CONCURRENCY = 5
|
||||
@@ -301,7 +301,7 @@ async def fetch_skill_from_github(url: str) -> SkillPackage:
|
||||
rf = _find_resource_files(tree_data.get("tree", []), skill_md_dir)
|
||||
resources = await _fetch_resource_contents(client, raw_base, rf)
|
||||
except httpx.HTTPError:
|
||||
logger.debug("Failed to fetch resource tree for %s/%s", owner, repo)
|
||||
log.debug("Failed to fetch resource tree for %s/%s", owner, repo)
|
||||
|
||||
# Build a per-skill source URL pointing to the specific subdirectory
|
||||
if skill_md_dir:
|
||||
@@ -418,7 +418,7 @@ async def fetch_skills_from_github_repo(url: str) -> list[SkillPackage]:
|
||||
try:
|
||||
parsed = parse_skill_md(content)
|
||||
except ValueError:
|
||||
logger.debug("Skipping invalid SKILL.md at %s", skill_md_path)
|
||||
log.debug("Skipping invalid SKILL.md at %s", skill_md_path)
|
||||
continue
|
||||
|
||||
# Collect resources for this skill (concurrent via helper)
|
||||
|
||||
@@ -2,11 +2,12 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
_MIGRATIONS_DIR = str(Path(__file__).parent / "migrations")
|
||||
|
||||
@@ -35,7 +36,14 @@ def run_migrations(storage: Any, backend: str) -> None:
|
||||
_bootstrap_existing_sqlite(engine, cfg)
|
||||
|
||||
if backend == "postgresql":
|
||||
_run_with_pg_lock(engine, cfg)
|
||||
try:
|
||||
_run_with_pg_lock(engine, cfg)
|
||||
except (OSError, EOFError) as exc:
|
||||
# Non-fatal for connection-class errors only (refused, reset,
|
||||
# timeout). DDL / migration errors still propagate. The Docker
|
||||
# entrypoint already runs migrations before the server starts, so
|
||||
# this second attempt is a safety net for stampede scenarios.
|
||||
log.warning("PostgreSQL migration failed (non-fatal): %s", exc)
|
||||
else:
|
||||
try:
|
||||
command.upgrade(cfg, "head")
|
||||
@@ -49,17 +57,42 @@ def _run_with_pg_lock(engine: Any, cfg: Any) -> None:
|
||||
Advisory lock ID 7_475_283 (arbitrary, derived from 'turnstone').
|
||||
``pg_advisory_lock`` blocks until the lock is available, so
|
||||
concurrent containers wait in line rather than racing.
|
||||
|
||||
Retries with jittered backoff if PostgreSQL is temporarily at
|
||||
max_connections (common during large-cluster startup stampedes).
|
||||
"""
|
||||
import random
|
||||
import time
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import command
|
||||
|
||||
with engine.connect() as conn:
|
||||
conn.execute(sa.text("SELECT pg_advisory_lock(7475283)"))
|
||||
max_retries = 10
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
command.upgrade(cfg, "head")
|
||||
finally:
|
||||
conn.execute(sa.text("SELECT pg_advisory_unlock(7475283)"))
|
||||
conn.commit()
|
||||
with engine.connect() as conn:
|
||||
conn.execute(sa.text("SELECT pg_advisory_lock(7475283)"))
|
||||
try:
|
||||
command.upgrade(cfg, "head")
|
||||
finally:
|
||||
conn.execute(sa.text("SELECT pg_advisory_unlock(7475283)"))
|
||||
conn.commit()
|
||||
return
|
||||
except Exception as exc:
|
||||
err_str = str(exc).lower()
|
||||
if "too many clients" not in err_str and "connection" not in err_str:
|
||||
raise
|
||||
if attempt == max_retries - 1:
|
||||
raise
|
||||
delay = min(2**attempt + random.uniform(0, 1), 30) # noqa: S311
|
||||
log.warning(
|
||||
"PG connection failed (attempt %d/%d), retrying in %.1fs: %s",
|
||||
attempt + 1,
|
||||
max_retries,
|
||||
delay,
|
||||
exc,
|
||||
)
|
||||
time.sleep(delay)
|
||||
|
||||
|
||||
def _bootstrap_existing_sqlite(engine: Any, cfg: Any) -> None:
|
||||
|
||||
@@ -2,23 +2,30 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Any
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.storage._schema import (
|
||||
api_tokens,
|
||||
audit_events,
|
||||
channel_routes,
|
||||
channel_users,
|
||||
conversations,
|
||||
intent_verdicts,
|
||||
mcp_servers,
|
||||
metadata,
|
||||
oidc_identities,
|
||||
oidc_pending_states,
|
||||
orgs,
|
||||
output_assessments,
|
||||
prompt_templates,
|
||||
roles,
|
||||
scheduled_task_runs,
|
||||
scheduled_tasks,
|
||||
services,
|
||||
skill_resources,
|
||||
skill_versions,
|
||||
structured_memories,
|
||||
@@ -27,6 +34,7 @@ from turnstone.core.storage._schema import (
|
||||
usage_events,
|
||||
user_roles,
|
||||
users,
|
||||
watches,
|
||||
workstream_config,
|
||||
workstreams,
|
||||
)
|
||||
@@ -61,7 +69,7 @@ from turnstone.core.storage._utils import (
|
||||
scan_skill_content as _scan_skill_content,
|
||||
)
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
|
||||
def _escape_ilike(s: str) -> str:
|
||||
@@ -73,7 +81,7 @@ class PostgreSQLBackend:
|
||||
"""PostgreSQL implementation of the StorageBackend protocol."""
|
||||
|
||||
def __init__(
|
||||
self, url: str, pool_size: int = 5, max_overflow: int = 10, *, create_tables: bool = True
|
||||
self, url: str, pool_size: int = 2, max_overflow: int = 3, *, create_tables: bool = True
|
||||
) -> None:
|
||||
self._engine = sa.create_engine(
|
||||
url,
|
||||
@@ -229,19 +237,17 @@ class PostgreSQLBackend:
|
||||
# -- Workstream config -----------------------------------------------------
|
||||
|
||||
def save_workstream_config(self, ws_id: str, config: dict[str, str]) -> None:
|
||||
if not config:
|
||||
return
|
||||
with self._engine.connect() as conn:
|
||||
for key, value in config.items():
|
||||
# Upsert: delete + insert
|
||||
conn.execute(
|
||||
sa.delete(workstream_config).where(
|
||||
workstream_config.c.ws_id == ws_id,
|
||||
workstream_config.c.key == key,
|
||||
)
|
||||
)
|
||||
conn.execute(
|
||||
sa.insert(workstream_config),
|
||||
{"ws_id": ws_id, "key": key, "value": value},
|
||||
)
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"INSERT INTO workstream_config (ws_id, key, value) "
|
||||
"VALUES (:ws_id, :key, :value) "
|
||||
"ON CONFLICT (ws_id, key) DO UPDATE SET value = EXCLUDED.value"
|
||||
),
|
||||
[{"ws_id": ws_id, "key": key, "value": value} for key, value in config.items()],
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
def load_workstream_config(self, ws_id: str) -> dict[str, str]:
|
||||
@@ -524,7 +530,6 @@ class PostgreSQLBackend:
|
||||
]
|
||||
|
||||
def delete_user(self, user_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import channel_users, oidc_identities
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(sa.delete(user_roles).where(user_roles.c.user_id == user_id))
|
||||
@@ -630,8 +635,6 @@ class PostgreSQLBackend:
|
||||
def create_channel_user(self, channel_type: str, channel_user_id: str, user_id: str) -> None:
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
from turnstone.core.storage._schema import channel_users
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(
|
||||
@@ -647,7 +650,6 @@ class PostgreSQLBackend:
|
||||
conn.commit()
|
||||
|
||||
def get_channel_user(self, channel_type: str, channel_user_id: str) -> dict[str, str] | None:
|
||||
from turnstone.core.storage._schema import channel_users
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
@@ -671,7 +673,6 @@ class PostgreSQLBackend:
|
||||
return None
|
||||
|
||||
def list_channel_users_by_user(self, user_id: str) -> list[dict[str, str]]:
|
||||
from turnstone.core.storage._schema import channel_users
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -695,7 +696,6 @@ class PostgreSQLBackend:
|
||||
]
|
||||
|
||||
def delete_channel_user(self, channel_type: str, channel_user_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import channel_users
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(
|
||||
@@ -714,8 +714,6 @@ class PostgreSQLBackend:
|
||||
) -> None:
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
from turnstone.core.storage._schema import channel_routes
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(
|
||||
@@ -732,7 +730,6 @@ class PostgreSQLBackend:
|
||||
conn.commit()
|
||||
|
||||
def get_channel_route(self, channel_type: str, channel_id: str) -> dict[str, str] | None:
|
||||
from turnstone.core.storage._schema import channel_routes
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
@@ -758,7 +755,6 @@ class PostgreSQLBackend:
|
||||
return None
|
||||
|
||||
def get_channel_route_by_ws(self, ws_id: str) -> dict[str, str] | None:
|
||||
from turnstone.core.storage._schema import channel_routes
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
@@ -781,7 +777,6 @@ class PostgreSQLBackend:
|
||||
return None
|
||||
|
||||
def list_channel_routes_by_type(self, channel_type: str) -> list[dict[str, str]]:
|
||||
from turnstone.core.storage._schema import channel_routes
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -807,7 +802,6 @@ class PostgreSQLBackend:
|
||||
]
|
||||
|
||||
def delete_channel_route(self, channel_type: str, channel_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import channel_routes
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(
|
||||
@@ -840,8 +834,6 @@ class PostgreSQLBackend:
|
||||
) -> None:
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
from turnstone.core.storage._schema import scheduled_tasks
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(
|
||||
@@ -870,7 +862,6 @@ class PostgreSQLBackend:
|
||||
conn.commit()
|
||||
|
||||
def get_scheduled_task(self, task_id: str) -> dict[str, Any] | None:
|
||||
from turnstone.core.storage._schema import scheduled_tasks
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
@@ -881,7 +872,6 @@ class PostgreSQLBackend:
|
||||
return dict(row._mapping)
|
||||
|
||||
def list_scheduled_tasks(self) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import scheduled_tasks
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -910,7 +900,6 @@ class PostgreSQLBackend:
|
||||
)
|
||||
|
||||
def update_scheduled_task(self, task_id: str, **fields: Any) -> bool:
|
||||
from turnstone.core.storage._schema import scheduled_tasks
|
||||
|
||||
fields = {k: v for k, v in fields.items() if k in self._UPDATABLE_TASK_FIELDS}
|
||||
fields["updated"] = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
@@ -930,7 +919,6 @@ class PostgreSQLBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def delete_scheduled_task(self, task_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import scheduled_task_runs, scheduled_tasks
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(
|
||||
@@ -943,7 +931,6 @@ class PostgreSQLBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def list_due_tasks(self, now: str) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import scheduled_tasks
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -969,7 +956,6 @@ class PostgreSQLBackend:
|
||||
status: str,
|
||||
error: str,
|
||||
) -> None:
|
||||
from turnstone.core.storage._schema import scheduled_task_runs
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(
|
||||
@@ -988,7 +974,6 @@ class PostgreSQLBackend:
|
||||
conn.commit()
|
||||
|
||||
def list_task_runs(self, task_id: str, limit: int = 50) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import scheduled_task_runs
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -1000,10 +985,6 @@ class PostgreSQLBackend:
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def prune_task_runs(self, retention_days: int = 90) -> int:
|
||||
from datetime import timedelta
|
||||
|
||||
from turnstone.core.storage._schema import scheduled_task_runs
|
||||
|
||||
cutoff = (datetime.now(UTC) - timedelta(days=retention_days)).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(
|
||||
@@ -1029,8 +1010,6 @@ class PostgreSQLBackend:
|
||||
) -> None:
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(
|
||||
@@ -1056,7 +1035,6 @@ class PostgreSQLBackend:
|
||||
conn.commit()
|
||||
|
||||
def get_watch(self, watch_id: str) -> dict[str, Any] | None:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(sa.select(watches).where(watches.c.watch_id == watch_id)).fetchone()
|
||||
@@ -1065,7 +1043,6 @@ class PostgreSQLBackend:
|
||||
return dict(row._mapping)
|
||||
|
||||
def list_watches_for_ws(self, ws_id: str) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -1076,7 +1053,6 @@ class PostgreSQLBackend:
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def list_watches_for_node(self, node_id: str) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -1087,7 +1063,6 @@ class PostgreSQLBackend:
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def list_due_watches(self, now: str) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -1116,7 +1091,6 @@ class PostgreSQLBackend:
|
||||
)
|
||||
|
||||
def update_watch(self, watch_id: str, **fields: Any) -> bool:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
fields = {k: v for k, v in fields.items() if k in self._UPDATABLE_WATCH_FIELDS}
|
||||
fields["updated"] = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
@@ -1130,7 +1104,6 @@ class PostgreSQLBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def delete_watch(self, watch_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(sa.delete(watches).where(watches.c.watch_id == watch_id))
|
||||
@@ -1138,7 +1111,6 @@ class PostgreSQLBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def delete_watches_for_ws(self, ws_id: str) -> int:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(sa.delete(watches).where(watches.c.ws_id == ws_id))
|
||||
@@ -1150,7 +1122,6 @@ class PostgreSQLBackend:
|
||||
def register_service(
|
||||
self, service_type: str, service_id: str, url: str, metadata: str = "{}"
|
||||
) -> None:
|
||||
from turnstone.core.storage._schema import services
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -1172,7 +1143,6 @@ class PostgreSQLBackend:
|
||||
conn.commit()
|
||||
|
||||
def heartbeat_service(self, service_type: str, service_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import services
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -1188,7 +1158,6 @@ class PostgreSQLBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def list_services(self, service_type: str, max_age_seconds: int = 120) -> list[dict[str, str]]:
|
||||
from turnstone.core.storage._schema import services
|
||||
|
||||
cutoff = (datetime.now(UTC) - timedelta(seconds=max_age_seconds)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
@@ -1205,7 +1174,6 @@ class PostgreSQLBackend:
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def deregister_service(self, service_type: str, service_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import services
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(
|
||||
@@ -1510,6 +1478,7 @@ class PostgreSQLBackend:
|
||||
allowed_tools: str = "[]",
|
||||
skill_license: str = "",
|
||||
compatibility: str = "",
|
||||
priority: int = 0,
|
||||
) -> None:
|
||||
# Sync is_default from activation when activation is explicitly set
|
||||
if activation == "default":
|
||||
@@ -1556,6 +1525,7 @@ class PostgreSQLBackend:
|
||||
"agent_max_turns": agent_max_turns,
|
||||
"notify_on_complete": notify_on_complete,
|
||||
"enabled": 1 if enabled else 0,
|
||||
"priority": priority,
|
||||
"created": now,
|
||||
"updated": now,
|
||||
},
|
||||
@@ -1609,7 +1579,7 @@ class PostgreSQLBackend:
|
||||
sa.select(prompt_templates)
|
||||
.where(prompt_templates.c.is_default == 1)
|
||||
.where(prompt_templates.c.enabled == 1)
|
||||
.order_by(prompt_templates.c.name)
|
||||
.order_by(prompt_templates.c.priority, prompt_templates.c.name)
|
||||
)
|
||||
if org_id:
|
||||
q = q.where(prompt_templates.c.org_id == org_id)
|
||||
@@ -1694,7 +1664,7 @@ class PostgreSQLBackend:
|
||||
q = (
|
||||
sa.select(prompt_templates)
|
||||
.where(prompt_templates.c.activation == activation)
|
||||
.order_by(prompt_templates.c.name)
|
||||
.order_by(prompt_templates.c.priority, prompt_templates.c.name)
|
||||
)
|
||||
if enabled_only:
|
||||
q = q.where(prompt_templates.c.enabled == 1)
|
||||
@@ -2426,6 +2396,32 @@ class PostgreSQLBackend:
|
||||
).fetchall()
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def touch_structured_memories(self, keys: list[tuple[str, str, str]]) -> int:
|
||||
"""Batch-touch multiple memories by (name, scope, scope_id)."""
|
||||
if not keys:
|
||||
return 0
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
total = 0
|
||||
with self._engine.connect() as conn:
|
||||
for name, scope, scope_id in keys:
|
||||
result = conn.execute(
|
||||
sa.update(structured_memories)
|
||||
.where(
|
||||
sa.and_(
|
||||
structured_memories.c.name == name,
|
||||
structured_memories.c.scope == scope,
|
||||
structured_memories.c.scope_id == scope_id,
|
||||
)
|
||||
)
|
||||
.values(
|
||||
last_accessed=now,
|
||||
access_count=structured_memories.c.access_count + 1,
|
||||
)
|
||||
)
|
||||
total += result.rowcount
|
||||
conn.commit()
|
||||
return total
|
||||
|
||||
def count_structured_memories(
|
||||
self, mem_type: str = "", scope: str = "", scope_id: str = ""
|
||||
) -> int:
|
||||
@@ -2650,8 +2646,6 @@ class PostgreSQLBackend:
|
||||
def create_oidc_identity(self, issuer: str, subject: str, user_id: str, email: str) -> None:
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
from turnstone.core.storage._schema import oidc_identities
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(
|
||||
@@ -2669,7 +2663,6 @@ class PostgreSQLBackend:
|
||||
conn.commit()
|
||||
|
||||
def get_oidc_identity(self, issuer: str, subject: str) -> dict[str, str] | None:
|
||||
from turnstone.core.storage._schema import oidc_identities
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
@@ -2696,7 +2689,6 @@ class PostgreSQLBackend:
|
||||
return None
|
||||
|
||||
def update_oidc_identity_login(self, issuer: str, subject: str) -> bool:
|
||||
from turnstone.core.storage._schema import oidc_identities
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -2711,7 +2703,6 @@ class PostgreSQLBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def list_oidc_identities_for_user(self, user_id: str) -> list[dict[str, str]]:
|
||||
from turnstone.core.storage._schema import oidc_identities
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -2739,7 +2730,6 @@ class PostgreSQLBackend:
|
||||
]
|
||||
|
||||
def delete_oidc_identity(self, issuer: str, subject: str) -> bool:
|
||||
from turnstone.core.storage._schema import oidc_identities
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(
|
||||
@@ -2755,7 +2745,6 @@ class PostgreSQLBackend:
|
||||
def create_oidc_pending_state(
|
||||
self, state: str, nonce: str, code_verifier: str, audience: str
|
||||
) -> None:
|
||||
from turnstone.core.storage._schema import oidc_pending_states
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -2774,7 +2763,6 @@ class PostgreSQLBackend:
|
||||
def pop_oidc_pending_state(
|
||||
self, state: str, max_age_seconds: int = 300
|
||||
) -> dict[str, str] | None:
|
||||
from turnstone.core.storage._schema import oidc_pending_states
|
||||
|
||||
cutoff = (datetime.now(UTC) - timedelta(seconds=max_age_seconds)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
@@ -2806,7 +2794,6 @@ class PostgreSQLBackend:
|
||||
}
|
||||
|
||||
def cleanup_expired_oidc_states(self, max_age_seconds: int = 300) -> int:
|
||||
from turnstone.core.storage._schema import oidc_pending_states
|
||||
|
||||
cutoff = (datetime.now(UTC) - timedelta(seconds=max_age_seconds)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
|
||||
@@ -131,6 +131,15 @@ class StorageBackend(Protocol):
|
||||
"""Search structured memories by query. Returns matching memory dicts."""
|
||||
...
|
||||
|
||||
def touch_structured_memories(self, keys: list[tuple[str, str, str]]) -> int:
|
||||
"""Batch-touch multiple memories.
|
||||
|
||||
Each key is ``(name, scope, scope_id)``. Callers should deduplicate
|
||||
before calling; each key increments ``access_count`` once per call.
|
||||
Returns count of rows found and updated.
|
||||
"""
|
||||
...
|
||||
|
||||
def count_structured_memories(
|
||||
self, mem_type: str = "", scope: str = "", scope_id: str = ""
|
||||
) -> int:
|
||||
@@ -581,6 +590,7 @@ class StorageBackend(Protocol):
|
||||
allowed_tools: str = "[]",
|
||||
skill_license: str = "",
|
||||
compatibility: str = "",
|
||||
priority: int = 0,
|
||||
) -> None:
|
||||
"""Create a prompt template (skill)."""
|
||||
...
|
||||
@@ -626,7 +636,7 @@ class StorageBackend(Protocol):
|
||||
enabled_only: bool = False,
|
||||
limit: int = 0,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Return prompt templates filtered by activation value, ordered by name."""
|
||||
"""Return prompt templates filtered by activation value, ordered by priority then name."""
|
||||
...
|
||||
|
||||
def get_skill_by_name(self, name: str) -> dict[str, Any] | None:
|
||||
|
||||
@@ -2,14 +2,15 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from turnstone.core.storage._protocol import StorageBackend
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
_storage: StorageBackend | None = None
|
||||
|
||||
@@ -19,7 +20,7 @@ def init_storage(
|
||||
*,
|
||||
path: str = "",
|
||||
url: str = "",
|
||||
pool_size: int = 5,
|
||||
pool_size: int = 2,
|
||||
run_migrations: bool = True,
|
||||
) -> StorageBackend:
|
||||
"""Initialize the storage backend singleton.
|
||||
|
||||
@@ -41,6 +41,8 @@ conversations = sa.Table(
|
||||
sa.Column("tool_calls", sa.Text),
|
||||
)
|
||||
|
||||
sa.Index("idx_conversations_timestamp", conversations.c.timestamp)
|
||||
|
||||
workstreams = sa.Table(
|
||||
"workstreams",
|
||||
metadata,
|
||||
@@ -328,6 +330,7 @@ prompt_templates = sa.Table(
|
||||
sa.Column("agent_max_turns", sa.Integer, nullable=True),
|
||||
sa.Column("notify_on_complete", sa.Text, nullable=False, server_default="{}"),
|
||||
sa.Column("enabled", sa.Integer, nullable=False, server_default="1"),
|
||||
sa.Column("priority", sa.Integer, nullable=False, server_default="0"),
|
||||
sa.Column("created", sa.Text, nullable=False),
|
||||
sa.Column("updated", sa.Text, nullable=False),
|
||||
)
|
||||
|
||||
@@ -2,23 +2,30 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Any
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.storage._schema import (
|
||||
api_tokens,
|
||||
audit_events,
|
||||
channel_routes,
|
||||
channel_users,
|
||||
conversations,
|
||||
intent_verdicts,
|
||||
mcp_servers,
|
||||
metadata,
|
||||
oidc_identities,
|
||||
oidc_pending_states,
|
||||
orgs,
|
||||
output_assessments,
|
||||
prompt_templates,
|
||||
roles,
|
||||
scheduled_task_runs,
|
||||
scheduled_tasks,
|
||||
services,
|
||||
skill_resources,
|
||||
skill_versions,
|
||||
structured_memories,
|
||||
@@ -27,6 +34,7 @@ from turnstone.core.storage._schema import (
|
||||
usage_events,
|
||||
user_roles,
|
||||
users,
|
||||
watches,
|
||||
workstream_config,
|
||||
workstreams,
|
||||
)
|
||||
@@ -61,7 +69,7 @@ from turnstone.core.storage._utils import (
|
||||
scan_skill_content as _scan_skill_content,
|
||||
)
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
|
||||
def _escape_like(s: str) -> str:
|
||||
@@ -87,8 +95,21 @@ class SQLiteBackend:
|
||||
self._engine = sa.create_engine(
|
||||
f"sqlite:///{path}",
|
||||
pool_pre_ping=True,
|
||||
connect_args={"check_same_thread": False},
|
||||
connect_args={"check_same_thread": False, "timeout": 30},
|
||||
)
|
||||
|
||||
# Enable WAL mode for better concurrent read/write performance.
|
||||
@sa.event.listens_for(self._engine, "connect")
|
||||
def _set_wal(dbapi_conn: Any, _rec: Any) -> None:
|
||||
try:
|
||||
cursor = dbapi_conn.execute("PRAGMA journal_mode=WAL")
|
||||
mode = cursor.fetchone()
|
||||
cursor.close()
|
||||
if mode and mode[0] != "wal":
|
||||
log.warning("SQLite WAL mode not enabled (got %s)", mode[0])
|
||||
except Exception:
|
||||
log.warning("Failed to set SQLite WAL mode", exc_info=True)
|
||||
|
||||
self._fts5_available = False
|
||||
if create_tables:
|
||||
self._init_schema()
|
||||
@@ -295,15 +316,16 @@ class SQLiteBackend:
|
||||
# -- Workstream config -----------------------------------------------------
|
||||
|
||||
def save_workstream_config(self, ws_id: str, config: dict[str, str]) -> None:
|
||||
if not config:
|
||||
return
|
||||
with self._engine.connect() as conn:
|
||||
for key, value in config.items():
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"INSERT OR REPLACE INTO workstream_config "
|
||||
"(ws_id, key, value) VALUES (:wid, :key, :value)"
|
||||
),
|
||||
{"wid": ws_id, "key": key, "value": value},
|
||||
)
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"INSERT OR REPLACE INTO workstream_config "
|
||||
"(ws_id, key, value) VALUES (:wid, :key, :value)"
|
||||
),
|
||||
[{"wid": ws_id, "key": key, "value": value} for key, value in config.items()],
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
def load_workstream_config(self, ws_id: str) -> dict[str, str]:
|
||||
@@ -573,7 +595,6 @@ class SQLiteBackend:
|
||||
]
|
||||
|
||||
def delete_user(self, user_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import channel_users, oidc_identities
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(sa.delete(user_roles).where(user_roles.c.user_id == user_id))
|
||||
@@ -677,7 +698,6 @@ class SQLiteBackend:
|
||||
# -- Channel user mapping ---------------------------------------------------
|
||||
|
||||
def create_channel_user(self, channel_type: str, channel_user_id: str, user_id: str) -> None:
|
||||
from turnstone.core.storage._schema import channel_users
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -693,7 +713,6 @@ class SQLiteBackend:
|
||||
conn.commit()
|
||||
|
||||
def get_channel_user(self, channel_type: str, channel_user_id: str) -> dict[str, str] | None:
|
||||
from turnstone.core.storage._schema import channel_users
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
@@ -717,7 +736,6 @@ class SQLiteBackend:
|
||||
return None
|
||||
|
||||
def list_channel_users_by_user(self, user_id: str) -> list[dict[str, str]]:
|
||||
from turnstone.core.storage._schema import channel_users
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -741,7 +759,6 @@ class SQLiteBackend:
|
||||
]
|
||||
|
||||
def delete_channel_user(self, channel_type: str, channel_user_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import channel_users
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(
|
||||
@@ -758,7 +775,6 @@ class SQLiteBackend:
|
||||
def create_channel_route(
|
||||
self, channel_type: str, channel_id: str, ws_id: str, node_id: str = ""
|
||||
) -> None:
|
||||
from turnstone.core.storage._schema import channel_routes
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -775,7 +791,6 @@ class SQLiteBackend:
|
||||
conn.commit()
|
||||
|
||||
def get_channel_route(self, channel_type: str, channel_id: str) -> dict[str, str] | None:
|
||||
from turnstone.core.storage._schema import channel_routes
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
@@ -801,7 +816,6 @@ class SQLiteBackend:
|
||||
return None
|
||||
|
||||
def get_channel_route_by_ws(self, ws_id: str) -> dict[str, str] | None:
|
||||
from turnstone.core.storage._schema import channel_routes
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
@@ -824,7 +838,6 @@ class SQLiteBackend:
|
||||
return None
|
||||
|
||||
def list_channel_routes_by_type(self, channel_type: str) -> list[dict[str, str]]:
|
||||
from turnstone.core.storage._schema import channel_routes
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -850,7 +863,6 @@ class SQLiteBackend:
|
||||
]
|
||||
|
||||
def delete_channel_route(self, channel_type: str, channel_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import channel_routes
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(
|
||||
@@ -881,7 +893,6 @@ class SQLiteBackend:
|
||||
next_run: str,
|
||||
skill: str = "",
|
||||
) -> None:
|
||||
from turnstone.core.storage._schema import scheduled_tasks
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -910,7 +921,6 @@ class SQLiteBackend:
|
||||
conn.commit()
|
||||
|
||||
def get_scheduled_task(self, task_id: str) -> dict[str, Any] | None:
|
||||
from turnstone.core.storage._schema import scheduled_tasks
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
@@ -921,7 +931,6 @@ class SQLiteBackend:
|
||||
return dict(row._mapping)
|
||||
|
||||
def list_scheduled_tasks(self) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import scheduled_tasks
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -950,7 +959,6 @@ class SQLiteBackend:
|
||||
)
|
||||
|
||||
def update_scheduled_task(self, task_id: str, **fields: Any) -> bool:
|
||||
from turnstone.core.storage._schema import scheduled_tasks
|
||||
|
||||
fields = {k: v for k, v in fields.items() if k in self._UPDATABLE_TASK_FIELDS}
|
||||
fields["updated"] = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
@@ -971,7 +979,6 @@ class SQLiteBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def delete_scheduled_task(self, task_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import scheduled_task_runs, scheduled_tasks
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(
|
||||
@@ -984,7 +991,6 @@ class SQLiteBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def list_due_tasks(self, now: str) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import scheduled_tasks
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -1010,7 +1016,6 @@ class SQLiteBackend:
|
||||
status: str,
|
||||
error: str,
|
||||
) -> None:
|
||||
from turnstone.core.storage._schema import scheduled_task_runs
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
conn.execute(
|
||||
@@ -1029,7 +1034,6 @@ class SQLiteBackend:
|
||||
conn.commit()
|
||||
|
||||
def list_task_runs(self, task_id: str, limit: int = 50) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import scheduled_task_runs
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -1041,10 +1045,6 @@ class SQLiteBackend:
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def prune_task_runs(self, retention_days: int = 90) -> int:
|
||||
from datetime import timedelta
|
||||
|
||||
from turnstone.core.storage._schema import scheduled_task_runs
|
||||
|
||||
cutoff = (datetime.now(UTC) - timedelta(days=retention_days)).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(
|
||||
@@ -1068,7 +1068,6 @@ class SQLiteBackend:
|
||||
created_by: str,
|
||||
next_poll: str,
|
||||
) -> None:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -1094,7 +1093,6 @@ class SQLiteBackend:
|
||||
conn.commit()
|
||||
|
||||
def get_watch(self, watch_id: str) -> dict[str, Any] | None:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(sa.select(watches).where(watches.c.watch_id == watch_id)).fetchone()
|
||||
@@ -1103,7 +1101,6 @@ class SQLiteBackend:
|
||||
return dict(row._mapping)
|
||||
|
||||
def list_watches_for_ws(self, ws_id: str) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -1114,7 +1111,6 @@ class SQLiteBackend:
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def list_watches_for_node(self, node_id: str) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -1125,7 +1121,6 @@ class SQLiteBackend:
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def list_due_watches(self, now: str) -> list[dict[str, Any]]:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -1154,7 +1149,6 @@ class SQLiteBackend:
|
||||
)
|
||||
|
||||
def update_watch(self, watch_id: str, **fields: Any) -> bool:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
fields = {k: v for k, v in fields.items() if k in self._UPDATABLE_WATCH_FIELDS}
|
||||
fields["updated"] = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
@@ -1168,7 +1162,6 @@ class SQLiteBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def delete_watch(self, watch_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(sa.delete(watches).where(watches.c.watch_id == watch_id))
|
||||
@@ -1176,7 +1169,6 @@ class SQLiteBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def delete_watches_for_ws(self, ws_id: str) -> int:
|
||||
from turnstone.core.storage._schema import watches
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(sa.delete(watches).where(watches.c.ws_id == ws_id))
|
||||
@@ -1190,8 +1182,6 @@ class SQLiteBackend:
|
||||
) -> None:
|
||||
from sqlalchemy.dialects.sqlite import insert as sqlite_insert
|
||||
|
||||
from turnstone.core.storage._schema import services
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
stmt = sqlite_insert(services).values(
|
||||
service_type=service_type,
|
||||
@@ -1210,7 +1200,6 @@ class SQLiteBackend:
|
||||
conn.commit()
|
||||
|
||||
def heartbeat_service(self, service_type: str, service_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import services
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -1226,7 +1215,6 @@ class SQLiteBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def list_services(self, service_type: str, max_age_seconds: int = 120) -> list[dict[str, str]]:
|
||||
from turnstone.core.storage._schema import services
|
||||
|
||||
cutoff = (datetime.now(UTC) - timedelta(seconds=max_age_seconds)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
@@ -1243,7 +1231,6 @@ class SQLiteBackend:
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def deregister_service(self, service_type: str, service_id: str) -> bool:
|
||||
from turnstone.core.storage._schema import services
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(
|
||||
@@ -1534,6 +1521,7 @@ class SQLiteBackend:
|
||||
allowed_tools: str = "[]",
|
||||
skill_license: str = "",
|
||||
compatibility: str = "",
|
||||
priority: int = 0,
|
||||
) -> None:
|
||||
# Sync is_default from activation when activation is explicitly set
|
||||
if activation == "default":
|
||||
@@ -1580,6 +1568,7 @@ class SQLiteBackend:
|
||||
"agent_max_turns": agent_max_turns,
|
||||
"notify_on_complete": notify_on_complete,
|
||||
"enabled": 1 if enabled else 0,
|
||||
"priority": priority,
|
||||
"created": now,
|
||||
"updated": now,
|
||||
},
|
||||
@@ -1633,7 +1622,7 @@ class SQLiteBackend:
|
||||
sa.select(prompt_templates)
|
||||
.where(prompt_templates.c.is_default == 1)
|
||||
.where(prompt_templates.c.enabled == 1)
|
||||
.order_by(prompt_templates.c.name)
|
||||
.order_by(prompt_templates.c.priority, prompt_templates.c.name)
|
||||
)
|
||||
if org_id:
|
||||
q = q.where(prompt_templates.c.org_id == org_id)
|
||||
@@ -1718,7 +1707,7 @@ class SQLiteBackend:
|
||||
q = (
|
||||
sa.select(prompt_templates)
|
||||
.where(prompt_templates.c.activation == activation)
|
||||
.order_by(prompt_templates.c.name)
|
||||
.order_by(prompt_templates.c.priority, prompt_templates.c.name)
|
||||
)
|
||||
if enabled_only:
|
||||
q = q.where(prompt_templates.c.enabled == 1)
|
||||
@@ -2457,6 +2446,32 @@ class SQLiteBackend:
|
||||
).fetchall()
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def touch_structured_memories(self, keys: list[tuple[str, str, str]]) -> int:
|
||||
"""Batch-touch multiple memories by (name, scope, scope_id)."""
|
||||
if not keys:
|
||||
return 0
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
total = 0
|
||||
with self._engine.connect() as conn:
|
||||
for name, scope, scope_id in keys:
|
||||
result = conn.execute(
|
||||
sa.update(structured_memories)
|
||||
.where(
|
||||
sa.and_(
|
||||
structured_memories.c.name == name,
|
||||
structured_memories.c.scope == scope,
|
||||
structured_memories.c.scope_id == scope_id,
|
||||
)
|
||||
)
|
||||
.values(
|
||||
last_accessed=now,
|
||||
access_count=structured_memories.c.access_count + 1,
|
||||
)
|
||||
)
|
||||
total += result.rowcount
|
||||
conn.commit()
|
||||
return total
|
||||
|
||||
def count_structured_memories(
|
||||
self, mem_type: str = "", scope: str = "", scope_id: str = ""
|
||||
) -> int:
|
||||
@@ -2678,7 +2693,6 @@ class SQLiteBackend:
|
||||
# -- OIDC identity ---------------------------------------------------------
|
||||
|
||||
def create_oidc_identity(self, issuer: str, subject: str, user_id: str, email: str) -> None:
|
||||
from turnstone.core.storage._schema import oidc_identities
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -2696,7 +2710,6 @@ class SQLiteBackend:
|
||||
conn.commit()
|
||||
|
||||
def get_oidc_identity(self, issuer: str, subject: str) -> dict[str, str] | None:
|
||||
from turnstone.core.storage._schema import oidc_identities
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
row = conn.execute(
|
||||
@@ -2723,7 +2736,6 @@ class SQLiteBackend:
|
||||
return None
|
||||
|
||||
def update_oidc_identity_login(self, issuer: str, subject: str) -> bool:
|
||||
from turnstone.core.storage._schema import oidc_identities
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -2738,7 +2750,6 @@ class SQLiteBackend:
|
||||
return result.rowcount > 0
|
||||
|
||||
def list_oidc_identities_for_user(self, user_id: str) -> list[dict[str, str]]:
|
||||
from turnstone.core.storage._schema import oidc_identities
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
rows = conn.execute(
|
||||
@@ -2766,7 +2777,6 @@ class SQLiteBackend:
|
||||
]
|
||||
|
||||
def delete_oidc_identity(self, issuer: str, subject: str) -> bool:
|
||||
from turnstone.core.storage._schema import oidc_identities
|
||||
|
||||
with self._engine.connect() as conn:
|
||||
result = conn.execute(
|
||||
@@ -2782,7 +2792,6 @@ class SQLiteBackend:
|
||||
def create_oidc_pending_state(
|
||||
self, state: str, nonce: str, code_verifier: str, audience: str
|
||||
) -> None:
|
||||
from turnstone.core.storage._schema import oidc_pending_states
|
||||
|
||||
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
|
||||
with self._engine.connect() as conn:
|
||||
@@ -2801,7 +2810,6 @@ class SQLiteBackend:
|
||||
def pop_oidc_pending_state(
|
||||
self, state: str, max_age_seconds: int = 300
|
||||
) -> dict[str, str] | None:
|
||||
from turnstone.core.storage._schema import oidc_pending_states
|
||||
|
||||
cutoff = (datetime.now(UTC) - timedelta(seconds=max_age_seconds)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
@@ -2835,7 +2843,6 @@ class SQLiteBackend:
|
||||
}
|
||||
|
||||
def cleanup_expired_oidc_states(self, max_age_seconds: int = 300) -> int:
|
||||
from turnstone.core.storage._schema import oidc_pending_states
|
||||
|
||||
cutoff = (datetime.now(UTC) - timedelta(seconds=max_age_seconds)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
|
||||
@@ -4,10 +4,11 @@ from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Row helper
|
||||
@@ -59,6 +60,7 @@ SKILL_MUTABLE = frozenset(
|
||||
"scan_version",
|
||||
"scan_status",
|
||||
"scan_report",
|
||||
"priority",
|
||||
}
|
||||
)
|
||||
STRUCTURED_MEMORY_MUTABLE = frozenset({"content", "description", "type"})
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
"""Add priority column to prompt_templates for skill ordering.
|
||||
|
||||
Multiple ``activation="default"`` skills are concatenated in priority
|
||||
order (ascending), falling back to name for ties. Previously the only
|
||||
ordering lever was the skill name itself.
|
||||
|
||||
Revision ID: 024
|
||||
Revises: 023
|
||||
Create Date: 2026-03-21
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "024"
|
||||
down_revision = "023"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"prompt_templates",
|
||||
sa.Column("priority", sa.Integer, nullable=False, server_default="0"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("prompt_templates", "priority")
|
||||
@@ -0,0 +1,29 @@
|
||||
"""Add index on conversations.timestamp for search_history_recent.
|
||||
|
||||
The ``search_history_recent`` query orders by ``timestamp DESC`` without
|
||||
an index, causing a full table scan. This migration adds the missing
|
||||
index to match the schema definition in ``_schema.py``.
|
||||
|
||||
Revision ID: 025
|
||||
Revises: 024
|
||||
Create Date: 2026-03-21
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "025"
|
||||
down_revision = "024"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_index(
|
||||
"idx_conversations_timestamp",
|
||||
"conversations",
|
||||
["timestamp"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("idx_conversations_timestamp", table_name="conversations")
|
||||
@@ -10,19 +10,19 @@ from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import subprocess
|
||||
import threading
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.safety import is_command_blocked, sanitize_command
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
log = get_logger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constants
|
||||
@@ -392,6 +392,8 @@ class WatchRunner:
|
||||
|
||||
def _run_command(self, command: str) -> tuple[str, int]:
|
||||
"""Run a shell command and return (stdout, exit_code)."""
|
||||
from turnstone.core.env import scrubbed_env
|
||||
|
||||
try:
|
||||
proc = subprocess.run(
|
||||
command,
|
||||
@@ -400,6 +402,7 @@ class WatchRunner:
|
||||
text=True,
|
||||
timeout=self._tool_timeout,
|
||||
start_new_session=True,
|
||||
env=scrubbed_env(),
|
||||
)
|
||||
output = proc.stdout
|
||||
if proc.stderr:
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
"""Pluggable web search backends.
|
||||
|
||||
``web_search`` is an abstract capability with swappable clients:
|
||||
|
||||
* **TavilyClient** — paid, high quality, requires API key
|
||||
* **DuckDuckGoClient** — free, no API key, uses ``duckduckgo-search``
|
||||
|
||||
Auto-detection (default): Tavily if key present, else DDG if installed,
|
||||
else ``None`` (tool removed from tool list).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Protocol
|
||||
|
||||
import httpx
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from turnstone.core.mcp_client import MCPClientManager
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
|
||||
class WebSearchClient(Protocol):
|
||||
"""Minimal interface for a web search backend."""
|
||||
|
||||
def search(self, query: str, max_results: int = 5, **kwargs: Any) -> str:
|
||||
"""Run a search and return formatted markdown results."""
|
||||
...
|
||||
|
||||
|
||||
class TavilyClient:
|
||||
"""Tavily search backend (paid, requires API key)."""
|
||||
|
||||
def __init__(self, api_key: str, timeout: float = 120) -> None:
|
||||
self._api_key = api_key
|
||||
self._timeout = timeout
|
||||
|
||||
def search(self, query: str, max_results: int = 5, **kwargs: Any) -> str:
|
||||
topic = kwargs.get("topic", "general")
|
||||
resp = httpx.post(
|
||||
"https://api.tavily.com/search",
|
||||
json={
|
||||
"query": query,
|
||||
"max_results": max_results,
|
||||
"topic": topic,
|
||||
"include_answer": True,
|
||||
},
|
||||
headers={"Authorization": f"Bearer {self._api_key}"},
|
||||
timeout=self._timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return _format_tavily(data, query)
|
||||
|
||||
|
||||
class DuckDuckGoClient:
|
||||
"""DuckDuckGo search backend (free, no API key)."""
|
||||
|
||||
def __init__(self, timeout: float = 120) -> None:
|
||||
self._timeout = timeout
|
||||
|
||||
def search(self, query: str, max_results: int = 5, **kwargs: Any) -> str:
|
||||
from duckduckgo_search import DDGS # type: ignore[import-not-found]
|
||||
|
||||
with DDGS(timeout=int(self._timeout)) as ddgs:
|
||||
raw = list(ddgs.text(query, max_results=max_results))
|
||||
return _format_ddg(raw, query)
|
||||
|
||||
|
||||
class MCPSearchClient:
|
||||
"""Delegates web_search to an MCP server tool."""
|
||||
|
||||
def __init__(self, mcp_client: MCPClientManager, tool_name: str, timeout: float = 120) -> None:
|
||||
self._mcp = mcp_client
|
||||
self._tool = tool_name
|
||||
self._timeout = timeout
|
||||
|
||||
def search(self, query: str, max_results: int = 5, **kwargs: Any) -> str:
|
||||
import math
|
||||
|
||||
args: dict[str, Any] = {"query": query}
|
||||
if max_results != 5:
|
||||
args["max_results"] = max_results
|
||||
topic = kwargs.get("topic")
|
||||
if topic:
|
||||
args["topic"] = topic
|
||||
return self._mcp.call_tool_sync(self._tool, args, timeout=max(1, math.ceil(self._timeout)))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Result formatters
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _format_tavily(data: dict[str, Any], query: str) -> str:
|
||||
parts: list[str] = []
|
||||
answer = (data.get("answer") or "").strip()
|
||||
if answer:
|
||||
parts.append(f"Answer: {answer}")
|
||||
results = data.get("results") or []
|
||||
if results:
|
||||
lines = []
|
||||
for i, r in enumerate(results, 1):
|
||||
title = r.get("title", "")
|
||||
url = r.get("url", "")
|
||||
content = (r.get("content") or "")[:500]
|
||||
lines.append(f"{i}. [{title}]({url})\n {content}")
|
||||
parts.append("\n".join(lines))
|
||||
return "\n\n".join(parts) if parts else f"No results for '{query}'."
|
||||
|
||||
|
||||
def _format_ddg(results: list[dict[str, Any]], query: str) -> str:
|
||||
if not results:
|
||||
return f"No results for '{query}'."
|
||||
lines = []
|
||||
for i, r in enumerate(results, 1):
|
||||
title = r.get("title", "")
|
||||
url = r.get("href", "")
|
||||
body = (r.get("body") or "")[:500]
|
||||
lines.append(f"{i}. [{title}]({url})\n {body}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Resolver
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _ddg_available() -> bool:
|
||||
"""Check if duckduckgo-search is installed."""
|
||||
try:
|
||||
import duckduckgo_search # noqa: F401
|
||||
|
||||
return True
|
||||
except ImportError:
|
||||
return False
|
||||
|
||||
|
||||
def resolve_web_search_client(
|
||||
backend: str,
|
||||
tavily_key: str | None,
|
||||
mcp_client: Any | None = None,
|
||||
timeout: float = 120,
|
||||
) -> WebSearchClient | None:
|
||||
"""Return a search client based on configuration, or None if unavailable.
|
||||
|
||||
Args:
|
||||
backend: ``""`` (auto), ``"tavily"``, ``"ddg"``, or ``"mcp:server:tool"``
|
||||
tavily_key: Tavily API key (None if not configured)
|
||||
mcp_client: MCPClientManager instance (for MCP backends)
|
||||
timeout: HTTP/tool timeout in seconds
|
||||
"""
|
||||
if backend == "tavily":
|
||||
if tavily_key:
|
||||
return TavilyClient(tavily_key, timeout=timeout)
|
||||
return None
|
||||
|
||||
if backend == "ddg":
|
||||
if _ddg_available():
|
||||
return DuckDuckGoClient(timeout=timeout)
|
||||
return None
|
||||
|
||||
if backend.startswith("mcp:"):
|
||||
parts = backend.split(":", 2)
|
||||
if len(parts) == 3 and mcp_client is not None:
|
||||
_, server, tool = parts
|
||||
prefixed = f"mcp__{server}__{tool}"
|
||||
if mcp_client.is_mcp_tool(prefixed):
|
||||
return MCPSearchClient(mcp_client, prefixed, timeout=timeout)
|
||||
return None
|
||||
|
||||
if backend == "":
|
||||
# Auto-detect: Tavily > DDG > None
|
||||
if tavily_key:
|
||||
return TavilyClient(tavily_key, timeout=timeout)
|
||||
if _ddg_available():
|
||||
return DuckDuckGoClient(timeout=timeout)
|
||||
return None
|
||||
|
||||
log.warning("Unknown web_search_backend %r — web search disabled", backend)
|
||||
return None
|
||||
@@ -78,7 +78,7 @@ class WorkstreamManager:
|
||||
self,
|
||||
session_factory: _SessionFactory,
|
||||
*,
|
||||
max_workstreams: int = 10,
|
||||
max_workstreams: int = 50,
|
||||
node_id: str | None = None,
|
||||
):
|
||||
"""
|
||||
@@ -107,6 +107,11 @@ class WorkstreamManager:
|
||||
self._evictions: int = 0
|
||||
self._last_evicted: Workstream | None = None
|
||||
|
||||
@property
|
||||
def max_workstreams(self) -> int:
|
||||
"""Configured maximum concurrent workstreams."""
|
||||
return self._max_workstreams
|
||||
|
||||
@property
|
||||
def eviction_count(self) -> int:
|
||||
"""Number of workstreams auto-evicted by ``create()``."""
|
||||
@@ -225,6 +230,9 @@ class WorkstreamManager:
|
||||
@staticmethod
|
||||
def _cleanup_ui(ws: Workstream) -> None:
|
||||
"""Unblock pending approval/plan/foreground events on a workstream."""
|
||||
# Cancel any in-flight generation so the worker thread stops promptly.
|
||||
if ws.session and hasattr(ws.session, "cancel"):
|
||||
ws.session.cancel()
|
||||
if ws.ui:
|
||||
if hasattr(ws.ui, "_approval_event"):
|
||||
ws.ui._approval_result = False, None # type: ignore[attr-defined]
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user