Files
turnstone/tests/test_console.py
T
Patrick Buckley 71d3ed6abe fix(ui): repair interactive/proxied workstream lifecycle (reload, create, launcher)
Workstream-lifecycle bugfixes on the L-shell:

- Node-proxied interactive panes now SURVIVE a browser reload.  On first
  activate a pane resolves its owning node and (re)opens the session there
  before streaming — the node /events stream 404s on a ws not loaded on its
  node, so a rehydrated pane could not just connect blind.  Resolution is
  origin-first via the new TS_APP.resolveInteractiveNode seam (POST /open with
  a rendezvous /route fallback).  PaneManager now persists a pane's resolved
  nodeId as opaque meta and hands it back on rehydrate, so a reload restores
  the pane onto the SAME node even before the Tier-1 snapshot has populated —
  the exact timing that used to strand it on base="" (the console, not a node).

- Both launcher personas open the new session as a PANE, not a full-page nav
  (coordinator -> coordinator pane; interactive -> node-proxied pane); the
  full-page nav stays only as the shell-absent fallback.  Every interactive
  entry point (create, active row, rail, saved row, child link, reload) now
  funnels through one resolve-open-connect path, folding away the bespoke
  restoreInteractiveSession helper.

- The interactive launcher gains a node-selection strategy (Least loaded |
  Specific node, with a live node picker fed from the cluster snapshot) and a
  persona-aware task hint — the shared composer no longer shows
  "...coordinator orchestrate?" when the interactive persona is selected.

Guards updated to pin the new wiring; the stale console landing test (asserting
the renovation-retired bottom-bar node picker) is corrected to the rail.
2026-06-08 10:08:30 -07:00

2563 lines
94 KiB
Python

