Compare commits

...

40 Commits

Author SHA1 Message Date
Patrick Buckley c2c6689a8e chore: bump version to 0.7.0 2026-03-15 17:15:27 -07:00
Patrick Buckley f5af4875ba fix: surface MCP server errors in admin UI instead of silent logging (#100)
* fix: surface MCP server errors in admin UI instead of silent logging

get_server_status() hardcoded error="" — connection and refresh failures
were logged but never surfaced to the admin panel.

Added _last_error dict to MCPClientManager: set on failure (connect,
refresh, periodic refresh, notification handler), cleared on success,
cleaned up on remove. Read in get_server_status().

Admin UI: error tooltip on list row status span, error text in red
in detail modal per-node list. Schema already had the field.

6 new tests for error tracking lifecycle.

* feat: add turnstone_mcp_server_errors Prometheus gauge

Exposes the count of MCP servers currently in error state via
/metrics for alerting and reliability tracking.

* fix: address copilot review — sanitize error strings, clear on notification success

- Add _set_error() helper: strips newlines, truncates to 256 chars
- All error-setting sites now use _set_error() for consistent sanitization
- Notification handler clears _last_error on successful refresh (fixes
  stale error for push-notification servers that skip _periodic_refresh)
2026-03-15 17:12:13 -07:00
Patrick Buckley 0d77d65266 test: add OIDC handler integration tests (22 tests) (#99)
TestClient-based integration tests for the 4 OIDC HTTP endpoints: authorize, callback, admin list identities, admin delete identity.

Uses real SQLite storage with mocked external OIDC calls (exchange_code, validate_id_token, provision_oidc_user) to exercise the full handler→module→storage contract. Covers happy paths, error flows, rate limiting, JWKS key rotation retry, and state expiration.
2026-03-15 16:44:44 -07:00
Patrick Buckley c11991819e fix: apt-get upgrade in Dockerfile to resolve CVE-2026-0861
Trivy scan fails on HIGH for libc-bin/libc6 (2.41-12+deb13u1).
The fix (2.41-12+deb13u2) is available in Debian repos but the
base python:3.14-slim image hasn't been rebuilt yet. Adding
apt-get upgrade pulls in all pending security patches at build time.
2026-03-15 16:40:03 -07:00
Patrick Buckley fc948d711d fix: use pgautoupgrade for seamless postgres major version upgrades
Replaces postgres:18-alpine with pgautoupgrade/pgautoupgrade:18-alpine
in compose.yaml. Sets PGDATA=/var/lib/postgresql/data so pgautoupgrade
detects existing pg17 data and runs pg_upgrade automatically on first
start. No manual migration needed.

Also increases healthcheck start_period to 30s to accommodate the
one-time upgrade process.
2026-03-15 16:26:19 -07:00
renovate[bot] d9722f3578 chore(config): migrate config .github/renovate.json (#97)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-03-15 16:24:39 -07:00
renovate[bot] f494553020 chore(deps): update helm release redis to v25 (#93)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-03-15 15:58:41 -07:00
renovate[bot] c766c81f25 chore(deps): update helm release postgresql to v18 (#92)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-03-15 15:58:38 -07:00
renovate[bot] 3b746fb28d chore(deps): update docker images (#90)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-03-15 15:58:34 -07:00
renovate[bot] b82fa4923c chore(deps): lock file maintenance (#94)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-03-15 15:56:18 -07:00
renovate[bot] 4ac316dc0e chore(deps): update github actions (#91)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-03-15 15:56:16 -07:00
renovate[bot] 387ef06da7 chore(deps): update helm release redis to ~20.13.0 (#89)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-03-15 15:56:13 -07:00
renovate[bot] e4e2200c33 chore(deps): update helm release postgresql to ~16.7.0 (#88)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-03-15 15:56:11 -07:00
renovate[bot] 8b94d553e4 chore(deps): update docker images (#87)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-03-15 15:55:53 -07:00
renovate[bot] cf16724137 chore(deps): pin dependencies (#86)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-03-15 15:53:03 -07:00
Patrick Buckley 22402e89de feat: add dependency management with Renovate, uv.lock, and security … (#83)
* feat: add dependency management with Renovate, uv.lock, and security scanning

Adds automated dependency update detection and vulnerability scanning
across all dependency layers (Python, vendored JS, TypeScript SDK, Docker,
GitHub Actions).

- Renovate config with 10 package groups and custom regex managers for
  vendored JS (KaTeX, Highlight.js, Mermaid) tracking via npm registry
- uv.lock for reproducible builds (80 packages)
- Dockerfile switched to uv sync --frozen with layer caching
- CI: pip-audit (via lock file), npm audit, lock-check jobs
- CI: lint job uses pre-commit for ruff version consistency
- Docker security scan workflow (weekly Trivy, HIGH/CRITICAL)
- Helper script for vendored JS library updates

* fix: resolve CI failures and address review feedback

- Update pre-commit hooks: ruff v0.9.10 -> v0.15.6 (fixes deprecated
  UP038 rule), mypy v1.14.1 -> v1.19.1
- Add per-file-ignore for N802 on sandbox.py (ast visitor convention)
- Fix pip-audit: install into uv venv so uv run can find it
- Pin uv-version in CI to match lock file generator (0.9.18)
- Upgrade vitest ^2.0 -> ^4.1 to fix esbuild GHSA-67mh-4wv8-2f99
- Vendored JS script: use grep -rl for auto-discovery of version refs
  (catches docs/architecture.md), fix LICENSE comment, portable grep
2026-03-15 15:46:48 -07:00
Patrick Buckley e7743fd079 feat: per-tool "Always" approve instead of blanket auto-approve (#82)
* feat: per-tool "Always" approve instead of blanket auto-approve

Interactive "Always" button now adds specific tool names to
auto_approve_tools instead of setting blanket auto_approve=True.
Only the tool types in the current batch are auto-approved going
forward — new tool types still prompt for approval.

Server uses approval_label (with func_name fallback) matching the
existing approve_tools() lookup. CLI and bridge use func_name.
Budget override excluded from all paths.

UI: dashed border on Always button signals persistent action,
dynamic tooltip/badge show tool names, aria-label for screen
readers, focus-visible outline fix, overflow-wrap on badge.

Bridge: seeds with DEFAULT_SAFE_TOOLS on first "always" to avoid
losing existing safe-tool auto-approvals.

16 new tests (10 unit + 6 TestClient integration). Updated tool
pipeline diagram and docs.

* fix: address copilot review — filter errored items, hide Always on budget-only

- Server/bridge/JS: add `not it.get("error")` filter so policy-denied
  items aren't added to auto_approve_tools
- Hide Always button when no eligible tools (budget-override-only batch)
- Docs: clarify CLI/bridge use func_name (coarser MCP granularity)
2026-03-15 15:25:40 -07:00
Patrick Buckley 27349e1c13 refactor: move bridge content buffer to server-side single source of truth
Eliminate dual accumulation by piggybacking assistant response text on
the server's ws_state:idle SSE event.  The bridge no longer maintains
its own _ws_content_buffer — it reads content directly from the idle
event and passes it through to TurnCompleteEvent unchanged.

Server-side: WebUI accumulates tokens in on_content_token(), joins and
includes in the idle broadcast, then resets (with 256 KB cap).

Downstream consumers (Discord bidi DM forwarding, catch-up) are
unaffected — TurnCompleteEvent.content is still populated.
2026-03-15 15:04:22 -07:00
Patrick Buckley 37a48bb30d fix: validate scope_id requires scope in memory API (#80)
* fix: validate scope_id requires scope in memory API

Prevent misleading scope_id usage: reject scope_id with global scope,
require scope when scope_id is provided, require scope_id for
workstream/user scopes on writes. Belt-and-suspenders guard in storage
backends ignores scope_id when scope is empty.

* fix: strip whitespace in scope validation, relax user scope_id requirement

Address Copilot review: .strip() whitespace-only values in all three
validation helpers; SaveMemoryRequest no longer requires scope_id for
user scope since the server auto-resolves it from auth context.
2026-03-15 14:28:03 -07:00
Patrick Buckley 730f5704ff fix: inject prompt template guardrails into plan agent system message (#79)
* fix: inject prompt template guardrails into plan agent system message

Safety/behavioral templates were silently bypassed by the plan agent,
which only used _PLAN_IDENTITY. Now _plan_system_content() prepends
_template_content (when present) so admin-configured guardrails apply
to both _exec_plan and _refine_plan, matching the task agent pattern.

* fix: address Copilot review — log truncation, comment clarity, test robustness

- Log warning on template truncation in _plan_system_content() for
  consistency with _init_system_messages()
- Clarify comment that prior plan pairs (not general history) are forwarded
- Use ChatSession._PLAN_IDENTITY for index assertions instead of substring
2026-03-15 14:24:10 -07:00
Patrick Buckley 9b3b1c1ddd fix: reorder new-workstream modal so Task is the primary field (#77)
* fix: reorder new-workstream modal so Task is the primary field

Users were typing their prompt into the Name field (first text input,
auto-focused) and leaving Task empty, creating idle workstreams. Move
Task textarea to the top of the form, auto-focus it, and add
Ctrl/Cmd+Enter submit shortcut. Accessibility fixes: cancel button
focus-visible, label-hint contrast raised to WCAG AA, platform-aware
keyboard hint, Ctrl+Enter added to shortcuts overlay.

* fix: Enter on Cancel button no longer triggers submit

Copilot review caught that pressing Enter while focused on the Cancel
button bypassed native click and called submitNewWs(). Skip the
Enter-to-submit handler for BUTTON elements so native activation fires.
Also make keyboard shortcuts overlay platform-aware (Ctrl vs ⌘).
2026-03-15 14:23:57 -07:00
Patrick Buckley a9ed8a954b fix: convert _pending_nudge from single-slot to list for defensive correctness
Every append site immediately drains via _init_system_messages(), so this
is defensive — ensures multiple nudges survive if the drain flow is ever
refactored to batch calls.
2026-03-15 14:17:34 -07:00
Patrick Buckley 1efcbcf2ba perf: parallelize _collect_mcp_status and _notify_nodes_mcp_reload wi… (#73)
* perf: parallelize _collect_mcp_status and _notify_nodes_mcp_reload with asyncio.gather

Both functions queried cluster nodes sequentially, making latency
O(N × timeout). Use asyncio.gather to query all nodes concurrently,
matching the existing admin_list_watches pattern. Also reuse the
shared proxy_client instead of creating throwaway httpx clients per
node, and add debug logging on MCP status fetch failures.

* perf: bound node fan-out concurrency and improve debug logging

Add _NODE_FAN_OUT_LIMIT (50) semaphore to all three gather fan-out
sites (_collect_mcp_status, _notify_nodes_mcp_reload, admin_list_watches)
to cap concurrent outbound connections below the httpx pool limit,
leaving headroom for other proxy traffic at 1000-node scale.

Add exc_info=True to all debug log calls for actionable diagnostics.

* test: add unit tests for _collect_mcp_status and _notify_nodes_mcp_reload

11 tests covering success, non-200, missing URL, exceptions, empty
cluster, and mixed multi-node scenarios for both fan-out helpers.
2026-03-15 13:54:49 -07:00
Patrick Buckley 1d36d80fe5 fix: reduce metacognition false positives with strong/weak pattern tiers (#76)
* fix: reduce metacognition false positives with strong/weak pattern tiers

Correction detection: split "no" handling — "no," and "no." are strong
(always fire), "no <word>" uses an allowlist of correction-context words
(pronouns, demonstratives, verbs) instead of a blocklist. Phrases like
"no problem", "no worries", "no rush" are excluded automatically.

Completion detection: move most patterns to weak tier, gated by message
length (<80 chars) and absence of continuation markers ("?", "can you",
"but", "now", "please", etc.). "thanks for X" excluded at regex level.
Strong tier (always fire): "that's all", "lgtm".

* fix: align allowlist comment with implementation (include articles)
2026-03-15 13:53:39 -07:00
Patrick Buckley e603a6a7d1 fix(oidc): pin redirect URI via TURNSTONE_OIDC_REDIRECT_BASE env var (#74)
* fix(oidc): pin redirect URI via TURNSTONE_OIDC_REDIRECT_BASE env var

OIDC redirect_uri was derived from the request Host header, which is
unreliable behind reverse proxies. Add TURNSTONE_OIDC_REDIRECT_BASE
(env var / config.toml) to pin the externally-reachable origin.

Extract _build_oidc_redirect_uri() helper to deduplicate the authorize
and callback handlers. Validate redirect_base at load time (must be
scheme://host[:port], rejects paths/query strings/invalid schemes).

* fix(oidc): reject redirect_base with missing hostname

Addresses Copilot review: values like `https://` or `https://:443`
passed validation but would produce invalid redirect URIs.

* fix(oidc): reject redirect_base with userinfo or invalid port

Addresses Copilot round 2: urlparse silently accepts user:pass@host
and non-numeric ports. Now explicitly rejects both.
2026-03-15 13:52:05 -07:00
Patrick Buckley 2e95f2ac73 test: add scope coverage for internal MCP/config reload endpoints (#75)
* test: add scope coverage for internal MCP/config reload endpoints

Verify required_scope() returns "approve" for _internal endpoints
across all access patterns (bare, /v1/-prefixed, console proxy with
and without /v1/), plus a GET negative test confirming only POST is
elevated. Closes the "internal endpoints accept read scope" item in
PROGRESS.md — the endpoints were already in APPROVE_PATHS.

* test: add config-reload v1/proxy scope tests per review feedback

Add /v1/-prefixed and console proxy variants for config-reload to
match the mcp-reload coverage, as flagged by Copilot review.
2026-03-15 13:45:35 -07:00
Patrick Buckley 5f27ed9fca feat: OIDC identity management inline in Users admin tab (#72)
Expandable user rows in the console Users tab reveal OIDC identities
linked to each user. Issuer badge, truncated subject, email, relative
last-login time, and unlink action with confirmation modal + audit trail.

Keyboard accessible (tabindex, Enter/Space, aria-expanded, focus-visible).
In-place refresh after unlink (no close/reopen flicker). Audit captures
user_id before delete. Mobile responsive (3-column at <700px).
Reduced-motion support. 2 new admin API endpoints reusing admin.users
permission and existing storage methods.
2026-03-15 03:52:42 -07:00
Patrick Buckley 20df7b3034 feat: OIDC SSO authentication with PKCE, auto-provisioning, and role … (#71)
* feat: OIDC SSO authentication with PKCE, auto-provisioning, and role mapping

Add OpenID Connect as a fourth authentication method, enabling single sign-on
via any OIDC provider (Okta, Azure AD, Google, Keycloak). Opt-in via env vars
(TURNSTONE_OIDC_ISSUER, CLIENT_ID, CLIENT_SECRET).

Security:
- Authorization Code Flow with PKCE (S256)
- State/nonce parameters with database-backed pending store (multi-node safe)
- JWKS signature validation with async fetch + key rotation retry
- Algorithm allowlist from JWKS key (not token header) prevents confusion
- Identity matching exclusively by (issuer, sub) — prevents account takeover
- password_enabled=false enforced server-side, not just UI
- Rate limiting on both authorize and callback endpoints
- OIDC users get "!oidc" password sentinel (bcrypt rejects naturally)
- ID token validated for iss, aud, exp, nonce

Features:
- Auto-provisioning with username deduplication on first login
- Claim-based role mapping with IdP demotion propagation (revokes stale roles)
- "Continue with [Provider]" SSO button on login page
- OIDC-only mode hides password form
- Setup wizard required before OIDC login (admin bootstrap)

Storage: migration 018 (oidc_identities + oidc_pending_states tables),
8 new protocol methods on both SQLite and PostgreSQL backends.
66 new tests (2273 total).

* fix: address PR #71 review feedback (18 items)

Bugs fixed:
- OIDC success redirect now fetches permissions via new /auth/whoami
  endpoint before completing login (fixes permission-gating in UI)
- Remove double decodeURIComponent on oidc_error (URLSearchParams
  already decodes; extra call throws on stray %)
- Authorize rate limiter returns redirect instead of JSON 429
  (endpoint reached via browser navigation, not fetch)
- Lazy JWKS fetch in callback when startup discovery failed (IdP
  recovery without restart)
- Startup exception handlers now log with exc_info=True
- PostgreSQL pop_oidc_pending_state uses DELETE...RETURNING for
  true atomicity (eliminates TOCTOU)

Behavior:
- New OIDC users without role mapping get builtin-viewer by default
  (assigned_by="oidc-default", not revoked by role sync)

Documentation fixes:
- Role mapping: sync semantics (add + revoke stale), not "additive only"
- PASSWORD_ENABLED=false blocks ALL password logins including admin
- Algorithm: asymmetric allowlist, not per-key derivation
- PlantUML diagram updated for role revocation

API spec fixes:
- Removed error_codes=[302] from callback (302 is success redirect)
- Added /auth/whoami to both server + console specs
- Regenerated TypeScript SDK OpenAPI snapshots (23 + 51 paths)

* fix: address PR #71 round 2 review feedback (10 items)

Rate limiting:
- Authorize endpoint now calls record() after check() so the rate
  limiter actually counts attempts (was a no-op before)

OIDC resilience:
- Split startup try/except: discovery failure disables OIDC, JWKS
  prefetch failure leaves OIDC enabled for lazy retry on first login
- JWKS unavailable message changed to "temporarily unavailable"
  (was misleadingly "not configured")
- create_oidc_pending_state raises on collision instead of OR IGNORE
  (prevents silent insert drop on state collision)
- SQLite pop_oidc_pending_state uses BEGIN IMMEDIATE for write lock
  (eliminates TOCTOU race)

Frontend:
- OIDC error display deferred 300ms so showLogin()'s async status
  fetch doesn't clear it via _switchMode → _clearError

API spec:
- OIDC authorize/callback endpoints now declare response_code=302
- Added AuthWhoamiResponse Pydantic model for /auth/whoami
- Regenerated TypeScript SDK OpenAPI snapshots

Documentation:
- Diagram: JWKS "cached at startup, refreshed on-demand" (was "hourly")
- Added TODO(tech-debt) comments on Host header redirect_uri sites
2026-03-15 03:44:18 -07:00
Patrick Buckley 68c991fbdd fix: restore safe HTML element rendering and suppress plantuml warning (#70)
* fix: restore safe HTML element rendering and suppress plantuml warning

- Add safe HTML tag allowlist in inlineMarkdown: br, hr, kbd, mark,
  sub, sup, ins, wbr, details, summary, abbr, small, u, s
  (attribute-free only — XSS safe, tags with attributes stay escaped)
- Add <details>/<summary> block-level protection pass with recursive
  markdown rendering of inner content
- Add plantuml to _NO_HIGHLIGHT_LANGS (suppresses highlight.js warning
  for unsupported language)
- CSS for details (collapsible, overflow hidden), kbd (mono font,
  key style), mark (yellow-glow token for theme adaptation)

* fix: restrict safe tags to inline-only, broaden details regex

- Remove hr, details, summary from inline _SAFE_TAGS allowlist (they
  are block-level and produce invalid HTML inside <p> wrappers)
- Make <details> regex newline-optional so same-line
  <details><summary>Title</summary> patterns are captured
2026-03-15 02:34:03 -07:00
Patrick Buckley 376da3d084 feat: prompt template tech debt — tests, read-only endpoints, double-… (#67)
* feat: prompt template tech debt — tests, read-only endpoints, double-load fix, server creation modal

Close test coverage gaps for prompt templates:
- Resume with deleted template: verifies graceful degradation (template_content=None, warning logged)
- Threading safety: concurrent set_template/init_system_messages with no race conditions
- Factory passthrough: template kwarg propagation through WorkstreamManager.create()

Add read-only template listing endpoints (read scope, no content exposed):
- GET /v1/api/templates — prompt template summaries (name, category, is_default, origin)
- GET /v1/api/ws-templates — enabled workstream template summaries (name, description, model)
- Available on both server and console; Python + TypeScript SDK methods added
- Console creation modal switched from admin endpoint to read-scope endpoint

Eliminate double-load inefficiency in workstream creation:
- Template validation moved before mgr.create() (no create-then-rollback on invalid template)
- template kwarg plumbed through WorkstreamManager.create() and session factory
- _SessionFactory Protocol added for proper mypy typing

Add workstream creation modal to server web UI:
- Name, model, template dropdown, ws_template/profile dropdown
- Instrument panel aesthetic: gradient top border, blur backdrop, amber accent
- Focus trap, Escape/Enter keyboard handling, loading state, error display
- WCAG AA contrast compliance, reduced-motion support

* fix: add list_ws_templates SDK methods + regenerate OpenAPI snapshots

Add list_ws_templates() to Python SDK (async + sync) and listWsTemplates()
to TypeScript SDK for the new GET /v1/api/ws-templates server endpoint.
Add WsTemplateSummary + ListWsTemplateSummaryResponse TypeScript types.
Regenerate openapi-server.json and openapi-console.json snapshots.

Addresses Copilot review feedback on PR #67.

* fix: skip template pre-validation when resuming a workstream

When resume_ws is set, the request's template field is irrelevant —
resume() restores the template from workstream_config. Pre-validating
a stale template name would incorrectly return 400 before the resume
even runs.

Addresses Copilot review feedback on PR #67.
2026-03-15 02:09:31 -07:00
Patrick Buckley e2a199c9c3 feat: mermaid diagram rendering with lazy loading and theme integration (#69)
* feat: mermaid diagram rendering with lazy loading and theme integration

Integrate mermaid.js 11.13.0 (self-hosted, MIT, ~2.9MB) for rendering
```mermaid code blocks as inline SVG diagrams. Covers flowcharts,
sequence, class, ER, state, gantt, pie, timeline, and mindmap.

- Lazy-loaded via dynamic script injection on first mermaid block
  detection (not eagerly loaded on every page view)
- 3-state loader (idle/loading/ready) with callback queue
- Serialized rendering to avoid mermaid internal state corruption
- Theme integration via getComputedStyle reading CSS design tokens;
  re-renders all diagrams on dark/light theme toggle
- Source preserved in data-mermaid-source for theme re-rendering
- Error handling with source code fallback display
- securityLevel: "strict" (DOMPurify) for SVG XSS prevention
- THIRD-PARTY-NOTICES updated with mermaid MIT license

* fix: mermaid render fixes from Copilot review

- Call result.bindFunctions(container) after SVG insertion for
  interactive diagram elements (click handlers, links, tooltips)
- Clear mermaid-error class on successful render (fixes stale error
  styling after theme toggle re-render)
- Clear mermaid-error in reRenderAllMermaid before re-render sequence
- Restructure postRenderMarkdown so mermaid rendering runs even when
  highlight.js is unavailable (hljs guard changed from early return
  to conditional block)
2026-03-15 02:09:10 -07:00
Patrick Buckley 4152ea2352 fix: widen code fence regex and skip auto-detect on unlabeled blocks
- Regex changed from (\w*) to ([^\s`]*) to capture language names with
  special chars (c++, c#, objective-c, shell-session)
- Alias map normalizes c++ → cpp, c# → csharp, f# → fsharp for CSS
  class names
- Empty language no longer emits class="language-", preventing
  highlight.js auto-detect across all 37 bundled languages on
  unlabeled code blocks (performance fix for large blocks)
2026-03-15 02:04:58 -07:00
Patrick Buckley 44cc14b46f feat: syntax highlighting via highlight.js with code block variants (#68)
Integrate highlight.js 11.11.1 (self-hosted, BSD-3-Clause, ~125KB) for
language-aware syntax highlighting on fenced code blocks.

- postRenderMarkdown() hook applies highlighting at stream_end and
  history load — not during streaming (innerHTML replaced per token)
- Custom theme using CSS design tokens (auto-adapts dark/light)
- Code block variants: diff (green/red line coloring), bash/shell
  (terminal left-border), ascii/text/plaintext (no highlighting)
- Class prefix changed from lang- to language- (CommonMark standard)
- Graceful degradation when highlight.js unavailable
- THIRD-PARTY-NOTICES file for bundled dependency attribution
- pyproject.toml package-data glob for vendored hljs directory
2026-03-15 01:36:57 -07:00
Patrick Buckley 83b0cde32f feat: GFM extended syntax renderers (callouts, footnotes, definition … (#66)
* feat: GFM extended syntax renderers (callouts, footnotes, definition lists)

Add three GFM extended syntax features to the server web UI markdown
renderer, with no external library dependencies (pure JS/CSS):

- Callouts/Alerts: > [!NOTE], [!TIP], [!IMPORTANT], [!WARNING], [!CAUTION]
  with color-coded left borders, icons, and recursive markdown body
- Definition Lists: Term + `: Definition` pattern with multi-term support
- Footnotes: [^id] inline superscript references, [^id]: definitions
  collected into a numbered section with bidirectional navigation

Design review fixes: scoped footnote IDs (prevent collisions across
messages), aria-hidden on callout icons, aria-label on callout containers,
focus-visible on footnote links, smooth-scroll footnote navigation.

* fix: use getElementById for footnote scroll to handle special chars in IDs

querySelector throws on fragment IDs containing &, . or : characters
(produced by escapeHtml on footnote labels). getElementById accepts any
string and is the correct API for ID-based element lookup.
2026-03-15 01:33:54 -07:00
Patrick Buckley 2ef8a8711b feat: rich markdown renderer with LaTeX support for server web UI (#65)
* feat: rich markdown renderer with LaTeX support for server web UI

Extract markdown rendering from app.js into dedicated renderer.js with
full GFM support: tables (alignment, hover, striping), nested lists,
task list checkboxes, nested blockquotes, images (click-to-load for
privacy), and inline/display LaTeX math via self-hosted KaTeX 0.16.38.

Security: escape image/link URLs to prevent attribute injection, block
javascript: scheme in links, add rel="noopener noreferrer", images
require explicit click to load (no automatic external requests).

Accessibility: scope="col" on table headers, tabindex on scrollable
table containers, aria-labels on task checkboxes and image placeholders,
KaTeX error color override for WCAG AA contrast, reduced-motion support.

* fix: address code review — XSS hardening and list type splitting

- Escape all text through escapeHtml() at start of inlineMarkdown()
  so only renderer-generated tags appear in innerHTML (prevents raw
  HTML/script injection from LLM output)
- Replace inline onclick handler on image placeholders with data-*
  attributes and delegated DOM event listeners (prevents entity
  decoding XSS in event handler attributes)
- Split list blocks into separate <ul>/<ol> when marker type changes
  at the same indent level (mixed ordered/unordered sequences)
2026-03-15 00:04:01 -07:00
Patrick Buckley 3658b77de8 feat: Discord content catch-up + bidirectional notification replies (… (#64)
* feat: Discord content catch-up + bidirectional notification replies (#64)

Two improvements to the Discord channel adapter:

1. Fix intermittent dropped responses caused by a race between the
   bridge's two independent SSE connections (global SSE detects idle
   before per-ws SSE delivers all content tokens). The bridge now
   accumulates content in _ws_content_buffer and attaches it to
   TurnCompleteEvent.content. The Discord bot uses this as a catch-up
   when streaming events were missed.

2. Bidirectional notification replies — when the notify tool sends a DM,
   the message is tracked with the originating ws_id. Users can reply to
   the DM and the reply is routed to the workstream. The response is
   forwarded back to the DM, with the response itself tracked for
   multi-turn conversations. Includes user identity verification,
   stale notification feedback, and FIFO-capped tracking (100 entries).

* fix: address Copilot review — re-insert on unlinked user, deque buffer

- Re-insert _notify_ws_map entry when resolve_user returns None so the
  user can retry after linking (same pattern as user-mismatch re-insert)
- Rename _MAX_CONTENT_BUFFER_BYTES → _MAX_CONTENT_BUFFER_CHARS (len()
  returns characters, not bytes)
- Use deque + running total for O(1) popleft instead of list.pop(0)
2026-03-14 23:58:27 -07:00
Patrick Buckley 83577739e0 chore: bump version to 0.6.2
- MCP admin tab: database-backed server management, hot-reload,
  reconcile, unified config view, paste-based import
- Catch-up migration for builtin-admin permissions (017)
- `[all]` optional dependency group (@Burhan-Q)
2026-03-14 17:25:58 -07:00
Burhan 71d13936fe add "all" optional dep (#61) 2026-03-14 17:05:49 -07:00
Patrick Buckley 0cd061196c fix: catch-up migration ensuring builtin-admin has all 20 permissions (#63)
Migrations 011-016 each appended a permission to the builtin-admin role
via conditional UPDATE, but on some deployments these never applied.
Migration 017 idempotently sets the complete permission string rather
than appending incrementally.

Must be merged after feat/admin-mcp-servers (migration 016).
2026-03-14 17:03:21 -07:00
Patrick Buckley 19abc0cc65 feat: admin MCP Servers tab — database-backed MCP server management w… (#62)
* feat: admin MCP Servers tab — database-backed MCP server management with live status

Add MCP Servers admin tab (14th tab, System group) for managing MCP server
definitions via the database instead of static JSON config files.

Storage: `mcp_servers` table (migration 016), 6 CRUD methods on both SQLite
and PostgreSQL backends, `MCP_SERVER_MUTABLE` field allowlist.

Config priority chain: DB rows (if any enabled) → CLI `--mcp-config` →
`mcp.config_path` setting → none. Nodes auto-load from DB on startup via
`load_mcp_config(storage=)`.

Hot-reload: `reconcile_sync(storage)` diffs running servers against DB —
adds missing, removes stale, reconnects changed. `_db_managed` set tracks
DB-sourced servers so config-file servers (MCP_CONFIG env) are never removed
by reconcile. Per-server `AsyncExitStack` for clean teardown.

Reload pattern: console writes to DB then signals nodes via
`POST /_internal/mcp-reload` (update by reference, no config payload).

Console admin API: 7 endpoints under `/v1/api/admin/mcp-servers` (CRUD +
reload + import), `admin.mcp` permission, secret masking (env/headers
replaced with *** unless ?reveal=true), audit log sanitization.

Unified view: tab merges DB-managed servers with config-sourced servers
detected on nodes. Config servers shown as read-only rows with "config"
badge — no edit/delete.

Admin UI: 7-column grid with magenta status dots, transport badges,
single-column create/edit modal, paste-based JSON import (mcpServers format),
detail modal with per-node status. Mobile 3-column collapse, reduced-motion
support, backdrop-click dismiss, focus trapping.

SDKs: 7 methods on Python (async+sync) and TypeScript SDKs.

Also fixes: Settings tab permission gate (admin.users → admin.settings),
_ALL_PERMISSIONS list in governance.js (5 missing permissions added),
_internal/mcp-reload added to APPROVE_PATHS.

Docs: architecture.md (14 tabs), api-reference.md (7 endpoints),
20-mcp-architecture.puml updated with admin-driven lifecycle.

66 new tests (2232 total).

* fix: address Copilot review feedback on MCP admin PR

- Docs: fix "merges both sources" → "first-match-wins priority" (architecture.md)
- Validation: require command for stdio, url for streamable-http transport
- Validation: check args/headers/env types in import handler before storing
- Schema: add transport/command/url to McpServerStatus, source to McpServerDetail
- Thread safety: move all remove_server_sync mutations onto MCP event loop thread
- Regenerate OpenAPI JSON snapshots for TypeScript SDK
2026-03-14 17:02:50 -07:00
137 changed files with 21151 additions and 1253 deletions
+208
View File
@@ -0,0 +1,208 @@
{
"$schema": "https://docs.renovatebot.com/renovate-schema.json",
"extends": [
"config:recommended",
"helpers:pinGitHubActionDigests",
":separateMajorReleases"
],
"labels": [
"dependencies"
],
"prConcurrentLimit": 5,
"prHourlyLimit": 2,
"schedule": [
"before 9am on Monday"
],
"timezone": "America/New_York",
"lockFileMaintenance": {
"enabled": true,
"schedule": [
"before 9am on Monday"
]
},
"customManagers": [
{
"customType": "regex",
"description": "Track vendored KaTeX version",
"managerFilePatterns": [
"/pyproject\\.toml$/"
],
"matchStrings": [
"katex-(?<currentValue>[\\d.]+)/"
],
"depNameTemplate": "katex",
"datasourceTemplate": "npm"
},
{
"customType": "regex",
"description": "Track vendored Highlight.js version",
"managerFilePatterns": [
"/pyproject\\.toml$/"
],
"matchStrings": [
"hljs-(?<currentValue>[\\d.]+)/"
],
"depNameTemplate": "highlight.js",
"datasourceTemplate": "npm"
},
{
"customType": "regex",
"description": "Track vendored Mermaid version",
"managerFilePatterns": [
"/pyproject\\.toml$/"
],
"matchStrings": [
"mermaid-(?<currentValue>[\\d.]+)/"
],
"depNameTemplate": "mermaid",
"datasourceTemplate": "npm"
}
],
"packageRules": [
{
"description": "LLM SDKs — always review manually",
"groupName": "LLM SDKs",
"matchPackageNames": [
"openai",
"anthropic",
"mcp"
],
"schedule": [
"before 9am on Monday"
],
"automerge": false
},
{
"description": "Web framework stack",
"groupName": "Web Framework",
"matchPackageNames": [
"starlette",
"uvicorn",
"sse-starlette",
"httpx",
"httpx-sse",
"pydantic"
],
"schedule": [
"before 9am on Wednesday"
],
"automerge": true,
"matchUpdateTypes": [
"patch"
]
},
{
"description": "Database layer",
"groupName": "Database",
"matchPackageNames": [
"sqlalchemy",
"alembic",
"psycopg"
],
"schedule": [
"before 9am on Wednesday"
],
"automerge": true,
"matchUpdateTypes": [
"patch"
]
},
{
"description": "Security-critical — always review manually",
"groupName": "Security",
"matchPackageNames": [
"PyJWT",
"pyjwt",
"bcrypt"
],
"automerge": false
},
{
"description": "Infrastructure dependencies",
"groupName": "Infrastructure",
"matchPackageNames": [
"structlog",
"redis",
"croniter",
"discord.py"
],
"schedule": [
"before 9am on the first day of the month"
],
"automerge": true,
"matchUpdateTypes": [
"patch"
]
},
{
"description": "Vendored JS — requires manual file download after merge",
"groupName": "Vendored JS",
"matchPackageNames": [
"katex",
"highlight.js",
"mermaid"
],
"schedule": [
"before 9am on the first day of the month"
],
"automerge": false,
"prBodyNotes": [
"This PR updates version references only.",
"After merging, run `scripts/update-vendored-js.sh <lib> <version>` to download the actual files."
]
},
{
"description": "Dev/test tooling",
"groupName": "Tooling",
"matchPackageNames": [
"ruff",
"mypy",
"types-redis",
"pytest",
"pytest-cov",
"pre-commit"
],
"schedule": [
"before 9am on the first day of the month"
],
"automerge": true,
"matchUpdateTypes": [
"patch"
]
},
{
"description": "Docker base images",
"groupName": "Docker Images",
"matchManagers": [
"dockerfile",
"docker-compose"
],
"schedule": [
"before 9am on the first day of the month"
],
"automerge": false
},
{
"description": "TypeScript SDK dev dependencies",
"groupName": "TypeScript SDK",
"matchFileNames": [
"sdk/typescript/**"
],
"schedule": [
"before 9am on the first day of the month"
],
"automerge": true,
"matchUpdateTypes": [
"patch"
]
},
{
"description": "GitHub Actions — group all action updates",
"groupName": "GitHub Actions",
"matchManagers": [
"github-actions"
],
"automerge": false
}
]
}
+48 -12
View File
@@ -10,21 +10,21 @@ jobs:
lint:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
- uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6
with:
python-version: "3.13"
- run: pip install ruff
- run: ruff check turnstone/ tests/
- run: ruff format --check turnstone/ tests/
python-version: "3.14"
- run: pip install pre-commit
# mypy runs separately in typecheck job with full project deps
- run: SKIP=mypy pre-commit run --all-files
typecheck:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
- uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6
with:
python-version: "3.13"
python-version: "3.14"
- run: pip install mypy types-redis
- run: pip install -e ".[mq]"
- run: mypy turnstone/
@@ -35,14 +35,50 @@ jobs:
matrix:
python-version: ["3.11", "3.12", "3.13"]
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
- uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6
with:
python-version: ${{ matrix.python-version }}
- run: pip install -e ".[test,mq]"
- run: pytest tests/ -m "not live" --cov=turnstone --cov-report=term-missing --cov-report=xml -q
- uses: actions/upload-artifact@v4
- uses: actions/upload-artifact@bbbca2ddaa5d8feaa63e36b76fdaad77386f024f # v7
if: always()
with:
name: coverage-${{ matrix.python-version }}
path: coverage.xml
lock-check:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
- uses: astral-sh/setup-uv@e06108dd0aef18192324c70427afc47652e63a82 # v7
with:
uv-version: "0.9.18"
- run: uv lock --check
security:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
- uses: astral-sh/setup-uv@e06108dd0aef18192324c70427afc47652e63a82 # v7
with:
uv-version: "0.9.18"
- uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6
with:
python-version: "3.14"
- run: uv sync --frozen --all-extras
- run: uv pip install pip-audit
- run: uv run pip-audit --strict --desc
security-ts:
runs-on: ubuntu-latest
defaults:
run:
working-directory: sdk/typescript
steps:
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
- uses: actions/setup-node@53b83947a5a98c8d113130e565377fae1a50d02f # v6
with:
node-version: "24"
- run: npm ci
- run: npm audit --audit-level=moderate
+22
View File
@@ -0,0 +1,22 @@
name: Docker Security Scan
on:
push:
branches: [main]
schedule:
- cron: "0 6 * * 1" # Weekly Monday 06:00 UTC
permissions:
contents: read
jobs:
scan:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
- run: docker build -t turnstone:scan .
- uses: aquasecurity/trivy-action@57a97c7e7821a5776cebc9bb87c984fa69cba8f1 # 0.35.0
with:
image-ref: "turnstone:scan"
severity: "HIGH,CRITICAL"
exit-code: "1"
+4 -4
View File
@@ -13,16 +13,16 @@ jobs:
runs-on: ubuntu-latest
environment: pypi
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6
- uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6
with:
python-version: "3.13"
python-version: "3.14"
- run: pip install build
- run: python -m build
- uses: pypa/gh-action-pypi-publish@release/v1
- name: Create GitHub Release
uses: softprops/action-gh-release@v2
uses: softprops/action-gh-release@b25b93d384199fc0fc8c2e126b2d937a0cbeb2ae # v2
with:
generate_release_notes: true
draft: false
+2 -2
View File
@@ -1,13 +1,13 @@
repos:
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.9.10
rev: v0.15.6
hooks:
- id: ruff
args: [--fix]
- id: ruff-format
- repo: https://github.com/pre-commit/mirrors-mypy
rev: v1.14.1
rev: v1.19.1
hooks:
- id: mypy
additional_dependencies: [types-redis>=4.6, redis>=7.2]
+22 -24
View File
@@ -1,41 +1,39 @@
# =============================================================================
# Turnstone — multi-stage Docker build
# Turnstone — Docker build with uv for reproducible, locked installs
# Single image for all services: server, bridge, console, sim, eval
# =============================================================================
# ----------------------------------------------------------------------------
# Stage 1: Builder — build the wheel
# ----------------------------------------------------------------------------
FROM python:3.13-slim AS builder
WORKDIR /build
RUN pip install --no-cache-dir hatchling
COPY pyproject.toml README.md LICENSE ./
COPY turnstone/ turnstone/
RUN pip wheel --no-deps --wheel-dir /build/wheels .
# ----------------------------------------------------------------------------
# Stage 2: Runtime — slim image with the installed package
# ----------------------------------------------------------------------------
FROM python:3.13-slim
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
# System dependencies for psycopg (PostgreSQL client library)
RUN apt-get update && apt-get install -y --no-install-recommends libpq5 \
RUN apt-get update && apt-get upgrade -y && apt-get install -y --no-install-recommends libpq5 \
&& rm -rf /var/lib/apt/lists/*
# Non-root user
RUN useradd --create-home --shell /bin/bash turnstone
# Install the wheel with all optional extras
COPY --from=builder /build/wheels/*.whl /tmp/wheels/
RUN pip install --no-cache-dir "$(ls /tmp/wheels/*.whl)[mq,console,sim,postgres,discord]" \
&& rm -rf /tmp/wheels
WORKDIR /app
# Compile bytecode for faster startup
ENV UV_COMPILE_BYTECODE=1
# Install dependencies first (cached layer — only re-runs when deps change)
COPY pyproject.toml uv.lock README.md LICENSE ./
RUN uv sync --frozen --no-install-project --no-dev \
--extra mq --extra console --extra sim --extra postgres --extra discord --extra anthropic
# Install the project itself
COPY turnstone/ turnstone/
RUN uv sync --frozen --no-dev \
--extra mq --extra console --extra sim --extra postgres --extra discord --extra anthropic
# Add venv to PATH so entry points are found
ENV PATH="/app/.venv/bin:$PATH"
# Health check script (stdlib only, no pip deps needed)
COPY docker/healthcheck.py /usr/local/bin/healthcheck.py
+3 -1
View File
@@ -18,7 +18,7 @@ Turnstone gives LLMs tools — shell, files, search, web, planning — and orche
- **Multi-node clusters** — generic work load-balances across nodes, directed work routes to a specific server
- **Cluster dashboard** — real-time view of all nodes and workstreams, reverse proxy for server UIs
- **Intent validation** — an LLM judge evaluates every tool call before approval, presenting risk assessments and evidence-based recommendations so users can make informed decisions instead of blindly approving raw tool calls
- **Governance & compliance** — RBAC, tool policies, prompt templates, workstream templates, usage tracking, and append-only audit logs
- **Governance & compliance** — RBAC, OIDC SSO (Okta, Azure AD, Google, Keycloak), tool policies, prompt templates, workstream templates, usage tracking, and append-only audit logs
- **Cluster simulator** — test the stack at scale (up to 1000 nodes) without an LLM backend
Works with any OpenAI-compatible API (vLLM, llama.cpp, NVIDIA NIM) or Anthropic's native Messages API. Supports [MCP](https://modelcontextprotocol.io/) for external tool servers with native deferred tool loading on Anthropic and OpenAI APIs (BM25 fallback for local models).
@@ -136,12 +136,14 @@ Detailed UML diagrams are available in [`docs/diagrams/`](docs/diagrams/):
| [Governance Architecture](docs/diagrams/png/19-governance-architecture.png) | RBAC, policies, audit, usage enforcement flow |
| [WS Template Architecture](docs/diagrams/png/21-ws-template-architecture.png) | Workstream template application and lifecycle |
| [Judge Architecture](docs/diagrams/png/22-judge-architecture.png) | Intent validation two-tier evaluation pipeline |
| [OIDC Architecture](docs/diagrams/png/25-oidc-architecture.png) | OIDC SSO authorization code flow with PKCE |
### Governance
Turnstone includes a built-in governance layer for enterprise deployments — manage who can do what, which tools run unattended, and where every token goes.
- **RBAC** — 15 granular permissions, 3 built-in roles (admin / operator / viewer), custom roles, privilege escalation prevention
- **OIDC SSO** — single sign-on via any OpenID Connect provider (Okta, Azure AD, Google, Keycloak); Authorization Code Flow with PKCE, auto-provisioning, claim-based role mapping with demotion propagation; see [docs/oidc.md](docs/oidc.md)
- **Tool policies** — glob-pattern rules (`allow` / `deny` / `ask`) with priority ordering; automate approvals or lock down dangerous tools
- **Prompt templates** — reusable system messages with `{{variable}}` substitution and categories
- **Usage tracking** — per-request token and tool metrics, aggregation by day / model / user, automatic 90-day pruning
+97
View File
@@ -0,0 +1,97 @@
Turnstone — Third-Party Notices
This file contains the licenses and notices for third-party software bundled
with Turnstone. Each bundled dependency retains its original license; the
Turnstone BUSL-1.1 license does not apply to these components.
================================================================================
KaTeX 0.16.38
https://katex.org/
https://github.com/KaTeX/KaTeX
The MIT License (MIT)
Copyright (c) 2013-2020 Khan Academy and other contributors
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
================================================================================
highlight.js 11.11.1
https://highlightjs.org/
https://github.com/highlightjs/highlight.js
BSD 3-Clause License
Copyright (c) 2006, Ivan Sagalaev.
All rights reserved.
Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are met:
1. Redistributions of source code must retain the above copyright notice, this
list of conditions and the following disclaimer.
2. Redistributions in binary form must reproduce the above copyright notice,
this list of conditions and the following disclaimer in the documentation
and/or other materials provided with the distribution.
3. Neither the name of the copyright holder nor the names of its
contributors may be used to endorse or promote products derived from
this software without specific prior written permission.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
================================================================================
Mermaid 11.13.0
https://mermaid.js.org/
https://github.com/mermaid-js/mermaid
The MIT License (MIT)
Copyright (c) 2014-2022 Knut Sveidqvist
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+5 -4
View File
@@ -26,7 +26,7 @@ services:
# PostgreSQL — production database (profile: production)
# -------------------------------------------------------------------
postgres:
image: postgres:17-alpine
image: pgautoupgrade/pgautoupgrade:18-alpine
profiles:
- production
- cluster
@@ -35,6 +35,7 @@ services:
POSTGRES_DB: turnstone
POSTGRES_USER: ${POSTGRES_USER:-turnstone}
POSTGRES_PASSWORD: ${POSTGRES_PASSWORD:?POSTGRES_PASSWORD is required for production profile}
PGDATA: /var/lib/postgresql/data
volumes:
- postgres-data:/var/lib/postgresql/data
networks:
@@ -44,7 +45,7 @@ services:
interval: 5s
timeout: 3s
retries: 5
start_period: 5s
start_period: 30s
deploy:
resources:
limits:
@@ -56,7 +57,7 @@ services:
# Redis — message broker, pub/sub, node registry
# -------------------------------------------------------------------
redis:
image: redis:7.4-alpine
image: redis:8.6-alpine
command:
- sh
- -c
@@ -257,7 +258,7 @@ services:
# docker compose --profile ddgCluster up
# -------------------------------------------------------------------
ddg-search:
image: python:3.13-slim
image: python:3.14-slim
profiles:
- ddgCluster
command:
+2 -2
View File
@@ -7,10 +7,10 @@ appVersion: "0.3.0"
dependencies:
- name: postgresql
version: ~16.0
version: ~18.5.0
repository: https://charts.bitnami.com/bitnami
condition: postgresql.enabled
- name: redis
version: ~20.0
version: ~25.3.0
repository: https://charts.bitnami.com/bitnami
condition: redis.enabled
+79
View File
@@ -618,6 +618,67 @@ Each saved workstream object:
---
### `GET /v1/api/templates`
Returns a summary list of all available prompt templates. This is a read-only
endpoint (requires `read` scope) that exposes template names and categories
without revealing template content. Useful for populating template selectors
in UIs or discovering available templates before creating a workstream.
**Response:**
```json
{
"templates": [
{"name": "safety-guidelines", "category": "safety", "is_default": true, "origin": "manual"},
{"name": "mcp__server__code", "category": "", "is_default": false, "origin": "mcp"}
]
}
```
Each template summary:
| Field | Type | Description |
|--------------|--------|------------------------------------------------------|
| `name` | string | Template name (used in `template` field on creation) |
| `category` | string | Template category |
| `is_default` | bool | Whether template is auto-applied to all sessions |
| `origin` | string | Template origin: `manual` or `mcp` |
> **Note:** For full template management (create, update, delete, view content),
> use the admin endpoints at `GET /v1/api/admin/templates` (requires `admin.templates` permission).
---
### `GET /v1/api/ws-templates`
Returns a summary list of enabled workstream templates. This is a read-only
endpoint (requires `read` scope) for populating template selectors in UIs.
**Response:**
```json
{
"ws_templates": [
{"name": "code-review", "description": "Code review profile", "model": "gpt-5"},
{"name": "ops-triage", "description": "On-call triage", "model": ""}
]
}
```
Each workstream template summary:
| Field | Type | Description |
|---------------|--------|-------------------------------------------------|
| `name` | string | Template name (used in `ws_template` on creation)|
| `description` | string | Human-readable description |
| `model` | string | Model alias override (empty = use default) |
> **Note:** For full workstream template management, use the admin endpoints at
> `GET /v1/api/admin/ws-templates` (requires `admin.templates` permission).
---
### `POST /v1/api/send`
Sends a user message to a workstream. Spawns a daemon worker thread that calls
@@ -1391,6 +1452,24 @@ the `admin.settings` permission.
---
### MCP Servers
| Method | Path | Description |
|--------|------|-------------|
| GET | `/v1/api/admin/mcp-servers` | List all MCP server definitions with live node status. Query: `?reveal=true` to show env/header secrets. |
| POST | `/v1/api/admin/mcp-servers` | Create an MCP server definition. Body: `{name, transport, command?, args?, url?, headers?, env?, auto_approve?, enabled?}` |
| GET | `/v1/api/admin/mcp-servers/{server_id}` | Get a single MCP server with per-node connection status. |
| PUT | `/v1/api/admin/mcp-servers/{server_id}` | Update an MCP server definition. Partial updates supported. |
| DELETE | `/v1/api/admin/mcp-servers/{server_id}` | Delete an MCP server definition. |
| POST | `/v1/api/admin/mcp-servers/reload` | Tell all cluster nodes to re-read the `mcp_servers` DB table and reconcile (add new, remove stale, reconnect changed). |
| POST | `/v1/api/admin/mcp-servers/import` | Import servers from a pasted JSON config. Body: `{config: {mcpServers: {...}}}`. Skips existing names. |
Permission: `admin.mcp`
Secrets (`env`, `headers` fields) are masked with `***` by default. Use `?reveal=true` on GET endpoints to see actual values.
---
### `OPTIONS` (any path)
Handles CORS preflight requests.
+44 -13
View File
@@ -91,14 +91,16 @@ 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)
ui/
colors.py ANSI color constants with NO_COLOR support
markdown.py Streaming terminal markdown renderer (line-buffered)
spinner.py Braille character spinner (daemon thread)
static/
index.html Single-page app shell (links to CSS and JS)
style.css Page-specific UI styles (dashboard layout, approval blocks)
app.js Page-specific client-side JavaScript (SSE, workstreams, markdown)
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)
tools/
*.json 15 tool schemas (OpenAI function-calling format + turnstone metadata)
```
@@ -506,8 +508,18 @@ independently, then returns the final content as the tool result.
and exposes their tools alongside built-in tools. The MCP SDK is fully async; turnstone
bridges this with a background asyncio event loop in a daemon thread.
**Configuration sources:** MCP servers can be defined in config files (TOML/JSON)
or in the database via the admin UI. Database-backed definitions are managed
through the console admin panel's MCP Servers tab and stored in the
`mcp_servers` table. On startup, `load_mcp_config(storage=)` uses
first-match-wins priority: DB rows (if any enabled) take precedence over
config files. The console can trigger a cluster-wide reload (`POST
/_internal/mcp-reload`) that causes each node to call `reconcile_sync()`,
which diffs the running MCP connections against the current DB state and
adds, removes, or reconnects servers as needed.
**Lifecycle:**
1. `create_mcp_client()` reads server configs from TOML or JSON
1. `create_mcp_client()` reads server configs from TOML/JSON and database
2. `MCPClientManager.start()` launches the background event loop thread
3. `_connect_all()` connects to each server (stdio subprocess or HTTP), runs
`initialize()` + `list_tools()`, converts schemas to OpenAI format, detects
@@ -1018,8 +1030,8 @@ Three hierarchical scopes control endpoint access:
- **Console** is the auth management hub — it hosts the admin endpoints for
creating users, issuing API tokens, and managing channel mappings. User
records and token hashes live in the shared storage backend. The console
dashboard includes an **admin panel** (13 tabs) for managing
credentials, governance, and runtime settings through the browser.
dashboard includes an **admin panel** (14 tabs) for managing
credentials, governance, MCP servers, and runtime settings through the browser.
- **Server** is a JWT validator only — it validates tokens on each request but
never creates users or tokens. Both processes share the same `jwt_secret`
(via `TURNSTONE_JWT_SECRET` env var or `[auth].jwt_secret` config).
@@ -1190,8 +1202,13 @@ The bridge dispatches it to `POST /v1/api/cancel` on the server owning the works
which sets the cooperative cancel flag and unblocks any pending approval/plan waits.
**Completion detection:** The bridge tracks which `correlation_id` maps to which
`ws_id` for active sends. When the global SSE reports `ws_state → idle` for a tracked
workstream, the bridge emits a synthetic `TurnCompleteEvent` with the correlation ID.
`ws_id` for active sends. The server accumulates content tokens in the WebUI and
piggybacks the full response text onto the `ws_state → idle` global SSE event.
When the bridge receives this event, it emits a synthetic `TurnCompleteEvent`
carrying the correlation ID and the server-provided `content`. This lets downstream
consumers (e.g. the Discord bot) recover the full response when individual
`ContentEvent`s were missed, and serves as the primary delivery path for
bidirectional notification DM forwarding.
**Multi-node routing:** Each bridge retrieves its `node_id` from the server's
`/health` endpoint on startup (with exponential backoff retry). The server
@@ -1355,11 +1372,23 @@ directly over HTTP for lower latency: `_exec_notify()` queries the
`services` database table for healthy channel gateways (heartbeat within
120 seconds), authenticates with a service JWT (`aud: turnstone-channel`),
and POSTs to `POST /v1/api/notify` on the first healthy gateway. The
gateway validates the JWT, resolves the target (username lookup via
payload includes the originating `ws_id` for reply routing. The gateway
validates the JWT, resolves the target (username lookup via
`channel_users` or direct `channel_type`+`channel_id`), and delegates to
the appropriate `ChannelAdapter.send()`. Delivery retries up to 3 times
with backoff, re-querying the service registry on each attempt. See
[Notification Flow diagram](diagrams/png/17-notify-flow.png).
`ChannelAdapter.send_notification()` which sends the message and tracks
the outgoing message ID → `(ws_id, target_user_id)` mapping. Delivery
retries up to 3 times with backoff, re-querying the service registry on
each attempt. See [Notification Flow diagram](diagrams/png/17-notify-flow.png).
**Bidirectional replies:** When a user replies to a notification DM, the
Discord bot looks up the originating `ws_id` from the tracked message ID,
verifies the replying user matches the notification recipient, and routes
the reply to the workstream via `router.send_message()`. The workstream's
response is forwarded back to the DM via a temporary entry in
`_notify_reply_channels`. On `TurnCompleteEvent`, the response message is
itself tracked for further replies, enabling multi-turn DM conversations
without requiring the user to open the web UI. Tracking entries are capped
at 100 (FIFO eviction) and cleaned up on workstream close.
---
@@ -1393,8 +1422,10 @@ enforcement tracks consumption in `session.send()` with 80% warning and
100% approval gate via the `__budget_override__` synthetic tool name.
The console admin panel adds 6 governance tabs (Roles, Policies, Templates,
WS Templates, Usage, Audit), a Memories tab, and a Settings tab (form-based
editor for all ConfigStore settings) for a total of 13 tabs, all permission-gated.
WS Templates, Usage, Audit), a Memories tab, a Settings tab (form-based
editor for all ConfigStore settings), and an MCP Servers tab (database-backed
server definitions with live connection status and cluster-wide reload) for a
total of 14 tabs, all permission-gated.
Both Python and TypeScript SDKs expose governance methods on the console
client.
+40 -4
View File
@@ -30,8 +30,8 @@ Key components:
- **ChannelAdapter protocol** (`turnstone/channels/_protocol.py`) — generic
interface for any messaging platform. Defines `start()`, `stop()`,
`send()`, `edit_message()`, `send_approval_request()`,
`send_plan_review()`, and `create_thread()`.
`send()`, `send_notification()`, `edit_message()`,
`send_approval_request()`, `send_plan_review()`, and `create_thread()`.
- **ChannelRouter** (`turnstone/channels/_routing.py`) — maps
channel/thread IDs to turnstone workstream IDs. Handles workstream
creation via MQ, stale route detection, and user identity resolution.
@@ -271,13 +271,43 @@ gateway directly over HTTP:
2. `_exec_notify()` queries the `services` table for healthy channel
gateways (heartbeat within the last 120 seconds)
3. The server mints a service JWT (`aud: turnstone-channel`) via
`ServiceTokenManager` and POSTs to the first healthy gateway
`ServiceTokenManager` and POSTs to the first healthy gateway. The
payload includes the originating `ws_id` for reply routing.
4. The gateway validates the JWT, resolves the target, and calls
`adapter.send()` on the appropriate platform adapter
`adapter.send_notification()` which sends the message and tracks
the outgoing message ID for reply routing
5. On failure, the server tries the next gateway. If all fail, it
retries up to 2 more times (delays: 1s, 3s), re-querying the
service registry on each attempt
### Bidirectional Replies
Notifications support multi-turn DM conversations. When a user replies
to a notification DM:
1. The bot looks up the originating `ws_id` from the tracked message ID
(`_notify_ws_map`)
2. Verifies the replying user matches the original notification
recipient (defence in depth — Discord DMs are already private)
3. Routes the reply to the workstream via `router.send_message()`
4. Registers the DM channel for response forwarding
(`_notify_reply_channels`)
5. When the workstream responds (`TurnCompleteEvent`), the response is
forwarded to the DM
6. The response message is itself tracked, so the user can reply again
for another turn
This enables scenarios like an oncall engineer responding to a CI/CD
failure notification from their phone before opening a laptop.
**Limits:**
- Tracking map capped at 100 entries (FIFO eviction of oldest)
- Entries cleaned up on workstream close/unsubscribe
- Replying to an expired notification sends
*"This notification is no longer active."*
- DM reply content capped at 4096 characters
### Service Registry
The channel gateway registers itself in the `services` database table
@@ -328,12 +358,18 @@ class ChannelAdapter(Protocol):
async def start(self) -> None: ...
async def stop(self) -> None: ...
async def send(self, channel_id: str, content: str) -> str: ...
async def send_notification(self, channel_id: str, content: str, ws_id: str) -> str: ...
async def edit_message(self, channel_id: str, message_id: str, content: str) -> None: ...
async def send_approval_request(self, channel_id: str, ws_id: str, correlation_id: str, items: list[dict]) -> None: ...
async def send_plan_review(self, channel_id: str, ws_id: str, correlation_id: str, content: str) -> None: ...
async def create_thread(self, parent_channel_id: str, name: str, message_id: str = "") -> str: ...
```
`send_notification()` is like `send()` but associates the outgoing
message with a `ws_id` so that user replies can be routed back to the
originating workstream. Adapters must track the mapping from outgoing
message ID to `(ws_id, target_user_id)` and handle DM replies.
To add a new platform:
1. Create `turnstone/channels/<platform>/` package
+1
View File
@@ -75,6 +75,7 @@ package "turnstone/ui/" <<Rectangle>> {
component [colors.py\nANSI colors] as colors <<ui>>
component [markdown.py\nMD rendering] as markdown <<ui>>
component [spinner.py\nTerminal spinner] as spinner <<ui>>
component [renderer.js\nBrowser MD + LaTeX] as renderer <<ui>>
}
' API schemas
+2 -2
View File
@@ -69,8 +69,8 @@ partition "Phase 2: Approve" #FFF3E0 {
**TerminalUI**: Print headers/previews,
prompt [y/n/a, optional message]
If user chose "always":
Set ui.auto_approve = True
(auto-approve all future tools in this session)
Add pending tool names to auto_approve_tools
(auto-approve these tool types going forward)
**WebUI**: Enqueue approve_request,
block on _approval_event.wait()
**NullUI**: Return (True, None)
+1
View File
@@ -174,6 +174,7 @@ package "Outbound Events (Bridge → Client)" #E3F2FD {
}
class TurnCompleteEvent {
type = "turn_complete"
+ content: str
}
}
+2 -2
View File
@@ -39,8 +39,8 @@ BridgeA -> Redis : PUBLISH turnstone:events:abc12345\nAckEvent(status:"ok")
... SSE events flow: content, tool_output_chunk, tool_result, status, state_change ...
BridgeA -> Redis : PUBLISH turnstone:events:abc12345\nContentEvent, ToolResultEvent, ...
BridgeA -> Redis : PUBLISH turnstone:events:global\nStateChangeEvent(state:"idle")
BridgeA -> Redis : PUBLISH turnstone:events:abc12345\nTurnCompleteEvent
BridgeA -> Redis : PUBLISH turnstone:events:global\nStateChangeEvent(state:"idle", content:"...")
BridgeA -> Redis : PUBLISH turnstone:events:abc12345\nTurnCompleteEvent(content:"...")
== Scenario B: Directed Message to Specific Node ==
+15 -2
View File
@@ -53,6 +53,7 @@ class "DiscordBot" as Bot <<service>> {
+on_message(msg)
+on_interaction(interaction)
+send(channel_id, content)
+send_notification(channel_id, content, ws_id)
+run(token)
--
discord.py Client
@@ -61,6 +62,9 @@ class "DiscordBot" as Bot <<service>> {
Creates threads for workstreams
Renders approval buttons
escape_mentions() on send
--
_notify_ws_map: msg_id → (ws_id, user_id)
_notify_reply_channels: ws_id → (dm, user_id)
}
class "ChannelRouter" as Router <<service>> {
@@ -240,11 +244,20 @@ note bottom of SVC
3. Queries services table for healthy gateways
4. Mints JWT (aud: turnstone-channel) via
ServiceTokenManager
5. POSTs to first healthy gateway
5. POSTs to first healthy gateway (incl. ws_id)
6. Gateway validates JWT, resolves target
7. adapter.send() → Discord API
7. adapter.send_notification() → Discord API
(tracks msg_id → ws_id for reply routing)
8. On failure: retry up to 3× (1s, 3s backoff)
9. SSRF: only http(s) URLs allowed
**Bidirectional DM Replies**
1. User replies to notification DM
2. Bot looks up ws_id from _notify_ws_map
3. Verifies author == notification recipient
4. Routes reply via router.send_message()
5. Response forwarded to DM on TurnCompleteEvent
6. Response tracked for multi-turn conversation
end note
@enduml
+35
View File
@@ -103,6 +103,41 @@ alt all retries exhausted
Session --> Session : "Error: notification delivery failed"
end
== Bidirectional Reply (User responds to notification DM) ==
Discord -> Adapter : user replies to\nnotification message
Adapter -> Adapter : lookup message_id\nin _notify_ws_map
note right
Maps message_id →
(ws_id, target_user_id)
Atomic pop prevents TOCTOU
end note
alt message not tracked
Adapter -> Discord : "This notification\nis no longer active."
else tracked
Adapter -> Adapter : verify author ==\ntarget_user_id
Adapter -> Adapter : resolve_user()\n(unlinked → drop)
Adapter -> Adapter : router.send_message(ws_id, content)
note right
Routes reply via MQ to
the originating workstream.
Registers DM channel in
_notify_reply_channels[ws_id]
end note
... workstream processes reply ...
Adapter <- Adapter : TurnCompleteEvent\n(with content)
Adapter -> Discord : forward response to DM
Adapter -> Adapter : track response message\nfor multi-turn replies
note right
Response message_id added
to _notify_ws_map — user can
reply again indefinitely
end note
end
== Service Registry (Background) ==
note over Gateway, Storage
@@ -85,6 +85,12 @@ tload --> trender : template content
trender --> tsys : rendered content
tset --> tload : name or None
note right of pt_db
Read-only listing:
GET /v1/api/templates
(read scope, summary only)
end note
govjs --> wt_db : /v1/api/admin/ws-templates
wtr --> wt_db : get_ws_template_by_name()
wtr --> wta : template settings
+19
View File
@@ -16,8 +16,27 @@ participant "ChatSession\n(session.py)" as Session <<session>>
participant "StorageBackend\n(governance)" as Storage <<storage>>
participant "Server / Console\n(health + UI)" as UI <<server>>
participant "Console Admin UI\n(admin panel)" as Admin <<ui>>
participant "Database\n(mcp_servers table)" as DB <<storage>>
== Admin-Driven Configuration ==
Admin -> DB : CRUD MCP server definitions\n(POST/PUT/DELETE /v1/api/admin/mcp-servers)
Admin -> UI : POST /v1/api/admin/mcp-servers/reload
UI -> MCPMgr : POST /_internal/mcp-reload\n(forwarded to each node)
MCPMgr -> MCPMgr : reconcile_sync()
note right
Diffs running servers against DB:
- New entries → connect
- Removed entries → disconnect
- Changed entries → reconnect
end note
== Startup: Connection & Discovery ==
MCPMgr -> DB : load_mcp_config(storage=)\n(merge config file + DB)
MCPMgr -> MCPSrv : initialize (stdio or HTTP)
MCPSrv --> MCPMgr : capabilities\n(tools, resources, prompts)
@@ -48,6 +48,14 @@ end note
Admin -> Server : GET /v1/api/admin/ws-templates
Server -> Storage : list_ws_templates()
Server <-- Server : GET /v1/api/ws-templates\n(read scope, summary only)
note right
**Read-only listing:**
name, description, model.
Used by creation UI dropdowns.
Available on both server + console.
end note
Admin -> Server : DELETE /v1/api/admin/ws-templates/{id}
Server -> Storage : delete_ws_template(id)
+147
View File
@@ -0,0 +1,147 @@
@startuml
!theme plain
title Turnstone — OIDC Authorization Code Flow with PKCE
skinparam participant {
BackgroundColor<<browser>> #E8EAF6
BackgroundColor<<server>> #FFE0B2
BackgroundColor<<storage>> #B3E5FC
BackgroundColor<<idp>> #C8E6C9
}
participant "Browser" as Browser <<browser>>
participant "Turnstone\n(Server / Console)" as Server <<server>>
database "SQLite /\nPostgreSQL" as DB <<storage>>
participant "Identity Provider\n(IdP)" as IdP <<idp>>
== Page Load ==
Browser -> Server : GET /v1/api/auth/status
Server --> Browser : {oidc_enabled: true,\noidc_provider_name: "...",\npassword_enabled: true}
note right of Browser
Login screen renders
"Continue with {provider_name}"
button alongside password form.
If password_enabled=false,
only the SSO button is shown.
end note
== Authorization Request ==
Browser -> Server : GET /v1/api/auth/oidc/authorize
Server -> Server : Generate state (random)\nnonce (random)\nPKCE code_verifier + code_challenge
Server -> DB : create_oidc_pending_state(\nstate, nonce, code_verifier, audience)
note right of DB
Stored with created_at timestamp.
Expires after 5 minutes.
end note
Server --> Browser : 302 Redirect to IdP\nauthorization_endpoint
Browser -> IdP : GET /authorize?\nresponse_type=code&\nclient_id=...&\nredirect_uri=...&\nscope=openid email profile&\nstate=...&nonce=...&\ncode_challenge=...&\ncode_challenge_method=S256
== User Authentication (at IdP) ==
IdP -> Browser : Login page (if no\nexisting IdP session)
Browser -> IdP : User authenticates\n(username/password, MFA, etc.)
IdP --> Browser : 302 Redirect to callback\n?code=AUTH_CODE&state=STATE
== Callback Processing ==
Browser -> Server : GET /v1/api/auth/oidc/callback\n?code=AUTH_CODE&state=STATE
Server -> Server : Rate limit check\n(5 per 5min per IP)
Server -> DB : cleanup_expired_oidc_states(300)
note right of DB
Lazy cleanup of states
older than 5 minutes.
end note
Server -> DB : pop_oidc_pending_state(state)
DB --> Server : {nonce, code_verifier, audience}
note right of Server
Atomic fetch-and-delete.
Returns None if state is
expired or unknown.
end note
== Token Exchange ==
Server -> IdP : POST /token\ngrant_type=authorization_code&\ncode=AUTH_CODE&\nclient_id=...&\nclient_secret=...&\ncode_verifier=...&\nredirect_uri=...
note right of Server
Client secret + PKCE verifier
sent server-side only.
Never exposed to browser.
end note
IdP --> Server : {id_token: "eyJ...",\naccess_token: "..."}
== ID Token Validation ==
Server -> IdP : Fetch JWKS public keys\n(cached at startup, refreshed\non-demand when unknown kid\nencountered — key rotation)
Server -> Server : Validate ID token:\n1. Verify signature (RS256/ES256)\n2. Check iss == configured issuer\n3. Check aud == client_id\n4. Check exp (not expired)\n5. Verify nonce matches
== User Provisioning ==
Server -> DB : get_oidc_identity(issuer, sub)
alt Existing identity found
DB --> Server : {user_id, ...}
Server -> DB : update_oidc_identity_login()\nupdate last_login timestamp
Server -> DB : get_user(user_id)
DB --> Server : user record
else New user (first login)
Server -> Server : Derive username from\npreferred_username / email
Server -> DB : create_user(user_id, username,\ndisplay_name, "!oidc")
note right of DB
Password hash set to sentinel
value "!oidc" — not a valid
bcrypt hash, so password login
is always rejected.
end note
Server -> DB : create_oidc_identity(\nissuer, sub, user_id, email)
end
opt Role mapping configured
Server -> Server : Read role_claim from ID token\nMap values via role_map
Server -> DB : Sync roles: add new,\nrevoke stale OIDC-assigned,\npreserve manually assigned
end
== Issue Turnstone JWT ==
Server -> Server : Load user permissions\nDerive scopes from permissions
Server -> Server : Create JWT (HS256)\nsub: user_id\nscopes: read,write,...\nsrc: "oidc"\naud: turnstone-server\nexp: +24h
Server --> Browser : 302 Redirect to /?oidc_success=1\nSet-Cookie: session=JWT\n(HttpOnly, SameSite=Lax, Secure)
== Browser Success Detection ==
Browser -> Browser : Detect ?oidc_success=1\nStrip param from URL\n(history.replaceState)
Browser -> Browser : Hide login overlay\nCall onLoginSuccess()
note right of Browser
Browser is now authenticated.
JWT cookie sent on all
subsequent requests.
end note
== Error Paths ==
note over Browser, IdP
**Error handling:**
- IdP returns error param → redirect to /?oidc_error=...
- State missing/expired → redirect to /?oidc_error=Login+session+expired
- Token exchange fails → redirect to /?oidc_error=...
- ID token validation fails → redirect to /?oidc_error=...
- No admin user exists → redirect to /?oidc_error=Initial+setup+required
- Rate limit exceeded → redirect to /?oidc_error=Too+many+login+attempts
All errors are shown as toast messages on the login screen.
end note
@enduml
+2 -2
View File
@@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:c9daca81971ba7a8ed6736d23d5373c69435158fa6240b9880d14fc4759ab580
size 329673
oid sha256:efcc7cbe8161a54b5ec24bdfd47e8a142f70029e6e66c707e811b99369f85ebf
size 310079
+2 -2
View File
@@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:6dd3c923d1e1c49b5f91d8d342fb4b0d49a46d432460379ad146a9e3b075a05a
size 277234
oid sha256:43844b07d36beb04db871f6795a3f3be17852a6a484fdc0ea207403bd7f512a6
size 274286
+2 -2
View File
@@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:2229801220548e4794baa67e27a0a39dc7968c826a28fa8144763e678c8ed733
size 192556
oid sha256:6e94a10f039a7f69517e84d0946e0c649035c15b38ebc2314e7b9cd501eb244d
size 192559
@@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:1380065cbb5f95b5ea7dc6b2a00986c455b82888af60784980dffbd936460dcf
size 431129
oid sha256:6fc99bb8d84d6e9f3dac9d5c12ac7f569a041b29431c57612c24b50f332982ed
size 462992
+2 -2
View File
@@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:f0f6097840fccdbfe16cd5e4c9f5d063b2c36942460a944df68b8ec947e63ea3
size 221452
oid sha256:cc4c511c34a2e5d286fd128c3509405a5b240ca02a4bafb395d2e94d002a5b8b
size 293203
@@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:f4dac4948d928b4705936d73b4d159aa1e89315ec0397616ca914bbf19e7a1ce
size 206479
oid sha256:3aaca1ae4c6c255dc9569f59e3ccc24f8b3bab0ac2a9b08c85e2af72d6a400c7
size 218575
+2 -2
View File
@@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:8e6dc5142c7908314ce01229b3c4f13bf9450adcbb62a178838bd4cf81d9f4da
size 250417
oid sha256:e4593873599342b2830fedd5d783e9a28eab0bb0d6589798ef6ef2649eeee80f
size 324518
@@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:c06d7086d7965eb9fe333396f027133d42507cf120bfe8dc851c009a8768ec48
size 284926
oid sha256:fadf5b07f8230ecf97805a86b308eaa9eb30516dd26900e5c9f70e6fb7562bab
size 296339
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:1c21910e3916be789b0377c8a0dcc8f47d66a967861a543d5bdd0c26da185259
size 309584
+429
View File
@@ -0,0 +1,429 @@
# OpenID Connect (OIDC) Single Sign-On
Turnstone supports OpenID Connect for federated authentication, allowing
users to log in with their existing corporate identity provider instead of
managing a separate password. OIDC is opt-in: when configured, the login
screen shows a "Continue with SSO" button alongside the existing
username/password form. When not configured, the login experience is
unchanged.
Any OIDC-compliant provider works: Google, Okta, Azure AD, Keycloak,
Auth0, OneLogin, and others that publish a
`.well-known/openid-configuration` discovery document.
---
## Prerequisites
1. A registered **confidential** OIDC client at your identity provider
2. The client's redirect URI must include:
`https://your-turnstone-host/v1/api/auth/oidc/callback`
3. A local admin user must exist in Turnstone (complete the initial setup
wizard before enabling OIDC)
---
## Configuration
OIDC is configured via environment variables (preferred) or the `[oidc]`
section of `config.toml`. Environment variables take precedence when both
are set.
| Variable | Required | Default | Description |
|----------|----------|---------|-------------|
| `TURNSTONE_OIDC_ISSUER` | Yes | — | Issuer URL (e.g. `https://accounts.google.com`). Must serve `/.well-known/openid-configuration`. |
| `TURNSTONE_OIDC_CLIENT_ID` | Yes | — | OAuth 2.0 client ID from your provider |
| `TURNSTONE_OIDC_CLIENT_SECRET` | Yes | — | OAuth 2.0 client secret (confidential client) |
| `TURNSTONE_OIDC_SCOPES` | No | `openid email profile` | Space-separated OAuth scopes to request |
| `TURNSTONE_OIDC_PROVIDER_NAME` | No | `SSO` | Display name for the login button (e.g. "Google", "Okta") |
| `TURNSTONE_OIDC_ROLE_CLAIM` | No | — | ID token claim containing role/group values (see [Role Mapping](#role-mapping)) |
| `TURNSTONE_OIDC_ROLE_MAP` | No | — | Mapping from claim values to Turnstone role IDs (see [Role Mapping](#role-mapping)) |
| `TURNSTONE_OIDC_PASSWORD_ENABLED` | No | `true` | Set to `false` to hide the password form and block all username/password logins (including admin). API tokens and config-file tokens still work. |
| `TURNSTONE_OIDC_REDIRECT_BASE` | No | — | Externally-reachable origin for the OIDC redirect URI (e.g. `https://app.example.com`). Recommended when running behind a reverse proxy. When unset, derived from the request Host header. |
OIDC is enabled when all three required fields (issuer, client ID, client
secret) are non-empty. If any is missing, OIDC is silently disabled and
the login screen shows only the password form.
### Reverse Proxy / Load Balancer
When Turnstone runs behind a reverse proxy, the internal `Host` header may
not match the externally-reachable URL. Set `TURNSTONE_OIDC_REDIRECT_BASE`
to the public origin so the redirect URI sent to the identity provider is
correct:
```bash
TURNSTONE_OIDC_REDIRECT_BASE=https://app.example.com
```
The resulting callback URL will be
`https://app.example.com/v1/api/auth/oidc/callback` — register this as the
authorized redirect URI in your identity provider.
### config.toml alternative
```toml
[oidc]
issuer = "https://accounts.google.com"
client_id = "your-client-id"
client_secret = "your-client-secret"
scopes = "openid email profile"
provider_name = "Google"
role_claim = "groups"
password_enabled = true
redirect_base = "https://app.example.com"
[oidc.role_map]
admin = "builtin-admin"
engineering = "builtin-operator"
```
---
## Provider-Specific Setup
### Google
1. Go to [Google Cloud Console](https://console.cloud.google.com/) >
**APIs & Services** > **Credentials**
2. Click **Create Credentials** > **OAuth 2.0 Client ID**
3. Application type: **Web application**
4. Add authorized redirect URI:
`https://your-turnstone-host/v1/api/auth/oidc/callback`
5. Copy the **Client ID** and **Client secret**
```bash
TURNSTONE_OIDC_ISSUER=https://accounts.google.com
TURNSTONE_OIDC_CLIENT_ID=123456789.apps.googleusercontent.com
TURNSTONE_OIDC_CLIENT_SECRET=GOCSPX-...
TURNSTONE_OIDC_PROVIDER_NAME=Google
```
### Okta
1. In the Okta Admin Console, go to **Applications** > **Create App
Integration**
2. Sign-in method: **OIDC - OpenID Connect**
3. Application type: **Web Application**
4. Add sign-in redirect URI:
`https://your-turnstone-host/v1/api/auth/oidc/callback`
5. Note the **Issuer** (your Okta domain, e.g.
`https://dev-123456.okta.com`)
```bash
TURNSTONE_OIDC_ISSUER=https://dev-123456.okta.com
TURNSTONE_OIDC_CLIENT_ID=0oaXXXXXXXXXXXXX
TURNSTONE_OIDC_CLIENT_SECRET=...
TURNSTONE_OIDC_PROVIDER_NAME=Okta
TURNSTONE_OIDC_ROLE_CLAIM=groups
TURNSTONE_OIDC_ROLE_MAP="admin:builtin-admin,everyone:builtin-operator"
```
### Azure AD (Entra ID)
1. In the Azure Portal, go to **App registrations** > **New registration**
2. Redirect URI: **Web** >
`https://your-turnstone-host/v1/api/auth/oidc/callback`
3. Under **Certificates & secrets**, create a new **Client secret** and
copy the value immediately
4. The issuer URL is
`https://login.microsoftonline.com/{tenant-id}/v2.0`
```bash
TURNSTONE_OIDC_ISSUER=https://login.microsoftonline.com/YOUR_TENANT_ID/v2.0
TURNSTONE_OIDC_CLIENT_ID=xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx
TURNSTONE_OIDC_CLIENT_SECRET=...
TURNSTONE_OIDC_PROVIDER_NAME="Azure AD"
TURNSTONE_OIDC_ROLE_CLAIM=roles
TURNSTONE_OIDC_ROLE_MAP="Admin:builtin-admin,User:builtin-operator"
```
### Keycloak
1. In the Keycloak Admin Console, select your **Realm**
2. Go to **Clients** > **Create client**
3. Client type: **OpenID Connect**
4. Set **Client authentication** to **On** (confidential)
5. Add valid redirect URI:
`https://your-turnstone-host/v1/api/auth/oidc/callback`
6. The issuer URL is
`https://keycloak.example.com/realms/your-realm`
```bash
TURNSTONE_OIDC_ISSUER=https://keycloak.example.com/realms/your-realm
TURNSTONE_OIDC_CLIENT_ID=turnstone
TURNSTONE_OIDC_CLIENT_SECRET=...
TURNSTONE_OIDC_PROVIDER_NAME=Keycloak
TURNSTONE_OIDC_ROLE_CLAIM=realm_access.roles
TURNSTONE_OIDC_ROLE_MAP="admin:builtin-admin,operator:builtin-operator"
```
---
## Role Mapping
OIDC role mapping assigns Turnstone roles to users based on claims in the
ID token. This is optional — without it, OIDC users are provisioned with
the `builtin-viewer` role (read-only access) by default.
### Configuration
Set `TURNSTONE_OIDC_ROLE_CLAIM` to the name of the claim in the ID token
that contains the user's group or role memberships. Then set
`TURNSTONE_OIDC_ROLE_MAP` to map claim values to Turnstone role IDs.
The role map is a comma-separated list of `claim_value:turnstone_role`
pairs:
```bash
TURNSTONE_OIDC_ROLE_CLAIM=groups
TURNSTONE_OIDC_ROLE_MAP="admin:builtin-admin,engineering:builtin-operator,viewer:builtin-viewer"
```
### Behavior
- **Synced on every login**: roles are added when new claim values appear,
and OIDC-assigned roles are revoked when the corresponding claim value
is no longer present. Roles assigned manually (or by other sources) are
never touched — only roles with `assigned_by="oidc"` are subject to
revocation.
- **List or string**: the claim value can be a JSON array
(`["admin", "engineering"]`) or a single string (`"admin"`). Both are
handled correctly.
- **Unknown values**: claim values not present in the role map are silently
ignored.
- **Missing roles**: if the role map references a Turnstone role ID that
does not exist in the database, the assignment is skipped (no error).
- **Evaluated on every login**: roles are checked and applied each time
the user authenticates via OIDC, so new group memberships are picked
up on the next login.
### Built-in Roles
| Role ID | Permissions |
|---------|-------------|
| `builtin-admin` | All permissions |
| `builtin-operator` | read, write, workstreams.create, workstreams.close |
| `builtin-viewer` | read |
---
## User Provisioning
When a user logs in via OIDC for the first time, Turnstone automatically
creates a local user account:
1. The OIDC identity (`issuer` + `sub` claim) is stored in the
`oidc_identities` table and linked to the new user
2. The **username** is derived from the `preferred_username` claim,
falling back to the email local part, with deduplication if needed
3. The **display name** comes from the `name` claim, falling back to
`preferred_username` or email
4. The user's password hash is set to a sentinel value (`!oidc`) — OIDC
users cannot log in with a password
On subsequent logins, the existing user is matched by `(issuer, sub)` and
the `last_login` timestamp is updated. Role mapping is re-evaluated on
every login.
---
## OIDC-Only Mode
To enforce OIDC for all logins and hide the password form, set:
```bash
TURNSTONE_OIDC_PASSWORD_ENABLED=false
```
In this mode the login screen shows only the "Continue with SSO" button.
The password form, token toggle, and sign-in button are all hidden.
All username/password logins are blocked at the API level, including
admin accounts.
The first admin account must be created via the setup wizard (with a
password) before OIDC is enabled. The setup wizard always works
regardless of this setting because it is only available when zero users
exist in the database.
API token login (`POST /v1/api/auth/login` with a `ts_` token) and
config-file tokens (`Authorization: Bearer tok_xxx`) continue to work
regardless of this setting. OIDC-only mode affects password-based
authentication only.
---
## Login Flow
Both the server and console support OIDC login. The flow is identical:
1. The browser fetches `GET /v1/api/auth/status` at page load
2. If the response includes `oidc_enabled: true`, the login screen shows
a "Continue with {provider_name}" button
3. Clicking the button navigates to `GET /v1/api/auth/oidc/authorize`
4. Turnstone generates a state token, nonce, and PKCE verifier, stores
them in the database, and redirects the browser to the identity
provider's authorization endpoint
5. The user authenticates at the identity provider
6. The IdP redirects back to
`GET /v1/api/auth/oidc/callback?code=...&state=...`
7. Turnstone validates the state, exchanges the authorization code for
tokens using the PKCE verifier, validates the ID token against the
provider's JWKS public keys, provisions or matches the user, and
issues a Turnstone JWT
8. The browser is redirected to `/?oidc_success=1` with the JWT set in
an `HttpOnly` session cookie
9. The browser JavaScript detects the `oidc_success` query parameter,
strips it from the URL, hides the login overlay, and calls
`onLoginSuccess()` to initialize the application
---
## API Endpoints
| Method | Path | Auth | Description |
|--------|------|------|-------------|
| GET | `/v1/api/auth/oidc/authorize` | Public | Redirects to identity provider |
| GET | `/v1/api/auth/oidc/callback` | Public | Handles IdP callback, issues JWT |
Both endpoints are public (no authentication required) because they are
part of the login flow itself.
### Auth status response
When OIDC is enabled, `GET /v1/api/auth/status` includes additional
fields:
```json
{
"auth_enabled": true,
"has_users": true,
"setup_required": false,
"oidc_enabled": true,
"oidc_provider_name": "Google",
"password_enabled": true
}
```
---
## Database Schema
Migration 018 creates two tables:
```sql
CREATE TABLE oidc_identities (
issuer TEXT NOT NULL,
subject TEXT NOT NULL,
user_id TEXT NOT NULL,
email TEXT NOT NULL DEFAULT '',
created TEXT NOT NULL,
last_login TEXT NOT NULL,
PRIMARY KEY (issuer, subject)
);
CREATE INDEX idx_oidc_identities_user_id ON oidc_identities(user_id);
CREATE TABLE oidc_pending_states (
state TEXT PRIMARY KEY,
nonce TEXT NOT NULL,
code_verifier TEXT NOT NULL,
audience TEXT NOT NULL,
created_at TEXT NOT NULL
);
```
The `oidc_identities` table links an OIDC subject (identified by
`issuer` + `subject`) to a Turnstone `user_id`. A single user can have
multiple OIDC identities (e.g. from different providers).
The `oidc_pending_states` table stores authorization flow state for
callback validation. Entries are automatically cleaned up after 5 minutes.
---
## Security Notes
- **Authorization Code Flow with PKCE**: the recommended OAuth 2.0 flow
for web applications. PKCE prevents authorization code interception
attacks even without a client secret (though the client secret is still
used for additional security).
- **ID token validation**: all tokens are validated using the provider's
JWKS public keys (RS256 or ES256). The signature, issuer, audience,
and expiry are all checked.
- **State parameter**: a cryptographically random state token prevents
CSRF attacks on the callback endpoint. The state is stored server-side
and verified on callback.
- **Nonce**: a random nonce is included in the authorization request and
verified in the ID token to prevent replay attacks.
- **Client secret**: never leaves the server — it is only used in the
server-to-IdP token exchange, not exposed to the browser.
- **OIDC users cannot use password login**: the sentinel password hash
(`!oidc`) ensures `verify_password()` always rejects password attempts
for OIDC-provisioned users.
- **Rate limiting**: the callback endpoint shares the login rate limiter
(5 attempts per 5-minute window per IP).
- **State TTL**: pending authorization states expire after 5 minutes.
Expired states are lazily cleaned up on each callback.
- **Setup guard**: OIDC login requires at least one local admin user to
exist. This ensures the initial admin account is always created via the
setup wizard with a password, not hijacked by an external identity.
---
## Troubleshooting
### "OIDC not configured"
All three required environment variables must be set:
`TURNSTONE_OIDC_ISSUER`, `TURNSTONE_OIDC_CLIENT_ID`, and
`TURNSTONE_OIDC_CLIENT_SECRET`. Check that none are empty or
whitespace-only.
### "Login session expired"
The authorization flow must complete within 5 minutes. If the user takes
too long at the identity provider, the pending state expires. Try again.
### "Initial setup required"
OIDC login is blocked until at least one local admin user exists.
Complete the setup wizard first (navigate to the Turnstone URL and follow
the prompts to create an admin user with a password).
### Discovery fails at startup
Check that the issuer URL is reachable from the Turnstone server and
serves a valid `/.well-known/openid-configuration` document. The server
logs the discovery attempt at startup:
```
OIDC discovery failed for https://your-issuer.example.com: ...
```
OIDC is automatically disabled when discovery fails. Restart the server
after fixing the connectivity issue.
### Redirect URI mismatch
The redirect URI configured at the identity provider must exactly match
`https://your-host/v1/api/auth/oidc/callback`. Common issues:
- **Scheme mismatch**: the redirect uses `https://` — make sure TLS is
configured or a reverse proxy sets the `X-Forwarded-Proto` header
- **Port mismatch**: if running on a non-standard port, include it in
the redirect URI
- **Path mismatch**: the path must include the `/v1` API version prefix
### User not assigned expected roles
Check that:
1. `TURNSTONE_OIDC_ROLE_CLAIM` matches the exact claim name in the ID
token (case-sensitive)
2. `TURNSTONE_OIDC_ROLE_MAP` maps the correct claim values to valid
Turnstone role IDs
3. The roles referenced in the map exist in the database (check the
admin panel > Roles tab)
4. The identity provider is configured to include the claim in the ID
token (some providers require explicit scope or claim configuration)
+101 -2
View File
@@ -54,7 +54,7 @@ Claims:
|-------|-------------|
| `sub` | User ID |
| `scopes` | Comma-separated scope list (`read,write,approve`) |
| `src` | Token source (`password`, `api_token`, `config`) |
| `src` | Token source (`password`, `api_token`, `config`, `oidc`) |
| `iss` | Issuer — always `turnstone` |
| `aud` | Audience — `turnstone-server` or `turnstone-console` |
| `iat` | Issued-at timestamp |
@@ -90,7 +90,8 @@ Scopes are hierarchical — higher scopes imply all lower ones.
Public paths bypass authentication entirely: `/`, `/health`, `/metrics`,
`/static/*`, `/shared/*`, `/docs`, `/openapi.json`, `/api/auth/login`,
`/api/auth/logout`, `/api/auth/status`, `/api/auth/setup`.
`/api/auth/logout`, `/api/auth/status`, `/api/auth/setup`,
`/api/auth/oidc/authorize`, `/api/auth/oidc/callback`.
### RBAC (Granular Permissions)
@@ -199,6 +200,94 @@ Response:
The response also sets an `HttpOnly` session cookie containing the JWT,
so the browser is immediately authenticated after setup completes.
### OIDC SSO (Single Sign-On)
Turnstone supports OIDC Authorization Code Flow with PKCE for
single sign-on with external identity providers (Okta, Azure AD,
Google, etc.). SSO is opt-in — enabled when the three required
environment variables are set. Users are auto-provisioned on first
login.
#### Configuration
| Variable | Required | Description |
|----------|----------|-------------|
| `TURNSTONE_OIDC_ISSUER` | Yes | OIDC issuer URL (e.g., `https://accounts.google.com`) |
| `TURNSTONE_OIDC_CLIENT_ID` | Yes | Client ID from the identity provider |
| `TURNSTONE_OIDC_CLIENT_SECRET` | Yes | Client secret (confidential client) |
| `TURNSTONE_OIDC_SCOPES` | No | OIDC scopes (default: `openid email profile`) |
| `TURNSTONE_OIDC_PROVIDER_NAME` | No | Display name for the SSO button (default: `SSO`) |
| `TURNSTONE_OIDC_ROLE_CLAIM` | No | Claim name in the ID token for role mapping (e.g., `groups`) |
| `TURNSTONE_OIDC_ROLE_MAP` | No | Comma-separated `claim_value:role_id` pairs (e.g., `admin:builtin-admin,eng:builtin-operator`) |
| `TURNSTONE_OIDC_PASSWORD_ENABLED` | No | Set to `false` to hide password login and force SSO-only |
OIDC is enabled when all three required variables (`ISSUER`,
`CLIENT_ID`, `CLIENT_SECRET`) are set.
#### Login flow
1. User clicks "Continue with [Provider]" on the login page
2. `GET /v1/api/auth/oidc/authorize` generates state, nonce, and PKCE
challenge, stores them in the database, and redirects to the IdP
3. User authenticates at the identity provider
4. IdP redirects to `/v1/api/auth/oidc/callback` with `code` + `state`
5. Server validates state, exchanges the authorization code (with PKCE
verifier), and validates the ID token (JWKS signature, issuer,
audience, nonce)
6. Provisions or matches the user by `(issuer, sub)` — never by
username or email
7. Issues a JWT (`src: oidc`), sets a session cookie, and redirects to
the application
#### Security measures
- **PKCE (S256)** — prevents authorization code interception
- **State parameter** — one-time use, 5-minute TTL, database-backed
(multi-node safe)
- **Nonce** — prevents ID token replay
- **JWKS validation** — asymmetric algorithm allowlist (RS/ES/PS
256-512), HMAC excluded
- **Algorithm allowlist enforced** — the signing key is resolved from
the JWKS by ``kid``; PyJWK infers the key's algorithm from the JWKS
``alg``/``kty`` fields; the token header's ``alg`` must be in the
allowlist AND match the key type, preventing algorithm confusion
- **Identity matching by (issuer, sub) only** — prevents account
takeover via email or username reuse
- **`password_enabled=false` enforced server-side** — not just a UI
toggle
- **Rate limiting** on both authorize and callback endpoints
- **OIDC-provisioned users cannot password-login** — the password hash
is set to the `!oidc` sentinel, which never matches bcrypt verify
#### Role mapping
When `TURNSTONE_OIDC_ROLE_CLAIM` is set (e.g., `groups`), the server
reads that claim from the ID token and maps values to Turnstone roles
via `TURNSTONE_OIDC_ROLE_MAP`. Roles are synced on every login:
matching claim values are added, and stale OIDC-assigned roles are
revoked. Roles assigned manually (not by OIDC) are never touched.
If no role mapping is configured, OIDC users are provisioned with the
`builtin-viewer` role by default.
#### OIDC-only mode
Setting `TURNSTONE_OIDC_PASSWORD_ENABLED=false` hides the password
form on the login page and blocks password-based login at the API
level. The setup wizard always works regardless of this setting — the
first admin user is created with a password before OIDC is relevant.
API tokens and config-file tokens are unaffected by this setting.
#### Known limitations
- **No session revocation** — deprovisioned IdP users retain their JWT
until the 24-hour expiry
- **Single IdP** — configuration supports one issuer (the database
schema supports multiple for future expansion)
- **Redirect URI** — defaults to request Host header; deployments behind
reverse proxies should set `TURNSTONE_OIDC_REDIRECT_BASE` to the
externally-reachable origin to pin the redirect URI
---
## Token Detection Order
@@ -483,3 +572,13 @@ and browsers enforce same-origin policy.
refresh, eliminating long-lived static tokens for inter-service auth.
- **Secret strength validation** — warning logged when JWT secret is
shorter than 32 characters.
- **OIDC PKCE enforcement** — S256 code challenge on every
authorization request prevents code interception in transit.
- **OIDC state/nonce in database** — one-time-use, TTL-bounded tokens
stored in the database, safe for multi-node deployments.
- **OIDC JWKS-only validation** — ID tokens are verified using the
provider's published JWKS keys with asymmetric algorithms only;
HMAC-based algorithms are rejected to prevent algorithm confusion.
- **OIDC identity binding by (issuer, sub)** — user matching uses the
immutable subject identifier, not email or username, preventing
account takeover via IdP attribute changes.
+10 -4
View File
@@ -87,8 +87,11 @@ All prepared items are sent to the UI via `ui.approve_tools(items)`.
but do not block execution.
- Items where `needs_approval` is `True` require the user to accept or deny.
- The user can provide feedback alongside their approval (e.g. "y, use full path").
- If `auto_approve` is `True` on the session (headless mode), all tools are
approved automatically.
- Choosing "always" (key `a`) adds the pending tool names to `auto_approve_tools`,
so that specific tool type is auto-approved going forward (other tool types still
prompt). This is per-tool, not blanket.
- If `auto_approve` is `True` on the session (via `--skip-permissions` or workstream
template), all tools are approved automatically.
### Phase 3: Execute
@@ -620,8 +623,11 @@ MCP-compatible service.
MCP tools **require user approval by default** (`needs_approval: True`). turnstone
does not auto-approve MCP tools based on their schema, since it cannot guarantee
that external tools are read-only. However, global overrides such as
`--skip-permissions` or the UI's "always allow" setting will auto-approve all
tools, including MCP tools.
`--skip-permissions` will auto-approve all tools, including MCP tools. The
interactive "Always" button adds specific tool types to the per-tool auto-approve
set. The web UI and server use `approval_label` for MCP tools, giving
per-prompt/per-resource granularity. The CLI and bridge use `func_name`, which
gives per-tool-type granularity (e.g., all `use_prompt` calls).
### Sub-agent availability
+6 -2
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "turnstone"
version = "0.6.1"
version = "0.7.0"
description = "Multi-node AI orchestration platform with tool use, agent routing, and cluster simulation."
readme = "README.md"
license = "BUSL-1.1"
@@ -51,7 +51,7 @@ sim = ["redis>=7.2"]
anthropic = ["anthropic>=0.39"]
postgres = ["psycopg[binary]>=3.2"]
discord = ["discord.py>=2.4", "redis>=7.2"]
all = ["turnstone[mq,console,sim,anthropic,postgres,discord]"]
[project.scripts]
turnstone = "turnstone.cli:main"
@@ -76,6 +76,9 @@ include = [
"turnstone/console/static/*.js",
"turnstone/shared_static/*.css",
"turnstone/shared_static/*.js",
"turnstone/shared_static/katex-0.16.38/**/*",
"turnstone/shared_static/hljs-11.11.1/**/*",
"turnstone/shared_static/mermaid-11.13.0/**/*",
"turnstone/sdk/py.typed",
]
@@ -90,6 +93,7 @@ line-length = 100
[tool.ruff.lint]
select = ["E", "F", "W", "I", "N", "UP", "B", "A", "SIM", "TCH"]
ignore = ["E501"]
per-file-ignores = { "turnstone/core/sandbox.py" = ["N802"] }
[tool.ruff.format]
quote-style = "double"
+137
View File
@@ -0,0 +1,137 @@
#!/usr/bin/env bash
# Update a vendored JavaScript library in turnstone/shared_static/.
#
# Usage:
# scripts/update-vendored-js.sh katex 0.16.39
# scripts/update-vendored-js.sh hljs 11.12.0
# scripts/update-vendored-js.sh mermaid 11.14.0
#
# This script:
# 1. Downloads the new version from CDN
# 2. Creates the new versioned directory
# 3. Updates all version references in source files
# 4. Removes the old versioned directory
set -euo pipefail
STATIC_DIR="turnstone/shared_static"
CDN="https://cdn.jsdelivr.net/npm"
usage() {
echo "Usage: $0 <katex|hljs|mermaid> <version>"
echo "Example: $0 katex 0.16.39"
exit 1
}
[[ $# -eq 2 ]] || usage
LIB="$1"
VERSION="$2"
# Detect current version from pyproject.toml
detect_old_version() {
local pattern="$1"
grep -oE "${pattern}-[0-9.]+" pyproject.toml | head -1 | sed "s/${pattern}-//"
}
# Update version references across all source files
update_refs() {
local old_pattern="$1" # e.g. katex-0.16.38
local new_pattern="$2" # e.g. katex-0.16.39
# Find all files with version references (excludes vendored JS and worktrees)
local files
files=$(grep -rl --include='*.toml' --include='*.html' --include='*.js' --include='*.md' \
-F "$old_pattern" . \
--exclude-dir='.claude' --exclude-dir='node_modules' --exclude-dir='shared_static' \
2>/dev/null || true)
for f in $files; do
sed -i "s|${old_pattern}|${new_pattern}|g" "$f"
echo " Updated $f"
done
}
case "$LIB" in
katex)
OLD_VERSION=$(detect_old_version "katex")
OLD_DIR="${STATIC_DIR}/katex-${OLD_VERSION}"
NEW_DIR="${STATIC_DIR}/katex-${VERSION}"
echo "Updating KaTeX ${OLD_VERSION} -> ${VERSION}"
mkdir -p "${NEW_DIR}/fonts"
echo " Downloading katex.min.js..."
curl -sSfL "${CDN}/katex@${VERSION}/dist/katex.min.js" -o "${NEW_DIR}/katex.min.js"
echo " Downloading katex.min.css..."
curl -sSfL "${CDN}/katex@${VERSION}/dist/katex.min.css" -o "${NEW_DIR}/katex.min.css"
echo " Downloading fonts..."
# Extract font filenames from the CSS
font_files=$(curl -sSfL "${CDN}/katex@${VERSION}/dist/katex.min.css" \
| grep -oE 'fonts/[^")]+' | sort -u)
for font in $font_files; do
if ! curl -sSfL "${CDN}/katex@${VERSION}/dist/${font}" -o "${NEW_DIR}/${font}" 2>/dev/null; then
echo " WARNING: Failed to download font: ${font}"
fi
done
# Copy LICENSE from old dir if present
if [[ -f "${OLD_DIR}/LICENSE" ]]; then
cp "${OLD_DIR}/LICENSE" "${NEW_DIR}/LICENSE"
fi
update_refs "katex-${OLD_VERSION}" "katex-${VERSION}"
rm -rf "${OLD_DIR}"
echo "Done. Old directory removed: ${OLD_DIR}"
;;
hljs)
OLD_VERSION=$(detect_old_version "hljs")
OLD_DIR="${STATIC_DIR}/hljs-${OLD_VERSION}"
NEW_DIR="${STATIC_DIR}/hljs-${VERSION}"
echo "Updating Highlight.js ${OLD_VERSION} -> ${VERSION}"
mkdir -p "${NEW_DIR}"
echo " Downloading highlight.min.js..."
curl -sSfL "${CDN}/@highlightjs/cdn-assets@${VERSION}/highlight.min.js" -o "${NEW_DIR}/highlight.min.js"
if [[ -f "${OLD_DIR}/LICENSE" ]]; then
cp "${OLD_DIR}/LICENSE" "${NEW_DIR}/LICENSE"
fi
update_refs "hljs-${OLD_VERSION}" "hljs-${VERSION}"
rm -rf "${OLD_DIR}"
echo "Done. Old directory removed: ${OLD_DIR}"
;;
mermaid)
OLD_VERSION=$(detect_old_version "mermaid")
OLD_DIR="${STATIC_DIR}/mermaid-${OLD_VERSION}"
NEW_DIR="${STATIC_DIR}/mermaid-${VERSION}"
echo "Updating Mermaid ${OLD_VERSION} -> ${VERSION}"
mkdir -p "${NEW_DIR}"
echo " Downloading mermaid.min.js..."
curl -sSfL "${CDN}/mermaid@${VERSION}/dist/mermaid.min.js" -o "${NEW_DIR}/mermaid.min.js"
if [[ -f "${OLD_DIR}/LICENSE" ]]; then
cp "${OLD_DIR}/LICENSE" "${NEW_DIR}/LICENSE"
fi
update_refs "mermaid-${OLD_VERSION}" "mermaid-${VERSION}"
rm -rf "${OLD_DIR}"
echo "Done. Old directory removed: ${OLD_DIR}"
;;
*)
echo "Unknown library: ${LIB}"
usage
;;
esac
echo ""
echo "Verify the update:"
echo " git diff --stat"
echo " python -m turnstone.server # test locally"
+847 -1
View File
@@ -2,7 +2,7 @@
"openapi": "3.1.0",
"info": {
"title": "turnstone Console API",
"version": "0.6.0",
"version": "0.6.2",
"description": "Cluster-wide visibility and control across all turnstone nodes."
},
"paths": {
@@ -460,6 +460,85 @@
}
}
},
"/v1/api/auth/oidc/authorize": {
"get": {
"summary": "Redirect to OIDC provider for SSO login",
"operationId": "v1_api_auth_oidc_authorize_get",
"tags": [
"Auth"
],
"responses": {
"302": {
"description": "Success"
},
"404": {
"description": "Error 404",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
},
"503": {
"description": "Error 503",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
}
}
}
},
"/v1/api/auth/oidc/callback": {
"get": {
"summary": "OIDC callback \u2014 validates code, provisions user, sets JWT cookie, redirects to app",
"operationId": "v1_api_auth_oidc_callback_get",
"tags": [
"Auth"
],
"responses": {
"302": {
"description": "Success"
}
}
}
},
"/v1/api/auth/whoami": {
"get": {
"summary": "Return authenticated user info and permissions",
"operationId": "v1_api_auth_whoami_get",
"tags": [
"Auth"
],
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/AuthWhoamiResponse"
}
}
}
},
"401": {
"description": "Error 401",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
}
}
}
},
"/v1/api/admin/users": {
"get": {
"summary": "List all users",
@@ -2052,6 +2131,27 @@
}
}
},
"/v1/api/templates": {
"get": {
"summary": "List available prompt templates (summary)",
"operationId": "v1_api_templates_get",
"tags": [
"Templates"
],
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ListPromptTemplateSummaryResponse"
}
}
}
}
}
}
},
"/v1/api/admin/usage": {
"get": {
"summary": "Aggregated usage data",
@@ -2661,6 +2761,298 @@
}
}
},
"/v1/api/admin/mcp-servers": {
"get": {
"summary": "List MCP server definitions with live status",
"operationId": "v1_api_admin_mcp-servers_get",
"tags": [
"Admin"
],
"parameters": [
{
"name": "reveal",
"in": "query",
"required": false,
"schema": {
"type": "boolean"
},
"description": "Show secret env/header values"
}
],
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ListMcpServersResponse"
}
}
}
}
}
},
"post": {
"summary": "Create an MCP server definition",
"operationId": "v1_api_admin_mcp-servers_post",
"tags": [
"Admin"
],
"requestBody": {
"required": true,
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/CreateMcpServerRequest"
}
}
}
},
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/McpServerDetail"
}
}
}
},
"400": {
"description": "Error 400",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
},
"409": {
"description": "Error 409",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
}
}
}
},
"/v1/api/admin/mcp-servers/{server_id}": {
"get": {
"summary": "Get a single MCP server with status",
"operationId": "v1_api_admin_mcp-servers_{server_id}_get",
"tags": [
"Admin"
],
"parameters": [
{
"name": "server_id",
"in": "path",
"required": true,
"schema": {
"type": "string"
}
}
],
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/McpServerDetail"
}
}
}
},
"404": {
"description": "Error 404",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
}
}
},
"put": {
"summary": "Update an MCP server definition",
"operationId": "v1_api_admin_mcp-servers_{server_id}_put",
"tags": [
"Admin"
],
"parameters": [
{
"name": "server_id",
"in": "path",
"required": true,
"schema": {
"type": "string"
}
}
],
"requestBody": {
"required": true,
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/UpdateMcpServerRequest"
}
}
}
},
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/McpServerDetail"
}
}
}
},
"400": {
"description": "Error 400",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
},
"404": {
"description": "Error 404",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
},
"409": {
"description": "Error 409",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
}
}
},
"delete": {
"summary": "Delete an MCP server definition",
"operationId": "v1_api_admin_mcp-servers_{server_id}_delete",
"tags": [
"Admin"
],
"parameters": [
{
"name": "server_id",
"in": "path",
"required": true,
"schema": {
"type": "string"
}
}
],
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/StatusResponse"
}
}
}
},
"404": {
"description": "Error 404",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
}
}
}
},
"/v1/api/admin/mcp-servers/reload": {
"post": {
"summary": "Tell all nodes to re-read MCP server config from DB and reconcile",
"operationId": "v1_api_admin_mcp-servers_reload_post",
"tags": [
"Admin"
],
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/McpReloadResponse"
}
}
}
}
}
}
},
"/v1/api/admin/mcp-servers/import": {
"post": {
"summary": "Import MCP servers from a JSON config file",
"operationId": "v1_api_admin_mcp-servers_import_post",
"tags": [
"Admin"
],
"requestBody": {
"required": true,
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ImportMcpConfigRequest"
}
}
}
},
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ImportMcpConfigResponse"
}
}
}
},
"400": {
"description": "Error 400",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
}
}
}
},
"/health": {
"get": {
"summary": "Console health check",
@@ -2866,6 +3258,21 @@
"setup_required": {
"title": "Setup Required",
"type": "boolean"
},
"oidc_enabled": {
"default": false,
"title": "Oidc Enabled",
"type": "boolean"
},
"oidc_provider_name": {
"default": "",
"title": "Oidc Provider Name",
"type": "string"
},
"password_enabled": {
"default": true,
"title": "Password Enabled",
"type": "boolean"
}
},
"required": [
@@ -5264,6 +5671,445 @@
],
"title": "UpdateSettingRequest",
"type": "object"
},
"McpServerDetail": {
"properties": {
"server_id": {
"title": "Server Id",
"type": "string"
},
"name": {
"title": "Name",
"type": "string"
},
"transport": {
"title": "Transport",
"type": "string"
},
"command": {
"default": "",
"title": "Command",
"type": "string"
},
"args": {
"default": "[]",
"title": "Args",
"type": "string"
},
"url": {
"default": "",
"title": "Url",
"type": "string"
},
"headers": {
"default": "{}",
"title": "Headers",
"type": "string"
},
"env": {
"default": "{}",
"title": "Env",
"type": "string"
},
"auto_approve": {
"default": false,
"title": "Auto Approve",
"type": "boolean"
},
"enabled": {
"default": true,
"title": "Enabled",
"type": "boolean"
},
"created_by": {
"default": "",
"title": "Created By",
"type": "string"
},
"created": {
"title": "Created",
"type": "string"
},
"updated": {
"title": "Updated",
"type": "string"
},
"status": {
"additionalProperties": {
"$ref": "#/components/schemas/McpServerStatus"
},
"title": "Status",
"type": "object"
},
"source": {
"default": "",
"title": "Source",
"type": "string"
}
},
"required": [
"server_id",
"name",
"transport",
"created",
"updated"
],
"title": "McpServerDetail",
"type": "object"
},
"McpServerStatus": {
"properties": {
"connected": {
"default": false,
"title": "Connected",
"type": "boolean"
},
"tools": {
"default": 0,
"title": "Tools",
"type": "integer"
},
"resources": {
"default": 0,
"title": "Resources",
"type": "integer"
},
"prompts": {
"default": 0,
"title": "Prompts",
"type": "integer"
},
"error": {
"default": "",
"title": "Error",
"type": "string"
},
"transport": {
"default": "",
"title": "Transport",
"type": "string"
},
"command": {
"default": "",
"title": "Command",
"type": "string"
},
"url": {
"default": "",
"title": "Url",
"type": "string"
}
},
"title": "McpServerStatus",
"type": "object"
},
"CreateMcpServerRequest": {
"properties": {
"name": {
"title": "Name",
"type": "string"
},
"transport": {
"title": "Transport",
"type": "string"
},
"command": {
"default": "",
"title": "Command",
"type": "string"
},
"args": {
"default": [],
"items": {
"type": "string"
},
"title": "Args",
"type": "array"
},
"url": {
"default": "",
"title": "Url",
"type": "string"
},
"headers": {
"additionalProperties": {
"type": "string"
},
"title": "Headers",
"type": "object"
},
"env": {
"additionalProperties": {
"type": "string"
},
"title": "Env",
"type": "object"
},
"auto_approve": {
"default": false,
"title": "Auto Approve",
"type": "boolean"
},
"enabled": {
"default": true,
"title": "Enabled",
"type": "boolean"
}
},
"required": [
"name",
"transport"
],
"title": "CreateMcpServerRequest",
"type": "object"
},
"UpdateMcpServerRequest": {
"properties": {
"name": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"default": null,
"title": "Name"
},
"transport": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"default": null,
"title": "Transport"
},
"command": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"default": null,
"title": "Command"
},
"args": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"default": null,
"title": "Args"
},
"url": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"default": null,
"title": "Url"
},
"headers": {
"anyOf": [
{
"additionalProperties": {
"type": "string"
},
"type": "object"
},
{
"type": "null"
}
],
"default": null,
"title": "Headers"
},
"env": {
"anyOf": [
{
"additionalProperties": {
"type": "string"
},
"type": "object"
},
{
"type": "null"
}
],
"default": null,
"title": "Env"
},
"auto_approve": {
"anyOf": [
{
"type": "boolean"
},
{
"type": "null"
}
],
"default": null,
"title": "Auto Approve"
},
"enabled": {
"anyOf": [
{
"type": "boolean"
},
{
"type": "null"
}
],
"default": null,
"title": "Enabled"
}
},
"title": "UpdateMcpServerRequest",
"type": "object"
},
"ListMcpServersResponse": {
"properties": {
"servers": {
"items": {
"$ref": "#/components/schemas/McpServerDetail"
},
"title": "Servers",
"type": "array"
}
},
"required": [
"servers"
],
"title": "ListMcpServersResponse",
"type": "object"
},
"ImportMcpConfigRequest": {
"properties": {
"config": {
"additionalProperties": true,
"description": "JSON config object with mcpServers key",
"title": "Config",
"type": "object"
}
},
"required": [
"config"
],
"title": "ImportMcpConfigRequest",
"type": "object"
},
"ImportMcpConfigResponse": {
"properties": {
"imported": {
"default": [],
"items": {
"type": "string"
},
"title": "Imported",
"type": "array"
},
"skipped": {
"default": [],
"items": {
"type": "string"
},
"title": "Skipped",
"type": "array"
},
"errors": {
"default": [],
"items": {
"type": "string"
},
"title": "Errors",
"type": "array"
}
},
"title": "ImportMcpConfigResponse",
"type": "object"
},
"McpReloadResponse": {
"properties": {
"status": {
"default": "ok",
"title": "Status",
"type": "string"
},
"results": {
"additionalProperties": true,
"title": "Results",
"type": "object"
}
},
"title": "McpReloadResponse",
"type": "object"
},
"PromptTemplateSummary": {
"properties": {
"name": {
"description": "Template name",
"title": "Name",
"type": "string"
},
"category": {
"default": "",
"description": "Template category",
"title": "Category",
"type": "string"
},
"is_default": {
"default": false,
"description": "Whether this template is applied by default",
"title": "Is Default",
"type": "boolean"
},
"origin": {
"default": "manual",
"description": "Template origin: manual or mcp",
"title": "Origin",
"type": "string"
}
},
"required": [
"name"
],
"title": "PromptTemplateSummary",
"type": "object"
},
"ListPromptTemplateSummaryResponse": {
"properties": {
"templates": {
"items": {
"$ref": "#/components/schemas/PromptTemplateSummary"
},
"title": "Templates",
"type": "array"
}
},
"required": [
"templates"
],
"title": "ListPromptTemplateSummaryResponse",
"type": "object"
}
}
}
+224 -1
View File
@@ -2,7 +2,7 @@
"openapi": "3.1.0",
"info": {
"title": "turnstone Server API",
"version": "0.6.0",
"version": "0.6.2",
"description": "Single-node workstream management, chat interaction, and real-time streaming."
},
"paths": {
@@ -437,6 +437,48 @@
}
}
},
"/v1/api/templates": {
"get": {
"summary": "List available prompt templates (summary)",
"operationId": "v1_api_templates_get",
"tags": [
"Templates"
],
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ListPromptTemplateSummaryResponse"
}
}
}
}
}
}
},
"/v1/api/ws-templates": {
"get": {
"summary": "List enabled workstream templates (summary)",
"operationId": "v1_api_ws-templates_get",
"tags": [
"Templates"
],
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ListWsTemplateSummaryResponse"
}
}
}
}
}
}
},
"/v1/api/auth/login": {
"post": {
"summary": "Authenticate with a token",
@@ -581,6 +623,85 @@
}
}
},
"/v1/api/auth/oidc/authorize": {
"get": {
"summary": "Redirect to OIDC provider for SSO login",
"operationId": "v1_api_auth_oidc_authorize_get",
"tags": [
"Auth"
],
"responses": {
"302": {
"description": "Success"
},
"404": {
"description": "Error 404",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
},
"503": {
"description": "Error 503",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
}
}
}
},
"/v1/api/auth/oidc/callback": {
"get": {
"summary": "OIDC callback \u2014 validates code, provisions user, sets JWT cookie, redirects to app",
"operationId": "v1_api_auth_oidc_callback_get",
"tags": [
"Auth"
],
"responses": {
"302": {
"description": "Success"
}
}
}
},
"/v1/api/auth/whoami": {
"get": {
"summary": "Return authenticated user info and permissions",
"operationId": "v1_api_auth_whoami_get",
"tags": [
"Auth"
],
"responses": {
"200": {
"description": "Success",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/AuthWhoamiResponse"
}
}
}
},
"401": {
"description": "Error 401",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
}
}
}
},
"/v1/api/memories": {
"get": {
"summary": "List structured memories",
@@ -975,6 +1096,21 @@
"setup_required": {
"title": "Setup Required",
"type": "boolean"
},
"oidc_enabled": {
"default": false,
"title": "Oidc Enabled",
"type": "boolean"
},
"oidc_provider_name": {
"default": "",
"title": "Oidc Provider Name",
"type": "string"
},
"password_enabled": {
"default": true,
"title": "Password Enabled",
"type": "boolean"
}
},
"required": [
@@ -1777,6 +1913,93 @@
],
"title": "SearchMemoriesRequest",
"type": "object"
},
"PromptTemplateSummary": {
"properties": {
"name": {
"description": "Template name",
"title": "Name",
"type": "string"
},
"category": {
"default": "",
"description": "Template category",
"title": "Category",
"type": "string"
},
"is_default": {
"default": false,
"description": "Whether this template is applied by default",
"title": "Is Default",
"type": "boolean"
},
"origin": {
"default": "manual",
"description": "Template origin: manual or mcp",
"title": "Origin",
"type": "string"
}
},
"required": [
"name"
],
"title": "PromptTemplateSummary",
"type": "object"
},
"ListPromptTemplateSummaryResponse": {
"properties": {
"templates": {
"items": {
"$ref": "#/components/schemas/PromptTemplateSummary"
},
"title": "Templates",
"type": "array"
}
},
"required": [
"templates"
],
"title": "ListPromptTemplateSummaryResponse",
"type": "object"
},
"WsTemplateSummary": {
"properties": {
"name": {
"title": "Name",
"type": "string"
},
"description": {
"title": "Description",
"type": "string"
},
"model": {
"title": "Model",
"type": "string"
}
},
"required": [
"name",
"description",
"model"
],
"title": "WsTemplateSummary",
"type": "object"
},
"ListWsTemplateSummaryResponse": {
"properties": {
"ws_templates": {
"items": {
"$ref": "#/components/schemas/WsTemplateSummary"
},
"title": "Ws Templates",
"type": "array"
}
},
"required": [
"ws_templates"
],
"title": "ListWsTemplateSummaryResponse",
"type": "object"
}
}
}
+773 -843
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -33,6 +33,6 @@
"license": "BUSL-1.1",
"devDependencies": {
"typescript": "^5.4",
"vitest": "^2.0"
"vitest": "^4.1"
}
}
+50
View File
@@ -16,16 +16,20 @@ import type {
ConsoleCreateWsRequest,
ConsoleCreateWsResponse,
ConsoleHealthResponse,
CreateMcpServerRequest,
CreatePolicyOptions,
CreateRoleOptions,
CreateScheduleRequest,
CreateTemplateOptions,
CreateWsTemplateOptions,
ImportMcpConfigResponse,
ListAdminMemoriesResponse,
ListMcpServersResponse,
ListScheduleRunsResponse,
ListSchedulesResponse,
ListSettingSchemaResponse,
ListSettingsResponse,
McpServerDetail,
NodeDetailResponse,
NodesOptions,
OrgInfo,
@@ -35,6 +39,7 @@ import type {
SettingInfo,
StatusResponse,
ToolPolicyInfo,
UpdateMcpServerRequest,
UpdateOrgOptions,
UpdatePolicyOptions,
UpdateRoleOptions,
@@ -410,4 +415,49 @@ export class TurnstoneConsole extends BaseClient {
params,
});
}
// -- MCP servers ----------------------------------------------------------
async listMcpServers(opts?: {
reveal?: boolean;
}): Promise<ListMcpServersResponse> {
const params: Record<string, string> = {};
if (opts?.reveal) params.reveal = "true";
return this.request("GET", "/v1/api/admin/mcp-servers", { params });
}
async createMcpServer(
body: CreateMcpServerRequest,
): Promise<McpServerDetail> {
return this.request("POST", "/v1/api/admin/mcp-servers", { json: body });
}
async getMcpServer(serverId: string): Promise<McpServerDetail> {
return this.request("GET", `/v1/api/admin/mcp-servers/${serverId}`);
}
async updateMcpServer(
serverId: string,
body: UpdateMcpServerRequest,
): Promise<McpServerDetail> {
return this.request("PUT", `/v1/api/admin/mcp-servers/${serverId}`, {
json: body,
});
}
async deleteMcpServer(serverId: string): Promise<StatusResponse> {
return this.request("DELETE", `/v1/api/admin/mcp-servers/${serverId}`);
}
async reloadMcpServers(): Promise<StatusResponse> {
return this.request("POST", "/v1/api/admin/mcp-servers/reload");
}
async importMcpConfig(
config: Record<string, unknown>,
): Promise<ImportMcpConfigResponse> {
return this.request("POST", "/v1/api/admin/mcp-servers/import", {
json: { config },
});
}
}
+2
View File
@@ -115,6 +115,8 @@ export interface WsStateEvent {
context_ratio: number;
activity: string;
activity_state: string;
/** Full assistant response text — populated on idle transitions only. */
content?: string;
}
export interface WsActivityEvent {
+7
View File
@@ -160,6 +160,13 @@ export type {
SettingSchemaInfo,
ListSettingSchemaResponse,
UpdateSettingOptions,
// MCP server types
McpServerStatus,
McpServerDetail,
ListMcpServersResponse,
CreateMcpServerRequest,
UpdateMcpServerRequest,
ImportMcpConfigResponse,
} from "./types.js";
// SSE parser (for advanced usage)
+12
View File
@@ -11,7 +11,9 @@ import type {
HealthResponse,
ListMemoriesOptions,
ListMemoriesResponse,
ListPromptTemplateSummaryResponse,
ListSavedWorkstreamsResponse,
ListWsTemplateSummaryResponse,
ListWorkstreamsResponse,
MemoryInfo,
SaveMemoryRequest,
@@ -196,6 +198,16 @@ export class TurnstoneServer extends BaseClient {
return this.request("GET", "/v1/api/workstreams/saved");
}
// -- Templates --------------------------------------------------------------
async listTemplates(): Promise<ListPromptTemplateSummaryResponse> {
return this.request("GET", "/v1/api/templates");
}
async listWsTemplates(): Promise<ListWsTemplateSummaryResponse> {
return this.request("GET", "/v1/api/ws-templates");
}
// -- Memories -------------------------------------------------------------
async listMemories(
+90
View File
@@ -143,6 +143,35 @@ export interface ListSavedWorkstreamsResponse {
workstreams: SavedWorkstreamInfo[];
}
// ---------------------------------------------------------------------------
// Server API — Prompt templates
// ---------------------------------------------------------------------------
export interface PromptTemplateSummary {
name: string;
category: string;
is_default: boolean;
origin: string;
}
export interface ListPromptTemplateSummaryResponse {
templates: PromptTemplateSummary[];
}
// ---------------------------------------------------------------------------
// Server API — Workstream templates
// ---------------------------------------------------------------------------
export interface WsTemplateSummary {
name: string;
description: string;
model: string;
}
export interface ListWsTemplateSummaryResponse {
ws_templates: WsTemplateSummary[];
}
// ---------------------------------------------------------------------------
// Server API — Health
// ---------------------------------------------------------------------------
@@ -728,6 +757,67 @@ export interface AdminSearchMemoriesOptions {
limit?: number;
}
// -- Console API: MCP Servers -----------------------------------------------
export interface McpServerStatus {
connected: boolean;
tools: number;
resources: number;
prompts: number;
error: string;
}
export interface McpServerDetail {
server_id: string;
name: string;
transport: string;
command: string;
args: string;
url: string;
headers: string;
env: string;
auto_approve: boolean;
enabled: boolean;
created_by: string;
created: string;
updated: string;
status: Record<string, McpServerStatus>;
}
export interface ListMcpServersResponse {
servers: McpServerDetail[];
}
export interface CreateMcpServerRequest {
name: string;
transport: string;
command?: string;
args?: string[];
url?: string;
headers?: Record<string, string>;
env?: Record<string, string>;
auto_approve?: boolean;
enabled?: boolean;
}
export interface UpdateMcpServerRequest {
name?: string;
transport?: string;
command?: string;
args?: string[];
url?: string;
headers?: Record<string, string>;
env?: Record<string, string>;
auto_approve?: boolean;
enabled?: boolean;
}
export interface ImportMcpConfigResponse {
imported: string[];
skipped: string[];
errors: string[];
}
// -- Console API: System Settings -------------------------------------------
export interface SettingInfo {
+75
View File
@@ -169,6 +169,35 @@ class TestRequiredScope:
def test_admin_memory_delete_needs_approve(self):
assert required_scope("DELETE", "/api/admin/memories/some-id") == "approve"
# Internal endpoints
def test_internal_mcp_reload_needs_approve(self):
assert required_scope("POST", "/api/_internal/mcp-reload") == "approve"
def test_v1_internal_mcp_reload_needs_approve(self):
assert required_scope("POST", "/v1/api/_internal/mcp-reload") == "approve"
def test_internal_config_reload_needs_approve(self):
assert required_scope("POST", "/api/_internal/config-reload") == "approve"
def test_v1_internal_config_reload_needs_approve(self):
assert required_scope("POST", "/v1/api/_internal/config-reload") == "approve"
def test_proxy_internal_config_reload_needs_approve(self):
assert required_scope("POST", "/node/n1/v1/api/_internal/config-reload") == "approve"
def test_proxy_no_v1_internal_config_reload_needs_approve(self):
assert required_scope("POST", "/node/n1/api/_internal/config-reload") == "approve"
def test_proxy_internal_mcp_reload_needs_approve(self):
assert required_scope("POST", "/node/n1/v1/api/_internal/mcp-reload") == "approve"
def test_proxy_no_v1_internal_mcp_reload_needs_approve(self):
assert required_scope("POST", "/node/n1/api/_internal/mcp-reload") == "approve"
def test_get_internal_mcp_reload_needs_read(self):
"""Only POST is elevated — GET falls through to read."""
assert required_scope("GET", "/api/_internal/mcp-reload") == "read"
# ---------------------------------------------------------------------------
# TestAuthConfig
@@ -1375,3 +1404,49 @@ class TestCorsConfigurable:
)
assert resp.headers.get("Access-Control-Allow-Origin") == "http://example.com"
client.close()
# ---------------------------------------------------------------------------
# TestVerifyPassword — OIDC sentinel handling
# ---------------------------------------------------------------------------
class TestVerifyPassword:
def test_valid_bcrypt_hash(self):
from turnstone.core.auth import hash_password, verify_password
hashed = hash_password("mypassword")
assert verify_password("mypassword", hashed) is True
assert verify_password("wrongpassword", hashed) is False
def test_oidc_sentinel_rejected(self):
from turnstone.core.auth import verify_password
# OIDC sentinel must return False, not crash with ValueError
assert verify_password("anypassword", "!oidc") is False
def test_non_bcrypt_hash_rejected(self):
from turnstone.core.auth import verify_password
assert verify_password("password", "not_a_hash") is False
assert verify_password("password", "") is False
def test_empty_password_against_oidc_sentinel(self):
from turnstone.core.auth import verify_password
assert verify_password("", "!oidc") is False
# ---------------------------------------------------------------------------
# TestOIDCPublicPaths — OIDC endpoints are public
# ---------------------------------------------------------------------------
class TestOIDCPublicPaths:
def test_oidc_authorize_is_public(self):
assert is_public_path("/api/auth/oidc/authorize") is True
assert is_public_path("/v1/api/auth/oidc/authorize") is True
def test_oidc_callback_is_public(self):
assert is_public_path("/api/auth/oidc/callback") is True
assert is_public_path("/v1/api/auth/oidc/callback") is True
+52 -1
View File
@@ -3,7 +3,7 @@
from unittest.mock import MagicMock, patch
from turnstone.mq.bridge import Bridge
from turnstone.mq.protocol import StateChangeEvent, TurnCompleteEvent
from turnstone.mq.protocol import ContentEvent, StateChangeEvent, TurnCompleteEvent
def _make_bridge():
@@ -67,3 +67,54 @@ class TestIdleTurnComplete:
assert len(state_changes) == 1
assert state_changes[0].state == "thinking"
assert len(turn_completes) == 0
class TestContentPassthrough:
"""Bridge should pass through content from the server's idle SSE event."""
def test_content_passed_through_in_turn_complete(self):
"""Content from idle event should be included in TurnCompleteEvent."""
bridge = _make_bridge()
published = []
with patch.object(
bridge, "_publish_ws", side_effect=lambda ws, ev: published.append((ws, ev))
):
bridge._handle_global_event(
{"type": "ws_state", "ws_id": "ws-1", "state": "idle", "content": "Hello world"}
)
turn_completes = [(ws, ev) for ws, ev in published if isinstance(ev, TurnCompleteEvent)]
assert len(turn_completes) == 1
_, ev = turn_completes[0]
assert ev.content == "Hello world"
def test_content_empty_when_not_in_event(self):
"""TurnCompleteEvent.content should be empty when idle event has no content."""
bridge = _make_bridge()
published = []
with patch.object(
bridge, "_publish_ws", side_effect=lambda ws, ev: published.append((ws, ev))
):
bridge._handle_global_event({"type": "ws_state", "ws_id": "ws-1", "state": "idle"})
turn_completes = [(ws, ev) for ws, ev in published if isinstance(ev, TurnCompleteEvent)]
assert len(turn_completes) == 1
_, ev = turn_completes[0]
assert ev.content == ""
def test_content_event_still_published(self):
"""Content events should still be published to per-ws channel."""
bridge = _make_bridge()
published = []
with patch.object(
bridge, "_publish_ws", side_effect=lambda ws, ev: published.append((ws, ev))
):
bridge._handle_ws_event("ws-1", {"type": "content", "text": "hello"})
content_events = [(ws, ev) for ws, ev in published if isinstance(ev, ContentEvent)]
assert len(content_events) == 1
_, ev = content_events[0]
assert ev.text == "hello"
+317 -1
View File
@@ -21,7 +21,7 @@ def _run(coro):
return asyncio.run(coro)
def _make_message(*, bot=False, guild=True, content="hello", channel=None):
def _make_message(*, bot=False, guild=True, content="hello", channel=None, reference=None):
"""Build a mock ``discord.Message``."""
msg = MagicMock(spec=discord.Message)
msg.author = MagicMock()
@@ -31,6 +31,7 @@ def _make_message(*, bot=False, guild=True, content="hello", channel=None):
msg.guild = MagicMock() if guild else None
msg.channel = channel or MagicMock()
msg.mentions = []
msg.reference = reference
return msg
@@ -204,6 +205,8 @@ class TestMessageCog:
ts.router.send_message = AsyncMock()
ts.config = MagicMock()
ts._ws_tasks = {}
ts._notify_ws_map = {}
ts._notify_reply_channels = {}
bot.turnstone = ts
cog = MessageCog(bot)
@@ -328,6 +331,7 @@ class TestWsEventFinalization:
bot.config.auto_approve_tools = []
bot._streaming = {}
bot._pending_approval_msgs = {}
bot._notify_reply_channels = {}
# Use the real _on_ws_event method
bot._on_ws_event = TurnstoneBot._on_ws_event.__get__(bot, TurnstoneBot)
@@ -356,6 +360,7 @@ class TestWsEventFinalization:
bot = MagicMock(spec=TurnstoneBot)
bot._streaming = {}
bot._pending_approval_msgs = {}
bot._notify_reply_channels = {}
bot._on_ws_event = TurnstoneBot._on_ws_event.__get__(bot, TurnstoneBot)
thread = AsyncMock()
@@ -387,6 +392,7 @@ class TestApprovalVerdictDisplay:
bot.config.auto_approve_tools = []
bot._streaming = {}
bot._pending_approval_msgs = {}
bot._notify_reply_channels = {}
bot._should_auto_approve = MagicMock(return_value=False)
bot._on_ws_event = TurnstoneBot._on_ws_event.__get__(bot, TurnstoneBot)
return bot
@@ -504,6 +510,7 @@ class TestApprovalVerdictDisplay:
bot = MagicMock(spec=TurnstoneBot)
bot._streaming = {}
bot._pending_approval_msgs = {"ws-1": MagicMock()}
bot._notify_reply_channels = {}
bot._on_ws_event = TurnstoneBot._on_ws_event.__get__(bot, TurnstoneBot)
thread = AsyncMock()
@@ -513,6 +520,315 @@ class TestApprovalVerdictDisplay:
assert "ws-1" not in bot._pending_approval_msgs
class TestContentCatchup:
"""TurnCompleteEvent with content field provides catch-up for missed ContentEvents."""
def _make_bot(self):
from turnstone.channels.discord.bot import TurnstoneBot
bot = MagicMock(spec=TurnstoneBot)
bot.config = MagicMock()
bot.config.max_message_length = 2000
bot.config.streaming_edit_interval = 1.5
bot.config.auto_approve = False
bot.config.auto_approve_tools = []
bot._streaming = {}
bot._pending_approval_msgs = {}
bot._notify_reply_channels = {}
bot._on_ws_event = TurnstoneBot._on_ws_event.__get__(bot, TurnstoneBot)
return bot
def test_catchup_sends_content_when_no_streaming(self):
"""TurnCompleteEvent with content but no SM sends catch-up message."""
from turnstone.mq.protocol import TurnCompleteEvent
bot = self._make_bot()
thread = AsyncMock()
raw = TurnCompleteEvent(
ws_id="ws-1", correlation_id="", content="Caught up response"
).to_json()
_run(bot._on_ws_event("ws-1", thread, raw))
thread.send.assert_awaited_once_with("Caught up response")
def test_catchup_skipped_when_streaming_exists(self):
"""TurnCompleteEvent with content and existing SM uses SM finalize, not catch-up."""
from turnstone.mq.protocol import ContentEvent, TurnCompleteEvent
bot = self._make_bot()
thread = AsyncMock()
# Feed content event to create SM
content_raw = ContentEvent(ws_id="ws-1", text="Streamed").to_json()
_run(bot._on_ws_event("ws-1", thread, content_raw))
assert "ws-1" in bot._streaming
# Now TurnCompleteEvent with content — SM should be finalized, not catch-up
complete_raw = TurnCompleteEvent(
ws_id="ws-1", correlation_id="", content="Streamed"
).to_json()
_run(bot._on_ws_event("ws-1", thread, complete_raw))
assert "ws-1" not in bot._streaming
def test_catchup_empty_content_no_message(self):
"""TurnCompleteEvent with empty content and no SM sends nothing."""
from turnstone.mq.protocol import TurnCompleteEvent
bot = self._make_bot()
thread = AsyncMock()
raw = TurnCompleteEvent(ws_id="ws-1", correlation_id="", content="").to_json()
_run(bot._on_ws_event("ws-1", thread, raw))
thread.send.assert_not_awaited()
class TestNotificationTracking:
"""Tests for notification message tracking and DM reply routing."""
def test_send_notification_tracks_message(self):
"""send_notification should store message_id -> (ws_id, target_user) mapping."""
from turnstone.channels.discord.bot import TurnstoneBot
bot = MagicMock(spec=TurnstoneBot)
bot._notify_ws_map = {}
bot._MAX_NOTIFY_TRACKING = 100
bot.send = AsyncMock(return_value="12345")
bot.send_notification = TurnstoneBot.send_notification.__get__(bot, TurnstoneBot)
bot._track_notification = TurnstoneBot._track_notification.__get__(bot, TurnstoneBot)
_run(bot.send_notification("chan-1", "Hello", "ws-abc"))
assert 12345 in bot._notify_ws_map
assert bot._notify_ws_map[12345] == ("ws-abc", "chan-1")
def test_send_notification_evicts_old_entries(self):
"""Oldest notification tracking entries are evicted when cap is reached."""
from turnstone.channels.discord.bot import TurnstoneBot
bot = MagicMock(spec=TurnstoneBot)
bot._MAX_NOTIFY_TRACKING = 3
bot._notify_ws_map = {
1: ("ws-1", "u1"),
2: ("ws-2", "u2"),
3: ("ws-3", "u3"),
}
bot.send = AsyncMock(return_value="4")
bot.send_notification = TurnstoneBot.send_notification.__get__(bot, TurnstoneBot)
bot._track_notification = TurnstoneBot._track_notification.__get__(bot, TurnstoneBot)
_run(bot.send_notification("chan-1", "Hello", "ws-4"))
assert 4 in bot._notify_ws_map
assert 1 not in bot._notify_ws_map # oldest evicted
assert len(bot._notify_ws_map) <= 3
def test_dm_reply_routes_to_workstream(self):
"""DM reply to a tracked notification routes the message to the workstream."""
from turnstone.channels.discord.cog import MessageCog
bot = MagicMock()
bot.user = MagicMock()
bot.user.id = 99999
ts = MagicMock()
ts._is_allowed_channel = MagicMock(return_value=True)
ts.storage = MagicMock()
ts.router = MagicMock()
ts.router.resolve_user = AsyncMock(return_value="u_abc")
ts.router.send_message = AsyncMock()
ts.config = MagicMock()
# Maps message_id -> (ws_id, target_discord_user_id)
ts._notify_ws_map = {77777: ("ws-target", "12345")}
ts._notify_reply_channels = {}
bot.turnstone = ts
cog = MessageCog(bot)
# Build a DM reply to the tracked notification message
ref = MagicMock()
ref.message_id = 77777
msg = _make_message(guild=False, content="additional context", reference=ref)
# msg.author.id defaults to 12345 from _make_message
_run(cog._on_message(msg))
ts.router.send_message.assert_awaited_once_with("ws-target", "additional context")
assert "ws-target" in ts._notify_reply_channels
dm_chan, target_uid = ts._notify_reply_channels["ws-target"]
assert target_uid == "12345"
assert 77777 not in ts._notify_ws_map # cleaned up
def test_dm_reply_user_mismatch_rejected_and_preserved(self):
"""DM reply from wrong user is rejected; entry re-inserted for legitimate user."""
from turnstone.channels.discord.cog import MessageCog
bot = MagicMock()
bot.user = MagicMock()
bot.user.id = 99999
ts = MagicMock()
ts.router = MagicMock()
ts.router.resolve_user = AsyncMock(return_value="u_abc")
ts.router.send_message = AsyncMock()
# Target user is "99999" but replying user has author.id = 12345
ts._notify_ws_map = {77777: ("ws-target", "99999")}
ts._notify_reply_channels = {}
bot.turnstone = ts
cog = MessageCog(bot)
ref = MagicMock()
ref.message_id = 77777
msg = _make_message(guild=False, content="impostor", reference=ref)
_run(cog._on_message(msg))
ts.router.send_message.assert_not_awaited()
# Entry should be re-inserted so the legitimate user can still reply.
assert 77777 in ts._notify_ws_map
assert ts._notify_ws_map[77777] == ("ws-target", "99999")
def test_dm_reply_stale_notification_feedback(self):
"""DM reply to an expired/unknown notification should inform the user."""
from turnstone.channels.discord.cog import MessageCog
bot = MagicMock()
bot.user = MagicMock()
bot.user.id = 99999
ts = MagicMock()
ts.router = MagicMock()
ts.router.send_message = AsyncMock()
ts._notify_ws_map = {} # empty — no tracked notifications
ts._notify_reply_channels = {}
bot.turnstone = ts
cog = MessageCog(bot)
ref = MagicMock()
ref.message_id = 99999 # not in map
dm_channel = AsyncMock()
msg = _make_message(guild=False, content="reply", reference=ref, channel=dm_channel)
_run(cog._on_message(msg))
# Should NOT route to any workstream
ts.router.send_message.assert_not_awaited()
# Should send feedback to the DM channel
dm_channel.send.assert_awaited_once_with("*This notification is no longer active.*")
def test_dm_without_reference_ignored(self):
"""DM without a message reference should be ignored."""
from turnstone.channels.discord.cog import MessageCog
bot = MagicMock()
bot.user = MagicMock()
bot.user.id = 99999
ts = MagicMock()
ts.router = MagicMock()
ts.router.send_message = AsyncMock()
ts._notify_ws_map = {77777: ("ws-target", "12345")}
ts._notify_reply_channels = {}
bot.turnstone = ts
cog = MessageCog(bot)
msg = _make_message(guild=False) # reference=None
_run(cog._on_message(msg))
ts.router.send_message.assert_not_awaited()
def test_dm_reply_unlinked_user_ignored(self):
"""DM reply from an unlinked user should be ignored."""
from turnstone.channels.discord.cog import MessageCog
bot = MagicMock()
bot.user = MagicMock()
bot.user.id = 99999
ts = MagicMock()
ts.router = MagicMock()
ts.router.resolve_user = AsyncMock(return_value=None)
ts.router.send_message = AsyncMock()
ts._notify_ws_map = {77777: ("ws-target", "12345")}
ts._notify_reply_channels = {}
bot.turnstone = ts
cog = MessageCog(bot)
ref = MagicMock()
ref.message_id = 77777
msg = _make_message(guild=False, content="reply", reference=ref)
_run(cog._on_message(msg))
ts.router.send_message.assert_not_awaited()
def test_turn_complete_forwards_to_dm(self):
"""TurnCompleteEvent should forward content to notification reply DM."""
from turnstone.channels.discord.bot import TurnstoneBot
from turnstone.mq.protocol import TurnCompleteEvent
bot = MagicMock(spec=TurnstoneBot)
bot.config = MagicMock()
bot.config.max_message_length = 2000
bot._streaming = {}
bot._pending_approval_msgs = {}
bot._notify_ws_map = {}
bot._MAX_NOTIFY_TRACKING = 100
dm_channel = AsyncMock()
sent_msg = MagicMock()
sent_msg.id = 88888
dm_channel.send = AsyncMock(return_value=sent_msg)
bot._notify_reply_channels = {"ws-1": (dm_channel, "u123")}
bot._on_ws_event = TurnstoneBot._on_ws_event.__get__(bot, TurnstoneBot)
bot._track_notification = TurnstoneBot._track_notification.__get__(bot, TurnstoneBot)
thread = AsyncMock()
raw = TurnCompleteEvent(
ws_id="ws-1", correlation_id="", content="Here's the response"
).to_json()
_run(bot._on_ws_event("ws-1", thread, raw))
# Should send to DM channel
dm_channel.send.assert_awaited_once_with("Here's the response")
# Should clean up forwarding
assert "ws-1" not in bot._notify_reply_channels
# Response message should be tracked for multi-turn replies
assert 88888 in bot._notify_ws_map
assert bot._notify_ws_map[88888] == ("ws-1", "u123")
def test_turn_complete_cleans_up_dm_even_without_content(self):
"""TurnCompleteEvent without content should still clean up DM tracking."""
from turnstone.channels.discord.bot import TurnstoneBot
from turnstone.mq.protocol import TurnCompleteEvent
bot = MagicMock(spec=TurnstoneBot)
bot._streaming = {}
bot._pending_approval_msgs = {}
bot._notify_ws_map = {}
dm_channel = AsyncMock()
bot._notify_reply_channels = {"ws-1": (dm_channel, "u123")}
bot._on_ws_event = TurnstoneBot._on_ws_event.__get__(bot, TurnstoneBot)
thread = AsyncMock()
raw = TurnCompleteEvent(ws_id="ws-1", correlation_id="", content="").to_json()
_run(bot._on_ws_event("ws-1", thread, raw))
# DM should not be sent to (no content)
dm_channel.send.assert_not_awaited()
# But should still be cleaned up
assert "ws-1" not in bot._notify_reply_channels
# No response tracked (nothing was sent)
assert len(bot._notify_ws_map) == 0
class TestChannelCLI:
"""Tests for the channel CLI entry point."""
+695
View File
@@ -0,0 +1,695 @@
"""Tests for MCP server admin API endpoints."""
from __future__ import annotations
import json
import uuid
from typing import TYPE_CHECKING, Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from starlette.applications import Starlette
from starlette.middleware import Middleware
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.routing import Mount, Route
from starlette.testclient import TestClient
if TYPE_CHECKING:
from starlette.requests import Request
from starlette.responses import Response
from turnstone.console.server import (
_collect_mcp_status,
_notify_nodes_mcp_reload,
admin_create_mcp_server,
admin_delete_mcp_server,
admin_get_mcp_server,
admin_import_mcp_config,
admin_list_mcp_servers,
admin_update_mcp_server,
)
from turnstone.core.auth import AuthResult
from turnstone.core.storage._sqlite import SQLiteBackend
# ---------------------------------------------------------------------------
# Auth middleware variants
# ---------------------------------------------------------------------------
class _InjectAuthMiddleware(BaseHTTPMiddleware):
"""Inject an admin auth result with admin.mcp permission."""
async def dispatch(self, request: Request, call_next: Any) -> Response:
request.state.auth_result = AuthResult(
user_id="test-user",
scopes=frozenset({"approve"}),
token_source="config",
permissions=frozenset(
{
"read",
"write",
"approve",
"admin.mcp",
}
),
)
resp: Response = await call_next(request)
return resp
class _InjectAuthNoMcpMiddleware(BaseHTTPMiddleware):
"""Inject an auth result WITHOUT admin.mcp permission."""
async def dispatch(self, request: Request, call_next: Any) -> Response:
request.state.auth_result = AuthResult(
user_id="test-user",
scopes=frozenset({"approve"}),
token_source="jwt",
permissions=frozenset(
{
"read",
"write",
"approve",
}
),
)
resp: Response = await call_next(request)
return resp
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
_ROUTES = [
Mount(
"/v1",
routes=[
Route("/api/admin/mcp-servers", admin_list_mcp_servers),
Route(
"/api/admin/mcp-servers",
admin_create_mcp_server,
methods=["POST"],
),
Route(
"/api/admin/mcp-servers/import",
admin_import_mcp_config,
methods=["POST"],
),
Route(
"/api/admin/mcp-servers/{server_id}",
admin_get_mcp_server,
),
Route(
"/api/admin/mcp-servers/{server_id}",
admin_update_mcp_server,
methods=["PUT"],
),
Route(
"/api/admin/mcp-servers/{server_id}",
admin_delete_mcp_server,
methods=["DELETE"],
),
],
),
]
@pytest.fixture
def storage(tmp_path):
return SQLiteBackend(str(tmp_path / "test.db"))
@pytest.fixture
def client(storage):
"""TestClient wired to console admin MCP endpoints with full permissions."""
app = Starlette(
routes=_ROUTES,
middleware=[Middleware(_InjectAuthMiddleware)],
)
app.state.auth_storage = storage
return TestClient(app)
@pytest.fixture
def client_no_perm(storage):
"""TestClient without admin.mcp permission."""
app = Starlette(
routes=_ROUTES,
middleware=[Middleware(_InjectAuthNoMcpMiddleware)],
)
app.state.auth_storage = storage
return TestClient(app)
def _create_server(
client: TestClient,
*,
name: str = "test-server",
transport: str = "stdio",
command: str = "npx",
args: list[str] | None = None,
env: dict[str, str] | None = None,
headers: dict[str, str] | None = None,
url: str = "",
) -> dict[str, Any]:
"""Helper to create a server via the API and return the response dict."""
body: dict[str, Any] = {"name": name, "transport": transport}
if transport == "stdio":
body["command"] = command
body["args"] = args or ["-y", "@modelcontextprotocol/server-test"]
else:
body["url"] = url or "http://localhost:8080/mcp"
if env is not None:
body["env"] = env
if headers is not None:
body["headers"] = headers
r = client.post("/v1/api/admin/mcp-servers", json=body)
assert r.status_code == 200
data: dict[str, Any] = r.json()
return data
# ---------------------------------------------------------------------------
# Mock _collect_mcp_status to avoid real HTTP calls
# ---------------------------------------------------------------------------
_PATCH_MCP_STATUS = patch(
"turnstone.console.server._collect_mcp_status",
new_callable=AsyncMock,
return_value={},
)
# ---------------------------------------------------------------------------
# List
# ---------------------------------------------------------------------------
class TestListMcpServers:
def test_list_empty(self, client):
with _PATCH_MCP_STATUS:
r = client.get("/v1/api/admin/mcp-servers")
assert r.status_code == 200
assert r.json()["servers"] == []
def test_list_returns_created_servers(self, client):
_create_server(client, name="server-a")
_create_server(client, name="server-b")
with _PATCH_MCP_STATUS:
r = client.get("/v1/api/admin/mcp-servers")
assert r.status_code == 200
names = [s["name"] for s in r.json()["servers"]]
assert "server-a" in names
assert "server-b" in names
# ---------------------------------------------------------------------------
# Create
# ---------------------------------------------------------------------------
class TestCreateMcpServer:
def test_create_stdio_server(self, client):
data = _create_server(client, name="my-mcp", transport="stdio", command="node")
assert data["name"] == "my-mcp"
assert data["transport"] == "stdio"
assert data["command"] == "node"
assert data["server_id"]
assert data["enabled"] is True
def test_create_http_server(self, client):
data = _create_server(
client,
name="remote-mcp",
transport="streamable-http",
url="http://mcp.example.com/sse",
)
assert data["name"] == "remote-mcp"
assert data["transport"] == "streamable-http"
assert data["url"] == "http://mcp.example.com/sse"
def test_create_invalid_name_spaces(self, client):
r = client.post(
"/v1/api/admin/mcp-servers",
json={"name": "bad name!", "transport": "stdio", "command": "x"},
)
assert r.status_code == 400
assert "name" in r.json()["error"].lower()
def test_create_invalid_name_double_underscore(self, client):
r = client.post(
"/v1/api/admin/mcp-servers",
json={"name": "bad__name", "transport": "stdio", "command": "x"},
)
assert r.status_code == 400
assert "__" in r.json()["error"]
def test_create_invalid_transport(self, client):
r = client.post(
"/v1/api/admin/mcp-servers",
json={"name": "ok-name", "transport": "grpc"},
)
assert r.status_code == 400
assert "transport" in r.json()["error"].lower()
def test_create_duplicate_name(self, client):
_create_server(client, name="dup-test")
r = client.post(
"/v1/api/admin/mcp-servers",
json={"name": "dup-test", "transport": "stdio", "command": "x"},
)
assert r.status_code == 409
assert "already exists" in r.json()["error"]
def test_create_missing_name(self, client):
r = client.post(
"/v1/api/admin/mcp-servers",
json={"transport": "stdio", "command": "x"},
)
assert r.status_code == 400
assert "name" in r.json()["error"].lower()
# ---------------------------------------------------------------------------
# Get single
# ---------------------------------------------------------------------------
class TestGetMcpServer:
def test_get_existing(self, client):
created = _create_server(client, name="get-test")
sid = created["server_id"]
with _PATCH_MCP_STATUS:
r = client.get(f"/v1/api/admin/mcp-servers/{sid}")
assert r.status_code == 200
assert r.json()["name"] == "get-test"
def test_get_not_found(self, client):
fake_id = uuid.uuid4().hex
with _PATCH_MCP_STATUS:
r = client.get(f"/v1/api/admin/mcp-servers/{fake_id}")
assert r.status_code == 404
assert "not found" in r.json()["error"].lower()
# ---------------------------------------------------------------------------
# Update
# ---------------------------------------------------------------------------
class TestUpdateMcpServer:
def test_update_name(self, client):
created = _create_server(client, name="old-name")
sid = created["server_id"]
r = client.put(
f"/v1/api/admin/mcp-servers/{sid}",
json={"name": "new-name"},
)
assert r.status_code == 200
assert r.json()["name"] == "new-name"
def test_update_transport(self, client):
created = _create_server(
client,
name="update-transport",
transport="streamable-http",
url="http://localhost/mcp",
)
sid = created["server_id"]
r = client.put(
f"/v1/api/admin/mcp-servers/{sid}",
json={"transport": "stdio", "command": "node"},
)
assert r.status_code == 200
assert r.json()["transport"] == "stdio"
def test_update_enabled(self, client):
created = _create_server(client, name="toggle-enabled")
sid = created["server_id"]
r = client.put(
f"/v1/api/admin/mcp-servers/{sid}",
json={"enabled": False},
)
assert r.status_code == 200
assert r.json()["enabled"] is False
def test_update_not_found(self, client):
fake_id = uuid.uuid4().hex
r = client.put(
f"/v1/api/admin/mcp-servers/{fake_id}",
json={"name": "x"},
)
assert r.status_code == 404
def test_update_invalid_transport(self, client):
created = _create_server(client, name="bad-transport-update")
sid = created["server_id"]
r = client.put(
f"/v1/api/admin/mcp-servers/{sid}",
json={"transport": "websocket"},
)
assert r.status_code == 400
assert "transport" in r.json()["error"].lower()
# ---------------------------------------------------------------------------
# Delete
# ---------------------------------------------------------------------------
class TestDeleteMcpServer:
def test_delete_existing(self, client):
created = _create_server(client, name="del-test")
sid = created["server_id"]
r = client.delete(f"/v1/api/admin/mcp-servers/{sid}")
assert r.status_code == 200
assert r.json()["status"] == "ok"
# Confirm it's gone
with _PATCH_MCP_STATUS:
r2 = client.get(f"/v1/api/admin/mcp-servers/{sid}")
assert r2.status_code == 404
def test_delete_not_found(self, client):
fake_id = uuid.uuid4().hex
r = client.delete(f"/v1/api/admin/mcp-servers/{fake_id}")
assert r.status_code == 404
assert "not found" in r.json()["error"].lower()
# ---------------------------------------------------------------------------
# Secret masking
# ---------------------------------------------------------------------------
class TestSecretMasking:
def test_list_masks_secrets(self, client):
_create_server(
client,
name="secret-test",
env={"API_KEY": "sk-real-secret-123"},
headers={"Authorization": "Bearer tok-xyz"},
transport="streamable-http",
url="http://localhost/mcp",
)
with _PATCH_MCP_STATUS:
r = client.get("/v1/api/admin/mcp-servers")
assert r.status_code == 200
server = r.json()["servers"][0]
env = json.loads(server["env"])
headers = json.loads(server["headers"])
assert env["API_KEY"] == "***"
assert headers["Authorization"] == "***"
def test_list_reveals_secrets(self, client):
_create_server(
client,
name="reveal-test",
env={"API_KEY": "sk-real-secret-123"},
headers={"Authorization": "Bearer tok-xyz"},
transport="streamable-http",
url="http://localhost/mcp",
)
with _PATCH_MCP_STATUS:
r = client.get("/v1/api/admin/mcp-servers?reveal=true")
assert r.status_code == 200
server = r.json()["servers"][0]
env = json.loads(server["env"])
headers = json.loads(server["headers"])
assert env["API_KEY"] == "sk-real-secret-123"
assert headers["Authorization"] == "Bearer tok-xyz"
def test_get_masks_secrets_by_default(self, client):
created = _create_server(
client,
name="mask-get-test",
env={"SECRET": "value"},
)
sid = created["server_id"]
with _PATCH_MCP_STATUS:
r = client.get(f"/v1/api/admin/mcp-servers/{sid}")
assert r.status_code == 200
env = json.loads(r.json()["env"])
assert env["SECRET"] == "***"
def test_get_reveals_secrets(self, client):
created = _create_server(
client,
name="reveal-get-test",
env={"SECRET": "real-value"},
)
sid = created["server_id"]
with _PATCH_MCP_STATUS:
r = client.get(f"/v1/api/admin/mcp-servers/{sid}?reveal=true")
assert r.status_code == 200
env = json.loads(r.json()["env"])
assert env["SECRET"] == "real-value"
# ---------------------------------------------------------------------------
# Import
# ---------------------------------------------------------------------------
class TestImportMcpConfig:
def test_import_inline_config(self, client):
config = {
"mcpServers": {
"filesystem": {
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-filesystem", "/tmp"],
},
"remote": {
"url": "http://remote.example.com/mcp",
},
},
}
r = client.post(
"/v1/api/admin/mcp-servers/import",
json={"config": config},
)
assert r.status_code == 200
data = r.json()
assert "filesystem" in data["imported"]
assert "remote" in data["imported"]
assert data["skipped"] == []
assert data["errors"] == []
def test_import_not_a_dict(self, client):
r = client.post(
"/v1/api/admin/mcp-servers/import",
json={"config": "not-a-dict"},
)
assert r.status_code == 400
def test_import_skips_duplicates(self, client):
_create_server(client, name="existing-srv")
config = {
"mcpServers": {
"existing-srv": {"command": "node", "args": []},
"new-srv": {"command": "node", "args": []},
},
}
r = client.post(
"/v1/api/admin/mcp-servers/import",
json={"config": config},
)
assert r.status_code == 200
data = r.json()
assert "new-srv" in data["imported"]
assert "existing-srv" in data["skipped"]
def test_import_empty_body(self, client):
r = client.post(
"/v1/api/admin/mcp-servers/import",
json={},
)
assert r.status_code == 400
assert "config" in r.json()["error"].lower()
def test_import_no_mcp_servers_key(self, client):
r = client.post(
"/v1/api/admin/mcp-servers/import",
json={"config": {"other": "data"}},
)
assert r.status_code == 400
assert "mcpServers" in r.json()["error"] or "No" in r.json()["error"]
# ---------------------------------------------------------------------------
# Permission check
# ---------------------------------------------------------------------------
class TestPermission:
def test_list_without_permission(self, client_no_perm):
with _PATCH_MCP_STATUS:
r = client_no_perm.get("/v1/api/admin/mcp-servers")
assert r.status_code == 403
assert "admin.mcp" in r.json()["error"]
def test_create_without_permission(self, client_no_perm):
r = client_no_perm.post(
"/v1/api/admin/mcp-servers",
json={"name": "test", "transport": "stdio", "command": "x"},
)
assert r.status_code == 403
def test_delete_without_permission(self, client_no_perm):
r = client_no_perm.delete(f"/v1/api/admin/mcp-servers/{uuid.uuid4().hex}")
assert r.status_code == 403
# ---------------------------------------------------------------------------
# Unit tests for _collect_mcp_status / _notify_nodes_mcp_reload
# ---------------------------------------------------------------------------
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))
req = MagicMock()
req.app.state.collector = collector
req.app.state.proxy_client = proxy_client or AsyncMock()
req.app.state.proxy_token_mgr = None
req.app.state.proxy_auth_token = "tok"
return req
def _mock_resp(status_code: int = 200, json_data: Any = None) -> MagicMock:
"""Build a mock httpx response (sync .json(), like the real thing)."""
resp = MagicMock()
resp.status_code = status_code
resp.json.return_value = json_data or {}
return resp
class TestCollectMcpStatus:
@pytest.mark.anyio
async def test_returns_servers_on_200(self):
resp = _mock_resp(200, {"servers": {"s1": {"status": "ok"}}})
client = AsyncMock()
client.get.return_value = resp
req = _fake_request(
{"node_id": "n1", "server_url": "http://n1:8000"},
proxy_client=client,
)
result = await _collect_mcp_status(req)
assert result == {"n1": {"s1": {"status": "ok"}}}
@pytest.mark.anyio
async def test_skips_non_200(self):
client = AsyncMock()
client.get.return_value = _mock_resp(503)
req = _fake_request(
{"node_id": "n1", "server_url": "http://n1:8000"},
proxy_client=client,
)
result = await _collect_mcp_status(req)
assert result == {}
@pytest.mark.anyio
async def test_skips_nodes_without_url(self):
client = AsyncMock()
req = _fake_request(
{"node_id": "n1", "server_url": ""},
{"node_id": "n2"},
proxy_client=client,
)
result = await _collect_mcp_status(req)
assert result == {}
client.get.assert_not_called()
@pytest.mark.anyio
async def test_handles_exception(self):
client = AsyncMock()
client.get.side_effect = ConnectionError("refused")
req = _fake_request(
{"node_id": "n1", "server_url": "http://n1:8000"},
proxy_client=client,
)
result = await _collect_mcp_status(req)
assert result == {}
@pytest.mark.anyio
async def test_empty_cluster(self):
req = _fake_request()
result = await _collect_mcp_status(req)
assert result == {}
@pytest.mark.anyio
async def test_multiple_nodes_mixed(self):
ok_resp = _mock_resp(200, {"servers": {"s1": {"status": "ok"}}})
err_resp = _mock_resp(500)
client = AsyncMock()
client.get.side_effect = [ok_resp, ConnectionError("down"), err_resp]
req = _fake_request(
{"node_id": "n1", "server_url": "http://n1:8000"},
{"node_id": "n2", "server_url": "http://n2:8000"},
{"node_id": "n3", "server_url": "http://n3:8000"},
proxy_client=client,
)
result = await _collect_mcp_status(req)
assert result == {"n1": {"s1": {"status": "ok"}}}
class TestNotifyNodesMcpReload:
@pytest.mark.anyio
async def test_returns_json_on_success(self):
client = AsyncMock()
client.post.return_value = _mock_resp(200, {"reloaded": 3})
req = _fake_request(
{"node_id": "n1", "server_url": "http://n1:8000"},
proxy_client=client,
)
result = await _notify_nodes_mcp_reload(req)
assert result == {"n1": {"reloaded": 3}}
@pytest.mark.anyio
async def test_skips_nodes_without_url(self):
client = AsyncMock()
req = _fake_request(
{"node_id": "n1", "server_url": ""},
proxy_client=client,
)
result = await _notify_nodes_mcp_reload(req)
assert result == {}
client.post.assert_not_called()
@pytest.mark.anyio
async def test_records_error_on_exception(self):
client = AsyncMock()
client.post.side_effect = ConnectionError("refused")
req = _fake_request(
{"node_id": "n1", "server_url": "http://n1:8000"},
proxy_client=client,
)
result = await _notify_nodes_mcp_reload(req)
assert "n1" in result
assert "error" in result["n1"]
assert "refused" in result["n1"]["error"]
@pytest.mark.anyio
async def test_empty_cluster(self):
req = _fake_request()
result = await _notify_nodes_mcp_reload(req)
assert result == {}
@pytest.mark.anyio
async def test_multiple_nodes_mixed(self):
client = AsyncMock()
client.post.side_effect = [
_mock_resp(200, {"reloaded": 2}),
TimeoutError("timeout"),
]
req = _fake_request(
{"node_id": "n1", "server_url": "http://n1:8000"},
{"node_id": "n2", "server_url": "http://n2:8000"},
proxy_client=client,
)
result = await _notify_nodes_mcp_reload(req)
assert result["n1"] == {"reloaded": 2}
assert "error" in result["n2"]
+403
View File
@@ -0,0 +1,403 @@
"""Tests for MCPClientManager hot-reload methods."""
from __future__ import annotations
from typing import Any
from turnstone.core.mcp_client import MCPClientManager
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _fake_openai_tool(name: str = "mcp__test__search") -> dict[str, Any]:
"""Create a fake OpenAI-format tool dict."""
return {
"type": "function",
"function": {
"name": name,
"description": "[MCP: test] Search stuff",
"parameters": {
"type": "object",
"properties": {"query": {"type": "string"}},
"required": ["query"],
},
},
}
def _fake_resource_dict(
uri: str = "file:///README.md",
name: str = "readme",
server: str = "test",
) -> dict[str, Any]:
"""Create a fake resource dict as stored in per-server state."""
return {
"uri": uri,
"name": name,
"description": "A resource",
"mimeType": "text/plain",
"server": server,
}
def _fake_prompt_dict(
name: str = "mcp__test__code_review",
original_name: str = "code_review",
server: str = "test",
) -> dict[str, Any]:
"""Create a fake prompt dict as stored in per-server state."""
return {
"name": name,
"original_name": original_name,
"server": server,
"description": "Generate a code review",
"arguments": [
{"name": "language", "description": "Programming language", "required": True}
],
}
# ---------------------------------------------------------------------------
# add_server_sync
# ---------------------------------------------------------------------------
class TestAddServerSync:
def test_rejects_double_underscore_name(self) -> None:
"""Names containing __ should be rejected."""
mgr = MCPClientManager({})
result = mgr.add_server_sync("bad__name", {"command": "echo"})
assert result["connected"] is False
assert "__" in result["error"]
assert result["tools"] == 0
assert result["resources"] == 0
assert result["prompts"] == 0
def test_fails_without_event_loop(self) -> None:
"""Adding a server without starting the event loop should fail gracefully."""
mgr = MCPClientManager({})
result = mgr.add_server_sync("test", {"command": "echo"})
assert result["connected"] is False
assert "loop" in result["error"].lower()
def test_config_removed_on_failure(self) -> None:
"""add_server_sync removes the config entry when connection fails."""
mgr = MCPClientManager({})
mgr.add_server_sync("new-srv", {"command": "echo"})
# Since the loop isn't running, it fails and config is cleaned up
assert "new-srv" not in mgr._server_configs
# ---------------------------------------------------------------------------
# remove_server_sync
# ---------------------------------------------------------------------------
class TestRemoveServerSync:
def test_returns_false_for_nonexistent(self) -> None:
"""Removing a non-connected server returns False."""
mgr = MCPClientManager({})
assert mgr.remove_server_sync("nonexistent") is False
def test_cleans_up_per_server_state(self) -> None:
"""remove_server_sync cleans up all per-server state dicts."""
mgr = MCPClientManager({"test": {"command": "echo"}})
# Simulate state as if the server was connected
mgr._per_server_tools["test"] = [_fake_openai_tool()]
mgr._per_server_resources["test"] = [_fake_resource_dict()]
mgr._per_server_prompts["test"] = [_fake_prompt_dict()]
mgr._supports_list_changed["test"] = True
mgr._supports_resources["test"] = True
mgr._supports_resource_list_changed["test"] = True
mgr._supports_prompts["test"] = True
mgr._supports_prompt_list_changed["test"] = True
mgr._rebuild_tools()
mgr._rebuild_resources()
mgr._rebuild_prompts()
# Verify preconditions
assert len(mgr.get_tools()) == 1
assert mgr.resource_count == 1
assert mgr.prompt_count == 1
mgr.remove_server_sync("test")
assert len(mgr.get_tools()) == 0
assert mgr.resource_count == 0
assert mgr.prompt_count == 0
assert "test" not in mgr._per_server_tools
assert "test" not in mgr._per_server_resources
assert "test" not in mgr._per_server_prompts
assert "test" not in mgr._supports_list_changed
assert "test" not in mgr._supports_resources
assert "test" not in mgr._supports_resource_list_changed
assert "test" not in mgr._supports_prompts
assert "test" not in mgr._supports_prompt_list_changed
def test_removes_config_to_prevent_reconnect(self) -> None:
"""remove_server_sync removes from _server_configs to prevent reconnect."""
mgr = MCPClientManager({"test": {"command": "echo"}})
assert "test" in mgr._server_configs
mgr.remove_server_sync("test")
assert "test" not in mgr._server_configs
def test_preserves_other_servers(self) -> None:
"""Removing one server does not affect another server's state."""
mgr = MCPClientManager({"srv_a": {}, "srv_b": {}})
mgr._per_server_tools["srv_a"] = [_fake_openai_tool("mcp__srv_a__foo")]
mgr._per_server_tools["srv_b"] = [_fake_openai_tool("mcp__srv_b__bar")]
mgr._rebuild_tools()
assert len(mgr.get_tools()) == 2
mgr.remove_server_sync("srv_a")
assert len(mgr.get_tools()) == 1
assert mgr.get_tools()[0]["function"]["name"] == "mcp__srv_b__bar"
assert "srv_b" in mgr._server_configs
# ---------------------------------------------------------------------------
# get_server_status
# ---------------------------------------------------------------------------
class TestGetServerStatus:
def test_disconnected_server_in_config(self) -> None:
"""Status of a configured but not connected server shows disconnected."""
mgr = MCPClientManager({"test": {"command": "echo"}})
status = mgr.get_server_status("test")
assert status["connected"] is False
assert status["tools"] == 0
assert status["resources"] == 0
assert status["prompts"] == 0
assert status["error"] == ""
def test_connected_server_with_tools(self) -> None:
"""Status of a connected server reports correct tool/resource/prompt counts."""
mgr = MCPClientManager({"test": {}})
# Simulate connected state
mgr._sessions["test"] = object() # any truthy value
mgr._per_server_tools["test"] = [
_fake_openai_tool("mcp__test__a"),
_fake_openai_tool("mcp__test__b"),
]
mgr._per_server_resources["test"] = [_fake_resource_dict()]
mgr._per_server_prompts["test"] = [_fake_prompt_dict()]
status = mgr.get_server_status("test")
assert status["connected"] is True
assert status["tools"] == 2
assert status["resources"] == 1
assert status["prompts"] == 1
def test_unknown_server(self) -> None:
"""Status of a server not in config or sessions shows disconnected."""
mgr = MCPClientManager({})
status = mgr.get_server_status("unknown")
assert status["connected"] is False
assert status["tools"] == 0
# ---------------------------------------------------------------------------
# get_all_server_status
# ---------------------------------------------------------------------------
class TestGetAllServerStatus:
def test_empty_manager(self) -> None:
"""Empty manager returns empty status dict."""
mgr = MCPClientManager({})
assert mgr.get_all_server_status() == {}
def test_multiple_servers(self) -> None:
"""Manager with configs but no connections returns status for each."""
mgr = MCPClientManager({"alpha": {}, "bravo": {}})
statuses = mgr.get_all_server_status()
assert len(statuses) == 2
assert "alpha" in statuses
assert "bravo" in statuses
assert statuses["alpha"]["connected"] is False
assert statuses["bravo"]["connected"] is False
def test_mixed_connected_and_disconnected(self) -> None:
"""Status correctly reflects a mix of connected and disconnected servers."""
mgr = MCPClientManager({"up": {}, "down": {}})
mgr._sessions["up"] = object()
mgr._per_server_tools["up"] = [_fake_openai_tool("mcp__up__x")]
statuses = mgr.get_all_server_status()
assert statuses["up"]["connected"] is True
assert statuses["up"]["tools"] == 1
assert statuses["down"]["connected"] is False
assert statuses["down"]["tools"] == 0
# ---------------------------------------------------------------------------
# Error tracking (_last_error)
# ---------------------------------------------------------------------------
class TestErrorTracking:
def test_get_server_status_returns_error(self) -> None:
"""Error stored in _last_error flows through get_server_status."""
mgr = MCPClientManager({"test": {"command": "echo"}})
mgr._last_error["test"] = "Connection refused"
status = mgr.get_server_status("test")
assert status["error"] == "Connection refused"
assert status["connected"] is False
def test_no_error_by_default(self) -> None:
"""Default error is empty string."""
mgr = MCPClientManager({"test": {"command": "echo"}})
status = mgr.get_server_status("test")
assert status["error"] == ""
def test_error_cleared_after_pop(self) -> None:
"""Clearing _last_error makes get_server_status return empty."""
mgr = MCPClientManager({"test": {"command": "echo"}})
mgr._last_error["test"] = "Connection refused"
mgr._last_error.pop("test", None)
status = mgr.get_server_status("test")
assert status["error"] == ""
def test_error_cleared_on_remove(self) -> None:
"""remove_server_sync cleans up _last_error entry."""
mgr = MCPClientManager({"test": {"command": "echo"}})
mgr._last_error["test"] = "Connection refused"
mgr.remove_server_sync("test")
assert "test" not in mgr._last_error
def test_all_server_status_includes_errors(self) -> None:
"""get_all_server_status propagates per-server errors."""
mgr = MCPClientManager({"alpha": {}, "bravo": {}})
mgr._last_error["alpha"] = "Timeout"
statuses = mgr.get_all_server_status()
assert statuses["alpha"]["error"] == "Timeout"
assert statuses["bravo"]["error"] == ""
def test_error_does_not_leak_across_servers(self) -> None:
"""Error on one server does not affect another."""
mgr = MCPClientManager({"a": {}, "b": {}})
mgr._last_error["a"] = "Failed"
assert mgr.get_server_status("b")["error"] == ""
# ---------------------------------------------------------------------------
# reconcile_sync
# ---------------------------------------------------------------------------
class _FakeStorage:
"""Minimal mock storage for reconcile tests."""
def __init__(self, rows: list[dict[str, Any]]) -> None:
self._rows = rows
def list_mcp_servers(self, enabled_only: bool = False) -> list[dict[str, Any]]:
if enabled_only:
return [r for r in self._rows if r.get("enabled", True)]
return list(self._rows)
def _db_row(
name: str,
transport: str = "stdio",
command: str = "echo",
args: str = "[]",
url: str = "",
headers: str = "{}",
env: str = "{}",
enabled: bool = True,
) -> dict[str, Any]:
return {
"name": name,
"transport": transport,
"command": command,
"args": args,
"url": url,
"headers": headers,
"env": env,
"enabled": enabled,
}
class TestReconcileSync:
def test_adds_new_servers(self) -> None:
mgr = MCPClientManager({})
storage = _FakeStorage([_db_row("new-srv")])
# Can't actually connect (no loop), but config should be attempted
result = mgr.reconcile_sync(storage)
# add_server_sync fails without a loop, but the method shouldn't crash
assert "new-srv" not in result["added"] # fails gracefully
assert result["removed"] == []
assert result["updated"] == []
def test_removes_stale_db_servers(self) -> None:
mgr = MCPClientManager({"old-srv": {"command": "echo"}})
mgr._db_managed.add("old-srv") # mark as DB-managed
storage = _FakeStorage([]) # DB is empty
result = mgr.reconcile_sync(storage)
assert "old-srv" in result["removed"]
assert "old-srv" not in mgr._server_configs
def test_preserves_config_file_servers(self) -> None:
"""Config-file servers (not in _db_managed) survive reconcile."""
mgr = MCPClientManager({"env-srv": {"command": "echo"}})
# NOT in _db_managed — loaded from MCP_CONFIG env
storage = _FakeStorage([]) # DB is empty
result = mgr.reconcile_sync(storage)
assert result["removed"] == []
assert "env-srv" in mgr._server_configs # still there
def test_config_server_not_overwritten_by_db_name_collision(self) -> None:
"""DB server with same name as config-file server does not replace it."""
original_cfg = {"type": "stdio", "command": "config-echo", "args": [], "env": {}}
mgr = MCPClientManager({"shared-name": dict(original_cfg)})
# NOT in _db_managed — this is a config-file server
# DB has a server with the same name but different config
storage = _FakeStorage([_db_row("shared-name", command="db-echo")])
result = mgr.reconcile_sync(storage)
# Config-file server should NOT be updated
assert result["updated"] == []
assert "shared-name" in mgr._server_configs
assert mgr._server_configs["shared-name"]["command"] == "config-echo"
def test_updates_changed_config(self) -> None:
original_cfg = {"type": "stdio", "command": "echo", "args": [], "env": {}}
mgr = MCPClientManager({"srv": dict(original_cfg)})
mgr._db_managed.add("srv") # mark as DB-managed
# DB has updated command — config differs
storage = _FakeStorage([_db_row("srv", command="cat")])
result = mgr.reconcile_sync(storage)
# remove_server_sync ran (old config cleared), add_server_sync attempted
# but fails without a running event loop — that's expected in unit tests.
# The key assertion: the old config was evicted (not left stale).
assert "srv" not in mgr._server_configs
# Not in "removed" (that's for servers absent from DB)
assert "srv" not in result["removed"]
def test_no_change_is_noop(self) -> None:
cfg = {"type": "stdio", "command": "echo", "args": [], "env": {}}
mgr = MCPClientManager({"srv": dict(cfg)})
storage = _FakeStorage([_db_row("srv", command="echo")])
result = mgr.reconcile_sync(storage)
assert result["added"] == []
assert result["removed"] == []
assert result["updated"] == []
# Config unchanged
assert "srv" in mgr._server_configs
def test_storage_failure_graceful(self) -> None:
mgr = MCPClientManager({"srv": {}})
class _BrokenStorage:
def list_mcp_servers(self, **kw: Any) -> list[dict[str, Any]]:
raise RuntimeError("DB down")
result = mgr.reconcile_sync(_BrokenStorage())
assert result == {"added": [], "removed": [], "updated": []}
# Existing server untouched
assert "srv" in mgr._server_configs
+152
View File
@@ -0,0 +1,152 @@
"""Tests for MCP server storage CRUD operations."""
from __future__ import annotations
import uuid
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_id() -> str:
return uuid.uuid4().hex
class TestMcpServerStorage:
def test_create_and_get(self, db: SQLiteBackend) -> None:
sid = _make_id()
db.create_mcp_server(
server_id=sid,
name="test-server",
transport="stdio",
command="echo",
args='["hello"]',
)
s = db.get_mcp_server(sid)
assert s is not None
assert s["name"] == "test-server"
assert s["transport"] == "stdio"
assert s["command"] == "echo"
assert s["args"] == '["hello"]'
assert s["enabled"] is True
assert s["auto_approve"] is False
def test_get_by_name(self, db: SQLiteBackend) -> None:
sid = _make_id()
db.create_mcp_server(server_id=sid, name="named-srv", transport="stdio")
s = db.get_mcp_server_by_name("named-srv")
assert s is not None
assert s["server_id"] == sid
def test_get_by_name_not_found(self, db: SQLiteBackend) -> None:
assert db.get_mcp_server_by_name("nope") is None
def test_get_not_found(self, db: SQLiteBackend) -> None:
assert db.get_mcp_server("nonexistent") is None
def test_list_empty(self, db: SQLiteBackend) -> None:
assert db.list_mcp_servers() == []
def test_list_all(self, db: SQLiteBackend) -> None:
db.create_mcp_server(server_id=_make_id(), name="alpha", transport="stdio")
db.create_mcp_server(
server_id=_make_id(), name="beta", transport="streamable-http", url="http://x"
)
servers = db.list_mcp_servers()
assert len(servers) == 2
assert servers[0]["name"] == "alpha" # ordered by name
assert servers[1]["name"] == "beta"
def test_list_enabled_only(self, db: SQLiteBackend) -> None:
sid1 = _make_id()
sid2 = _make_id()
db.create_mcp_server(server_id=sid1, name="enabled-srv", transport="stdio", enabled=True)
db.create_mcp_server(server_id=sid2, name="disabled-srv", transport="stdio", enabled=False)
enabled = db.list_mcp_servers(enabled_only=True)
assert len(enabled) == 1
assert enabled[0]["name"] == "enabled-srv"
def test_update_basic_fields(self, db: SQLiteBackend) -> None:
sid = _make_id()
db.create_mcp_server(server_id=sid, name="orig", transport="stdio", command="echo")
ok = db.update_mcp_server(sid, name="renamed", command="cat")
assert ok is True
s = db.get_mcp_server(sid)
assert s is not None
assert s["name"] == "renamed"
assert s["command"] == "cat"
def test_update_boolean_conversion(self, db: SQLiteBackend) -> None:
sid = _make_id()
db.create_mcp_server(server_id=sid, name="booltest", transport="stdio")
db.update_mcp_server(sid, auto_approve=True, enabled=False)
s = db.get_mcp_server(sid)
assert s is not None
assert s["auto_approve"] is True
assert s["enabled"] is False
def test_update_not_found(self, db: SQLiteBackend) -> None:
ok = db.update_mcp_server("nonexistent", name="x")
assert ok is False
def test_update_ignores_disallowed_fields(self, db: SQLiteBackend) -> None:
sid = _make_id()
db.create_mcp_server(server_id=sid, name="guard", transport="stdio", created_by="admin")
original = db.get_mcp_server(sid)
assert original is not None
original_created = original["created"]
# created_by and created are not in the mutable allowlist
db.update_mcp_server(sid, created_by="evil", created="2000-01-01T00:00:00")
s = db.get_mcp_server(sid)
assert s is not None
assert s["created_by"] == "admin" # unchanged
assert s["created"] == original_created # unchanged
def test_delete(self, db: SQLiteBackend) -> None:
sid = _make_id()
db.create_mcp_server(server_id=sid, name="delme", transport="stdio")
ok = db.delete_mcp_server(sid)
assert ok is True
assert db.get_mcp_server(sid) is None
def test_delete_not_found(self, db: SQLiteBackend) -> None:
ok = db.delete_mcp_server("nonexistent")
assert ok is False
def test_create_duplicate_name(self, db: SQLiteBackend) -> None:
db.create_mcp_server(server_id=_make_id(), name="unique", transport="stdio")
# Second create with same name but different ID should be no-op (OR IGNORE)
sid2 = _make_id()
db.create_mcp_server(server_id=sid2, name="unique", transport="stdio")
# OR IGNORE silently drops the conflicting insert
assert db.get_mcp_server(sid2) is None
def test_create_idempotent_same_id(self, db: SQLiteBackend) -> None:
sid = _make_id()
db.create_mcp_server(server_id=sid, name="idem", transport="stdio", command="v1")
db.create_mcp_server(server_id=sid, name="idem", transport="stdio", command="v2")
s = db.get_mcp_server(sid)
assert s is not None
assert s["command"] == "v1" # original preserved, second ignored
def test_http_transport_fields(self, db: SQLiteBackend) -> None:
sid = _make_id()
db.create_mcp_server(
server_id=sid,
name="http-srv",
transport="streamable-http",
url="https://example.com/mcp",
headers='{"Authorization":"Bearer xyz"}',
)
s = db.get_mcp_server(sid)
assert s is not None
assert s["transport"] == "streamable-http"
assert s["url"] == "https://example.com/mcp"
assert "Authorization" in s["headers"]
+87
View File
@@ -292,6 +292,69 @@ class TestServerUserScopeSecurity:
assert r.status_code == 403
class TestServerScopeScopeIdValidation:
"""scope_id requires scope; global scope rejects scope_id."""
def test_save_global_with_scope_id_rejected(self, server_client):
r = server_client.post(
"/v1/api/memories",
json={"name": "k", "content": "c", "scope": "global", "scope_id": "ws1"},
)
assert r.status_code == 400
assert "scope_id" in r.json()["error"]
def test_save_workstream_without_scope_id_rejected(self, server_client):
r = server_client.post(
"/v1/api/memories",
json={"name": "k", "content": "c", "scope": "workstream"},
)
assert r.status_code == 400
assert "scope_id is required" in r.json()["error"]
def test_save_workstream_with_scope_id_ok(self, server_client):
r = server_client.post(
"/v1/api/memories",
json={"name": "k", "content": "c", "scope": "workstream", "scope_id": "ws1"},
)
assert r.status_code == 201
def test_list_scope_id_without_scope_rejected(self, server_client):
r = server_client.get("/v1/api/memories?scope_id=ws1")
assert r.status_code == 400
assert "scope is required" in r.json()["error"]
def test_list_global_with_scope_id_rejected(self, server_client):
r = server_client.get("/v1/api/memories?scope=global&scope_id=ws1")
assert r.status_code == 400
assert "scope_id" in r.json()["error"]
def test_search_scope_id_without_scope_rejected(self, server_client):
r = server_client.post(
"/v1/api/memories/search",
json={"query": "test", "scope_id": "ws1"},
)
assert r.status_code == 400
assert "scope is required" in r.json()["error"]
def test_search_global_with_scope_id_rejected(self, server_client):
r = server_client.post(
"/v1/api/memories/search",
json={"query": "test", "scope": "global", "scope_id": "ws1"},
)
assert r.status_code == 400
assert "scope_id" in r.json()["error"]
def test_delete_global_with_scope_id_rejected(self, server_client):
r = server_client.delete("/v1/api/memories/k?scope=global&scope_id=ws1")
assert r.status_code == 400
assert "scope_id" in r.json()["error"]
def test_delete_workstream_without_scope_id_rejected(self, server_client):
r = server_client.delete("/v1/api/memories/k?scope=workstream")
assert r.status_code == 400
assert "scope_id is required" in r.json()["error"]
class TestServerSearchMemories:
def test_search(self, server_client, storage):
_seed_memory(storage, "db_config", "postgresql host", description="database")
@@ -367,6 +430,30 @@ class TestAdminListMemories:
assert r.json()["total"] == 1
class TestAdminScopeScopeIdValidation:
"""Console admin: scope_id requires scope; global scope rejects scope_id."""
def test_list_scope_id_without_scope_rejected(self, admin_client):
r = admin_client.get("/v1/api/admin/memories?scope_id=ws1")
assert r.status_code == 400
assert "scope is required" in r.json()["error"]
def test_list_global_with_scope_id_rejected(self, admin_client):
r = admin_client.get("/v1/api/admin/memories?scope=global&scope_id=ws1")
assert r.status_code == 400
assert "scope_id" in r.json()["error"]
def test_search_scope_id_without_scope_rejected(self, admin_client):
r = admin_client.get("/v1/api/admin/memories/search?q=test&scope_id=ws1")
assert r.status_code == 400
assert "scope is required" in r.json()["error"]
def test_search_global_with_scope_id_rejected(self, admin_client):
r = admin_client.get("/v1/api/admin/memories/search?q=test&scope=global&scope_id=ws1")
assert r.status_code == 400
assert "scope_id" in r.json()["error"]
class TestAdminSearchMemories:
def test_search(self, admin_client, storage):
_seed_memory(storage, "db_config", "pg host", description="database")
+116 -9
View File
@@ -14,15 +14,16 @@ from turnstone.core.metacognition import (
class TestDetectCorrection:
"""Strong patterns always fire; weak 'no <word>' uses allowlist."""
# -- strong patterns (always fire) --
def test_no_comma(self):
assert detect_correction("no, that's wrong") is True
def test_no_period(self):
assert detect_correction("no. do it differently") is True
def test_no_space(self):
assert detect_correction("no I meant the other one") is True
def test_dont(self):
assert detect_correction("don't use tabs") is True
@@ -47,6 +48,57 @@ class TestDetectCorrection:
def test_please_dont(self):
assert detect_correction("please don't mock the database") is True
# -- weak pattern: "no" + allowlisted context word --
def test_no_space(self):
assert detect_correction("no I meant the other one") is True
def test_no_that(self):
assert detect_correction("no that's wrong") is True
def test_no_it(self):
assert detect_correction("no it should be different") is True
def test_no_the(self):
assert detect_correction("no the other one") is True
def test_no_not(self):
assert detect_correction("no not that file") is True
def test_no_you(self):
assert detect_correction("no you should use pytest") is True
# -- negatives: "no <word>" not in allowlist --
def test_negative_no_problem(self):
assert detect_correction("no problem") is False
def test_negative_no_worries(self):
assert detect_correction("no worries") is False
def test_negative_no_rush(self):
assert detect_correction("no rush") is False
def test_negative_no_one(self):
assert detect_correction("no one knows") is False
def test_negative_no_thanks(self):
assert detect_correction("no thanks") is False
def test_negative_no_doubt(self):
assert detect_correction("no doubt about it") is False
def test_negative_no_idea(self):
assert detect_correction("no idea what you mean") is False
def test_negative_no_kidding(self):
assert detect_correction("no kidding") is False
def test_negative_no_luck(self):
assert detect_correction("no luck finding the bug") is False
# -- negatives: unrelated messages --
def test_negative_notice(self):
assert detect_correction("I noticed the test passes") is False
@@ -70,27 +122,82 @@ class TestDetectCorrection:
class TestDetectCompletion:
def test_thanks(self):
assert detect_completion("thanks, that's perfect") is True
"""Strong patterns always fire; weak patterns gated by length + continuation."""
# -- strong patterns (always fire) --
def test_thats_all(self):
assert detect_completion("that's all for now") is True
def test_lgtm(self):
assert detect_completion("lgtm") is True
# -- weak patterns: short message, no continuation --
def test_thanks(self):
assert detect_completion("thanks, that's perfect") is True
def test_thanks_standalone(self):
assert detect_completion("thanks") is True
def test_thanks_exclaim(self):
assert detect_completion("thanks!") is True
def test_looks_good(self):
assert detect_completion("looks good to me") is True
def test_perfect(self):
assert detect_completion("perfect") is True
def test_lgtm(self):
assert detect_completion("lgtm") is True
def test_done(self):
assert detect_completion("done") is True
def test_negative_normal(self):
def test_great_job(self):
assert detect_completion("great job") is True
def test_that_works(self):
assert detect_completion("that works") is True
# -- negatives: "thanks for" is acknowledgment --
def test_negative_thanks_for(self):
assert detect_completion("thanks for the update") is False
def test_negative_thanks_for_looking(self):
assert detect_completion("thanks for looking into this") is False
# -- negatives: continuation markers suppress weak patterns --
def test_negative_thanks_but(self):
assert detect_completion("thanks but can you also add tests") is False
def test_negative_thanks_though(self):
assert detect_completion("thanks though I have one more question") is False
def test_negative_looks_good_but(self):
assert detect_completion("looks good but can you also add validation") is False
def test_negative_perfect_now(self):
assert detect_completion("perfect, now add error handling") is False
def test_negative_done_can_you(self):
assert detect_completion("done with that, can you start on the tests?") is False
def test_negative_question_mark(self):
assert detect_completion("can you add error handling?") is False
# -- negatives: long messages suppress weak patterns --
def test_negative_thanks_long(self):
msg = "thanks, this is really helpful — I was also wondering about the deployment pipeline and whether we need to update the CI config"
assert detect_completion(msg) is False
def test_negative_looks_good_long(self):
msg = "looks good overall, there are a few things I'd like to tweak though — the error messages could be more descriptive and the retry logic needs a backoff"
assert detect_completion(msg) is False
# -- negatives: unrelated --
def test_negative_empty(self):
assert detect_completion("") is False
+6 -2
View File
@@ -530,7 +530,9 @@ class TestWorkstreamModelParam:
captured_alias = None
def factory(ui: Any, model_alias: str | None = None, ws_id: str | None = None) -> Any:
def factory(
ui: Any, model_alias: str | None = None, ws_id: str | None = None, **kwargs: Any
) -> Any:
nonlocal captured_alias
captured_alias = model_alias
mock_session = MagicMock()
@@ -544,7 +546,9 @@ class TestWorkstreamModelParam:
def test_create_without_model(self) -> None:
captured_alias = None
def factory(ui: Any, model_alias: str | None = None, ws_id: str | None = None) -> Any:
def factory(
ui: Any, model_alias: str | None = None, ws_id: str | None = None, **kwargs: Any
) -> Any:
nonlocal captured_alias
captured_alias = model_alias
mock_session = MagicMock()
+967
View File
@@ -0,0 +1,967 @@
"""Tests for turnstone.core.oidc — OIDC authentication support."""
from __future__ import annotations
import asyncio
import base64
import hashlib
import urllib.parse
from unittest.mock import MagicMock, patch
import httpx
import jwt as pyjwt
import pytest
from turnstone.core.oidc import (
OIDCConfig,
OIDCError,
apply_role_mapping,
build_authorize_url,
discover_oidc,
generate_pkce_pair,
load_oidc_config,
provision_oidc_user,
validate_id_token,
)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_config(**overrides) -> OIDCConfig:
"""Build a test OIDCConfig with sensible defaults."""
defaults = {
"enabled": True,
"issuer": "https://idp.example.com",
"client_id": "my-client",
"client_secret": "my-secret",
"scopes": "openid email profile",
"provider_name": "TestIDP",
"role_claim": "",
"role_map": {},
"password_enabled": True,
"authorization_endpoint": "https://idp.example.com/authorize",
"token_endpoint": "https://idp.example.com/token",
"userinfo_endpoint": "https://idp.example.com/userinfo",
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
}
defaults.update(overrides)
return OIDCConfig(**defaults)
def _mock_storage(**overrides):
"""Build a MagicMock with sensible storage defaults."""
s = MagicMock()
s.get_oidc_identity.return_value = overrides.get("identity")
s.get_user.return_value = overrides.get("user")
s.get_user_by_username.return_value = overrides.get("user_by_username")
s.get_role.return_value = overrides.get("role")
return s
def _mock_async_client(mock_get):
"""Build a patched httpx.AsyncClient context manager for async tests."""
class _AsyncCtx:
async def __aenter__(self):
return self
async def __aexit__(self, *args):
pass
async def get(self, url):
return await mock_get(url)
return _AsyncCtx()
# ---------------------------------------------------------------------------
# Config Loading
# ---------------------------------------------------------------------------
class TestLoadOIDCConfig:
def test_load_oidc_config_from_env(self, monkeypatch):
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_SCOPES", "openid")
monkeypatch.setenv("TURNSTONE_OIDC_PROVIDER_NAME", "Okta")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.enabled is True
assert cfg.issuer == "https://auth.example.com"
assert cfg.client_id == "cid"
assert cfg.client_secret == "csecret"
assert cfg.scopes == "openid"
assert cfg.provider_name == "Okta"
def test_load_oidc_config_disabled_when_missing(self, monkeypatch):
monkeypatch.delenv("TURNSTONE_OIDC_ISSUER", raising=False)
monkeypatch.delenv("TURNSTONE_OIDC_CLIENT_ID", raising=False)
monkeypatch.delenv("TURNSTONE_OIDC_CLIENT_SECRET", raising=False)
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.enabled is False
def test_load_oidc_config_partial_env(self, monkeypatch):
"""Only issuer set, no client_id -> enabled=False."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.delenv("TURNSTONE_OIDC_CLIENT_ID", raising=False)
monkeypatch.delenv("TURNSTONE_OIDC_CLIENT_SECRET", raising=False)
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.enabled is False
assert cfg.issuer == "https://auth.example.com"
assert cfg.client_id == ""
def test_load_oidc_config_role_map_parsing(self, monkeypatch):
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_ROLE_CLAIM", "roles")
monkeypatch.setenv("TURNSTONE_OIDC_ROLE_MAP", "admin:builtin-admin,eng:builtin-operator")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.role_claim == "roles"
assert cfg.role_map == {"admin": "builtin-admin", "eng": "builtin-operator"}
def test_load_oidc_config_password_enabled_false(self, monkeypatch):
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_PASSWORD_ENABLED", "false")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.enabled is True
assert cfg.password_enabled is False
def test_load_oidc_config_password_enabled_true(self, monkeypatch):
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_PASSWORD_ENABLED", "true")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.password_enabled is True
def test_load_oidc_config_role_map_empty_entries(self, monkeypatch):
"""Role map with empty/whitespace entries should be silently skipped."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_ROLE_MAP", "admin:builtin-admin, , :, foo:")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.role_map == {"admin": "builtin-admin"}
def test_load_oidc_config_defaults(self, monkeypatch):
"""Defaults for scopes and provider_name when not set."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.delenv("TURNSTONE_OIDC_SCOPES", raising=False)
monkeypatch.delenv("TURNSTONE_OIDC_PROVIDER_NAME", raising=False)
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.scopes == "openid email profile"
assert cfg.provider_name == "SSO"
def test_load_oidc_config_redirect_base_from_env(self, monkeypatch):
"""TURNSTONE_OIDC_REDIRECT_BASE populates redirect_base."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://app.example.com")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.redirect_base == "https://app.example.com"
def test_load_oidc_config_redirect_base_strips_trailing_slash(self, monkeypatch):
"""Trailing slashes are stripped from redirect_base."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://app.example.com/")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.redirect_base == "https://app.example.com"
def test_load_oidc_config_redirect_base_default_empty(self, monkeypatch):
"""redirect_base defaults to empty string when not set."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.delenv("TURNSTONE_OIDC_REDIRECT_BASE", raising=False)
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.redirect_base == ""
def test_load_oidc_config_redirect_base_rejects_path(self, monkeypatch):
"""redirect_base with a path component is rejected (falls back to empty)."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://app.example.com/subpath")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.redirect_base == ""
def test_load_oidc_config_redirect_base_rejects_no_scheme(self, monkeypatch):
"""redirect_base without a scheme is rejected."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "app.example.com")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.redirect_base == ""
def test_load_oidc_config_redirect_base_rejects_userinfo(self, monkeypatch):
"""redirect_base with userinfo (user:pass@host) is rejected."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://user:pass@app.example.com")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.redirect_base == ""
def test_load_oidc_config_redirect_base_rejects_invalid_port(self, monkeypatch):
"""redirect_base with non-numeric port is rejected."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://app.example.com:abc")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.redirect_base == ""
def test_load_oidc_config_redirect_base_rejects_missing_hostname(self, monkeypatch):
"""redirect_base without a hostname is rejected."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.redirect_base == ""
def test_load_oidc_config_redirect_base_allows_http(self, monkeypatch):
"""http:// redirect_base is allowed (with warning) for local dev."""
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "http://localhost:8000")
with patch("turnstone.core.config.load_config", return_value={}):
cfg = load_oidc_config()
assert cfg.redirect_base == "http://localhost:8000"
# ---------------------------------------------------------------------------
# Redirect URI Builder
# ---------------------------------------------------------------------------
class TestBuildOIDCRedirectURI:
"""Tests for ``_build_oidc_redirect_uri`` in auth.py."""
def _make_request(self, host="app.example.com", scheme="https", forwarded_proto=""):
"""Build a minimal mock Starlette Request."""
req = MagicMock()
headers = {"host": host}
if forwarded_proto:
headers["x-forwarded-proto"] = forwarded_proto
req.headers = headers
req.url.scheme = scheme
return req
def test_pinned_redirect_base(self):
"""When redirect_base is set, Host header is ignored."""
from turnstone.core.auth import _build_oidc_redirect_uri
config = _make_config(redirect_base="https://public.example.com")
req = self._make_request(host="internal-host:8080", scheme="http")
result = _build_oidc_redirect_uri(req, config)
assert result == "https://public.example.com/v1/api/auth/oidc/callback"
def test_fallback_to_host_header(self):
"""When redirect_base is empty, redirect URI uses Host header."""
from turnstone.core.auth import _build_oidc_redirect_uri
config = _make_config(redirect_base="")
req = self._make_request(host="app.example.com", scheme="https")
result = _build_oidc_redirect_uri(req, config)
assert result == "https://app.example.com/v1/api/auth/oidc/callback"
def test_fallback_x_forwarded_proto(self):
"""When redirect_base is empty and X-Forwarded-Proto is https, scheme is https."""
from turnstone.core.auth import _build_oidc_redirect_uri
config = _make_config(redirect_base="")
req = self._make_request(host="app.example.com", scheme="http", forwarded_proto="https")
result = _build_oidc_redirect_uri(req, config)
assert result == "https://app.example.com/v1/api/auth/oidc/callback"
# ---------------------------------------------------------------------------
# PKCE
# ---------------------------------------------------------------------------
class TestPKCE:
def test_generate_pkce_pair(self):
verifier, challenge = generate_pkce_pair()
# Verifier should be URL-safe base64
assert isinstance(verifier, str)
assert len(verifier) > 40 # 48 bytes -> ~64 chars
# Challenge should be base64url SHA-256 of verifier
expected_digest = hashlib.sha256(verifier.encode("ascii")).digest()
expected_challenge = base64.urlsafe_b64encode(expected_digest).rstrip(b"=").decode("ascii")
assert challenge == expected_challenge
def test_pkce_challenge_matches_verifier(self):
"""Manually compute challenge and verify it matches."""
verifier, challenge = generate_pkce_pair()
digest = hashlib.sha256(verifier.encode("ascii")).digest()
manual_challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
assert challenge == manual_challenge
def test_pkce_pair_uniqueness(self):
"""Each call should produce a unique pair."""
v1, c1 = generate_pkce_pair()
v2, c2 = generate_pkce_pair()
assert v1 != v2
assert c1 != c2
# ---------------------------------------------------------------------------
# Authorization URL
# ---------------------------------------------------------------------------
class TestBuildAuthorizeURL:
def test_build_authorize_url_contains_required_params(self):
config = _make_config()
verifier, _ = generate_pkce_pair()
url = build_authorize_url(
config=config,
redirect_uri="https://app.example.com/callback",
state="test-state",
nonce="test-nonce",
code_verifier=verifier,
)
assert url.startswith("https://idp.example.com/authorize?")
assert "response_type=code" in url
assert "client_id=my-client" in url
assert "redirect_uri=" in url
assert "scope=openid" in url
assert "state=test-state" in url
assert "nonce=test-nonce" in url
assert "code_challenge=" in url
assert "code_challenge_method=S256" in url
def test_build_authorize_url_pkce(self):
"""code_challenge in URL should be correct S256 of the verifier."""
config = _make_config()
verifier, _ = generate_pkce_pair()
url = build_authorize_url(
config=config,
redirect_uri="https://app.example.com/callback",
state="s",
nonce="n",
code_verifier=verifier,
)
# Extract code_challenge from URL
parsed = urllib.parse.urlparse(url)
params = urllib.parse.parse_qs(parsed.query)
actual_challenge = params["code_challenge"][0]
# Compute expected challenge
digest = hashlib.sha256(verifier.encode("ascii")).digest()
expected = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
assert actual_challenge == expected
def test_build_authorize_url_redirect_uri_encoded(self):
config = _make_config()
verifier, _ = generate_pkce_pair()
redirect = "https://app.example.com/callback?extra=1"
url = build_authorize_url(
config=config,
redirect_uri=redirect,
state="s",
nonce="n",
code_verifier=verifier,
)
# The redirect_uri should be URL-encoded
parsed = urllib.parse.urlparse(url)
params = urllib.parse.parse_qs(parsed.query)
assert params["redirect_uri"][0] == redirect
# ---------------------------------------------------------------------------
# ID Token Validation
# ---------------------------------------------------------------------------
class TestValidateIDToken:
_FAKE_JWKS = {"keys": [{"kid": "key1", "kty": "RSA", "n": "abc", "e": "AQAB"}]}
def test_validate_id_token_nonce_mismatch(self):
"""Nonce mismatch should raise OIDCError."""
config = _make_config()
mock_pyjwk = MagicMock()
mock_pyjwk.return_value.key = "fake-key"
with (
patch("jwt.get_unverified_header", return_value={"kid": "key1", "alg": "RS256"}),
patch("jwt.PyJWK", mock_pyjwk),
patch("jwt.decode", return_value={"sub": "user1", "nonce": "wrong-nonce"}),
pytest.raises(OIDCError, match="nonce mismatch"),
):
validate_id_token(
raw_token="fake.jwt.token",
jwks_data=self._FAKE_JWKS,
config=config,
nonce="expected-nonce",
)
def test_validate_id_token_success(self):
"""Successful validation returns decoded claims."""
config = _make_config()
mock_pyjwk = MagicMock()
mock_pyjwk.return_value.key = "fake-key"
expected_claims = {
"sub": "user1",
"email": "user@example.com",
"nonce": "test-nonce",
}
with (
patch("jwt.get_unverified_header", return_value={"kid": "key1", "alg": "RS256"}),
patch("jwt.PyJWK", mock_pyjwk),
patch("jwt.decode", return_value=expected_claims) as mock_decode,
):
claims = validate_id_token(
raw_token="fake.jwt.token",
jwks_data=self._FAKE_JWKS,
config=config,
nonce="test-nonce",
)
assert claims == expected_claims
mock_decode.assert_called_once_with(
"fake.jwt.token",
"fake-key",
algorithms=[
"RS256",
"RS384",
"RS512",
"ES256",
"ES384",
"ES512",
"PS256",
"PS384",
"PS512",
],
audience="my-client",
issuer="https://idp.example.com",
)
def test_validate_id_token_kid_not_found(self):
"""Unknown kid raises OIDCError with descriptive message."""
config = _make_config()
jwks_data = {"keys": [{"kid": "other-key", "kty": "RSA"}]}
with (
patch("jwt.get_unverified_header", return_value={"kid": "unknown", "alg": "RS256"}),
pytest.raises(OIDCError, match="not found in JWKS"),
):
validate_id_token(
raw_token="bad.token",
jwks_data=jwks_data,
config=config,
nonce="n",
)
def test_validate_id_token_invalid_jwt(self):
"""Invalid JWT raises OIDCError."""
config = _make_config()
mock_pyjwk = MagicMock()
mock_pyjwk.return_value.key = "fake-key"
with (
patch("jwt.get_unverified_header", return_value={"kid": "key1", "alg": "RS256"}),
patch("jwt.PyJWK", mock_pyjwk := MagicMock(return_value=MagicMock(key="fake-key"))),
patch("jwt.decode", side_effect=pyjwt.InvalidTokenError("expired")),
pytest.raises(OIDCError, match="ID token validation failed"),
):
validate_id_token(
raw_token="expired.token",
jwks_data=self._FAKE_JWKS,
config=config,
nonce="n",
)
def test_validate_id_token_invalid_header(self):
"""Malformed token header raises OIDCError."""
config = _make_config()
with (
patch("jwt.get_unverified_header", side_effect=pyjwt.DecodeError("bad header")),
pytest.raises(OIDCError, match="Invalid ID token header"),
):
validate_id_token(
raw_token="garbage",
jwks_data=self._FAKE_JWKS,
config=config,
nonce="n",
)
# ---------------------------------------------------------------------------
# User Provisioning
# ---------------------------------------------------------------------------
class TestProvisionOIDCUser:
def test_provision_oidc_user_existing(self):
"""Existing identity -> returns existing user, updates last_login."""
config = _make_config()
existing_user = {
"user_id": "u1",
"username": "alice",
"display_name": "Alice",
"password_hash": "!oidc",
}
existing_identity = {
"issuer": "https://idp.example.com",
"subject": "sub-123",
"user_id": "u1",
"email": "alice@example.com",
"created": "2024-01-01T00:00:00",
"last_login": "2024-01-01T00:00:00",
}
storage = _mock_storage(identity=existing_identity, user=existing_user)
claims = {"sub": "sub-123", "email": "alice@example.com", "name": "Alice"}
user = provision_oidc_user(storage, config, claims)
assert user["user_id"] == "u1"
assert user["username"] == "alice"
storage.update_oidc_identity_login.assert_called_once()
# Should not create a new user
storage.create_user.assert_not_called()
storage.create_oidc_identity.assert_not_called()
def test_provision_oidc_user_new(self):
"""No identity -> creates user + identity."""
config = _make_config()
storage = _mock_storage()
# After create_user, get_user should return the new user
new_user = {
"user_id": "u-new",
"username": "bob",
"display_name": "Bob",
"password_hash": "!oidc",
}
storage.get_user.return_value = new_user
claims = {"sub": "sub-456", "preferred_username": "bob", "email": "bob@example.com"}
with patch("turnstone.core.oidc.uuid") as mock_uuid:
mock_uuid.uuid4.return_value = MagicMock(hex="u-new-hex-00000000000000000000")
user = provision_oidc_user(storage, config, claims)
assert user["username"] == "bob"
storage.create_user.assert_called_once()
storage.create_oidc_identity.assert_called_once()
# Verify create_oidc_identity was called with correct issuer and sub
call_args = storage.create_oidc_identity.call_args
assert call_args[0][0] == "https://idp.example.com" # issuer
assert call_args[0][1] == "sub-456" # subject
def test_provision_oidc_user_username_dedup(self):
"""First username taken -> appends suffix."""
config = _make_config()
storage = _mock_storage()
# First call: username "bob" exists; second call: "bob2" doesn't exist
storage.get_user_by_username.side_effect = [
{"user_id": "u-other", "username": "bob"}, # "bob" taken
None, # "bob2" available
]
new_user = {
"user_id": "u-new",
"username": "bob2",
"display_name": "Bob",
"password_hash": "!oidc",
}
storage.get_user.return_value = new_user
claims = {"sub": "sub-789", "preferred_username": "bob", "email": "bob@example.com"}
user = provision_oidc_user(storage, config, claims)
assert user["username"] == "bob2"
# create_user should have been called with "bob2" as username
call_args = storage.create_user.call_args
assert call_args[0][1] == "bob2"
def test_provision_oidc_user_email_prefix(self):
"""No preferred_username -> uses email prefix."""
config = _make_config()
storage = _mock_storage()
new_user = {
"user_id": "u-new",
"username": "charlie",
"display_name": "charlie@example.com",
"password_hash": "!oidc",
}
storage.get_user.return_value = new_user
claims = {"sub": "sub-abc", "email": "charlie@example.com"}
provision_oidc_user(storage, config, claims)
# create_user should have been called with "charlie" (email prefix)
call_args = storage.create_user.call_args
assert call_args[0][1] == "charlie"
def test_provision_oidc_user_missing_user_raises(self):
"""Identity references missing user -> raises OIDCError."""
config = _make_config()
existing_identity = {
"issuer": "https://idp.example.com",
"subject": "sub-orphan",
"user_id": "u-gone",
"email": "gone@example.com",
"created": "2024-01-01T00:00:00",
"last_login": "2024-01-01T00:00:00",
}
storage = _mock_storage(identity=existing_identity, user=None)
claims = {"sub": "sub-orphan", "email": "gone@example.com"}
with pytest.raises(OIDCError, match="missing user"):
provision_oidc_user(storage, config, claims)
def test_provision_oidc_user_fallback_username(self):
"""No preferred_username and no email -> falls back to 'user'."""
config = _make_config()
storage = _mock_storage()
new_user = {
"user_id": "u-new",
"username": "user",
"display_name": "",
"password_hash": "!oidc",
}
storage.get_user.return_value = new_user
claims = {"sub": "sub-noemail"}
provision_oidc_user(storage, config, claims)
call_args = storage.create_user.call_args
assert call_args[0][1] == "user"
# ---------------------------------------------------------------------------
# Role Mapping
# ---------------------------------------------------------------------------
class TestApplyRoleMapping:
def test_apply_role_mapping_basic(self):
"""Maps claim value to role."""
config = _make_config(
role_claim="groups",
role_map={"admin": "builtin-admin"},
)
storage = _mock_storage(role={"role_id": "builtin-admin", "name": "Admin"})
claims = {"sub": "u1", "groups": "admin"}
apply_role_mapping(storage, "u1", claims, config)
storage.assign_role.assert_called_once_with("u1", "builtin-admin", "oidc")
def test_apply_role_mapping_list_claim(self):
"""Claim is a list of strings -> maps each."""
config = _make_config(
role_claim="roles",
role_map={"admin": "builtin-admin", "editor": "builtin-operator"},
)
storage = _mock_storage()
# get_role returns non-None for both roles
storage.get_role.return_value = {"role_id": "some-role"}
claims = {"sub": "u1", "roles": ["admin", "editor"]}
apply_role_mapping(storage, "u1", claims, config)
assert storage.assign_role.call_count == 2
def test_apply_role_mapping_no_config(self):
"""No role_claim configured -> no-op."""
config = _make_config(role_claim="", role_map={})
storage = _mock_storage()
claims = {"sub": "u1", "roles": "admin"}
apply_role_mapping(storage, "u1", claims, config)
storage.assign_role.assert_not_called()
def test_apply_role_mapping_unknown_role(self):
"""Claim maps to nonexistent role -> skipped."""
config = _make_config(
role_claim="groups",
role_map={"admin": "nonexistent-role"},
)
storage = _mock_storage(role=None) # role doesn't exist
claims = {"sub": "u1", "groups": "admin"}
apply_role_mapping(storage, "u1", claims, config)
storage.assign_role.assert_not_called()
def test_apply_role_mapping_no_matching_claim_value(self):
"""Claim value not in role_map -> no assignment."""
config = _make_config(
role_claim="groups",
role_map={"admin": "builtin-admin"},
)
storage = _mock_storage()
claims = {"sub": "u1", "groups": "viewer"} # "viewer" not in role_map
apply_role_mapping(storage, "u1", claims, config)
storage.assign_role.assert_not_called()
def test_apply_role_mapping_claim_missing(self):
"""Claim key not present in claims -> no-op."""
config = _make_config(
role_claim="groups",
role_map={"admin": "builtin-admin"},
)
storage = _mock_storage()
claims = {"sub": "u1"} # no "groups" key
apply_role_mapping(storage, "u1", claims, config)
storage.assign_role.assert_not_called()
def test_apply_role_mapping_no_role_map(self):
"""role_claim set but role_map empty -> no-op (early return)."""
config = _make_config(role_claim="groups", role_map={})
storage = _mock_storage()
claims = {"sub": "u1", "groups": "admin"}
apply_role_mapping(storage, "u1", claims, config)
storage.assign_role.assert_not_called()
def test_apply_role_mapping_revokes_stale_oidc_roles(self):
"""Roles previously assigned by OIDC but no longer in claims are revoked."""
config = _make_config(
role_claim="groups",
role_map={"admin": "builtin-admin", "eng": "builtin-operator"},
)
storage = _mock_storage()
storage.get_role.return_value = {"role_id": "some-role"}
# User currently has admin (via OIDC) and a manual role
storage.list_user_roles.return_value = [
{"role_id": "builtin-admin", "assigned_by": "oidc"},
{"role_id": "custom-role", "assigned_by": "admin-ui"},
]
# IdP now only says "eng", not "admin"
claims = {"sub": "u1", "groups": ["eng"]}
apply_role_mapping(storage, "u1", claims, config)
# builtin-admin should be revoked (OIDC-assigned, no longer in claims)
storage.unassign_role.assert_called_once_with("u1", "builtin-admin")
# custom-role should NOT be revoked (not assigned by OIDC)
def test_apply_role_mapping_preserves_manual_roles(self):
"""Manually assigned roles are never revoked by OIDC sync."""
config = _make_config(
role_claim="groups",
role_map={"admin": "builtin-admin"},
)
storage = _mock_storage()
storage.get_role.return_value = {"role_id": "some-role"}
storage.list_user_roles.return_value = [
{"role_id": "builtin-admin", "assigned_by": "admin-ui"},
]
# Claims have no groups at all
claims = {"sub": "u1"}
apply_role_mapping(storage, "u1", claims, config)
# Manual admin role must NOT be revoked
storage.unassign_role.assert_not_called()
def test_apply_role_mapping_revokes_all_oidc_roles_when_claim_absent(self):
"""When the claim is absent from the token, all OIDC-assigned roles are revoked."""
config = _make_config(
role_claim="groups",
role_map={"admin": "builtin-admin"},
)
storage = _mock_storage()
storage.get_role.return_value = {"role_id": "some-role"}
storage.list_user_roles.return_value = [
{"role_id": "builtin-admin", "assigned_by": "oidc"},
{"role_id": "builtin-operator", "assigned_by": "oidc"},
]
claims = {"sub": "u1"} # no "groups" key
apply_role_mapping(storage, "u1", claims, config)
assert storage.unassign_role.call_count == 2
# ---------------------------------------------------------------------------
# Discovery (async)
# ---------------------------------------------------------------------------
class TestDiscoverOIDC:
def test_discover_oidc_success(self):
"""Mock httpx response, verify endpoints populated."""
config = _make_config(
authorization_endpoint="",
token_endpoint="",
userinfo_endpoint="",
jwks_uri="",
)
discovery_doc = {
"authorization_endpoint": "https://idp.example.com/authorize",
"token_endpoint": "https://idp.example.com/token",
"userinfo_endpoint": "https://idp.example.com/userinfo",
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
}
mock_response = MagicMock()
mock_response.json.return_value = discovery_doc
mock_response.raise_for_status = MagicMock()
async def _run():
client = _mock_async_client(lambda url: _async_return(mock_response))
with patch("httpx.AsyncClient", return_value=client):
result = await discover_oidc(config)
assert result.authorization_endpoint == "https://idp.example.com/authorize"
assert result.token_endpoint == "https://idp.example.com/token"
assert result.userinfo_endpoint == "https://idp.example.com/userinfo"
assert result.jwks_uri == "https://idp.example.com/.well-known/jwks.json"
assert result.enabled is True
asyncio.run(_run())
def test_discover_oidc_failure(self):
"""Mock httpx error -> enabled=False returned."""
config = _make_config(
authorization_endpoint="",
token_endpoint="",
userinfo_endpoint="",
jwks_uri="",
)
async def _failing_get(url):
raise httpx.ConnectError("connection refused")
async def _run():
client = _mock_async_client(_failing_get)
with patch("httpx.AsyncClient", return_value=client):
result = await discover_oidc(config)
assert result.enabled is False
asyncio.run(_run())
def test_discover_oidc_no_issuer(self):
"""Empty issuer -> enabled=False."""
config = _make_config(issuer="")
async def _run():
result = await discover_oidc(config)
assert result.enabled is False
asyncio.run(_run())
def test_discover_oidc_missing_required_endpoints(self):
"""Discovery doc missing authorization_endpoint -> enabled=False."""
config = _make_config(
authorization_endpoint="",
token_endpoint="",
userinfo_endpoint="",
jwks_uri="",
)
# Document missing authorization_endpoint
discovery_doc = {
"token_endpoint": "https://idp.example.com/token",
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
}
mock_response = MagicMock()
mock_response.json.return_value = discovery_doc
mock_response.raise_for_status = MagicMock()
async def _run():
client = _mock_async_client(lambda url: _async_return(mock_response))
with patch("httpx.AsyncClient", return_value=client):
result = await discover_oidc(config)
assert result.enabled is False
asyncio.run(_run())
async def _async_return(value):
"""Helper: return a value from an async function."""
return value
+581
View File
@@ -0,0 +1,581 @@
"""Integration tests for OIDC HTTP handlers (authorize, callback, admin endpoints).
Uses Starlette TestClient with real SQLiteBackend storage. External OIDC
functions (exchange_code, validate_id_token, etc.) are mocked the focus is
on the HTTP handler logic, request/response wiring, and storage side-effects.
"""
from __future__ import annotations
import urllib.parse
from typing import TYPE_CHECKING, Any
from unittest.mock import AsyncMock, patch
import pytest
from starlette.applications import Starlette
from starlette.middleware import Middleware
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.routing import Mount, Route
from starlette.testclient import TestClient
if TYPE_CHECKING:
from starlette.requests import Request
from starlette.responses import Response
from turnstone.console.server import (
admin_delete_oidc_identity,
admin_list_oidc_identities,
)
from turnstone.core.auth import (
AuthResult,
LoginRateLimiter,
handle_oidc_authorize,
handle_oidc_callback,
)
from turnstone.core.oidc import OIDCConfig, OIDCError
from turnstone.core.storage._sqlite import SQLiteBackend
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_oidc_config(**overrides: Any) -> OIDCConfig:
"""Build a test OIDCConfig with sensible defaults."""
defaults: dict[str, Any] = {
"enabled": True,
"issuer": "https://idp.example.com",
"client_id": "my-client",
"client_secret": "my-secret",
"scopes": "openid email profile",
"provider_name": "TestIDP",
"role_claim": "",
"role_map": {},
"password_enabled": True,
"authorization_endpoint": "https://idp.example.com/authorize",
"token_endpoint": "https://idp.example.com/token",
"userinfo_endpoint": "https://idp.example.com/userinfo",
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
}
defaults.update(overrides)
return OIDCConfig(**defaults)
# ---------------------------------------------------------------------------
# Thin handler wrappers — match the pattern used in server.py / console
# ---------------------------------------------------------------------------
async def _oidc_authorize(request: Request) -> Response:
return await handle_oidc_authorize(request, "test-audience")
async def _oidc_callback(request: Request) -> Response:
return await handle_oidc_callback(request, "test-audience")
# ---------------------------------------------------------------------------
# Auth bypass middleware for admin endpoints
# ---------------------------------------------------------------------------
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.users",
}
),
)
return await call_next(request)
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def storage(tmp_path: Any) -> SQLiteBackend:
"""Fresh SQLite backend with a seeded admin user."""
backend = SQLiteBackend(str(tmp_path / "test.db"))
backend.create_user("test-admin", "testadmin", "Test Admin", "hash")
return backend
@pytest.fixture
def oidc_config() -> OIDCConfig:
return _make_oidc_config()
@pytest.fixture
def authorize_client(storage: SQLiteBackend, oidc_config: OIDCConfig) -> TestClient:
"""TestClient wired to the OIDC authorize + callback handlers."""
app = Starlette(
routes=[
Mount(
"/v1",
routes=[
Route("/api/auth/oidc/authorize", _oidc_authorize),
Route("/api/auth/oidc/callback", _oidc_callback),
],
),
],
)
app.state.oidc_config = oidc_config
app.state.auth_storage = storage
app.state.jwt_secret = "test-jwt-secret"
app.state.jwks_data = {"keys": []}
app.state.login_limiter = None
return TestClient(app, raise_server_exceptions=False)
@pytest.fixture
def admin_client(storage: SQLiteBackend) -> TestClient:
"""TestClient wired to the admin OIDC identity endpoints."""
app = Starlette(
routes=[
Mount(
"/v1",
routes=[
Route(
"/api/admin/users/{user_id}/oidc-identities",
admin_list_oidc_identities,
),
Route(
"/api/admin/oidc-identities",
admin_delete_oidc_identity,
methods=["DELETE"],
),
],
),
],
middleware=[Middleware(_InjectAuthMiddleware)],
)
app.state.auth_storage = storage
return TestClient(app, raise_server_exceptions=False)
# ---------------------------------------------------------------------------
# /authorize tests
# ---------------------------------------------------------------------------
class TestOIDCAuthorize:
"""Tests for GET /v1/api/auth/oidc/authorize."""
def test_happy_path_redirects_to_idp(self, authorize_client: TestClient) -> None:
resp = authorize_client.get("/v1/api/auth/oidc/authorize", follow_redirects=False)
assert resp.status_code == 302
location = resp.headers["location"]
assert location.startswith("https://idp.example.com/authorize?")
parsed = urllib.parse.urlparse(location)
params = urllib.parse.parse_qs(parsed.query)
assert params["response_type"] == ["code"]
assert params["client_id"] == ["my-client"]
assert params["scope"] == ["openid email profile"]
assert "state" in params
assert "nonce" in params
assert "code_challenge" in params
assert params["code_challenge_method"] == ["S256"]
def test_oidc_not_configured_returns_404(
self,
storage: SQLiteBackend,
) -> None:
app = Starlette(
routes=[Mount("/v1", routes=[Route("/api/auth/oidc/authorize", _oidc_authorize)])]
)
app.state.auth_storage = storage
# No oidc_config at all
client = TestClient(app, raise_server_exceptions=False)
resp = client.get("/v1/api/auth/oidc/authorize")
assert resp.status_code == 404
assert resp.json()["error"] == "OIDC not configured"
def test_oidc_not_enabled_returns_404(
self,
storage: SQLiteBackend,
) -> None:
app = Starlette(
routes=[Mount("/v1", routes=[Route("/api/auth/oidc/authorize", _oidc_authorize)])]
)
app.state.oidc_config = _make_oidc_config(enabled=False)
app.state.auth_storage = storage
client = TestClient(app, raise_server_exceptions=False)
resp = client.get("/v1/api/auth/oidc/authorize")
assert resp.status_code == 404
def test_no_storage_returns_503(self) -> None:
app = Starlette(
routes=[Mount("/v1", routes=[Route("/api/auth/oidc/authorize", _oidc_authorize)])]
)
app.state.oidc_config = _make_oidc_config()
app.state.login_limiter = None
# No auth_storage
client = TestClient(app, raise_server_exceptions=False)
resp = client.get("/v1/api/auth/oidc/authorize")
assert resp.status_code == 503
def test_no_users_returns_403(self, tmp_path: Any) -> None:
backend = SQLiteBackend(str(tmp_path / "empty.db"))
app = Starlette(
routes=[Mount("/v1", routes=[Route("/api/auth/oidc/authorize", _oidc_authorize)])]
)
app.state.oidc_config = _make_oidc_config()
app.state.auth_storage = backend
app.state.login_limiter = None
client = TestClient(app, raise_server_exceptions=False)
resp = client.get("/v1/api/auth/oidc/authorize")
assert resp.status_code == 403
assert "setup" in resp.json()["error"].lower()
def test_pending_state_persisted(
self,
authorize_client: TestClient,
storage: SQLiteBackend,
) -> None:
resp = authorize_client.get("/v1/api/auth/oidc/authorize", follow_redirects=False)
assert resp.status_code == 302
location = resp.headers["location"]
parsed = urllib.parse.urlparse(location)
params = urllib.parse.parse_qs(parsed.query)
state = params["state"][0]
# The pending state should be retrievable from storage
pending = storage.pop_oidc_pending_state(state, max_age_seconds=300)
assert pending is not None
assert pending["audience"] == "test-audience"
assert pending["nonce"] != ""
assert pending["code_verifier"] != ""
def test_rate_limited_redirects_with_error(self, storage: SQLiteBackend) -> None:
app = Starlette(
routes=[Mount("/v1", routes=[Route("/api/auth/oidc/authorize", _oidc_authorize)])]
)
app.state.oidc_config = _make_oidc_config()
app.state.auth_storage = storage
limiter = LoginRateLimiter(max_attempts=1, window_seconds=300)
# Exhaust the rate limit
limiter.record("ip:testclient")
app.state.login_limiter = limiter
client = TestClient(app, raise_server_exceptions=False)
resp = client.get("/v1/api/auth/oidc/authorize", follow_redirects=False)
assert resp.status_code == 302
assert "oidc_error" in resp.headers["location"]
assert "Too+many" in resp.headers["location"]
# ---------------------------------------------------------------------------
# /callback tests
# ---------------------------------------------------------------------------
class TestOIDCCallback:
"""Tests for GET /v1/api/auth/oidc/callback."""
def _seed_pending_state(
self,
storage: SQLiteBackend,
state: str = "valid-state",
nonce: str = "test-nonce",
code_verifier: str = "test-verifier",
audience: str = "test-audience",
) -> None:
storage.create_oidc_pending_state(state, nonce, code_verifier, audience)
@patch("turnstone.core.oidc.provision_oidc_user")
@patch("turnstone.core.oidc.validate_id_token")
@patch("turnstone.core.oidc.exchange_code", new_callable=AsyncMock)
def test_happy_path(
self,
mock_exchange: AsyncMock,
mock_validate: Any,
mock_provision: Any,
authorize_client: TestClient,
storage: SQLiteBackend,
) -> None:
self._seed_pending_state(storage)
mock_exchange.return_value = {"id_token": "fake.jwt.token", "access_token": "at"}
mock_validate.return_value = {
"sub": "user123",
"email": "u@example.com",
"nonce": "test-nonce",
}
mock_provision.return_value = {"user_id": "test-admin", "username": "testadmin"}
resp = authorize_client.get(
"/v1/api/auth/oidc/callback?code=authcode&state=valid-state",
follow_redirects=False,
)
assert resp.status_code == 302
assert "oidc_success=1" in resp.headers["location"]
assert "set-cookie" in resp.headers
assert "turnstone_auth=" in resp.headers["set-cookie"]
def test_oidc_not_configured_returns_404(self, storage: SQLiteBackend) -> None:
app = Starlette(
routes=[Mount("/v1", routes=[Route("/api/auth/oidc/callback", _oidc_callback)])]
)
app.state.auth_storage = storage
# No oidc_config
client = TestClient(app, raise_server_exceptions=False)
resp = client.get("/v1/api/auth/oidc/callback?code=x&state=y")
assert resp.status_code == 404
def test_idp_error_param_redirects(
self,
authorize_client: TestClient,
) -> None:
resp = authorize_client.get(
"/v1/api/auth/oidc/callback?error=access_denied&error_description=User+cancelled",
follow_redirects=False,
)
assert resp.status_code == 302
location = resp.headers["location"]
assert "oidc_error" in location
assert "User" in urllib.parse.unquote(location)
def test_invalid_state_redirects_expired(
self,
authorize_client: TestClient,
) -> None:
resp = authorize_client.get(
"/v1/api/auth/oidc/callback?code=authcode&state=nonexistent",
follow_redirects=False,
)
assert resp.status_code == 302
assert "Login+session+expired" in resp.headers["location"]
def test_expired_state_redirects(
self,
authorize_client: TestClient,
storage: SQLiteBackend,
) -> None:
# Insert state then backdate created_at via raw SQL so that
# pop_oidc_pending_state's max_age_seconds=300 check rejects it.
self._seed_pending_state(storage, state="old-state")
import sqlalchemy as sa
with storage._engine.connect() as conn:
conn.execute(
sa.text(
"UPDATE oidc_pending_states SET created_at = '2020-01-01T00:00:00' "
"WHERE state = 'old-state'"
)
)
conn.commit()
resp = authorize_client.get(
"/v1/api/auth/oidc/callback?code=authcode&state=old-state",
follow_redirects=False,
)
assert resp.status_code == 302
assert "Login+session+expired" in resp.headers["location"]
@patch("turnstone.core.oidc.exchange_code", new_callable=AsyncMock)
def test_code_exchange_failure(
self,
mock_exchange: AsyncMock,
authorize_client: TestClient,
storage: SQLiteBackend,
) -> None:
self._seed_pending_state(storage)
mock_exchange.side_effect = OIDCError("Token endpoint error")
resp = authorize_client.get(
"/v1/api/auth/oidc/callback?code=badcode&state=valid-state",
follow_redirects=False,
)
assert resp.status_code == 302
assert "Authentication+failed" in resp.headers["location"]
@patch("turnstone.core.oidc.validate_id_token")
@patch("turnstone.core.oidc.exchange_code", new_callable=AsyncMock)
def test_token_validation_failure(
self,
mock_exchange: AsyncMock,
mock_validate: Any,
authorize_client: TestClient,
storage: SQLiteBackend,
) -> None:
self._seed_pending_state(storage)
mock_exchange.return_value = {"id_token": "fake.jwt.token"}
mock_validate.side_effect = OIDCError("Signature invalid")
resp = authorize_client.get(
"/v1/api/auth/oidc/callback?code=authcode&state=valid-state",
follow_redirects=False,
)
assert resp.status_code == 302
assert "Authentication+failed" in resp.headers["location"]
@patch("turnstone.core.oidc.provision_oidc_user")
@patch("turnstone.core.oidc.validate_id_token")
@patch("turnstone.core.oidc.fetch_jwks", new_callable=AsyncMock)
@patch("turnstone.core.oidc.exchange_code", new_callable=AsyncMock)
def test_jwks_key_rotation_retry(
self,
mock_exchange: AsyncMock,
mock_fetch_jwks: AsyncMock,
mock_validate: Any,
mock_provision: Any,
authorize_client: TestClient,
storage: SQLiteBackend,
) -> None:
"""First validate raises 'kid not found in JWKS', fetch_jwks retried, second validate succeeds."""
self._seed_pending_state(storage)
mock_exchange.return_value = {"id_token": "fake.jwt.token"}
# First call raises kid-not-found; second call (after JWKS refresh) succeeds
mock_validate.side_effect = [
OIDCError("Signing key 'new-kid' not found in JWKS"),
{"sub": "user123", "email": "u@example.com", "nonce": "test-nonce"},
]
mock_fetch_jwks.return_value = {"keys": [{"kid": "new-kid", "kty": "RSA"}]}
mock_provision.return_value = {"user_id": "test-admin", "username": "testadmin"}
resp = authorize_client.get(
"/v1/api/auth/oidc/callback?code=authcode&state=valid-state",
follow_redirects=False,
)
assert resp.status_code == 302
assert "oidc_success=1" in resp.headers["location"]
mock_fetch_jwks.assert_called_once()
assert mock_validate.call_count == 2
@patch("turnstone.core.oidc.provision_oidc_user")
@patch("turnstone.core.oidc.validate_id_token")
@patch("turnstone.core.oidc.exchange_code", new_callable=AsyncMock)
def test_no_users_after_oidc_success_redirects_setup(
self,
mock_exchange: AsyncMock,
mock_validate: Any,
mock_provision: Any,
tmp_path: Any,
) -> None:
"""When OIDC succeeds but no users exist (edge case), redirect with setup error."""
# Use a fresh empty-user storage
backend = SQLiteBackend(str(tmp_path / "empty.db"))
app = Starlette(
routes=[Mount("/v1", routes=[Route("/api/auth/oidc/callback", _oidc_callback)])]
)
app.state.oidc_config = _make_oidc_config()
app.state.auth_storage = backend
app.state.jwt_secret = "secret"
app.state.jwks_data = {"keys": []}
app.state.login_limiter = None
# Seed a pending state in the empty database
backend.create_oidc_pending_state("state1", "nonce1", "verifier1", "test-audience")
mock_exchange.return_value = {"id_token": "fake.jwt.token"}
mock_validate.return_value = {"sub": "u1", "email": "u@example.com", "nonce": "nonce1"}
client = TestClient(app, raise_server_exceptions=False)
resp = client.get(
"/v1/api/auth/oidc/callback?code=authcode&state=state1",
follow_redirects=False,
)
assert resp.status_code == 302
assert "Initial+setup+required" in resp.headers["location"]
def test_rate_limited_redirects_with_error(self, storage: SQLiteBackend) -> None:
app = Starlette(
routes=[Mount("/v1", routes=[Route("/api/auth/oidc/callback", _oidc_callback)])]
)
app.state.oidc_config = _make_oidc_config()
app.state.auth_storage = storage
app.state.jwt_secret = "secret"
app.state.jwks_data = {"keys": []}
limiter = LoginRateLimiter(max_attempts=1, window_seconds=300)
limiter.record("ip:testclient")
app.state.login_limiter = limiter
client = TestClient(app, raise_server_exceptions=False)
resp = client.get(
"/v1/api/auth/oidc/callback?code=authcode&state=x",
follow_redirects=False,
)
assert resp.status_code == 302
assert "oidc_error" in resp.headers["location"]
assert "Too+many" in resp.headers["location"]
# ---------------------------------------------------------------------------
# Admin OIDC identity endpoint tests
# ---------------------------------------------------------------------------
class TestAdminOIDCIdentities:
"""Tests for admin OIDC identity management endpoints."""
def test_list_identities(
self,
admin_client: TestClient,
storage: SQLiteBackend,
) -> None:
storage.create_oidc_identity(
"https://idp.example.com",
"sub-123",
"test-admin",
"admin@example.com",
)
resp = admin_client.get("/v1/api/admin/users/test-admin/oidc-identities")
assert resp.status_code == 200
data = resp.json()
assert len(data["oidc_identities"]) == 1
identity = data["oidc_identities"][0]
assert identity["issuer"] == "https://idp.example.com"
assert identity["subject"] == "sub-123"
assert identity["user_id"] == "test-admin"
assert identity["email"] == "admin@example.com"
def test_list_empty(self, admin_client: TestClient) -> None:
resp = admin_client.get("/v1/api/admin/users/test-admin/oidc-identities")
assert resp.status_code == 200
assert resp.json()["oidc_identities"] == []
def test_delete_identity(
self,
admin_client: TestClient,
storage: SQLiteBackend,
) -> None:
storage.create_oidc_identity(
"https://idp.example.com",
"sub-456",
"test-admin",
"admin@example.com",
)
resp = admin_client.delete(
"/v1/api/admin/oidc-identities?issuer=https://idp.example.com&subject=sub-456",
)
assert resp.status_code == 200
assert resp.json()["status"] == "ok"
# Verify it's gone
assert storage.get_oidc_identity("https://idp.example.com", "sub-456") is None
def test_delete_nonexistent_returns_404(self, admin_client: TestClient) -> None:
resp = admin_client.delete(
"/v1/api/admin/oidc-identities?issuer=https://no.such&subject=nope",
)
assert resp.status_code == 404
assert "not found" in resp.json()["error"].lower()
def test_delete_missing_params_returns_400(self, admin_client: TestClient) -> None:
# Missing subject
resp = admin_client.delete(
"/v1/api/admin/oidc-identities?issuer=https://idp.example.com",
)
assert resp.status_code == 400
assert "required" in resp.json()["error"].lower()
# Missing both
resp = admin_client.delete("/v1/api/admin/oidc-identities")
assert resp.status_code == 400
+315
View File
@@ -0,0 +1,315 @@
"""Tests for OIDC identity and pending state storage CRUD (SQLite backend)."""
from __future__ import annotations
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
# ---------------------------------------------------------------------------
class TestOIDCIdentityCRUD:
def test_create_and_get_oidc_identity(self, db):
db.create_oidc_identity("https://idp.example.com", "sub-123", "u1", "alice@example.com")
identity = db.get_oidc_identity("https://idp.example.com", "sub-123")
assert identity is not None
assert identity["issuer"] == "https://idp.example.com"
assert identity["subject"] == "sub-123"
assert identity["user_id"] == "u1"
assert identity["email"] == "alice@example.com"
assert identity["created"] != ""
assert identity["last_login"] != ""
def test_get_oidc_identity_not_found(self, db):
assert db.get_oidc_identity("https://unknown.example.com", "sub-999") is None
def test_create_oidc_identity_idempotent(self, db):
"""Creating twice with same (issuer, subject) does not error (OR IGNORE)."""
db.create_oidc_identity("https://idp.example.com", "sub-123", "u1", "alice@example.com")
db.create_oidc_identity("https://idp.example.com", "sub-123", "u2", "bob@example.com")
identity = db.get_oidc_identity("https://idp.example.com", "sub-123")
assert identity is not None
# OR IGNORE preserves the first insert
assert identity["user_id"] == "u1"
assert identity["email"] == "alice@example.com"
def test_update_oidc_identity_login(self, db):
db.create_oidc_identity("https://idp.example.com", "sub-123", "u1", "alice@example.com")
before = db.get_oidc_identity("https://idp.example.com", "sub-123")
assert before is not None
original_login = before["last_login"]
# Small sleep to ensure timestamp differs
time.sleep(0.05)
result = db.update_oidc_identity_login("https://idp.example.com", "sub-123")
assert result is True
after = db.get_oidc_identity("https://idp.example.com", "sub-123")
assert after is not None
assert after["last_login"] >= original_login
def test_update_oidc_identity_login_nonexistent(self, db):
result = db.update_oidc_identity_login("https://idp.example.com", "sub-999")
assert result is False
def test_list_oidc_identities_for_user(self, db):
"""Two identities for same user, list returns both."""
db.create_oidc_identity("https://idp1.example.com", "sub-A", "u1", "alice@idp1.com")
db.create_oidc_identity("https://idp2.example.com", "sub-B", "u1", "alice@idp2.com")
identities = db.list_oidc_identities_for_user("u1")
assert len(identities) == 2
issuers = {i["issuer"] for i in identities}
assert issuers == {"https://idp1.example.com", "https://idp2.example.com"}
def test_list_oidc_identities_for_user_empty(self, db):
assert db.list_oidc_identities_for_user("u-none") == []
def test_list_oidc_identities_excludes_other_users(self, db):
db.create_oidc_identity("https://idp.example.com", "sub-1", "u1", "alice@example.com")
db.create_oidc_identity("https://idp.example.com", "sub-2", "u2", "bob@example.com")
identities = db.list_oidc_identities_for_user("u1")
assert len(identities) == 1
assert identities[0]["user_id"] == "u1"
def test_delete_oidc_identity(self, db):
db.create_oidc_identity("https://idp.example.com", "sub-123", "u1", "alice@example.com")
assert db.delete_oidc_identity("https://idp.example.com", "sub-123") is True
assert db.get_oidc_identity("https://idp.example.com", "sub-123") is None
def test_delete_oidc_identity_nonexistent(self, db):
assert db.delete_oidc_identity("https://idp.example.com", "sub-999") is False
def test_delete_oidc_identity_only_deletes_target(self, db):
"""Deleting one identity does not affect others."""
db.create_oidc_identity("https://idp.example.com", "sub-1", "u1", "a@example.com")
db.create_oidc_identity("https://idp.example.com", "sub-2", "u1", "b@example.com")
db.delete_oidc_identity("https://idp.example.com", "sub-1")
assert db.get_oidc_identity("https://idp.example.com", "sub-1") is None
assert db.get_oidc_identity("https://idp.example.com", "sub-2") is not None
# ---------------------------------------------------------------------------
# OIDC Pending State
# ---------------------------------------------------------------------------
class TestOIDCPendingState:
def test_create_and_pop_pending_state(self, db):
db.create_oidc_pending_state(
state="state-abc",
nonce="nonce-xyz",
code_verifier="verifier-123",
audience="server",
)
result = db.pop_oidc_pending_state("state-abc")
assert result is not None
assert result["state"] == "state-abc"
assert result["nonce"] == "nonce-xyz"
assert result["code_verifier"] == "verifier-123"
assert result["audience"] == "server"
assert result["created_at"] != ""
def test_pop_pending_state_not_found(self, db):
assert db.pop_oidc_pending_state("nonexistent-state") is None
def test_pop_pending_state_expired(self, db):
"""Create with old timestamp, pop returns None."""
# Insert a row with an old created_at timestamp directly
import sqlalchemy as sa
from turnstone.core.storage._schema import oidc_pending_states
with db._engine.connect() as conn:
conn.execute(
sa.insert(oidc_pending_states),
{
"state": "state-old",
"nonce": "nonce-old",
"code_verifier": "verifier-old",
"audience": "server",
"created_at": "2020-01-01T00:00:00",
},
)
conn.commit()
# Default max_age_seconds=300, so a 2020 timestamp is expired
result = db.pop_oidc_pending_state("state-old")
assert result is None
def test_pop_pending_state_consumed(self, db):
"""Pop twice -> second returns None (one-time use)."""
db.create_oidc_pending_state(
state="state-once",
nonce="nonce-1",
code_verifier="verifier-1",
audience="server",
)
first = db.pop_oidc_pending_state("state-once")
assert first is not None
second = db.pop_oidc_pending_state("state-once")
assert second is None
def test_pop_pending_state_custom_max_age(self, db):
"""Custom max_age_seconds allows longer-lived states."""
db.create_oidc_pending_state(
state="state-long",
nonce="nonce-long",
code_verifier="verifier-long",
audience="server",
)
# With very short max_age, it might still be valid since we just created it
result = db.pop_oidc_pending_state("state-long", max_age_seconds=600)
assert result is not None
def test_create_pending_state_duplicate_raises(self, db):
"""Duplicate state insertion raises IntegrityError (no silent drop)."""
import sqlalchemy.exc
db.create_oidc_pending_state("state-dup", "nonce-1", "verifier-1", "server")
with pytest.raises(sqlalchemy.exc.IntegrityError):
db.create_oidc_pending_state("state-dup", "nonce-2", "verifier-2", "server")
def test_cleanup_expired_states(self, db):
"""Create expired + fresh, cleanup removes only expired."""
import sqlalchemy as sa
from turnstone.core.storage._schema import oidc_pending_states
# Insert an expired state directly with old timestamp
with db._engine.connect() as conn:
conn.execute(
sa.insert(oidc_pending_states),
{
"state": "state-expired",
"nonce": "nonce-old",
"code_verifier": "verifier-old",
"audience": "server",
"created_at": "2020-01-01T00:00:00",
},
)
conn.commit()
# Insert a fresh state via normal API
db.create_oidc_pending_state("state-fresh", "nonce-new", "verifier-new", "server")
# Cleanup with default 300s max age
deleted = db.cleanup_expired_oidc_states()
assert deleted == 1
# Fresh state should still exist
result = db.pop_oidc_pending_state("state-fresh")
assert result is not None
def test_cleanup_expired_states_none_expired(self, db):
"""Cleanup with no expired states returns 0."""
db.create_oidc_pending_state("state-1", "nonce-1", "verifier-1", "server")
deleted = db.cleanup_expired_oidc_states()
assert deleted == 0
def test_cleanup_expired_states_all_expired(self, db):
"""Cleanup with all expired states removes all."""
import sqlalchemy as sa
from turnstone.core.storage._schema import oidc_pending_states
with db._engine.connect() as conn:
for i in range(3):
conn.execute(
sa.insert(oidc_pending_states),
{
"state": f"state-{i}",
"nonce": f"nonce-{i}",
"code_verifier": f"verifier-{i}",
"audience": "server",
"created_at": "2020-01-01T00:00:00",
},
)
conn.commit()
deleted = db.cleanup_expired_oidc_states()
assert deleted == 3
def test_cleanup_expired_states_custom_max_age(self, db):
"""Custom max_age_seconds affects what counts as expired."""
from datetime import UTC, datetime, timedelta
import sqlalchemy as sa
from turnstone.core.storage._schema import oidc_pending_states
# Insert a state created 60 seconds ago
old_ts = (datetime.now(UTC) - timedelta(seconds=60)).strftime("%Y-%m-%dT%H:%M:%S")
with db._engine.connect() as conn:
conn.execute(
sa.insert(oidc_pending_states),
{
"state": "state-1",
"nonce": "nonce-1",
"code_verifier": "verifier-1",
"audience": "server",
"created_at": old_ts,
},
)
conn.commit()
# With default max_age=300s the 60s-old state is NOT expired
deleted = db.cleanup_expired_oidc_states(max_age_seconds=300)
assert deleted == 0
# With max_age=30s the 60s-old state IS expired
deleted = db.cleanup_expired_oidc_states(max_age_seconds=30)
assert deleted == 1
def test_pop_expired_cleans_up_row(self, db):
"""Popping an expired state should delete the row (not leave orphan)."""
import sqlalchemy as sa
from turnstone.core.storage._schema import oidc_pending_states
with db._engine.connect() as conn:
conn.execute(
sa.insert(oidc_pending_states),
{
"state": "state-cleanup",
"nonce": "nonce-c",
"code_verifier": "verifier-c",
"audience": "server",
"created_at": "2020-01-01T00:00:00",
},
)
conn.commit()
# Pop returns None (expired)
assert db.pop_oidc_pending_state("state-cleanup") is None
# Row should be gone (cleaned up even though expired)
with db._engine.connect() as conn:
count = conn.execute(
sa.select(sa.func.count())
.select_from(oidc_pending_states)
.where(oidc_pending_states.c.state == "state-cleanup")
).scalar()
assert count == 0
+129
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import threading
from unittest.mock import MagicMock
from turnstone.core.session import ChatSession, _render_template
@@ -391,3 +392,131 @@ class TestMCPTemplates:
session = _make_session(template="mcp__server__code")
content = _sys_content(session)
assert "MCP_EXPLICIT" in content
# ---------------------------------------------------------------------------
# Resume with deleted template
# ---------------------------------------------------------------------------
class TestResumeDeletedTemplate:
def test_resume_with_deleted_template_degrades_gracefully(self, tmp_db, capsys):
from turnstone.core.memory import save_message
from turnstone.core.storage import get_storage
db = get_storage()
_create_template(db, "t1", "ephemeral-tpl", "EPHEMERAL_CONTENT", is_default=False)
# Create session with template, save a message so resume has history
session1 = _make_session(template="ephemeral-tpl")
ws_id = session1.ws_id
save_message(ws_id, "user", "hello")
assert "EPHEMERAL_CONTENT" in _sys_content(session1)
# Delete the template from storage
db.delete_prompt_template("t1")
# Resume into a new session
session2 = _make_session()
resumed = session2.resume(ws_id)
assert resumed
assert session2._template_name == "ephemeral-tpl"
assert session2._template_content is None
# System message should not contain the deleted template content
content = _sys_content(session2)
assert "EPHEMERAL_CONTENT" not in content
# Warning should be logged via structlog
captured = capsys.readouterr()
assert "not_found" in captured.out or "not_found" in captured.err
# ---------------------------------------------------------------------------
# Threading safety
# ---------------------------------------------------------------------------
class TestTemplateFactoryPassthrough:
def test_template_passed_through_workstream_create(self, tmp_db):
"""WorkstreamManager.create(template=...) propagates to session factory."""
from turnstone.core.storage import get_storage
from turnstone.core.workstream import WorkstreamManager
db = get_storage()
_create_template(db, "t1", "factory-tpl", "FACTORY_CONTENT", is_default=False)
captured_template = None
def factory(ui, model_alias=None, ws_id=None, *, template=None):
nonlocal captured_template
captured_template = template
return _make_session(template=template)
mgr = WorkstreamManager(factory)
ws = mgr.create(name="test", template="factory-tpl")
assert captured_template == "factory-tpl"
assert ws.session is not None
assert ws.session._template_name == "factory-tpl"
assert "FACTORY_CONTENT" in _sys_content(ws.session)
def test_template_none_uses_defaults(self, tmp_db):
"""WorkstreamManager.create() without template passes None."""
captured_template = "sentinel"
def factory(ui, model_alias=None, ws_id=None, *, template=None):
nonlocal captured_template
captured_template = template
return _make_session(template=template)
from turnstone.core.workstream import WorkstreamManager
mgr = WorkstreamManager(factory)
mgr.create(name="test")
assert captured_template is None
class TestTemplateThreadSafety:
def test_concurrent_template_and_system_message_init(self, tmp_db):
from turnstone.core.storage import get_storage
db = get_storage()
_create_template(db, "t1", "thread-tpl", "THREAD_TEMPLATE", is_default=False)
session = _make_session(template="thread-tpl")
errors: list[Exception] = []
stop = threading.Event()
iterations = 200
def init_loop():
"""Simulate MCP callback repeatedly calling _init_system_messages."""
try:
for _ in range(iterations):
if stop.is_set():
break
session._init_system_messages()
# system_messages must always be a valid list
msgs = session.system_messages
assert isinstance(msgs, list)
assert len(msgs) > 0
except Exception as exc:
errors.append(exc)
t = threading.Thread(target=init_loop, daemon=True)
t.start()
# Main thread toggles template on/off
try:
for i in range(iterations):
if i % 2 == 0:
session.set_template("thread-tpl")
else:
session.set_template(None)
finally:
stop.set()
t.join(timeout=5)
assert not errors, f"Thread raised: {errors}"
# Final state: system_messages is a valid list
msgs = session.system_messages
assert isinstance(msgs, list)
assert len(msgs) > 0
+43
View File
@@ -338,6 +338,28 @@ class TestPlanExec:
# Last user message in second call is the coaching message
assert "did not follow" in captured_messages[1][-1]["content"]
def test_plan_includes_template_content(self, tmp_db, tmp_path, monkeypatch):
"""Plan agent system message includes template guardrails."""
monkeypatch.chdir(tmp_path)
session = _make_session()
session._template_content = "SAFETY: Do not produce harmful plans."
_, _, messages = self._run_plan(session, "build something")
sys_content = messages[0]["content"]
assert "SAFETY: Do not produce harmful plans." in sys_content
assert ChatSession._PLAN_IDENTITY in sys_content
# Template appears before plan identity
tpl_pos = sys_content.index("SAFETY:")
identity_pos = sys_content.index(ChatSession._PLAN_IDENTITY)
assert tpl_pos < identity_pos
def test_plan_no_template_is_identity_only(self, tmp_db, tmp_path, monkeypatch):
"""Without templates, plan system message is exactly _PLAN_IDENTITY."""
monkeypatch.chdir(tmp_path)
session = _make_session()
assert session._template_content is None
_, _, messages = self._run_plan(session, "build something")
assert messages[0]["content"] == ChatSession._PLAN_IDENTITY
# ---------------------------------------------------------------------------
# Plan validation
@@ -557,6 +579,27 @@ class TestPlanRefinement:
assert msgs[3]["role"] == "user"
assert "add tests too" in msgs[3]["content"]
def test_refine_plan_includes_template_content(self, tmp_db, tmp_path, monkeypatch):
"""_refine_plan system message includes template guardrails."""
monkeypatch.chdir(tmp_path)
session = _make_session()
session._template_content = "SAFETY: guardrails here"
captured = {}
def fake_run_agent(messages, **kwargs):
captured["messages"] = list(messages)
return self.GOOD_PLAN
with patch.object(session, "_run_agent", side_effect=fake_run_agent):
session._refine_plan(self.GOOD_PLAN, "add auth", "add tests too")
sys_content = captured["messages"][0]["content"]
assert "SAFETY: guardrails here" in sys_content
assert ChatSession._PLAN_IDENTITY in sys_content
tpl_pos = sys_content.index("SAFETY:")
identity_pos = sys_content.index(ChatSession._PLAN_IDENTITY)
assert tpl_pos < identity_pos
# ---------------------------------------------------------------------------
# Vision / image support
+155
View File
@@ -0,0 +1,155 @@
"""Tests for WebUI content accumulation — server-side single source of truth."""
import queue
import pytest
from turnstone.server import WebUI
@pytest.fixture(autouse=True)
def _reset_global_queue():
"""Ensure WebUI._global_queue is set for tests and cleaned up after."""
WebUI._global_queue = queue.Queue()
yield
WebUI._global_queue = None
def _make_ui() -> WebUI:
"""Create a WebUI with a global queue for capturing broadcast events."""
return WebUI(ws_id="ws-test")
def _drain_global() -> list[dict]:
"""Drain all events from the global queue."""
events = []
assert WebUI._global_queue is not None
while not WebUI._global_queue.empty():
events.append(WebUI._global_queue.get_nowait())
return events
class TestContentAccumulation:
"""WebUI should accumulate content tokens and include in idle broadcast."""
def test_content_token_accumulates(self):
"""on_content_token should append to _ws_turn_content."""
ui = _make_ui()
ui.on_content_token("Hello ")
ui.on_content_token("world")
assert ui._ws_turn_content == ["Hello ", "world"]
def test_idle_broadcast_includes_content(self):
"""_broadcast_state('idle') should include joined content and reset."""
ui = _make_ui()
ui.on_content_token("Hello ")
ui.on_content_token("world")
ui._broadcast_state("idle")
events = _drain_global()
idle_events = [e for e in events if e.get("state") == "idle"]
assert len(idle_events) == 1
assert idle_events[0]["content"] == "Hello world"
# Accumulator should be reset
assert ui._ws_turn_content == []
assert ui._ws_turn_content_size == 0
def test_error_broadcast_resets_without_content(self):
"""_broadcast_state('error') should reset accumulator without content in event."""
ui = _make_ui()
ui.on_content_token("partial")
ui._broadcast_state("error")
events = _drain_global()
error_events = [e for e in events if e.get("state") == "error"]
assert len(error_events) == 1
assert "content" not in error_events[0]
assert ui._ws_turn_content == []
assert ui._ws_turn_content_size == 0
def test_thinking_broadcast_does_not_touch_accumulator(self):
"""_broadcast_state('thinking') should not affect the accumulator."""
ui = _make_ui()
ui.on_content_token("in progress")
ui._broadcast_state("thinking")
assert ui._ws_turn_content == ["in progress"]
events = _drain_global()
thinking_events = [e for e in events if e.get("state") == "thinking"]
assert len(thinking_events) == 1
assert "content" not in thinking_events[0]
def test_multi_round_accumulation(self):
"""Content from multiple streaming rounds accumulates before idle."""
ui = _make_ui()
# Round 1
ui.on_content_token("I'll check ")
ui.on_content_token("that. ")
# Round 2 (after tool execution)
ui.on_content_token("Here's ")
ui.on_content_token("the result.")
ui._broadcast_state("idle")
events = _drain_global()
idle_events = [e for e in events if e.get("state") == "idle"]
assert len(idle_events) == 1
assert idle_events[0]["content"] == "I'll check that. Here's the result."
def test_empty_content_on_idle_without_tokens(self):
"""idle with no content tokens should include empty content string."""
ui = _make_ui()
ui._broadcast_state("idle")
events = _drain_global()
idle_events = [e for e in events if e.get("state") == "idle"]
assert len(idle_events) == 1
assert idle_events[0]["content"] == ""
def test_cancellation_preserves_partial_content(self):
"""Partial content accumulated before cancel should appear in idle event."""
ui = _make_ui()
ui.on_content_token("I'll ")
ui.on_content_token("start by...")
# Cancellation triggers idle broadcast with partial content
ui._broadcast_state("idle")
events = _drain_global()
idle_events = [e for e in events if e.get("state") == "idle"]
assert len(idle_events) == 1
assert idle_events[0]["content"] == "I'll start by..."
def test_consecutive_turns_isolated(self):
"""Content from turn 1 should not leak into turn 2."""
ui = _make_ui()
# Turn 1
ui.on_content_token("first response")
ui._broadcast_state("idle")
_drain_global()
# Turn 2
ui.on_content_token("second response")
ui._broadcast_state("idle")
events = _drain_global()
idle_events = [e for e in events if e.get("state") == "idle"]
assert len(idle_events) == 1
assert idle_events[0]["content"] == "second response"
def test_content_cap_prevents_unbounded_growth(self):
"""Content exceeding the cap should stop accumulating."""
from turnstone.server import _MAX_TURN_CONTENT_CHARS
ui = _make_ui()
# Fill to capacity
chunk = "x" * 1024
for _ in range(_MAX_TURN_CONTENT_CHARS // 1024 + 10):
ui.on_content_token(chunk)
assert ui._ws_turn_content_size <= _MAX_TURN_CONTENT_CHARS + 1024
ui._broadcast_state("idle")
events = _drain_global()
idle_events = [e for e in events if e.get("state") == "idle"]
assert len(idle_events) == 1
# Content should be capped, not contain everything
assert len(idle_events[0]["content"]) <= _MAX_TURN_CONTENT_CHARS + 1024
+1 -1
View File
@@ -20,7 +20,7 @@ class FakeSession:
self.messages = []
def _fake_factory(ui, model_alias=None, ws_id=None):
def _fake_factory(ui, model_alias=None, ws_id=None, **kwargs):
return FakeSession()
+375
View File
@@ -333,6 +333,381 @@ def test_webui_auto_approve_tools_empty_no_effect():
webui._approval_event.wait.assert_called_once()
# ---------------------------------------------------------------------------
# Per-tool "always approve" — interactive "Always" adds to auto_approve_tools
# ---------------------------------------------------------------------------
def test_server_always_approve_adds_tool_names():
"""POST /approve with always=True adds pending tool names to auto_approve_tools."""
webui = WebUI(ws_id="ws-1")
webui._pending_approval = {
"type": "approve_request",
"items": [
{"func_name": "bash", "needs_approval": True, "preview": "ls"},
{"func_name": "read_file", "needs_approval": False, "preview": "/tmp"},
],
}
items = webui._pending_approval.get("items", [])
tool_names = {
it.get("approval_label", "") or it.get("func_name", "")
for it in items
if it.get("needs_approval") and it.get("func_name")
}
tool_names.discard("")
tool_names.discard("__budget_override__")
webui.auto_approve_tools.update(tool_names)
assert webui.auto_approve_tools == {"bash"}
assert webui.auto_approve is False # blanket flag NOT set
def test_server_always_approve_uses_approval_label():
"""When approval_label differs from func_name, approval_label is stored."""
webui = WebUI(ws_id="ws-1")
webui._pending_approval = {
"type": "approve_request",
"items": [
{
"func_name": "use_prompt",
"approval_label": "mcp__git__commit_msg",
"needs_approval": True,
"preview": "",
},
],
}
items = webui._pending_approval.get("items", [])
tool_names = {
it.get("approval_label", "") or it.get("func_name", "")
for it in items
if it.get("needs_approval") and it.get("func_name")
}
tool_names.discard("")
tool_names.discard("__budget_override__")
webui.auto_approve_tools.update(tool_names)
assert "mcp__git__commit_msg" in webui.auto_approve_tools
assert "use_prompt" not in webui.auto_approve_tools
def test_server_always_approve_excludes_budget_override():
"""__budget_override__ should never be added to auto_approve_tools."""
webui = WebUI(ws_id="ws-1")
webui._pending_approval = {
"type": "approve_request",
"items": [
{"func_name": "__budget_override__", "needs_approval": True, "preview": ""},
{"func_name": "bash", "needs_approval": True, "preview": "ls"},
],
}
items = webui._pending_approval.get("items", [])
tool_names = {
it.get("approval_label", "") or it.get("func_name", "")
for it in items
if it.get("needs_approval") and it.get("func_name")
}
tool_names.discard("")
tool_names.discard("__budget_override__")
webui.auto_approve_tools.update(tool_names)
assert "__budget_override__" not in webui.auto_approve_tools
assert webui.auto_approve_tools == {"bash"}
def test_server_always_approve_accumulates():
"""Successive 'always' approvals accumulate tool names."""
webui = WebUI(ws_id="ws-1")
# First always-approve: bash
webui._pending_approval = {
"type": "approve_request",
"items": [{"func_name": "bash", "needs_approval": True, "preview": "ls"}],
}
items = webui._pending_approval["items"]
names = {
it.get("approval_label", "") or it["func_name"] for it in items if it.get("needs_approval")
}
names.discard("__budget_override__")
webui.auto_approve_tools.update(names)
# Second always-approve: write_file
webui._pending_approval = {
"type": "approve_request",
"items": [{"func_name": "write_file", "needs_approval": True, "preview": ""}],
}
items = webui._pending_approval["items"]
names = {
it.get("approval_label", "") or it["func_name"] for it in items if it.get("needs_approval")
}
names.discard("__budget_override__")
webui.auto_approve_tools.update(names)
assert webui.auto_approve_tools == {"bash", "write_file"}
def test_server_always_approve_no_pending_is_noop():
"""If _pending_approval is None, always=True does nothing."""
webui = WebUI(ws_id="ws-1")
webui._pending_approval = None
# The guard `if always and approved and ui._pending_approval:` prevents action
assert webui.auto_approve_tools == set()
assert webui.auto_approve is False
# ---------------------------------------------------------------------------
# CLI per-tool "always approve"
# ---------------------------------------------------------------------------
def test_cli_always_adds_tool_names():
"""CLI 'a' adds pending tool names to auto_approve_tools, not blanket flag."""
from turnstone.cli import TerminalUI
ui = TerminalUI()
items = [
{"func_name": "bash", "header": "bash: ls", "needs_approval": True, "preview": "ls"},
]
with patch("builtins.input", return_value="a"):
approved, _ = ui.approve_tools(items)
assert approved is True
assert ui.auto_approve is False
assert "bash" in ui.auto_approve_tools
def test_cli_per_tool_auto_approves_subsequent():
"""After 'always' for bash, subsequent bash calls auto-approve silently."""
from turnstone.cli import TerminalUI
ui = TerminalUI()
ui.auto_approve_tools = {"bash"}
items = [
{"func_name": "bash", "header": "bash: ls", "needs_approval": True, "preview": "ls"},
]
# Should auto-approve without prompting
approved, _ = ui.approve_tools(items)
assert approved is True
def test_cli_per_tool_does_not_approve_unknown():
"""Per-tool set for bash does NOT auto-approve write_file."""
from turnstone.cli import TerminalUI
ui = TerminalUI()
ui.auto_approve_tools = {"bash"}
items = [
{
"func_name": "write_file",
"header": "write_file: /tmp/x",
"needs_approval": True,
"preview": "",
},
]
with patch("builtins.input", return_value="n"):
approved, _ = ui.approve_tools(items)
assert approved is False
def test_cli_always_excludes_budget_override():
"""CLI 'always' should not add __budget_override__ to auto_approve_tools."""
from turnstone.cli import TerminalUI
ui = TerminalUI()
items = [
{
"func_name": "__budget_override__",
"header": "budget",
"needs_approval": True,
"preview": "",
},
]
with patch("builtins.input", return_value="a"):
approved, _ = ui.approve_tools(items)
assert approved is True
assert "__budget_override__" not in ui.auto_approve_tools
# ---------------------------------------------------------------------------
# Bridge per-tool "always approve"
# ---------------------------------------------------------------------------
def test_bridge_always_adds_to_approve_tools():
"""Bridge 'always' adds tool names to _ws_approve_tools, not _ws_auto_approve."""
import threading
from turnstone.mq.bridge import DEFAULT_SAFE_TOOLS, Bridge
bridge = Bridge.__new__(Bridge)
bridge._lock = threading.Lock()
bridge._ws_auto_approve = {}
bridge._ws_approve_tools = {}
ws_id = "ws-1"
items = [
{"func_name": "bash", "needs_approval": True},
{"func_name": "read_file", "needs_approval": False},
]
# Simulate the always-approve extraction logic from _wait_approval
tool_names = {
it.get("func_name", "") for it in items if it.get("needs_approval") and it.get("func_name")
}
tool_names.discard("")
tool_names.discard("__budget_override__")
if tool_names:
with bridge._lock:
existing = bridge._ws_approve_tools.get(ws_id, set(DEFAULT_SAFE_TOOLS))
bridge._ws_approve_tools[ws_id] = existing | tool_names
# bash added, and DEFAULT_SAFE_TOOLS preserved
assert "bash" in bridge._ws_approve_tools[ws_id]
for name in DEFAULT_SAFE_TOOLS:
assert name in bridge._ws_approve_tools[ws_id]
assert ws_id not in bridge._ws_auto_approve
# ---------------------------------------------------------------------------
# Integration tests — POST /v1/api/approve with always=True
# ---------------------------------------------------------------------------
class TestApproveEndpointAlways:
"""Integration tests for the approve handler's per-tool 'always' logic."""
@staticmethod
def _make_client(webui):
import queue
import threading
from starlette.testclient import TestClient
from turnstone.core.auth import AuthConfig
from turnstone.server import create_app
mock_ws = MagicMock()
mock_ws.ui = webui
mock_mgr = MagicMock()
mock_mgr.get.return_value = mock_ws
mock_mgr.list_all.return_value = []
app = create_app(
workstreams=mock_mgr,
global_queue=queue.Queue(),
global_listeners=[],
global_listeners_lock=threading.Lock(),
skip_permissions=False,
auth_config=AuthConfig(),
)
return TestClient(app, raise_server_exceptions=False)
def test_always_adds_tool_to_auto_approve_tools(self):
webui = WebUI(ws_id="ws-1")
webui._pending_approval = {
"type": "approve_request",
"items": [
{"func_name": "bash", "needs_approval": True, "preview": "ls"},
],
}
client = self._make_client(webui)
resp = client.post(
"/v1/api/approve",
json={"approved": True, "always": True, "ws_id": "ws-1"},
)
assert resp.status_code == 200
assert "bash" in webui.auto_approve_tools
assert webui.auto_approve is False
def test_always_uses_approval_label_over_func_name(self):
webui = WebUI(ws_id="ws-1")
webui._pending_approval = {
"type": "approve_request",
"items": [
{
"func_name": "use_prompt",
"approval_label": "mcp__git__commit_msg",
"needs_approval": True,
"preview": "",
},
],
}
client = self._make_client(webui)
resp = client.post(
"/v1/api/approve",
json={"approved": True, "always": True, "ws_id": "ws-1"},
)
assert resp.status_code == 200
assert "mcp__git__commit_msg" in webui.auto_approve_tools
assert "use_prompt" not in webui.auto_approve_tools
def test_always_excludes_budget_override(self):
webui = WebUI(ws_id="ws-1")
webui._pending_approval = {
"type": "approve_request",
"items": [
{"func_name": "__budget_override__", "needs_approval": True, "preview": ""},
{"func_name": "bash", "needs_approval": True, "preview": "ls"},
],
}
client = self._make_client(webui)
resp = client.post(
"/v1/api/approve",
json={"approved": True, "always": True, "ws_id": "ws-1"},
)
assert resp.status_code == 200
assert "__budget_override__" not in webui.auto_approve_tools
assert "bash" in webui.auto_approve_tools
def test_always_skips_non_pending_items(self):
webui = WebUI(ws_id="ws-1")
webui._pending_approval = {
"type": "approve_request",
"items": [
{"func_name": "bash", "needs_approval": True, "preview": "ls"},
{"func_name": "read_file", "needs_approval": False, "preview": "/tmp"},
],
}
client = self._make_client(webui)
resp = client.post(
"/v1/api/approve",
json={"approved": True, "always": True, "ws_id": "ws-1"},
)
assert resp.status_code == 200
assert webui.auto_approve_tools == {"bash"}
def test_always_false_does_not_add_tools(self):
webui = WebUI(ws_id="ws-1")
webui._pending_approval = {
"type": "approve_request",
"items": [
{"func_name": "bash", "needs_approval": True, "preview": "ls"},
],
}
client = self._make_client(webui)
resp = client.post(
"/v1/api/approve",
json={"approved": True, "always": False, "ws_id": "ws-1"},
)
assert resp.status_code == 200
assert webui.auto_approve_tools == set()
def test_deny_with_always_does_not_add_tools(self):
webui = WebUI(ws_id="ws-1")
webui._pending_approval = {
"type": "approve_request",
"items": [
{"func_name": "bash", "needs_approval": True, "preview": "ls"},
],
}
client = self._make_client(webui)
resp = client.post(
"/v1/api/approve",
json={"approved": False, "always": True, "ws_id": "ws-1"},
)
assert resp.status_code == 200
assert webui.auto_approve_tools == set()
# ---------------------------------------------------------------------------
# Protocol round-trip — CreateWorkstreamMessage
# ---------------------------------------------------------------------------
+1 -1
View File
@@ -1,3 +1,3 @@
"""turnstone - Multi-node AI orchestration platform with tool use, agent routing, and cluster simulation."""
__version__ = "0.6.1"
__version__ = "0.7.0"
+80
View File
@@ -570,3 +570,83 @@ class ListSettingSchemaResponse(BaseModel):
class UpdateSettingRequest(BaseModel):
value: Any
node_id: str = ""
# ---------------------------------------------------------------------------
# Admin: MCP Servers
# ---------------------------------------------------------------------------
class McpServerInfo(BaseModel):
server_id: str
name: str
transport: str
command: str = ""
args: str = "[]"
url: str = ""
headers: str = "{}"
env: str = "{}"
auto_approve: bool = False
enabled: bool = True
created_by: str = ""
created: str
updated: str
class McpServerStatus(BaseModel):
connected: bool = False
tools: int = 0
resources: int = 0
prompts: int = 0
error: str = ""
transport: str = ""
command: str = ""
url: str = ""
class McpServerDetail(McpServerInfo):
status: dict[str, McpServerStatus] = Field(default_factory=dict)
source: str = "" # "config" for config-file servers, empty for DB-managed
class CreateMcpServerRequest(BaseModel):
name: str
transport: str # "stdio" | "streamable-http"
command: str = ""
args: list[str] = []
url: str = ""
headers: dict[str, str] = Field(default_factory=dict)
env: dict[str, str] = Field(default_factory=dict)
auto_approve: bool = False
enabled: bool = True
class UpdateMcpServerRequest(BaseModel):
name: str | None = None
transport: str | None = None
command: str | None = None
args: list[str] | None = None
url: str | None = None
headers: dict[str, str] | None = None
env: dict[str, str] | None = None
auto_approve: bool | None = None
enabled: bool | None = None
class ListMcpServersResponse(BaseModel):
servers: list[McpServerDetail]
class ImportMcpConfigRequest(BaseModel):
config: dict[str, Any] = Field(..., description="JSON config object with mcpServers key")
class ImportMcpConfigResponse(BaseModel):
imported: list[str] = []
skipped: list[str] = []
errors: list[str] = []
class McpReloadResponse(BaseModel):
status: str = "ok"
results: dict[str, Any] = Field(default_factory=dict)
+124
View File
@@ -20,13 +20,17 @@ from turnstone.api.console_schemas import (
ConsoleCreateWsResponse,
ConsoleHealthResponse,
CreateChannelUserRequest,
CreateMcpServerRequest,
CreatePromptTemplateRequest,
CreateRoleRequest,
CreateToolPolicyRequest,
CreateWsTemplateRequest,
ImportMcpConfigRequest,
ImportMcpConfigResponse,
ListAdminMemoriesResponse,
ListAuditEventsResponse,
ListChannelUsersResponse,
ListMcpServersResponse,
ListOrgsResponse,
ListPromptTemplatesResponse,
ListRolesResponse,
@@ -38,6 +42,8 @@ from turnstone.api.console_schemas import (
ListWsTemplatesResponse,
ListWsTemplateSummaryResponse,
ListWsTemplateVersionsResponse,
McpReloadResponse,
McpServerDetail,
NodeDetailResponse,
OrgInfo,
PromptTemplateInfo,
@@ -45,6 +51,7 @@ from turnstone.api.console_schemas import (
SettingInfo,
SettingSchemaInfo,
ToolPolicyInfo,
UpdateMcpServerRequest,
UpdateOrgRequest,
UpdatePromptTemplateRequest,
UpdateRoleRequest,
@@ -64,6 +71,7 @@ from turnstone.api.schemas import (
AuthSetupRequest,
AuthSetupResponse,
AuthStatusResponse,
AuthWhoamiResponse,
CreateScheduleRequest,
CreateTokenRequest,
CreateTokenResponse,
@@ -78,6 +86,7 @@ from turnstone.api.schemas import (
UpdateScheduleRequest,
UserInfo,
)
from turnstone.api.server_schemas import ListPromptTemplateSummaryResponse, PromptTemplateSummary
CONSOLE_ENDPOINTS: list[EndpointSpec] = [
# --- Cluster ---
@@ -190,6 +199,29 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [
response_model=StatusResponse,
tags=["Auth"],
),
EndpointSpec(
"/v1/api/auth/oidc/authorize",
"GET",
"Redirect to OIDC provider for SSO login",
response_code=302,
error_codes=[404, 503],
tags=["Auth"],
),
EndpointSpec(
"/v1/api/auth/oidc/callback",
"GET",
"OIDC callback — validates code, provisions user, sets JWT cookie, redirects to app",
response_code=302,
tags=["Auth"],
),
EndpointSpec(
"/v1/api/auth/whoami",
"GET",
"Return authenticated user info and permissions",
response_model=AuthWhoamiResponse,
error_codes=[401],
tags=["Auth"],
),
# --- Admin ---
EndpointSpec(
"/v1/api/admin/users",
@@ -262,6 +294,20 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [
error_codes=[404],
tags=["Admin"],
),
# --- OIDC Identities ---
EndpointSpec(
"/v1/api/admin/users/{user_id}/oidc-identities",
"GET",
"List OIDC identities linked to a user",
tags=["Admin"],
),
EndpointSpec(
"/v1/api/admin/oidc-identities",
"DELETE",
"Unlink an OIDC identity (issuer + subject as query params)",
error_codes=[400, 404],
tags=["Admin"],
),
# --- Schedules ---
EndpointSpec(
"/v1/api/admin/schedules",
@@ -525,6 +571,14 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [
response_model=ListWsTemplateSummaryResponse,
tags=["Workstreams"],
),
# --- Prompt templates ---
EndpointSpec(
"/v1/api/templates",
"GET",
"List available prompt templates (summary)",
response_model=ListPromptTemplateSummaryResponse,
tags=["Templates"],
),
# --- Governance: Usage & Audit ---
EndpointSpec(
"/v1/api/admin/usage",
@@ -661,6 +715,67 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [
error_codes=[400, 404],
tags=["Admin"],
),
# --- Admin: MCP Servers ---
EndpointSpec(
"/v1/api/admin/mcp-servers",
"GET",
"List MCP server definitions with live status",
response_model=ListMcpServersResponse,
query_params=[
QueryParam("reveal", "Show secret env/header values", schema_type="boolean"),
],
tags=["Admin"],
),
EndpointSpec(
"/v1/api/admin/mcp-servers",
"POST",
"Create an MCP server definition",
request_model=CreateMcpServerRequest,
response_model=McpServerDetail,
error_codes=[400, 409],
tags=["Admin"],
),
EndpointSpec(
"/v1/api/admin/mcp-servers/{server_id}",
"GET",
"Get a single MCP server with status",
response_model=McpServerDetail,
error_codes=[404],
tags=["Admin"],
),
EndpointSpec(
"/v1/api/admin/mcp-servers/{server_id}",
"PUT",
"Update an MCP server definition",
request_model=UpdateMcpServerRequest,
response_model=McpServerDetail,
error_codes=[400, 404, 409],
tags=["Admin"],
),
EndpointSpec(
"/v1/api/admin/mcp-servers/{server_id}",
"DELETE",
"Delete an MCP server definition",
response_model=StatusResponse,
error_codes=[404],
tags=["Admin"],
),
EndpointSpec(
"/v1/api/admin/mcp-servers/reload",
"POST",
"Tell all nodes to re-read MCP server config from DB and reconcile",
response_model=McpReloadResponse,
tags=["Admin"],
),
EndpointSpec(
"/v1/api/admin/mcp-servers/import",
"POST",
"Import MCP servers from a JSON config file",
request_model=ImportMcpConfigRequest,
response_model=ImportMcpConfigResponse,
error_codes=[400],
tags=["Admin"],
),
# --- Observability ---
EndpointSpec(
"/health",
@@ -732,6 +847,15 @@ _ALL_MODELS: list[type[BaseModel]] = [
SettingSchemaInfo,
ListSettingSchemaResponse,
UpdateSettingRequest,
McpServerDetail,
CreateMcpServerRequest,
UpdateMcpServerRequest,
ListMcpServersResponse,
ImportMcpConfigRequest,
ImportMcpConfigResponse,
McpReloadResponse,
PromptTemplateSummary,
ListPromptTemplateSummaryResponse,
]
+10
View File
@@ -155,6 +155,16 @@ class AuthStatusResponse(BaseModel):
auth_enabled: bool
has_users: bool
setup_required: bool
oidc_enabled: bool = False
oidc_provider_name: str = ""
password_enabled: bool = True
class AuthWhoamiResponse(BaseModel):
"""GET /v1/api/auth/whoami response."""
user_id: str
permissions: str = ""
# ---------------------------------------------------------------------------
+40 -2
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
from typing import Literal
from pydantic import BaseModel, Field
from pydantic import BaseModel, Field, model_validator
# ---------------------------------------------------------------------------
# Workstream management
@@ -23,7 +23,9 @@ class SendResponse(BaseModel):
class ApproveRequest(BaseModel):
approved: bool = Field(description="True to approve, false to deny")
feedback: str | None = Field(default=None, description="Optional denial reason")
always: bool = Field(default=False, description="Enable auto-approve for this tool")
always: bool = Field(
default=False, description="Auto-approve the tools in this batch going forward"
)
ws_id: str = Field(description="Target workstream ID")
@@ -186,6 +188,15 @@ class SaveMemoryRequest(BaseModel):
description="Scope identifier (ws_id for workstream, user_id for user scope)",
)
@model_validator(mode="after")
def _validate_scope_scope_id(self) -> SaveMemoryRequest:
scope_id = self.scope_id.strip()
if self.scope == "global" and scope_id:
raise ValueError("scope_id is not allowed with global scope")
if self.scope == "workstream" and not scope_id:
raise ValueError("scope_id is required for workstream scope")
return self
class MemoryInfo(BaseModel):
memory_id: str
@@ -214,3 +225,30 @@ class SearchMemoriesRequest(BaseModel):
scope: MemoryScopeFilter = Field(default="", description="Filter by scope")
scope_id: str = Field(default="", description="Filter by scope_id")
limit: int = Field(default=20, description="Max results (1-50)", ge=1, le=50)
@model_validator(mode="after")
def _validate_scope_scope_id(self) -> SearchMemoriesRequest:
scope_id = self.scope_id.strip()
if self.scope == "global" and scope_id:
raise ValueError("scope_id is not allowed with global scope")
if scope_id and not self.scope:
raise ValueError("scope is required when scope_id is provided")
return self
# ---------------------------------------------------------------------------
# Prompt templates (read-only listing)
# ---------------------------------------------------------------------------
class PromptTemplateSummary(BaseModel):
name: str = Field(description="Template name")
category: str = Field(default="", description="Template category")
is_default: bool = Field(
default=False, description="Whether this template is applied by default"
)
origin: str = Field(default="manual", description="Template origin: manual or mcp")
class ListPromptTemplateSummaryResponse(BaseModel):
templates: list[PromptTemplateSummary]
+47
View File
@@ -4,6 +4,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Any
from turnstone.api.console_schemas import ListWsTemplateSummaryResponse, WsTemplateSummary
from turnstone.api.openapi import EndpointSpec, QueryParam, build_openapi
if TYPE_CHECKING:
@@ -14,6 +15,7 @@ from turnstone.api.schemas import (
AuthSetupRequest,
AuthSetupResponse,
AuthStatusResponse,
AuthWhoamiResponse,
ErrorResponse,
StatusResponse,
)
@@ -27,10 +29,12 @@ from turnstone.api.server_schemas import (
DashboardResponse,
HealthResponse,
ListMemoriesResponse,
ListPromptTemplateSummaryResponse,
ListSavedWorkstreamsResponse,
ListWorkstreamsResponse,
MemoryInfo,
PlanFeedbackRequest,
PromptTemplateSummary,
SaveMemoryRequest,
SearchMemoriesRequest,
SendRequest,
@@ -144,6 +148,22 @@ SERVER_ENDPOINTS: list[EndpointSpec] = [
response_model=ListSavedWorkstreamsResponse,
tags=["Workstreams"],
),
# --- Prompt templates ---
EndpointSpec(
"/v1/api/templates",
"GET",
"List available prompt templates (summary)",
response_model=ListPromptTemplateSummaryResponse,
tags=["Templates"],
),
# --- Workstream templates ---
EndpointSpec(
"/v1/api/ws-templates",
"GET",
"List enabled workstream templates (summary)",
response_model=ListWsTemplateSummaryResponse,
tags=["Templates"],
),
# --- Auth ---
EndpointSpec(
"/v1/api/auth/login",
@@ -177,6 +197,29 @@ SERVER_ENDPOINTS: list[EndpointSpec] = [
response_model=StatusResponse,
tags=["Auth"],
),
EndpointSpec(
"/v1/api/auth/oidc/authorize",
"GET",
"Redirect to OIDC provider for SSO login",
response_code=302,
error_codes=[404, 503],
tags=["Auth"],
),
EndpointSpec(
"/v1/api/auth/oidc/callback",
"GET",
"OIDC callback — validates code, provisions user, sets JWT cookie, redirects to app",
response_code=302,
tags=["Auth"],
),
EndpointSpec(
"/v1/api/auth/whoami",
"GET",
"Return authenticated user info and permissions",
response_model=AuthWhoamiResponse,
error_codes=[401],
tags=["Auth"],
),
# --- Memories ---
EndpointSpec(
"/v1/api/memories",
@@ -257,6 +300,10 @@ _ALL_MODELS: list[type[BaseModel]] = [
MemoryInfo,
ListMemoriesResponse,
SearchMemoriesRequest,
PromptTemplateSummary,
ListPromptTemplateSummaryResponse,
WsTemplateSummary,
ListWsTemplateSummaryResponse,
]
+18 -1
View File
@@ -93,6 +93,16 @@ For commercial providers (OpenAI, Anthropic-via-proxy), use the real key.
- `TURNSTONE_JWT_SECRET` JWT signing secret (required if auth enabled)
- `TURNSTONE_AUTH_TOKEN` Static bearer token for inter-service auth
### OIDC SSO (optional)
- `TURNSTONE_OIDC_ISSUER` OIDC issuer URL (e.g., https://accounts.google.com). Setting this + CLIENT_ID + CLIENT_SECRET enables SSO.
- `TURNSTONE_OIDC_CLIENT_ID` Client ID from the identity provider
- `TURNSTONE_OIDC_CLIENT_SECRET` Client secret (confidential client)
- `TURNSTONE_OIDC_PROVIDER_NAME` Display name for the SSO button (default: "SSO")
- `TURNSTONE_OIDC_SCOPES` OIDC scopes (default: "openid email profile")
- `TURNSTONE_OIDC_ROLE_CLAIM` Claim name for role mapping (e.g., "groups")
- `TURNSTONE_OIDC_ROLE_MAP` Comma-separated claim_value:role_id pairs (e.g., "admin:builtin-admin,eng:builtin-operator")
- `TURNSTONE_OIDC_PASSWORD_ENABLED` Set to "false" to hide password login and force SSO-only
### Ports
- `SERVER_PORT` Server port (default: 8080)
- `CONSOLE_PORT` Console port (default: 8090)
@@ -122,6 +132,10 @@ This is a one-time endpoint that only works when zero users exist.
Subsequent governance setup (roles, policies, templates) uses the console admin API \
with the JWT returned from setup.
If OIDC is configured, users can also log in via the "Continue with [Provider]" button on the login page.
The first admin user must still be created via the setup wizard (OIDC login requires at least one user to exist).
OIDC users are auto-provisioned on first login with a default viewer role unless role mapping is configured.
## Runtime Settings (ConfigStore)
After the stack is running, ~40 runtime settings (model, temperature, max_tokens, \
reasoning_effort, tool timeout, rate limiting, health probes, judge config, memory \
@@ -157,7 +171,10 @@ Walk the user through setting up their deployment step by step:
PostgreSQL is required for cluster mode.
5. **Security**: Recommend enabling auth for any non-local deployment. \
Use `generate_secret` for JWT secret, Redis password, auth token, and Postgres password. \
Ask for initial admin username and password.
Ask for initial admin username and password. \
If the user's deployment will use an external identity provider (Okta, Azure AD, Google, etc.), \
offer to configure OIDC SSO. Ask for the issuer URL, client ID, and client secret. \
Optionally configure role mapping and OIDC-only mode.
6. **Ports**: Check defaults with `check_port`, suggest alternatives if conflicts.
7. **Optional features**: Discord integration, web search (Tavily key), \
DuckDuckGo Search MCP (for cluster uses `ddgCluster` profile with \
+11 -1
View File
@@ -8,6 +8,7 @@ from __future__ import annotations
import asyncio
import json
import re
import socket
import uuid
from typing import TYPE_CHECKING, Any
@@ -26,6 +27,9 @@ if TYPE_CHECKING:
log = get_logger(__name__)
# ws_id is a hex string (832 chars depending on entry point).
_WS_ID_RE = re.compile(r"^[0-9a-f]{8,32}$")
async def _handle_health(request: Request) -> JSONResponse:
return JSONResponse({"status": "ok", "service": "channel"})
@@ -81,6 +85,9 @@ async def _handle_notify(request: Request) -> JSONResponse:
target = body.get("target")
message = body.get("message", "").strip() if isinstance(body.get("message"), str) else ""
title = body.get("title", "").strip() if isinstance(body.get("title"), str) else ""
ws_id = body.get("ws_id", "").strip() if isinstance(body.get("ws_id"), str) else ""
if ws_id and not _WS_ID_RE.match(ws_id):
return JSONResponse({"error": "invalid ws_id format"}, status_code=400)
if not target or not message:
return JSONResponse({"error": "target and message are required"}, status_code=400)
@@ -132,7 +139,10 @@ async def _handle_notify(request: Request) -> JSONResponse:
)
continue
try:
msg_id = await adapter.send(channel_id, content)
if ws_id:
msg_id = await adapter.send_notification(channel_id, content, ws_id)
else:
msg_id = await adapter.send(channel_id, content)
results.append(
{
"channel_type": channel_type,
+8
View File
@@ -41,6 +41,14 @@ class ChannelAdapter(Protocol):
"""Send a message to a channel. Returns the platform message ID."""
...
async def send_notification(self, channel_id: str, content: str, ws_id: str) -> str:
"""Send a notification and track the reply mapping. Returns message ID.
Like :meth:`send` but associates the outgoing message with *ws_id*
so that replies can be routed back to the originating workstream.
"""
...
async def edit_message(self, channel_id: str, message_id: str, content: str) -> None:
"""Edit an existing message in a channel."""
...
+68
View File
@@ -124,6 +124,7 @@ class TurnstoneBot:
"""
channel_type: str = "discord"
_MAX_NOTIFY_TRACKING: int = 100
def __init__(
self,
@@ -151,6 +152,16 @@ class TurnstoneBot:
# workstream so that IntentVerdictEvent can update it with LLM judge
# results.
self._pending_approval_msgs: dict[str, discord.Message] = {}
# Notification reply tracking: maps Discord message ID →
# (ws_id, target_discord_user_id) so that DM replies can be routed
# back to the originating workstream. The target user ID is checked
# on reply to prevent cross-user message injection.
self._notify_ws_map: dict[int, tuple[str, str]] = {}
# Temporary DM forwarding: maps ws_id → (DM channel, target_user_id)
# for forwarding the workstream's next response back to the
# notification reply DM. The target_user_id is carried so the
# response message can be re-tracked for multi-turn DM conversations.
self._notify_reply_channels: dict[str, tuple[discord.abc.Messageable, str]] = {}
intents = discord.Intents.default()
intents.message_content = True
@@ -256,6 +267,11 @@ class TurnstoneBot:
self._subscribed_ws.discard(ws_id)
self._streaming.pop(ws_id, None)
self._pending_approval_msgs.pop(ws_id, None)
self._notify_reply_channels.pop(ws_id, None)
# Purge stale notification tracking entries for this workstream.
stale = [mid for mid, entry in self._notify_ws_map.items() if entry[0] == ws_id]
for mid in stale:
del self._notify_ws_map[mid]
log.info("discord.unsubscribed", ws_id=ws_id)
# -- event dispatch ------------------------------------------------------
@@ -360,6 +376,26 @@ class TurnstoneBot:
sm = self._streaming.pop(ws_id, None)
if sm is not None:
await sm.finalize()
elif event.content:
# Catch-up: content events were missed (race between global
# SSE and per-ws SSE) — send the full response directly.
for chunk in chunk_message(event.content, self.config.max_message_length):
await thread.send(chunk)
# Forward response to notification reply DM if active.
dm_entry = self._notify_reply_channels.pop(ws_id, None)
if dm_entry is not None and event.content:
dm_channel, target_user_id = dm_entry
last_msg: discord.Message | None = None
for chunk in chunk_message(event.content, self.config.max_message_length):
try:
last_msg = await dm_channel.send(chunk)
except Exception:
log.debug("discord.notify_reply_dm_failed", ws_id=ws_id)
break
# Track the response message so the user can reply again
# for multi-turn DM conversations.
if last_msg is not None:
self._track_notification(last_msg.id, ws_id, target_user_id)
# Clean up pending approval message tracking.
self._pending_approval_msgs.pop(ws_id, None)
@@ -390,6 +426,18 @@ class TurnstoneBot:
return False
return True
def _track_notification(self, message_id: int, ws_id: str, target_user_id: str) -> None:
"""Record a notification message for reply routing.
Evicts the oldest entry when the map exceeds
``_MAX_NOTIFY_TRACKING``. Relies on dict insertion order
(Python 3.7+).
"""
while len(self._notify_ws_map) >= self._MAX_NOTIFY_TRACKING:
oldest = next(iter(self._notify_ws_map))
del self._notify_ws_map[oldest]
self._notify_ws_map[message_id] = (ws_id, target_user_id)
def _is_allowed_channel(self, channel_id: int) -> bool:
"""Return True if *channel_id* is in the allowed list (or list is empty)."""
if not self.config.allowed_channels:
@@ -430,6 +478,26 @@ class TurnstoneBot:
return str(msg.id) if msg else ""
async def send_notification(self, channel_id: str, content: str, ws_id: str) -> str:
"""Send a notification DM and track the message for reply routing.
Like :meth:`send` but records a mapping from the outgoing Discord
message ID to ``(ws_id, channel_id)`` so that a user reply can be
routed back to the originating workstream. The *channel_id* is the
Discord user ID the notification was sent to verified on reply to
prevent cross-user message injection.
"""
msg_id_str = await self.send(channel_id, content)
if msg_id_str and ws_id:
self._track_notification(int(msg_id_str), ws_id, channel_id)
log.debug(
"discord.notification_tracked",
message_id=msg_id_str,
ws_id=ws_id,
target_user=channel_id,
)
return msg_id_str
async def stop(self) -> None:
"""Disconnect the bot and clean up subscriptions."""
for ws_id in list(self._subscribed_ws):
+60 -1
View File
@@ -21,6 +21,7 @@ if TYPE_CHECKING:
log = get_logger(__name__)
_THREAD_NAME_MAX = 100
_DM_REPLY_MAX_LENGTH = 4096 # Discord's own message limit
class MessageCog:
@@ -102,8 +103,9 @@ class MessageCog:
if message.author == self.bot.user or message.author.bot:
return
# Ignore DMs.
# DM handling — route replies to tracked notifications.
if message.guild is None:
await self._handle_dm(message)
return
channel = message.channel
@@ -198,6 +200,63 @@ class MessageCog:
author=str(message.author),
)
# -- DM reply handling ---------------------------------------------------
async def _handle_dm(self, message: discord.Message) -> None:
"""Route DM replies to tracked notification workstreams."""
# Only handle explicit replies to a tracked notification message.
ref = message.reference
if ref is None or ref.message_id is None:
return
# Atomic pop prevents TOCTOU race across await points.
entry = self.ts._notify_ws_map.pop(ref.message_id, None)
if entry is None:
# NOTE: This also fires for replies to non-notification bot
# messages in DMs (false positive). Acceptable because DM
# interactions are almost exclusively notification-driven.
await message.channel.send("*This notification is no longer active.*")
return
ws_id, target_user_id = entry
# Defence in depth: verify the replying user is the notification
# recipient. Discord enforces this (DMs are private), but a
# server-side check prevents cross-user injection via compromised
# accounts or API-level forgery.
if str(message.author.id) != target_user_id:
# Re-insert so the legitimate user can still reply.
self.ts._notify_ws_map[ref.message_id] = entry
log.warning(
"discord.notification_reply_user_mismatch",
expected=target_user_id,
actual=str(message.author.id),
)
return
# Resolve user identity — unlinked users are silently ignored.
# Re-insert the tracking entry so the user can retry after linking.
user_id = await self.ts.router.resolve_user("discord", str(message.author.id))
if user_id is None:
self.ts._notify_ws_map[ref.message_id] = entry
return
# Route the reply to the originating workstream.
content = message.content[:_DM_REPLY_MAX_LENGTH]
await self.ts.router.send_message(ws_id, content)
# Register the DM channel for response forwarding. The bot's
# _on_ws_event handler will send the next turn's response here,
# track the response for further replies, and clean up on
# TurnCompleteEvent.
self.ts._notify_reply_channels[ws_id] = (message.channel, target_user_id)
log.info(
"discord.notification_reply_routed",
ws_id=ws_id,
author=str(message.author),
)
# -- slash commands ------------------------------------------------------
async def _cmd_link(self, interaction: discord.Interaction, token: str) -> None:
+19 -3
View File
@@ -98,6 +98,7 @@ class TerminalUI(SessionUI):
self.spinner: Spinner | None = None
self._print_lock = threading.Lock()
self.auto_approve = False
self.auto_approve_tools: set[str] = set()
def on_thinking_start(self) -> None:
self.spinner = Spinner("Thinking")
@@ -159,6 +160,12 @@ class TerminalUI(SessionUI):
if not pending or self.auto_approve:
return True, None
# 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")}
if pending_names and pending_names.issubset(self.auto_approve_tools):
return True, None
# Prompt
try:
if len(pending) == 1:
@@ -188,7 +195,10 @@ class TerminalUI(SessionUI):
break
if decision in ("a", "always"):
self.auto_approve = True
tool_names = {it.get("func_name", "") for it in pending if it.get("func_name")}
tool_names.discard("")
tool_names.discard("__budget_override__")
self.auto_approve_tools.update(tool_names)
return True, feedback
elif decision in ("y", "yes"):
return True, feedback
@@ -1006,15 +1016,21 @@ def main() -> None:
# Initialize MCP client (connects to configured MCP servers, if any)
from turnstone.core.mcp_client import create_mcp_client
from turnstone.core.storage._registry import get_storage as _get_storage
mcp_client = create_mcp_client(
getattr(args, "mcp_config", None),
refresh_interval=getattr(args, "mcp_refresh_interval", 14400),
storage=_get_storage(),
)
# ChatSession factory — captures shared config for creating workstreams
def session_factory(
ui: SessionUI | None, model_alias: str | None = None, ws_id: str | None = None
ui: SessionUI | None,
model_alias: str | None = None,
ws_id: str | None = None,
*,
template: str | None = None,
) -> ChatSession:
assert ui is not None, "session_factory requires a non-None UI"
r_client, r_model, r_cfg = registry.resolve(model_alias)
@@ -1038,7 +1054,7 @@ def main() -> None:
tool_search=args.tool_search,
tool_search_threshold=args.tool_search_threshold,
tool_search_max_results=args.tool_search_max_results,
template=args.template,
template=template if template is not None else args.template,
)
# Create workstream manager and initial workstream
+759 -12
View File
@@ -323,6 +323,27 @@ async def auth_setup(request: Request) -> Response:
return await handle_auth_setup(request, JWT_AUD_CONSOLE)
async def auth_whoami(request: Request) -> Response:
"""GET /v1/api/auth/whoami — return authenticated user info."""
from turnstone.core.auth import handle_auth_whoami
return await handle_auth_whoami(request)
async def oidc_authorize(request: Request) -> Response:
"""GET /v1/api/auth/oidc/authorize — redirect to OIDC provider."""
from turnstone.core.auth import handle_oidc_authorize
return await handle_oidc_authorize(request, JWT_AUD_CONSOLE)
async def oidc_callback(request: Request) -> Response:
"""GET /v1/api/auth/oidc/callback — OIDC callback, exchange code for JWT."""
from turnstone.core.auth import handle_oidc_callback
return await handle_oidc_callback(request, JWT_AUD_CONSOLE)
# ---------------------------------------------------------------------------
# Route handlers — workstream creation
# ---------------------------------------------------------------------------
@@ -674,6 +695,31 @@ async def _lifespan(app: Starlette) -> AsyncGenerator[None, None]:
scheduler = getattr(app.state, "scheduler", None)
if scheduler is not None:
scheduler.start()
# OIDC discovery (if configured)
oidc_config = app.state.oidc_config
if oidc_config.enabled:
from turnstone.core.oidc import discover_oidc
try:
oidc_config = await discover_oidc(oidc_config)
app.state.oidc_config = oidc_config
except Exception:
log.warning("OIDC discovery failed — OIDC login disabled", exc_info=True)
if oidc_config.enabled and oidc_config.jwks_uri:
try:
from turnstone.core.oidc import fetch_jwks
app.state.jwks_data = await fetch_jwks(oidc_config.jwks_uri)
log.info(
"OIDC enabled: %s (%s)",
oidc_config.provider_name,
oidc_config.issuer,
)
except Exception:
log.warning(
"OIDC JWKS prefetch failed — will retry on first login",
exc_info=True,
)
yield
# Shutdown
if scheduler is not None:
@@ -1052,6 +1098,67 @@ async def admin_delete_channel(request: Request) -> JSONResponse:
return JSONResponse({"error": "Channel link not found"}, status_code=404)
# ---------------------------------------------------------------------------
# Admin API endpoints — OIDC identities
# ---------------------------------------------------------------------------
async def admin_list_oidc_identities(request: Request) -> JSONResponse:
"""GET /v1/api/admin/users/{user_id}/oidc-identities — list OIDC links for a user."""
from turnstone.core.auth import require_permission
from turnstone.core.web_helpers import require_storage_or_503
storage, err = require_storage_or_503(request)
if err:
return err
err = require_permission(request, "admin.users")
if err:
return err
user_id = request.path_params["user_id"]
identities = storage.list_oidc_identities_for_user(user_id)
return JSONResponse({"oidc_identities": identities})
async def admin_delete_oidc_identity(request: Request) -> JSONResponse:
"""DELETE /v1/api/admin/oidc-identities?issuer=...&subject=... — unlink OIDC identity."""
from turnstone.core.audit import record_audit
from turnstone.core.auth import require_permission
from turnstone.core.web_helpers import require_storage_or_503
storage, err = require_storage_or_503(request)
if err:
return err
err = require_permission(request, "admin.users")
if err:
return err
issuer = request.query_params.get("issuer", "")
subject = request.query_params.get("subject", "")
if not issuer or not subject:
return JSONResponse({"error": "issuer and subject required"}, status_code=400)
# Look up before delete so audit captures which user was affected
identity = storage.get_oidc_identity(issuer, subject)
if not identity:
return JSONResponse({"error": "Identity not found"}, status_code=404)
storage.delete_oidc_identity(issuer, subject)
audit_uid, ip = _audit_context(request)
record_audit(
storage,
audit_uid,
"oidc_identity.delete",
"oidc_identity",
f"{issuer}:{subject}",
{"user_id": identity["user_id"]},
ip,
)
return JSONResponse({"status": "ok"})
# ---------------------------------------------------------------------------
# Admin API endpoints — scheduled tasks
# ---------------------------------------------------------------------------
@@ -1385,23 +1492,29 @@ async def admin_list_watches(request: Request) -> JSONResponse:
nodes, _ = collector.get_nodes(limit=500)
client: httpx.AsyncClient = request.app.state.proxy_client
headers = _proxy_auth_headers(request)
sem = asyncio.Semaphore(_NODE_FAN_OUT_LIMIT)
async def _fetch_node(node: dict[str, Any]) -> list[dict[str, Any]]:
server_url = (node.get("server_url") or "").rstrip("/")
if not server_url:
return []
try:
resp = await client.get(f"{server_url}/v1/api/watches", headers=headers)
if resp.status_code == 200:
data = resp.json()
watches: list[dict[str, Any]] = data.get("watches", [])
# Tag each watch with node_id in case the server omits it
for w in watches:
if not w.get("node_id"):
w["node_id"] = node["node_id"]
return watches
except Exception:
log.debug("Failed to fetch watches from node %s", node.get("node_id"))
async with sem:
try:
resp = await client.get(f"{server_url}/v1/api/watches", headers=headers)
if resp.status_code == 200:
data = resp.json()
watches: list[dict[str, Any]] = data.get("watches", [])
# Tag each watch with node_id in case the server omits it
for w in watches:
if not w.get("node_id"):
w["node_id"] = node["node_id"]
return watches
except Exception:
log.debug(
"Failed to fetch watches from node %s",
node.get("node_id"),
exc_info=True,
)
return []
tasks = [_fetch_node(n) for n in nodes]
@@ -1417,6 +1530,11 @@ 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
async def admin_cancel_watch(request: Request) -> Response:
"""POST /v1/api/admin/watches/{watch_id}/cancel — proxy cancel to the owning node."""
@@ -1513,6 +1631,7 @@ _VALID_PERMISSIONS = frozenset(
"admin.judge",
"admin.memories",
"admin.settings",
"admin.mcp",
"tools.approve",
"workstreams.create",
"workstreams.close",
@@ -2479,6 +2598,26 @@ async def list_ws_templates_summary(request: Request) -> JSONResponse:
return JSONResponse({"ws_templates": summary})
async def list_templates_summary(request: Request) -> JSONResponse:
"""GET /v1/api/templates — list available prompt templates (read scope)."""
from turnstone.core.web_helpers import require_storage_or_503
storage, err = require_storage_or_503(request)
if err:
return err
templates = storage.list_prompt_templates()
summaries = [
{
"name": t["name"],
"category": t.get("category", ""),
"is_default": bool(t.get("is_default")),
"origin": t.get("origin", "manual"),
}
for t in templates
]
return JSONResponse({"templates": summaries})
async def admin_usage(request: Request) -> JSONResponse:
"""GET /v1/api/admin/usage — query usage data."""
from datetime import UTC, datetime, timedelta
@@ -2614,6 +2753,19 @@ async def admin_list_verdicts(request: Request) -> JSONResponse:
# ---------------------------------------------------------------------------
def _validate_memory_scope_filter(scope: str, scope_id: str) -> JSONResponse | None:
"""Validate scope/scope_id consistency for memory queries."""
scope = scope.strip()
scope_id = scope_id.strip()
if scope == "global" and scope_id:
return JSONResponse({"error": "scope_id is not allowed with global scope"}, status_code=400)
if scope_id and not scope:
return JSONResponse(
{"error": "scope is required when scope_id is provided"}, status_code=400
)
return None
async def admin_list_memories(request: Request) -> JSONResponse:
"""GET /v1/api/admin/memories — list structured memories with filters."""
from turnstone.core.auth import require_permission
@@ -2629,6 +2781,9 @@ async def admin_list_memories(request: Request) -> JSONResponse:
mem_type = request.query_params.get("type", "")
scope = request.query_params.get("scope", "")
scope_id = request.query_params.get("scope_id", "")
err = _validate_memory_scope_filter(scope, scope_id)
if err:
return err
try:
limit = min(int(request.query_params.get("limit", "100")), 200)
except (ValueError, TypeError):
@@ -2659,6 +2814,9 @@ async def admin_search_memories(request: Request) -> JSONResponse:
mem_type = request.query_params.get("type", "")
scope = request.query_params.get("scope", "")
scope_id = request.query_params.get("scope_id", "")
err = _validate_memory_scope_filter(scope, scope_id)
if err:
return err
try:
limit = min(int(request.query_params.get("limit", "20")), 50)
except (ValueError, TypeError):
@@ -2968,6 +3126,544 @@ async def admin_delete_setting(request: Request) -> JSONResponse:
return JSONResponse({"status": "ok", "key": key})
# ---------------------------------------------------------------------------
# Admin: MCP Servers
# ---------------------------------------------------------------------------
_MCP_NAME_RE = re.compile(r"^[a-zA-Z0-9._-]+$")
_MCP_MAX_SERVERS = 50
def _mask_mcp_secrets(server: dict[str, Any], reveal: bool = False) -> dict[str, Any]:
"""Replace env/headers values with '***' unless reveal is True."""
if reveal:
return server
s = dict(server)
if s.get("env") and s["env"] != "{}":
try:
env_dict = json.loads(s["env"]) if isinstance(s["env"], str) else s["env"]
s["env"] = json.dumps({k: "***" for k in env_dict})
except (json.JSONDecodeError, TypeError):
s["env"] = "{}"
if s.get("headers") and s["headers"] != "{}":
try:
hdr_dict = json.loads(s["headers"]) if isinstance(s["headers"], str) else s["headers"]
s["headers"] = json.dumps({k: "***" for k in hdr_dict})
except (json.JSONDecodeError, TypeError):
s["headers"] = "{}"
return s
def _mcp_server_to_detail(
server: dict[str, Any],
node_statuses: dict[str, dict[str, Any]] | None = None,
) -> dict[str, Any]:
"""Convert a storage dict to a McpServerDetail-shaped dict."""
d = dict(server)
d["status"] = node_statuses or {}
return d
async def _collect_mcp_status(
request: Request,
) -> 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)
client: httpx.AsyncClient = request.app.state.proxy_client
headers = _proxy_auth_headers(request)
sem = asyncio.Semaphore(_NODE_FAN_OUT_LIMIT)
async def _fetch(node: dict[str, Any]) -> tuple[str, dict[str, dict[str, Any]] | None]:
node_id = node.get("node_id", "")
url = node.get("server_url", "")
if not url:
return node_id, None
async with sem:
try:
resp = await client.get(
f"{url.rstrip('/')}/v1/api/_internal/mcp-status",
headers=headers,
timeout=10,
)
if resp.status_code == 200:
return node_id, resp.json().get("servers", {})
except Exception:
log.debug("Failed to fetch MCP status from node %s", node_id, exc_info=True)
return node_id, None
results = await asyncio.gather(*[_fetch(n) for n in nodes])
return {nid: servers for nid, servers in results if servers is not None}
async def admin_list_mcp_servers(request: Request) -> JSONResponse:
"""GET /v1/api/admin/mcp-servers — list all MCP server definitions."""
from turnstone.core.auth import require_permission
from turnstone.core.web_helpers import require_storage_or_503
storage, err = require_storage_or_503(request)
if err:
return err
err = require_permission(request, "admin.mcp")
if err:
return err
reveal = str(request.query_params.get("reveal", "")).lower() in ("true", "1")
servers = storage.list_mcp_servers()
# Collect live status from all nodes
node_statuses = await _collect_mcp_status(request)
db_names: set[str] = set()
result = []
for s in servers:
db_names.add(s["name"])
# Build per-node status for this server
per_node: dict[str, dict[str, Any]] = {}
for node_id, node_servers in node_statuses.items():
status = node_servers.get(s["name"])
if status:
per_node[node_id] = status
s = _mask_mcp_secrets(s, reveal)
result.append(_mcp_server_to_detail(s, per_node))
# Merge config-sourced servers visible on nodes but not in DB
config_names: set[str] = set()
for node_servers in node_statuses.values():
for name in node_servers:
if name not in db_names:
config_names.add(name)
for name in sorted(config_names):
# Build a synthetic read-only entry from node-reported data
per_node = {}
transport = "stdio"
command = ""
url = ""
for node_id, node_servers in node_statuses.items():
ns = node_servers.get(name)
if ns:
per_node[node_id] = ns
transport = ns.get("transport", "stdio")
command = ns.get("command", "")
url = ns.get("url", "")
result.append(
{
"server_id": "",
"name": name,
"transport": transport,
"command": command,
"args": "[]",
"url": url,
"headers": "{}",
"env": "{}",
"auto_approve": False,
"enabled": True,
"created_by": "",
"created": "",
"updated": "",
"source": "config",
"status": per_node,
}
)
return JSONResponse({"servers": result})
async def admin_create_mcp_server(request: Request) -> JSONResponse:
"""POST /v1/api/admin/mcp-servers — create an MCP server definition."""
import uuid
from turnstone.core.audit import record_audit
from turnstone.core.auth import require_permission
from turnstone.core.web_helpers import read_json_or_400, require_storage_or_503
storage, err = require_storage_or_503(request)
if err:
return err
err = require_permission(request, "admin.mcp")
if err:
return err
body = await read_json_or_400(request)
if isinstance(body, JSONResponse):
return body
name = str(body.get("name", "")).strip()[:64]
transport = str(body.get("transport", "")).strip()
if not name:
return JSONResponse({"error": "name is required"}, status_code=400)
if not _MCP_NAME_RE.match(name):
return JSONResponse(
{"error": "name must match [a-zA-Z0-9._-]+"},
status_code=400,
)
if "__" in name:
return JSONResponse(
{"error": "name must not contain '__' (reserved delimiter)"},
status_code=400,
)
if transport not in ("stdio", "streamable-http"):
return JSONResponse(
{"error": "transport must be 'stdio' or 'streamable-http'"},
status_code=400,
)
if transport == "stdio" and not str(body.get("command", "")).strip():
return JSONResponse({"error": "command is required for stdio transport"}, status_code=400)
if transport == "streamable-http" and not str(body.get("url", "")).strip():
return JSONResponse(
{"error": "url is required for streamable-http transport"}, status_code=400
)
# Check max servers
existing = storage.list_mcp_servers()
if len(existing) >= _MCP_MAX_SERVERS:
return JSONResponse(
{"error": f"Maximum {_MCP_MAX_SERVERS} servers"},
status_code=400,
)
# Check name uniqueness
if storage.get_mcp_server_by_name(name):
return JSONResponse(
{"error": f"Server '{name}' already exists"},
status_code=409,
)
server_id = uuid.uuid4().hex
audit_uid, ip = _audit_context(request)
args_list = body.get("args", [])
headers_dict = body.get("headers", {})
env_dict = body.get("env", {})
storage.create_mcp_server(
server_id=server_id,
name=name,
transport=transport,
command=str(body.get("command", "")).strip(),
args=json.dumps(args_list) if isinstance(args_list, list) else "[]",
url=str(body.get("url", "")).strip(),
headers=json.dumps(headers_dict) if isinstance(headers_dict, dict) else "{}",
env=json.dumps(env_dict) if isinstance(env_dict, dict) else "{}",
auto_approve=bool(body.get("auto_approve", False)),
enabled=bool(body.get("enabled", True)),
created_by=audit_uid,
)
record_audit(
storage,
audit_uid,
"mcp_server.create",
"mcp_server",
server_id,
{"name": name},
ip,
)
server = storage.get_mcp_server(server_id)
return JSONResponse(_mcp_server_to_detail(_mask_mcp_secrets(server or {})))
async def admin_get_mcp_server(request: Request) -> JSONResponse:
"""GET /v1/api/admin/mcp-servers/{server_id} — get single MCP server."""
from turnstone.core.auth import require_permission
from turnstone.core.web_helpers import require_storage_or_503
storage, err = require_storage_or_503(request)
if err:
return err
err = require_permission(request, "admin.mcp")
if err:
return err
server_id = request.path_params["server_id"]
server = storage.get_mcp_server(server_id)
if server is None:
return JSONResponse({"error": "MCP server not found"}, status_code=404)
node_statuses = await _collect_mcp_status(request)
per_node: dict[str, dict[str, Any]] = {}
for node_id, node_servers in node_statuses.items():
status = node_servers.get(server["name"])
if status:
per_node[node_id] = status
reveal = str(request.query_params.get("reveal", "")).lower() in ("true", "1")
server = _mask_mcp_secrets(server, reveal)
return JSONResponse(_mcp_server_to_detail(server, per_node))
async def admin_update_mcp_server(request: Request) -> JSONResponse:
"""PUT /v1/api/admin/mcp-servers/{server_id} — update an MCP server."""
from turnstone.core.audit import record_audit
from turnstone.core.auth import require_permission
from turnstone.core.web_helpers import read_json_or_400, require_storage_or_503
storage, err = require_storage_or_503(request)
if err:
return err
err = require_permission(request, "admin.mcp")
if err:
return err
server_id = request.path_params["server_id"]
existing = storage.get_mcp_server(server_id)
if existing is None:
return JSONResponse({"error": "MCP server not found"}, status_code=404)
body = await read_json_or_400(request)
if isinstance(body, JSONResponse):
return body
updates: dict[str, Any] = {}
if "name" in body:
name = str(body["name"]).strip()[:64]
if not name:
return JSONResponse({"error": "name cannot be empty"}, status_code=400)
if not _MCP_NAME_RE.match(name):
return JSONResponse(
{"error": "name must match [a-zA-Z0-9._-]+"},
status_code=400,
)
if "__" in name:
return JSONResponse(
{"error": "name must not contain '__'"},
status_code=400,
)
if name != existing["name"] and storage.get_mcp_server_by_name(name):
return JSONResponse(
{"error": f"Server '{name}' already exists"},
status_code=409,
)
updates["name"] = name
if "transport" in body:
transport = str(body["transport"]).strip()
if transport not in ("stdio", "streamable-http"):
return JSONResponse(
{"error": "transport must be 'stdio' or 'streamable-http'"},
status_code=400,
)
updates["transport"] = transport
if "command" in body:
updates["command"] = str(body["command"]).strip()
if "args" in body:
updates["args"] = json.dumps(body["args"]) if isinstance(body["args"], list) else "[]"
if "url" in body:
updates["url"] = str(body["url"]).strip()
if "headers" in body:
updates["headers"] = (
json.dumps(body["headers"]) if isinstance(body["headers"], dict) else "{}"
)
if "env" in body:
updates["env"] = json.dumps(body["env"]) if isinstance(body["env"], dict) else "{}"
if "auto_approve" in body:
updates["auto_approve"] = bool(body["auto_approve"])
if "enabled" in body:
updates["enabled"] = bool(body["enabled"])
if updates:
storage.update_mcp_server(server_id, **updates)
audit_uid, ip = _audit_context(request)
audit_detail = dict(updates)
for _secret_key in ("env", "headers"):
if _secret_key in audit_detail:
audit_detail[_secret_key] = "(updated)"
record_audit(
storage,
audit_uid,
"mcp_server.update",
"mcp_server",
server_id,
audit_detail,
ip,
)
server = storage.get_mcp_server(server_id)
return JSONResponse(_mcp_server_to_detail(_mask_mcp_secrets(server or {})))
async def admin_delete_mcp_server(request: Request) -> JSONResponse:
"""DELETE /v1/api/admin/mcp-servers/{server_id}."""
from turnstone.core.audit import record_audit
from turnstone.core.auth import require_permission
from turnstone.core.web_helpers import require_storage_or_503
storage, err = require_storage_or_503(request)
if err:
return err
err = require_permission(request, "admin.mcp")
if err:
return err
server_id = request.path_params["server_id"]
existing = storage.get_mcp_server(server_id)
if existing is None:
return JSONResponse({"error": "MCP server not found"}, status_code=404)
storage.delete_mcp_server(server_id)
audit_uid, ip = _audit_context(request)
record_audit(
storage,
audit_uid,
"mcp_server.delete",
"mcp_server",
server_id,
{"name": existing.get("name", "")},
ip,
)
return JSONResponse({"status": "ok"})
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)
client: httpx.AsyncClient = request.app.state.proxy_client
headers = _proxy_auth_headers(request)
sem = asyncio.Semaphore(_NODE_FAN_OUT_LIMIT)
async def _notify(node: dict[str, Any]) -> tuple[str, Any]:
node_id = node.get("node_id", "")
url = node.get("server_url", "")
if not url:
return node_id, None
async with sem:
try:
resp = await client.post(
f"{url.rstrip('/')}/v1/api/_internal/mcp-reload",
headers=headers,
timeout=30,
)
return node_id, resp.json()
except Exception as exc:
log.debug("Failed to notify node %s for MCP reload", node_id, exc_info=True)
return node_id, {"error": str(exc)}
results = await asyncio.gather(*[_notify(n) for n in nodes])
return {nid: data for nid, data in results if data is not None}
async def admin_mcp_reload(request: Request) -> JSONResponse:
"""POST /v1/api/admin/mcp-servers/reload — tell nodes to re-read DB."""
from turnstone.core.auth import require_permission
from turnstone.core.web_helpers import require_storage_or_503
storage, err = require_storage_or_503(request)
if err:
return err
err = require_permission(request, "admin.mcp")
if err:
return err
results = await _notify_nodes_mcp_reload(request)
return JSONResponse({"status": "ok", "results": results})
async def admin_import_mcp_config(request: Request) -> JSONResponse:
"""POST /v1/api/admin/mcp-servers/import — import from pasted JSON config."""
import uuid
from turnstone.core.audit import record_audit
from turnstone.core.auth import require_permission
from turnstone.core.web_helpers import read_json_or_400, require_storage_or_503
storage, err = require_storage_or_503(request)
if err:
return err
err = require_permission(request, "admin.mcp")
if err:
return err
body = await read_json_or_400(request)
if isinstance(body, JSONResponse):
return body
data = body.get("config")
if not isinstance(data, dict):
return JSONResponse(
{"error": "config is required (JSON object with mcpServers key)"}, status_code=400
)
servers = data.get("mcpServers", {})
if not isinstance(servers, dict) or not servers:
return JSONResponse(
{"error": "No mcpServers found in config"},
status_code=400,
)
imported: list[str] = []
skipped: list[str] = []
errors: list[str] = []
audit_uid, ip = _audit_context(request)
current_count = len(storage.list_mcp_servers())
for srv_name, cfg in servers.items():
srv_name = str(srv_name).strip()[:64]
if not srv_name or not _MCP_NAME_RE.match(srv_name) or "__" in srv_name:
errors.append(f"{srv_name}: invalid server name")
continue
if storage.get_mcp_server_by_name(srv_name):
skipped.append(srv_name)
continue
if current_count >= _MCP_MAX_SERVERS:
errors.append(f"{srv_name}: max servers reached")
break
transport = "stdio"
if "url" in cfg or cfg.get("type") in ("http", "streamable-http"):
transport = "streamable-http"
# Coerce fields to expected types
raw_args = cfg.get("args", [])
raw_headers = cfg.get("headers", {})
raw_env = cfg.get("env", {})
if not isinstance(raw_args, list):
errors.append(f"{srv_name}: args must be a list")
continue
if not isinstance(raw_headers, dict):
errors.append(f"{srv_name}: headers must be an object")
continue
if not isinstance(raw_env, dict):
errors.append(f"{srv_name}: env must be an object")
continue
server_id = uuid.uuid4().hex
try:
storage.create_mcp_server(
server_id=server_id,
name=srv_name,
transport=transport,
command=str(cfg.get("command", "")),
args=json.dumps(raw_args),
url=str(cfg.get("url", "")),
headers=json.dumps(raw_headers),
env=json.dumps(raw_env),
auto_approve=False,
enabled=True,
created_by=audit_uid,
)
imported.append(srv_name)
current_count += 1
except Exception as exc:
errors.append(f"{srv_name}: {exc}")
if imported:
record_audit(
storage,
audit_uid,
"mcp_server.import",
"mcp_server",
"",
{"imported": imported, "skipped": skipped},
ip,
)
return JSONResponse({"imported": imported, "skipped": skipped, "errors": errors})
# ---------------------------------------------------------------------------
# App factory
# ---------------------------------------------------------------------------
@@ -3003,10 +3699,14 @@ def create_app(
Route("/api/cluster/snapshot", cluster_snapshot),
Route("/api/cluster/events", cluster_events_sse),
Route("/api/ws-templates", list_ws_templates_summary),
Route("/api/templates", list_templates_summary),
Route("/api/auth/login", auth_login, methods=["POST"]),
Route("/api/auth/logout", auth_logout, methods=["POST"]),
Route("/api/auth/status", auth_status),
Route("/api/auth/setup", auth_setup, methods=["POST"]),
Route("/api/auth/whoami", auth_whoami),
Route("/api/auth/oidc/authorize", oidc_authorize),
Route("/api/auth/oidc/callback", oidc_callback),
Route("/api/admin/users", admin_list_users),
Route("/api/admin/users", admin_create_user, methods=["POST"]),
Route("/api/admin/users/{user_id}", admin_delete_user, methods=["DELETE"]),
@@ -3029,6 +3729,15 @@ def create_app(
admin_delete_channel,
methods=["DELETE"],
),
Route(
"/api/admin/users/{user_id}/oidc-identities",
admin_list_oidc_identities,
),
Route(
"/api/admin/oidc-identities",
admin_delete_oidc_identity,
methods=["DELETE"],
),
Route("/api/admin/schedules", admin_list_schedules),
Route("/api/admin/schedules", admin_create_schedule, methods=["POST"]),
Route("/api/admin/schedules/{task_id}", admin_get_schedule),
@@ -3131,6 +3840,37 @@ def create_app(
admin_delete_setting,
methods=["DELETE"],
),
# System: MCP Servers
Route("/api/admin/mcp-servers", admin_list_mcp_servers),
Route(
"/api/admin/mcp-servers",
admin_create_mcp_server,
methods=["POST"],
),
Route(
"/api/admin/mcp-servers/import",
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,
),
Route(
"/api/admin/mcp-servers/{server_id}",
admin_update_mcp_server,
methods=["PUT"],
),
Route(
"/api/admin/mcp-servers/{server_id}",
admin_delete_mcp_server,
methods=["DELETE"],
),
# Governance: Usage & Audit
Route("/api/admin/usage", admin_usage),
Route("/api/admin/audit", admin_audit),
@@ -3166,6 +3906,13 @@ def create_app(
app.state.login_limiter = LoginRateLimiter()
# OIDC configuration (opt-in via env vars)
from turnstone.core.oidc import load_oidc_config
oidc_config = load_oidc_config()
app.state.oidc_config = oidc_config
app.state.jwks_data = None # populated after async discovery
# Scheduler — start background thread if storage is available
if auth_storage is not None:
from turnstone.console.scheduler import TaskScheduler
+861 -2
View File
@@ -64,7 +64,8 @@ function showAdmin() {
usage: "admin.usage",
audit: "admin.audit",
memories: "admin.memories",
settings: "admin.users",
settings: "admin.settings",
mcp: "admin.mcp",
};
if (perms) {
var permSet = perms.split(",");
@@ -194,6 +195,7 @@ function switchAdminTab(tab) {
"audit",
"memories",
"settings",
"mcp",
];
for (var p = 0; p < panels.length; p++) {
var el = document.getElementById("admin-" + panels[p]);
@@ -216,6 +218,7 @@ function switchAdminTab(tab) {
}
if (tab === "memories") loadAdminMemories();
if (tab === "settings") loadSettings();
if (tab === "mcp") loadAdminMcp();
// Update breadcrumb with active tab label
var activeNav = document.querySelector('.admin-nav[data-tab="' + tab + '"]');
@@ -264,8 +267,13 @@ function _renderUsers(users) {
for (var i = 0; i < users.length; i++) {
var u = users[i];
html +=
'<div class="admin-row" role="listitem">' +
'<div class="admin-row" role="listitem" data-expandable data-user-id="' +
escapeHtml(u.user_id) +
'" data-username="' +
escapeHtml(u.username) +
'" tabindex="0" aria-expanded="false">' +
'<span class="admin-col admin-col-username">' +
'<span class="admin-expand-indicator" aria-hidden="true">\u25b8</span>' +
escapeHtml(u.username) +
"</span>" +
'<span class="admin-col admin-col-name">' +
@@ -304,6 +312,31 @@ function _renderUsers(users) {
);
});
}
// Bind expandable row click + keyboard handlers for OIDC detail panel
var rows = container.querySelectorAll(".admin-row[data-expandable]");
for (var k = 0; k < rows.length; k++) {
(function (row) {
var _expand = function () {
var uid = row.getAttribute("data-user-id");
var uname = row.getAttribute("data-username");
_toggleOidcPanel(uid, uname, row);
};
row.addEventListener("click", function (e) {
if (
e.target.closest(".admin-btn-danger") ||
e.target.closest(".admin-btn-action")
)
return;
_expand();
});
row.addEventListener("keydown", function (e) {
if (e.key === "Enter" || e.key === " ") {
e.preventDefault();
_expand();
}
});
})(rows[k]);
}
}
function confirmDeleteUser(userId, username) {
@@ -329,6 +362,253 @@ function confirmDeleteUser(userId, username) {
);
}
// ---------------------------------------------------------------------------
// OIDC identity expansion in Users tab
// ---------------------------------------------------------------------------
function _toggleOidcPanel(userId, username, rowEl) {
var existing = rowEl.nextElementSibling;
if (existing && existing.classList.contains("oidc-detail-panel")) {
// Collapse
existing.style.maxHeight = "0";
var indicator = rowEl.querySelector(".admin-expand-indicator");
if (indicator) indicator.classList.remove("expanded");
rowEl.setAttribute("aria-expanded", "false");
setTimeout(function () {
if (existing.parentNode) existing.remove();
}, 160);
return;
}
// Collapse any other open panel first
var openPanels = document.querySelectorAll(
"#admin-users-table .oidc-detail-panel",
);
for (var i = 0; i < openPanels.length; i++) {
openPanels[i].style.maxHeight = "0";
var prevRow = openPanels[i].previousElementSibling;
if (prevRow) {
var ind = prevRow.querySelector(".admin-expand-indicator");
if (ind) ind.classList.remove("expanded");
prevRow.setAttribute("aria-expanded", "false");
}
(function (panel) {
setTimeout(function () {
if (panel.parentNode) panel.remove();
}, 160);
})(openPanels[i]);
}
// Mark expanded
var indicator = rowEl.querySelector(".admin-expand-indicator");
if (indicator) indicator.classList.add("expanded");
rowEl.setAttribute("aria-expanded", "true");
// Create panel (role="none" so it doesn't break the parent role="list")
var panel = document.createElement("div");
panel.className = "oidc-detail-panel";
panel.setAttribute("role", "none");
panel.innerHTML =
'<div class="oidc-detail-inner">' +
'<div class="oidc-detail-header">OIDC Identities</div>' +
'<div class="oidc-detail-body"><span class="oidc-detail-empty">Loading\u2026</span></div>' +
"</div>";
rowEl.after(panel);
// Animate open
requestAnimationFrame(function () {
panel.style.maxHeight = panel.scrollHeight + "px";
});
// Fetch identities
authFetch(
"/v1/api/admin/users/" + encodeURIComponent(userId) + "/oidc-identities",
)
.then(function (r) {
if (!r.ok) throw new Error("Failed");
return r.json();
})
.then(function (data) {
_renderOidcDetail(panel, data.oidc_identities || [], userId, username);
})
.catch(function () {
var body = panel.querySelector(".oidc-detail-body");
if (body)
body.innerHTML =
'<span class="oidc-detail-empty">Failed to load</span>';
});
}
function _renderOidcDetail(panel, identities, userId, username) {
var body = panel.querySelector(".oidc-detail-body");
if (!body) return;
if (!identities.length) {
body.innerHTML =
'<span class="oidc-detail-empty">No OIDC identities linked</span>';
panel.style.maxHeight = panel.scrollHeight + "px";
return;
}
var html = "";
for (var i = 0; i < identities.length; i++) {
var oid = identities[i];
var shortIssuer = _issuerShortName(oid.issuer || "");
var shortSubject =
(oid.subject || "").length > 12
? (oid.subject || "").slice(0, 12) + "\u2026"
: oid.subject || "";
var lastLogin = oid.last_login ? _relativeTime(oid.last_login) : "never";
html +=
'<div class="oidc-identity-row">' +
'<span class="oidc-identity-issuer"><span class="scope-badge">' +
escapeHtml(shortIssuer) +
"</span></span>" +
'<span class="oidc-identity-subject" title="' +
escapeHtml(oid.subject || "") +
'">' +
escapeHtml(shortSubject) +
"</span>" +
'<span class="oidc-identity-email" title="' +
escapeHtml(oid.email || "") +
'">' +
escapeHtml(oid.email || "\u2014") +
"</span>" +
'<span class="oidc-identity-time">' +
escapeHtml(lastLogin) +
"</span>" +
'<span class="oidc-identity-actions">' +
'<button class="admin-btn-danger" aria-label="Unlink ' +
escapeHtml(shortIssuer) +
" identity " +
escapeHtml(shortSubject) +
'" data-oidc-issuer="' +
escapeHtml(oid.issuer || "") +
'" data-oidc-subject="' +
escapeHtml(oid.subject || "") +
'" data-oidc-username="' +
escapeHtml(username) +
'" data-oidc-user-id="' +
escapeHtml(userId) +
'">unlink</button>' +
"</span></div>";
}
body.innerHTML = html;
// Update panel height for animation
panel.style.maxHeight = panel.scrollHeight + "px";
// Bind unlink buttons
var btns = body.querySelectorAll("[data-oidc-issuer]");
for (var j = 0; j < btns.length; j++) {
btns[j].addEventListener("click", function (e) {
e.stopPropagation();
var issuer = this.getAttribute("data-oidc-issuer");
var subject = this.getAttribute("data-oidc-subject");
var uname = this.getAttribute("data-oidc-username");
var uid = this.getAttribute("data-oidc-user-id");
_confirmUnlinkOidc(issuer, subject, uname, uid);
});
}
}
function _confirmUnlinkOidc(issuer, subject, username, userId) {
var shortIssuer = _issuerShortName(issuer);
var shortSubject =
subject.length > 16 ? subject.slice(0, 16) + "\u2026" : subject;
showConfirmModal(
"Unlink OIDC Identity",
"Unlink " +
shortIssuer +
" identity \u2018" +
shortSubject +
"\u2019 from user " +
username +
"?\n\nThe user will need to log in via OIDC again to re-link.",
"Unlink",
function () {
authFetch(
"/v1/api/admin/oidc-identities?issuer=" +
encodeURIComponent(issuer) +
"&subject=" +
encodeURIComponent(subject),
{ method: "DELETE" },
)
.then(function (r) {
if (!r.ok) throw new Error("Unlink failed");
showToast("OIDC identity unlinked");
// Refresh the panel content in place (no close/reopen flicker)
var allRows = document.querySelectorAll(
"#admin-users-table .admin-row[data-expandable]",
);
var targetRow = null;
for (var ri = 0; ri < allRows.length; ri++) {
if (allRows[ri].getAttribute("data-user-id") === userId) {
targetRow = allRows[ri];
break;
}
}
if (targetRow) {
var panel = targetRow.nextElementSibling;
if (panel && panel.classList.contains("oidc-detail-panel")) {
var body = panel.querySelector(".oidc-detail-body");
if (body)
body.innerHTML =
'<span class="oidc-detail-empty">Loading\u2026</span>';
authFetch(
"/v1/api/admin/users/" +
encodeURIComponent(userId) +
"/oidc-identities",
)
.then(function (r2) {
if (!r2.ok) throw new Error("Failed");
return r2.json();
})
.then(function (data) {
_renderOidcDetail(
panel,
data.oidc_identities || [],
userId,
username,
);
})
.catch(function () {
if (body)
body.innerHTML =
'<span class="oidc-detail-empty">Failed to load</span>';
});
}
}
})
.catch(function () {
showToast("Failed to unlink OIDC identity");
});
},
);
}
function _issuerShortName(issuer) {
try {
var host = new URL(issuer).hostname;
if (host.includes("google")) return "google";
if (host.includes("microsoftonline") || host.includes("azure"))
return "azure";
if (host.includes("okta")) return "okta";
if (host.includes("auth0")) return "auth0";
if (host.includes("keycloak")) return "keycloak";
return host.replace(/^(login|accounts|auth|id|sso)\./, "");
} catch (e) {
return issuer || "unknown";
}
}
function _relativeTime(isoStr) {
try {
var then = new Date(
isoStr + (isoStr.includes("Z") || isoStr.includes("+") ? "" : "Z"),
);
var diff = (Date.now() - then.getTime()) / 1000;
if (diff < 60) return "just now";
if (diff < 3600) return Math.floor(diff / 60) + "m ago";
if (diff < 86400) return Math.floor(diff / 3600) + "h ago";
if (diff < 2592000) return Math.floor(diff / 86400) + "d ago";
return isoStr.slice(0, 10);
} catch (e) {
return isoStr || "unknown";
}
}
// ---------------------------------------------------------------------------
// Tokens
// ---------------------------------------------------------------------------
@@ -1593,6 +1873,9 @@ function _installTrap(overlayId, boxId, trapRef) {
else if (overlayId === "edit-wst-overlay") hideEditWsTemplateModal();
else if (overlayId === "wst-history-overlay") hideWstHistoryModal();
else if (overlayId === "memory-detail-overlay") hideMemoryDetailModal();
else if (overlayId === "mcp-create-overlay") hideCreateMcpModal();
else if (overlayId === "mcp-import-overlay") hideImportMcpModal();
else if (overlayId === "mcp-detail-overlay") hideMcpDetailModal();
}
};
}
@@ -1679,6 +1962,9 @@ document.addEventListener("keydown", function (e) {
["edit-wst-overlay", hideEditWsTemplateModal],
["wst-history-overlay", hideWstHistoryModal],
["memory-detail-overlay", hideMemoryDetailModal],
["mcp-detail-overlay", hideMcpDetailModal],
["mcp-import-overlay", hideImportMcpModal],
["mcp-create-overlay", hideCreateMcpModal],
];
for (var gi = 0; gi < govOverlays.length; gi++) {
var govEl = document.getElementById(govOverlays[gi][0]);
@@ -2354,3 +2640,576 @@ function _showModalError(el, msg) {
el.textContent = msg;
el.style.display = "block";
}
/* ── MCP Servers tab ─────────────────────────────────────────────────────── */
var _mcpServers = [];
var _mcpCreateTrap = null;
var _mcpCreateTrigger = null;
var _mcpImportTrap = null;
var _mcpImportTrigger = null;
var _mcpDetailTrap = null;
var _mcpDetailTrigger = null;
function loadAdminMcp() {
authFetch("/v1/api/admin/mcp-servers")
.then(function (r) {
if (!r.ok) throw new Error("Failed");
return r.json();
})
.then(function (data) {
_mcpServers = data.servers || [];
_renderMcpServers(_mcpServers);
})
.catch(function () {
document.getElementById("admin-mcp-table").innerHTML =
'<div class="dashboard-empty">Failed to load MCP servers</div>';
});
}
function _renderMcpServers(items) {
var el = document.getElementById("admin-mcp-table");
if (!items.length) {
el.innerHTML =
'<div class="dashboard-empty">No MCP servers configured</div>';
return;
}
var html = "";
for (var i = 0; i < items.length; i++) {
var s = items[i];
var statusEntries = s.status || {};
var nodeIds = Object.keys(statusEntries);
var anyConnected = false;
var anyError = false;
var firstError = "";
var totalTools = 0,
totalRes = 0,
totalPrompts = 0;
for (var j = 0; j < nodeIds.length; j++) {
var ns = statusEntries[nodeIds[j]];
if (ns.connected) {
anyConnected = true;
totalTools += ns.tools || 0;
totalRes += ns.resources || 0;
totalPrompts += ns.prompts || 0;
}
if (ns.error) {
anyError = true;
if (!firstError) firstError = ns.error;
}
}
var dotClass = "mcp-status-dot disabled";
var rowClass = "mcp-row-disabled";
var statusText = "disabled";
if (!s.enabled) {
statusText = "disabled";
} else if (anyConnected) {
dotClass = "mcp-status-dot connected";
rowClass = "mcp-row-connected";
statusText = "connected";
} else if (anyError) {
dotClass = "mcp-status-dot error";
rowClass = "mcp-row-error";
statusText = "error";
} else {
dotClass = "mcp-status-dot disabled";
rowClass = "mcp-row-disabled";
statusText = "idle";
}
var transportCls =
s.transport === "stdio" ? "mcp-transport-stdio" : "mcp-transport-http";
var toolsVal = anyConnected
? totalTools
: '<span class="mcp-count-dim">--</span>';
var resVal = anyConnected
? totalRes
: '<span class="mcp-count-dim">--</span>';
var promptsVal = anyConnected
? totalPrompts
: '<span class="mcp-count-dim">--</span>';
var isConfig = s.source === "config";
var nameBadge = isConfig
? ' <span class="scope-badge scope-channel">config</span>'
: "";
var detailAttr = isConfig
? 'data-mcp-detail-name="' + escapeHtml(s.name) + '"'
: 'data-mcp-detail="' + escapeHtml(s.server_id) + '"';
var actions = isConfig
? ""
: '<button class="admin-btn-action" data-mcp-edit="' +
escapeHtml(s.server_id) +
'">edit</button>' +
'<button class="admin-btn-danger" data-mcp-delete="' +
escapeHtml(s.server_id) +
'" data-mcp-name="' +
escapeHtml(s.name) +
'">del</button>';
html +=
'<div class="admin-row mcp-grid ' +
rowClass +
'" role="listitem">' +
'<span class="admin-col admin-col-mname"><a href="#" ' +
detailAttr +
">" +
escapeHtml(s.name) +
"</a>" +
nameBadge +
"</span>" +
'<span class="admin-col admin-col-mtransport"><span class="mcp-transport-badge ' +
transportCls +
'">' +
escapeHtml(s.transport) +
"</span></span>" +
'<span class="admin-col admin-col-mtools">' +
toolsVal +
"</span>" +
'<span class="admin-col admin-col-mres">' +
resVal +
"</span>" +
'<span class="admin-col admin-col-mprompts">' +
promptsVal +
"</span>" +
'<span class="admin-col admin-col-mstatus"' +
(firstError ? ' title="' + escapeHtml(firstError) + '"' : "") +
'><span class="' +
dotClass +
'" aria-hidden="true"></span>' +
escapeHtml(statusText) +
"</span>" +
'<span class="admin-col admin-col-mactions">' +
actions +
"</span></div>";
}
el.innerHTML = html;
// Bind event handlers
el.querySelectorAll("[data-mcp-detail]").forEach(function (a) {
a.addEventListener("click", function (e) {
e.preventDefault();
showMcpDetailModal(this.getAttribute("data-mcp-detail"));
});
});
el.querySelectorAll("[data-mcp-detail-name]").forEach(function (a) {
a.addEventListener("click", function (e) {
e.preventDefault();
showMcpDetailByName(this.getAttribute("data-mcp-detail-name"));
});
});
el.querySelectorAll("[data-mcp-edit]").forEach(function (btn) {
btn.addEventListener("click", function () {
showEditMcpModal(this.getAttribute("data-mcp-edit"));
});
});
el.querySelectorAll("[data-mcp-delete]").forEach(function (btn) {
btn.addEventListener("click", function () {
var sid = this.getAttribute("data-mcp-delete");
var sname = this.getAttribute("data-mcp-name");
showConfirmModal(
"Delete MCP Server",
'Delete server "' + sname + '"?',
"Delete",
function () {
authFetch("/v1/api/admin/mcp-servers/" + sid, { method: "DELETE" })
.then(function (r) {
if (!r.ok) throw new Error();
return r.json();
})
.then(function () {
showToast("Server deleted");
loadAdminMcp();
})
.catch(function () {
showToast("Failed to delete server");
});
},
);
});
});
}
function toggleMcpTransport() {
var v = document.getElementById("mcp-transport").value;
document.getElementById("mcp-stdio-fields").style.display =
v === "stdio" ? "" : "none";
document.getElementById("mcp-http-fields").style.display =
v === "streamable-http" ? "" : "none";
}
function showCreateMcpModal() {
_mcpCreateTrigger = document.activeElement;
var ov = document.getElementById("mcp-create-overlay");
ov.style.display = "flex";
document.getElementById("mcp-edit-id").value = "";
document.getElementById("mcp-create-title").textContent = "Add MCP Server";
document.getElementById("mcp-create-submit").textContent = "Create";
document.getElementById("mcp-name").value = "";
document.getElementById("mcp-transport").value = "stdio";
document.getElementById("mcp-command").value = "";
document.getElementById("mcp-args").value = "";
document.getElementById("mcp-env").value = "";
document.getElementById("mcp-url").value = "";
document.getElementById("mcp-headers").value = "";
document.getElementById("mcp-auto-approve").checked = false;
document.getElementById("mcp-enabled").checked = true;
document.getElementById("mcp-create-error").style.display = "none";
toggleMcpTransport();
document.getElementById("mcp-name").focus();
_mcpCreateTrap = _installTrap("mcp-create-overlay", "mcp-create-box");
}
function showEditMcpModal(serverId) {
// Fetch with reveal=true to get actual secret values for editing
authFetch("/v1/api/admin/mcp-servers/" + serverId + "?reveal=true")
.then(function (r) {
if (!r.ok) throw new Error("Failed to load server");
return r.json();
})
.then(function (s) {
showCreateMcpModal();
document.getElementById("mcp-edit-id").value = serverId;
document.getElementById("mcp-create-title").textContent =
"Edit MCP Server";
document.getElementById("mcp-create-submit").textContent = "Save";
document.getElementById("mcp-name").value = s.name;
document.getElementById("mcp-transport").value = s.transport;
document.getElementById("mcp-command").value = s.command || "";
try {
var argsList = JSON.parse(s.args || "[]");
document.getElementById("mcp-args").value = argsList.join("\n");
} catch (e) {
document.getElementById("mcp-args").value = "";
}
try {
var envObj = JSON.parse(s.env || "{}");
document.getElementById("mcp-env").value = Object.keys(envObj)
.map(function (k) {
return k + "=" + envObj[k];
})
.join("\n");
} catch (e) {
document.getElementById("mcp-env").value = "";
}
document.getElementById("mcp-url").value = s.url || "";
try {
var hdrObj = JSON.parse(s.headers || "{}");
document.getElementById("mcp-headers").value = Object.keys(hdrObj)
.map(function (k) {
return k + ": " + hdrObj[k];
})
.join("\n");
} catch (e) {
document.getElementById("mcp-headers").value = "";
}
document.getElementById("mcp-auto-approve").checked =
s.auto_approve || false;
document.getElementById("mcp-enabled").checked = s.enabled !== false;
toggleMcpTransport();
})
.catch(function () {
showToast("Failed to load server details");
});
}
function hideCreateMcpModal() {
document.getElementById("mcp-create-overlay").style.display = "none";
_mcpCreateTrap = _removeTrap(_mcpCreateTrap);
if (_mcpCreateTrigger && _mcpCreateTrigger.focus) _mcpCreateTrigger.focus();
_mcpCreateTrigger = null;
}
function _parseMcpForm() {
var name = document.getElementById("mcp-name").value.trim();
var transport = document.getElementById("mcp-transport").value;
if (!name) return { error: "Name is required" };
if (!/^[a-zA-Z0-9._-]+$/.test(name))
return { error: "Name must match [a-zA-Z0-9._-]+" };
if (name.indexOf("__") >= 0) return { error: "Name must not contain '__'" };
var payload = {
name: name,
transport: transport,
auto_approve: document.getElementById("mcp-auto-approve").checked,
enabled: document.getElementById("mcp-enabled").checked,
};
if (transport === "stdio") {
payload.command = document.getElementById("mcp-command").value.trim();
var argsText = document.getElementById("mcp-args").value.trim();
payload.args = argsText
? argsText
.split("\n")
.map(function (l) {
return l.trim();
})
.filter(Boolean)
: [];
var envText = document.getElementById("mcp-env").value.trim();
var envObj = {};
if (envText) {
envText.split("\n").forEach(function (line) {
var eq = line.indexOf("=");
if (eq > 0)
envObj[line.substring(0, eq).trim()] = line.substring(eq + 1).trim();
});
}
payload.env = envObj;
} else {
payload.url = document.getElementById("mcp-url").value.trim();
var hdrText = document.getElementById("mcp-headers").value.trim();
var hdrObj = {};
if (hdrText) {
hdrText.split("\n").forEach(function (line) {
var colon = line.indexOf(":");
if (colon > 0)
hdrObj[line.substring(0, colon).trim()] = line
.substring(colon + 1)
.trim();
});
}
payload.headers = hdrObj;
}
return payload;
}
function submitCreateMcp() {
var form = _parseMcpForm();
if (form.error) {
var e = document.getElementById("mcp-create-error");
e.textContent = form.error;
e.style.display = "";
return;
}
var editId = document.getElementById("mcp-edit-id").value;
var method = editId ? "PUT" : "POST";
var url = editId
? "/v1/api/admin/mcp-servers/" + editId
: "/v1/api/admin/mcp-servers";
document.getElementById("mcp-create-submit").disabled = true;
authFetch(url, {
method: method,
headers: { "Content-Type": "application/json" },
body: JSON.stringify(form),
})
.then(function (r) {
if (!r.ok)
return r.json().then(function (d) {
throw new Error(d.error || "Failed");
});
return r.json();
})
.then(function () {
hideCreateMcpModal();
showToast(editId ? "Server updated" : "Server created");
loadAdminMcp();
})
.catch(function (e) {
var el = document.getElementById("mcp-create-error");
el.textContent = e.message;
el.style.display = "";
})
.finally(function () {
document.getElementById("mcp-create-submit").disabled = false;
});
}
function reloadMcpNodes() {
authFetch("/v1/api/admin/mcp-servers/reload", { method: "POST" })
.then(function (r) {
if (!r.ok) throw new Error();
return r.json();
})
.then(function (data) {
var results = data.results || {};
var nodeIds = Object.keys(results);
var totalAdded = 0,
totalRemoved = 0;
for (var i = 0; i < nodeIds.length; i++) {
var nr = results[nodeIds[i]];
totalAdded += (nr.added || []).length;
totalRemoved += (nr.removed || []).length;
}
var msg = "Reload sent to " + nodeIds.length + " node(s)";
if (totalAdded) msg += ", +" + totalAdded + " added";
if (totalRemoved) msg += ", -" + totalRemoved + " removed";
showToast(msg);
setTimeout(loadAdminMcp, 1500);
})
.catch(function () {
showToast("Failed to reload nodes");
});
}
function showMcpDetailByName(name) {
for (var i = 0; i < _mcpServers.length; i++) {
if (_mcpServers[i].name === name) {
return _openMcpDetail(_mcpServers[i]);
}
}
}
function showMcpDetailModal(serverId) {
for (var i = 0; i < _mcpServers.length; i++) {
if (_mcpServers[i].server_id === serverId) {
return _openMcpDetail(_mcpServers[i]);
}
}
}
function _openMcpDetail(s) {
if (!s) return;
_mcpDetailTrigger = document.activeElement;
var html = '<div class="modal-columns">';
html += '<div class="modal-col">';
html += '<div class="mcp-detail-section"><h3>Configuration</h3>';
html +=
'<p style="font-size:12px;color:var(--fg-dim)">Transport: <span class="mcp-transport-badge ' +
(s.transport === "stdio" ? "mcp-transport-stdio" : "mcp-transport-http") +
'">' +
escapeHtml(s.transport) +
"</span></p>";
if (s.transport === "stdio") {
html +=
'<p style="font-size:12px;color:var(--fg-dim)">Command: <code>' +
escapeHtml(s.command || "") +
"</code></p>";
try {
var a = JSON.parse(s.args || "[]");
if (a.length)
html +=
'<p style="font-size:12px;color:var(--fg-dim)">Args: <code>' +
escapeHtml(a.join(" ")) +
"</code></p>";
} catch (e) {}
} else {
html +=
'<p style="font-size:12px;color:var(--fg-dim)">URL: <code>' +
escapeHtml(s.url || "") +
"</code></p>";
}
html += "</div></div>";
html += '<div class="modal-col">';
var statusEntries = s.status || {};
var nodeIds = Object.keys(statusEntries);
html += '<div class="mcp-detail-section"><h3>Node Status</h3>';
if (nodeIds.length === 0) {
html +=
'<p style="font-size:12px;color:var(--fg-dim)">Not connected on any node</p>';
} else {
html += '<ul class="mcp-detail-list">';
for (var j = 0; j < nodeIds.length; j++) {
var ns = statusEntries[nodeIds[j]];
var dot = ns.connected
? '<span class="mcp-status-dot connected"></span>'
: '<span class="mcp-status-dot error"></span>';
var nodeInfo =
escapeHtml(nodeIds[j]) +
" — " +
(ns.tools || 0) +
" tools, " +
(ns.resources || 0) +
" resources, " +
(ns.prompts || 0) +
" prompts";
if (ns.error) {
nodeInfo +=
'<br><span style="color:var(--red);font-size:11px">' +
escapeHtml(ns.error) +
"</span>";
}
html += "<li>" + dot + nodeInfo + "</li>";
}
html += "</ul>";
}
html += "</div></div></div>";
document.getElementById("mcp-detail-title").textContent = s.name;
document.getElementById("mcp-detail-content").innerHTML = html;
document.getElementById("mcp-detail-overlay").style.display = "flex";
_mcpDetailTrap = _installTrap("mcp-detail-overlay", "mcp-detail-box");
}
function hideMcpDetailModal() {
document.getElementById("mcp-detail-overlay").style.display = "none";
_mcpDetailTrap = _removeTrap(_mcpDetailTrap);
if (_mcpDetailTrigger && _mcpDetailTrigger.focus) _mcpDetailTrigger.focus();
_mcpDetailTrigger = null;
}
function showImportMcpModal() {
_mcpImportTrigger = document.activeElement;
document.getElementById("mcp-import-overlay").style.display = "flex";
document.getElementById("mcp-import-json").value = "";
document.getElementById("mcp-import-error").style.display = "none";
document.getElementById("mcp-import-json").focus();
_mcpImportTrap = _installTrap("mcp-import-overlay", "mcp-import-box");
}
function hideImportMcpModal() {
document.getElementById("mcp-import-overlay").style.display = "none";
_mcpImportTrap = _removeTrap(_mcpImportTrap);
if (_mcpImportTrigger && _mcpImportTrigger.focus) _mcpImportTrigger.focus();
_mcpImportTrigger = null;
}
function submitImportMcp() {
var raw = document.getElementById("mcp-import-json").value.trim();
if (!raw) {
var e = document.getElementById("mcp-import-error");
e.textContent = "Paste a JSON config";
e.style.display = "";
return;
}
var parsed;
try {
parsed = JSON.parse(raw);
} catch (ex) {
var e2 = document.getElementById("mcp-import-error");
e2.textContent = "Invalid JSON: " + ex.message;
e2.style.display = "";
return;
}
if (!parsed.mcpServers || typeof parsed.mcpServers !== "object") {
var e3 = document.getElementById("mcp-import-error");
e3.textContent = 'No "mcpServers" key found in JSON';
e3.style.display = "";
return;
}
document.getElementById("mcp-import-submit").disabled = true;
authFetch("/v1/api/admin/mcp-servers/import", {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({ config: parsed }),
})
.then(function (r) {
if (!r.ok)
return r.json().then(function (d) {
throw new Error(d.error || "Failed");
});
return r.json();
})
.then(function (data) {
hideImportMcpModal();
var msg = "Imported " + (data.imported || []).length;
if ((data.skipped || []).length)
msg += ", skipped " + data.skipped.length;
if ((data.errors || []).length)
msg += ", " + data.errors.length + " error(s)";
showToast(msg);
loadAdminMcp();
})
.catch(function (e) {
var el = document.getElementById("mcp-import-error");
el.textContent = e.message;
el.style.display = "";
})
.finally(function () {
document.getElementById("mcp-import-submit").disabled = false;
});
}
+14 -8
View File
@@ -1263,7 +1263,7 @@ function showNewWsModal() {
// Populate template dropdown
var tplSelect = document.getElementById("new-ws-template");
tplSelect.innerHTML = '<option value="">Use defaults</option>';
authFetch("/v1/api/admin/templates")
authFetch("/v1/api/templates")
.then(function (r) {
return r.json();
})
@@ -1303,7 +1303,14 @@ function showNewWsModal() {
});
document.getElementById("new-ws-name").value = "";
document.getElementById("new-ws-model").value = "";
document.getElementById("new-ws-task").value = "";
var taskEl = document.getElementById("new-ws-task");
taskEl.value = "";
var mod =
navigator.platform && navigator.platform.indexOf("Mac") > -1
? "\u2318"
: "Ctrl";
taskEl.placeholder =
"What should this workstream work on? (" + mod + "+Enter to create)";
var errEl = document.getElementById("new-ws-error");
errEl.style.display = "none";
errEl.textContent = "";
@@ -1336,7 +1343,7 @@ function showNewWsModal() {
document.addEventListener("keydown", _newWsTrapHandler);
setTimeout(function () {
document.getElementById("new-ws-name").focus();
document.getElementById("new-ws-task").focus();
}, 50);
}
@@ -1412,11 +1419,10 @@ document.addEventListener("keydown", function (e) {
e.preventDefault();
hideNewWsModal();
}
if (
e.key === "Enter" &&
e.target.tagName !== "SELECT" &&
e.target.tagName !== "TEXTAREA"
) {
if (e.key === "Enter") {
if (e.target.tagName === "SELECT") return;
if (e.target.tagName === "BUTTON") return; // let native click fire
if (e.target.tagName === "TEXTAREA" && !(e.ctrlKey || e.metaKey)) return;
e.preventDefault();
var btn = document.getElementById("new-ws-submit");
if (btn && !btn.disabled) submitNewWs();
+5
View File
@@ -152,10 +152,15 @@ var _ALL_PERMISSIONS = [
"admin.orgs",
"admin.policies",
"admin.templates",
"admin.ws_templates",
"admin.audit",
"admin.usage",
"admin.schedules",
"admin.watches",
"admin.judge",
"admin.memories",
"admin.settings",
"admin.mcp",
"tools.approve",
"workstreams.create",
"workstreams.close",
+86 -2
View File
@@ -107,6 +107,7 @@
<div class="admin-sidebar-group" data-group="system" role="group" aria-label="System">
<div class="admin-sidebar-group-label" aria-hidden="true">System</div>
<button id="tab-settings" class="admin-nav" data-tab="settings" role="tab" aria-selected="false" aria-controls="admin-settings" tabindex="-1" onclick="switchAdminTab('settings')">Settings</button>
<button id="tab-mcp" class="admin-nav" data-tab="mcp" role="tab" aria-selected="false" aria-controls="admin-mcp" tabindex="-1" onclick="switchAdminTab('mcp')">MCP Servers</button>
</div>
</nav>
<div id="admin-sidebar-backdrop" class="admin-sidebar-backdrop" aria-hidden="true"></div>
@@ -393,6 +394,27 @@
</div>
</div>
<div id="admin-mcp" class="admin-panel" role="tabpanel" aria-labelledby="tab-mcp" style="display:none">
<div class="admin-toolbar">
<span class="section-header">MCP SERVERS</span>
<button class="admin-action-btn admin-action-btn-ghost" onclick="reloadMcpNodes()" title="Push MCP server config to all cluster nodes and reconnect">Sync to Nodes</button>
<button class="admin-action-btn admin-action-btn-ghost" onclick="showImportMcpModal()">Import JSON</button>
<button class="admin-action-btn" onclick="showCreateMcpModal()">+ Add Server</button>
</div>
<div class="admin-colheaders mcp-grid" aria-hidden="true">
<span class="admin-col admin-col-mname">NAME</span>
<span class="admin-col admin-col-mtransport">TRANSPORT</span>
<span class="admin-col admin-col-mtools">TOOLS</span>
<span class="admin-col admin-col-mres">RES</span>
<span class="admin-col admin-col-mprompts">PROMPTS</span>
<span class="admin-col admin-col-mstatus">STATUS</span>
<span class="admin-col admin-col-mactions">ACTIONS</span>
</div>
<div id="admin-mcp-table" role="list" aria-label="MCP servers">
<div class="dashboard-empty">Loading...</div>
</div>
</div>
</div><!-- /admin-content -->
</div><!-- /admin-layout -->
</div>
@@ -415,6 +437,7 @@ window.TURNSTONE_KB_SHORTCUTS = [
{ desc: "Navigate rows", badge: '<span class="kb-key">\u2191</span> <span class="kb-key">\u2193</span>' }
]},
{ title: "General", keys: [
{ desc: "Submit form from textarea", badge: '<span class="kb-key">' + (navigator.platform && navigator.platform.indexOf("Mac") > -1 ? "\u2318" : "Ctrl") + '</span>+<span class="kb-key">Enter</span>' },
{ desc: "Show this help", badge: '<span class="kb-key">?</span>' },
{ desc: "Close overlay", badge: '<span class="kb-key">Esc</span>' }
]}
@@ -430,6 +453,8 @@ window.TURNSTONE_KB_SHORTCUTS = [
<div id="new-ws-box">
<h2 id="new-ws-title">New Workstream</h2>
<div id="new-ws-error" role="alert" aria-live="assertive"></div>
<label for="new-ws-task">Task <span class="label-hint">optional &mdash; sent as first message</span></label>
<textarea id="new-ws-task" rows="4" placeholder="What should this workstream work on?"></textarea>
<label for="new-ws-node">Node</label>
<select id="new-ws-node">
<option value="">Auto (best available)</option>
@@ -446,8 +471,6 @@ window.TURNSTONE_KB_SHORTCUTS = [
<select id="new-ws-profile">
<option value="">None</option>
</select>
<label for="new-ws-task">Task <span class="label-hint">optional &mdash; sent as first message</span></label>
<textarea id="new-ws-task" rows="3" placeholder="What should this workstream work on?"></textarea>
<div id="new-ws-buttons">
<button id="new-ws-cancel" onclick="hideNewWsModal()">Cancel</button>
<button id="new-ws-submit" onclick="submitNewWs()">Create</button>
@@ -972,6 +995,67 @@ window.TURNSTONE_KB_SHORTCUTS = [
</div>
</div>
<div id="mcp-create-overlay" style="display:none" role="dialog" aria-modal="true" aria-labelledby="mcp-create-title">
<div id="mcp-create-box" class="admin-modal">
<h2 id="mcp-create-title">Add MCP Server</h2>
<div id="mcp-create-error" role="alert" aria-live="assertive" style="display:none"></div>
<input type="hidden" id="mcp-edit-id" value="">
<label for="mcp-name">Server Name</label>
<input type="text" id="mcp-name" placeholder="e.g. filesystem" maxlength="64" pattern="[a-zA-Z0-9._-]+">
<label for="mcp-transport">Transport</label>
<select id="mcp-transport" onchange="toggleMcpTransport()">
<option value="stdio">stdio</option>
<option value="streamable-http">streamable-http</option>
</select>
<div id="mcp-stdio-fields">
<label for="mcp-command">Command</label>
<input type="text" id="mcp-command" placeholder="e.g. npx">
<label for="mcp-args">Arguments <span style="font-weight:400;text-transform:none">(one per line)</span></label>
<textarea id="mcp-args" rows="3" placeholder="-y&#10;@modelcontextprotocol/server-filesystem&#10;/tmp"></textarea>
<label for="mcp-env">Environment Variables <span style="font-weight:400;text-transform:none">(KEY=VALUE, one per line)</span></label>
<textarea id="mcp-env" rows="2" placeholder="API_KEY=..."></textarea>
</div>
<div id="mcp-http-fields" style="display:none">
<label for="mcp-url">URL</label>
<input type="text" id="mcp-url" placeholder="https://...">
<label for="mcp-headers">Headers <span style="font-weight:400;text-transform:none">(KEY: VALUE, one per line)</span></label>
<textarea id="mcp-headers" rows="2" placeholder="Authorization: Bearer ..."></textarea>
</div>
<div style="display:flex;gap:20px;margin-top:14px">
<label style="margin:0;font-size:12px;color:var(--fg-dim)"><input type="checkbox" id="mcp-auto-approve" style="margin-right:5px">Auto-approve tools</label>
<label style="margin:0;font-size:12px;color:var(--fg-dim)"><input type="checkbox" id="mcp-enabled" checked style="margin-right:5px">Enabled</label>
</div>
<div class="modal-buttons">
<button class="modal-cancel" onclick="hideCreateMcpModal()">Cancel</button>
<button id="mcp-create-submit" class="modal-submit" onclick="submitCreateMcp()">Create</button>
</div>
</div>
</div>
<div id="mcp-import-overlay" style="display:none" role="dialog" aria-modal="true" aria-labelledby="mcp-import-title">
<div id="mcp-import-box" class="admin-modal">
<h2 id="mcp-import-title">Import MCP Config</h2>
<div id="mcp-import-error" role="alert" aria-live="assertive" style="display:none"></div>
<label for="mcp-import-json">Paste JSON</label>
<textarea id="mcp-import-json" rows="10" placeholder='{"mcpServers":{"filesystem":{"command":"npx","args":["-y","@modelcontextprotocol/server-filesystem","/tmp"]}}}' style="font-family:var(--font-mono);font-size:11px"></textarea>
<p style="font-size:11px;color:var(--fg-dim);margin-top:8px">Paste a JSON object with a <code>mcpServers</code> key (Claude Desktop / VS Code / Cursor format). Existing servers with the same name will be skipped.</p>
<div class="modal-buttons">
<button class="modal-cancel" onclick="hideImportMcpModal()">Cancel</button>
<button id="mcp-import-submit" class="modal-submit" onclick="submitImportMcp()">Import</button>
</div>
</div>
</div>
<div id="mcp-detail-overlay" style="display:none" role="dialog" aria-modal="true" aria-labelledby="mcp-detail-title">
<div id="mcp-detail-box" class="admin-modal admin-modal-wide mcp-detail-modal">
<h2 id="mcp-detail-title">MCP Server Detail</h2>
<div id="mcp-detail-content"></div>
<div class="modal-buttons">
<button class="modal-cancel" onclick="hideMcpDetailModal()">Close</button>
</div>
</div>
</div>
<script src="/static/admin.js"></script>
<script src="/static/governance.js"></script>
<script src="/static/app.js"></script>
+145 -3
View File
@@ -627,8 +627,9 @@
font-weight: 400;
text-transform: none;
letter-spacing: 0;
opacity: 0.6;
opacity: 0.75;
}
#new-ws-box label:first-of-type { margin-top: 0; }
#new-ws-box select,
#new-ws-box input[type="text"],
#new-ws-box textarea {
@@ -642,7 +643,7 @@
font-size: 13px;
transition: border-color 0.15s, box-shadow 0.15s;
}
#new-ws-box textarea { resize: vertical; min-height: 40px; }
#new-ws-box textarea { resize: vertical; min-height: 60px; }
#new-ws-box textarea::placeholder { color: var(--fg-dim); opacity: 0.6; }
#new-ws-box select:focus,
#new-ws-box input:focus,
@@ -688,6 +689,7 @@
letter-spacing: 0.02em;
}
#new-ws-cancel:hover { background: var(--bg-elevated); border-color: var(--border-strong); }
#new-ws-cancel:focus-visible { outline: 2px solid var(--accent); outline-offset: 2px; }
#new-ws-submit {
background: var(--accent);
color: var(--bg);
@@ -1173,7 +1175,8 @@
#create-policy-overlay, #edit-policy-overlay,
#create-template-overlay, #edit-template-overlay,
#create-wst-overlay, #edit-wst-overlay, #wst-history-overlay,
#memory-detail-overlay {
#memory-detail-overlay,
#mcp-create-overlay, #mcp-import-overlay, #mcp-detail-overlay {
position: fixed;
inset: 0;
background: rgba(0, 0, 0, 0.7);
@@ -1851,6 +1854,143 @@
.settings-input select { max-width: 100%; }
}
/* -- MCP Servers grid ----------------------------------------------------- */
.admin-col-mname a{color:var(--fg);text-decoration:none;transition:color .15s}
.admin-col-mname a:hover{color:var(--magenta)}
.admin-col-mname a:focus-visible{outline:2px solid var(--magenta);outline-offset:2px}
.mcp-grid{grid-template-columns:1.5fr 80px 55px 45px 80px 95px 120px;gap:0 6px}
@media(max-width:700px){
.mcp-grid{grid-template-columns:1fr 100px 130px}
.admin-col-mtransport,.admin-col-mtools,.admin-col-mres,.admin-col-mprompts{display:none}
}
.mcp-status-dot{display:inline-block;width:8px;height:8px;border-radius:50%;vertical-align:middle;margin-right:6px}
.mcp-status-dot.connected{background:var(--magenta);box-shadow:0 0 6px var(--magenta-glow, rgba(192,132,252,.45))}
.mcp-status-dot.error{background:var(--red);box-shadow:0 0 6px var(--red-glow);border-radius:1px}
.mcp-status-dot.disabled{background:var(--fg-dim);opacity:.35}
.mcp-status-dot.connecting{background:var(--magenta);animation:mcp-pulse 1.2s ease-in-out infinite}
@keyframes mcp-pulse{0%,100%{opacity:.3}50%{opacity:1}}
.mcp-row-connected{border-left:3px solid var(--magenta)}
.mcp-row-error{border-left:3px solid var(--red)}
.mcp-row-disabled{border-left:3px solid transparent}
.mcp-transport-badge{display:inline-block;font-size:9px;font-weight:600;text-transform:uppercase;letter-spacing:.06em;padding:1px 6px;border-radius:2px;background:var(--bg-highlight);border:1px solid var(--border)}
.mcp-transport-stdio{color:var(--cyan);border-color:rgba(103,232,249,.2)}
.mcp-transport-http{color:var(--magenta);border-color:rgba(192,132,252,.25)}
.admin-col-mtools,.admin-col-mres,.admin-col-mprompts{text-align:right;font-variant-numeric:tabular-nums}
.mcp-count-dim{opacity:.4}
.mcp-detail-modal::before{background:linear-gradient(90deg,transparent,var(--magenta),transparent)!important}
.mcp-detail-modal h2{color:var(--magenta)!important}
.mcp-detail-section{margin-top:16px}
.mcp-detail-section h3{font-size:11px;font-weight:600;text-transform:uppercase;letter-spacing:.08em;color:var(--magenta);margin-bottom:8px}
.mcp-detail-list{list-style:none;padding:0;margin:0}
.mcp-detail-list li{font-size:12px;padding:3px 0;border-bottom:1px solid var(--border);color:var(--fg-dim)}
.mcp-detail-list li:last-child{border-bottom:none}
.admin-action-btn-ghost{background:transparent;color:var(--fg-dim);border:1px solid var(--border-strong)}
.admin-action-btn-ghost:hover{color:var(--fg);background:var(--bg-highlight)}
/* ==========================================================================
OIDC detail panel (inline expansion below user row)
========================================================================== */
.oidc-detail-panel {
max-height: 0;
overflow: hidden;
transition: max-height 150ms ease;
margin: 0 8px 0 24px;
}
.oidc-detail-inner {
border: 1px dashed var(--border);
border-radius: var(--radius-sm);
padding: 12px 16px;
margin-bottom: 8px;
background: var(--row-alt);
}
.oidc-detail-header {
font-family: var(--font-display);
font-size: 10px;
text-transform: uppercase;
letter-spacing: 0.08em;
color: var(--fg-dim);
margin-bottom: 8px;
}
.oidc-detail-header::before {
content: "\25c6 ";
color: var(--accent);
}
.oidc-identity-row {
display: grid;
grid-template-columns: 70px 100px 1fr 60px 50px;
gap: 8px;
padding: 5px 0;
font-size: 12px;
align-items: center;
}
.oidc-identity-row + .oidc-identity-row {
border-top: 1px solid var(--border);
}
.oidc-identity-issuer {
overflow: hidden;
}
.oidc-identity-issuer .scope-badge {
font-size: 10px;
max-width: 100%;
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
}
.oidc-identity-subject {
font-family: var(--font-mono);
font-size: 11px;
color: var(--fg-dim);
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
}
.oidc-identity-email {
color: var(--fg-dim);
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
}
.oidc-identity-time {
color: var(--fg-dim);
font-size: 11px;
}
.oidc-identity-actions .admin-btn-danger { font-size: 11px; }
.oidc-detail-empty {
color: var(--fg-dim);
font-size: 12px;
font-style: italic;
}
/* Expand indicator on user rows */
.admin-row[data-expandable] { cursor: pointer; }
.admin-row[data-expandable]:hover { background: var(--bg-highlight); }
.admin-row[data-expandable]:hover .admin-expand-indicator { color: var(--fg); }
.admin-row[data-expandable]:focus-visible {
outline: 2px solid var(--accent);
outline-offset: -2px;
}
.admin-expand-indicator {
display: inline-block;
width: 14px;
font-size: 10px;
color: var(--fg-dim);
transition: transform 150ms ease;
transform-origin: center;
}
.admin-expand-indicator.expanded { transform: rotate(90deg); }
@media (max-width: 700px) {
.oidc-identity-row { grid-template-columns: 70px 1fr 50px; }
.oidc-identity-email, .oidc-identity-time { display: none; }
.oidc-detail-panel { margin-left: 8px; }
}
/* ==========================================================================
Reduced motion console-specific
========================================================================== */
@@ -1868,4 +2008,6 @@
#view-admin { animation: none; }
.admin-action-btn, .modal-cancel, .modal-submit { transition: none; }
.admin-modal input, .admin-modal select { transition: none; }
.mcp-status-dot.connecting { animation: none; }
.oidc-detail-panel, .admin-expand-indicator { transition: none; }
}
+255 -10
View File
@@ -28,6 +28,7 @@ import re
import secrets
import threading
import time
import urllib.parse
import uuid
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
@@ -37,6 +38,8 @@ if TYPE_CHECKING:
from starlette.responses import JSONResponse, Response
from starlette.types import ASGIApp, Receive, Scope, Send
from turnstone.core.oidc import OIDCConfig
log = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
@@ -150,6 +153,8 @@ PUBLIC_PATHS: frozenset[str] = frozenset(
"/api/auth/logout",
"/api/auth/status",
"/api/auth/setup",
"/api/auth/oidc/authorize",
"/api/auth/oidc/callback",
}
)
PUBLIC_PREFIXES: tuple[str, ...] = ("/static/", "/shared/")
@@ -167,7 +172,9 @@ WRITE_PATHS: frozenset[str] = frozenset(
}
)
APPROVE_PATHS: frozenset[str] = frozenset({"/api/approve", "/api/_internal/config-reload"})
APPROVE_PATHS: frozenset[str] = frozenset(
{"/api/approve", "/api/_internal/config-reload", "/api/_internal/mcp-reload"}
)
ADMIN_PREFIX = "/api/admin/"
@@ -256,10 +263,20 @@ def hash_password(password: str) -> str:
def verify_password(password: str, password_hash: str) -> bool:
"""Verify a password against a bcrypt hash."""
"""Verify a password against a bcrypt hash.
Returns ``False`` immediately for non-bcrypt hashes (e.g. the ``!oidc``
sentinel used for OIDC-provisioned users) to avoid ``ValueError`` from
``bcrypt.checkpw``.
"""
import bcrypt
return bcrypt.checkpw(password.encode("utf-8"), password_hash.encode("utf-8"))
if not password_hash.startswith("$2"):
return False # Not a bcrypt hash (e.g. OIDC sentinel)
try:
return bcrypt.checkpw(password.encode("utf-8"), password_hash.encode("utf-8"))
except (ValueError, TypeError):
return False
def parse_scopes(scopes_str: str) -> frozenset[str]:
@@ -907,6 +924,13 @@ async def handle_auth_login(request: Request, audience: str) -> Response:
password = body.get("password", "")
if username and password and storage is not None:
# Enforce OIDC-only mode: reject password login when disabled
oidc_config = getattr(request.app.state, "oidc_config", None)
if oidc_config and oidc_config.enabled and not oidc_config.password_enabled:
return JSONResponse(
{"error": "Password login is disabled — use SSO"},
status_code=403,
)
user = storage.get_user_by_username(username)
if user and verify_password(password, user["password_hash"]):
# Derive scopes and permissions from assigned roles
@@ -988,13 +1012,21 @@ async def handle_auth_status(request: Request) -> Response:
except Exception:
pass
return JSONResponse(
{
"auth_enabled": auth_config.enabled,
"has_users": has_users,
"setup_required": auth_config.enabled and not has_users,
}
)
# OIDC configuration
oidc_config = getattr(request.app.state, "oidc_config", None)
oidc_enabled = bool(oidc_config and oidc_config.enabled)
resp: dict[str, Any] = {
"auth_enabled": auth_config.enabled,
"has_users": has_users,
"setup_required": auth_config.enabled and not has_users,
}
if oidc_enabled and oidc_config is not None:
resp["oidc_enabled"] = True
resp["oidc_provider_name"] = oidc_config.provider_name
resp["password_enabled"] = oidc_config.password_enabled
return JSONResponse(resp)
async def handle_auth_setup(request: Request, audience: str) -> Response:
@@ -1095,3 +1127,216 @@ async def handle_auth_setup(request: Request, audience: str) -> Response:
if jwt_token:
response.headers["Set-Cookie"] = make_set_cookie(jwt_token, secure=secure)
return response
async def handle_auth_whoami(request: Request) -> Response:
"""Shared ``GET /api/auth/whoami`` handler — return authenticated user info."""
from starlette.responses import JSONResponse
auth_result: AuthResult | None = getattr(request.state, "auth_result", None)
if not auth_result or not auth_result.user_id:
return JSONResponse({"error": "Not authenticated"}, status_code=401)
resp: dict[str, str] = {
"user_id": auth_result.user_id,
}
if auth_result.permissions:
resp["permissions"] = ",".join(sorted(auth_result.permissions))
return JSONResponse(resp)
def _build_oidc_redirect_uri(request: Request, oidc_config: OIDCConfig) -> str:
"""Build the OIDC callback redirect URI.
Uses ``redirect_base`` from OIDC config when set (recommended for
reverse-proxy deployments), otherwise falls back to the request Host header.
"""
if oidc_config.redirect_base:
return f"{oidc_config.redirect_base}/v1/api/auth/oidc/callback"
scheme = "https" if is_secure_request(dict(request.headers), request.url.scheme) else "http"
host = request.headers.get("host", "localhost")
return f"{scheme}://{host}/v1/api/auth/oidc/callback"
async def handle_oidc_authorize(request: Request, audience: str) -> Response:
"""Shared ``GET /api/auth/oidc/authorize`` handler — redirect to IdP."""
from starlette.responses import JSONResponse, RedirectResponse
oidc_config = getattr(request.app.state, "oidc_config", None)
if not oidc_config or not oidc_config.enabled:
return JSONResponse({"error": "OIDC not configured"}, status_code=404)
# Rate limit — prevents flooding oidc_pending_states table
login_limiter: LoginRateLimiter | None = getattr(request.app.state, "login_limiter", None)
client_ip = request.client.host if request.client else "unknown"
if login_limiter is not None:
ip_ok, _ip_retry = login_limiter.check(f"ip:{client_ip}")
if not ip_ok:
return RedirectResponse("/?oidc_error=Too+many+login+attempts", status_code=302)
login_limiter.record(f"ip:{client_ip}") # Count every authorize to bound pending states
storage = getattr(request.app.state, "auth_storage", None)
if storage is None:
return JSONResponse({"error": "Storage not available"}, status_code=503)
# Require setup to be complete before allowing OIDC login
try:
users = storage.list_users()
except Exception:
return JSONResponse({"error": "Storage unavailable"}, status_code=503)
if not users:
return JSONResponse(
{"error": "Initial setup required before OIDC login"},
status_code=403,
)
from turnstone.core.oidc import build_authorize_url, generate_pkce_pair
state = secrets.token_urlsafe(32)
nonce = secrets.token_urlsafe(32)
code_verifier, _code_challenge = generate_pkce_pair()
# Store pending state in database
storage.create_oidc_pending_state(state, nonce, code_verifier, audience)
# Build redirect URI (pinned by TURNSTONE_OIDC_REDIRECT_BASE when set)
redirect_uri = _build_oidc_redirect_uri(request, oidc_config)
url = build_authorize_url(oidc_config, redirect_uri, state, nonce, code_verifier)
return RedirectResponse(url, status_code=302)
async def handle_oidc_callback(request: Request, audience: str) -> Response:
"""Shared ``GET /api/auth/oidc/callback`` handler — exchange code, provision user, issue JWT."""
from starlette.responses import JSONResponse, RedirectResponse
oidc_config = getattr(request.app.state, "oidc_config", None)
if not oidc_config or not oidc_config.enabled:
return JSONResponse({"error": "OIDC not configured"}, status_code=404)
storage = getattr(request.app.state, "auth_storage", None)
jwt_secret = getattr(request.app.state, "jwt_secret", "")
if storage is None:
return JSONResponse({"error": "Storage not available"}, status_code=503)
# Rate limiting
login_limiter: LoginRateLimiter | None = getattr(request.app.state, "login_limiter", None)
client_ip = request.client.host if request.client else "unknown"
if login_limiter is not None:
ip_ok, ip_retry = login_limiter.check(f"ip:{client_ip}")
if not ip_ok:
return RedirectResponse("/?oidc_error=Too+many+login+attempts", status_code=302)
# Lazy cleanup of expired pending states
with contextlib.suppress(Exception):
storage.cleanup_expired_oidc_states(300)
def _record_oidc_failure() -> None:
if login_limiter is not None:
login_limiter.record(f"ip:{client_ip}")
# Check for IdP error
error = request.query_params.get("error", "")
if error:
_record_oidc_failure()
desc = request.query_params.get("error_description", error)
return RedirectResponse(f"/?oidc_error={urllib.parse.quote(desc)}", status_code=302)
# Validate state
state = request.query_params.get("state", "")
pending = storage.pop_oidc_pending_state(state, max_age_seconds=300)
if not pending:
_record_oidc_failure()
return RedirectResponse("/?oidc_error=Login+session+expired", status_code=302)
# Build redirect URI (must match what was sent in authorize)
redirect_uri = _build_oidc_redirect_uri(request, oidc_config)
try:
from turnstone.core.oidc import (
OIDCError,
exchange_code,
fetch_jwks,
provision_oidc_user,
validate_id_token,
)
# Exchange code for tokens
code = request.query_params.get("code", "")
tokens = await exchange_code(oidc_config, code, redirect_uri, pending["code_verifier"])
# Validate ID token against cached JWKS keys (no I/O).
# On unknown kid, refresh JWKS once (async) for key rotation.
jwks_data: dict[str, Any] | None = getattr(request.app.state, "jwks_data", None)
if jwks_data is None and oidc_config.jwks_uri:
# Lazy fetch: JWKS may have failed at startup but IdP recovered
try:
jwks_data = await fetch_jwks(oidc_config.jwks_uri)
request.app.state.jwks_data = jwks_data
except OIDCError:
pass
if jwks_data is None:
return RedirectResponse("/?oidc_error=OIDC+temporarily+unavailable", status_code=302)
try:
id_claims = validate_id_token(
tokens["id_token"],
jwks_data,
oidc_config,
pending["nonce"],
)
except OIDCError as first_err:
if "not found in JWKS" not in str(first_err):
raise
# Key rotation: re-fetch JWKS and retry once.
log.info("JWKS key not found — refreshing for possible key rotation")
jwks_data = await fetch_jwks(oidc_config.jwks_uri)
request.app.state.jwks_data = jwks_data
id_claims = validate_id_token(
tokens["id_token"],
jwks_data,
oidc_config,
pending["nonce"],
)
# Verify setup is complete
users = storage.list_users()
if not users:
return RedirectResponse("/?oidc_error=Initial+setup+required", status_code=302)
# Provision or match user
user = provision_oidc_user(storage, oidc_config, id_claims)
except OIDCError as exc:
log.warning("OIDC callback failed: %s", exc)
_record_oidc_failure()
return RedirectResponse("/?oidc_error=Authentication+failed", status_code=302)
except Exception:
log.exception("OIDC callback error")
_record_oidc_failure()
return RedirectResponse("/?oidc_error=Authentication+failed", status_code=302)
# Load permissions and issue Turnstone JWT
perms = _load_user_permissions(storage, user["user_id"])
scopes = _permissions_to_scopes(perms)
jwt_token = ""
if jwt_secret:
# Use the audience stored during authorize (not the handler param)
# to bind the JWT to the service that initiated the flow
jwt_audience = pending.get("audience", audience)
jwt_token = create_jwt(
user_id=user["user_id"],
scopes=scopes,
source="oidc",
secret=jwt_secret,
audience=jwt_audience,
permissions=frozenset(perms),
)
# Set cookie and redirect to app
response = RedirectResponse("/?oidc_success=1", status_code=302)
if jwt_token:
secure = is_secure_request(dict(request.headers), request.url.scheme)
response.headers["Set-Cookie"] = make_set_cookie(jwt_token, secure=secure)
return response
+358 -35
View File
@@ -97,6 +97,7 @@ class MCPClientManager:
self._loop: asyncio.AbstractEventLoop | None = None
self._thread: threading.Thread | None = None
self._exit_stack: AsyncExitStack | None = None
self._per_server_stacks: dict[str, AsyncExitStack] = {}
self._sessions: dict[str, Any] = {}
self._tools: list[dict[str, Any]] = []
@@ -104,6 +105,13 @@ class MCPClientManager:
self._tool_map: dict[str, tuple[str, str]] = {}
self._connected = threading.Event()
self._error: str | None = None
# Names managed by the DB (added via reconcile_sync / add_server_sync).
# Config-file servers loaded at startup are NOT in this set and
# will never be removed by reconcile_sync.
self._db_managed: set[str] = set()
# Per-server last-error tracking (set on failure, cleared on success)
self._last_error: dict[str, str] = {}
self._MAX_ERROR_LEN = 256
# Per-server tool storage for surgical refresh
self._per_server_tools: dict[str, list[dict[str, Any]]] = {}
@@ -166,8 +174,9 @@ class MCPClientManager:
for name, cfg in self._server_configs.items():
try:
await self._connect_one(name, cfg)
except Exception:
except Exception as exc:
log.warning("Failed to connect MCP server '%s'", name, exc_info=True)
self._set_error(name, f"{type(exc).__name__}: {exc}")
self._connected.set()
@@ -189,30 +198,37 @@ class MCPClientManager:
async def _connect_one(self, name: str, cfg: dict[str, Any]) -> None:
"""Connect to a single MCP server and discover its tools."""
assert self._exit_stack is not None
if "__" in name:
log.error("MCP server name '%s' contains '__' (reserved delimiter), skipping", name)
return
# Per-server exit stack for clean per-server lifecycle management
stack = AsyncExitStack()
await stack.__aenter__()
transport = cfg.get("type", "stdio")
if transport in ("http", "streamable-http") or "url" in cfg:
read, write, _ = await self._exit_stack.enter_async_context(
streamablehttp_client(url=cfg["url"], headers=cfg.get("headers"))
)
else:
# Default: stdio transport
command = cfg.get("command", "")
if not command:
log.warning("MCP server '%s' has no command configured", name)
return
env = {**os.environ, **cfg.get("env", {})}
params = StdioServerParameters(
command=command,
args=cfg.get("args", []),
env=env,
)
read, write = await self._exit_stack.enter_async_context(stdio_client(params))
try:
if transport in ("http", "streamable-http") or "url" in cfg:
read, write, _ = await stack.enter_async_context(
streamablehttp_client(url=cfg["url"], headers=cfg.get("headers"))
)
else:
# Default: stdio transport
command = cfg.get("command", "")
if not command:
log.warning("MCP server '%s' has no command configured", name)
await stack.aclose()
return
env = {**os.environ, **cfg.get("env", {})}
params = StdioServerParameters(
command=command,
args=cfg.get("args", []),
env=env,
)
read, write = await stack.enter_async_context(stdio_client(params))
except Exception:
await stack.aclose()
raise
# Register notification handler — dispatches tool, resource, and
# prompt list-change notifications to the appropriate refresh method.
@@ -232,13 +248,27 @@ class MCPClientManager:
elif isinstance(root, mcp_types.PromptListChangedNotification):
log.info("Received prompts/list_changed from '%s'", name)
await self._refresh_server_prompts(name)
except Exception:
self._last_error.pop(name, None)
except Exception as exc:
log.warning("Refresh after notification failed for '%s'", name, exc_info=True)
self._set_error(name, f"Refresh failed: {exc}")
session = await self._exit_stack.enter_async_context(
ClientSession(read, write, message_handler=_on_notification) # type: ignore[arg-type]
)
await session.initialize()
try:
session = await stack.enter_async_context(
ClientSession(read, write, message_handler=_on_notification) # type: ignore[arg-type]
)
except Exception:
await stack.aclose()
raise
self._per_server_stacks[name] = stack
try:
await session.initialize()
except Exception:
self._per_server_stacks.pop(name, None)
with contextlib.suppress(Exception):
await stack.aclose()
raise
self._sessions[name] = session
# Check push notification support for each capability
@@ -348,6 +378,9 @@ class MCPClientManager:
except Exception:
log.warning("Prompt sync after connect failed for '%s'", name, exc_info=True)
# Connection succeeded — clear any previous error
self._last_error.pop(name, None)
# -- tool refresh --------------------------------------------------------
def _rebuild_tools(self) -> None:
@@ -404,6 +437,7 @@ class MCPClientManager:
added, removed = await self._refresh_server_tools(name)
await self._refresh_server_resources(name)
await self._refresh_server_prompts(name)
self._last_error.pop(name, None)
return added, removed
async def _refresh_all(
@@ -432,8 +466,9 @@ class MCPClientManager:
continue
added, removed = await self._refresh_server(name)
results[name] = (added, removed)
except Exception:
except Exception as exc:
log.warning("Refresh failed for MCP server '%s'", name, exc_info=True)
self._set_error(name, f"Refresh failed: {exc}")
results[name] = ([], [])
# Final sync to clean up templates from servers that are no longer connected
@@ -473,8 +508,10 @@ class MCPClientManager:
await self._refresh_server_resources(name)
if not self._supports_prompt_list_changed.get(name, False):
await self._refresh_server_prompts(name)
except Exception:
self._last_error.pop(name, None)
except Exception as exc:
log.warning("Periodic refresh failed for '%s'", name, exc_info=True)
self._set_error(name, f"Periodic refresh failed: {exc}")
await asyncio.sleep(self._refresh_interval)
# -- resource refresh ----------------------------------------------------
@@ -807,12 +844,27 @@ class MCPClientManager:
if self._refresh_task and self._loop:
self._loop.call_soon_threadsafe(self._refresh_task.cancel)
# Close all per-server stacks (transports + sessions)
if self._loop and self._per_server_stacks:
async def _close_all_stacks() -> None:
for stack in self._per_server_stacks.values():
with contextlib.suppress(Exception):
await stack.aclose()
future = asyncio.run_coroutine_threadsafe(_close_all_stacks(), self._loop)
try:
future.result(timeout=10)
except Exception:
log.debug("Error closing MCP sessions", exc_info=True)
# Close legacy shared stack (if any resources were registered on it)
if self._loop and self._exit_stack:
future = asyncio.run_coroutine_threadsafe(self._exit_stack.aclose(), self._loop)
try:
future.result(timeout=10)
except Exception:
log.debug("Error closing MCP sessions", exc_info=True)
log.debug("Error closing MCP exit stack", exc_info=True)
if self._loop:
self._loop.call_soon_threadsafe(self._loop.stop)
@@ -821,6 +873,8 @@ class MCPClientManager:
# Clear all state
self._sessions.clear()
self._per_server_stacks.clear()
self._db_managed.clear()
self._tools = []
self._tool_map = {}
self._per_server_tools.clear()
@@ -843,6 +897,216 @@ class MCPClientManager:
log.info("MCP client shut down")
# -- hot-reload (add/remove servers) ------------------------------------
def add_server_sync(self, name: str, cfg: dict[str, Any], timeout: int = 30) -> dict[str, Any]:
"""Connect a new MCP server at runtime (blocks the calling thread).
Returns status dict with keys: connected, tools, resources, prompts, error.
"""
if "__" in name:
return {
"connected": False,
"tools": 0,
"resources": 0,
"prompts": 0,
"error": f"Server name '{name}' contains '__' (reserved delimiter)",
}
if self._loop is None:
return {
"connected": False,
"tools": 0,
"resources": 0,
"prompts": 0,
"error": "MCP event loop not running",
}
# Add to config so _refresh_all can reconnect on failure
self._server_configs[name] = cfg
future = asyncio.run_coroutine_threadsafe(self._connect_one(name, cfg), self._loop)
try:
future.result(timeout=timeout)
except Exception as exc:
# Remove from configs on failure
self._server_configs.pop(name, None)
return {"connected": False, "tools": 0, "resources": 0, "prompts": 0, "error": str(exc)}
return {
"connected": name in self._sessions,
"tools": len(self._per_server_tools.get(name, [])),
"resources": len(self._per_server_resources.get(name, [])),
"prompts": len(self._per_server_prompts.get(name, [])),
"error": "",
}
def remove_server_sync(self, name: str, timeout: int = 15) -> bool:
"""Disconnect and remove an MCP server at runtime (blocks the calling thread).
All state mutations run on the MCP event loop thread to avoid races
with notification handlers and refresh tasks.
Returns True if the server was connected and successfully removed.
"""
was_connected = name in self._sessions
# Remove from config to prevent reconnection
self._server_configs.pop(name, None)
if self._loop is not None:
async def _remove() -> None:
# Close session + transport via per-server stack
self._sessions.pop(name, None)
stack = self._per_server_stacks.pop(name, None)
if stack is not None:
with contextlib.suppress(Exception):
await stack.aclose()
# Clean up per-server state (on the event loop thread)
self._per_server_tools.pop(name, None)
self._per_server_resources.pop(name, None)
self._per_server_prompts.pop(name, None)
self._supports_list_changed.pop(name, None)
self._supports_resources.pop(name, None)
self._supports_resource_list_changed.pop(name, None)
self._supports_prompts.pop(name, None)
self._supports_prompt_list_changed.pop(name, None)
self._last_error.pop(name, None)
# Rebuild merged state (serialized with notification handlers)
self._rebuild_tools()
self._rebuild_resources()
self._rebuild_prompts()
future = asyncio.run_coroutine_threadsafe(_remove(), self._loop)
try:
future.result(timeout=timeout)
except Exception:
log.warning("Error removing MCP server '%s'", name, exc_info=True)
else:
# No event loop (tests / pre-start) — mutate directly
self._sessions.pop(name, None)
self._per_server_tools.pop(name, None)
self._per_server_resources.pop(name, None)
self._per_server_prompts.pop(name, None)
self._supports_list_changed.pop(name, None)
self._supports_resources.pop(name, None)
self._supports_resource_list_changed.pop(name, None)
self._supports_prompts.pop(name, None)
self._supports_prompt_list_changed.pop(name, None)
self._last_error.pop(name, None)
self._rebuild_tools()
self._rebuild_resources()
self._rebuild_prompts()
# Clean up governance templates from this server
try:
self.sync_prompts_to_storage()
except Exception:
log.warning("Prompt sync after remove failed for '%s'", name, exc_info=True)
log.info("Removed MCP server '%s'", name)
return was_connected
def _set_error(self, name: str, msg: str) -> None:
"""Store a sanitized error string for a server."""
clean = msg.replace("\n", " ").replace("\r", "")
self._last_error[name] = clean[: self._MAX_ERROR_LEN]
def get_server_status(self, name: str) -> dict[str, Any]:
"""Return live status for a single server, including config details."""
connected = name in self._sessions
cfg = self._server_configs.get(name, {})
transport = cfg.get("type", "stdio")
return {
"connected": connected,
"tools": len(self._per_server_tools.get(name, [])) if connected else 0,
"resources": len(self._per_server_resources.get(name, [])) if connected else 0,
"prompts": len(self._per_server_prompts.get(name, [])) if connected else 0,
"error": self._last_error.get(name, ""),
"transport": transport,
"command": cfg.get("command", "") if transport == "stdio" else "",
"url": cfg.get("url", "") if transport != "stdio" else "",
}
def get_all_server_status(self) -> dict[str, dict[str, Any]]:
"""Return live status for all configured servers."""
result: dict[str, dict[str, Any]] = {}
for name in list(self._server_configs):
result[name] = self.get_server_status(name)
return result
def reconcile_sync(self, storage: Any, timeout: int = 30) -> dict[str, Any]:
"""Reconcile DB-managed servers against DB state.
Reads enabled ``mcp_servers`` rows from *storage*, then:
- Connects servers in DB but not currently running.
- Disconnects DB-managed servers no longer in DB (or disabled).
- Reconnects DB-managed servers whose config has changed.
Config-file servers (loaded at startup, not in ``_db_managed``)
are never touched only servers previously added via DB are
eligible for removal.
Returns ``{"added": [...], "removed": [...], "updated": [...]}``.
"""
try:
rows = storage.list_mcp_servers(enabled_only=True)
except Exception:
log.warning("reconcile_sync: failed to read mcp_servers table", exc_info=True)
return {"added": [], "removed": [], "updated": []}
desired = _db_servers_to_config(rows)
desired_names = set(desired)
added: list[str] = []
removed: list[str] = []
updated: list[str] = []
# Remove DB-managed servers no longer in DB (or disabled).
# Config-file servers (not in _db_managed) are left untouched.
for name in list(self._db_managed - desired_names):
self.remove_server_sync(name, timeout=timeout)
self._db_managed.discard(name)
removed.append(name)
# Add servers in DB but not running
for name in desired_names - set(self._server_configs):
result = self.add_server_sync(name, desired[name], timeout=timeout)
if result.get("connected"):
added.append(name)
self._db_managed.add(name)
else:
log.warning("reconcile_sync: failed to add '%s': %s", name, result.get("error", ""))
# Update DB-managed servers whose config has changed (cycle: remove + add).
# Config-file servers with the same name as a DB server are left untouched.
for name in desired_names & set(self._server_configs):
if name not in self._db_managed:
continue # config-file server — DB doesn't own it
if desired[name] != self._server_configs.get(name):
log.info("Config changed for MCP server '%s', reconnecting", name)
self.remove_server_sync(name, timeout=timeout)
result = self.add_server_sync(name, desired[name], timeout=timeout)
if result.get("connected"):
updated.append(name)
self._db_managed.add(name)
else:
self._db_managed.discard(name)
log.warning(
"reconcile_sync: failed to reconnect '%s': %s",
name,
result.get("error", ""),
)
if added or removed or updated:
log.info(
"MCP reconcile: +%d added, -%d removed, ~%d updated",
len(added),
len(removed),
len(updated),
)
return {"added": added, "removed": removed, "updated": updated}
# -- query methods -------------------------------------------------------
def get_tools(self) -> list[dict[str, Any]]:
@@ -879,6 +1143,11 @@ class MCPClientManager:
def server_count(self) -> int:
return len(self._sessions)
@property
def error_count(self) -> int:
"""Number of servers currently in error state."""
return len(self._last_error)
@property
def server_names(self) -> list[str]:
"""Return configured server names."""
@@ -1024,23 +1293,64 @@ class MCPClientManager:
# ---------------------------------------------------------------------------
def load_mcp_config(config_path: str | None = None) -> dict[str, dict[str, Any]]:
def _db_servers_to_config(rows: list[dict[str, Any]]) -> dict[str, dict[str, Any]]:
"""Convert mcp_servers DB rows to the config dict format."""
result: dict[str, dict[str, Any]] = {}
for row in rows:
name = row["name"]
cfg: dict[str, Any] = {"type": row["transport"]}
if row["transport"] == "stdio":
cfg["command"] = row.get("command", "")
try:
cfg["args"] = json.loads(row.get("args", "[]"))
except (json.JSONDecodeError, TypeError):
cfg["args"] = []
try:
cfg["env"] = json.loads(row.get("env", "{}"))
except (json.JSONDecodeError, TypeError):
cfg["env"] = {}
else:
cfg["url"] = row.get("url", "")
try:
cfg["headers"] = json.loads(row.get("headers", "{}"))
except (json.JSONDecodeError, TypeError):
cfg["headers"] = {}
result[name] = cfg
return result
def load_mcp_config(
config_path: str | None = None,
storage: Any = None,
) -> dict[str, dict[str, Any]]:
"""Load MCP server configurations.
Sources (first match wins):
1. Explicit *config_path* (standard MCP JSON format).
2. ``[mcp.servers.*]`` sections in ``config.toml``.
1. DB ``mcp_servers`` table (if *storage* provided and has enabled rows).
2. Explicit *config_path* (standard MCP JSON format).
3. ``[mcp.servers.*]`` sections in ``config.toml``.
Returns an empty dict if nothing is configured.
"""
# 1. Explicit JSON file
# 1. Database
if storage is not None:
try:
rows = storage.list_mcp_servers(enabled_only=True)
if rows:
servers = _db_servers_to_config(rows)
log.info("Loaded MCP config from database (%d server(s))", len(servers))
return servers
except Exception:
log.debug("DB MCP config lookup failed (table may not exist yet)", exc_info=True)
# 2. Explicit JSON file
if config_path:
path = Path(config_path).expanduser()
if path.is_file():
try:
data = json.loads(path.read_text(encoding="utf-8"))
servers: dict[str, Any] = data.get("mcpServers", {})
servers = data.get("mcpServers", {})
if isinstance(servers, dict) and servers:
log.info("Loaded MCP config from %s (%d server(s))", path, len(servers))
return servers
@@ -1049,7 +1359,7 @@ def load_mcp_config(config_path: str | None = None) -> dict[str, dict[str, Any]]
else:
log.warning("MCP config file not found: %s", path)
# 2. TOML config
# 3. TOML config
mcp_section = load_config("mcp")
servers_section = mcp_section.get("servers", {})
@@ -1069,15 +1379,28 @@ def create_mcp_client(
config_path: str | None = None,
*,
refresh_interval: float = _DEFAULT_REFRESH_INTERVAL,
storage: Any = None,
) -> MCPClientManager | None:
"""Create and start an MCP client manager.
Returns *None* if no servers are configured.
"""
servers = load_mcp_config(config_path)
# Check DB first to know which servers are DB-managed
db_names: set[str] = set()
if storage is not None:
try:
rows = storage.list_mcp_servers(enabled_only=True)
if rows:
db_names = {r["name"] for r in rows}
except Exception:
pass
servers = load_mcp_config(config_path, storage=storage)
if not servers:
return None
mgr = MCPClientManager(servers, refresh_interval=refresh_interval)
# Mark DB-sourced servers so reconcile_sync won't remove config-file servers
mgr._db_managed = {name for name in servers if name in db_names}
mgr.start()
return mgr
+47 -8
View File
@@ -53,11 +53,16 @@ _NUDGE_MAP: dict[str, str] = {
}
# ---------------------------------------------------------------------------
# Detection heuristics
# Detection heuristics — strong/weak tiers
#
# Strong patterns fire unconditionally. Weak patterns carry inherent
# ambiguity ("no …", "thanks …") and only fire when the surrounding
# message looks like a genuine correction/completion rather than normal
# conversation.
# ---------------------------------------------------------------------------
_CORRECTION_PATTERNS: list[re.Pattern[str]] = [
re.compile(r"(?i)^no[,.\s]"),
_STRONG_CORRECTION: list[re.Pattern[str]] = [
re.compile(r"(?i)^no[,.]"), # "no," / "no." — clear rejection
re.compile(r"(?i)\bdon'?t\b"),
re.compile(r"(?i)^stop\b"),
re.compile(r"(?i)^actually[,\s]"),
@@ -71,30 +76,64 @@ _CORRECTION_PATTERNS: list[re.Pattern[str]] = [
re.compile(r"(?i)^please don'?t\b"),
]
_COMPLETION_PATTERNS: list[re.Pattern[str]] = [
re.compile(r"(?i)^thanks\b"),
# "no <word>" is ambiguous — only match when the next word is a pronoun,
# demonstrative, article, or verb that signals the user is redirecting,
# not a fixed phrase like "no problem" or "no worries". Allowlist >
# blocklist: we don't need to enumerate every benign "no X" phrase.
_WEAK_CORRECTION: list[re.Pattern[str]] = [
re.compile(
r"(?i)^no\s+(?:I\b|you\b|we\b|they\b|it\b|he\b|she\b"
r"|that\b|this\b|those\b|these\b"
r"|the\b|a\b|an\b"
r"|not\b|do\b|did\b|but\b)"
),
]
_STRONG_COMPLETION: list[re.Pattern[str]] = [
re.compile(r"(?i)\bthat'?s all\b"),
re.compile(r"(?i)^lgtm\b"),
]
# These patterns are common in both completion AND mid-conversation
# acknowledgment. Only fire when the message is short and has no
# continuation markers (question marks, follow-up requests).
_WEAK_COMPLETION: list[re.Pattern[str]] = [
re.compile(r"(?i)^thanks\b(?!\s+for\b)"), # "thanks for X" = acknowledgment
re.compile(r"(?i)\blooks good\b"),
re.compile(r"(?i)^perfect\b"),
re.compile(r"(?i)^great job\b"),
re.compile(r"(?i)\bthat works\b"),
re.compile(r"(?i)^done\b"),
re.compile(r"(?i)^lgtm\b"),
]
_WEAK_MSG_CAP = 80 # weak completion patterns suppressed above this length
_CONTINUATION = re.compile(
r"(?i)(?:\?|(?:can you|could you|please\s|also\s|but\s|now\s|next\s"
r"|and\s+then|after\s+that|one\s+more|however))"
)
def detect_correction(message: str) -> bool:
"""Return True if the message looks like a user correction."""
if not message:
return False
return any(p.search(message) for p in _CORRECTION_PATTERNS)
if any(p.search(message) for p in _STRONG_CORRECTION):
return True
return any(p.search(message) for p in _WEAK_CORRECTION)
def detect_completion(message: str) -> bool:
"""Return True if the message signals session completion."""
if not message:
return False
return any(p.search(message) for p in _COMPLETION_PATTERNS)
if any(p.search(message) for p in _STRONG_COMPLETION):
return True
if len(message) > _WEAK_MSG_CAP:
return False
if _CONTINUATION.search(message):
return False
return any(p.search(message) for p in _WEAK_COMPLETION)
def should_nudge(
+5
View File
@@ -402,6 +402,11 @@ class MetricsCollector:
"Number of MCP prompts available",
mcp_info.get("prompts", 0),
)
gauge(
"turnstone_mcp_server_errors",
"Number of MCP servers currently in error state",
mcp_info.get("errors", 0),
)
lines.append("") # trailing newline
return "\n".join(lines)
+583
View File
@@ -0,0 +1,583 @@
"""OpenID Connect (OIDC) authentication support for Turnstone.
Implements the Authorization Code Flow with PKCE for secure SSO login.
All external HTTP calls use ``httpx.AsyncClient`` to avoid blocking the
event loop.
"""
from __future__ import annotations
import base64
import dataclasses
import hashlib
import logging
import os
import re
import secrets
import urllib.parse
import uuid
from dataclasses import dataclass, field
from typing import Any
import httpx
log = logging.getLogger(__name__)
# Sentinel password hash for OIDC-provisioned users.
# Not a valid bcrypt hash -- verify_password() always rejects it.
OIDC_PASSWORD_SENTINEL = "!oidc"
# Sanitisation pattern: only keep safe username characters.
_USERNAME_SAFE_RE = re.compile(r"[^a-zA-Z0-9._-]")
# Asymmetric algorithms accepted for ID token signatures.
# Symmetric (HMAC) algorithms are deliberately excluded to prevent
# algorithm confusion attacks where the IdP's public key is used as
# an HMAC secret.
_ALLOWED_ID_TOKEN_ALGS = [
"RS256",
"RS384",
"RS512",
"ES256",
"ES384",
"ES512",
"PS256",
"PS384",
"PS512",
]
# ---------------------------------------------------------------------------
# Exception
# ---------------------------------------------------------------------------
class OIDCError(Exception):
"""Raised when an OIDC operation fails."""
# ---------------------------------------------------------------------------
# Configuration
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class OIDCConfig:
"""OIDC provider configuration -- immutable after startup."""
enabled: bool = False
issuer: str = ""
client_id: str = ""
client_secret: str = ""
scopes: str = "openid email profile"
provider_name: str = "SSO"
role_claim: str = ""
role_map: dict[str, str] = field(default_factory=dict)
password_enabled: bool = True
redirect_base: str = ""
# Discovered from .well-known/openid-configuration
authorization_endpoint: str = ""
token_endpoint: str = ""
userinfo_endpoint: str = ""
jwks_uri: str = ""
def _parse_role_map(raw: str) -> dict[str, str]:
"""Parse ``"admin:builtin-admin,eng:builtin-operator"`` into a dict."""
result: dict[str, str] = {}
for pair in raw.split(","):
pair = pair.strip()
if ":" in pair:
k, v = pair.split(":", 1)
k, v = k.strip(), v.strip()
if k and v:
result[k] = v
return result
def load_oidc_config() -> OIDCConfig:
"""Build :class:`OIDCConfig` from env vars with config.toml fallback.
Returns ``OIDCConfig(enabled=False)`` when the required fields
(issuer, client_id, client_secret) are not all present.
"""
from turnstone.core.config import load_config
cfg = load_config("oidc")
# Start with config.toml values, then override with env vars.
issuer = os.environ.get("TURNSTONE_OIDC_ISSUER", "").strip()
if not issuer:
issuer = str(cfg.get("issuer", "")).strip()
client_id = os.environ.get("TURNSTONE_OIDC_CLIENT_ID", "").strip()
if not client_id:
client_id = str(cfg.get("client_id", "")).strip()
client_secret = os.environ.get("TURNSTONE_OIDC_CLIENT_SECRET", "").strip()
if not client_secret:
client_secret = str(cfg.get("client_secret", "")).strip()
scopes = os.environ.get("TURNSTONE_OIDC_SCOPES", "").strip()
if not scopes:
scopes = str(cfg.get("scopes", "openid email profile")).strip()
provider_name = os.environ.get("TURNSTONE_OIDC_PROVIDER_NAME", "").strip()
if not provider_name:
provider_name = str(cfg.get("provider_name", "SSO")).strip()
role_claim = os.environ.get("TURNSTONE_OIDC_ROLE_CLAIM", "").strip()
if not role_claim:
role_claim = str(cfg.get("role_claim", "")).strip()
# Role map: env var is "admin:builtin-admin,eng:builtin-operator"
role_map_raw = os.environ.get("TURNSTONE_OIDC_ROLE_MAP", "").strip()
if role_map_raw:
role_map = _parse_role_map(role_map_raw)
else:
cfg_role_map = cfg.get("role_map", {})
role_map = dict(cfg_role_map) if isinstance(cfg_role_map, dict) else {}
password_raw = os.environ.get("TURNSTONE_OIDC_PASSWORD_ENABLED", "").strip().lower()
if password_raw:
password_enabled = password_raw in ("true", "1", "yes")
else:
password_enabled = bool(cfg.get("password_enabled", True))
redirect_base = os.environ.get("TURNSTONE_OIDC_REDIRECT_BASE", "").strip()
if not redirect_base:
redirect_base = str(cfg.get("redirect_base", "")).strip()
redirect_base = redirect_base.rstrip("/")
if redirect_base:
parsed = urllib.parse.urlparse(redirect_base)
if parsed.scheme not in ("https", "http"):
log.warning(
"TURNSTONE_OIDC_REDIRECT_BASE has invalid scheme, ignoring: %s",
redirect_base,
)
redirect_base = ""
elif not parsed.hostname:
log.warning(
"TURNSTONE_OIDC_REDIRECT_BASE missing hostname, ignoring: %s",
redirect_base,
)
redirect_base = ""
elif parsed.username or parsed.password:
log.warning(
"TURNSTONE_OIDC_REDIRECT_BASE must not contain userinfo, ignoring: %s",
redirect_base,
)
redirect_base = ""
elif parsed.path or parsed.query or parsed.fragment:
log.warning(
"TURNSTONE_OIDC_REDIRECT_BASE must be scheme://host[:port] only, ignoring: %s",
redirect_base,
)
redirect_base = ""
else:
# Validate port is numeric (urlparse accepts "host:abc" silently).
try:
parsed.port # noqa: B018 — triggers ValueError on non-numeric port
except ValueError:
log.warning(
"TURNSTONE_OIDC_REDIRECT_BASE has invalid port, ignoring: %s",
redirect_base,
)
redirect_base = ""
if redirect_base and parsed.scheme != "https":
log.warning(
"TURNSTONE_OIDC_REDIRECT_BASE should use https:// in production: %s",
redirect_base,
)
# OIDC is enabled when all three required fields are non-empty.
enabled = bool(issuer and client_id and client_secret)
if enabled:
log.info("OIDC enabled: issuer=%s provider=%s", issuer, provider_name)
else:
log.debug("OIDC not configured (issuer/client_id/client_secret incomplete)")
return OIDCConfig(
enabled=enabled,
issuer=issuer,
client_id=client_id,
client_secret=client_secret,
scopes=scopes,
provider_name=provider_name,
role_claim=role_claim,
role_map=role_map,
password_enabled=password_enabled,
redirect_base=redirect_base,
)
# ---------------------------------------------------------------------------
# Discovery
# ---------------------------------------------------------------------------
async def discover_oidc(config: OIDCConfig) -> OIDCConfig:
"""Fetch OIDC discovery document and return updated config with endpoints.
On failure, logs a warning and returns config with ``enabled=False``.
"""
if not config.issuer:
return dataclasses.replace(config, enabled=False)
url = config.issuer.rstrip("/") + "/.well-known/openid-configuration"
try:
async with httpx.AsyncClient(timeout=10.0) as client:
resp = await client.get(url)
resp.raise_for_status()
doc = resp.json()
except Exception as exc:
log.warning("OIDC discovery failed for %s: %s", config.issuer, exc)
return dataclasses.replace(config, enabled=False)
authorization_endpoint = str(doc.get("authorization_endpoint", ""))
token_endpoint = str(doc.get("token_endpoint", ""))
userinfo_endpoint = str(doc.get("userinfo_endpoint", ""))
jwks_uri = str(doc.get("jwks_uri", ""))
if not authorization_endpoint or not token_endpoint or not jwks_uri:
log.warning(
"OIDC discovery document missing required endpoints for %s",
config.issuer,
)
return dataclasses.replace(config, enabled=False)
log.info("OIDC discovery complete: %s", config.issuer)
return dataclasses.replace(
config,
authorization_endpoint=authorization_endpoint,
token_endpoint=token_endpoint,
userinfo_endpoint=userinfo_endpoint,
jwks_uri=jwks_uri,
)
# ---------------------------------------------------------------------------
# JWKS key management
# ---------------------------------------------------------------------------
async def fetch_jwks(jwks_uri: str) -> dict[str, Any]:
"""Fetch the JWKS key set from the IdP.
Returns the parsed JSON document (``{"keys": [...]}``) . Called during
startup discovery and on-demand when an unknown ``kid`` is encountered
(key rotation). Uses ``httpx.AsyncClient`` never blocks the event loop.
Raises :class:`OIDCError` on network failures or malformed responses.
"""
try:
async with httpx.AsyncClient(timeout=10.0) as client:
resp = await client.get(jwks_uri)
resp.raise_for_status()
result: dict[str, Any] = resp.json()
except Exception as exc:
raise OIDCError(f"JWKS fetch failed: {exc}") from exc
if not isinstance(result.get("keys"), list):
raise OIDCError("JWKS document missing 'keys' array")
return result
# ---------------------------------------------------------------------------
# PKCE helpers
# ---------------------------------------------------------------------------
def generate_pkce_pair() -> tuple[str, str]:
"""Generate a PKCE code_verifier and code_challenge pair."""
code_verifier = secrets.token_urlsafe(48)
digest = hashlib.sha256(code_verifier.encode("ascii")).digest()
code_challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
return code_verifier, code_challenge
# ---------------------------------------------------------------------------
# Authorization URL
# ---------------------------------------------------------------------------
def build_authorize_url(
config: OIDCConfig,
redirect_uri: str,
state: str,
nonce: str,
code_verifier: str,
) -> str:
"""Build the OIDC authorization URL with PKCE."""
digest = hashlib.sha256(code_verifier.encode("ascii")).digest()
code_challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
params = {
"response_type": "code",
"client_id": config.client_id,
"redirect_uri": redirect_uri,
"scope": config.scopes,
"state": state,
"nonce": nonce,
"code_challenge": code_challenge,
"code_challenge_method": "S256",
}
return config.authorization_endpoint + "?" + urllib.parse.urlencode(params)
# ---------------------------------------------------------------------------
# Token exchange
# ---------------------------------------------------------------------------
async def exchange_code(
config: OIDCConfig,
code: str,
redirect_uri: str,
code_verifier: str,
) -> dict[str, Any]:
"""Exchange authorization code for tokens at the token endpoint.
Raises :class:`OIDCError` on non-200 response.
"""
data = {
"grant_type": "authorization_code",
"code": code,
"redirect_uri": redirect_uri,
"client_id": config.client_id,
"client_secret": config.client_secret,
"code_verifier": code_verifier,
}
try:
async with httpx.AsyncClient(timeout=10.0) as client:
resp = await client.post(config.token_endpoint, data=data)
except Exception as exc:
raise OIDCError(f"Token exchange request failed: {exc}") from exc
if resp.status_code != 200:
raise OIDCError(f"Token endpoint returned {resp.status_code}: {resp.text[:500]}")
result: dict[str, Any] = resp.json()
return result
# ---------------------------------------------------------------------------
# ID token validation
# ---------------------------------------------------------------------------
def validate_id_token(
raw_token: str,
jwks_data: dict[str, Any],
config: OIDCConfig,
nonce: str,
) -> dict[str, Any]:
"""Validate and decode an OIDC ID token. Returns decoded claims.
*jwks_data* is the pre-fetched JWKS document (the ``{"keys": [...]}``
dict). No network I/O happens here the signing key is resolved
locally from the cached key set.
Raises :class:`OIDCError` on validation failure.
"""
import jwt
from jwt import PyJWK
# Extract kid from the token header to find the matching key.
try:
header = jwt.get_unverified_header(raw_token)
except jwt.DecodeError as exc:
raise OIDCError(f"Invalid ID token header: {exc}") from exc
kid = header.get("kid") # None if absent, not ""
# Find matching key in the JWKS by kid.
# PyJWK infers the key's algorithm from the JWKS ``alg``/``kty``
# fields. jwt.decode() requires the token header's ``alg`` to be in
# our _ALLOWED_ID_TOKEN_ALGS allowlist (asymmetric only) AND to match
# the key type — preventing algorithm confusion attacks.
signing_key = None
for key_dict in jwks_data.get("keys", []):
if kid is not None and key_dict.get("kid") == kid:
try:
signing_key = PyJWK(key_dict)
except Exception as exc:
raise OIDCError(f"Failed to parse signing key: {exc}") from exc
break
# Fallback: if token has no kid and JWKS has exactly one key, use it.
if signing_key is None and kid is None:
keys = jwks_data.get("keys", [])
if len(keys) == 1:
try:
signing_key = PyJWK(keys[0])
except Exception as exc:
raise OIDCError(f"Failed to parse signing key: {exc}") from exc
if signing_key is None:
raise OIDCError(f"Signing key '{kid}' not found in JWKS")
try:
claims: dict[str, Any] = jwt.decode(
raw_token,
signing_key.key,
algorithms=_ALLOWED_ID_TOKEN_ALGS,
audience=config.client_id,
issuer=config.issuer,
)
except jwt.InvalidTokenError as exc:
raise OIDCError(f"ID token validation failed: {exc}") from exc
if claims.get("nonce") != nonce:
raise OIDCError("ID token nonce mismatch")
return claims
# ---------------------------------------------------------------------------
# User provisioning
# ---------------------------------------------------------------------------
def provision_oidc_user(
storage: Any,
config: OIDCConfig,
claims: dict[str, Any],
) -> dict[str, str]:
"""Match or create a user from OIDC claims. Returns user dict.
Looks up an existing OIDC identity by (issuer, sub). If found,
updates ``last_login`` and applies role mapping. Otherwise creates
a new user and OIDC identity record.
Raises :class:`OIDCError` if user creation fails.
"""
issuer = config.issuer
sub = str(claims["sub"])
email = str(claims.get("email", ""))
display_name = str(claims.get("name", "") or claims.get("preferred_username", "") or email)
# Try to find existing identity
identity = storage.get_oidc_identity(issuer, sub)
if identity is not None:
user_id = identity["user_id"]
storage.update_oidc_identity_login(issuer, sub)
apply_role_mapping(storage, user_id, claims, config)
user: dict[str, str] | None = storage.get_user(user_id)
if user is None:
raise OIDCError(f"OIDC identity references missing user: {user_id}")
return user
# New user -- derive username
username = _derive_username(storage, claims)
user_id = uuid.uuid4().hex
storage.create_user(user_id, username, display_name, OIDC_PASSWORD_SENTINEL)
storage.create_oidc_identity(issuer, sub, user_id, email)
apply_role_mapping(storage, user_id, claims, config)
# Ensure new OIDC users have at least a default role so they can
# access the application. builtin-viewer grants read-only access.
user_roles = storage.list_user_roles(user_id)
if not user_roles and storage.get_role("builtin-viewer") is not None:
storage.assign_role(user_id, "builtin-viewer", "oidc-default")
created_user: dict[str, str] | None = storage.get_user(user_id)
if created_user is None:
raise OIDCError(f"Failed to retrieve newly created user: {user_id}")
log.info("Provisioned OIDC user: %s (%s) from %s", username, user_id, issuer)
return created_user
def _derive_username(storage: Any, claims: dict[str, Any]) -> str:
"""Derive a unique, valid username from OIDC claims."""
from turnstone.core.auth import is_valid_username
raw = str(claims.get("preferred_username", ""))
if not raw:
email = str(claims.get("email", ""))
raw = email.split("@")[0] if email else ""
if not raw:
raw = "user"
# Sanitise: keep only safe chars, truncate.
sanitised = _USERNAME_SAFE_RE.sub("", raw)[:64]
if not sanitised:
sanitised = "user"
# Check validity and uniqueness.
if is_valid_username(sanitised) and storage.get_user_by_username(sanitised) is None:
return sanitised
# Deduplicate: append suffix.
for suffix in range(2, 11):
candidate = f"{sanitised[:60]}{suffix}"
if is_valid_username(candidate) and storage.get_user_by_username(candidate) is None:
return candidate
# Last resort: full UUID suffix with validation + uniqueness check.
for _ in range(3):
candidate = f"{sanitised[:32]}{uuid.uuid4().hex}"
if not is_valid_username(candidate):
candidate = f"user{uuid.uuid4().hex}"
if storage.get_user_by_username(candidate) is None:
return candidate
raise OIDCError("Failed to generate unique username")
# ---------------------------------------------------------------------------
# Role mapping
# ---------------------------------------------------------------------------
def apply_role_mapping(
storage: Any,
user_id: str,
claims: dict[str, Any],
config: OIDCConfig,
) -> None:
"""Sync Turnstone roles from OIDC claims.
If ``config.role_claim`` is set, reads the corresponding claim value,
normalises it to a list, and maps each value via ``config.role_map``
to a Turnstone role ID. Roles assigned by OIDC on previous logins
that are no longer present in the claims are revoked (IdP demotions
propagate). Roles assigned manually or by other sources are never
touched.
"""
if not config.role_claim or not config.role_map:
return
claim_value = claims.get(config.role_claim)
# Normalise to list (could be string, list, or absent from IdP).
if claim_value is None:
values: list[str] = []
elif isinstance(claim_value, str):
values = [claim_value]
elif isinstance(claim_value, list):
values = [str(v) for v in claim_value]
else:
values = [str(claim_value)]
# Compute the set of roles the IdP says this user should have.
desired_role_ids: set[str] = set()
for value in values:
role_id = config.role_map.get(value)
if role_id and storage.get_role(role_id) is not None:
desired_role_ids.add(role_id)
# Add new roles from claims.
for role_id in desired_role_ids:
storage.assign_role(user_id, role_id, "oidc")
log.debug("Assigned role %s to user %s via OIDC claim", role_id, user_id)
# Revoke OIDC-assigned roles no longer present in claims.
current_roles = storage.list_user_roles(user_id)
for role in current_roles:
if role.get("assigned_by") == "oidc" and role["role_id"] not in desired_role_ids:
storage.unassign_role(user_id, role["role_id"])
log.info(
"Revoked role %s from user %s (removed from IdP claims)", role["role_id"], user_id
)
+24 -11
View File
@@ -296,7 +296,7 @@ class ChatSession:
self._watch_dispatch_depth = 0
# Metacognitive nudges: ephemeral prompts for proactive memory use
self._metacog_state: dict[str, float] = {}
self._pending_nudge: str | None = None
self._pending_nudge: list[str] = []
# Cooperative cancellation: set from outside to stop generation
self._cancel_event = threading.Event()
self._cancelled_partial_msg: dict[str, Any] | None = None
@@ -666,7 +666,7 @@ class ChatSession:
memory_count=self._visible_memory_count(),
cooldown_secs=self._memory_config.nudge_cooldown,
):
self._pending_nudge = format_nudge("resume")
self._pending_nudge.append(format_nudge("resume"))
self._init_system_messages()
return True
@@ -807,9 +807,10 @@ class ChatSession:
"Use memory(action='search') or memory(action='list') for more."
)
if self._pending_nudge:
dev_parts.append("")
dev_parts.append(self._pending_nudge)
self._pending_nudge = None
for nudge in self._pending_nudge:
dev_parts.append("")
dev_parts.append(nudge)
self._pending_nudge.clear()
new_system_messages.append({"role": "system", "content": "\n".join(dev_parts)})
# Atomic swap — readers see either old or new, never partial
self.system_messages = new_system_messages
@@ -997,7 +998,7 @@ class ChatSession:
# Metacognitive nudge: check for correction/completion signals
nudge = self._check_metacognitive_nudge(user_input)
if nudge:
self._pending_nudge = nudge
self._pending_nudge.append(nudge)
self._init_system_messages()
try:
@@ -1873,7 +1874,7 @@ class ChatSession:
memory_count=self._visible_memory_count(),
cooldown_secs=self._memory_config.nudge_cooldown,
):
self._pending_nudge = format_nudge("denial")
self._pending_nudge.append(format_nudge("denial"))
self._init_system_messages()
# Phase 3: execute (check cancellation before starting)
@@ -3548,6 +3549,16 @@ class ChatSession:
"and functions in every step."
)
def _plan_system_content(self) -> str:
"""Plan agent system message: template guardrails + plan identity."""
if not self._template_content:
return self._PLAN_IDENTITY
tpl = self._template_content
if len(tpl) > _MAX_TEMPLATE_CONTENT:
log.warning("template_content.truncated", length=len(tpl), agent="plan")
tpl = tpl[:_MAX_TEMPLATE_CONTENT]
return tpl + "\n\n" + self._PLAN_IDENTITY
_MIN_PLAN_LENGTH = 100
_PLAN_REQUIRED_SECTIONS = ("## goal", "## current state", "## plan", "## risks")
_MIN_PLAN_SECTIONS = 2
@@ -3625,10 +3636,11 @@ class ChatSession:
prior_plan_msgs = [msg, self.messages[j]]
break
# Plan agent gets its own identity only — no main session system
# prompt or conversation history. It's an autonomous sub-agent.
# Plan agent gets template guardrails + its own identity — no tool
# patterns, MCP resources, or general conversation history (only
# prior plan tool_call/result pairs are forwarded for refinement).
agent_messages: list[dict[str, Any]] = [
{"role": "system", "content": self._PLAN_IDENTITY},
{"role": "system", "content": self._plan_system_content()},
]
agent_messages.extend(prior_plan_msgs)
agent_messages.append({"role": "user", "content": prompt})
@@ -3699,7 +3711,7 @@ class ChatSession:
"""Re-run the plan agent incorporating user feedback."""
tc_id = f"plan_refine_{uuid.uuid4().hex[:8]}"
agent_messages: list[dict[str, Any]] = [
{"role": "system", "content": self._PLAN_IDENTITY},
{"role": "system", "content": self._plan_system_content()},
{
"role": "assistant",
"content": None,
@@ -3963,6 +3975,7 @@ class ChatSession:
"target": target,
"message": item["message"],
"title": item.get("title", ""),
"ws_id": self._ws_id,
}
# Build auth headers for service-to-service call
+2 -1
View File
@@ -229,7 +229,8 @@ def _build_registry() -> dict[str, SettingDef]:
"mcp",
restart_required=True,
help="Model Context Protocol (MCP) lets the AI connect to external tool servers. "
"This points to a JSON file listing which MCP servers to connect to on startup.",
"This points to a JSON file listing which MCP servers to connect to on startup. "
"Tip: use the MCP Servers tab to manage servers via the database instead.",
reference_url="https://modelcontextprotocol.io",
),
SettingDef(
+274 -4
View File
@@ -14,6 +14,7 @@ from turnstone.core.storage._schema import (
audit_events,
conversations,
intent_verdicts,
mcp_servers,
metadata,
orgs,
prompt_templates,
@@ -29,6 +30,9 @@ from turnstone.core.storage._schema import (
workstream_templates,
workstreams,
)
from turnstone.core.storage._utils import (
MCP_SERVER_MUTABLE as _MCP_SERVER_MUTABLE,
)
from turnstone.core.storage._utils import (
ORG_MUTABLE as _ORG_MUTABLE,
)
@@ -536,12 +540,13 @@ class PostgreSQLBackend:
]
def delete_user(self, user_id: str) -> bool:
from turnstone.core.storage._schema import channel_users
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))
conn.execute(sa.delete(channel_users).where(channel_users.c.user_id == user_id))
conn.execute(sa.delete(api_tokens).where(api_tokens.c.user_id == user_id))
conn.execute(sa.delete(oidc_identities).where(oidc_identities.c.user_id == user_id))
result = conn.execute(sa.delete(users).where(users.c.user_id == user_id))
conn.commit()
return result.rowcount > 0
@@ -2202,7 +2207,7 @@ class PostgreSQLBackend:
q = q.where(structured_memories.c.type == mem_type)
if scope:
q = q.where(structured_memories.c.scope == scope)
if scope_id:
if scope_id and scope:
q = q.where(structured_memories.c.scope_id == scope_id)
q = q.limit(limit)
rows = conn.execute(q).fetchall()
@@ -2241,7 +2246,7 @@ class PostgreSQLBackend:
if scope:
where += " AND scope = :scope_filter"
params["scope_filter"] = scope
if scope_id:
if scope_id and scope:
where += " AND scope_id = :scope_id_filter"
params["scope_id_filter"] = scope_id
rows = conn.execute(
@@ -2262,7 +2267,7 @@ class PostgreSQLBackend:
q = q.where(structured_memories.c.type == mem_type)
if scope:
q = q.where(structured_memories.c.scope == scope)
if scope_id:
if scope_id and scope:
q = q.where(structured_memories.c.scope_id == scope_id)
result = conn.execute(q).scalar()
return int(result or 0)
@@ -2364,6 +2369,271 @@ class PostgreSQLBackend:
).fetchall()
return {r.key: r.value for r in rows}
# -- MCP server definitions ------------------------------------------------
def create_mcp_server(
self,
server_id: str,
name: str,
transport: str,
command: str = "",
args: str = "[]",
url: str = "",
headers: str = "{}",
env: str = "{}",
auto_approve: bool = False,
enabled: bool = True,
created_by: str = "",
) -> None:
from sqlalchemy.dialects import postgresql
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
with self._engine.connect() as conn:
conn.execute(
postgresql.insert(mcp_servers)
.values(
server_id=server_id,
name=name,
transport=transport,
command=command,
args=args,
url=url,
headers=headers,
env=env,
auto_approve=1 if auto_approve else 0,
enabled=1 if enabled else 0,
created_by=created_by,
created=now,
updated=now,
)
.on_conflict_do_nothing()
)
conn.commit()
def get_mcp_server(self, server_id: str) -> dict[str, Any] | None:
with self._engine.connect() as conn:
row = conn.execute(
sa.select(mcp_servers).where(mcp_servers.c.server_id == server_id)
).fetchone()
if row is None:
return None
return _row_to_dict(row, "auto_approve", "enabled")
def get_mcp_server_by_name(self, name: str) -> dict[str, Any] | None:
with self._engine.connect() as conn:
row = conn.execute(sa.select(mcp_servers).where(mcp_servers.c.name == name)).fetchone()
if row is None:
return None
return _row_to_dict(row, "auto_approve", "enabled")
def list_mcp_servers(self, enabled_only: bool = False) -> list[dict[str, Any]]:
with self._engine.connect() as conn:
q = sa.select(mcp_servers).order_by(mcp_servers.c.name)
if enabled_only:
q = q.where(mcp_servers.c.enabled == 1)
rows = conn.execute(q).fetchall()
return [_row_to_dict(r, "auto_approve", "enabled") for r in rows]
def update_mcp_server(self, server_id: str, **fields: Any) -> bool:
fields = {k: v for k, v in fields.items() if k in _MCP_SERVER_MUTABLE}
fields["updated"] = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
if "auto_approve" in fields:
fields["auto_approve"] = 1 if fields["auto_approve"] else 0
if "enabled" in fields:
fields["enabled"] = 1 if fields["enabled"] else 0
with self._engine.connect() as conn:
result = conn.execute(
sa.update(mcp_servers).where(mcp_servers.c.server_id == server_id).values(**fields)
)
conn.commit()
return result.rowcount > 0
def delete_mcp_server(self, server_id: str) -> bool:
with self._engine.connect() as conn:
result = conn.execute(
sa.delete(mcp_servers).where(mcp_servers.c.server_id == server_id)
)
conn.commit()
return result.rowcount > 0
# -- OIDC identity ---------------------------------------------------------
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(
postgresql.insert(oidc_identities)
.values(
issuer=issuer,
subject=subject,
user_id=user_id,
email=email,
created=now,
last_login=now,
)
.on_conflict_do_nothing()
)
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(
sa.select(
oidc_identities.c.issuer,
oidc_identities.c.subject,
oidc_identities.c.user_id,
oidc_identities.c.email,
oidc_identities.c.created,
oidc_identities.c.last_login,
).where(
(oidc_identities.c.issuer == issuer) & (oidc_identities.c.subject == subject)
)
).fetchone()
if row:
return {
"issuer": row[0],
"subject": row[1],
"user_id": row[2],
"email": row[3],
"created": row[4],
"last_login": row[5],
}
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:
result = conn.execute(
sa.update(oidc_identities)
.where(
(oidc_identities.c.issuer == issuer) & (oidc_identities.c.subject == subject)
)
.values(last_login=now)
)
conn.commit()
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(
sa.select(
oidc_identities.c.issuer,
oidc_identities.c.subject,
oidc_identities.c.user_id,
oidc_identities.c.email,
oidc_identities.c.created,
oidc_identities.c.last_login,
)
.where(oidc_identities.c.user_id == user_id)
.order_by(oidc_identities.c.created.desc())
).fetchall()
return [
{
"issuer": r[0],
"subject": r[1],
"user_id": r[2],
"email": r[3],
"created": r[4],
"last_login": r[5],
}
for r in rows
]
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(
sa.delete(oidc_identities).where(
(oidc_identities.c.issuer == issuer) & (oidc_identities.c.subject == subject)
)
)
conn.commit()
return result.rowcount > 0
# -- OIDC pending state ----------------------------------------------------
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:
conn.execute(
sa.insert(oidc_pending_states),
{
"state": state,
"nonce": nonce,
"code_verifier": code_verifier,
"audience": audience,
"created_at": now,
},
)
conn.commit()
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"
)
with self._engine.connect() as conn:
# Atomic DELETE...RETURNING for true one-time consumption
row = conn.execute(
sa.text(
"DELETE FROM oidc_pending_states "
"WHERE state = :state AND created_at > :cutoff "
"RETURNING state, nonce, code_verifier, audience, created_at"
),
{"state": state, "cutoff": cutoff},
).fetchone()
# Also clean up the row if it existed but was expired
if not row:
conn.execute(
sa.delete(oidc_pending_states).where(oidc_pending_states.c.state == state)
)
conn.commit()
if not row:
return None
return {
"state": row[0],
"nonce": row[1],
"code_verifier": row[2],
"audience": row[3],
"created_at": row[4],
}
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"
)
with self._engine.connect() as conn:
result = conn.execute(
sa.delete(oidc_pending_states).where(oidc_pending_states.c.created_at < cutoff)
)
conn.commit()
return result.rowcount
# -- Lifecycle -------------------------------------------------------------
def close(self) -> None:
+79
View File
@@ -259,6 +259,46 @@ class StorageBackend(Protocol):
"""Remove a channel user mapping. Returns True if existed."""
...
# -- OIDC identity ---------------------------------------------------------
def create_oidc_identity(self, issuer: str, subject: str, user_id: str, email: str) -> None:
"""Link an OIDC subject to a turnstone user. No-op if exists."""
...
def get_oidc_identity(self, issuer: str, subject: str) -> dict[str, str] | None:
"""Lookup turnstone user by OIDC issuer+subject. Returns dict or None."""
...
def update_oidc_identity_login(self, issuer: str, subject: str) -> bool:
"""Update last_login timestamp. Returns True if row existed."""
...
def list_oidc_identities_for_user(self, user_id: str) -> list[dict[str, str]]:
"""List all OIDC identities linked to a turnstone user."""
...
def delete_oidc_identity(self, issuer: str, subject: str) -> bool:
"""Remove an OIDC identity link. Returns True if existed."""
...
# -- OIDC pending state ----------------------------------------------------
def create_oidc_pending_state(
self, state: str, nonce: str, code_verifier: str, audience: str
) -> None:
"""Store OIDC authorization flow state for callback validation."""
...
def pop_oidc_pending_state(
self, state: str, max_age_seconds: int = 300
) -> dict[str, str] | None:
"""Fetch and delete pending state atomically. Returns None if expired or missing."""
...
def cleanup_expired_oidc_states(self, max_age_seconds: int = 300) -> int:
"""Delete expired pending states. Returns count of deleted rows."""
...
# -- Channel routing -------------------------------------------------------
def create_channel_route(
@@ -784,6 +824,45 @@ class StorageBackend(Protocol):
"""
...
# -- MCP server definitions ------------------------------------------------
def create_mcp_server(
self,
server_id: str,
name: str,
transport: str,
command: str = "",
args: str = "[]",
url: str = "",
headers: str = "{}",
env: str = "{}",
auto_approve: bool = False,
enabled: bool = True,
created_by: str = "",
) -> None:
"""Create an MCP server definition. No-op if server_id already exists."""
...
def get_mcp_server(self, server_id: str) -> dict[str, Any] | None:
"""Return MCP server dict or None."""
...
def get_mcp_server_by_name(self, name: str) -> dict[str, Any] | None:
"""Return MCP server dict by name or None."""
...
def list_mcp_servers(self, enabled_only: bool = False) -> list[dict[str, Any]]:
"""Return MCP servers ordered by name."""
...
def update_mcp_server(self, server_id: str, **fields: Any) -> bool:
"""Update specified fields on an MCP server. Returns True if found."""
...
def delete_mcp_server(self, server_id: str) -> bool:
"""Delete an MCP server definition. Returns True if existed."""
...
# -- Lifecycle -------------------------------------------------------------
def close(self) -> None:
+52
View File
@@ -445,3 +445,55 @@ system_settings = sa.Table(
)
sa.Index("idx_system_settings_node", system_settings.c.node_id)
# ---------------------------------------------------------------------------
# MCP server definitions — database-backed MCP configuration
# ---------------------------------------------------------------------------
mcp_servers = sa.Table(
"mcp_servers",
metadata,
sa.Column("server_id", sa.Text, primary_key=True),
sa.Column("name", sa.Text, nullable=False, unique=True),
sa.Column("transport", sa.Text, nullable=False), # "stdio" | "streamable-http"
sa.Column("command", sa.Text, nullable=False, server_default=""),
sa.Column("args", sa.Text, nullable=False, server_default="[]"), # JSON array
sa.Column("url", sa.Text, nullable=False, server_default=""),
sa.Column("headers", sa.Text, nullable=False, server_default="{}"), # JSON object
sa.Column("env", sa.Text, nullable=False, server_default="{}"), # JSON object
sa.Column("auto_approve", sa.Integer, nullable=False, server_default="0"),
sa.Column("enabled", sa.Integer, nullable=False, server_default="1"),
sa.Column("created_by", sa.Text, nullable=False, server_default=""),
sa.Column("created", sa.Text, nullable=False),
sa.Column("updated", sa.Text, nullable=False),
)
sa.Index("idx_mcp_servers_enabled", mcp_servers.c.enabled)
# ---------------------------------------------------------------------------
# OIDC identity tables
# ---------------------------------------------------------------------------
oidc_identities = sa.Table(
"oidc_identities",
metadata,
sa.Column("issuer", sa.Text, nullable=False),
sa.Column("subject", sa.Text, nullable=False),
sa.Column("user_id", sa.Text, nullable=False),
sa.Column("email", sa.Text, nullable=False, server_default=""),
sa.Column("created", sa.Text, nullable=False),
sa.Column("last_login", sa.Text, nullable=False),
sa.PrimaryKeyConstraint("issuer", "subject"),
)
sa.Index("idx_oidc_identities_user_id", oidc_identities.c.user_id)
oidc_pending_states = sa.Table(
"oidc_pending_states",
metadata,
sa.Column("state", sa.Text, primary_key=True),
sa.Column("nonce", sa.Text, nullable=False),
sa.Column("code_verifier", sa.Text, nullable=False),
sa.Column("audience", sa.Text, nullable=False),
sa.Column("created_at", sa.Text, nullable=False),
)
+271 -4
View File
@@ -14,6 +14,7 @@ from turnstone.core.storage._schema import (
audit_events,
conversations,
intent_verdicts,
mcp_servers,
metadata,
orgs,
prompt_templates,
@@ -29,6 +30,9 @@ from turnstone.core.storage._schema import (
workstream_templates,
workstreams,
)
from turnstone.core.storage._utils import (
MCP_SERVER_MUTABLE as _MCP_SERVER_MUTABLE,
)
from turnstone.core.storage._utils import (
ORG_MUTABLE as _ORG_MUTABLE,
)
@@ -585,12 +589,13 @@ class SQLiteBackend:
]
def delete_user(self, user_id: str) -> bool:
from turnstone.core.storage._schema import channel_users
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))
conn.execute(sa.delete(channel_users).where(channel_users.c.user_id == user_id))
conn.execute(sa.delete(api_tokens).where(api_tokens.c.user_id == user_id))
conn.execute(sa.delete(oidc_identities).where(oidc_identities.c.user_id == user_id))
result = conn.execute(sa.delete(users).where(users.c.user_id == user_id))
conn.commit()
return result.rowcount > 0
@@ -2226,7 +2231,7 @@ class SQLiteBackend:
q = q.where(structured_memories.c.type == mem_type)
if scope:
q = q.where(structured_memories.c.scope == scope)
if scope_id:
if scope_id and scope:
q = q.where(structured_memories.c.scope_id == scope_id)
q = q.limit(limit)
rows = conn.execute(q).fetchall()
@@ -2265,7 +2270,7 @@ class SQLiteBackend:
if scope:
where += " AND scope = :scope_filter"
params["scope_filter"] = scope
if scope_id:
if scope_id and scope:
where += " AND scope_id = :scope_id_filter"
params["scope_id_filter"] = scope_id
rows = conn.execute(
@@ -2286,7 +2291,7 @@ class SQLiteBackend:
q = q.where(structured_memories.c.type == mem_type)
if scope:
q = q.where(structured_memories.c.scope == scope)
if scope_id:
if scope_id and scope:
q = q.where(structured_memories.c.scope_id == scope_id)
result = conn.execute(q).scalar()
return int(result or 0)
@@ -2389,6 +2394,268 @@ class SQLiteBackend:
).fetchall()
return {r.key: r.value for r in rows}
# -- MCP server definitions ------------------------------------------------
def create_mcp_server(
self,
server_id: str,
name: str,
transport: str,
command: str = "",
args: str = "[]",
url: str = "",
headers: str = "{}",
env: str = "{}",
auto_approve: bool = False,
enabled: bool = True,
created_by: str = "",
) -> None:
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
with self._engine.connect() as conn:
conn.execute(
sa.insert(mcp_servers).prefix_with("OR IGNORE"),
{
"server_id": server_id,
"name": name,
"transport": transport,
"command": command,
"args": args,
"url": url,
"headers": headers,
"env": env,
"auto_approve": 1 if auto_approve else 0,
"enabled": 1 if enabled else 0,
"created_by": created_by,
"created": now,
"updated": now,
},
)
conn.commit()
def get_mcp_server(self, server_id: str) -> dict[str, Any] | None:
with self._engine.connect() as conn:
row = conn.execute(
sa.select(mcp_servers).where(mcp_servers.c.server_id == server_id)
).fetchone()
if row is None:
return None
return _row_to_dict(row, "auto_approve", "enabled")
def get_mcp_server_by_name(self, name: str) -> dict[str, Any] | None:
with self._engine.connect() as conn:
row = conn.execute(sa.select(mcp_servers).where(mcp_servers.c.name == name)).fetchone()
if row is None:
return None
return _row_to_dict(row, "auto_approve", "enabled")
def list_mcp_servers(self, enabled_only: bool = False) -> list[dict[str, Any]]:
with self._engine.connect() as conn:
q = sa.select(mcp_servers).order_by(mcp_servers.c.name)
if enabled_only:
q = q.where(mcp_servers.c.enabled == 1)
rows = conn.execute(q).fetchall()
return [_row_to_dict(r, "auto_approve", "enabled") for r in rows]
def update_mcp_server(self, server_id: str, **fields: Any) -> bool:
fields = {k: v for k, v in fields.items() if k in _MCP_SERVER_MUTABLE}
fields["updated"] = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
if "auto_approve" in fields:
fields["auto_approve"] = 1 if fields["auto_approve"] else 0
if "enabled" in fields:
fields["enabled"] = 1 if fields["enabled"] else 0
with self._engine.connect() as conn:
result = conn.execute(
sa.update(mcp_servers).where(mcp_servers.c.server_id == server_id).values(**fields)
)
conn.commit()
return result.rowcount > 0
def delete_mcp_server(self, server_id: str) -> bool:
with self._engine.connect() as conn:
result = conn.execute(
sa.delete(mcp_servers).where(mcp_servers.c.server_id == server_id)
)
conn.commit()
return result.rowcount > 0
# -- 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:
conn.execute(
sa.insert(oidc_identities).prefix_with("OR IGNORE"),
{
"issuer": issuer,
"subject": subject,
"user_id": user_id,
"email": email,
"created": now,
"last_login": now,
},
)
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(
sa.select(
oidc_identities.c.issuer,
oidc_identities.c.subject,
oidc_identities.c.user_id,
oidc_identities.c.email,
oidc_identities.c.created,
oidc_identities.c.last_login,
).where(
(oidc_identities.c.issuer == issuer) & (oidc_identities.c.subject == subject)
)
).fetchone()
if row:
return {
"issuer": row[0],
"subject": row[1],
"user_id": row[2],
"email": row[3],
"created": row[4],
"last_login": row[5],
}
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:
result = conn.execute(
sa.update(oidc_identities)
.where(
(oidc_identities.c.issuer == issuer) & (oidc_identities.c.subject == subject)
)
.values(last_login=now)
)
conn.commit()
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(
sa.select(
oidc_identities.c.issuer,
oidc_identities.c.subject,
oidc_identities.c.user_id,
oidc_identities.c.email,
oidc_identities.c.created,
oidc_identities.c.last_login,
)
.where(oidc_identities.c.user_id == user_id)
.order_by(oidc_identities.c.created.desc())
).fetchall()
return [
{
"issuer": r[0],
"subject": r[1],
"user_id": r[2],
"email": r[3],
"created": r[4],
"last_login": r[5],
}
for r in rows
]
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(
sa.delete(oidc_identities).where(
(oidc_identities.c.issuer == issuer) & (oidc_identities.c.subject == subject)
)
)
conn.commit()
return result.rowcount > 0
# -- OIDC pending state ----------------------------------------------------
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:
conn.execute(
sa.insert(oidc_pending_states),
{
"state": state,
"nonce": nonce,
"code_verifier": code_verifier,
"audience": audience,
"created_at": now,
},
)
conn.commit()
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"
)
with self._engine.connect() as conn:
# Acquire write lock before SELECT to prevent TOCTOU race
conn.execute(sa.text("BEGIN IMMEDIATE"))
row = conn.execute(
sa.select(
oidc_pending_states.c.state,
oidc_pending_states.c.nonce,
oidc_pending_states.c.code_verifier,
oidc_pending_states.c.audience,
oidc_pending_states.c.created_at,
).where(
(oidc_pending_states.c.state == state)
& (oidc_pending_states.c.created_at > cutoff)
)
).fetchone()
# Always delete the row (whether valid, expired, or missing is fine)
conn.execute(sa.delete(oidc_pending_states).where(oidc_pending_states.c.state == state))
conn.commit()
if not row:
return None
return {
"state": row[0],
"nonce": row[1],
"code_verifier": row[2],
"audience": row[3],
"created_at": row[4],
}
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"
)
with self._engine.connect() as conn:
result = conn.execute(
sa.delete(oidc_pending_states).where(oidc_pending_states.c.created_at < cutoff)
)
conn.commit()
return result.rowcount
# -- Lifecycle -------------------------------------------------------------
def close(self) -> None:
+13
View File
@@ -48,6 +48,19 @@ WS_TEMPLATE_MUTABLE = frozenset(
}
)
STRUCTURED_MEMORY_MUTABLE = frozenset({"content", "description", "type"})
MCP_SERVER_MUTABLE = frozenset(
{
"name",
"transport",
"command",
"args",
"url",
"headers",
"env",
"auto_approve",
"enabled",
}
)
VERDICT_MUTABLE = frozenset(
{
"user_decision",
@@ -0,0 +1,55 @@
"""Create mcp_servers table and grant admin.mcp permission.
Revision ID: 016
Revises: 015
Create Date: 2026-03-14
"""
import sqlalchemy as sa
from alembic import op
revision = "016"
down_revision = "015"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"mcp_servers",
sa.Column("server_id", sa.Text, primary_key=True),
sa.Column("name", sa.Text, nullable=False, unique=True),
sa.Column("transport", sa.Text, nullable=False),
sa.Column("command", sa.Text, nullable=False, server_default=""),
sa.Column("args", sa.Text, nullable=False, server_default="[]"),
sa.Column("url", sa.Text, nullable=False, server_default=""),
sa.Column("headers", sa.Text, nullable=False, server_default="{}"),
sa.Column("env", sa.Text, nullable=False, server_default="{}"),
sa.Column("auto_approve", sa.Integer, nullable=False, server_default="0"),
sa.Column("enabled", sa.Integer, nullable=False, server_default="1"),
sa.Column("created_by", sa.Text, nullable=False, server_default=""),
sa.Column("created", sa.Text, nullable=False),
sa.Column("updated", sa.Text, nullable=False),
)
op.create_index("idx_mcp_servers_enabled", "mcp_servers", ["enabled"])
# Grant admin.mcp permission to the built-in admin role
conn = op.get_bind()
conn.execute(
sa.text(
"UPDATE roles SET permissions = permissions || ',admin.mcp' "
"WHERE role_id = 'builtin-admin' "
"AND permissions NOT LIKE '%admin.mcp%'"
)
)
def downgrade() -> None:
conn = op.get_bind()
conn.execute(
sa.text(
"UPDATE roles SET permissions = REPLACE(permissions, ',admin.mcp', '') "
"WHERE role_id = 'builtin-admin'"
)
)
op.drop_table("mcp_servers")
@@ -0,0 +1,46 @@
"""Catch-up: ensure builtin-admin role has all current permissions.
Migrations 011-016 each appended a permission to the builtin-admin role,
but on some deployments these UPDATE statements did not take effect
(e.g. due to version stamping without running, or create_all bypassing
Alembic). This migration idempotently ensures the builtin-admin role
has the complete permission set.
Revision ID: 017
Revises: 016
Create Date: 2026-03-14
"""
import sqlalchemy as sa
from alembic import op
revision = "017"
down_revision = "016"
branch_labels = None
depends_on = None
# The complete set of permissions the builtin-admin role should have.
# Must stay in sync with _VALID_PERMISSIONS in console/server.py.
_EXPECTED_ADMIN_PERMS = (
"read,write,approve,"
"admin.users,admin.roles,admin.orgs,"
"admin.policies,admin.templates,admin.ws_templates,"
"admin.audit,admin.usage,"
"admin.schedules,admin.watches,"
"admin.judge,admin.memories,admin.settings,admin.mcp,"
"tools.approve,workstreams.create,workstreams.close"
)
def upgrade() -> None:
conn = op.get_bind()
conn.execute(
sa.text("UPDATE roles SET permissions = :perms WHERE role_id = 'builtin-admin'"),
{"perms": _EXPECTED_ADMIN_PERMS},
)
def downgrade() -> None:
# No-op: we don't remove permissions on downgrade since we can't
# know which subset the deployment originally had.
pass
@@ -0,0 +1,42 @@
"""Create OIDC identity and pending state tables.
Revision ID: 018
Revises: 017
Create Date: 2026-03-15
"""
import sqlalchemy as sa
from alembic import op
revision = "018"
down_revision = "017"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"oidc_identities",
sa.Column("issuer", sa.Text, nullable=False),
sa.Column("subject", sa.Text, nullable=False),
sa.Column("user_id", sa.Text, nullable=False),
sa.Column("email", sa.Text, nullable=False, server_default=""),
sa.Column("created", sa.Text, nullable=False),
sa.Column("last_login", sa.Text, nullable=False),
sa.PrimaryKeyConstraint("issuer", "subject"),
)
op.create_index("idx_oidc_identities_user_id", "oidc_identities", ["user_id"])
op.create_table(
"oidc_pending_states",
sa.Column("state", sa.Text, primary_key=True),
sa.Column("nonce", sa.Text, nullable=False),
sa.Column("code_verifier", sa.Text, nullable=False),
sa.Column("audience", sa.Text, nullable=False),
sa.Column("created_at", sa.Text, nullable=False),
)
def downgrade() -> None:
op.drop_table("oidc_pending_states")
op.drop_table("oidc_identities")
+18 -6
View File
@@ -16,9 +16,20 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING:
from collections.abc import Callable
from typing import Protocol
from turnstone.core.session import ChatSession, SessionUI
class _SessionFactory(Protocol):
def __call__(
self,
ui: SessionUI | None,
model_alias: str | None = ...,
ws_id: str | None = ...,
*,
template: str | None = ...,
) -> ChatSession: ...
# ---------------------------------------------------------------------------
# State enum
@@ -65,14 +76,14 @@ class WorkstreamManager:
def __init__(
self,
session_factory: Callable[[SessionUI | None, str | None, str | None], ChatSession],
session_factory: _SessionFactory,
*,
max_workstreams: int = 10,
node_id: str | None = None,
):
"""
Args:
session_factory: callable(ui, model_alias, ws_id) -> ChatSession.
session_factory: callable(ui, model_alias, ws_id, *, template) -> ChatSession.
Captures shared config (registry, temperature, ) so the
manager can create ChatSession instances without knowing
those details. *model_alias* selects a model from the
@@ -85,9 +96,7 @@ class WorkstreamManager:
"""
if max_workstreams < 1:
raise ValueError(f"max_workstreams must be >= 1, got {max_workstreams}")
self._session_factory: Callable[[SessionUI | None, str | None, str | None], ChatSession] = (
session_factory
)
self._session_factory: _SessionFactory = session_factory
self._node_id = node_id
self._max_workstreams: int = max_workstreams
self._workstreams: dict[str, Workstream] = {}
@@ -115,6 +124,7 @@ class WorkstreamManager:
name: str = "",
ui_factory: Callable[..., SessionUI] | None = None,
model: str | None = None,
template: str | None = None,
) -> Workstream:
"""Create a new workstream. Returns the new ws.
@@ -125,6 +135,8 @@ class WorkstreamManager:
Args:
model: Optional model alias from the registry. ``None`` uses the
default model.
template: Optional prompt template name passed through to session
factory.
"""
# Fast-fail capacity check (avoids expensive ChatSession creation when full).
first_evicted: Workstream | None = None
@@ -147,7 +159,7 @@ class WorkstreamManager:
ws = Workstream(name=name)
if ui_factory:
ws.ui = ui_factory(ws.id)
ws.session = self._session_factory(ws.ui, model, ws.id)
ws.session = self._session_factory(ws.ui, model, ws.id, template=template)
# Authoritative insert under lock with re-check (another thread may
# have filled capacity while we were unlocked).
+24 -4
View File
@@ -701,9 +701,22 @@ class Bridge:
feedback = getattr(resp_msg, "feedback", None)
always = getattr(resp_msg, "always", False)
self._api_approve(ws_id, approved=approved, feedback=feedback)
if always:
with self._lock:
self._ws_auto_approve[ws_id] = True
if always and approved:
tool_names = {
it.get("func_name", "")
for it in items
if it.get("needs_approval")
and it.get("func_name")
and not it.get("error")
}
tool_names.discard("")
tool_names.discard("__budget_override__")
if tool_names:
with self._lock:
existing = self._ws_approve_tools.get(
ws_id, set(DEFAULT_SAFE_TOOLS)
)
self._ws_approve_tools[ws_id] = existing | tool_names
else:
log.warning("Approval timeout for ws %s — denying", ws_id)
self._api_approve(ws_id, approved=False, feedback="Approval timed out")
@@ -825,7 +838,14 @@ class Bridge:
if state == "idle":
with self._lock:
cid = self._active_sends.pop(ws_id, None)
self._publish_ws(ws_id, TurnCompleteEvent(ws_id=ws_id, correlation_id=cid or ""))
self._publish_ws(
ws_id,
TurnCompleteEvent(
ws_id=ws_id,
correlation_id=cid or "",
content=data.get("content", ""),
),
)
elif etype == "ws_rename":
self._publish_global(WorkstreamRenameEvent(ws_id=ws_id, name=data.get("name", "")))
+7
View File
@@ -268,9 +268,16 @@ class TurnCompleteEvent(OutboundEvent):
This is a synthetic event produced by the bridge when it detects
the ws_state transition to 'idle'. ``correlation_id`` is set for
MQ-initiated turns and empty for turns initiated from the server UI.
``content`` carries the full assistant response text piggybacked on
the server's idle SSE event (accumulated server-side in WebUI).
Downstream consumers (e.g. Discord bot) use it for catch-up when the
streaming path missed events, and as the primary delivery path for
bidirectional notification DM forwarding.
"""
type: str = "turn_complete"
content: str = ""
@dataclass
+142
View File
@@ -21,8 +21,10 @@ from turnstone.api.console_schemas import (
ClusterWorkstreamsResponse,
ConsoleCreateWsResponse,
ConsoleHealthResponse,
ImportMcpConfigResponse,
ListAdminMemoriesResponse,
ListAuditEventsResponse,
ListMcpServersResponse,
ListOrgsResponse,
ListPromptTemplatesResponse,
ListRolesResponse,
@@ -32,6 +34,7 @@ from turnstone.api.console_schemas import (
ListUserRolesResponse,
ListWsTemplatesResponse,
ListWsTemplateVersionsResponse,
McpServerDetail,
NodeDetailResponse,
OrgInfo,
PromptTemplateInfo,
@@ -666,6 +669,98 @@ class AsyncTurnstoneConsole(_BaseClient):
"DELETE", f"/v1/api/admin/settings/{key}", params=params, response_model=StatusResponse
)
# -- MCP servers -------------------------------------------------------
async def list_mcp_servers(self, reveal: bool = False) -> ListMcpServersResponse:
"""List MCP server definitions with live status."""
params: dict[str, str] = {}
if reveal:
params["reveal"] = "true"
return await self._request(
"GET",
"/v1/api/admin/mcp-servers",
params=params,
response_model=ListMcpServersResponse,
)
async def create_mcp_server(
self,
name: str,
transport: str,
*,
command: str = "",
args: list[str] | None = None,
url: str = "",
headers: dict[str, str] | None = None,
env: dict[str, str] | None = None,
auto_approve: bool = False,
enabled: bool = True,
) -> McpServerDetail:
"""Create an MCP server definition."""
body: dict[str, Any] = {"name": name, "transport": transport}
if command:
body["command"] = command
if args:
body["args"] = args
if url:
body["url"] = url
if headers:
body["headers"] = headers
if env:
body["env"] = env
if auto_approve:
body["auto_approve"] = True
if not enabled:
body["enabled"] = False
return await self._request(
"POST",
"/v1/api/admin/mcp-servers",
json_body=body,
response_model=McpServerDetail,
)
async def get_mcp_server(self, server_id: str) -> McpServerDetail:
"""Get a single MCP server with status."""
return await self._request(
"GET",
f"/v1/api/admin/mcp-servers/{server_id}",
response_model=McpServerDetail,
)
async def update_mcp_server(self, server_id: str, **fields: Any) -> McpServerDetail:
"""Update an MCP server definition."""
return await self._request(
"PUT",
f"/v1/api/admin/mcp-servers/{server_id}",
json_body=fields,
response_model=McpServerDetail,
)
async def delete_mcp_server(self, server_id: str) -> StatusResponse:
"""Delete an MCP server definition."""
return await self._request(
"DELETE",
f"/v1/api/admin/mcp-servers/{server_id}",
response_model=StatusResponse,
)
async def reload_mcp_servers(self) -> StatusResponse:
"""Tell all nodes to re-read MCP server config from DB."""
return await self._request(
"POST",
"/v1/api/admin/mcp-servers/reload",
response_model=StatusResponse,
)
async def import_mcp_config(self, config: dict[str, Any]) -> ImportMcpConfigResponse:
"""Import MCP servers from a config dict with mcpServers key."""
return await self._request(
"POST",
"/v1/api/admin/mcp-servers/import",
json_body={"config": config},
response_model=ImportMcpConfigResponse,
)
class TurnstoneConsole:
"""Synchronous client for the turnstone console API.
@@ -1041,6 +1136,53 @@ class TurnstoneConsole:
def delete_setting(self, key: str, *, node_id: str = "") -> StatusResponse:
return self._runner.run(self._async.delete_setting(key, node_id=node_id))
# -- MCP servers -------------------------------------------------------
def list_mcp_servers(self, reveal: bool = False) -> ListMcpServersResponse:
return self._runner.run(self._async.list_mcp_servers(reveal=reveal))
def create_mcp_server(
self,
name: str,
transport: str,
*,
command: str = "",
args: list[str] | None = None,
url: str = "",
headers: dict[str, str] | None = None,
env: dict[str, str] | None = None,
auto_approve: bool = False,
enabled: bool = True,
) -> McpServerDetail:
return self._runner.run(
self._async.create_mcp_server(
name,
transport,
command=command,
args=args,
url=url,
headers=headers,
env=env,
auto_approve=auto_approve,
enabled=enabled,
)
)
def get_mcp_server(self, server_id: str) -> McpServerDetail:
return self._runner.run(self._async.get_mcp_server(server_id))
def update_mcp_server(self, server_id: str, **fields: Any) -> McpServerDetail:
return self._runner.run(self._async.update_mcp_server(server_id, **fields))
def delete_mcp_server(self, server_id: str) -> StatusResponse:
return self._runner.run(self._async.delete_mcp_server(server_id))
def reload_mcp_servers(self) -> StatusResponse:
return self._runner.run(self._async.reload_mcp_servers())
def import_mcp_config(self, config: dict[str, Any]) -> ImportMcpConfigResponse:
return self._runner.run(self._async.import_mcp_config(config))
# -- lifecycle -----------------------------------------------------------
def close(self) -> None:

Some files were not shown because too many files have changed in this diff Show More