"""Tests for turnstone.console — collector and HTTP server."""
import asyncio
import json
import queue
from typing import Any
from unittest.mock import MagicMock
import pytest
from turnstone.console.collector import ClusterCollector, NodeSnapshot
from turnstone.console.server import _PROXY_AUTH_LOCAL_HANDLERS
# Shared test auth — JWT-based
_TEST_JWT_SECRET = "test-jwt-secret-minimum-32-chars!"
def _test_jwt() -> str:
from turnstone.core.auth import JWT_AUD_CONSOLE, create_jwt
return create_jwt(
user_id="test-console",
scopes=frozenset({"read", "write", "approve", "service"}),
source="test",
secret=_TEST_JWT_SECRET,
audience=JWT_AUD_CONSOLE,
)
_TEST_AUTH_HEADERS = {"Authorization": f"Bearer {_test_jwt()}"}
# ---------------------------------------------------------------------------
# Mock storage for collector tests
# ---------------------------------------------------------------------------
from tests._coord_test_helpers import MockStorage # noqa: E402, F401
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_collector(storage=None, discovery_interval=999):
"""Create a collector for tests (discovery disabled by default)."""
s = storage or MockStorage()
return ClusterCollector(
storage=s,
discovery_interval=discovery_interval,
)
# ---------------------------------------------------------------------------
# ClusterCollector — unit tests
# ---------------------------------------------------------------------------
class TestCollectorDiscovery:
"""Node discovery from service registry."""
def test_discover_new_nodes(self):
storage = MockStorage()
storage.services = [
{"service_id": "node-a", "url": "http://a:8080", "metadata": "{}"},
{"service_id": "node-b", "url": "http://b:8080", "metadata": "{}"},
]
c = _make_collector(storage)
c._discover_nodes()
overview = c.get_overview()
assert overview["nodes"] == 2
def test_discover_removes_lost_nodes(self):
storage = MockStorage()
storage.services = [
{"service_id": "node-a", "url": "http://a:8080", "metadata": "{}"},
]
c = _make_collector(storage)
c._discover_nodes()
assert c.get_overview()["nodes"] == 1
# Node disappears
storage.services = []
c._discover_nodes()
assert c.get_overview()["nodes"] == 0
def test_discover_updates_server_url(self):
storage = MockStorage()
storage.services = [
{"service_id": "node-a", "url": "http://a:8080", "metadata": "{}"},
]
c = _make_collector(storage)
c._discover_nodes()
storage.services = [
{"service_id": "node-a", "url": "http://a:9090", "metadata": "{}"},
]
c._discover_nodes()
detail = c.get_node_detail("node-a")
assert detail["server_url"] == "http://a:9090"
def test_discover_emits_node_joined_event(self):
storage = MockStorage()
c = _make_collector(storage)
q = queue.Queue()
c.register_listener(q)
storage.services = [
{"service_id": "node-a", "url": "http://a:8080", "metadata": "{}"},
]
c._discover_nodes()
event = q.get_nowait()
assert event["type"] == "node_joined"
assert event["node_id"] == "node-a"
def test_discover_emits_node_lost_event(self):
storage = MockStorage()
storage.services = [
{"service_id": "node-a", "url": "http://a:8080", "metadata": "{}"},
]
c = _make_collector(storage)
c._discover_nodes()
q = queue.Queue()
c.register_listener(q)
storage.services = []
c._discover_nodes()
event = q.get_nowait()
assert event["type"] == "node_lost"
assert event["node_id"] == "node-a"
def test_discover_parses_metadata(self):
storage = MockStorage()
storage.services = [
{
"service_id": "node-a",
"url": "http://a:8080",
"metadata": '{"max_ws": 20, "started": 1234567890.0}',
},
]
c = _make_collector(storage)
c._discover_nodes()
detail = c.get_node_detail("node-a")
assert detail is not None
# Verify metadata was parsed into the NodeSnapshot
assert c._nodes["node-a"].max_ws == 20
assert c._nodes["node-a"].started == 1234567890.0
class TestCollectorNotifyWireIn:
"""NotifyDispatcher-driven discovery — reactive node visibility."""
def test_start_subscribes_to_services_channel(self):
# Stub dispatcher records subscriptions without spawning threads.
class _StubDispatcher:
def __init__(self):
self.subscriptions: list[tuple[str, Any]] = []
def subscribe(self, channel, handler):
self.subscriptions.append((channel, handler))
return lambda: None
stub = _StubDispatcher()
storage = MockStorage()
c = ClusterCollector(
storage=storage,
discovery_interval=999,
notify_dispatcher=stub,
)
try:
c.start()
assert len(stub.subscriptions) == 1
channel, handler = stub.subscriptions[0]
assert channel == "services"
assert handler == c._on_services_notify
finally:
c.stop()
def test_no_dispatcher_means_no_subscribe(self):
# Collector without a dispatcher (single-node / SQLite dev) just
# falls back to the 60 s discovery-loop polling — no error.
c = _make_collector(MockStorage())
try:
c.start()
assert c._notify_unsubscribe is None
finally:
c.stop()
def test_on_notify_runs_discovery(self):
# Construct a synthetic Notify and invoke the handler directly —
# asserts the wire-in delegates back to ``_discover_nodes``.
from turnstone.core.storage._notify import Notify
storage = MockStorage()
c = _make_collector(storage)
c._running = True # bypass start() so we don't spawn threads
q: queue.Queue[dict[str, Any]] = queue.Queue()
c.register_listener(q)
storage.services = [
{"service_id": "node-z", "url": "http://z:8080", "metadata": "{}"},
]
c._on_services_notify(Notify(channel="services", payload="{}", pid=0))
event = q.get_nowait()
assert event["type"] == "node_joined"
assert event["node_id"] == "node-z"
def test_on_notify_when_not_running_is_noop(self):
# If a stray notify arrives after stop, the handler doesn't run
# discovery on a half-torn-down collector.
from turnstone.core.storage._notify import Notify
storage = MockStorage()
storage.services = [{"service_id": "node-y", "url": "http://y:8080", "metadata": "{}"}]
c = _make_collector(storage)
# _running stays False (never called start()).
c._on_services_notify(Notify(channel="services", payload="{}", pid=0))
assert c.get_overview()["nodes"] == 0
class TestCollectorSnapshot:
"""Applying node_snapshot SSE events."""
def test_apply_snapshot_populates_workstreams(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
c._apply_snapshot(
"node-a",
{
"type": "node_snapshot",
"node_id": "node-a",
"workstreams": [
{
"id": "ws1",
"name": "test",
"state": "running",
"tokens": 1000,
"context_ratio": 0.15,
"activity": "bash: ls",
"activity_state": "tool",
"tool_calls": 3,
"title": "My task",
},
],
"health": {"status": "ok"},
"aggregate": {"total_tokens": 1000, "total_tool_calls": 3},
},
)
detail = c.get_node_detail("node-a")
assert len(detail["workstreams"]) == 1
assert detail["workstreams"][0]["name"] == "test"
assert detail["workstreams"][0]["node"] == "node-a"
assert detail["workstreams"][0]["server_url"] == "http://a:8080"
assert detail["health"]["status"] == "ok"
def test_apply_snapshot_replaces_stale_workstreams(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"old-ws": {"id": "old-ws", "name": "old", "state": "idle"}},
)
c._apply_snapshot(
"node-a",
{
"type": "node_snapshot",
"node_id": "node-a",
"workstreams": [{"id": "new-ws", "name": "new", "state": "running"}],
"health": {},
"aggregate": {},
},
)
detail = c.get_node_detail("node-a")
assert len(detail["workstreams"]) == 1
assert detail["workstreams"][0]["id"] == "new-ws"
def test_apply_snapshot_ignores_unknown_node(self):
c = _make_collector()
# Should not raise
c._apply_snapshot(
"unknown", {"type": "node_snapshot", "workstreams": [], "health": {}, "aggregate": {}}
)
def test_apply_snapshot_emits_ws_created_for_new_workstream(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_snapshot(
"node-a",
{
"type": "node_snapshot",
"node_id": "node-a",
"workstreams": [{"id": "ws1", "name": "new-task", "state": "idle"}],
"health": {},
"aggregate": {},
},
)
event = q.get_nowait()
assert event["type"] == "ws_created"
assert event["ws_id"] == "ws1"
assert event["name"] == "new-task"
assert event["node_id"] == "node-a"
def test_apply_snapshot_emits_ws_closed_for_removed_workstream(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "old", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_snapshot(
"node-a",
{
"type": "node_snapshot",
"node_id": "node-a",
"workstreams": [],
"health": {},
"aggregate": {},
},
)
event = q.get_nowait()
assert event["type"] == "ws_closed"
assert event["ws_id"] == "ws1"
def test_apply_snapshot_no_events_when_unchanged(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "same", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_snapshot(
"node-a",
{
"type": "node_snapshot",
"node_id": "node-a",
"workstreams": [{"id": "ws1", "name": "same", "state": "idle"}],
"health": {},
"aggregate": {},
},
)
assert q.empty()
def test_apply_snapshot_emits_state_change_as_cluster_state(self):
"""State change events must use type 'cluster_state' for the frontend."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "same", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_snapshot(
"node-a",
{
"type": "node_snapshot",
"node_id": "node-a",
"workstreams": [{"id": "ws1", "name": "same", "state": "running"}],
"health": {},
"aggregate": {},
},
)
event = q.get_nowait()
assert event["type"] == "cluster_state"
assert event["ws_id"] == "ws1"
assert event["state"] == "running"
def test_apply_snapshot_state_change_does_not_carry_pending_approval_detail(self):
"""Stage 3 cleanup — the snapshot-resync cluster_state event no
longer piggybacks ``pending_approval_detail`` (the field is
gone from cluster_state entirely). On reconnect the browser's
bulk fetch — triggered by the ``activity_state="approval"``
transition in the reducer — pulls the items directly from
``ui.serialize_pending_approval_detail()`` via the dashboard
endpoint."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "same", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_snapshot(
"node-a",
{
"type": "node_snapshot",
"node_id": "node-a",
"workstreams": [
{
"id": "ws1",
"name": "same",
"state": "running",
"activity_state": "approval",
}
],
"health": {},
"aggregate": {},
},
)
event = q.get_nowait()
assert event["type"] == "cluster_state"
assert event["activity_state"] == "approval"
assert "pending_approval_detail" not in event
def test_apply_snapshot_skips_empty_id_workstream(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_snapshot(
"node-a",
{
"type": "node_snapshot",
"node_id": "node-a",
"workstreams": [{"name": "no-id", "state": "idle"}],
"health": {},
"aggregate": {},
},
)
assert q.empty()
assert len(c._nodes["node-a"].workstreams) == 0
class TestCollectorDelta:
"""Applying individual SSE delta events."""
def test_apply_delta_ws_state_fans_out_as_cluster_state(self):
"""Server emits ws_state; collector must translate to cluster_state."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_delta(
"node-a", {"type": "ws_state", "ws_id": "ws1", "state": "running", "tokens": 500}
)
event = q.get_nowait()
assert event["type"] == "cluster_state"
assert event["state"] == "running"
# Verify in-memory state was updated
assert c._nodes["node-a"].workstreams["ws1"]["state"] == "running"
def test_apply_delta_ws_state_does_not_carry_pending_approval_detail(self):
"""Stage 3 cleanup — ``cluster_state`` no longer carries the
``pending_approval_detail`` piggyback. Approval items now arrive
via bulk fetch on activity_state transition; verdicts via the
explicit ``intent_verdict`` event class. Symmetric event flow,
no piggyback to dedupe against."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_delta(
"node-a",
{
"type": "ws_state",
"ws_id": "ws1",
"state": "running",
"activity_state": "approval",
},
)
event = q.get_nowait()
assert event["type"] == "cluster_state"
assert event["activity_state"] == "approval"
assert "pending_approval_detail" not in event
def test_apply_delta_ws_created(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_delta("node-a", {"type": "ws_created", "ws_id": "ws1", "name": "new"})
event = q.get_nowait()
assert event["type"] == "ws_created"
assert "ws1" in c._nodes["node-a"].workstreams
def test_apply_delta_ws_closed(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "old", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_delta("node-a", {"type": "ws_closed", "ws_id": "ws1"})
event = q.get_nowait()
assert event["type"] == "ws_closed"
assert "ws1" not in c._nodes["node-a"].workstreams
def test_apply_delta_ws_rename(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "old-name", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_delta("node-a", {"type": "ws_rename", "ws_id": "ws1", "name": "new-name"})
event = q.get_nowait()
assert event["type"] == "ws_rename"
assert event["name"] == "new-name"
assert c._nodes["node-a"].workstreams["ws1"]["name"] == "new-name"
def test_apply_delta_intent_verdict_forwards_verbatim(self):
"""Stage 3 Step 5 — node-emitted intent_verdict events flow
through _apply_delta to cluster fan-out so coord adapters can
re-emit as child_ws_intent_verdict on the parent's SSE."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
verdict = {
"call_id": "c1",
"risk_level": "low",
"confidence": 0.9,
"recommendation": "approve",
}
c._apply_delta(
"node-a",
{"type": "intent_verdict", "ws_id": "ws1", "verdict": verdict},
)
event = q.get_nowait()
assert event["type"] == "intent_verdict"
assert event["ws_id"] == "ws1"
assert event["node_id"] == "node-a"
assert event["verdict"] == verdict
def test_apply_delta_intent_verdict_drops_when_ws_id_missing(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_delta("node-a", {"type": "intent_verdict", "verdict": {}})
assert q.empty()
def test_apply_delta_approval_resolved_forwards_verbatim(self):
"""Stage 3 Step 5 — paired with intent_verdict; clears the
coord tree's pending-approval pill in lockstep with the
actual decision rather than waiting for the state-change
piggyback."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_delta(
"node-a",
{
"type": "approval_resolved",
"ws_id": "ws1",
"approved": True,
"feedback": "lgtm",
"always": False,
},
)
event = q.get_nowait()
assert event["type"] == "approval_resolved"
assert event["ws_id"] == "ws1"
assert event["node_id"] == "node-a"
assert event["approved"] is True
assert event["feedback"] == "lgtm"
assert event["always"] is False
def test_apply_delta_approve_request_forwards_detail(self):
"""Push path for the initial approval items — eliminates the
bulk-fetch race that left the coord row stuck on a loading
placeholder when the bulk fetch landed in the gap between
_emit_state(ATTENTION) and approve_tools setting _pending_approval."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
detail = {
"type": "approve_request",
"items": [{"call_id": "c1", "header": "tool x"}],
"judge_pending": True,
}
c._apply_delta(
"node-a",
{"type": "approve_request", "ws_id": "ws1", "detail": detail},
)
event = q.get_nowait()
assert event["type"] == "approve_request"
assert event["ws_id"] == "ws1"
assert event["node_id"] == "node-a"
assert event["detail"] == detail
def test_apply_delta_approve_request_drops_when_ws_id_missing(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_delta("node-a", {"type": "approve_request", "detail": {}})
assert q.empty()
def test_apply_delta_approval_resolved_coerces_missing_fields(self):
"""Defensive: ``approved`` / ``always`` / ``feedback`` may be
omitted by older nodes mid-rolling-upgrade; collector coerces
to safe defaults."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
)
q: queue.Queue[dict] = queue.Queue()
c.register_listener(q)
c._apply_delta("node-a", {"type": "approval_resolved", "ws_id": "ws1"})
event = q.get_nowait()
assert event["approved"] is False
assert event["feedback"] == ""
assert event["always"] is False
def test_apply_delta_health_changed(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
health={"status": "ok", "backend": {"status": "up"}},
)
c._apply_delta("node-a", {"type": "health_changed", "backend_status": "degraded"})
health = c._nodes["node-a"].health
assert health["backend"]["status"] == "down"
assert health["status"] == "degraded"
def test_apply_delta_aggregate(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
c._apply_delta(
"node-a",
{"type": "aggregate", "total_tokens": 5000, "total_tool_calls": 42, "active_count": 3},
)
assert c._nodes["node-a"].aggregate["total_tokens"] == 5000
assert c._nodes["node-a"].aggregate["total_tool_calls"] == 42
def test_mark_unreachable_preserves_workstreams(self):
"""Disconnection marks unreachable but preserves workstream data."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
reachable=True,
workstreams={"ws1": {"id": "ws1", "name": "existing", "state": "idle"}},
)
c._mark_unreachable("node-a")
assert c._nodes["node-a"].reachable is False
assert "ws1" in c._nodes["node-a"].workstreams
assert c._nodes["node-a"].workstreams["ws1"]["name"] == "existing"
class TestCollectorFanout:
"""SSE fan-out to registered listeners."""
def test_unregister_listener_stops_fanout(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a")
q = queue.Queue()
c.register_listener(q)
c.unregister_listener(q)
c._fanout({"type": "test"})
assert q.empty()
class TestCollectorQueries:
"""Query methods: get_overview, get_nodes, get_workstreams, get_node_detail."""
@pytest.fixture()
def populated_collector(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
workstreams={
"ws1": {
"id": "ws1",
"name": "alpha",
"state": "running",
"node": "node-a",
"title": "Task A",
"tokens": 5000,
"context_ratio": 0.2,
"activity": "",
"activity_state": "",
"tool_calls": 10,
},
"ws2": {
"id": "ws2",
"name": "beta",
"state": "idle",
"node": "node-a",
"title": "Task B",
"tokens": 2000,
"context_ratio": 0.1,
"activity": "",
"activity_state": "",
"tool_calls": 5,
},
},
aggregate={"total_tokens": 7000, "total_tool_calls": 15},
)
c._nodes["node-b"] = NodeSnapshot(
node_id="node-b",
server_url="http://b:8080",
workstreams={
"ws3": {
"id": "ws3",
"name": "gamma",
"state": "attention",
"node": "node-b",
"title": "Task C",
"tokens": 10000,
"context_ratio": 0.5,
"activity": "awaiting approval",
"activity_state": "approval",
"tool_calls": 20,
},
},
aggregate={"total_tokens": 10000, "total_tool_calls": 20},
)
return c
def test_get_overview(self, populated_collector):
o = populated_collector.get_overview()
assert o["nodes"] == 2
assert o["workstreams"] == 3
assert o["states"]["running"] == 1
assert o["states"]["idle"] == 1
assert o["states"]["attention"] == 1
assert o["aggregate"]["total_tokens"] == 17000
assert o["aggregate"]["total_tool_calls"] == 35
def test_get_nodes_sorted_by_activity(self, populated_collector):
nodes, total = populated_collector.get_nodes(sort_by="activity")
assert total == 2
# node-a has 1 running, node-b has 1 attention — both have activity=1
# order depends on tie-breaking but both should be present
ids = [n["node_id"] for n in nodes]
assert "node-a" in ids
assert "node-b" in ids
def test_get_nodes_pagination(self, populated_collector):
nodes, total = populated_collector.get_nodes(limit=1, offset=0)
assert len(nodes) == 1
assert total == 2
nodes2, _ = populated_collector.get_nodes(limit=1, offset=1)
assert len(nodes2) == 1
assert nodes2[0]["node_id"] != nodes[0]["node_id"]
def test_get_workstreams_no_filter(self, populated_collector):
ws, total = populated_collector.get_workstreams()
assert total == 3
assert len(ws) == 3
def test_get_workstreams_filter_by_state(self, populated_collector):
ws, total = populated_collector.get_workstreams(state="running")
assert total == 1
assert ws[0]["name"] == "alpha"
def test_get_workstreams_filter_by_node(self, populated_collector):
ws, total = populated_collector.get_workstreams(node="node-b")
assert total == 1
assert ws[0]["name"] == "gamma"
def test_get_workstreams_filter_by_search(self, populated_collector):
ws, total = populated_collector.get_workstreams(search="Task C")
assert total == 1
assert ws[0]["name"] == "gamma"
def test_get_workstreams_search_case_insensitive(self, populated_collector):
ws, total = populated_collector.get_workstreams(search="task c")
assert total == 1
def test_get_workstreams_pagination(self, populated_collector):
ws, total = populated_collector.get_workstreams(page=1, per_page=2)
assert len(ws) == 2
assert total == 3
ws2, _ = populated_collector.get_workstreams(page=2, per_page=2)
assert len(ws2) == 1
def test_get_workstreams_sorted_by_state(self, populated_collector):
ws, _ = populated_collector.get_workstreams(sort_by="state")
states = [w["state"] for w in ws]
# running before attention before idle
assert states.index("running") < states.index("attention") < states.index("idle")
def test_get_workstreams_combined_filters(self, populated_collector):
ws, total = populated_collector.get_workstreams(state="idle", node="node-a")
assert total == 1
assert ws[0]["name"] == "beta"
def test_get_node_detail_found(self, populated_collector):
detail = populated_collector.get_node_detail("node-a")
assert detail is not None
assert detail["node_id"] == "node-a"
assert len(detail["workstreams"]) == 2
def test_get_node_detail_not_found(self, populated_collector):
assert populated_collector.get_node_detail("nonexistent") is None
def test_get_snapshot_empty(self):
c = _make_collector()
snap = c.get_snapshot()
assert snap["nodes"] == []
assert snap["overview"]["nodes"] == 0
assert snap["overview"]["workstreams"] == 0
assert snap["overview"]["states"]["running"] == 0
assert "timestamp" in snap
def test_get_snapshot_with_nodes(self, populated_collector):
snap = populated_collector.get_snapshot()
assert len(snap["nodes"]) == 2
assert snap["overview"]["nodes"] == 2
assert snap["overview"]["workstreams"] == 3
assert snap["overview"]["states"]["running"] == 1
assert snap["overview"]["states"]["attention"] == 1
assert snap["overview"]["states"]["idle"] == 1
assert snap["overview"]["aggregate"]["total_tokens"] == 17000
assert snap["timestamp"] > 0
# Each node should embed its workstreams
node_ids = {n["node_id"] for n in snap["nodes"]}
assert node_ids == {"node-a", "node-b"}
for n in snap["nodes"]:
if n["node_id"] == "node-a":
assert len(n["workstreams"]) == 2
elif n["node_id"] == "node-b":
assert len(n["workstreams"]) == 1
def test_get_snapshot_consistency(self, populated_collector):
"""Snapshot overview should match get_overview()."""
snap = populated_collector.get_snapshot()
overview = populated_collector.get_overview()
assert snap["overview"]["nodes"] == overview["nodes"]
assert snap["overview"]["workstreams"] == overview["workstreams"]
assert snap["overview"]["states"] == overview["states"]
assert snap["overview"]["aggregate"] == overview["aggregate"]
assert snap["overview"]["version_drift"] == overview["version_drift"]
# ---------------------------------------------------------------------------
# Console HTTP server tests
# ---------------------------------------------------------------------------
class TestConsoleHTTPEndpoints:
"""Test console HTTP API endpoints with a mock collector."""
@pytest.fixture()
def mock_collector(self):
collector = MagicMock(spec=ClusterCollector)
collector.get_overview.return_value = {
"nodes": 3,
"workstreams": 15,
"states": {
"running": 5,
"thinking": 2,
"attention": 1,
"idle": 6,
"error": 1,
},
"aggregate": {"total_tokens": 50000, "total_tool_calls": 200},
}
collector.get_nodes.return_value = (
[
{
"node_id": "node-a",
"ws_total": 5,
"ws_running": 3,
"total_tokens": 20000,
}
],
1,
)
collector.get_workstreams.return_value = (
[{"id": "ws1", "name": "test", "state": "running", "node": "node-a"}],
1,
)
collector.get_node_detail.return_value = {
"node_id": "node-a",
"server_url": "http://a:8080",
"health": {},
"workstreams": [],
"aggregate": {},
}
collector.get_snapshot.return_value = {
"nodes": [
{
"node_id": "node-a",
"server_url": "http://a:8080",
"max_ws": 10,
"reachable": True,
"version": "0.5.0",
"health": {},
"aggregate": {"total_tokens": 50000, "total_tool_calls": 200},
"workstreams": [
{"id": "ws1", "name": "test", "state": "running", "node": "node-a"},
],
},
],
"overview": {
"nodes": 3,
"workstreams": 15,
"states": {"running": 5, "thinking": 2, "attention": 1, "idle": 6, "error": 1},
"aggregate": {"total_tokens": 50000, "total_tool_calls": 200},
"version_drift": False,
"versions": ["0.5.0"],
},
"timestamp": 1234567890.0,
}
return collector
@pytest.fixture()
def client(self, mock_collector):
from starlette.testclient import TestClient
from turnstone.console.server import _load_static, create_app
_load_static()
app = create_app(
collector=mock_collector,
jwt_secret=_TEST_JWT_SECRET,
)
client = TestClient(app, raise_server_exceptions=False, headers=_TEST_AUTH_HEADERS)
yield client
client.close()
def _get(self, client, path):
resp = client.get(path)
return resp.status_code, resp.json()
def _get_raw(self, client, path):
resp = client.get(path)
return resp.status_code, resp.text, resp.headers.get("content-type")
def test_get_overview(self, client, mock_collector):
status, data = self._get(client, "/v1/api/cluster/overview")
assert status == 200
assert data["nodes"] == 3
assert data["workstreams"] == 15
assert data["states"]["running"] == 5
mock_collector.get_overview.assert_called_once()
def test_get_nodes(self, client, mock_collector):
status, data = self._get(client, "/v1/api/cluster/nodes?sort=activity&limit=10&offset=0")
assert status == 200
assert len(data["nodes"]) == 1
assert data["total"] == 1
mock_collector.get_nodes.assert_called_once_with(
sort_by="activity", limit=10, offset=0, node_ids=None
)
def test_get_workstreams(self, client, mock_collector):
status, data = self._get(
client, "/v1/api/cluster/workstreams?state=running&page=1&per_page=25"
)
assert status == 200
assert len(data["workstreams"]) == 1
assert data["total"] == 1
assert data["page"] == 1
assert data["pages"] == 1
mock_collector.get_workstreams.assert_called_once_with(
state="running",
node=None,
search=None,
sort_by="state",
page=1,
per_page=25,
extra_rows=[],
)
def test_get_workstreams_per_page_capped(self, client, mock_collector):
self._get(client, "/v1/api/cluster/workstreams?per_page=999")
call_kwargs = mock_collector.get_workstreams.call_args
assert call_kwargs.kwargs["per_page"] == 200
def test_get_node_detail(self, client, mock_collector):
status, data = self._get(client, "/v1/api/cluster/node/node-a")
assert status == 200
assert data["node_id"] == "node-a"
mock_collector.get_node_detail.assert_called_once_with("node-a")
def test_get_node_detail_not_found(self, client, mock_collector):
mock_collector.get_node_detail.return_value = None
status, data = self._get(client, "/v1/api/cluster/node/nonexistent")
assert status == 404
assert "error" in data
def test_get_snapshot(self, client, mock_collector):
status, data = self._get(client, "/v1/api/cluster/snapshot")
assert status == 200
assert len(data["nodes"]) == 1
assert data["nodes"][0]["node_id"] == "node-a"
assert data["overview"]["nodes"] == 3
assert data["overview"]["workstreams"] == 15
assert data["timestamp"] == 1234567890.0
mock_collector.get_snapshot.assert_called_once()
def test_health_endpoint(self, client, mock_collector):
status, data = self._get(client, "/health")
assert status == 200
assert data["status"] == "ok"
assert data["service"] == "turnstone-console"
assert data["nodes"] == 3
def test_index_html(self, client):
status, body, ct = self._get_raw(client, "/")
assert status == 200
assert "text/html" in ct
assert "turnstone console" in body
def test_static_css(self, client):
status, body, ct = self._get_raw(client, "/static/style.css")
assert status == 200
assert "text/css" in ct
def test_static_js(self, client):
status, body, ct = self._get_raw(client, "/static/app.js")
assert status == 200
assert "javascript" in ct
def test_404(self, client):
resp = client.get("/nonexistent")
assert resp.status_code == 404
def test_index_landing_surfaces(self, client):
status, body, ct = self._get_raw(client, "/")
assert status == 200
# Node discovery moved to the L-shell RAIL; the legacy bottom-bar node
# picker (and #cluster-status-bar) was retired by the renovation — guard
# against reintroduction (mirrors test_shell_js bottom-bar-retired).
assert 'id="csb-node-picker"' not in body
assert 'id="cluster-status-bar"' not in body
# The landing now boots the shared shell module, which builds the rail +
# tab-bar + pane host and hands off to the legacy boot.
assert "/shared/shell.js" in body
# Removed in the 1.5.0 landing-page cleanup — guard against
# accidental reintroduction.
assert 'id="new-ws-overlay"' not in body
assert 'id="new-ws-btn"' not in body
assert 'id="cluster-summary-compact"' not in body
assert 'id="view-node"' not in body
# Replaced by the rail — guard against reintroduction.
assert 'id="view-overview"' not in body
assert 'id="node-table"' not in body
# ---------------------------------------------------------------------------
# Version tracking / drift detection
# ---------------------------------------------------------------------------
class TestCollectorVersionInfo:
"""Version extraction and drift detection."""
def test_get_overview_no_drift(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a", health={"status": "ok", "version": "0.3.0"}
)
c._nodes["node-b"] = NodeSnapshot(
node_id="node-b", health={"status": "ok", "version": "0.3.0"}
)
overview = c.get_overview()
assert overview["version_drift"] is False
assert overview["versions"] == ["0.3.0"]
def test_get_overview_drift_detected(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a", health={"status": "ok", "version": "0.3.0"}
)
c._nodes["node-b"] = NodeSnapshot(
node_id="node-b", health={"status": "ok", "version": "0.3.1"}
)
overview = c.get_overview()
assert overview["version_drift"] is True
assert sorted(overview["versions"]) == ["0.3.0", "0.3.1"]
def test_get_overview_no_version_in_health(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", health={"status": "ok"})
overview = c.get_overview()
assert overview["version_drift"] is False
assert overview["versions"] == []
def test_get_overview_single_node_no_drift(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", health={"version": "0.3.0"})
overview = c.get_overview()
assert overview["version_drift"] is False
assert overview["versions"] == ["0.3.0"]
def test_get_nodes_includes_version(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
health={"status": "ok", "version": "0.3.0"},
)
nodes, _ = c.get_nodes()
assert nodes[0]["version"] == "0.3.0"
def test_get_nodes_version_empty_when_missing(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080", health={})
nodes, _ = c.get_nodes()
assert nodes[0]["version"] == ""
def test_get_version_info(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", health={"version": "0.3.0"})
c._nodes["node-b"] = NodeSnapshot(node_id="node-b", health={"version": "0.3.1"})
info = c.get_version_info()
assert info["drift"] is True
assert info["versions"]["node-a"] == "0.3.0"
assert info["versions"]["node-b"] == "0.3.1"
assert sorted(info["unique_versions"]) == ["0.3.0", "0.3.1"]
def test_get_version_info_no_drift(self):
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", health={"version": "0.3.0"})
c._nodes["node-b"] = NodeSnapshot(node_id="node-b", health={"version": "0.3.0"})
info = c.get_version_info()
assert info["drift"] is False
assert info["unique_versions"] == ["0.3.0"]
# ---------------------------------------------------------------------------
# Workstream creation tests
# ---------------------------------------------------------------------------
class TestConsoleWorkstreamCreation:
"""Tests for POST /v1/api/cluster/workstreams/new (HTTP dispatch)."""
@pytest.fixture()
def mock_collector(self):
collector = MagicMock(spec=ClusterCollector)
collector.get_overview.return_value = {
"nodes": 2,
"workstreams": 5,
"states": {"running": 1, "idle": 4, "thinking": 0, "attention": 0, "error": 0},
"aggregate": {"total_tokens": 0, "total_tool_calls": 0},
}
collector.get_node_detail.return_value = {
"node_id": "node-a",
"server_url": "http://a:8080",
"health": {},
"workstreams": [],
"aggregate": {},
"reachable": True,
}
collector.get_nodes.return_value = (
[
{"node_id": "node-a", "reachable": True, "max_ws": 10, "ws_total": 8},
{"node_id": "node-b", "reachable": True, "max_ws": 10, "ws_total": 3},
],
2,
)
# get_all_nodes delegates to get_nodes (mirrors real implementation)
collector.get_all_nodes.side_effect = lambda: collector.get_nodes.return_value[0]
return collector
@pytest.fixture()
def client_and_mock(self, mock_collector):
"""Returns (TestClient, mock_proxy_post) where mock_proxy_post is the
patched proxy_client.post that captures outgoing HTTP calls."""
import httpx
from starlette.testclient import TestClient
from turnstone.console.server import _load_static, create_app
_load_static()
app = create_app(
collector=mock_collector,
jwt_secret=_TEST_JWT_SECRET,
)
# Set up a mock proxy_client (lifespan doesn't run in TestClient)
async def _mock_post(*args, **kwargs):
return httpx.Response(
200,
json={"ws_id": "ws_new_123", "name": "test"},
request=httpx.Request("POST", args[0] if args else "http://test"),
)
mock_post = MagicMock(side_effect=_mock_post)
mock_proxy = MagicMock(spec=httpx.AsyncClient)
mock_proxy.post = mock_post
app.state.proxy_client = mock_proxy
client = TestClient(app, raise_server_exceptions=False, headers=_TEST_AUTH_HEADERS)
yield client, mock_post
client.close()
def test_create_with_explicit_node(self, client_and_mock, mock_collector):
client, mock_post = client_and_mock
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"node_id": "node-a", "name": "test-ws"},
)
assert resp.status_code == 200
data = resp.json()
assert data["status"] == "ok"
assert data["target_node"] == "node-a"
assert "correlation_id" in data
mock_post.assert_called_once()
# Verify the HTTP call was to the right node
call_args = mock_post.call_args
assert "http://a:8080/v1/api/workstreams/new" in call_args[0]
body = call_args[1]["json"]
assert body["name"] == "test-ws"
def test_create_with_model(self, client_and_mock, mock_collector):
client, mock_post = client_and_mock
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"node_id": "node-a", "model": "gpt-5"},
)
assert resp.status_code == 200
body = mock_post.call_args[1]["json"]
assert body["model"] == "gpt-5"
def test_create_with_initial_message_directed(self, client_and_mock, mock_collector):
client, mock_post = client_and_mock
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"node_id": "node-a", "initial_message": "Do the thing"},
)
assert resp.status_code == 200
body = mock_post.call_args[1]["json"]
assert body["initial_message"] == "Do the thing"
def test_create_with_initial_message_pool(self, client_and_mock, mock_collector):
"""Pool mode picks the best node and dispatches via HTTP."""
client, mock_post = client_and_mock
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"node_id": "pool", "initial_message": "Pool task"},
)
assert resp.status_code == 200
body = mock_post.call_args[1]["json"]
assert body["initial_message"] == "Pool task"
def test_create_auto_selects_best_node(self, client_and_mock, mock_collector):
client, mock_post = client_and_mock
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"name": "auto-test"},
)
assert resp.status_code == 200
data = resp.json()
# node-b has more headroom (10-3=7 vs 10-8=2)
assert data["target_node"] == "node-b"
def test_create_no_reachable_nodes(self, client_and_mock, mock_collector):
client, mock_post = client_and_mock
mock_collector.get_nodes.return_value = ([], 0)
resp = client.post("/v1/api/cluster/workstreams/new", json={})
assert resp.status_code == 503
assert "No reachable nodes" in resp.json()["error"]
def test_create_unknown_node(self, client_and_mock, mock_collector):
client, mock_post = client_and_mock
mock_collector.get_node_detail.return_value = None
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"node_id": "nonexistent"},
)
assert resp.status_code == 404
def test_create_invalid_json(self, client_and_mock):
client, mock_post = client_and_mock
resp = client.post(
"/v1/api/cluster/workstreams/new",
content=b"not json",
headers={"Content-Type": "application/json"},
)
assert resp.status_code == 400
def test_create_dispatches_to_correct_node_url(self, client_and_mock, mock_collector):
client, mock_post = client_and_mock
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"node_id": "node-a"},
)
assert resp.status_code == 200
call_args = mock_post.call_args
assert "http://a:8080/v1/api/workstreams/new" in call_args[0]
def test_create_pool_picks_best_node(self, client_and_mock, mock_collector):
"""Pool mode dispatches to the best available node."""
client, mock_post = client_and_mock
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"node_id": "pool", "name": "pool-task"},
)
assert resp.status_code == 200
data = resp.json()
assert data["status"] == "ok"
# Pool picks best node (node-b has most headroom)
assert data["target_node"] == "node-b"
def test_create_pool_no_nodes_returns_503(self, client_and_mock, mock_collector):
"""Pool mode with no reachable nodes returns 503."""
client, mock_post = client_and_mock
mock_collector.get_nodes.return_value = ([], 0)
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"node_id": "pool"},
)
assert resp.status_code == 503
def test_create_with_resume_ws_directed(self, client_and_mock, mock_collector):
"""resume_ws is forwarded in directed dispatch."""
client, mock_post = client_and_mock
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"node_id": "node-a", "resume_ws": "old-ws-id-123"},
)
assert resp.status_code == 200
body = mock_post.call_args[1]["json"]
assert body["resume_ws"] == "old-ws-id-123"
def test_create_with_resume_ws_auto(self, client_and_mock, mock_collector):
"""resume_ws is forwarded in auto-select dispatch."""
client, mock_post = client_and_mock
resp = client.post(
"/v1/api/cluster/workstreams/new",
json={"resume_ws": "old-ws-id-789"},
)
assert resp.status_code == 200
body = mock_post.call_args[1]["json"]
assert body["resume_ws"] == "old-ws-id-789"
# ---------------------------------------------------------------------------
# Proxy tests
# ---------------------------------------------------------------------------
class TestConsoleProxy:
"""Tests for /node/{node_id}/ reverse proxy."""
@pytest.fixture()
def mock_collector(self):
collector = MagicMock(spec=ClusterCollector)
collector.get_overview.return_value = {
"nodes": 1,
"workstreams": 2,
"states": {"running": 0, "idle": 2, "thinking": 0, "attention": 0, "error": 0},
"aggregate": {"total_tokens": 0, "total_tool_calls": 0},
}
collector.get_node_detail.return_value = {
"node_id": "node-a",
"server_url": "http://a:8080",
"health": {},
"workstreams": [],
"aggregate": {},
"reachable": True,
}
return collector
@pytest.fixture()
def client(self, mock_collector):
from starlette.testclient import TestClient
from turnstone.console.server import _load_static, create_app
_load_static()
app = create_app(
collector=mock_collector,
jwt_secret=_TEST_JWT_SECRET,
)
client = TestClient(app, raise_server_exceptions=False, headers=_TEST_AUTH_HEADERS)
yield client
client.close()
def test_proxy_unknown_node_returns_404(self, client, mock_collector):
mock_collector.get_node_detail.return_value = None
resp = client.get("/node/unknown/")
assert resp.status_code == 404
def test_proxy_static_unknown_node_returns_404(self, client, mock_collector):
mock_collector.get_node_detail.return_value = None
resp = client.get("/node/unknown/static/app.js")
assert resp.status_code == 404
def test_proxy_api_unknown_node_returns_404(self, client, mock_collector):
mock_collector.get_node_detail.return_value = None
resp = client.get("/node/unknown/api/workstreams")
assert resp.status_code == 404
def test_proxy_api_post_unknown_node_returns_404(self, client, mock_collector):
mock_collector.get_node_detail.return_value = None
resp = client.post(
"/node/unknown/api/send",
json={"message": "hello", "ws_id": "ws1"},
)
assert resp.status_code == 404
def test_proxy_api_per_ws_events_routes_to_sse_handler(self, client, mock_collector):
"""``/node/{node_id}/v1/api/workstreams/{ws_id}/events`` is the
per-workstream SSE stream the interactive WebUI subscribes to.
Without explicit detection, the path falls through to the
regular GET branch and the EventSource API can't consume the
one-shot response — Firefox surfaces it as "can't establish a
connection". Regression guard for the legacy URL surface
removal (#422) that moved per-ws SSE under
``/workstreams/{ws_id}/events`` without updating the proxy."""
from unittest.mock import AsyncMock, patch
from starlette.responses import Response
mock_collector.get_node_detail.return_value = {
"node_id": "node-a",
"server_url": "http://a:8080",
"reachable": True,
}
ws_id = "a" * 32
with (
patch(
"turnstone.console.server._proxy_sse",
new_callable=AsyncMock,
return_value=Response("ok", status_code=200),
) as sse_mock,
patch(
"turnstone.console.server._proxy_get",
new_callable=AsyncMock,
return_value=Response("ok", status_code=200),
) as get_mock,
):
client.get(f"/node/node-a/v1/api/workstreams/{ws_id}/events")
assert sse_mock.await_count == 1, (
"per-ws events path must route to _proxy_sse, not _proxy_get"
)
assert get_mock.await_count == 0
# Path passed to _proxy_sse must be the workstreams-prefixed
# form so the upstream URL is reconstructed correctly.
sse_args = sse_mock.await_args
assert sse_args.args[2] == f"workstreams/{ws_id}/events"
def test_proxy_api_global_events_still_routes_to_sse(self, client, mock_collector):
"""The bare ``events/global`` path was the only SSE path the
proxy recognized before the per-ws fix. Verify it still routes
correctly so the new branch didn't regress the existing case."""
from unittest.mock import AsyncMock, patch
from starlette.responses import Response
mock_collector.get_node_detail.return_value = {
"node_id": "node-a",
"server_url": "http://a:8080",
"reachable": True,
}
with patch(
"turnstone.console.server._proxy_sse",
new_callable=AsyncMock,
return_value=Response("ok", status_code=200),
) as sse_mock:
client.get("/node/node-a/v1/api/events/global")
assert sse_mock.await_count == 1
# events/global must use the console's service token —
# the upstream gates this path on `service` scope and
# end-user JWTs don't carry it. Without this, the
# browser's interactive UI 403-loops on every retry.
assert sse_mock.await_args.kwargs.get("use_service_auth") is True
def test_proxy_api_per_ws_events_uses_user_auth_not_service(self, client, mock_collector):
"""Per-ws events route uses the user's re-minted JWT, not the
service token — the upstream per-ws SSE handler scopes by
user identity for tenant filtering, and a service-scoped
call would bypass that gate. Only ``events/global``
(cross-tenant inventory by design) opts into service auth."""
from unittest.mock import AsyncMock, patch
from starlette.responses import Response
mock_collector.get_node_detail.return_value = {
"node_id": "node-a",
"server_url": "http://a:8080",
"reachable": True,
}
ws_id = "b" * 32
with patch(
"turnstone.console.server._proxy_sse",
new_callable=AsyncMock,
return_value=Response("ok", status_code=200),
) as sse_mock:
client.get(f"/node/node-a/v1/api/workstreams/{ws_id}/events")
assert sse_mock.await_count == 1
assert sse_mock.await_args.kwargs.get("use_service_auth") is False
# -------------------------------------------------------------------
# Proxied auth endpoints — handled locally by the console, not
# forwarded to the upstream node. Cases derive directly from
# ``_PROXY_AUTH_LOCAL_HANDLERS`` so a new dispatch entry can't be
# added without a matching test (or vice versa). See proxy_api's
# docstring for the JWT-audience reasoning.
# -------------------------------------------------------------------
@pytest.mark.parametrize(
("method", "path", "handler_name"),
[
(method, path, handler_name)
for (method, path), handler_name in sorted(_PROXY_AUTH_LOCAL_HANDLERS.items())
],
)
def test_proxy_auth_endpoint_dispatches_to_local_handler(
self, client, method, path, handler_name
):
"""Every entry in ``_PROXY_AUTH_LOCAL_HANDLERS`` must route to its
local console handler and never reach the upstream proxy. The
lockout class of bug this dispatch was added to fix is exactly
what a regression here would reintroduce silently — covering all
eight branches keeps each path tied to its handler."""
from unittest.mock import AsyncMock, patch
from starlette.responses import JSONResponse
with (
patch(
f"turnstone.console.server.{handler_name}",
new_callable=AsyncMock,
return_value=JSONResponse({"status": "ok"}),
) as local_mock,
patch(
"turnstone.console.server._proxy_post",
new_callable=AsyncMock,
return_value=JSONResponse({"status": "should-not-be-called"}),
) as post_mock,
patch(
"turnstone.console.server._proxy_get",
new_callable=AsyncMock,
return_value=JSONResponse({"status": "should-not-be-called"}),
) as get_mock,
):
resp = client.request(method, f"/node/node-a/v1/api/{path}")
assert resp.status_code == 200
assert local_mock.await_count == 1
assert post_mock.await_count == 0
assert get_mock.await_count == 0
def test_proxy_auth_login_works_without_cookie(self, mock_collector):
"""Without this fix the AuthMiddleware 401s before any handler
runs — the user is locked out of the proxied UI once the cookie
expires. Test bypasses _TEST_AUTH_HEADERS to reproduce."""
from unittest.mock import AsyncMock, patch
from starlette.responses import JSONResponse
from starlette.testclient import TestClient
from turnstone.console.server import _load_static, create_app
_load_static()
app = create_app(collector=mock_collector, jwt_secret=_TEST_JWT_SECRET)
unauth_client = TestClient(app, raise_server_exceptions=False)
try:
with patch(
"turnstone.console.server.auth_login",
new_callable=AsyncMock,
return_value=JSONResponse({"status": "ok"}),
) as local_mock:
resp = unauth_client.post(
"/node/node-a/v1/api/auth/login",
json={"username": "x", "password": "y"},
)
# AuthMiddleware must classify the proxied login path as
# public (is_public_path change) AND proxy_api must
# dispatch to the local handler (proxy_api change).
assert resp.status_code == 200, (
f"login locked out: got {resp.status_code}, body={resp.text}"
)
assert local_mock.await_count == 1
finally:
unauth_client.close()
def test_proxy_auth_wrong_method_returns_405_not_forwarded(self, client):
"""A non-canonical method on an auth path (e.g. PUT on auth/login)
must short-circuit with 405 instead of falling through to the
upstream proxy — falling through would forward the request
authenticated as the console's service token (``_proxy_auth_headers``
fallback)."""
from unittest.mock import AsyncMock, patch
from starlette.responses import JSONResponse
with (
patch(
"turnstone.console.server._proxy_post",
new_callable=AsyncMock,
return_value=JSONResponse({"status": "should-not-be-called"}),
) as post_mock,
patch(
"turnstone.console.server._proxy_get",
new_callable=AsyncMock,
return_value=JSONResponse({"status": "should-not-be-called"}),
) as get_mock,
):
# PUT on a POST-only auth path → 405
put_resp = client.put("/node/node-a/v1/api/auth/login")
assert put_resp.status_code == 405
# POST on a GET-only auth path → 405
post_resp = client.post("/node/node-a/v1/api/auth/status")
assert post_resp.status_code == 405
assert post_mock.await_count == 0
assert get_mock.await_count == 0
def test_proxy_non_auth_endpoint_still_forwarded(self, client, mock_collector):
"""Sanity: only auth/* paths intercept. Other API paths still
forward to the upstream node."""
from unittest.mock import AsyncMock, patch
from starlette.responses import JSONResponse
mock_collector.get_node_detail.return_value = {
"node_id": "node-a",
"server_url": "http://a:8080",
"reachable": True,
}
with patch(
"turnstone.console.server._proxy_get",
new_callable=AsyncMock,
return_value=JSONResponse({"ok": True}),
) as proxy_mock:
resp = client.get("/node/node-a/v1/api/workstreams")
assert resp.status_code == 200
assert proxy_mock.await_count == 1
# ---------------------------------------------------------------------------
# Proxy URL rewriting unit tests (no HTTP needed)
# ---------------------------------------------------------------------------
class TestProxyRewriting:
"""Test the JS shim and HTML rewriting logic."""
def test_js_shim_contains_prefix_placeholder(self):
from turnstone.console.server import _JS_PROXY_SHIM
assert "PREFIX_PLACEHOLDER" in _JS_PROXY_SHIM
replaced = _JS_PROXY_SHIM.replace("PREFIX_PLACEHOLDER", "/node/my-node")
assert "/node/my-node" in replaced
assert "PREFIX_PLACEHOLDER" not in replaced
def test_js_shim_overrides_fetch_and_eventsource(self):
from turnstone.console.server import _JS_PROXY_SHIM
assert "window.fetch" in _JS_PROXY_SHIM
assert "window.EventSource" in _JS_PROXY_SHIM
def test_js_shim_carries_node_id_placeholder(self):
"""The picker reads the current node_id from the shim's _nodeId
closure variable; the placeholder must be present and substitutable."""
from turnstone.console.server import _JS_PROXY_SHIM
assert "NODE_ID_PLACEHOLDER" in _JS_PROXY_SHIM
replaced = _JS_PROXY_SHIM.replace("NODE_ID_PLACEHOLDER", "node-a")
assert "node-a" in replaced
assert "NODE_ID_PLACEHOLDER" not in replaced
def test_js_shim_includes_picker_pieces(self):
"""Picker logic ships in the same IIFE as the prefix shim — verify
the moving parts are present so a future refactor doesn't silently
drop them. /v1/api/cluster/nodes is the lazy-fetch target;
#ui-header is the DOM anchor; console-node-pill is the trigger
class; ws-tab-dropdown is the menu shell we share with the
workstream chevron menu (style + behaviour parity); ArrowDown is
the keyboard-nav primitive that disambiguates this from a plain
click-only menu."""
from turnstone.console.server import _JS_PROXY_SHIM
# limit=1000 matches the collector's hard cap; without it the
# picker would silently drop nodes past the 100-default in
# clusters with >100 nodes.
assert "/v1/api/cluster/nodes?limit=1000" in _JS_PROXY_SHIM
assert "ui-header" in _JS_PROXY_SHIM
assert "console-node-pill" in _JS_PROXY_SHIM
assert "ws-tab-dropdown" in _JS_PROXY_SHIM
assert "ArrowDown" in _JS_PROXY_SHIM
assert "DOMContentLoaded" in _JS_PROXY_SHIM
def test_proxy_style_drops_banner_styles(self):
"""The legacy banner CSS classes (.console-banner, .ts-header-back-link
offsets, .dashboard-overlay top:32px hack) should be gone — the new
picker lives inside #ui-header and doesn't need overlay offsets."""
from turnstone.console.server import _CONSOLE_PROXY_STYLE
assert ".console-banner" not in _CONSOLE_PROXY_STYLE
assert "dashboard-overlay" not in _CONSOLE_PROXY_STYLE
assert ".console-node-pill" in _CONSOLE_PROXY_STYLE
assert ".console-node-menu" in _CONSOLE_PROXY_STYLE
def test_proxy_style_uses_canonical_degraded_color(self):
"""Degraded health dot must use --accent (the canonical "needs
attention" token used by the cluster-overview node table at
console/static/style.css:548) and not --yellow. Yellow is reserved
for the dash-state attention dot, a stronger signal."""
from turnstone.console.server import _CONSOLE_PROXY_STYLE
assert "console-node-menu-item-dot--degraded" in _CONSOLE_PROXY_STYLE
# The degraded rule sits on its own line; assert it uses --accent
# by checking the CSS substring has --accent and not --yellow.
idx = _CONSOLE_PROXY_STYLE.find("console-node-menu-item-dot--degraded")
rule = _CONSOLE_PROXY_STYLE[idx : idx + 200]
assert "var(--accent)" in rule
assert "var(--yellow)" not in rule
def test_html_rewriting_changes_static_paths(self):
"""Simulate the proxy_index rewriting logic."""
sample_html = (
'<link rel="stylesheet" href="/static/style.css">\n'
'<script src="/static/app.js"></script>'
)
prefix = "/node/test-node"
rewritten = sample_html.replace('href="/static/', f'href="{prefix}/static/')
rewritten = rewritten.replace('src="/static/', f'src="{prefix}/static/')
assert "/node/test-node/static/style.css" in rewritten
assert "/node/test-node/static/app.js" in rewritten
# Originals should be gone
assert 'href="/static/' not in rewritten
assert 'src="/static/' not in rewritten
def test_shim_injection_after_body(self):
"""Simulate the proxy shim injection — the shim ships the node-id
and prefix as JS literals and renders the picker at runtime, so
we assert the substituted JS literals land in the page."""
from turnstone.console.server import _CONSOLE_PROXY_STYLE, _JS_PROXY_SHIM
sample_html = "<html><body><div>content</div></body></html>"
prefix = "/node/node-a"
shim_js = _JS_PROXY_SHIM.replace('"PREFIX_PLACEHOLDER"', json.dumps(prefix)).replace(
'"NODE_ID_PLACEHOLDER"', json.dumps("node-a")
)
injection = _CONSOLE_PROXY_STYLE + "<script>" + shim_js + "</script>"
result = sample_html.replace("<body>", "<body>" + injection, 1)
assert '"node-a"' in result
assert '"/node/node-a"' in result
assert "PREFIX_PLACEHOLDER" not in result
assert "NODE_ID_PLACEHOLDER" not in result
assert result.startswith("<html><body><style>")
# ---------------------------------------------------------------------------
# _pick_best_node unit tests
# ---------------------------------------------------------------------------
class TestPickBestNode:
"""Test the _pick_best_node helper."""
@staticmethod
def _mock_collector(nodes: list) -> MagicMock:
collector = MagicMock(spec=ClusterCollector)
collector.get_nodes.return_value = (nodes, len(nodes))
collector.get_all_nodes.side_effect = lambda: collector.get_nodes.return_value[0]
return collector
def test_picks_node_with_most_headroom(self):
from turnstone.console.server import _pick_best_node
collector = self._mock_collector(
[
{"node_id": "busy", "reachable": True, "max_ws": 10, "ws_total": 9},
{"node_id": "free", "reachable": True, "max_ws": 10, "ws_total": 2},
{"node_id": "mid", "reachable": True, "max_ws": 10, "ws_total": 5},
]
)
assert _pick_best_node(collector) == "free"
def test_skips_unreachable_nodes(self):
from turnstone.console.server import _pick_best_node
collector = self._mock_collector(
[
{"node_id": "down", "reachable": False, "max_ws": 10, "ws_total": 0},
{"node_id": "up", "reachable": True, "max_ws": 10, "ws_total": 5},
]
)
assert _pick_best_node(collector) == "up"
def test_returns_empty_when_no_nodes(self):
from turnstone.console.server import _pick_best_node
collector = self._mock_collector([])
assert _pick_best_node(collector) == ""
def test_returns_empty_when_all_unreachable(self):
from turnstone.console.server import _pick_best_node
collector = self._mock_collector(
[
{"node_id": "down", "reachable": False, "max_ws": 10, "ws_total": 0},
]
)
assert _pick_best_node(collector) == ""
# ---------------------------------------------------------------------------
# Version tracking endpoint tests
# ---------------------------------------------------------------------------
class TestConsoleVersionEndpoints:
"""HTTP endpoint tests for version drift fields."""
@pytest.fixture()
def mock_collector(self):
collector = MagicMock(spec=ClusterCollector)
collector.get_overview.return_value = {
"nodes": 2,
"workstreams": 5,
"states": {"running": 1, "thinking": 0, "attention": 0, "idle": 4, "error": 0},
"aggregate": {"total_tokens": 10000, "total_tool_calls": 50},
"version_drift": True,
"versions": ["0.3.0", "0.3.1"],
}
return collector
@pytest.fixture()
def client(self, mock_collector):
from starlette.testclient import TestClient
from turnstone.console.server import _load_static, create_app
_load_static()
app = create_app(
collector=mock_collector,
jwt_secret=_TEST_JWT_SECRET,
)
client = TestClient(app, raise_server_exceptions=False, headers=_TEST_AUTH_HEADERS)
yield client
client.close()
def _get(self, client, path):
resp = client.get(path)
return resp.status_code, resp.json()
def test_overview_includes_version_drift(self, client, mock_collector):
status, data = self._get(client, "/v1/api/cluster/overview")
assert status == 200
assert data["version_drift"] is True
assert "0.3.0" in data["versions"]
assert "0.3.1" in data["versions"]
def test_health_includes_version_drift(self, client, mock_collector):
status, data = self._get(client, "/health")
assert status == 200
assert data["version_drift"] is True
assert "0.3.0" in data["versions"]
# ---------------------------------------------------------------------------
# Shared static serving
# ---------------------------------------------------------------------------
class TestSharedStatic:
"""Tests for /shared/ static file serving."""
@pytest.fixture()
def client(self):
from starlette.testclient import TestClient
from turnstone.console.server import _load_static, create_app
_load_static()
collector = MagicMock(spec=ClusterCollector)
collector.get_overview.return_value = {
"nodes": 0,
"workstreams": 0,
"states": {},
"aggregate": {},
}
app = create_app(
collector=collector,
jwt_secret=_TEST_JWT_SECRET,
)
client = TestClient(app, raise_server_exceptions=False, headers=_TEST_AUTH_HEADERS)
yield client
client.close()
def test_shared_base_css(self, client):
resp = client.get("/shared/base.css")
assert resp.status_code == 200
assert "text/css" in resp.headers.get("content-type", "")
def test_shared_utils_js(self, client):
resp = client.get("/shared/utils.js")
assert resp.status_code == 200
assert "javascript" in resp.headers.get("content-type", "")
def test_shared_auth_js(self, client):
resp = client.get("/shared/auth.js")
assert resp.status_code == 200
assert "javascript" in resp.headers.get("content-type", "")
def test_shared_toast_js(self, client):
resp = client.get("/shared/toast.js")
assert resp.status_code == 200
def test_shared_theme_js(self, client):
resp = client.get("/shared/theme.js")
assert resp.status_code == 200
def test_shared_kb_js(self, client):
resp = client.get("/shared/kb.js")
assert resp.status_code == 200
def test_shared_nonexistent_returns_404(self, client):
resp = client.get("/shared/nonexistent.js")
assert resp.status_code == 404
def test_index_imports_shared_base_css(self, client):
resp = client.get("/")
assert resp.status_code == 200
assert "/shared/base.css?v=" in resp.text
def test_index_imports_shared_scripts(self, client):
resp = client.get("/")
body = resp.text
assert "/shared/utils.js" in body
assert "/shared/toast.js" in body
assert "/shared/theme.js" in body
assert "/shared/auth.js" in body
assert "/shared/kb.js" in body
def test_shared_scripts_load_before_app_js(self, client):
"""Shared scripts must appear before page-specific app.js."""
body = client.get("/").text
shared_pos = body.find("/shared/utils.js")
app_pos = body.find("/static/app.js")
assert shared_pos < app_pos
def test_index_cache_control_no_cache(self, client):
resp = client.get("/")
assert resp.headers.get("cache-control") == "no-cache"
def test_index_etag_present(self, client):
resp = client.get("/")
assert resp.headers.get("etag")
def test_index_etag_304(self, client):
resp = client.get("/")
etag = resp.headers.get("etag")
resp2 = client.get("/", headers={"If-None-Match": etag})
assert resp2.status_code == 304
class TestProxySharedStatic:
"""Tests for proxy rewriting of /shared/ paths."""
def test_html_rewriting_includes_shared_paths(self):
"""Verify proxy_index rewrites /shared/ paths like /static/ paths."""
sample_html = (
'<link rel="stylesheet" href="/shared/base.css">\n'
'<link rel="stylesheet" href="/static/style.css">\n'
'<script src="/shared/utils.js"></script>\n'
'<script src="/static/app.js"></script>'
)
prefix = "/node/test-node"
rewritten = sample_html.replace('href="/static/', f'href="{prefix}/static/')
rewritten = rewritten.replace('src="/static/', f'src="{prefix}/static/')
rewritten = rewritten.replace('href="/shared/', f'href="{prefix}/shared/')
rewritten = rewritten.replace('src="/shared/', f'src="{prefix}/shared/')
assert "/node/test-node/shared/base.css" in rewritten
assert "/node/test-node/shared/utils.js" in rewritten
assert "/node/test-node/static/style.css" in rewritten
assert "/node/test-node/static/app.js" in rewritten
assert 'href="/shared/' not in rewritten
assert 'src="/shared/' not in rewritten
def test_proxy_shim_injected_in_html(self):
"""Verify shim is injected as inline script in proxied HTML."""
from turnstone.console.server import _JS_PROXY_SHIM
sample_html = "<html><body><div>content</div></body></html>"
prefix = "/node/test-node"
shim_js = _JS_PROXY_SHIM.replace('"PREFIX_PLACEHOLDER"', json.dumps(prefix)).replace(
'"NODE_ID_PLACEHOLDER"', json.dumps("test-node")
)
shim = "<script>" + shim_js + "</script>"
result = sample_html.replace("<body>", "<body>" + shim, 1)
assert "<script>" in result
assert "/node/test-node" in result
assert "window.fetch" in result
assert "window.EventSource" in result
def test_proxy_shared_static_unknown_node_returns_404(self):
from starlette.testclient import TestClient
from turnstone.console.server import _load_static, create_app
_load_static()
collector = MagicMock(spec=ClusterCollector)
collector.get_overview.return_value = {
"nodes": 0,
"workstreams": 0,
"states": {},
"aggregate": {},
}
collector.get_node_detail.return_value = None
app = create_app(
collector=collector,
jwt_secret=_TEST_JWT_SECRET,
)
client = TestClient(app, raise_server_exceptions=False, headers=_TEST_AUTH_HEADERS)
resp = client.get("/node/unknown/shared/base.css")
assert resp.status_code == 404
client.close()
# ---------------------------------------------------------------------------
# SSE proxy — raw byte passthrough
# ---------------------------------------------------------------------------
class TestSSEProxy:
"""Verify _proxy_sse forwards raw bytes including ping comments."""
def test_proxy_sse_preserves_pings_and_events(self):
"""SSE proxy should forward ping comments and events verbatim."""
from turnstone.console.server import _proxy_sse
# Simulate an upstream SSE response with a ping comment and a real event
sse_payload = b': ping - 2026-03-08T12:00:00Z\n\nevent: message\ndata: {"type": "test"}\n\n'
class FakeResponse:
status_code = 200
headers = {"content-type": "text/event-stream"}
async def aiter_bytes(self):
yield sse_payload
async def aclose(self):
pass
async def __aenter__(self):
return self
async def __aexit__(self, *args):
pass
class FakeClient:
def stream(self, method, url, **kwargs):
return FakeResponse()
class FakeRequest:
class url: # noqa: N801
query = "ws_id=test123"
class app: # noqa: N801
class state: # noqa: N801
proxy_sse_client = FakeClient()
proxy_auth_token = ""
headers = {}
async def is_disconnected(self):
return False
async def _run():
response = await _proxy_sse(
FakeRequest(), "http://fake:8080", "events", api_prefix="v1/api"
)
assert response.media_type == "text/event-stream"
# Collect the streamed bytes
chunks: list[bytes] = []
async for chunk in response.body_iterator:
chunks.append(chunk if isinstance(chunk, bytes) else chunk.encode())
body = b"".join(chunks)
# Ping comment must be preserved (not filtered)
assert b": ping" in body
# Real event must be preserved
assert b"event: message" in body
assert b'"type": "test"' in body
asyncio.run(_run())
def test_proxy_sse_upstream_error_status(self):
"""Non-200 upstream status should yield an error event."""
from turnstone.console.server import _proxy_sse
class FakeResponse:
status_code = 502
async def aiter_bytes(self):
return
yield # make it an async generator
async def aclose(self):
pass
async def __aenter__(self):
return self
async def __aexit__(self, *args):
pass
class FakeClient:
def stream(self, method, url, **kwargs):
return FakeResponse()
class FakeRequest:
class url: # noqa: N801
query = ""
class app: # noqa: N801
class state: # noqa: N801
proxy_sse_client = FakeClient()
proxy_auth_token = ""
headers = {}
async def is_disconnected(self):
return False
async def _run():
response = await _proxy_sse(FakeRequest(), "http://fake:8080", "events")
chunks: list[bytes] = []
async for chunk in response.body_iterator:
chunks.append(chunk if isinstance(chunk, bytes) else chunk.encode())
body = b"".join(chunks)
assert b"event: error" in body
assert b"502" in body
asyncio.run(_run())
def test_proxy_sse_disconnect_handling(self):
"""Proxy should stop when browser disconnects."""
from turnstone.console.server import _proxy_sse
class FakeResponse:
status_code = 200
async def aiter_bytes(self):
yield b"data: chunk1\n\n"
yield b"data: chunk2\n\n" # should not be reached
yield b"data: chunk3\n\n"
async def aclose(self):
pass
async def __aenter__(self):
return self
async def __aexit__(self, *args):
pass
class FakeClient:
def stream(self, method, url, **kwargs):
return FakeResponse()
call_count = 0
class FakeRequest:
class url: # noqa: N801
query = ""
class app: # noqa: N801
class state: # noqa: N801
proxy_sse_client = FakeClient()
proxy_auth_token = ""
headers = {}
async def is_disconnected(self):
nonlocal call_count
call_count += 1
return call_count > 1 # disconnect after first chunk
async def _run():
response = await _proxy_sse(FakeRequest(), "http://fake:8080", "events")
chunks: list[bytes] = []
async for chunk in response.body_iterator:
chunks.append(chunk if isinstance(chunk, bytes) else chunk.encode())
body = b"".join(chunks)
assert b"chunk1" in body
# Should have stopped before chunk3
assert b"chunk3" not in body
asyncio.run(_run())
# ---------------------------------------------------------------------------
# Proxy auth header propagation
# ---------------------------------------------------------------------------
class TestProxyAuthHeaders:
"""Verify _proxy_auth_headers mints user-scoped JWTs for proxy requests."""
SECRET = "test-secret-that-is-at-least-32-chars"
def _make_request(
self, *, auth_result=None, jwt_secret="", proxy_token_mgr=None, proxy_auth_token=""
):
"""Build a minimal fake request for _proxy_auth_headers."""
class _State:
pass
class _AppState:
pass
class _App:
state = _AppState()
class _Request:
state = _State()
app = _App()
req = _Request()
req.state.auth_result = auth_result
req.app.state.jwt_secret = jwt_secret
req.app.state.proxy_token_mgr = proxy_token_mgr
req.app.state.proxy_auth_token = proxy_auth_token
return req
def test_mints_user_jwt(self):
"""Real user auth_result → JWT with correct sub, scopes, src, aud, permissions."""
import jwt as pyjwt
from turnstone.console.server import _proxy_auth_headers
from turnstone.core.auth import JWT_AUD_SERVER, AuthResult
auth = AuthResult(
user_id="alice",
scopes=frozenset({"read", "write"}),
token_source="jwt",
permissions=frozenset({"admin.users"}),
)
req = self._make_request(auth_result=auth, jwt_secret=self.SECRET)
headers = _proxy_auth_headers(req)
assert "Authorization" in headers
token = headers["Authorization"].removeprefix("Bearer ")
payload = pyjwt.decode(token, self.SECRET, algorithms=["HS256"], audience=JWT_AUD_SERVER)
assert payload["sub"] == "alice"
assert set(payload["scopes"].split(",")) == {"read", "write"}
assert payload["src"] == "console-proxy"
assert payload["aud"] == JWT_AUD_SERVER
assert payload["permissions"] == "admin.users"
def test_narrows_scopes(self):
"""Read-only user → JWT carries only read scope, not full {read,write,approve}."""
import jwt as pyjwt
from turnstone.console.server import _proxy_auth_headers
from turnstone.core.auth import JWT_AUD_SERVER, AuthResult
auth = AuthResult(
user_id="viewer",
scopes=frozenset({"read"}),
token_source="jwt",
)
req = self._make_request(auth_result=auth, jwt_secret=self.SECRET)
headers = _proxy_auth_headers(req)
token = headers["Authorization"].removeprefix("Bearer ")
payload = pyjwt.decode(token, self.SECRET, algorithms=["HS256"], audience=JWT_AUD_SERVER)
assert payload["scopes"] == "read"
def test_short_expiry(self):
"""Minted JWT expires in 300 seconds, not hours."""
import jwt as pyjwt
from turnstone.console.server import _proxy_auth_headers
from turnstone.core.auth import AuthResult
auth = AuthResult(
user_id="alice",
scopes=frozenset({"read"}),
token_source="jwt",
)
req = self._make_request(auth_result=auth, jwt_secret=self.SECRET)
headers = _proxy_auth_headers(req)
token = headers["Authorization"].removeprefix("Bearer ")
payload = pyjwt.decode(
token, self.SECRET, algorithms=["HS256"], options={"verify_aud": False}
)
assert payload["exp"] - payload["iat"] == 300
def test_fallback_no_user(self):
"""No auth_result → falls back to ServiceTokenManager."""
from turnstone.console.server import _proxy_auth_headers
from turnstone.core.auth import ServiceTokenManager
mgr = ServiceTokenManager(
user_id="console-proxy",
scopes=frozenset({"read", "write", "approve"}),
source="console",
secret=self.SECRET,
)
req = self._make_request(proxy_token_mgr=mgr)
headers = _proxy_auth_headers(req)
assert "Authorization" in headers
assert headers["Authorization"] == f"Bearer {mgr.token}"
def test_fallback_no_secret(self):
"""auth_result present but empty jwt_secret → falls back to ServiceTokenManager."""
from turnstone.console.server import _proxy_auth_headers
from turnstone.core.auth import AuthResult, ServiceTokenManager
auth = AuthResult(
user_id="alice",
scopes=frozenset({"read"}),
token_source="jwt",
)
mgr = ServiceTokenManager(
user_id="console-proxy",
scopes=frozenset({"read", "write", "approve"}),
source="console",
secret=self.SECRET,
)
req = self._make_request(auth_result=auth, jwt_secret="", proxy_token_mgr=mgr)
headers = _proxy_auth_headers(req)
# Should use ServiceTokenManager, not mint a user JWT
assert headers["Authorization"] == f"Bearer {mgr.token}"
def test_no_mgr_no_user_returns_empty(self):
"""No auth_result, no ServiceTokenManager → empty headers."""
from turnstone.console.server import _proxy_auth_headers
req = self._make_request()
headers = _proxy_auth_headers(req)
assert headers == {}
# ---------------------------------------------------------------------------
# Server: trusted user_id forwarding on create_workstream
# ---------------------------------------------------------------------------
class TestCreateWorkstreamUserIdTrust:
"""Verify that only trusted service tokens can forward user_id in create_workstream."""
def _extract_uid(self, body: dict, auth_result) -> str:
"""Replicate the trust check from server.py:create_workstream."""
auth = auth_result
uid: str = getattr(auth, "user_id", "") or ""
trusted_sources = {"console"}
if (
body.get("user_id")
and isinstance(body["user_id"], str)
and auth is not None
and auth.token_source in trusted_sources
):
uid = body["user_id"]
return uid
def test_console_can_forward_user_id(self):
from turnstone.core.auth import AuthResult
auth = AuthResult(
user_id="console",
scopes=frozenset({"approve"}),
token_source="console",
)
uid = self._extract_uid({"user_id": "real-user-abc"}, auth)
assert uid == "real-user-abc"
def test_console_service_can_forward_user_id(self):
from turnstone.core.auth import AuthResult
auth = AuthResult(
user_id="console",
scopes=frozenset({"approve"}),
token_source="console",
)
uid = self._extract_uid({"user_id": "real-user-abc"}, auth)
assert uid == "real-user-abc"
def test_console_proxy_user_cannot_override_user_id(self):
"""End-user tokens via console-proxy must NOT override user_id."""
from turnstone.core.auth import AuthResult
auth = AuthResult(
user_id="real-user-abc",
scopes=frozenset({"read", "write"}),
token_source="console-proxy",
)
uid = self._extract_uid({"user_id": "impersonated-user"}, auth)
# Should use JWT identity, NOT the body override
assert uid == "real-user-abc"
def test_direct_user_cannot_override_user_id(self):
"""Direct JWT login must NOT override user_id."""
from turnstone.core.auth import AuthResult
auth = AuthResult(
user_id="real-user-abc",
scopes=frozenset({"read", "write"}),
token_source="password",
)
uid = self._extract_uid({"user_id": "impersonated-user"}, auth)
assert uid == "real-user-abc"
def test_no_body_user_id_uses_jwt(self):
from turnstone.core.auth import AuthResult
auth = AuthResult(
user_id="bridge",
scopes=frozenset({"approve"}),
token_source="bridge",
)
uid = self._extract_uid({"name": "test-ws"}, auth)
assert uid == "bridge"
# ---------------------------------------------------------------------------
# Collector — MCP aggregation in get_overview()
# ---------------------------------------------------------------------------
class TestCollectorMCPAggregation:
"""Verify MCP server/resource/prompt aggregation in overview and snapshot."""
def test_overview_mcp_aggregation(self):
"""Two nodes with MCP data produce correct sums in the overview."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
health={"mcp": {"servers": 2, "resources": 5, "prompts": 3}},
)
c._nodes["node-b"] = NodeSnapshot(
node_id="node-b",
server_url="http://b:8080",
health={"mcp": {"servers": 1, "resources": 4, "prompts": 2}},
)
overview = c.get_overview()
assert overview["mcp_servers"] == 3
assert overview["mcp_resources"] == 9
assert overview["mcp_prompts"] == 5
def test_overview_mcp_absent_when_zero(self):
"""Nodes without MCP data produce no mcp_servers key in the overview."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
health={"status": "ok"},
)
c._nodes["node-b"] = NodeSnapshot(
node_id="node-b",
server_url="http://b:8080",
health={},
)
overview = c.get_overview()
assert "mcp_servers" not in overview
assert "mcp_resources" not in overview
assert "mcp_prompts" not in overview
def test_overview_mcp_mixed_nodes(self):
"""One node with MCP, one without — only the MCP node contributes."""
c = _make_collector()
c._nodes["node-a"] = NodeSnapshot(
node_id="node-a",
server_url="http://a:8080",
health={"mcp": {"servers": 3, "resources": 10, "prompts": 7}},
)
c._nodes["node-b"] = NodeSnapshot(
node_id="node-b",
server_url="http://b:8080",
health={"status": "ok"},
)
overview = c.get_overview()
assert overview["mcp_servers"] == 3
assert overview["mcp_resources"] == 10
assert overview["mcp_prompts"] == 7