Quality overhaul: code tooling, CI/CD, architecture diagrams, UI rede… (#1)

* Quality overhaul: code tooling, CI/CD, architecture diagrams, UI redesign, and legacy cleanup

- Add ruff (lint+format) and mypy (strict) with zero errors across 37 source files
- Add GitHub Actions CI (lint, typecheck, test matrix 3.11/3.12/3.13) and PyPI publish workflow
- Create 12 PlantUML architecture diagrams with PNG renders covering all subsystems
- Refresh README and docs with badges, diagram links, and current descriptions
- Refactor test_server_live.py with mock streaming helpers for deterministic CI testing
- Update dependencies to current versions (openai>=2.24, httpx>=0.28, redis>=7.2)

Console dashboard:
- Move state indicators from top cards to fixed bottom status bar with cluster metrics
- Replace flat 50-node list with hostname-prefix grouped nodes (expand/collapse, up to 1000)
- Apply "Instrument Panel" visual redesign: IBM Plex Mono + Outfit fonts, warm amber accent,
  LED glow state indicators, deep charcoal surfaces, WCAG AA contrast compliance
- Add render cache, stale indicator, active filter highlight, loading states

Server web UI:
- Apply matching Instrument Panel aesthetic for visual consistency with console
- Fix branding (pcode → turnstone), extract inline styles to CSS classes
- Rename pcode localStorage keys and history state to turnstone

Legacy cleanup:
- Remove persona-model-specific --persona flag and /persona slash command
- Remove model_identity from chat_template_kwargs (vLLM-specific mechanism)
- Refactor plan agent to use standard developer message instead of model_identity
- Remove dead code (unused date/has_tools variables, noqa suppressions)

* Fix CI typecheck: add mypy overrides for optional sympy/numpy imports

The math sandbox optionally imports sympy and numpy at runtime (try/except
ImportError). In CI these packages are not installed, so mypy raises
import-not-found rather than import-untyped. Add mypy overrides to
ignore missing imports for these optional dependencies.

* Fix Copilot review findings: ARIA role, status bar cache, and pulse opacity

- Change #node-table from role="tree" to role="list" and group elements
  from role="treeitem" to role="listitem" (proper ARIA semantics)
- Include currentView and currentFilter.state in renderStatusBar cache key
  so active pill highlight updates when switching views
- Align pulse animation to 0.35 opacity (already applied in CSS)
This commit is contained in:
Patrick Buckley
2026-03-02 16:55:12 -08:00
committed by GitHub
parent 0d6252dd7d
commit 9be155b97a
86 changed files with 5038 additions and 1698 deletions
+43
View File
@@ -0,0 +1,43 @@
name: CI
on:
push:
branches: [main]
pull_request:
branches: [main]
jobs:
lint:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.13"
- run: pip install ruff
- run: ruff check turnstone/ tests/
- run: ruff format --check turnstone/ tests/
typecheck:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.13"
- run: pip install mypy types-redis
- run: pip install -e ".[mq]"
- run: mypy turnstone/
test:
runs-on: ubuntu-latest
strategy:
matrix:
python-version: ["3.11", "3.12", "3.13"]
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
- run: pip install -e ".[test,mq]"
- run: pytest tests/ -m "not live" -q
+21
View File
@@ -0,0 +1,21 @@
name: Publish to PyPI
on:
push:
tags: ["v*"]
permissions:
id-token: write
jobs:
publish:
runs-on: ubuntu-latest
environment: pypi
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.13"
- run: pip install build
- run: python -m build
- uses: pypa/gh-action-pypi-publish@release/v1
+27 -1
View File
@@ -1,5 +1,10 @@
# Turnstone # Turnstone
[![CI](https://github.com/turnstonelabs/turnstone/actions/workflows/ci.yml/badge.svg)](https://github.com/turnstonelabs/turnstone/actions/workflows/ci.yml)
[![PyPI](https://img.shields.io/pypi/v/turnstone)](https://pypi.org/project/turnstone/)
[![Python](https://img.shields.io/pypi/pyversions/turnstone)](https://pypi.org/project/turnstone/)
[![License](https://img.shields.io/badge/license-BSL--1.1-blue)](LICENSE)
Multi-node AI orchestration platform. Deploy tool-using AI agents across a cluster of servers, driven by message queues or interactive interfaces. Multi-node AI orchestration platform. Deploy tool-using AI agents across a cluster of servers, driven by message queues or interactive interfaces.
Named after the [Ruddy Turnstone](https://en.wikipedia.org/wiki/Ruddy_turnstone) — a bird that flips rocks to expose what's hiding underneath. Named after the [Ruddy Turnstone](https://en.wikipedia.org/wiki/Ruddy_turnstone) — a bird that flips rocks to expose what's hiding underneath.
@@ -141,9 +146,30 @@ docs/
├── docker.md # Docker Compose deployment and configuration ├── docker.md # Docker Compose deployment and configuration
├── simulator.md # Cluster simulator usage and scenarios ├── simulator.md # Cluster simulator usage and scenarios
├── tools.md # Tool schemas, execution pipeline, approval flow ├── tools.md # Tool schemas, execution pipeline, approval flow
── eval.md # Evaluation harness internals ── eval.md # Evaluation harness internals
└── diagrams/ # UML architecture diagrams (PlantUML sources + PNGs)
└── png/ # Pre-rendered diagram images
``` ```
### Architecture Diagrams
Detailed UML diagrams are available in [`docs/diagrams/`](docs/diagrams/):
| Diagram | Description |
|---------|-------------|
| [System Context](docs/diagrams/png/01-system-context.png) | Top-level components and external dependencies |
| [Package Structure](docs/diagrams/png/02-package-structure.png) | Python modules and dependency graph |
| [Core Engine Classes](docs/diagrams/png/03-core-engine-classes.png) | SessionUI protocol, ChatSession, WorkstreamManager |
| [Conversation Turn](docs/diagrams/png/04-conversation-turn.png) | Full message lifecycle through the engine |
| [Tool Pipeline](docs/diagrams/png/05-tool-pipeline.png) | Three-phase prepare/approve/execute |
| [MQ Protocol](docs/diagrams/png/06-mq-protocol.png) | 9 inbound + 19 outbound message types |
| [Message Routing](docs/diagrams/png/07-message-routing.png) | Multi-node routing scenarios |
| [Redis Key Schema](docs/diagrams/png/08-redis-key-schema.png) | All Redis keys, types, and TTLs |
| [Workstream States](docs/diagrams/png/09-workstream-states.png) | State machine transitions |
| [Simulator](docs/diagrams/png/10-simulator-architecture.png) | SimCluster, dispatchers, scenarios |
| [Console Data Flow](docs/diagrams/png/11-console-data-flow.png) | Dashboard data collection threads |
| [Deployment](docs/diagrams/png/12-deployment.png) | Docker Compose service topology |
## Multi-node routing ## Multi-node routing
Each Turnstone server runs a bridge process. Bridges share a Redis instance for coordination: Each Turnstone server runs a bridge process. Bridges share a Redis instance for coordination:
+1 -1
View File
@@ -23,7 +23,7 @@ services:
# Redis — message broker, pub/sub, node registry # Redis — message broker, pub/sub, node registry
# ------------------------------------------------------------------- # -------------------------------------------------------------------
redis: redis:
image: redis:7-alpine image: redis:7.4-alpine
command: command:
- sh - sh
- -c - -c
+2
View File
@@ -2,6 +2,8 @@
## Overview ## Overview
> See also: [MQ Protocol diagram](diagrams/png/06-mq-protocol.png) | [Message Routing diagram](diagrams/png/07-message-routing.png) | [Redis Key Schema diagram](diagrams/png/08-redis-key-schema.png)
`turnstone-server` exposes a browser-based chat UI backed by a Python stdlib HTTP `turnstone-server` exposes a browser-based chat UI backed by a Python stdlib HTTP
server (`socketserver.ThreadingMixIn` + `http.server.HTTPServer`). The server server (`socketserver.ThreadingMixIn` + `http.server.HTTPServer`). The server
uses **Server-Sent Events (SSE)** for real-time streaming and **HTTP POST** for uses **Server-Sent Events (SSE)** for real-time streaming and **HTTP POST** for
+10 -2
View File
@@ -65,6 +65,8 @@ turnstone/
## Core Loop ## Core Loop
> See also: [Conversation Turn diagram](diagrams/png/04-conversation-turn.png)
A user message flows through the system as follows: A user message flows through the system as follows:
``` ```
@@ -110,6 +112,8 @@ A user message flows through the system as follows:
### Tool Execution Pipeline ### Tool Execution Pipeline
> See also: [Tool Pipeline diagram](diagrams/png/05-tool-pipeline.png)
Tool execution is a three-phase process: Tool execution is a three-phase process:
``` ```
@@ -168,6 +172,8 @@ The engine emits state changes via `_emit_state()` which calls
## SessionUI Protocol ## SessionUI Protocol
> See also: [Core Engine Classes diagram](diagrams/png/03-core-engine-classes.png)
Defined in `turnstone.core.session.SessionUI` as a `typing.Protocol` with 13 Defined in `turnstone.core.session.SessionUI` as a `typing.Protocol` with 13
methods. Every frontend must implement all of them. methods. Every frontend must implement all of them.
@@ -225,6 +231,8 @@ Workstreams are parallel, independent chat sessions. Each has its own
### WorkstreamState ### WorkstreamState
> See also: [Workstream States diagram](diagrams/png/09-workstream-states.png)
Defined in `turnstone.core.workstream.WorkstreamState` (5 states): Defined in `turnstone.core.workstream.WorkstreamState` (5 states):
``` ```
@@ -418,8 +426,8 @@ independently, then returns the final content as the tool result.
to `.plan-<session_id>.md` — unique per `ChatSession` so concurrent workstreams to `.plan-<session_id>.md` — unique per `ChatSession` so concurrent workstreams
don't collide. On repeat invocations the prior `plan` tool call and its result don't collide. On repeat invocations the prior `plan` tool call and its result
are forwarded from `self.messages` so the agent refines the existing plan rather are forwarded from `self.messages` so the agent refines the existing plan rather
than starting over. Planning instructions are passed via `model_identity` in than starting over. Planning instructions are injected as a developer message
`chat_template_kwargs` rather than as a developer message. prepended to the agent's conversation.
- **Turn limit**: controlled by `agent_max_turns` (default: `-1`, unlimited). - **Turn limit**: controlled by `agent_max_turns` (default: `-1`, unlimited).
When a limit is set and reached, the agent is forced to synthesize a final When a limit is set and reached, the agent is forced to synthesize a final
response without tools. When unlimited, the loop only exits when the model response without tools. When unlimited, the loop only exits when the model
+2
View File
@@ -6,6 +6,8 @@ The console is read-only — it observes but does not own workstreams or drive L
## Architecture ## Architecture
> See also: [Console Data Flow diagram](diagrams/png/11-console-data-flow.png)
``` ```
turnstone-server ──→ turnstone-bridge ──→ Redis ──→ turnstone-console ──→ Browser turnstone-server ──→ turnstone-bridge ──→ Redis ──→ turnstone-console ──→ Browser
(per node) (per node) (shared) (one instance) (per node) (per node) (shared) (one instance)
+65
View File
@@ -0,0 +1,65 @@
@startuml
!theme plain
title Turnstone — System Context
' Actors/Users
actor "CLI User" as cli_user
actor "Browser User" as browser_user
actor "External Client\n(Python / CI)" as ext_client
actor "Eval Harness" as eval_user
' External Systems
cloud "LLM Provider\n(OpenAI-compatible API)" as llm
database "Redis" as redis
database "SQLite\n(.turnstone.db)" as sqlite
' Turnstone System Boundary
package "Turnstone Platform" {
component [turnstone\n(CLI)] as cli <<entry point>>
component [turnstone-server\n(HTTP + SSE)] as server <<entry point>>
component [turnstone-bridge\n(Queue ↔ HTTP)] as bridge <<service>>
component [turnstone-console\n(Dashboard)] as console <<service>>
component [turnstone-eval\n(Headless)] as eval <<entry point>>
component [turnstone-sim\n(Simulator)] as sim <<service>>
}
' User connections
cli_user --> cli : stdin / stdout
browser_user --> server : HTTP + SSE\n(port 8080)
browser_user --> console : HTTP + SSE\n(port 8090)
ext_client --> redis : Redis LIST\n(push commands)
eval_user --> eval : Python API
' Internal connections
cli --> llm : OpenAI Streaming API\n(HTTPS)
cli --> sqlite : SQLite
server --> llm : OpenAI Streaming API\n(HTTPS)
server --> sqlite : SQLite
eval --> llm : OpenAI API\n(non-streaming)
eval --> sqlite : SQLite
bridge --> server : HTTP REST\n(POST /api/send, etc.)
bridge <-- server : SSE\n(GET /api/events)
bridge --> redis : Redis LIST + PUBSUB\n+ STRING (routing, heartbeats)
console --> redis : Redis PUBSUB + STRING\n(cluster channel, heartbeats)
console --> server : HTTP polling\n(GET /api/dashboard)
sim --> redis : Redis LIST + PUBSUB\n+ STRING (heartbeats)
' Notes
note right of sim
Simulator replaces Server+Bridge
with lightweight SimNodes that
publish to the same Redis channels.
end note
note right of redis
Shared message broker:
- LIST: command queues
- STRING: heartbeats, routing
- PUBSUB: event broadcast
end note
@enduml
+140
View File
@@ -0,0 +1,140 @@
@startuml
!theme plain
title Turnstone — Package & Module Structure
' Color definitions
skinparam component {
BackgroundColor<<entry>> #B8D4E3
BackgroundColor<<core>> #C8E6C9
BackgroundColor<<mq>> #FFE0B2
BackgroundColor<<sim>> #E1BEE7
BackgroundColor<<console>> #B2EBF2
BackgroundColor<<ui>> #F0F4C3
BackgroundColor<<artifact>> #ECEFF1
}
' Entry points
package "Entry Points" <<Rectangle>> {
component [cli.py\nturnstone] as cli <<entry>>
component [server.py\nturnstone-server] as server <<entry>>
component [eval.py\nturnstone-eval] as eval <<entry>>
component [chat.py\n(re-exports)] as chat <<entry>>
}
' Core engine
package "turnstone/core/" <<Rectangle>> {
component [session.py\nChatSession, SessionUI] as session <<core>>
component [workstream.py\nWorkstreamManager] as workstream <<core>>
component [tools.py\nTool loader] as tools <<core>>
component [memory.py\nSQLite + FTS5] as memory <<core>>
component [metrics.py\nPrometheus metrics] as metrics <<core>>
component [config.py\nTOML config] as config <<core>>
component [safety.py\nPath validation] as safety <<core>>
component [sandbox.py\nCommand sandbox] as sandbox <<core>>
component [edit.py\nFile editing] as edit <<core>>
component [web.py\nWeb helpers] as web <<core>>
component [auth.py\nAuthentication] as auth <<core>>
}
' MQ subsystem
package "turnstone/mq/" <<Rectangle>> {
component [protocol.py\n28 message types] as protocol <<mq>>
component [broker.py\nMessageBroker, RedisBroker] as broker <<mq>>
component [bridge.py\nturnstone-bridge] as bridge <<mq>>
component [client.py\nTurnstoneClient] as client <<mq>>
}
' Simulator
package "turnstone/sim/" <<Rectangle>> {
component [cluster.py\nSimCluster] as simcluster <<sim>>
component [node.py\nSimNode, SimWorkstream] as simnode <<sim>>
component [engine.py\nSimEngine] as simengine <<sim>>
component [scenario.py\n5 scenarios] as scenario <<sim>>
component [sim/config.py\nSimConfig] as simconfig <<sim>>
component [sim/metrics.py\nSim metrics] as simmetrics <<sim>>
component [sim/cli.py\nturnstone-sim] as simcli <<sim>>
}
' Console
package "turnstone/console/" <<Rectangle>> {
component [collector.py\nClusterCollector] as collector <<console>>
component [console/server.py\nDashboard HTTP+SSE] as consoleserver <<console>>
}
' UI
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>>
}
' Tool schemas
package "turnstone/tools/" <<Rectangle>> {
component [*.json\n14 tool schemas] as schemas <<artifact>>
}
' Entry point dependencies
cli --> session
cli --> workstream
cli --> config
cli --> memory
cli --> colors
cli --> markdown
cli --> spinner
cli --> tools
server --> session
server --> workstream
server --> config
server --> memory
server --> metrics
server --> auth
server --> tools
eval --> session
eval --> memory
eval --> config
eval --> tools
chat --> session
' Core internal deps
session --> tools
session --> memory
session --> safety
session --> sandbox
session --> edit
session --> web
tools --> schemas
' MQ dependencies
bridge --> protocol
bridge --> broker
bridge --> config
client --> protocol
client --> broker
' Sim dependencies
simcli --> simcluster
simcli --> simconfig
simcli --> scenario
simcluster --> simnode
simcluster --> broker
simcluster --> simmetrics
simcluster --> simconfig
simnode --> simengine
simnode --> protocol
simnode --> simconfig
simnode --> simmetrics
scenario --> broker
scenario --> protocol
scenario --> simconfig
scenario --> simmetrics
' Console dependencies
consoleserver --> collector
consoleserver --> config
consoleserver --> auth
collector --> broker
@enduml
+167
View File
@@ -0,0 +1,167 @@
@startuml
!theme plain
title Turnstone — Core Engine Classes
skinparam classAttributeIconSize 0
' SessionUI Protocol
interface "SessionUI" as SessionUI <<Protocol>> {
+ on_thinking_start()
+ on_thinking_stop()
+ on_reasoning_token(text: str)
+ on_content_token(text: str)
+ on_stream_end()
+ approve_tools(items: list) → (bool, str|None)
+ on_tool_result(name: str, output: str)
+ on_status(usage: dict, ctx_window: int, effort: str)
+ on_plan_review(content: str) → str
+ on_info(message: str)
+ on_error(message: str)
+ on_state_change(state: str)
+ on_rename(name: str)
}
' Implementations
class "TerminalUI" as TerminalUI {
Writes to stdout with ANSI colors
Prompts for approval via input()
--
cli.py
}
class "WorkstreamTerminalUI" as WsTermUI {
- _output_buffer: list[tuple]
- ws_id: str
- manager: WorkstreamManager
+ flush_buffer()
--
Buffers output when workstream
is not foregrounded
}
class "WebUI" as WebUI {
- _event_queue: Queue
- _approval_event: Event
- _plan_event: Event
- _ws_prompt_tokens: int
- _ws_tool_calls: dict
+ resolve_approval(approved, feedback)
+ resolve_plan(feedback)
--
Enqueues JSON events for SSE.
Blocks on threading.Event for
approval/plan review.
--
server.py
}
class "NullUI" as NullUI {
approve_tools() → (True, None)
All other methods: no-op
--
eval.py
}
' ChatSession
class "ChatSession" as ChatSession {
- client: OpenAI
- model: str
- ui: SessionUI
- messages: list[dict]
- _msg_tokens: list[int]
- _session_id: str
- _read_files: set[str]
- system_messages: list[dict]
--
+ send(user_input: str)
+ handle_command(command: str)
+ resume_session(session_id: str)
- _stream_response(stream) → dict
- _create_stream_with_retry(msgs) → Stream
- _execute_tools(tool_calls) → (results, feedback)
- _prepare_tool(tc) → item dict
- _run_agent(messages, tools, ...) → str
- _compact_messages(auto: bool)
- _full_messages() → list[dict]
- _update_token_table(msg)
- _emit_state(state: str)
- _generate_title()
}
' HeadlessSession
class "HeadlessSession" as HeadlessSession {
+ tool_call_log: list[dict]
+ auto_approve: bool = True
+ send_headless(input, max_turns, ...)
- _override_system_prompt(content)
--
eval.py: non-streaming,
records all tool calls
}
' WorkstreamManager
class "WorkstreamManager" as WsMgr {
- _session_factory: Callable[[SessionUI], ChatSession]
- _workstreams: dict[str, Workstream]
- _order: list[str]
- _active_id: str
- _on_state_change: Callable
--
+ create(name, ui_factory) → Workstream
+ close(ws_id)
+ get(ws_id) → Workstream
+ get_active() → Workstream
+ switch(ws_id)
+ set_state(ws_id, state)
+ close_idle(max_age_seconds)
--
MAX_WORKSTREAMS = 10
}
' Workstream
class "Workstream" as Ws <<dataclass>> {
+ id: str
+ name: str
+ state: WorkstreamState
+ session: ChatSession
+ ui: SessionUI
+ worker_thread: Thread
+ error_message: str
+ last_active: float
- _lock: Lock
}
' WorkstreamState
enum "WorkstreamState" as WsState {
IDLE
THINKING
RUNNING
ATTENTION
ERROR
}
' Relationships
SessionUI <|.. TerminalUI
TerminalUI <|-- WsTermUI
SessionUI <|.. WebUI
SessionUI <|.. NullUI
ChatSession --> SessionUI : uses
ChatSession <|-- HeadlessSession
WsMgr --> "*" Ws : manages
Ws --> "1" ChatSession : wraps
Ws --> "1" SessionUI : wraps
Ws --> "1" WsState : has
WsMgr ..> ChatSession : creates via\nsession_factory(ui)
note bottom of ChatSession
Central engine: multi-turn LLM loop
with tool dispatch, agent sub-sessions,
context compaction, and memory persistence.
core/session.py (~2700 lines)
end note
@enduml
+147
View File
@@ -0,0 +1,147 @@
@startuml
!theme plain
title Turnstone — Conversation Turn Lifecycle
skinparam sequenceArrowThickness 1.5
skinparam sequenceLifeLineBackgroundColor #F5F5F5
participant "User /\nHTTP Client" as User
participant "ChatSession" as CS
participant "SessionUI" as UI
participant "OpenAI API\n(LLM)" as LLM
participant "Tool Executor\n(ThreadPool)" as TP
database "SQLite" as DB
== User Input ==
User -> CS : send(user_input)
activate CS
CS -> CS : messages.append({role: "user", content: input})
CS -> DB : save_message(session_id, "user", input)
== LLM Call Loop ==
group loop [while tool_calls present]
CS -> UI : on_state_change("thinking")
CS -> UI : on_thinking_start()
CS -> LLM : client.chat.completions.create(\n model, messages, tools,\n stream=True, stream_options={include_usage})
activate LLM
note right of CS
Retry up to 3× on transient errors:
RateLimitError, APITimeoutError,
APIConnectionError, InternalServerError,
ServiceUnavailableError, APIError
Backoff: 1s, 2s, 4s
end note
== Streaming Response ==
loop for each chunk in stream
LLM --> CS : delta
note right of CS
on_thinking_stop() called on first
delta token via _stop_spinner_once()
end note
alt reasoning_content present
CS -> UI : on_reasoning_token(text)
else content present
CS -> UI : on_content_token(text)
else tool_call delta
CS -> CS : accumulate in tool_calls_acc
end
end
LLM --> CS : stream complete (usage stats)
deactivate LLM
CS -> UI : on_thinking_stop() (no-op guard: already called by _stop_spinner_once)
CS -> UI : on_stream_end()
CS -> CS : _update_token_table()\ncalibrate chars_per_token ratio
CS -> CS : messages.append(assistant_msg)
CS -> DB : save_message(session_id, "assistant", content)
CS -> DB : save_message(session_id, "tool_call", ...) ×N
== Tool Dispatch (if tool_calls) ==
alt no tool_calls
CS -> UI : on_status(usage, context_window, effort)
opt prompt_tokens > context_window × auto_compact_pct
CS -> CS : _compact_messages(auto=True)
CS -> LLM : Non-streaming summarization call
CS -> CS : Replace messages with [summary]
end
opt first exchange & no title
CS -> CS : Background thread: _generate_title()
end
CS -> UI : on_state_change("idle")
CS --> User : return
else has tool_calls
CS -> UI : on_state_change("running")
== Phase 1: Prepare ==
CS -> CS : [_prepare_tool(tc) for tc in tool_calls]\nParse JSON args, validate,\nbuild preview + header
== Phase 2: Approve ==
CS -> UI : on_state_change("attention")
CS -> UI : approve_tools(items)
activate UI
note right of UI
TerminalUI: input() prompt
WebUI: _approval_event.wait()
NullUI: returns (True, None)
end note
UI --> CS : (approved: bool, feedback: str?)
deactivate UI
CS -> UI : on_state_change("running")
== Phase 3: Execute ==
CS -> TP : ThreadPoolExecutor(max_workers=4)\nrun_one(item) for each tool
activate TP
note right of TP
Parallel execution:
bash → subprocess.run()
read_file → open().read()
search → grep subprocess
edit_file → string replace
task/plan → _run_agent() sub-loop
math → sandboxed subprocess
web_fetch → httpx + LLM summarize
web_search → Tavily API
remember/recall/forget → SQLite
end note
note right of TP
on_tool_result() called
inside each _exec_* handler
end note
TP --> CS : [(call_id, output), ...]
deactivate TP
loop for each result
CS -> CS : messages.append({role: "tool", ...})
CS -> DB : save_message(session_id, "tool_result", ...)
end
opt user_feedback from approval
CS -> CS : messages.append({role: "user", content: feedback})
end
note right of CS : Loop back for next LLM call
end
end
deactivate CS
@enduml
+130
View File
@@ -0,0 +1,130 @@
@startuml
!theme plain
title Turnstone — Tool Execution Pipeline (Three Phases)
start
partition "Phase 1: Prepare" #E8F5E9 {
:Receive tool_calls list from LLM response;
while (more tool_calls?) is (yes)
:Extract call_id, func_name, raw_args;
if (json.loads(raw_args) succeeds?) then (yes)
:parsed_args = JSON dict;
else (no)
:Fallback 1: regex extraction;
if (regex found keys?) then (yes)
:parsed_args = extracted dict;
else (no)
:Fallback 2: bare string →\nPRIMARY_KEY_MAP[func_name];
endif
endif
:Dispatch to _prepare_{func_name}();
note right
**Dispatch table (14 tools):**
┌─────────────┬──────────────────┐
│ Tool │ Needs Approval? │
├─────────────┼──────────────────┤
│ bash │ ✓ Yes │
│ read_file │ ✗ Auto-approve │
│ write_file │ ✓ Yes │
│ edit_file │ ✓ Yes │
│ search │ ✗ Auto-approve │
│ math │ ✓ Yes │
│ man │ ✗ Auto-approve │
│ web_fetch │ ✓ Yes │
│ web_search │ ✓ Yes │
│ task │ ✓ Yes │
│ plan │ ✓ Yes │
│ remember │ ✗ Auto-approve │
│ recall │ ✗ Auto-approve │
│ forget │ ✗ Auto-approve │
└─────────────┴──────────────────┘
end note
:Build item dict:
{call_id, func_name, header,
preview, needs_approval,
approval_label, execute: Callable};
endwhile (no)
}
partition "Phase 2: Approve" #FFF3E0 {
if (any items need approval?) then (yes)
:_emit_state("attention");
:ui.approve_tools(items);
note right
**auto_approve check is handled
internally by ui.approve_tools()**
**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)
**WebUI**: Enqueue approve_request,
block on _approval_event.wait()
**NullUI**: Return (True, None)
end note
if (user approved?) then (yes)
:_emit_state("running");
else (denied)
:Mark all pending items as denied;
:denial_msg = "Denied by user";
:_emit_state("running");
endif
else (all auto-approved)
:ui enqueues tool_info event\n(no blocking);
endif
}
partition "Phase 3: Execute" #E3F2FD {
if (single tool call?) then (yes)
:Execute sequentially:\nrun_one(items[0]);
else (multiple)
:Execute in parallel:\nThreadPoolExecutor(max_workers=4)\npool.map(run_one, items);
endif
note right
**run_one(item):**
if item.error → return error string
if item.denied → return denial message
else → item["execute"](item)
├─ _exec_bash: subprocess.run(["bash", script.sh])
├─ _exec_read_file: open().readlines()
├─ _exec_write_file: makedirs + write
├─ _exec_edit_file: find_occurrences + replace
├─ _exec_search: grep subprocess
├─ _exec_math: sandboxed subprocess
├─ _exec_man: man/info subprocess
├─ _exec_web_fetch: httpx.get + LLM summary
├─ _exec_web_search: Tavily API POST
├─ _exec_task: _run_agent(TASK_AGENT_TOOLS)
├─ _exec_plan: _run_agent(AGENT_TOOLS, read-only)
├─ _exec_remember: SQLite INSERT OR REPLACE
├─ _exec_recall: SQLite FTS5/LIKE search
└─ _exec_forget: SQLite DELETE
end note
:Collect results: [(call_id, output), ...];
:_truncate_output() on each result\n(max context_window × chars_per_token × 0.5 chars\ndefault: ~context_window × 2 chars);
:ui.on_tool_result(name, output) for each;
if (plan tool was executed?) then (yes)
:ui.on_plan_review(output);
:Block for user review/feedback;
endif
}
:Return (results, user_feedback);
stop
@enduml
+242
View File
@@ -0,0 +1,242 @@
@startuml
!theme plain
title Turnstone — Message Queue Protocol Types
skinparam classAttributeIconSize 0
skinparam packageStyle rectangle
package "Inbound Messages (Client → Bridge)" #FFF3E0 {
abstract class "InboundMessage" as IM {
+ type: str
+ correlation_id: str {auto: uuid4().hex[:12]}
+ timestamp: float {auto: time.time()}
--
+ to_json() → str
+ {static} from_json(raw) → InboundMessage
}
class SendMessage {
type = "send"
--
+ ws_id: str
+ message: str
+ auto_approve: bool = False
+ auto_approve_tools: list[str] = []
+ name: str = ""
+ target_node: str = ""
}
class ApproveMessage {
type = "approve"
--
+ ws_id: str
+ request_id: str
+ approved: bool = True
+ feedback: str | None
+ always: bool = False
}
class PlanFeedbackMessage {
type = "plan_feedback"
--
+ ws_id: str
+ request_id: str
+ feedback: str
}
class CommandMessage {
type = "command"
--
+ ws_id: str
+ command: str
}
class CreateWorkstreamMessage {
type = "create_workstream"
--
+ name: str = ""
+ auto_approve: bool = False
+ auto_approve_tools: list[str] = []
+ target_node: str = ""
}
class CloseWorkstreamMessage {
type = "close_workstream"
--
+ ws_id: str
}
class ListWorkstreamsMessage {
type = "list_workstreams"
}
class HealthMessage {
type = "health"
}
class ListNodesMessage {
type = "list_nodes"
}
IM <|-- SendMessage
IM <|-- ApproveMessage
IM <|-- PlanFeedbackMessage
IM <|-- CommandMessage
IM <|-- CreateWorkstreamMessage
IM <|-- CloseWorkstreamMessage
IM <|-- ListWorkstreamsMessage
IM <|-- HealthMessage
IM <|-- ListNodesMessage
}
package "Outbound Events (Bridge → Client)" #E3F2FD {
abstract class "OutboundEvent" as OE {
+ type: str
+ ws_id: str
+ correlation_id: str
+ timestamp: float
--
+ to_json() → str
+ {static} from_json(raw) → OutboundEvent
}
package "Streaming" #BBDEFB {
class ContentEvent {
type = "content"
+ text: str
}
class ReasoningEvent {
type = "reasoning"
+ text: str
}
class StreamEndEvent {
type = "stream_end"
}
}
package "Tools" #C8E6C9 {
class ToolInfoEvent {
type = "tool_info"
+ items: list
}
class ApprovalRequestEvent {
type = "approval_request"
+ items: list
..
correlation_id = request_id
}
class ToolResultEvent {
type = "tool_result"
+ name: str
+ output: str
}
class PlanReviewEvent {
type = "plan_review"
+ content: str
}
}
package "Status" #FFF9C4 {
class AckEvent {
type = "ack"
+ status: str
+ detail: str
}
class StatusEvent {
type = "status"
+ prompt_tokens: int
+ completion_tokens: int
+ total_tokens: int
+ context_window: int
+ pct: float
+ effort: str
}
class StateChangeEvent {
type = "state_change"
+ state: str
}
class TurnCompleteEvent {
type = "turn_complete"
}
}
package "Lifecycle" #F8BBD0 {
class WorkstreamCreatedEvent {
type = "ws_created"
+ name: str
}
class WorkstreamClosedEvent {
type = "ws_closed"
}
class WorkstreamListEvent {
type = "ws_list"
+ workstreams: list
}
class WorkstreamRenameEvent {
type = "ws_rename"
+ name: str
}
}
package "System" #E0E0E0 {
class HealthResponseEvent {
type = "health_response"
+ data: dict
}
class ErrorEvent {
type = "error"
+ message: str
}
class InfoEvent {
type = "info"
+ message: str
}
class NodeListEvent {
type = "node_list"
+ nodes: list
}
class ClusterStateEvent {
type = "cluster_state"
+ state: str
+ node_id: str
+ tokens: int
+ context_ratio: float
+ activity: str
+ activity_state: str
}
}
OE <|-- ContentEvent
OE <|-- ReasoningEvent
OE <|-- StreamEndEvent
OE <|-- ToolInfoEvent
OE <|-- ApprovalRequestEvent
OE <|-- ToolResultEvent
OE <|-- PlanReviewEvent
OE <|-- AckEvent
OE <|-- StatusEvent
OE <|-- StateChangeEvent
OE <|-- TurnCompleteEvent
OE <|-- WorkstreamCreatedEvent
OE <|-- WorkstreamClosedEvent
OE <|-- WorkstreamListEvent
OE <|-- WorkstreamRenameEvent
OE <|-- HealthResponseEvent
OE <|-- ErrorEvent
OE <|-- InfoEvent
OE <|-- NodeListEvent
OE <|-- ClusterStateEvent
}
note bottom of IM
**Deserialization**: Strict type-dispatch via _INBOUND_REGISTRY.
Unknown type raises ValueError.
end note
note bottom of OE
**Deserialization**: Lenient type-dispatch via _OUTBOUND_REGISTRY.
Unknown type falls back to base OutboundEvent.
end note
@enduml
+105
View File
@@ -0,0 +1,105 @@
@startuml
!theme plain
title Turnstone — Multi-Node Message Routing
skinparam sequenceArrowThickness 1.5
participant "TurnstoneClient" as Client
collections "Redis" as Redis
participant "Bridge-A\n(node_id: nodeA)" as BridgeA
participant "Bridge-B\n(node_id: nodeB)" as BridgeB
participant "Server-A" as ServerA
== Scenario A: New Message — No Workstream Affinity ==
Client -> Redis : RPUSH turnstone:inbound\n{type:"send", message:"...", ws_id:""}
note right of Redis : Shared queue — any bridge can pick up
BridgeA -> Redis : BLPOP [turnstone:inbound:nodeA,\n turnstone:inbound]
Redis --> BridgeA : SendMessage (from shared queue)
BridgeA -> ServerA : POST /api/workstreams/new\n{name:"", auto_approve:false}
ServerA --> BridgeA : {ws_id:"abc12345", name:"ws-abc1"}
BridgeA -> Redis : SET turnstone:ws:abc12345 "nodeA"
note right : Register workstream ownership
BridgeA -> ServerA : GET /api/events?ws_id=abc12345
note right : Start per-WS SSE thread
BridgeA -> Redis : PUBLISH turnstone:events:global\nWorkstreamCreatedEvent
BridgeA -> Redis : PUBLISH turnstone:events:cluster\nClusterStateEvent(ws_id, state:"idle", node_id:"nodeA")
BridgeA -> ServerA : POST /api/send\n{message:"...", ws_id:"abc12345"}
ServerA --> BridgeA : {status:"ok"}
BridgeA -> Redis : PUBLISH turnstone:events:abc12345\nAckEvent(status:"ok")
... SSE events flow: content, 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
== Scenario B: Directed Message to Specific Node ==
Client -> Redis : RPUSH turnstone:inbound:nodeB\n{type:"send", target_node:"nodeB", ...}
note right : Per-node queue — only nodeB picks up
BridgeB -> Redis : BLPOP [turnstone:inbound:nodeB,\n turnstone:inbound]
Redis --> BridgeB : SendMessage (from per-node queue, priority)
note right of BridgeB : Process locally on nodeB
== Scenario C: Re-routing (Lands on Wrong Node) ==
Client -> Redis : RPUSH turnstone:inbound\n{type:"send", ws_id:"abc12345"}
BridgeB -> Redis : BLPOP [..., turnstone:inbound]
Redis --> BridgeB : SendMessage (ws_id: abc12345)
BridgeB -> Redis : GET turnstone:ws:abc12345
Redis --> BridgeB : "nodeA"
note right of BridgeB : Owner is nodeA, not me — re-route
BridgeB -> Redis : RPUSH turnstone:inbound:nodeA\n(re-routed message)
BridgeA -> Redis : BLPOP [turnstone:inbound:nodeA, ...]
Redis --> BridgeA : SendMessage (from per-node queue)
note right of BridgeA : Process locally — I own this workstream
== Scenario D: Approval via Response Queue ==
BridgeA <- ServerA : SSE: {type:"approve_request", items:[...]}
note right of BridgeA
Bridge checks auto-approve:
1. _ws_auto_approve[ws_id]? → auto
2. All tools in safe set? → auto
(read_file, search, man,
remember, recall, forget)
3. Otherwise → manual approval
end note
BridgeA -> Redis : PUBLISH turnstone:events:abc12345\nApprovalRequestEvent(correlation_id: req_xyz)
Client <- Redis : (subscribed) ApprovalRequestEvent
Client -> Redis : RPUSH turnstone:resp:req_xyz\nApproveMessage(approved:true)
note right : Response queue — bypasses inbound queue
BridgeA -> Redis : BLPOP turnstone:resp:req_xyz\n(spawned approval thread, timeout 300s)
Redis --> BridgeA : ApproveMessage
BridgeA -> ServerA : POST /api/approve\n{approved:true, ws_id:"abc12345"}
== Heartbeat (continuous) ==
BridgeA -> Redis : SET turnstone:node:nodeA\n{server_url, started} EX 60
note right : Every 30s — TTL 60s
BridgeB -> Redis : SET turnstone:node:nodeB\n{server_url, started} EX 60
@enduml
+98
View File
@@ -0,0 +1,98 @@
@startuml
!theme plain
title Turnstone — Redis Key Schema
skinparam component {
BackgroundColor<<LIST>> #BBDEFB
BackgroundColor<<STRING>> #C8E6C9
BackgroundColor<<PUBSUB>> #FFE0B2
}
skinparam note {
BackgroundColor #FAFAFA
}
package "Queues (Redis LIST)" #E3F2FD {
component [**turnstone:inbound**\n\nShared command queue.\nAny bridge can consume.\n\nOps: RPUSH (write), BLPOP (read)] as inbound <<LIST>>
component [**turnstone:inbound:{node_id}**\n\nPer-node directed queue.\nPriority over shared queue.\n\nOps: RPUSH (write), BLPOP (read)] as inbound_node <<LIST>>
component [**turnstone:resp:{request_id}**\n\nPer-request response queue.\nFor approval / plan feedback.\nTTL: 600s\n\nOps: RPUSH + EXPIRE (write), BLPOP (read)] as resp <<LIST>>
}
package "Routing (Redis STRING)" #E8F5E9 {
component [**turnstone:ws:{ws_id}**\n\nWorkstream → node ownership.\nValue: node_id string.\nNo TTL.\n\nOps: SET, GET, DEL] as ws_owner <<STRING>>
component [**turnstone:node:{node_id}**\n\nNode heartbeat + metadata.\nValue: JSON {server_url, started, ...}\nTTL: 60s (refreshed every 30s)\n\nOps: SET with EX, GET, SCAN] as node_hb <<STRING>>
}
package "Event Channels (Redis PUBSUB)" #FFF3E0 {
component [**turnstone:events:global**\n\nGlobal event broadcast.\nAll state changes, ws lifecycle.\n\nOps: PUBLISH, SUBSCRIBE] as evt_global <<PUBSUB>>
component [**turnstone:events:{ws_id}**\n\nPer-workstream events.\nContent, tools, status.\n\nOps: PUBLISH, SUBSCRIBE] as evt_ws <<PUBSUB>>
component [**turnstone:events:cluster**\n\nCluster-wide state changes.\nUsed by Console dashboard.\n\nOps: PUBLISH, SUBSCRIBE] as evt_cluster <<PUBSUB>>
}
' Readers / Writers
actor "TurnstoneClient" as client
actor "Bridge" as bridge
actor "SimNode" as sim
actor "Console\nCollector" as console
actor "Scenario\n(injector)" as scenario
' Queue interactions
client --> inbound : RPUSH\n(send commands)
client --> inbound_node : RPUSH\n(directed)
scenario --> inbound : RPUSH\n(inject load)
scenario --> inbound_node : RPUSH\n(directed scenario)
bridge --> inbound : BLPOP\n(consume)
bridge --> inbound_node : BLPOP\n(priority)
bridge --> inbound_node : RPUSH\n(re-route)
sim --> inbound_node : BLPOP\n(via dispatcher)
client --> resp : RPUSH\n(approval response)
bridge --> resp : BLPOP\n(wait for approval)
' Routing interactions
bridge --> ws_owner : SET / GET / DEL
client --> ws_owner : GET\n(route lookup)
sim --> ws_owner : SET / DEL
bridge --> node_hb : SET with EX\n(heartbeat)
sim --> node_hb : SET with EX\n(heartbeat)
console --> node_hb : SCAN + GET\n(discovery)
client --> node_hb : SCAN + GET\n(list_nodes)
' Pub/sub interactions
bridge --> evt_global : PUBLISH
bridge --> evt_ws : PUBLISH
bridge --> evt_cluster : PUBLISH
client --> evt_global : SUBSCRIBE
client --> evt_ws : SUBSCRIBE
sim --> evt_global : PUBLISH
sim --> evt_ws : PUBLISH
sim --> evt_cluster : PUBLISH
console --> evt_cluster : SUBSCRIBE
note bottom of inbound
**BLPOP priority**: Bridges call
BLPOP [per-node, shared] so the
per-node queue is always checked first.
end note
note bottom of resp
**Bypasses inbound queue**: Approval
responses go directly to the response
queue, not through the inbound queue.
Auto-cleaned after 600s TTL.
end note
note bottom of evt_cluster
**ClusterStateEvent** includes node_id,
tokens, and context_ratio — enriched
data not available on the global channel.
end note
@enduml
+81
View File
@@ -0,0 +1,81 @@
@startuml
!theme plain
title Turnstone — Workstream State Machine
skinparam state {
BackgroundColor<<idle>> #E8F5E9
BackgroundColor<<thinking>> #E3F2FD
BackgroundColor<<running>> #FFF3E0
BackgroundColor<<attention>> #FCE4EC
BackgroundColor<<error>> #FFCDD2
}
state "IDLE" as idle <<idle>> : Waiting for user input.\nNo active LLM call or tool execution.
state "THINKING" as thinking <<thinking>> : LLM streaming response.\nTokens flowing (reasoning + content).
state "RUNNING" as running <<running>> : Tools executing.\nThreadPoolExecutor active.
state "ATTENTION" as attention <<attention>> : Blocked on user action.\nTool approval or plan review needed.
state "ERROR" as error <<error>> : Exception occurred.\nRecoverable on next send().
[*] --> idle : Session created
idle --> thinking : send() called\n_emit_state("thinking")
thinking --> running : Tool calls detected\nin LLM response\n_emit_state("running")
thinking --> idle : No tool calls\n(final answer)\n_emit_state("idle")
thinking --> error : Exception during\nstreaming
running --> attention : approve_tools() called\n_emit_state("attention")
attention --> running : User approves\n(tools execute)\n_emit_state("running")
attention --> running : User denies\n(denial recorded)\n_emit_state("running")
running --> thinking : Tool results appended,\nnext LLM call\n_emit_state("thinking")
running --> attention : Plan tool complete,\non_plan_review()\n_emit_state("attention")
running --> error : Exception during\ntool execution
error --> thinking : New send() call\n_emit_state("thinking")
note right of thinking
**Emitted via:**
session._emit_state(state)
→ ui.on_state_change(state)
**Propagation:**
• WebUI → global SSE queue (ws_state)
• Bridge → PUBLISH to global + cluster channels
• CLI → WorkstreamManager.set_state()
end note
note left of attention
**Blocking mechanisms:**
• TerminalUI: input() prompt
• WebUI: threading.Event.wait()
• Bridge: BLPOP on response queue
• NullUI: auto-approve (never reaches)
end note
state "SimWorkstream (simplified)" as sim_group {
state "sim_idle" as si <<idle>>
state "sim_thinking" as st <<thinking>>
state "sim_running" as sr <<running>>
state "sim_error" as se <<error>>
[*] --> si
si --> st : process_turn() called
st --> sr : Tool calls generated
sr --> st : More rounds
st --> si : No tools / max rounds
st --> se : Uncaught exception
}
note right of sim_group
SimWorkstream has no ATTENTION state —
tool approval is not simulated.
end note
@enduml
@@ -0,0 +1,113 @@
@startuml
!theme plain
title Turnstone — Simulator Architecture
skinparam component {
BackgroundColor<<cluster>> #E1BEE7
BackgroundColor<<node>> #CE93D8
BackgroundColor<<engine>> #F3E5F5
BackgroundColor<<scenario>> #FFF3E0
BackgroundColor<<metrics>> #E8F5E9
BackgroundColor<<redis>> #FFCDD2
}
package "SimCluster" as cluster <<cluster>> {
component [**ThreadPoolExecutor**\nmax_workers=64\n(blocking Redis ops)] as executor <<cluster>>
component [**redis.ConnectionPool**\nmax_connections=64\ndecode_responses=True\n(shared across all nodes)] as pool <<redis>>
package "InboundDispatchers" {
component [**Dispatcher 0**\nnodes 0-49] as d0
component [**Dispatcher 1**\nnodes 50-99] as d1
component [**...**\n(ceil(N/50) total)] as dn
note bottom of d0
Each dispatcher calls BLPOP on a single Redis
connection for up to 50 node queues + shared queue.
Keys: [prefix:inbound:sim-0000, ..., prefix:inbound]
Per-node keys have BLPOP priority over shared.
end note
}
package "SimNodes (N instances)" {
component [**SimNode sim-0000**] as n0 <<node>>
component [**SimNode sim-0001**] as n1 <<node>>
component [**...**] as nn <<node>>
component [**SimEngine**\n(per node, seeded RNG)\n\nLLM simulation:\n gaussian(μ=2s, σ=0.5s) latency\n gaussian(μ=200, σ=50) tokens\n random word content\n P(tool_calls) = 0.6/0.3\n\nTool simulation:\n gaussian(μ=0.5s, σ=0.2s) latency\n P(failure) = 0.02] as engine <<engine>>
component [**SimWorkstream**\n(0..max_ws per node)\n\nState: idle→thinking→running→idle\nToken accounting: word_count × 3\nContent: 8-chunk streaming] as ws <<node>>
}
component [**MetricsCollector**\n(thread-safe, shared)\n\nTracks: turn latencies,\nthroughput, utilization,\nerrors, node kills] as metrics <<metrics>>
}
package "Scenarios (5 workload patterns)" <<scenario>> {
component [**SteadyState**\nConstant rate:\n1/mps interval\nfor duration secs] as steady <<scenario>>
component [**Burst**\nburst_size messages\nas fast as possible\nthen wait] as burst <<scenario>>
component [**NodeFailure**\nSteadyState + periodic\nnode kills (up to N/2)] as failure <<scenario>>
component [**Directed**\nMessages targeted to\nspecific nodes via\ntarget_node field] as directed <<scenario>>
component [**Lifecycle**\n3 phases:\n1. Create workstreams\n2. Send messages\n3. Close half] as lifecycle <<scenario>>
}
database "Redis" as redis <<redis>>
' Scenario -> Redis
steady --> redis : RPUSH prefix:inbound\n(SendMessage)
burst --> redis : RPUSH prefix:inbound\n(burst)
failure --> redis : RPUSH prefix:inbound
directed --> redis : RPUSH prefix:inbound:{node}\n(directed)
lifecycle --> redis : RPUSH prefix:inbound\n(Create/Send/Close)
' Dispatchers -> Redis -> Nodes
d0 --> redis : BLPOP [per-node..., shared]
d1 --> redis : BLPOP [per-node..., shared]
d0 --> n0 : handle_message(raw)
d0 --> n1 : handle_message(raw)
' Nodes internal
n0 --> engine : simulate_llm_response()\nsimulate_tool_execution()
n0 --> ws : process_turn()
' Nodes -> Redis (events)
n0 --> redis : PUBLISH prefix:events:global\n(StateChangeEvent)
n0 --> redis : PUBLISH prefix:events:{ws_id}\n(ContentEvent, ToolResultEvent, ...)
n0 --> redis : PUBLISH prefix:events:cluster\n(ClusterStateEvent)
n0 --> redis : SET prefix:node:sim-0000\nEX 60 (heartbeat)
n0 --> redis : SET prefix:ws:{ws_id}\n(ownership)
' Shared pool
n0 ..> pool : PooledBroker\n(shared connection)
n1 ..> pool : PooledBroker
d0 ..> pool
d0 ..> executor : asyncio.to_thread()
' Metrics
ws --> metrics : record_turn(ws_id, node_id, latency)
steady --> metrics : record_inject()
burst --> metrics : record_inject()
directed --> metrics : record_inject()
lifecycle --> metrics : record_inject()
cluster --> metrics : record_node_kill(node_id)
cluster --> metrics : snapshot_utilization()\n(every metrics_interval)
note bottom of cluster
**SimConfig** controls all simulation parameters:
num_nodes, max_ws_per_node, redis settings,
llm_latency_mean/stddev, tool_failure_rate,
scenario, duration, messages_per_second, seed
end note
note right of redis
Simulator uses **real Redis** —
not a mock. Console dashboard
can monitor a running simulation
via the same cluster channel.
end note
@enduml
+124
View File
@@ -0,0 +1,124 @@
@startuml
!theme plain
title Turnstone — Console Dashboard Data Collection
skinparam sequenceArrowThickness 1.5
participant "Browser" as Browser
participant "Console\nHTTP Server" as Server
participant "ClusterCollector" as CC
collections "Redis" as Redis
participant "Node-A\n(real server)" as NodeA
participant "Node-B\n(sim node)" as NodeB
== Thread 1: Cluster Event Subscriber (real-time) ==
CC -> Redis : SUBSCRIBE turnstone:events:cluster
activate CC #E1BEE7
Redis --> CC : ClusterStateEvent\n{ws_id, state:"thinking",\nnode_id:"nodeA", tokens:500,\ncontext_ratio:0.05}
CC -> CC : Update NodeSnapshot["nodeA"]\n.workstreams["ws123"].state = "thinking"
CC -> CC : _fanout(event) → all SSE listeners
Redis --> CC : {"type":"ws_created",\nws_id:"ws456", name:"task-1",\nnode_id:"sim-0003"}
CC -> CC : Add workstream to\nNodeSnapshot["sim-0003"]
CC -> CC : _fanout(event)
Redis --> CC : ClusterStateEvent\n{ws_id:"ws456", state:"idle"}
CC -> CC : Update workstream state
note right of CC
Handles: cluster_state,
ws_created, ws_closed, ws_rename
Thread runs continuously.
All updates are thread-safe
via threading.Lock.
end note
deactivate CC
== Thread 2: Node Discovery (every 15s) ==
CC -> Redis : SCAN 0 MATCH turnstone:node:*
activate CC #B2EBF2
Redis --> CC : [turnstone:node:nodeA, turnstone:node:sim-0003, ...]
loop for each discovered key
CC -> Redis : GET turnstone:node:{id}
Redis --> CC : JSON: {server_url, started, max_ws, sim:true/false}
end
CC -> CC : Create new NodeSnapshot\nfor newly discovered nodes
CC -> CC : Remove NodeSnapshot\nfor disappeared nodes
CC -> CC : _fanout({type: "node_joined", ...})\n_fanout({type: "node_lost", ...})
deactivate CC
== Thread 3: HTTP Polling (every 10s, real nodes only) ==
CC -> CC : Filter nodes where\nserver_url.startswith("http")
activate CC #C8E6C9
note right of CC
sim:// nodes are SKIPPED.
Their data comes exclusively
from the cluster event channel.
end note
CC -> NodeA : GET /api/dashboard
activate NodeA
NodeA --> CC : {workstreams: [...],\naggregate: {total_tokens, ...}}
deactivate NodeA
CC -> NodeA : GET /health
activate NodeA
NodeA --> CC : {status:"ok", model:"...",\nworkstreams:{total, idle, ...}}
deactivate NodeA
CC -> CC : Replace NodeSnapshot["nodeA"]\n.workstreams, .health, .aggregate
CC -x NodeB : (SKIPPED: sim:// URL)
deactivate CC
== Browser SSE Stream ==
Browser -> Server : GET /api/cluster/events
activate Server
Server -> CC : register_listener(queue)
note right : Per-client queue.Queue(maxsize=500)
loop continuous
CC -> Server : event via listener queue\n(from any of the 3 threads)
Server -> Browser : data: {"type":"cluster_state",...}\n\n
end
alt timeout (5s no events)
Server -> Browser : : keepalive\n\n
end
Browser -> Server : connection closed
Server -> CC : unregister_listener(queue)
deactivate Server
== Browser REST Requests ==
Browser -> Server : GET /api/cluster/overview
Server -> CC : get_overview()
CC --> Server : {nodes: 10, workstreams: 47,\nstates: {running:5, thinking:3, ...},\naggregate: {total_tokens: 50000}}
Server --> Browser : JSON response
Browser -> Server : GET /api/cluster/nodes?sort=activity
Server -> CC : get_nodes(sort_by="activity")
CC --> Server : {nodes: [...], total: 10}
Server --> Browser : JSON response
Browser -> Server : GET /api/cluster/workstreams\n?state=running&node=sim-0003
Server -> CC : get_workstreams(state="running",\nnode="sim-0003")
CC --> Server : {workstreams: [...], total: 5,\npage: 1, per_page: 50, pages: 1}
Server --> Browser : JSON response
@enduml
+111
View File
@@ -0,0 +1,111 @@
@startuml
!theme plain
title Turnstone — Docker Compose Deployment
skinparam node {
BackgroundColor #F5F5F5
}
cloud "LLM Provider" as llm {
component [vLLM / OpenAI API\nport 8000] as llm_api
}
node "Docker Host" as host {
frame "turnstone-net (bridge network)" as net {
node "redis" <<redis:7.4-alpine>> as redis_node {
component [Redis Server\nport 6379] as redis
note bottom of redis
Healthcheck: redis-cli ping
Volume: redis-data
end note
}
node "server" <<turnstone image>> as server_node {
component [turnstone-server\nport 8080] as server
note bottom of server
Command: turnstone-server
--host 0.0.0.0
--port 8080
Depends: redis (healthy)
Volume: turnstone-data
(/data)
end note
}
node "bridge ×N" <<turnstone image>> as bridge_node {
component [turnstone-bridge] as bridge
note bottom of bridge
Command: turnstone-bridge
--server-url http://server:8080
--redis-host redis
Depends: server + redis
Scalable: --scale bridge=N
node_id: auto from hostname
end note
}
node "console" <<turnstone image>> as console_node {
component [turnstone-console\nport 8090] as console
note bottom of console
Command: turnstone-console
--redis-host redis
--port 8090
Depends: redis
end note
}
node "sim (profile: sim)" <<turnstone image>> as sim_node {
component [turnstone-sim] as sim
note bottom of sim
Command: turnstone-sim
--redis-host redis
--nodes 100
--scenario steady
Depends: redis
Optional: only with
--profile sim
end note
}
}
}
actor "Browser\nUser" as browser
actor "MQ Client" as mqclient
' External connections
browser --> server : HTTP + SSE\nport 8080
browser --> console : HTTP + SSE\nport 8090
mqclient --> redis : Redis protocol\nport 6379
' Internal connections
server --> redis : Redis protocol\n(6379)
server --> llm_api : OpenAI API\n(HTTPS/HTTP)
bridge --> server : HTTP REST\n(POST /api/send, etc.)
bridge <-- server : SSE\n(GET /api/events)
bridge --> redis : Redis protocol\n(queues + pubsub)
console --> redis : Redis PUBSUB\n(cluster channel)
console --> server : HTTP polling\n(GET /api/dashboard)
sim --> redis : Redis protocol\n(queues + pubsub + keys)
' Environment variables
note right of host
**Environment Variables:**
• LLM_BASE_URL — LLM endpoint
• OPENAI_API_KEY — API key
• REDIS_PASSWORD — Redis auth
• TURNSTONE_AUTH_TOKEN — API auth
end note
' Volumes
database "redis-data" as rv
database "turnstone-data" as tv
redis_node --> rv
server_node --> tv
@enduml
+3
View File
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:341a8ab1483b1e0146878bd384a11d56bc78d29262de8262d06ef924317e2762
size 139969
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:d4b1db039e9edbed8b2b7246328b49ae87afbe58c1b0fca0812712678faee366
size 252572
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:30deb9eec4cb61d9865611f3a6b10c696638f540f93683654ea8d903b2e2ac0b
size 227161
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:3bd265f9e3ecb55e93b88039f363cfd7953fd29b4a68e8830d924a34b431fa29
size 264506
+3
View File
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:32bb7e8aa409d872a3e519a878601649b569f4af8e1ee77da940b785e834effe
size 232985
+3
View File
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:3846bb799587c6b1ff8335f10eeeef874248500b36c6be0a54ee532c4772c459
size 190947
+3
View File
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:9d77472902935937f04420b35375b3a869a4cb51f8ac08dab3c1d097d549de2d
size 221103
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:d33b9b3affcdb07086b5aebca8a3b9c2b009cdfc6f360950a0e72e65fbcb8f17
size 201602
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:adac93a0bb062d7199b819a600a0983ff011a75d16928fb80322cbb41f9284ea
size 158866
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:69f201cff948cb0a19810b7c4ad26d346f869ee2dd3141eba4f353332efa2e21
size 373649
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:793d7c2b28a751c6f467f2de788fcd462d3b8fd9cd5cb7adb5b32d78fb185394
size 236004
+3
View File
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:0e6c1dfaef840d5228645aaad3637c973b2f71c372595814f3b743a991f5c6fc
size 239128
+2
View File
@@ -17,6 +17,8 @@ docker compose --profile sim up redis console sim
Console dashboard: http://localhost:8090 Console dashboard: http://localhost:8090
> See also: [Deployment diagram](diagrams/png/12-deployment.png)
## Services ## Services
| Service | Port | Profile | Description | | Service | Port | Profile | Description |
+2
View File
@@ -164,6 +164,8 @@ Open http://localhost:8090 to see simulated nodes, workstream states, token coun
## Architecture ## Architecture
> See also: [Simulator Architecture diagram](diagrams/png/10-simulator-architecture.png)
``` ```
turnstone/sim/ turnstone/sim/
├── __init__.py # Public API: SimCluster, SimConfig ├── __init__.py # Public API: SimCluster, SimConfig
+2
View File
@@ -55,6 +55,8 @@ schema plus turnstone-specific metadata keys:
## Execution Pipeline ## Execution Pipeline
> See also: [Tool Pipeline diagram](diagrams/png/05-tool-pipeline.png)
Tool execution follows a three-phase pipeline inside `ChatSession._execute_tools()`: Tool execution follows a three-phase pipeline inside `ChatSession._execute_tools()`:
### Phase 1: Prepare ### Phase 1: Prepare
+56 -8
View File
@@ -1,21 +1,39 @@
[build-system] [build-system]
requires = ["hatchling"] requires = ["hatchling>=1.29"]
build-backend = "hatchling.build" build-backend = "hatchling.build"
[project] [project]
name = "turnstone" name = "turnstone"
version = "0.2.0" version = "0.2.1"
description = "AI chat client with tool use, agent tools, and persistent memory." description = "Multi-node AI orchestration platform with tool use, agent routing, and cluster simulation."
readme = "README.md" readme = "README.md"
license = "BUSL-1.1" license = "BUSL-1.1"
requires-python = ">=3.11" requires-python = ">=3.11"
dependencies = ["openai>=1.0", "httpx>=0.24"] authors = [{name = "Patrick Buckley", email = "buckleypm@gmail.com"}]
keywords = ["ai", "chat", "llm", "agent", "tools", "openai"]
classifiers = [
"Development Status :: 4 - Beta",
"Environment :: Console",
"Intended Audience :: Developers",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Programming Language :: Python :: 3.13",
"Topic :: Scientific/Engineering :: Artificial Intelligence",
]
dependencies = ["openai>=2.24", "httpx>=0.28"]
[project.urls]
Homepage = "https://github.com/turnstonelabs/turnstone"
Repository = "https://github.com/turnstonelabs/turnstone"
Issues = "https://github.com/turnstonelabs/turnstone/issues"
[project.optional-dependencies] [project.optional-dependencies]
test = ["pytest>=7.0"] test = ["pytest>=9.0"]
mq = ["redis>=5.0"] dev = ["ruff>=0.9", "mypy>=1.14", "types-redis>=4.6"]
console = ["redis>=5.0"] mq = ["redis>=7.2"]
sim = ["redis>=5.0"] console = ["redis>=7.2"]
sim = ["redis>=7.2"]
[project.scripts] [project.scripts]
turnstone = "turnstone.cli:main" turnstone = "turnstone.cli:main"
@@ -39,3 +57,33 @@ include = [
[tool.pytest.ini_options] [tool.pytest.ini_options]
testpaths = ["tests"] testpaths = ["tests"]
markers = ["live: requires a running LLM backend"]
[tool.ruff]
target-version = "py311"
line-length = 100
[tool.ruff.lint]
select = ["E", "F", "W", "I", "N", "UP", "B", "A", "SIM", "TCH"]
ignore = ["E501"]
[tool.ruff.format]
quote-style = "double"
[tool.mypy]
python_version = "3.11"
strict = true
warn_return_any = true
warn_unused_configs = true
disallow_untyped_defs = true
disallow_incomplete_defs = true
check_untyped_defs = true
no_implicit_optional = true
[[tool.mypy.overrides]]
module = ["sympy", "sympy.*", "numpy", "numpy.*"]
ignore_missing_imports = true
[[tool.mypy.overrides]]
module = "tests.*"
disallow_untyped_defs = false
+1 -1
View File
@@ -1,5 +1,5 @@
{ {
"description": "pcode behavior tests — tool selection, sequencing, and multi-step reasoning", "description": "turnstone behavior tests — tool selection, sequencing, and multi-step reasoning",
"defaults": { "defaults": {
"n_runs": 5, "n_runs": 5,
"max_turns": 15 "max_turns": 15
+2 -1
View File
@@ -1,6 +1,7 @@
import pytest
from unittest.mock import MagicMock from unittest.mock import MagicMock
import pytest
@pytest.fixture @pytest.fixture
def tmp_db(tmp_path, monkeypatch): def tmp_db(tmp_path, monkeypatch):
+10 -29
View File
@@ -6,8 +6,8 @@ from unittest.mock import patch
import pytest import pytest
from turnstone.core.auth import ( from turnstone.core.auth import (
AuthConfig,
WRITE_PATHS, WRITE_PATHS,
AuthConfig,
_extract_bearer, _extract_bearer,
_extract_cookie, _extract_cookie,
check_request, check_request,
@@ -18,7 +18,6 @@ from turnstone.core.auth import (
required_role, required_role,
) )
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# TestIsPublicPath # TestIsPublicPath
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -186,9 +185,7 @@ class TestExtractCookie:
assert _extract_cookie("", "turnstone_auth") is None assert _extract_cookie("", "turnstone_auth") is None
def test_spaces_around_value(self): def test_spaces_around_value(self):
assert ( assert _extract_cookie("turnstone_auth = tok_abc ", "turnstone_auth") == "tok_abc"
_extract_cookie("turnstone_auth = tok_abc ", "turnstone_auth") == "tok_abc"
)
def test_no_equals(self): def test_no_equals(self):
assert _extract_cookie("malformed", "turnstone_auth") is None assert _extract_cookie("malformed", "turnstone_auth") is None
@@ -288,44 +285,32 @@ class TestCheckRequest:
assert status == 401 assert status == 401
def test_api_read_token_ok(self, enabled): def test_api_read_token_ok(self, enabled):
allowed, status, msg = check_request( allowed, status, msg = check_request(enabled, "GET", "/api/workstreams", "Bearer tok_read")
enabled, "GET", "/api/workstreams", "Bearer tok_read"
)
assert allowed is True assert allowed is True
assert status == 200 assert status == 200
def test_api_full_token_ok(self, enabled): def test_api_full_token_ok(self, enabled):
allowed, status, msg = check_request( allowed, status, msg = check_request(enabled, "GET", "/api/workstreams", "Bearer tok_full")
enabled, "GET", "/api/workstreams", "Bearer tok_full"
)
assert allowed is True assert allowed is True
def test_write_read_token_403(self, enabled): def test_write_read_token_403(self, enabled):
allowed, status, msg = check_request( allowed, status, msg = check_request(enabled, "POST", "/api/send", "Bearer tok_read")
enabled, "POST", "/api/send", "Bearer tok_read"
)
assert allowed is False assert allowed is False
assert status == 403 assert status == 403
assert "Forbidden" in msg assert "Forbidden" in msg
def test_write_full_token_ok(self, enabled): def test_write_full_token_ok(self, enabled):
allowed, status, msg = check_request( allowed, status, msg = check_request(enabled, "POST", "/api/send", "Bearer tok_full")
enabled, "POST", "/api/send", "Bearer tok_full"
)
assert allowed is True assert allowed is True
assert status == 200 assert status == 200
def test_approve_read_token_403(self, enabled): def test_approve_read_token_403(self, enabled):
allowed, status, msg = check_request( allowed, status, msg = check_request(enabled, "POST", "/api/approve", "Bearer tok_read")
enabled, "POST", "/api/approve", "Bearer tok_read"
)
assert allowed is False assert allowed is False
assert status == 403 assert status == 403
def test_approve_full_token_ok(self, enabled): def test_approve_full_token_ok(self, enabled):
allowed, status, msg = check_request( allowed, status, msg = check_request(enabled, "POST", "/api/approve", "Bearer tok_full")
enabled, "POST", "/api/approve", "Bearer tok_full"
)
assert allowed is True assert allowed is True
def test_no_auth_header_string(self, enabled): def test_no_auth_header_string(self, enabled):
@@ -576,9 +561,7 @@ class TestServerAuth:
mock_mgr = MagicMock() mock_mgr = MagicMock()
mock_mgr.list_all.return_value = [mock_ws] mock_mgr.list_all.return_value = [mock_ws]
cls.server = srv_mod.ThreadedHTTPServer( cls.server = srv_mod.ThreadedHTTPServer(("127.0.0.1", 0), srv_mod.TurnstoneHTTPHandler)
("127.0.0.1", 0), srv_mod.TurnstoneHTTPHandler
)
cls.server.workstreams = mock_mgr cls.server.workstreams = mock_mgr
cls.server.skip_permissions = False cls.server.skip_permissions = False
cls.server.global_listeners = [] cls.server.global_listeners = []
@@ -835,9 +818,7 @@ class TestServerLogin:
mock_mgr = MagicMock() mock_mgr = MagicMock()
mock_mgr.list_all.return_value = [mock_ws] mock_mgr.list_all.return_value = [mock_ws]
cls.server = srv_mod.ThreadedHTTPServer( cls.server = srv_mod.ThreadedHTTPServer(("127.0.0.1", 0), srv_mod.TurnstoneHTTPHandler)
("127.0.0.1", 0), srv_mod.TurnstoneHTTPHandler
)
cls.server.workstreams = mock_mgr cls.server.workstreams = mock_mgr
cls.server.skip_permissions = False cls.server.skip_permissions = False
cls.server.global_listeners = [] cls.server.global_listeners = []
+7 -24
View File
@@ -3,10 +3,7 @@
import json import json
import queue import queue
import threading import threading
import time from unittest.mock import MagicMock
from http.server import HTTPServer, BaseHTTPRequestHandler
from socketserver import ThreadingMixIn
from unittest.mock import MagicMock, patch
import httpx import httpx
import pytest import pytest
@@ -14,12 +11,8 @@ import pytest
from turnstone.console.collector import ClusterCollector, NodeSnapshot from turnstone.console.collector import ClusterCollector, NodeSnapshot
from turnstone.mq.protocol import ( from turnstone.mq.protocol import (
ClusterStateEvent, ClusterStateEvent,
WorkstreamClosedEvent,
WorkstreamCreatedEvent,
WorkstreamRenameEvent,
) )
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Mock broker for collector tests # Mock broker for collector tests
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -128,7 +121,7 @@ class TestCollectorDiscovery:
def test_discover_emits_node_joined_event(self): def test_discover_emits_node_joined_event(self):
broker = MockBroker() broker = MockBroker()
c = _make_collector(broker) c = _make_collector(broker)
events = [] _events = []
q = queue.Queue() q = queue.Queue()
c.register_listener(q) c.register_listener(q)
@@ -218,9 +211,7 @@ class TestCollectorEvents:
c._nodes["node-a"] = NodeSnapshot( c._nodes["node-a"] = NodeSnapshot(
node_id="node-a", node_id="node-a",
server_url="http://a:8080", server_url="http://a:8080",
workstreams={ workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle", "node": "node-a"}},
"ws1": {"id": "ws1", "name": "test", "state": "idle", "node": "node-a"}
},
) )
event = ClusterStateEvent( event = ClusterStateEvent(
@@ -276,9 +267,7 @@ class TestCollectorEvents:
workstreams={"ws1": {"id": "ws1", "name": "old-name", "state": "idle"}}, workstreams={"ws1": {"id": "ws1", "name": "old-name", "state": "idle"}},
) )
event_json = json.dumps( event_json = json.dumps({"type": "ws_rename", "ws_id": "ws1", "name": "new-name"})
{"type": "ws_rename", "ws_id": "ws1", "name": "new-name"}
)
c._on_cluster_event(event_json) c._on_cluster_event(event_json)
assert c._nodes["node-a"].workstreams["ws1"]["name"] == "new-name" assert c._nodes["node-a"].workstreams["ws1"]["name"] == "new-name"
@@ -440,9 +429,7 @@ class TestCollectorQueries:
ws, _ = populated_collector.get_workstreams(sort_by="state") ws, _ = populated_collector.get_workstreams(sort_by="state")
states = [w["state"] for w in ws] states = [w["state"] for w in ws]
# running before attention before idle # running before attention before idle
assert ( assert states.index("running") < states.index("attention") < states.index("idle")
states.index("running") < states.index("attention") < states.index("idle")
)
def test_get_workstreams_combined_filters(self, populated_collector): def test_get_workstreams_combined_filters(self, populated_collector):
ws, total = populated_collector.get_workstreams(state="idle", node="node-a") ws, total = populated_collector.get_workstreams(state="idle", node="node-a")
@@ -588,15 +575,11 @@ class TestConsoleHTTPEndpoints:
mock_collector.get_overview.assert_called_once() mock_collector.get_overview.assert_called_once()
def test_get_nodes(self, server, mock_collector): def test_get_nodes(self, server, mock_collector):
status, data = self._get( status, data = self._get(server, "/api/cluster/nodes?sort=activity&limit=10&offset=0")
server, "/api/cluster/nodes?sort=activity&limit=10&offset=0"
)
assert status == 200 assert status == 200
assert len(data["nodes"]) == 1 assert len(data["nodes"]) == 1
assert data["total"] == 1 assert data["total"] == 1
mock_collector.get_nodes.assert_called_once_with( mock_collector.get_nodes.assert_called_once_with(sort_by="activity", limit=10, offset=0)
sort_by="activity", limit=10, offset=0
)
def test_get_workstreams(self, server, mock_collector): def test_get_workstreams(self, server, mock_collector):
status, data = self._get( status, data = self._get(
+1 -2
View File
@@ -1,12 +1,11 @@
"""Tests for turnstone.core.memory — database operations.""" """Tests for turnstone.core.memory — database operations."""
import turnstone.core.memory as memory
from turnstone.core.memory import ( from turnstone.core.memory import (
normalize_key,
open_db, open_db,
save_message, save_message,
search_history, search_history,
search_history_recent, search_history_recent,
normalize_key,
) )
+1 -1
View File
@@ -1,6 +1,6 @@
"""Tests for turnstone.core.memory — fts5_query and escape_like.""" """Tests for turnstone.core.memory — fts5_query and escape_like."""
from turnstone.core.memory import fts5_query, escape_like from turnstone.core.memory import escape_like, fts5_query
class TestFts5Query: class TestFts5Query:
+1 -1
View File
@@ -1,7 +1,7 @@
"""Tests for turnstone.ui.markdown — MarkdownRenderer.""" """Tests for turnstone.ui.markdown — MarkdownRenderer."""
from turnstone.ui.colors import BOLD, CYAN, DIM, ITALIC, MAGENTA
from turnstone.ui.markdown import MarkdownRenderer from turnstone.ui.markdown import MarkdownRenderer
from turnstone.ui.colors import BOLD, MAGENTA, CYAN, DIM, ITALIC, RESET
class TestMarkdownRenderer: class TestMarkdownRenderer:
-1
View File
@@ -37,7 +37,6 @@ from turnstone.mq.protocol import (
WorkstreamRenameEvent, WorkstreamRenameEvent,
) )
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Inbound message round-trip tests # Inbound message round-trip tests
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
+1 -1
View File
@@ -1,6 +1,6 @@
"""Tests for turnstone.core.sandbox — validate_math_code and auto_print_wrap.""" """Tests for turnstone.core.sandbox — validate_math_code and auto_print_wrap."""
from turnstone.core.sandbox import validate_math_code, auto_print_wrap from turnstone.core.sandbox import auto_print_wrap, validate_math_code
class TestValidateMathCode: class TestValidateMathCode:
+1 -1
View File
@@ -1,6 +1,6 @@
"""Tests for turnstone.eval — score_run and _match_action.""" """Tests for turnstone.eval — score_run and _match_action."""
from turnstone.eval import score_run, _match_action from turnstone.eval import _match_action, score_run
class TestMatchAction: class TestMatchAction:
+351 -138
View File
@@ -1,15 +1,21 @@
"""Integration tests against a live llama.cpp backend on port 8000. """Tests for turnstone ChatSession and server endpoints.
These tests use turnstone's HeadlessSession to run actual LLM inference Mock-based tests verify streaming, tool calling, multi-turn conversation,
and tool execution against the backend. They verify end-to-end behavior: and session configuration WITHOUT a running LLM backend. The mocks replace
model connectivity, tool calling, response quality, and session mechanics. only the OpenAI streaming layer -- tool execution (bash, math, read_file)
still runs real subprocesses.
Requires: llama-server (or compatible OpenAI API) running on localhost:8000. The TestBackendConnectivity class is marked @pytest.mark.live and requires a
running llama-server (or compatible OpenAI API) on localhost:8000.
Run with: pytest tests/test_server_live.py -v --timeout=120 The TestServerHealthMetrics class spins up an in-process HTTP server and
needs no LLM backend at all.
The TestServerHealthMetrics class does NOT require a live LLM and can be run Run all non-live tests:
independently: pytest tests/test_server_live.py::TestServerHealthMetrics -v pytest tests/test_server_live.py -v -m "not live"
Run everything (needs backend):
pytest tests/test_server_live.py -v --timeout=120
""" """
import json import json
@@ -17,15 +23,15 @@ import os
import queue import queue
import tempfile import tempfile
import threading import threading
import time from types import SimpleNamespace
from unittest.mock import MagicMock
import httpx import httpx
import pytest import pytest
from openai import OpenAI from openai import OpenAI
from turnstone.core.session import ChatSession
from turnstone.core.tools import TOOLS
import turnstone.core.memory as _memory_module import turnstone.core.memory as _memory_module
from turnstone.core.session import ChatSession
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Fixtures # Fixtures
@@ -35,8 +41,8 @@ BASE_URL = os.environ.get("TURNSTONE_TEST_BASE_URL", "http://localhost:8000/v1")
@pytest.fixture(scope="module") @pytest.fixture(scope="module")
def client(): def live_client():
"""Create an OpenAI client pointed at the local backend.""" """Create an OpenAI client pointed at the local backend (live tests only)."""
return OpenAI( return OpenAI(
base_url=BASE_URL, base_url=BASE_URL,
api_key=os.environ.get("TURNSTONE_TEST_API_KEY", "not-needed"), api_key=os.environ.get("TURNSTONE_TEST_API_KEY", "not-needed"),
@@ -44,9 +50,9 @@ def client():
@pytest.fixture(scope="module") @pytest.fixture(scope="module")
def model_id(client): def live_model_id(live_client):
"""Auto-detect the model name from the backend.""" """Auto-detect the model name from the backend (live tests only)."""
models = client.models.list() models = live_client.models.list()
ids = [m.id for m in models.data] ids = [m.id for m in models.data]
assert len(ids) > 0, "No models found on the backend" assert len(ids) > 0, "No models found on the backend"
return ids[0] return ids[0]
@@ -125,16 +131,13 @@ def tmp_db():
os.unlink(path) os.unlink(path)
def _make_session( def _make_session(client, model_id, tmp_db, **kwargs) -> tuple[ChatSession, RecordingUI]:
client, model_id, tmp_db, **kwargs
) -> tuple[ChatSession, RecordingUI]:
"""Create a ChatSession with RecordingUI and sensible test defaults.""" """Create a ChatSession with RecordingUI and sensible test defaults."""
ui = RecordingUI() ui = RecordingUI()
defaults = dict( defaults = dict(
client=client, client=client,
model=model_id, model=model_id,
ui=ui, ui=ui,
persona=None,
instructions=None, instructions=None,
temperature=0.3, temperature=0.3,
max_tokens=2048, max_tokens=2048,
@@ -148,174 +151,339 @@ def _make_session(
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Tests — Backend connectivity # Mock streaming helpers
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def _make_chunk(
*,
content=None,
reasoning_content=None,
tool_calls=None,
finish_reason=None,
usage=None,
):
"""Build a single mock streaming chunk matching the OpenAI format.
The chunk structure mirrors openai.types.chat.ChatCompletionChunk:
chunk.choices[0].delta.content
chunk.choices[0].delta.reasoning_content
chunk.choices[0].delta.tool_calls
chunk.choices[0].finish_reason
chunk.usage
"""
delta = SimpleNamespace(
content=content,
reasoning_content=reasoning_content,
reasoning=None,
tool_calls=tool_calls,
role=None,
model_extra=None,
)
choice = SimpleNamespace(delta=delta, finish_reason=finish_reason)
chunk = SimpleNamespace(choices=[choice], usage=usage)
return chunk
def _make_tool_call_deltas(call_id, name, arguments):
"""Build a list of tool_call delta objects for a single tool call.
Returns a list with one element (single tool call at index 0).
"""
fn = SimpleNamespace(name=name, arguments=arguments)
return [SimpleNamespace(index=0, id=call_id, function=fn)]
def _usage(prompt=100, completion=50, total=None):
"""Build a mock usage object."""
return SimpleNamespace(
prompt_tokens=prompt,
completion_tokens=completion,
total_tokens=total or (prompt + completion),
)
def make_mock_stream(
content_tokens=None,
reasoning_tokens=None,
tool_calls=None,
finish_reason="stop",
usage=None,
):
"""Create an iterable of mock chunks simulating an OpenAI streaming response.
Parameters
----------
content_tokens : list[str] | None
Content token strings, each emitted as a separate chunk.
reasoning_tokens : list[str] | None
Reasoning token strings, emitted before content.
tool_calls : list[tuple[str, str, str]] | None
Each entry is (call_id, function_name, arguments_json).
When provided, finish_reason defaults to "tool_calls".
finish_reason : str
Finish reason on the last content/tool chunk.
usage : SimpleNamespace | None
Usage object for the final chunk. Defaults to a sensible value.
"""
chunks = []
if reasoning_tokens:
for token in reasoning_tokens:
chunks.append(_make_chunk(reasoning_content=token))
if content_tokens:
for i, token in enumerate(content_tokens):
is_last = (i == len(content_tokens) - 1) and not tool_calls
chunks.append(
_make_chunk(
content=token,
finish_reason=finish_reason if is_last else None,
)
)
if tool_calls:
for i, (call_id, name, arguments) in enumerate(tool_calls):
is_last = i == len(tool_calls) - 1
tc_deltas = _make_tool_call_deltas(call_id, name, arguments)
chunks.append(
_make_chunk(
tool_calls=tc_deltas,
finish_reason="tool_calls" if is_last else None,
)
)
# Final usage-only chunk (no choices)
if usage is None:
usage = _usage()
chunks.append(SimpleNamespace(choices=[], usage=usage))
return iter(chunks)
def _mock_client():
"""Create a mock OpenAI client with a patchable chat.completions.create."""
client = MagicMock(spec=OpenAI)
client.chat = MagicMock()
client.chat.completions = MagicMock()
client.chat.completions.create = MagicMock()
return client
# ---------------------------------------------------------------------------
# Tests -- Backend connectivity (live, requires running LLM)
# ---------------------------------------------------------------------------
@pytest.mark.live
class TestBackendConnectivity: class TestBackendConnectivity:
"""Verify the LLM backend is reachable and returns valid responses.""" """Verify the LLM backend is reachable and returns valid responses."""
def test_models_endpoint(self, client): def test_models_endpoint(self, live_client):
models = client.models.list() models = live_client.models.list()
assert len(models.data) > 0 assert len(models.data) > 0
def test_model_id_detected(self, model_id): def test_model_id_detected(self, live_model_id):
assert isinstance(model_id, str) assert isinstance(live_model_id, str)
assert len(model_id) > 0 assert len(live_model_id) > 0
def test_basic_completion(self, client, model_id): def test_basic_completion(self, live_client, live_model_id):
"""Raw API call no turnstone involved.""" """Raw API call -- no turnstone involved."""
resp = client.chat.completions.create( resp = live_client.chat.completions.create(
model=model_id, model=live_model_id,
messages=[{"role": "user", "content": "Say 'hello'"}], messages=[{"role": "user", "content": "Say 'hello'"}],
max_completion_tokens=200, max_completion_tokens=200,
temperature=0.0, temperature=0.0,
stream=False, stream=False,
) )
assert ( assert resp.choices[0].message.content or resp.choices[0].message.reasoning_content
resp.choices[0].message.content or resp.choices[0].message.reasoning_content
)
assert resp.usage.total_tokens > 0 assert resp.usage.total_tokens > 0
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Tests Streaming session # Tests -- Streaming session (mocked)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
class TestStreamingSession: class TestStreamingSession:
"""Test ChatSession.send() with streaming against the live backend.""" """Test ChatSession.send() with mocked streaming responses."""
def test_simple_response(self, client, model_id, tmp_db): def test_simple_response(self, tmp_db):
"""Model responds to a basic prompt via streaming.""" """Mock returns content tokens; verify RecordingUI captures them."""
session, ui = _make_session(client, model_id, tmp_db) client = _mock_client()
session.send("Reply with exactly: PONG") client.chat.completions.create.return_value = make_mock_stream(
content_tokens=["Hello", " ", "world"],
)
# Should have gotten some content or reasoning session, ui = _make_session(client, "mock-model", tmp_db)
total = ui.full_content + ui.full_reasoning session._title_generated = True # prevent background title generation
assert len(total) > 0, "No output from model"
session.send("Say hello")
assert "Hello world" in ui.full_content
def test_reasoning_tokens_appear(self, tmp_db):
"""Mock returns reasoning tokens then content; verify both captured."""
client = _mock_client()
client.chat.completions.create.return_value = make_mock_stream(
reasoning_tokens=["Let me", " think..."],
content_tokens=["The answer", " is 56"],
)
session, ui = _make_session(client, "mock-model", tmp_db)
session._title_generated = True
def test_reasoning_tokens_appear(self, client, model_id, tmp_db):
"""Model produces reasoning tokens (extended thinking)."""
session, ui = _make_session(client, model_id, tmp_db)
session.send("What is 7 * 8?") session.send("What is 7 * 8?")
# This model uses reasoning_content, so we expect reasoning tokens assert len(ui.reasoning_tokens) > 0
assert len(ui.reasoning_tokens) > 0, "No reasoning tokens received" assert "think" in ui.full_reasoning.lower()
assert "56" in ui.full_content
def test_stream_end_event(self, client, model_id, tmp_db): def test_stream_end_event(self, tmp_db):
"""stream_end event is emitted after response.""" """stream_end event is emitted after response."""
session, ui = _make_session(client, model_id, tmp_db) client = _mock_client()
client.chat.completions.create.return_value = make_mock_stream(
content_tokens=["Hi"],
)
session, ui = _make_session(client, "mock-model", tmp_db)
session._title_generated = True
session.send("Say hi") session.send("Say hi")
event_types = [e[0] for e in ui.events] event_types = [e[0] for e in ui.events]
assert "stream_end" in event_types assert "stream_end" in event_types
def test_thinking_lifecycle(self, client, model_id, tmp_db): def test_thinking_lifecycle(self, tmp_db):
"""thinking_start and thinking_stop bracket the response.""" """thinking_start and thinking_stop bracket the response."""
session, ui = _make_session(client, model_id, tmp_db) client = _mock_client()
client.chat.completions.create.return_value = make_mock_stream(
content_tokens=["Hi", " there"],
)
session, ui = _make_session(client, "mock-model", tmp_db)
session._title_generated = True
session.send("Say hi") session.send("Say hi")
event_types = [e[0] for e in ui.events] event_types = [e[0] for e in ui.events]
assert "thinking_start" in event_types assert "thinking_start" in event_types
assert "thinking_stop" in event_types assert "thinking_stop" in event_types
# thinking_start should come before thinking_stop
start_idx = event_types.index("thinking_start") start_idx = event_types.index("thinking_start")
stop_idx = event_types.index("thinking_stop") stop_idx = event_types.index("thinking_stop")
assert start_idx < stop_idx assert start_idx < stop_idx
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Tests Tool calling # Tests -- Tool calling (mocked LLM, real tool execution)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
class TestToolCalling: class TestToolCalling:
"""Test that the model can invoke tools and turnstone executes them.""" """Test that mocked tool_calls trigger real tool execution."""
def test_math_tool(self, client, model_id, tmp_db): def test_math_tool(self, tmp_db):
"""Model uses the math tool for computation.""" """First call returns tool_call for math(code='2+2'), second returns content."""
session, ui = _make_session( client = _mock_client()
client,
model_id, # First create() call: model requests math tool
tmp_db, stream1 = make_mock_stream(
instructions="You have tools. Use the math tool to compute results. Always use tools when asked to calculate.", tool_calls=[("call_math_1", "math", json.dumps({"code": "2+2"}))],
) )
session.send("Use the math tool to calculate: 17 * 23. Report the result.") # Second create() call: model produces final answer
stream2 = make_mock_stream(
content_tokens=["The result is ", "4"],
)
client.chat.completions.create.side_effect = [stream1, stream2]
# Check if math tool was invoked session, ui = _make_session(client, "mock-model", tmp_db)
session._title_generated = True
session.send("Calculate 2+2")
# math tool was invoked and returned a result
math_results = [r for r in ui.tool_results if r[0] == "math"] math_results = [r for r in ui.tool_results if r[0] == "math"]
if math_results: assert len(math_results) > 0
# Verify the result contains 391 assert "4" in math_results[0][1]
assert "391" in math_results[0][1], (
f"Expected 391, got: {math_results[0][1]}"
)
else:
# Model may have answered directly — check content
total = ui.full_content + ui.full_reasoning
assert "391" in total, f"Expected 391 somewhere in output"
def test_bash_tool(self, client, model_id, tmp_db): # Final content contains the answer
"""Model uses bash to answer a system question.""" assert "4" in ui.full_content
session, ui = _make_session(
client, def test_bash_tool(self, tmp_db):
model_id, """First call returns tool_call for bash, second returns content."""
tmp_db, client = _mock_client()
instructions="You have tools. Use the bash tool to run commands. Always use bash when asked about system info.",
stream1 = make_mock_stream(
tool_calls=[("call_bash_1", "bash", json.dumps({"command": "echo hello"}))],
) )
session.send( stream2 = make_mock_stream(
"Use the bash tool to run 'echo hello_from_test' and report what it prints." content_tokens=["The command printed: ", "hello"],
) )
client.chat.completions.create.side_effect = [stream1, stream2]
session, ui = _make_session(client, "mock-model", tmp_db)
session._title_generated = True
session.send("Run echo hello")
bash_results = [r for r in ui.tool_results if r[0] == "bash"] bash_results = [r for r in ui.tool_results if r[0] == "bash"]
if bash_results: assert len(bash_results) > 0
assert "hello_from_test" in bash_results[0][1] assert "hello" in bash_results[0][1]
else:
total = ui.full_content + ui.full_reasoning
assert "hello_from_test" in total, "Expected bash output in response"
def test_read_file_tool(self, client, model_id, tmp_db): def test_read_file_tool(self, tmp_db):
"""Model uses read_file to read a known file.""" """First call returns tool_call for read_file, second returns content."""
# Create a temp file for the model to read
with tempfile.NamedTemporaryFile(mode="w", suffix=".txt", delete=False) as f: with tempfile.NamedTemporaryFile(mode="w", suffix=".txt", delete=False) as f:
f.write("SECRET_CONTENT_42\n") f.write("SECRET_CONTENT_42\n")
path = f.name path = f.name
try: try:
session, ui = _make_session( client = _mock_client()
client,
model_id, stream1 = make_mock_stream(
tmp_db, tool_calls=[("call_read_1", "read_file", json.dumps({"path": path}))],
instructions="You have tools. Use the read_file tool to read files. Always use read_file when asked to read a file.", )
) stream2 = make_mock_stream(
session.send( content_tokens=["The file says: SECRET_CONTENT_42"],
f"Use the read_file tool to read {path} and tell me what it says." )
) client.chat.completions.create.side_effect = [stream1, stream2]
session, ui = _make_session(client, "mock-model", tmp_db)
session._title_generated = True
session.send(f"Read {path}")
# read_file was invoked (UI gets a summary like "1 lines")
read_results = [r for r in ui.tool_results if r[0] == "read_file"] read_results = [r for r in ui.tool_results if r[0] == "read_file"]
assert len(read_results) > 0, "read_file tool was not called" assert len(read_results) > 0
# The model sees the actual file content and should relay it # Model relays the content
total = ui.full_content + ui.full_reasoning assert "SECRET_CONTENT_42" in ui.full_content
assert "SECRET_CONTENT_42" in total, (
f"Model didn't relay file content. Got: {total[:500]}"
)
finally: finally:
os.unlink(path) os.unlink(path)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Tests Multi-turn conversation # Tests -- Multi-turn conversation (mocked)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
class TestMultiTurn: class TestMultiTurn:
"""Test multi-turn conversation state.""" """Test multi-turn conversation state with mocked responses."""
def test_context_retained(self, tmp_db):
"""Second send references context from the first."""
client = _mock_client()
stream1 = make_mock_stream(
content_tokens=["I'll remember ", "Zephyr"],
)
stream2 = make_mock_stream(
content_tokens=["Your name is ", "Zephyr"],
)
client.chat.completions.create.side_effect = [stream1, stream2]
session, ui = _make_session(client, "mock-model", tmp_db, max_tokens=1024)
session._title_generated = True
def test_context_retained(self, client, model_id, tmp_db):
"""Second message can reference the first."""
session, ui = _make_session(client, model_id, tmp_db, max_tokens=1024)
session.send("My name is Zephyr. Remember it.") session.send("My name is Zephyr. Remember it.")
# Reset UI tracking for second turn # Reset UI tracking for second turn
@@ -324,58 +492,86 @@ class TestMultiTurn:
session.send("What is my name?") session.send("What is my name?")
total = ui.full_content + ui.full_reasoning assert "zephyr" in ui.full_content.lower()
assert "zephyr" in total.lower(), f"Model forgot the name. Got: {total[:300]}"
def test_message_list_grows(self, client, model_id, tmp_db): def test_message_list_grows(self, tmp_db):
"""Each send adds user + assistant messages.""" """Each send adds user + assistant messages."""
session, ui = _make_session(client, model_id, tmp_db, max_tokens=512) client = _mock_client()
stream1 = make_mock_stream(content_tokens=["Hello"])
stream2 = make_mock_stream(content_tokens=["World"])
client.chat.completions.create.side_effect = [stream1, stream2]
session, ui = _make_session(client, "mock-model", tmp_db, max_tokens=512)
session._title_generated = True
initial_count = len(session.messages) initial_count = len(session.messages)
session.send("Hello") session.send("Hello")
after_first = len(session.messages)
assert after_first >= initial_count + 2
# Should have at least user + assistant session.send("World")
assert len(session.messages) >= initial_count + 2 after_second = len(session.messages)
assert after_second >= after_first + 2
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Tests Session configuration # Tests -- Session configuration (mocked)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
class TestSessionConfig: class TestSessionConfig:
"""Test session construction and configuration.""" """Test session construction and configuration with mocked responses."""
def test_creative_mode_no_tools(self, client, model_id, tmp_db): def test_creative_mode_no_tools(self, tmp_db):
"""In creative mode, tools are not sent to the API.""" """In creative mode, create() is called WITHOUT tools kwarg."""
session, ui = _make_session(client, model_id, tmp_db, max_tokens=256) client = _mock_client()
client.chat.completions.create.return_value = make_mock_stream(
content_tokens=["A haiku about code"],
)
session, ui = _make_session(client, "mock-model", tmp_db, max_tokens=256)
session._title_generated = True
session.creative_mode = True session.creative_mode = True
# Re-init system messages so creative_mode takes effect
session._init_system_messages()
session.send("Write a haiku about code.") session.send("Write a haiku about code.")
# Verify create() was called without 'tools' in kwargs
call_kwargs = client.chat.completions.create.call_args
assert "tools" not in call_kwargs.kwargs, "tools should not be passed in creative mode"
# Should get content back without tool calls # Should get content back without tool calls
total = ui.full_content + ui.full_reasoning assert len(ui.full_content) > 0
assert len(total) > 0
assert len(ui.tool_results) == 0 assert len(ui.tool_results) == 0
def test_custom_instructions(self, client, model_id, tmp_db): def test_custom_instructions(self, tmp_db):
"""Custom instructions are included in the session.""" """Custom instructions appear in system messages."""
client = _mock_client()
client.chat.completions.create.return_value = make_mock_stream(
content_tokens=["Hello. ENDMARKER"],
)
session, ui = _make_session( session, ui = _make_session(
client, client,
model_id, "mock-model",
tmp_db, tmp_db,
instructions="Always end your response with ENDMARKER.", instructions="Always end your response with ENDMARKER.",
max_tokens=512, max_tokens=512,
) )
session.send("Say hello briefly.") session._title_generated = True
total = ui.full_content # Verify custom instructions appear in system messages
# We can't strictly guarantee the model follows instructions, dev_msg = session.system_messages[0]
# but we verify the session didn't error out assert "ENDMARKER" in dev_msg["content"]
session.send("Say hello briefly.")
assert len(ui.errors) == 0 assert len(ui.errors) == 0
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Tests /health and /metrics endpoints (no live LLM required) # Tests -- /health and /metrics endpoints (no live LLM required)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -391,23 +587,40 @@ class TestServerHealthMetrics:
@classmethod @classmethod
def setup_class(cls): def setup_class(cls):
from unittest.mock import MagicMock from unittest.mock import MagicMock
import turnstone.server as srv_mod import turnstone.server as srv_mod
from turnstone.core.workstream import WorkstreamState
# Reset module-level metrics so each test run starts fresh # Reset module-level metrics so each test run starts fresh
srv_mod._metrics = srv_mod.MetricsCollector() from turnstone.core.metrics import MetricsCollector
from turnstone.core.workstream import WorkstreamState
srv_mod._metrics = MetricsCollector()
srv_mod._metrics.model = "test-model" srv_mod._metrics.model = "test-model"
# Mock WorkstreamManager.list_all() to return one idle workstream # Mock WorkstreamManager.list_all() to return one idle workstream
mock_ui = MagicMock()
mock_ui._ws_lock = threading.Lock()
mock_ui._ws_prompt_tokens = 0
mock_ui._ws_completion_tokens = 0
mock_ui._ws_messages = 0
mock_ui._ws_tool_calls = {}
mock_ui._ws_context_ratio = 0.0
mock_session = MagicMock()
mock_session.session_id = "test-session-id"
mock_ws = MagicMock() mock_ws = MagicMock()
mock_ws.id = "test-ws"
mock_ws.name = "test"
mock_ws.state = WorkstreamState.IDLE mock_ws.state = WorkstreamState.IDLE
mock_ws.ui = mock_ui
mock_ws.session = mock_session
mock_mgr = MagicMock() mock_mgr = MagicMock()
mock_mgr.list_all.return_value = [mock_ws] mock_mgr.list_all.return_value = [mock_ws]
# Start a server on a random port (port 0 OS assigns free port) # Start a server on a random port (port 0 -> OS assigns free port)
cls.server = srv_mod.ThreadedHTTPServer( cls.server = srv_mod.ThreadedHTTPServer(("127.0.0.1", 0), srv_mod.TurnstoneHTTPHandler)
("127.0.0.1", 0), srv_mod.TurnstoneHTTPHandler
)
from turnstone.core.auth import AuthConfig from turnstone.core.auth import AuthConfig
cls.server.workstreams = mock_mgr cls.server.workstreams = mock_mgr
+4 -15
View File
@@ -1,7 +1,6 @@
"""Tests for turnstone.core.session — ChatSession construction.""" """Tests for turnstone.core.session — ChatSession construction."""
import json import json
import os
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
from turnstone.core.session import ChatSession from turnstone.core.session import ChatSession
@@ -46,10 +45,12 @@ class NullUI:
def on_state_change(self, state): def on_state_change(self, state):
pass pass
def on_rename(self, name):
pass
def _make_session( def _make_session(
mock_openai_client=None, mock_openai_client=None,
persona=None,
instructions=None, instructions=None,
**kwargs, **kwargs,
): ):
@@ -59,7 +60,6 @@ def _make_session(
client=client, client=client,
model="test-model", model="test-model",
ui=NullUI(), ui=NullUI(),
persona=persona,
instructions=instructions, instructions=instructions,
temperature=0.5, temperature=0.5,
max_tokens=4096, max_tokens=4096,
@@ -77,15 +77,6 @@ class TestChatSessionConstruction:
roles = [m["role"] for m in session.system_messages] roles = [m["role"] for m in session.system_messages]
assert "developer" in roles assert "developer" in roles
def test_persona_injected_into_chat_template_kwargs(self, tmp_db):
session = _make_session(persona="Helpful assistant")
assert "model_identity" in session._chat_template_kwargs
assert "Helpful assistant" in session._chat_template_kwargs["model_identity"]
def test_no_persona_no_model_identity(self, tmp_db):
session = _make_session(persona=None)
assert "model_identity" not in session._chat_template_kwargs
def test_instructions_appended_to_developer_message(self, tmp_db): def test_instructions_appended_to_developer_message(self, tmp_db):
session = _make_session(instructions="Always be concise.") session = _make_session(instructions="Always be concise.")
dev_msgs = [m for m in session.system_messages if m["role"] == "developer"] dev_msgs = [m for m in session.system_messages if m["role"] == "developer"]
@@ -268,8 +259,6 @@ class TestPlanExec:
monkeypatch.chdir(tmp_path) monkeypatch.chdir(tmp_path)
session = _make_session() session = _make_session()
agent_output = "## Goal\n\nBuild it." agent_output = "## Goal\n\nBuild it."
call_id, content, _ = self._run_plan( call_id, content, _ = self._run_plan(session, "do stuff", agent_return=agent_output)
session, "do stuff", agent_return=agent_output
)
assert call_id == "test-call-1" assert call_id == "test-call-1"
assert content == agent_output assert content == agent_output
+13 -28
View File
@@ -1,22 +1,20 @@
"""Tests for session persistence and resume functionality.""" """Tests for session persistence and resume functionality."""
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock
import turnstone.core.memory as memory
from turnstone.core.memory import ( from turnstone.core.memory import (
register_session, delete_session,
update_session_title,
set_session_alias,
resolve_session,
list_sessions, list_sessions,
load_session_messages, load_session_messages,
delete_session,
save_message,
open_db, open_db,
register_session,
resolve_session,
save_message,
set_session_alias,
update_session_title,
) )
from turnstone.core.session import ChatSession from turnstone.core.session import ChatSession
# ── Session registration ────────────────────────────────────────────── # ── Session registration ──────────────────────────────────────────────
@@ -154,12 +152,8 @@ class TestLoadSessionMessages:
def test_tool_calls_with_ids(self, tmp_db): def test_tool_calls_with_ids(self, tmp_db):
save_message("s1", "user", "run ls") save_message("s1", "user", "run ls")
save_message("s1", "assistant", "Let me check.") save_message("s1", "assistant", "Let me check.")
save_message( save_message("s1", "tool_call", None, "bash", '{"command":"ls"}', tool_call_id="call_abc")
"s1", "tool_call", None, "bash", '{"command":"ls"}', tool_call_id="call_abc" save_message("s1", "tool_result", "file1.txt\nfile2.txt", "bash", tool_call_id="call_abc")
)
save_message(
"s1", "tool_result", "file1.txt\nfile2.txt", "bash", tool_call_id="call_abc"
)
msgs = load_session_messages("s1") msgs = load_session_messages("s1")
assert len(msgs) == 3 # user, assistant+tool_calls, tool assert len(msgs) == 3 # user, assistant+tool_calls, tool
# Assistant should have content merged with tool_calls # Assistant should have content merged with tool_calls
@@ -186,12 +180,8 @@ class TestLoadSessionMessages:
def test_parallel_tool_calls(self, tmp_db): def test_parallel_tool_calls(self, tmp_db):
save_message("s1", "user", "search two things") save_message("s1", "user", "search two things")
save_message( save_message("s1", "tool_call", None, "search", '{"query":"a"}', tool_call_id="call_1")
"s1", "tool_call", None, "search", '{"query":"a"}', tool_call_id="call_1" save_message("s1", "tool_call", None, "search", '{"query":"b"}', tool_call_id="call_2")
)
save_message(
"s1", "tool_call", None, "search", '{"query":"b"}', tool_call_id="call_2"
)
save_message("s1", "tool_result", "result a", "search", tool_call_id="call_1") save_message("s1", "tool_result", "result a", "search", tool_call_id="call_1")
save_message("s1", "tool_result", "result b", "search", tool_call_id="call_2") save_message("s1", "tool_result", "result b", "search", tool_call_id="call_2")
msgs = load_session_messages("s1") msgs = load_session_messages("s1")
@@ -231,9 +221,7 @@ class TestDeleteSession:
class TestSaveMessageToolCallId: class TestSaveMessageToolCallId:
def test_tool_call_id_stored(self, tmp_db): def test_tool_call_id_stored(self, tmp_db):
save_message( save_message("s1", "tool_call", None, "bash", '{"cmd":"ls"}', tool_call_id="call_xyz")
"s1", "tool_call", None, "bash", '{"cmd":"ls"}', tool_call_id="call_xyz"
)
conn = open_db() conn = open_db()
try: try:
row = conn.execute( row = conn.execute(
@@ -293,7 +281,6 @@ class TestResumeSession:
client=mock_openai_client, client=mock_openai_client,
model="test-model", model="test-model",
ui=MagicMock(), ui=MagicMock(),
persona=None,
instructions=None, instructions=None,
temperature=0.5, temperature=0.5,
max_tokens=1000, max_tokens=1000,
@@ -314,7 +301,6 @@ class TestResumeSession:
client=mock_openai_client, client=mock_openai_client,
model="test-model", model="test-model",
ui=MagicMock(), ui=MagicMock(),
persona=None,
instructions=None, instructions=None,
temperature=0.5, temperature=0.5,
max_tokens=1000, max_tokens=1000,
@@ -327,7 +313,6 @@ class TestResumeSession:
client=mock_openai_client, client=mock_openai_client,
model="test-model", model="test-model",
ui=MagicMock(), ui=MagicMock(),
persona=None,
instructions=None, instructions=None,
temperature=0.5, temperature=0.5,
max_tokens=1000, max_tokens=1000,
@@ -347,7 +332,7 @@ class TestSaveMessageUpdatesSession:
register_session("s1") register_session("s1")
save_message("s1", "user", "first") save_message("s1", "user", "first")
rows = list_sessions() rows = list_sessions()
original_updated = rows[0][4] _original_updated = rows[0][4]
import time import time
+1 -5
View File
@@ -4,18 +4,14 @@ from __future__ import annotations
import asyncio import asyncio
import random import random
from unittest.mock import MagicMock, call from unittest.mock import MagicMock
import pytest import pytest
from turnstone.mq.protocol import ( from turnstone.mq.protocol import (
ContentEvent,
InboundMessage,
OutboundEvent, OutboundEvent,
SendMessage, SendMessage,
StateChangeEvent, StateChangeEvent,
TurnCompleteEvent,
WorkstreamCreatedEvent,
) )
from turnstone.sim.config import SimConfig from turnstone.sim.config import SimConfig
from turnstone.sim.engine import SimEngine, ToolSimulationError from turnstone.sim.engine import SimEngine, ToolSimulationError
+11 -15
View File
@@ -1,13 +1,13 @@
"""Tests for turnstone.core.tools — JSON auto-loading and schema validation.""" """Tests for turnstone.core.tools — JSON auto-loading and schema validation."""
from turnstone.core.tools import ( from turnstone.core.tools import (
TOOLS,
AGENT_TOOLS,
TASK_AGENT_TOOLS,
AGENT_AUTO_TOOLS,
TASK_AUTO_TOOLS,
PRIMARY_KEY_MAP,
_META, _META,
AGENT_AUTO_TOOLS,
AGENT_TOOLS,
PRIMARY_KEY_MAP,
TASK_AGENT_TOOLS,
TASK_AUTO_TOOLS,
TOOLS,
) )
@@ -23,9 +23,7 @@ class TestToolsSchema:
def test_all_tools_have_description(self): def test_all_tools_have_description(self):
for tool in TOOLS: for tool in TOOLS:
assert "description" in tool["function"], ( assert "description" in tool["function"], f"Tool missing description: {tool}"
f"Tool missing description: {tool}"
)
assert len(tool["function"]["description"]) > 0 assert len(tool["function"]["description"]) > 0
def test_all_tools_have_parameters(self): def test_all_tools_have_parameters(self):
@@ -84,8 +82,8 @@ class TestToolsMetadata:
def test_auto_approve_sets_match(self): def test_auto_approve_sets_match(self):
expected = {"read_file", "search", "math", "man", "web_fetch", "web_search"} expected = {"read_file", "search", "math", "man", "web_fetch", "web_search"}
assert AGENT_AUTO_TOOLS == expected assert expected == AGENT_AUTO_TOOLS
assert TASK_AUTO_TOOLS == expected assert expected == TASK_AUTO_TOOLS
def test_primary_key_map(self): def test_primary_key_map(self):
expected = { expected = {
@@ -104,7 +102,7 @@ class TestToolsMetadata:
"recall": "query", "recall": "query",
"forget": "key", "forget": "key",
} }
assert PRIMARY_KEY_MAP == expected assert expected == PRIMARY_KEY_MAP
def test_no_metadata_in_function_dicts(self): def test_no_metadata_in_function_dicts(self):
"""Ensure turnstone metadata keys are stripped from the OpenAI schema.""" """Ensure turnstone metadata keys are stripped from the OpenAI schema."""
@@ -112,9 +110,7 @@ class TestToolsMetadata:
for tool in TOOLS: for tool in TOOLS:
func = tool["function"] func = tool["function"]
leaked = meta_keys & set(func) leaked = meta_keys & set(func)
assert not leaked, ( assert not leaked, f"Tool '{func['name']}' leaks metadata into function dict: {leaked}"
f"Tool '{func['name']}' leaks metadata into function dict: {leaked}"
)
def test_meta_has_all_tools(self): def test_meta_has_all_tools(self):
tool_names = {t["function"]["name"] for t in TOOLS} tool_names = {t["function"]["name"] for t in TOOLS}
+13 -13
View File
@@ -4,10 +4,8 @@ import threading
import time import time
import pytest import pytest
from unittest.mock import MagicMock
from turnstone.core.workstream import WorkstreamManager, WorkstreamState, Workstream
from turnstone.core.workstream import Workstream, WorkstreamManager, WorkstreamState
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Helpers # Helpers
@@ -130,7 +128,7 @@ class TestManagerCreation:
def test_create_second_does_not_change_active(self): def test_create_second_does_not_change_active(self):
mgr = WorkstreamManager(_fake_factory) mgr = WorkstreamManager(_fake_factory)
ws1 = mgr.create(ui_factory=lambda wid: FakeUI(wid)) ws1 = mgr.create(ui_factory=lambda wid: FakeUI(wid))
ws2 = mgr.create(ui_factory=lambda wid: FakeUI(wid)) _ws2 = mgr.create(ui_factory=lambda wid: FakeUI(wid))
assert mgr.active_id == ws1.id assert mgr.active_id == ws1.id
def test_create_assigns_session(self): def test_create_assigns_session(self):
@@ -176,9 +174,9 @@ class TestManagerLookup:
def test_list_all_creation_order(self): def test_list_all_creation_order(self):
mgr = WorkstreamManager(_fake_factory) mgr = WorkstreamManager(_fake_factory)
ws1 = mgr.create(name="a", ui_factory=lambda wid: FakeUI(wid)) _ws1 = mgr.create(name="a", ui_factory=lambda wid: FakeUI(wid))
ws2 = mgr.create(name="b", ui_factory=lambda wid: FakeUI(wid)) _ws2 = mgr.create(name="b", ui_factory=lambda wid: FakeUI(wid))
ws3 = mgr.create(name="c", ui_factory=lambda wid: FakeUI(wid)) _ws3 = mgr.create(name="c", ui_factory=lambda wid: FakeUI(wid))
result = mgr.list_all() result = mgr.list_all()
assert [w.name for w in result] == ["a", "b", "c"] assert [w.name for w in result] == ["a", "b", "c"]
@@ -222,7 +220,7 @@ class TestManagerSwitching:
def test_switch_by_index(self): def test_switch_by_index(self):
mgr = WorkstreamManager(_fake_factory) mgr = WorkstreamManager(_fake_factory)
ws1 = mgr.create(ui_factory=lambda wid: FakeUI(wid)) _ws1 = mgr.create(ui_factory=lambda wid: FakeUI(wid))
ws2 = mgr.create(ui_factory=lambda wid: FakeUI(wid)) ws2 = mgr.create(ui_factory=lambda wid: FakeUI(wid))
result = mgr.switch_by_index(2) result = mgr.switch_by_index(2)
@@ -244,7 +242,7 @@ class TestManagerSwitching:
class TestManagerClose: class TestManagerClose:
def test_close_removes_workstream(self): def test_close_removes_workstream(self):
mgr = WorkstreamManager(_fake_factory) mgr = WorkstreamManager(_fake_factory)
ws1 = mgr.create(ui_factory=lambda wid: FakeUI(wid)) _ws1 = mgr.create(ui_factory=lambda wid: FakeUI(wid))
ws2 = mgr.create(ui_factory=lambda wid: FakeUI(wid)) ws2 = mgr.create(ui_factory=lambda wid: FakeUI(wid))
assert mgr.close(ws2.id) is True assert mgr.close(ws2.id) is True
@@ -274,9 +272,9 @@ class TestManagerClose:
def test_close_updates_order(self): def test_close_updates_order(self):
mgr = WorkstreamManager(_fake_factory) mgr = WorkstreamManager(_fake_factory)
ws1 = mgr.create(name="a", ui_factory=lambda wid: FakeUI(wid)) _ws1 = mgr.create(name="a", ui_factory=lambda wid: FakeUI(wid))
ws2 = mgr.create(name="b", ui_factory=lambda wid: FakeUI(wid)) ws2 = mgr.create(name="b", ui_factory=lambda wid: FakeUI(wid))
ws3 = mgr.create(name="c", ui_factory=lambda wid: FakeUI(wid)) _ws3 = mgr.create(name="c", ui_factory=lambda wid: FakeUI(wid))
mgr.close(ws2.id) mgr.close(ws2.id)
names = [w.name for w in mgr.list_all()] names = [w.name for w in mgr.list_all()]
@@ -285,7 +283,7 @@ class TestManagerClose:
def test_close_unblocks_approval_event(self): def test_close_unblocks_approval_event(self):
"""Closing a workstream whose UI has a pending approval should unblock it.""" """Closing a workstream whose UI has a pending approval should unblock it."""
mgr = WorkstreamManager(_fake_factory) mgr = WorkstreamManager(_fake_factory)
ws1 = mgr.create(ui_factory=lambda wid: FakeUI(wid)) _ws1 = mgr.create(ui_factory=lambda wid: FakeUI(wid))
# Create a workstream with a WebUI-like approval mechanism # Create a workstream with a WebUI-like approval mechanism
from turnstone.server import WebUI from turnstone.server import WebUI
@@ -300,7 +298,7 @@ class TestManagerClose:
def test_close_unblocks_plan_event(self): def test_close_unblocks_plan_event(self):
"""Closing a workstream with pending plan review should unblock it.""" """Closing a workstream with pending plan review should unblock it."""
mgr = WorkstreamManager(_fake_factory) mgr = WorkstreamManager(_fake_factory)
ws1 = mgr.create(ui_factory=lambda wid: FakeUI(wid)) _ws1 = mgr.create(ui_factory=lambda wid: FakeUI(wid))
from turnstone.server import WebUI from turnstone.server import WebUI
@@ -551,6 +549,7 @@ class TestWebUI:
def test_on_state_change_broadcasts(self): def test_on_state_change_broadcasts(self):
"""on_state_change should put an event on the global queue.""" """on_state_change should put an event on the global queue."""
import queue import queue
from turnstone.server import WebUI from turnstone.server import WebUI
gq = queue.Queue() gq = queue.Queue()
@@ -742,6 +741,7 @@ class TestNoColor:
def test_no_color_env_disables_ansi(self): def test_no_color_env_disables_ansi(self):
import importlib import importlib
import os import os
import turnstone.ui.colors as colors_mod import turnstone.ui.colors as colors_mod
old_env = os.environ.get("NO_COLOR") old_env = os.environ.get("NO_COLOR")
+2 -2
View File
@@ -1,3 +1,3 @@
"""turnstone - Single-file AI chat client with tool use.""" """turnstone - Multi-node AI orchestration platform with tool use, agent routing, and cluster simulation."""
__version__ = "0.1.0" __version__ = "0.2.1"
+23 -29
View File
@@ -15,37 +15,31 @@ All functionality has been moved to submodules:
""" """
# Re-export public API for backward compatibility # Re-export public API for backward compatibility
from turnstone.core.session import ChatSession, SessionUI # noqa: F401 from turnstone.cli import detect_model, main # noqa: F401
from turnstone.core.tools import TOOLS, AGENT_TOOLS, TASK_AGENT_TOOLS # noqa: F401
from turnstone.core.edit import (
find_occurrences as _find_occurrences,
pick_nearest as _pick_nearest,
) # noqa: F401
from turnstone.core.sandbox import (
validate_math_code as _validate_math_code,
auto_print_wrap as _auto_print_wrap,
execute_math_sandboxed as _execute_math_sandboxed,
) # noqa: F401
from turnstone.core.safety import (
is_command_blocked,
sanitize_command as _sanitize_command,
BLOCKED_PATTERNS,
) # noqa: F401
from turnstone.core.web import strip_html as _strip_html # noqa: F401
from turnstone.core.memory import ( # noqa: F401 from turnstone.core.memory import ( # noqa: F401
open_db as _open_db, open_db as _open_db,
load_memories as _load_memories,
save_message as _save_message,
normalize_key as _normalize_key,
search_history as _search_history,
search_history_recent as _search_history_recent,
escape_like as _escape_like,
fts5_query as _fts5_query,
get_tavily_key as _get_tavily_key,
db_override as _db_override,
db_initialized as _db_initialized,
) )
from turnstone.ui.colors import * # noqa: F401, F403 from turnstone.core.session import ChatSession, SessionUI # noqa: F401
from turnstone.core.tools import AGENT_TOOLS, TASK_AGENT_TOOLS, TOOLS # noqa: F401
from turnstone.core.web import strip_html as _strip_html # noqa: F401
from turnstone.ui.colors import ( # noqa: F401
BLUE,
BOLD,
CYAN,
DIM,
GRAY,
GREEN,
ITALIC,
MAGENTA,
RED,
RESET,
YELLOW,
bold,
cyan,
dim,
green,
red,
yellow,
)
from turnstone.ui.markdown import MarkdownRenderer # noqa: F401 from turnstone.ui.markdown import MarkdownRenderer # noqa: F401
from turnstone.ui.spinner import Spinner # noqa: F401 from turnstone.ui.spinner import Spinner # noqa: F401
from turnstone.cli import main, detect_model # noqa: F401
+85 -111
View File
@@ -4,24 +4,23 @@ Provides TerminalUI (implementing the SessionUI protocol), readline setup,
model auto-detection, workstream management, and the main() REPL entry point. model auto-detection, workstream management, and the main() REPL entry point.
""" """
from __future__ import annotations
import argparse import argparse
import os import os
import readline import readline
import sys import sys
import textwrap import textwrap
import threading import threading
from typing import TYPE_CHECKING, Any
from openai import OpenAI from openai import OpenAI
from turnstone.core.session import ChatSession, SessionUI from turnstone.core.session import ChatSession, SessionUI
from turnstone.core.tools import TOOLS from turnstone.core.workstream import Workstream, WorkstreamManager, WorkstreamState
from turnstone.core.workstream import WorkstreamManager, WorkstreamState
from turnstone.ui.colors import ( from turnstone.ui.colors import (
BOLD, BOLD,
CYAN,
DIM, DIM,
GRAY,
GREEN,
RED, RED,
RESET, RESET,
YELLOW, YELLOW,
@@ -35,11 +34,12 @@ from turnstone.ui.colors import (
from turnstone.ui.markdown import MarkdownRenderer from turnstone.ui.markdown import MarkdownRenderer
from turnstone.ui.spinner import Spinner from turnstone.ui.spinner import Spinner
if TYPE_CHECKING:
from collections.abc import Callable
# ─── Readline ───────────────────────────────────────────────────────────── # ─── Readline ─────────────────────────────────────────────────────────────
SLASH_COMMANDS = [ SLASH_COMMANDS = [
"/persona",
"/instructions", "/instructions",
"/clear", "/clear",
"/new", "/new",
@@ -63,18 +63,15 @@ SLASH_COMMANDS = [
] ]
def _completer(text, state): def _completer(text: str, state: int) -> str | None:
"""Tab-complete slash commands.""" """Tab-complete slash commands."""
if text.startswith("/"): matches = [c for c in SLASH_COMMANDS if c.startswith(text)] if text.startswith("/") else []
matches = [c for c in SLASH_COMMANDS if c.startswith(text)]
else:
matches = []
if state < len(matches): if state < len(matches):
return matches[state] + " " return matches[state] + " "
return None return None
def setup_readline(): def setup_readline() -> None:
"""Set up readline with tab completion.""" """Set up readline with tab completion."""
readline.set_history_length(1000) readline.set_history_length(1000)
readline.set_completer(_completer) readline.set_completer(_completer)
@@ -88,32 +85,32 @@ def setup_readline():
class TerminalUI(SessionUI): class TerminalUI(SessionUI):
"""Terminal-based UI using ANSI colors, MarkdownRenderer, and Spinner.""" """Terminal-based UI using ANSI colors, MarkdownRenderer, and Spinner."""
def __init__(self): def __init__(self) -> None:
self.md = MarkdownRenderer() self.md = MarkdownRenderer()
self.spinner = None self.spinner: Spinner | None = None
self._print_lock = threading.Lock() self._print_lock = threading.Lock()
self.auto_approve = False self.auto_approve = False
def on_thinking_start(self): def on_thinking_start(self) -> None:
self.spinner = Spinner("Thinking") self.spinner = Spinner("Thinking")
self.spinner.start() self.spinner.start()
def on_thinking_stop(self): def on_thinking_stop(self) -> None:
if self.spinner: if self.spinner:
self.spinner.stop() self.spinner.stop()
self.spinner = None self.spinner = None
def on_reasoning_token(self, text): def on_reasoning_token(self, text: str) -> None:
sys.stdout.write(f"{DIM}{text}{RESET}") sys.stdout.write(f"{DIM}{text}{RESET}")
sys.stdout.flush() sys.stdout.flush()
def on_content_token(self, text): def on_content_token(self, text: str) -> None:
rendered = self.md.feed(text) rendered = self.md.feed(text)
if rendered: if rendered:
sys.stdout.write(rendered) sys.stdout.write(rendered)
sys.stdout.flush() sys.stdout.flush()
def on_stream_end(self): def on_stream_end(self) -> None:
remainder = self.md.flush() remainder = self.md.flush()
if remainder: if remainder:
sys.stdout.write(remainder) sys.stdout.write(remainder)
@@ -121,14 +118,12 @@ class TerminalUI(SessionUI):
sys.stdout.write("\n") sys.stdout.write("\n")
sys.stdout.flush() sys.stdout.flush()
def approve_tools(self, items): def approve_tools(self, items: list[dict[str, Any]]) -> tuple[bool, str | None]:
"""Display tool previews and prompt for batch approval. """Display tool previews and prompt for batch approval.
Returns (approved: bool, feedback: str | None). Returns (approved: bool, feedback: str | None).
""" """
pending = [ pending = [it for it in items if it.get("needs_approval") and not it.get("error")]
it for it in items if it.get("needs_approval") and not it.get("error")
]
with self._print_lock: with self._print_lock:
# Print all headers and previews # Print all headers and previews
@@ -153,9 +148,7 @@ class TerminalUI(SessionUI):
f"\001{DIM}\002[y/n/a(lways), optional message]\001{RESET}\002 " f"\001{DIM}\002[y/n/a(lways), optional message]\001{RESET}\002 "
) )
else: else:
labels = ", ".join( labels = ", ".join(it.get("approval_label", it["func_name"]) for it in pending)
it.get("approval_label", it["func_name"]) for it in pending
)
prompt_text = ( prompt_text = (
f" \001{BOLD}\002Allow {len(pending)} tools ({labels})?\001{RESET}\002 " f" \001{BOLD}\002Allow {len(pending)} tools ({labels})?\001{RESET}\002 "
f"\001{DIM}\002[y/n/a(lways), optional message]\001{RESET}\002 " f"\001{DIM}\002[y/n/a(lways), optional message]\001{RESET}\002 "
@@ -188,10 +181,10 @@ class TerminalUI(SessionUI):
item["denial_msg"] = denial_msg item["denial_msg"] = denial_msg
return False, None return False, None
def on_tool_result(self, name, output): def on_tool_result(self, name: str, output: str) -> None:
pass # Optional: display summary pass # Optional: display summary
def on_status(self, usage, context_window, effort): def on_status(self, usage: dict[str, Any], context_window: int, effort: str) -> None:
total_tok = usage["prompt_tokens"] + usage["completion_tokens"] total_tok = usage["prompt_tokens"] + usage["completion_tokens"]
pct = total_tok / context_window * 100 if context_window > 0 else 0 pct = total_tok / context_window * 100 if context_window > 0 else 0
parts = [f"{total_tok:,} / {context_window:,} tokens ({pct:.0f}%)"] parts = [f"{total_tok:,} / {context_window:,} tokens ({pct:.0f}%)"]
@@ -200,7 +193,7 @@ class TerminalUI(SessionUI):
sys.stdout.write(f"\n {DIM}[{' · '.join(parts)}]{RESET}\n") sys.stdout.write(f"\n {DIM}[{' · '.join(parts)}]{RESET}\n")
sys.stdout.flush() sys.stdout.flush()
def on_plan_review(self, content): def on_plan_review(self, content: str) -> str:
sys.stdout.write(f"\n{DIM}{'' * 60}{RESET}\n") sys.stdout.write(f"\n{DIM}{'' * 60}{RESET}\n")
for line in content.splitlines(): for line in content.splitlines():
sys.stdout.write(f" {line}\n") sys.stdout.write(f" {line}\n")
@@ -218,17 +211,17 @@ class TerminalUI(SessionUI):
resp = "reject" resp = "reject"
return resp return resp
def on_info(self, message): def on_info(self, message: str) -> None:
print(message) print(message)
def on_error(self, message): def on_error(self, message: str) -> None:
sys.stdout.write(f"{RED}{message}{RESET}\n") sys.stdout.write(f"{RED}{message}{RESET}\n")
sys.stdout.flush() sys.stdout.flush()
def on_state_change(self, state): def on_state_change(self, state: str) -> None:
pass # base TerminalUI ignores state changes pass # base TerminalUI ignores state changes
def on_rename(self, name: str): def on_rename(self, name: str) -> None:
pass # base TerminalUI ignores renames pass # base TerminalUI ignores renames
@@ -236,7 +229,7 @@ class TerminalUI(SessionUI):
# State display config: (symbol, color_fn, label) # State display config: (symbol, color_fn, label)
_STATE_DISPLAY = { _STATE_DISPLAY: dict[WorkstreamState, tuple[str, Callable[[str], str], str]] = {
WorkstreamState.IDLE: ("·", dim, "idle"), WorkstreamState.IDLE: ("·", dim, "idle"),
WorkstreamState.THINKING: ("", cyan, "thinking"), WorkstreamState.THINKING: ("", cyan, "thinking"),
WorkstreamState.RUNNING: ("", green, "running"), WorkstreamState.RUNNING: ("", green, "running"),
@@ -249,7 +242,7 @@ class WorkstreamTerminalUI(TerminalUI):
"""TerminalUI with workstream awareness: buffers output when in background, """TerminalUI with workstream awareness: buffers output when in background,
blocks on approval until foregrounded.""" blocks on approval until foregrounded."""
def __init__(self, ws_id: str, manager: WorkstreamManager): def __init__(self, ws_id: str, manager: WorkstreamManager) -> None:
super().__init__() super().__init__()
self.ws_id = ws_id self.ws_id = ws_id
self.manager = manager self.manager = manager
@@ -261,13 +254,13 @@ class WorkstreamTerminalUI(TerminalUI):
def is_foreground(self) -> bool: def is_foreground(self) -> bool:
return self.manager.active_id == self.ws_id return self.manager.active_id == self.ws_id
def set_foreground(self, fg: bool): def set_foreground(self, fg: bool) -> None:
if fg: if fg:
self._fg_event.set() self._fg_event.set()
else: else:
self._fg_event.clear() self._fg_event.clear()
def on_state_change(self, state: str): def on_state_change(self, state: str) -> None:
try: try:
ws_state = WorkstreamState(state) ws_state = WorkstreamState(state)
except ValueError: except ValueError:
@@ -276,61 +269,61 @@ class WorkstreamTerminalUI(TerminalUI):
# -- output buffering when in background -------------------------------- # -- output buffering when in background --------------------------------
def on_thinking_start(self): def on_thinking_start(self) -> None:
if self.is_foreground: if self.is_foreground:
super().on_thinking_start() super().on_thinking_start()
def on_thinking_stop(self): def on_thinking_stop(self) -> None:
if self.is_foreground: if self.is_foreground:
super().on_thinking_stop() super().on_thinking_stop()
elif self.spinner: elif self.spinner:
self.spinner.stop() self.spinner.stop()
self.spinner = None self.spinner = None
def _buffer(self, event_type: str, text: str): def _buffer(self, event_type: str, text: str) -> None:
with self._print_lock: with self._print_lock:
self._output_buffer.append((event_type, text)) self._output_buffer.append((event_type, text))
def on_reasoning_token(self, text): def on_reasoning_token(self, text: str) -> None:
if self.is_foreground: if self.is_foreground:
super().on_reasoning_token(text) super().on_reasoning_token(text)
else: else:
self._buffer("reasoning", text) self._buffer("reasoning", text)
def on_content_token(self, text): def on_content_token(self, text: str) -> None:
if self.is_foreground: if self.is_foreground:
super().on_content_token(text) super().on_content_token(text)
else: else:
self._buffer("content", text) self._buffer("content", text)
def on_stream_end(self): def on_stream_end(self) -> None:
if self.is_foreground: if self.is_foreground:
super().on_stream_end() super().on_stream_end()
else: else:
self._buffer("stream_end", "") self._buffer("stream_end", "")
def on_status(self, usage, context_window, effort): def on_status(self, usage: dict[str, Any], context_window: int, effort: str) -> None:
if self.is_foreground: if self.is_foreground:
super().on_status(usage, context_window, effort) super().on_status(usage, context_window, effort)
# silently drop status for background streams # silently drop status for background streams
def on_info(self, message): def on_info(self, message: str) -> None:
if self.is_foreground: if self.is_foreground:
super().on_info(message) super().on_info(message)
else: else:
self._buffer("info", message) self._buffer("info", message)
def on_error(self, message): def on_error(self, message: str) -> None:
if self.is_foreground: if self.is_foreground:
super().on_error(message) super().on_error(message)
else: else:
self._buffer("error", message) self._buffer("error", message)
def on_tool_result(self, name, output): def on_tool_result(self, name: str, output: str) -> None:
if self.is_foreground: if self.is_foreground:
super().on_tool_result(name, output) super().on_tool_result(name, output)
def on_plan_review(self, content): def on_plan_review(self, content: str) -> str:
# Must wait until foregrounded to show plan review # Must wait until foregrounded to show plan review
if not self.is_foreground: if not self.is_foreground:
self._buffer( self._buffer(
@@ -340,7 +333,7 @@ class WorkstreamTerminalUI(TerminalUI):
self._fg_event.wait() self._fg_event.wait()
return super().on_plan_review(content) return super().on_plan_review(content)
def approve_tools(self, items): def approve_tools(self, items: list[dict[str, Any]]) -> tuple[bool, str | None]:
"""Block until foregrounded if in background, then show approval prompt.""" """Block until foregrounded if in background, then show approval prompt."""
if not self.is_foreground: if not self.is_foreground:
tool_names = ", ".join( tool_names = ", ".join(
@@ -349,22 +342,18 @@ class WorkstreamTerminalUI(TerminalUI):
if it.get("needs_approval") and not it.get("error") if it.get("needs_approval") and not it.get("error")
) )
if tool_names: if tool_names:
self._buffer( self._buffer("info", f"{YELLOW}Waiting for approval: {tool_names}{RESET}")
"info", f"{YELLOW}Waiting for approval: {tool_names}{RESET}"
)
self._fg_event.wait() self._fg_event.wait()
return super().approve_tools(items) return super().approve_tools(items)
def flush_buffer(self): def flush_buffer(self) -> None:
"""Replay buffered output when switching to foreground.""" """Replay buffered output when switching to foreground."""
with self._print_lock: with self._print_lock:
if not self._output_buffer: if not self._output_buffer:
return return
buf = list(self._output_buffer) buf = list(self._output_buffer)
self._output_buffer.clear() self._output_buffer.clear()
sys.stdout.write( sys.stdout.write(f"\n {DIM}--- buffered output ({len(buf)} events) ---{RESET}\n")
f"\n {DIM}--- buffered output ({len(buf)} events) ---{RESET}\n"
)
replay_md = MarkdownRenderer() replay_md = MarkdownRenderer()
for event_type, text in buf: for event_type, text in buf:
if event_type == "reasoning": if event_type == "reasoning":
@@ -390,7 +379,7 @@ class WorkstreamTerminalUI(TerminalUI):
# ─── Workstream commands ────────────────────────────────────────────────── # ─── Workstream commands ──────────────────────────────────────────────────
def _print_ws_status_line(manager: WorkstreamManager): def _print_ws_status_line(manager: WorkstreamManager) -> None:
"""Print a one-line status of background workstreams that are active.""" """Print a one-line status of background workstreams that are active."""
active_id = manager.active_id active_id = manager.active_id
parts = [] parts = []
@@ -411,7 +400,7 @@ def _handle_ws_command(
manager: WorkstreamManager, manager: WorkstreamManager,
cmd_line: str, cmd_line: str,
skip_permissions: bool, skip_permissions: bool,
): ) -> bool:
"""Handle /ws subcommands. Returns (switched: bool).""" """Handle /ws subcommands. Returns (switched: bool)."""
parts = cmd_line.strip().split() parts = cmd_line.strip().split()
sub = parts[1] if len(parts) > 1 else "list" sub = parts[1] if len(parts) > 1 else "list"
@@ -438,27 +427,27 @@ def _handle_ws_command(
except RuntimeError as e: except RuntimeError as e:
print(red(str(e))) print(red(str(e)))
return False return False
if skip_permissions: if skip_permissions and isinstance(ws.ui, TerminalUI):
ws.ui.auto_approve = True ws.ui.auto_approve = True
# Mark old active as background # Mark old active as background
old = manager.get_active() old = manager.get_active()
if old and old.ui and hasattr(old.ui, "set_foreground"): if old and isinstance(old.ui, WorkstreamTerminalUI):
old.ui.set_foreground(False) old.ui.set_foreground(False)
manager.switch(ws.id) manager.switch(ws.id)
ws.ui.set_foreground(True) if isinstance(ws.ui, WorkstreamTerminalUI):
ws.ui.set_foreground(True)
print(f"Created workstream {cyan(ws.name)} (#{manager.index_of(ws.id)})") print(f"Created workstream {cyan(ws.name)} (#{manager.index_of(ws.id)})")
return True return True
elif sub.isdigit(): elif sub.isdigit():
idx = int(sub) idx = int(sub)
old = manager.get_active() old = manager.get_active()
ws = manager.switch_by_index(idx) ws: Workstream | None = manager.switch_by_index(idx) # type: ignore[no-redef]
if ws: if ws:
if old and old.ui and hasattr(old.ui, "set_foreground"): if old and isinstance(old.ui, WorkstreamTerminalUI):
old.ui.set_foreground(False) old.ui.set_foreground(False)
if hasattr(ws.ui, "set_foreground"): if isinstance(ws.ui, WorkstreamTerminalUI):
ws.ui.set_foreground(True) ws.ui.set_foreground(True)
if hasattr(ws.ui, "flush_buffer"):
ws.ui.flush_buffer() ws.ui.flush_buffer()
print(f"Switched to {cyan(ws.name)}") print(f"Switched to {cyan(ws.name)}")
return True return True
@@ -468,6 +457,7 @@ def _handle_ws_command(
elif sub == "close": elif sub == "close":
target_idx = int(parts[2]) if len(parts) > 2 and parts[2].isdigit() else None target_idx = int(parts[2]) if len(parts) > 2 and parts[2].isdigit() else None
ws_id: str | None = None
if target_idx is not None: if target_idx is not None:
all_ws = manager.list_all() all_ws = manager.list_all()
if 1 <= target_idx <= len(all_ws): if 1 <= target_idx <= len(all_ws):
@@ -477,13 +467,17 @@ def _handle_ws_command(
return False return False
else: else:
ws_id = manager.active_id ws_id = manager.active_id
if ws_id is None:
return False
ws_name = manager.get(ws_id).name if manager.get(ws_id) else "?" assert ws_id is not None
ws_obj = manager.get(ws_id)
ws_name = ws_obj.name if ws_obj else "?"
if manager.close(ws_id): if manager.close(ws_id):
print(f"Closed workstream {ws_name}") print(f"Closed workstream {ws_name}")
# Ensure new active is foregrounded # Ensure new active is foregrounded
new_active = manager.get_active() new_active = manager.get_active()
if new_active and hasattr(new_active.ui, "set_foreground"): if new_active and isinstance(new_active.ui, WorkstreamTerminalUI):
new_active.ui.set_foreground(True) new_active.ui.set_foreground(True)
return True return True
else: else:
@@ -495,34 +489,28 @@ def _handle_ws_command(
if not new_name: if not new_name:
print(red("Usage: /ws rename <name>")) print(red("Usage: /ws rename <name>"))
return False return False
ws = manager.get_active() ws_active: Workstream | None = manager.get_active()
if ws: if ws_active:
old_name = ws.name old_name = ws_active.name
ws.name = new_name ws_active.name = new_name
print(f"Renamed {old_name} -> {cyan(new_name)}") print(f"Renamed {old_name} -> {cyan(new_name)}")
return False return False
else: else:
print(f"Unknown /ws subcommand: {sub}") print(f"Unknown /ws subcommand: {sub}")
print(f"Usage: /ws [list|new [name]|<N>|close [N]|rename <name>]") print("Usage: /ws [list|new [name]|<N>|close [N]|rename <name>]")
return False return False
# ─── Cluster commands ───────────────────────────────────────────────────── # ─── Cluster commands ─────────────────────────────────────────────────────
def _handle_cluster_command( def _handle_cluster_command(cmd_line: str, console_url: str | None, auth_token: str = "") -> None:
cmd_line: str, console_url: str | None, auth_token: str = ""
):
"""Handle /cluster subcommands querying the turnstone-console API.""" """Handle /cluster subcommands querying the turnstone-console API."""
import httpx import httpx
if not console_url: if not console_url:
print( print(red("No console URL configured. Use --console-url or set [console] url in config."))
red(
"No console URL configured. Use --console-url or set [console] url in config."
)
)
return return
headers: dict[str, str] = {} headers: dict[str, str] = {}
@@ -534,9 +522,7 @@ def _handle_cluster_command(
try: try:
if sub == "status": if sub == "status":
resp = httpx.get( resp = httpx.get(f"{console_url}/api/cluster/overview", timeout=5, headers=headers)
f"{console_url}/api/cluster/overview", timeout=5, headers=headers
)
data = resp.json() data = resp.json()
states = data.get("states", {}) states = data.get("states", {})
agg = data.get("aggregate", {}) agg = data.get("aggregate", {})
@@ -584,9 +570,7 @@ def _handle_cluster_command(
print( print(
f"\n {'NODE'.ljust(max_name)} {'WS':>4} {'RUN':>4} {'ATTN':>4} {'TOKENS':>8}" f"\n {'NODE'.ljust(max_name)} {'WS':>4} {'RUN':>4} {'ATTN':>4} {'TOKENS':>8}"
) )
print( print(f" {'-' * max_name} {'----':>4} {'----':>4} {'----':>4} {'--------':>8}")
f" {'-' * max_name} {'----':>4} {'----':>4} {'----':>4} {'--------':>8}"
)
for n in nodes: for n in nodes:
name = n["node_id"].ljust(max_name) name = n["node_id"].ljust(max_name)
ws = str(n.get("ws_total", 0)) ws = str(n.get("ws_total", 0))
@@ -596,9 +580,7 @@ def _handle_cluster_command(
tok_str = f"{tok / 1000:.1f}k" if tok >= 1000 else str(tok) tok_str = f"{tok / 1000:.1f}k" if tok >= 1000 else str(tok)
run_str = green(str(run)) if run else dim("0") run_str = green(str(run)) if run else dim("0")
attn_str = yellow(str(attn)) if attn else dim("0") attn_str = yellow(str(attn)) if attn else dim("0")
print( print(f" {cyan(name)} {ws:>4} {run_str:>4} {attn_str:>4} {dim(tok_str):>8}")
f" {cyan(name)} {ws:>4} {run_str:>4} {attn_str:>4} {dim(tok_str):>8}"
)
if total > len(nodes): if total > len(nodes):
print(dim(f"\n Showing {len(nodes)} of {total} nodes")) print(dim(f"\n Showing {len(nodes)} of {total} nodes"))
print() print()
@@ -723,14 +705,13 @@ def detect_model(client: OpenAI) -> str:
# ─── Main ────────────────────────────────────────────────────────────────── # ─── Main ──────────────────────────────────────────────────────────────────
def main(): def main() -> None:
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
description="Interactive CLI for vLLM models with tool calling.", description="Interactive CLI for vLLM models with tool calling.",
formatter_class=argparse.RawDescriptionHelpFormatter, formatter_class=argparse.RawDescriptionHelpFormatter,
epilog=textwrap.dedent("""\ epilog=textwrap.dedent("""\
Examples: Examples:
python3 chat.py # auto-detect model python3 chat.py # auto-detect model
python3 chat.py --persona lawful_evil # with persona
python3 chat.py --model kappa_20b_131k # explicit model python3 chat.py --model kappa_20b_131k # explicit model
python3 chat.py --temperature 0.7 # lower temperature python3 chat.py --temperature 0.7 # lower temperature
"""), """),
@@ -745,11 +726,6 @@ def main():
default=None, default=None,
help="Model name (default: auto-detect from server)", help="Model name (default: auto-detect from server)",
) )
parser.add_argument(
"--persona",
default=None,
help="Persona name injected as system message",
)
parser.add_argument( parser.add_argument(
"--instructions", "--instructions",
default=None, default=None,
@@ -863,18 +839,15 @@ def main():
) )
# Detect or use provided model # Detect or use provided model
if args.model: model = args.model or detect_model(client)
model = args.model
else:
model = detect_model(client)
# Session factory — captures shared config for creating workstream sessions # Session factory — captures shared config for creating workstream sessions
def session_factory(ui): def session_factory(ui: SessionUI | None) -> ChatSession:
assert ui is not None, "session_factory requires a non-None UI"
return ChatSession( return ChatSession(
client=client, client=client,
model=model, model=model,
ui=ui, ui=ui,
persona=args.persona,
instructions=args.instructions, instructions=args.instructions,
temperature=args.temperature, temperature=args.temperature,
max_tokens=args.max_tokens, max_tokens=args.max_tokens,
@@ -892,7 +865,7 @@ def main():
ws = manager.create( ws = manager.create(
ui_factory=lambda wid: WorkstreamTerminalUI(wid, manager), ui_factory=lambda wid: WorkstreamTerminalUI(wid, manager),
) )
if args.skip_permissions: if args.skip_permissions and isinstance(ws.ui, TerminalUI):
ws.ui.auto_approve = True ws.ui.auto_approve = True
# Handle --resume # Handle --resume
@@ -903,15 +876,16 @@ def main():
if not target_id: if not target_id:
print(red(f"Session not found: {args.resume}")) print(red(f"Session not found: {args.resume}"))
sys.exit(1) sys.exit(1)
if ws.session is None:
print(red("No session available."))
sys.exit(1)
if not ws.session.resume_session(target_id): if not ws.session.resume_session(target_id):
print(red(f"Session '{args.resume}' has no messages.")) print(red(f"Session '{args.resume}' has no messages."))
sys.exit(1) sys.exit(1)
print( print(f"Resumed session {bold(target_id)} ({len(ws.session.messages)} messages)")
f"Resumed session {bold(target_id)} ({len(ws.session.messages)} messages)"
)
# Background attention notification — write to stderr while user types # Background attention notification — write to stderr while user types
def _bg_attention_notify(ws_id, state): def _bg_attention_notify(ws_id: str, state: WorkstreamState) -> None:
if state == WorkstreamState.ATTENTION and ws_id != manager.active_id: if state == WorkstreamState.ATTENTION and ws_id != manager.active_id:
bg_ws = manager.get(ws_id) bg_ws = manager.get(ws_id)
if bg_ws: if bg_ws:
@@ -927,9 +901,7 @@ def main():
# Print banner # Print banner
print(f"\n{bold('Chat')} with {cyan(model)}") print(f"\n{bold('Chat')} with {cyan(model)}")
if args.persona: print("Type /help for commands, /ws for workstreams, /exit or Ctrl+D to quit.\n")
print(f"Persona: {cyan(args.persona)}")
print(f"Type /help for commands, /ws for workstreams, /exit or Ctrl+D to quit.\n")
# Prompt string -- use a short display name # Prompt string -- use a short display name
display_name = model.split("/")[-1] # strip path prefixes if any display_name = model.split("/")[-1] # strip path prefixes if any
@@ -945,7 +917,7 @@ def main():
# Build prompt with workstream info # Build prompt with workstream info
active = manager.get_active() active = manager.get_active()
if manager.count > 1: if manager.count > 1 and active is not None:
idx = manager.index_of(active.id) idx = manager.index_of(active.id)
prompt_str = f"\001{BOLD}\002{idx}:{active.name}\001{RESET}\002 > " prompt_str = f"\001{BOLD}\002{idx}:{active.name}\001{RESET}\002 > "
else: else:
@@ -968,6 +940,8 @@ def main():
continue continue
active = manager.get_active() active = manager.get_active()
if active is None or active.session is None:
continue
if user_input.startswith("/"): if user_input.startswith("/"):
should_exit = active.session.handle_command(user_input) should_exit = active.session.handle_command(user_input)
if should_exit: if should_exit:
+22 -26
View File
@@ -7,6 +7,7 @@ for real-time state changes.
from __future__ import annotations from __future__ import annotations
import contextlib
import json import json
import logging import logging
import queue import queue
@@ -14,10 +15,12 @@ import threading
import time import time
from concurrent.futures import ThreadPoolExecutor, as_completed from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
import httpx import httpx
from turnstone.mq.broker import RedisBroker if TYPE_CHECKING:
from turnstone.mq.broker import RedisBroker
log = logging.getLogger("turnstone.console.collector") log = logging.getLogger("turnstone.console.collector")
@@ -31,9 +34,9 @@ class NodeSnapshot:
started: float = 0.0 started: float = 0.0
last_seen: float = 0.0 # monotonic time of last successful poll last_seen: float = 0.0 # monotonic time of last successful poll
max_ws: int = 10 # max workstreams (capacity) max_ws: int = 10 # max workstreams (capacity)
workstreams: dict[str, dict] = field(default_factory=dict) workstreams: dict[str, dict[str, Any]] = field(default_factory=dict)
health: dict = field(default_factory=dict) health: dict[str, Any] = field(default_factory=dict)
aggregate: dict = field(default_factory=dict) aggregate: dict[str, Any] = field(default_factory=dict)
reachable: bool = True reachable: bool = True
@@ -74,7 +77,7 @@ class ClusterCollector:
self._http_client = httpx.Client(timeout=http_timeout, headers=headers) self._http_client = httpx.Client(timeout=http_timeout, headers=headers)
# SSE fan-out to browser clients # SSE fan-out to browser clients
self._listeners: list[queue.Queue] = [] self._listeners: list[queue.Queue[dict[str, Any]]] = []
self._listeners_lock = threading.Lock() self._listeners_lock = threading.Lock()
# -- lifecycle ----------------------------------------------------------- # -- lifecycle -----------------------------------------------------------
@@ -165,14 +168,12 @@ class ClusterCollector:
# Fan out to SSE listeners # Fan out to SSE listeners
self._fanout(data) self._fanout(data)
def _fanout(self, event: dict) -> None: def _fanout(self, event: dict[str, Any]) -> None:
"""Copy an event to all registered SSE listener queues.""" """Copy an event to all registered SSE listener queues."""
with self._listeners_lock: with self._listeners_lock:
for q in self._listeners: for q in self._listeners:
try: with contextlib.suppress(queue.Full):
q.put_nowait(event) q.put_nowait(event)
except queue.Full:
pass
# -- node discovery ------------------------------------------------------ # -- node discovery ------------------------------------------------------
@@ -242,10 +243,7 @@ class ClusterCollector:
if not targets: if not targets:
return return
futures = { futures = {self._poll_pool.submit(self._fetch_node, nid, url): nid for nid, url in targets}
self._poll_pool.submit(self._fetch_node, nid, url): nid
for nid, url in targets
}
for future in as_completed(futures): for future in as_completed(futures):
nid = futures[future] nid = futures[future]
try: try:
@@ -257,19 +255,19 @@ class ClusterCollector:
if nid in self._nodes: if nid in self._nodes:
self._nodes[nid].reachable = False self._nodes[nid].reachable = False
def _fetch_node(self, node_id: str, server_url: str) -> tuple[dict, dict]: def _fetch_node(self, node_id: str, server_url: str) -> tuple[dict[str, Any], dict[str, Any]]:
"""Fetch /api/dashboard and /health from a single node.""" """Fetch /api/dashboard and /health from a single node."""
base = server_url.rstrip("/") base = server_url.rstrip("/")
dash_resp = self._http_client.get(f"{base}/api/dashboard") dash_resp = self._http_client.get(f"{base}/api/dashboard")
dash_data = dash_resp.json() dash_data: dict[str, Any] = dash_resp.json()
try: try:
health_resp = self._http_client.get(f"{base}/health") health_resp = self._http_client.get(f"{base}/health")
health_data = health_resp.json() health_data: dict[str, Any] = health_resp.json()
except Exception: except Exception:
health_data = {} health_data = {}
return dash_data, health_data return dash_data, health_data
def _apply_poll(self, node_id: str, dashboard: dict, health: dict) -> None: def _apply_poll(self, node_id: str, dashboard: dict[str, Any], health: dict[str, Any]) -> None:
"""Apply polled data to the in-memory node snapshot.""" """Apply polled data to the in-memory node snapshot."""
ws_list = dashboard.get("workstreams", []) ws_list = dashboard.get("workstreams", [])
aggregate = dashboard.get("aggregate", {}) aggregate = dashboard.get("aggregate", {})
@@ -289,7 +287,7 @@ class ClusterCollector:
# -- query methods (thread-safe) ----------------------------------------- # -- query methods (thread-safe) -----------------------------------------
def get_overview(self) -> dict: def get_overview(self) -> dict[str, Any]:
"""Return cluster overview: state counts, totals, aggregate stats.""" """Return cluster overview: state counts, totals, aggregate stats."""
states = {"running": 0, "thinking": 0, "attention": 0, "idle": 0, "error": 0} states = {"running": 0, "thinking": 0, "attention": 0, "idle": 0, "error": 0}
total_tokens = 0 total_tokens = 0
@@ -316,7 +314,7 @@ class ClusterCollector:
def get_nodes( def get_nodes(
self, sort_by: str = "activity", limit: int = 100, offset: int = 0 self, sort_by: str = "activity", limit: int = 100, offset: int = 0
) -> tuple[list[dict], int]: ) -> tuple[list[dict[str, Any]], int]:
"""Return sorted, paginated node list with per-node counts.""" """Return sorted, paginated node list with per-node counts."""
with self._lock: with self._lock:
items = [] items = []
@@ -334,9 +332,7 @@ class ClusterCollector:
# Use aggregate tokens if available, else sum from workstreams # Use aggregate tokens if available, else sum from workstreams
agg_tokens = node.aggregate.get("total_tokens", 0) agg_tokens = node.aggregate.get("total_tokens", 0)
if not agg_tokens: if not agg_tokens:
agg_tokens = sum( agg_tokens = sum(ws.get("tokens", 0) for ws in node.workstreams.values())
ws.get("tokens", 0) for ws in node.workstreams.values()
)
items.append( items.append(
{ {
"node_id": node.node_id, "node_id": node.node_id,
@@ -376,7 +372,7 @@ class ClusterCollector:
sort_by: str = "state", sort_by: str = "state",
page: int = 1, page: int = 1,
per_page: int = 50, per_page: int = 50,
) -> tuple[list[dict], int]: ) -> tuple[list[dict[str, Any]], int]:
"""Return filtered, sorted, paginated workstreams + total count.""" """Return filtered, sorted, paginated workstreams + total count."""
with self._lock: with self._lock:
all_ws = [] all_ws = []
@@ -419,7 +415,7 @@ class ClusterCollector:
page_ws = all_ws[start : start + per_page] page_ws = all_ws[start : start + per_page]
return page_ws, total return page_ws, total
def get_node_detail(self, node_id: str) -> dict | None: def get_node_detail(self, node_id: str) -> dict[str, Any] | None:
"""Return a single node's workstreams and health.""" """Return a single node's workstreams and health."""
with self._lock: with self._lock:
node = self._nodes.get(node_id) node = self._nodes.get(node_id)
@@ -436,12 +432,12 @@ class ClusterCollector:
# -- SSE listener management --------------------------------------------- # -- SSE listener management ---------------------------------------------
def register_listener(self, q: queue.Queue) -> None: def register_listener(self, q: queue.Queue[dict[str, Any]]) -> None:
"""Register a queue for SSE event fan-out.""" """Register a queue for SSE event fan-out."""
with self._listeners_lock: with self._listeners_lock:
self._listeners.append(q) self._listeners.append(q)
def unregister_listener(self, q: queue.Queue) -> None: def unregister_listener(self, q: queue.Queue[dict[str, Any]]) -> None:
"""Unregister a queue from SSE event fan-out.""" """Unregister a queue from SSE event fan-out."""
with self._listeners_lock: with self._listeners_lock:
if q in self._listeners: if q in self._listeners:
+22 -27
View File
@@ -4,6 +4,8 @@ Serves the cluster-level dashboard UI and provides REST/SSE APIs
backed by the ClusterCollector. backed by the ClusterCollector.
""" """
from __future__ import annotations
import argparse import argparse
import json import json
import logging import logging
@@ -11,12 +13,11 @@ import math
import os import os
import queue import queue
import textwrap import textwrap
import threading from http.server import BaseHTTPRequestHandler, HTTPServer
import time
from http.server import HTTPServer, BaseHTTPRequestHandler
from pathlib import Path from pathlib import Path
from socketserver import ThreadingMixIn from socketserver import ThreadingMixIn
from urllib.parse import urlparse, parse_qs from typing import Any
from urllib.parse import ParseResult, parse_qs, urlparse
from turnstone.console.collector import ClusterCollector from turnstone.console.collector import ClusterCollector
from turnstone.mq.broker import RedisBroker from turnstone.mq.broker import RedisBroker
@@ -48,17 +49,17 @@ def _load_static() -> None:
class ConsoleHTTPHandler(BaseHTTPRequestHandler): class ConsoleHTTPHandler(BaseHTTPRequestHandler):
"""HTTP handler for the cluster dashboard.""" """HTTP handler for the cluster dashboard."""
def log_message(self, format, *args): def log_message(self, fmt: str, *args: object) -> None: # noqa: N802
pass # suppress default logging pass # suppress default logging
def _set_headers(self, status=200, content_type="application/json"): def _set_headers(self, status: int = 200, content_type: str = "application/json") -> None:
self.send_response(status) self.send_response(status)
self.send_header("Content-Type", content_type) self.send_header("Content-Type", content_type)
self.send_header("Cache-Control", "no-cache") self.send_header("Cache-Control", "no-cache")
self.send_header("Access-Control-Allow-Origin", "*") self.send_header("Access-Control-Allow-Origin", "*")
self.end_headers() self.end_headers()
def _send_json(self, data: dict, status=200): def _send_json(self, data: dict[str, Any], status: int = 200) -> None:
self._set_headers(status, "application/json") self._set_headers(status, "application/json")
self.wfile.write(json.dumps(data).encode("utf-8")) self.wfile.write(json.dumps(data).encode("utf-8"))
@@ -69,24 +70,22 @@ class ConsoleHTTPHandler(BaseHTTPRequestHandler):
auth_config = self.server.auth_config # type: ignore[attr-defined] auth_config = self.server.auth_config # type: ignore[attr-defined]
auth_header = self.headers.get("Authorization") auth_header = self.headers.get("Authorization")
cookie_header = self.headers.get("Cookie") cookie_header = self.headers.get("Cookie")
allowed, status, msg = check_request( allowed, status, msg = check_request(auth_config, method, path, auth_header, cookie_header)
auth_config, method, path, auth_header, cookie_header
)
if not allowed: if not allowed:
self._send_json({"error": msg}, status) self._send_json({"error": msg}, status)
return allowed return allowed
def _read_body(self) -> dict: def _read_body(self) -> dict[str, Any]:
length = int(self.headers.get("Content-Length", 0)) length = int(self.headers.get("Content-Length", 0))
if length == 0: if length == 0:
return {} return {}
raw = self.rfile.read(length) raw = self.rfile.read(length)
try: try:
return json.loads(raw.decode("utf-8")) return json.loads(raw.decode("utf-8")) # type: ignore[no-any-return]
except (json.JSONDecodeError, UnicodeDecodeError, ValueError): except (json.JSONDecodeError, UnicodeDecodeError, ValueError):
return {} return {}
def do_POST(self): def do_POST(self) -> None:
# Login/logout pass through _check_auth because they are in PUBLIC_PATHS. # Login/logout pass through _check_auth because they are in PUBLIC_PATHS.
if not self._check_auth("POST", self.path): if not self._check_auth("POST", self.path):
return return
@@ -103,9 +102,7 @@ class ConsoleHTTPHandler(BaseHTTPRequestHandler):
self.send_header("Set-Cookie", make_set_cookie(token)) self.send_header("Set-Cookie", make_set_cookie(token))
self.send_header("Cache-Control", "no-cache") self.send_header("Cache-Control", "no-cache")
self.end_headers() self.end_headers()
self.wfile.write( self.wfile.write(json.dumps({"status": "ok", "role": role}).encode("utf-8"))
json.dumps({"status": "ok", "role": role}).encode("utf-8")
)
else: else:
self._send_json({"error": "Invalid token"}, 401) self._send_json({"error": "Invalid token"}, 401)
@@ -122,7 +119,7 @@ class ConsoleHTTPHandler(BaseHTTPRequestHandler):
else: else:
self._send_json({"error": "Not found"}, 404) self._send_json({"error": "Not found"}, 404)
def do_GET(self): def do_GET(self) -> None:
parsed = urlparse(self.path) parsed = urlparse(self.path)
try: try:
if not self._check_auth("GET", parsed.path): if not self._check_auth("GET", parsed.path):
@@ -134,7 +131,7 @@ class ConsoleHTTPHandler(BaseHTTPRequestHandler):
@staticmethod @staticmethod
def _parse_int( def _parse_int(
qs: dict, name: str, default: int, minimum: int = 0, maximum: int = 10000 qs: dict[str, list[str]], name: str, default: int, minimum: int = 0, maximum: int = 10000
) -> int: ) -> int:
try: try:
val = int(qs.get(name, [str(default)])[0]) val = int(qs.get(name, [str(default)])[0])
@@ -142,7 +139,7 @@ class ConsoleHTTPHandler(BaseHTTPRequestHandler):
val = default val = default
return max(minimum, min(val, maximum)) return max(minimum, min(val, maximum))
def _do_GET(self, parsed): def _do_GET(self, parsed: ParseResult) -> None: # noqa: N802
collector: ClusterCollector = self.server.collector # type: ignore[attr-defined] collector: ClusterCollector = self.server.collector # type: ignore[attr-defined]
if parsed.path == "/": if parsed.path == "/":
@@ -171,9 +168,7 @@ class ConsoleHTTPHandler(BaseHTTPRequestHandler):
sort_by = qs.get("sort", ["activity"])[0] sort_by = qs.get("sort", ["activity"])[0]
limit = self._parse_int(qs, "limit", 100, minimum=1, maximum=1000) limit = self._parse_int(qs, "limit", 100, minimum=1, maximum=1000)
offset = self._parse_int(qs, "offset", 0) offset = self._parse_int(qs, "offset", 0)
nodes, total = collector.get_nodes( nodes, total = collector.get_nodes(sort_by=sort_by, limit=limit, offset=offset)
sort_by=sort_by, limit=limit, offset=offset
)
self._send_json({"nodes": nodes, "total": total}) self._send_json({"nodes": nodes, "total": total})
elif parsed.path == "/api/cluster/workstreams": elif parsed.path == "/api/cluster/workstreams":
@@ -232,7 +227,7 @@ class ConsoleHTTPHandler(BaseHTTPRequestHandler):
self._set_headers(404, "text/plain") self._set_headers(404, "text/plain")
self.wfile.write(b"Not found") self.wfile.write(b"Not found")
def _handle_sse(self, collector: ClusterCollector): def _handle_sse(self, collector: ClusterCollector) -> None:
"""Server-Sent Events stream for cluster updates.""" """Server-Sent Events stream for cluster updates."""
self.send_response(200) self.send_response(200)
self.send_header("Content-Type", "text/event-stream") self.send_header("Content-Type", "text/event-stream")
@@ -241,14 +236,14 @@ class ConsoleHTTPHandler(BaseHTTPRequestHandler):
self.send_header("Access-Control-Allow-Origin", "*") self.send_header("Access-Control-Allow-Origin", "*")
self.end_headers() self.end_headers()
client_queue: queue.Queue = queue.Queue(maxsize=500) client_queue: queue.Queue[dict[str, Any]] = queue.Queue(maxsize=500)
collector.register_listener(client_queue) collector.register_listener(client_queue)
try: try:
while True: while True:
try: try:
event = client_queue.get(timeout=5) event = client_queue.get(timeout=5)
data = json.dumps(event) data = json.dumps(event)
self.wfile.write(f"data: {data}\n\n".encode("utf-8")) self.wfile.write(f"data: {data}\n\n".encode())
self.wfile.flush() self.wfile.flush()
except queue.Empty: except queue.Empty:
self.wfile.write(b": keepalive\n\n") self.wfile.write(b": keepalive\n\n")
@@ -258,7 +253,7 @@ class ConsoleHTTPHandler(BaseHTTPRequestHandler):
finally: finally:
collector.unregister_listener(client_queue) collector.unregister_listener(client_queue)
def do_OPTIONS(self): def do_OPTIONS(self) -> None:
"""Handle CORS preflight.""" """Handle CORS preflight."""
self.send_response(200) self.send_response(200)
self.send_header("Access-Control-Allow-Origin", "*") self.send_header("Access-Control-Allow-Origin", "*")
@@ -281,7 +276,7 @@ class ThreadedHTTPServer(ThreadingMixIn, HTTPServer):
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def main(): def main() -> None:
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
description="turnstone console — cluster dashboard service.", description="turnstone console — cluster dashboard service.",
formatter_class=argparse.RawDescriptionHelpFormatter, formatter_class=argparse.RawDescriptionHelpFormatter,
+375 -117
View File
@@ -38,6 +38,9 @@ function authFetch(url, opts) {
var currentView = "overview"; // "overview" | "node" | "filtered" var currentView = "overview"; // "overview" | "node" | "filtered"
var currentNodeId = null; var currentNodeId = null;
var currentFilter = { state: null, node: null, page: 1, per_page: 50 }; var currentFilter = { state: null, node: null, page: 1, per_page: 50 };
var expandedGroups = {};
var _lastOverviewJson = "";
var _lastNodesJson = "";
var evtSource = null; var evtSource = null;
var retryDelay = 1000; var retryDelay = 1000;
@@ -95,6 +98,8 @@ function connectSSE() {
retryDelay = 1000; retryDelay = 1000;
statusBar.classList.remove("disconnected"); statusBar.classList.remove("disconnected");
statusBar.textContent = ""; statusBar.textContent = "";
var csb = document.getElementById("cluster-status-bar");
if (csb) csb.classList.remove("stale");
try { try {
var data = JSON.parse(e.data); var data = JSON.parse(e.data);
handleClusterEvent(data); handleClusterEvent(data);
@@ -107,6 +112,8 @@ function connectSSE() {
evtSource = null; evtSource = null;
statusBar.textContent = "Reconnecting\u2026"; statusBar.textContent = "Reconnecting\u2026";
statusBar.classList.add("disconnected"); statusBar.classList.add("disconnected");
var csb = document.getElementById("cluster-status-bar");
if (csb) csb.classList.add("stale");
// Raw fetch (not authFetch) — need to inspect status before throwing // Raw fetch (not authFetch) — need to inspect status before throwing
fetch("/api/cluster/overview") fetch("/api/cluster/overview")
.then(function (r) { .then(function (r) {
@@ -167,16 +174,15 @@ function loadOverview() {
var overviewP = authFetch("/api/cluster/overview").then(function (r) { var overviewP = authFetch("/api/cluster/overview").then(function (r) {
return r.json(); return r.json();
}); });
var nodesP = authFetch("/api/cluster/nodes?sort=activity&limit=50").then( var nodesP = authFetch("/api/cluster/nodes?sort=activity&limit=1000").then(
function (r) { function (r) {
return r.json(); return r.json();
}, },
); );
Promise.all([overviewP, nodesP]) Promise.all([overviewP, nodesP])
.then(function (res) { .then(function (res) {
renderStateCards(res[0].states); renderStatusBar(res[0]);
renderAggregateBar(res[0]); renderNodeGroups(res[1].nodes, res[1].total);
renderNodeTable(res[1].nodes, res[1].total);
document.getElementById("cluster-summary").textContent = document.getElementById("cluster-summary").textContent =
res[0].nodes + res[0].nodes +
" nodes \u00b7 " + " nodes \u00b7 " +
@@ -189,87 +195,307 @@ function loadOverview() {
}); });
} }
function renderStateCards(states) { // --- Status Bar ---
var container = document.getElementById("state-cards"); function renderStatusBar(overview) {
container.innerHTML = ""; var cacheKey =
JSON.stringify(overview) +
"|" +
currentView +
"|" +
(currentFilter.state || "");
if (cacheKey === _lastOverviewJson) return;
_lastOverviewJson = cacheKey;
var states = overview.states || {};
var agg = overview.aggregate || {};
var statesContainer = document.getElementById("csb-states");
statesContainer.innerHTML = "";
STATE_ORDER.forEach(function (state) { STATE_ORDER.forEach(function (state) {
var count = states[state] || 0; var count = states[state] || 0;
var sd = STATE_DISPLAY[state] || STATE_DISPLAY.idle; var sd = STATE_DISPLAY[state] || STATE_DISPLAY.idle;
var card = document.createElement("div"); var pill = document.createElement("button");
card.className = "state-card"; pill.className = "csb-state";
card.dataset.state = state; if (currentView === "filtered" && currentFilter.state === state) {
card.setAttribute("role", "button"); pill.classList.add("active");
card.setAttribute("tabindex", "0"); }
card.setAttribute("aria-label", sd.label + ": " + count + " workstreams"); pill.setAttribute("aria-label", sd.label + ": " + count + " workstreams");
card.innerHTML = pill.innerHTML =
'<div class="state-card-count">' + '<span class="csb-state-dot" data-state="' +
escapeHtml(state) +
'" aria-hidden="true"></span>' +
'<span class="csb-state-count' +
(count === 0 ? " zero" : "") +
'">' +
formatCount(count) + formatCount(count) +
"</div>" + "</span>" +
'<div class="state-card-label">' + '<span class="csb-state-label">' +
sd.symbol +
" " +
sd.label + sd.label +
"</div>"; "</span>";
card.onclick = function () { pill.onclick = function () {
drillDownByState(state); drillDownByState(state);
}; };
card.onkeydown = function (e) { statesContainer.appendChild(pill);
if (e.key === "Enter" || e.key === " ") { });
e.preventDefault();
drillDownByState(state); var metricsContainer = document.getElementById("csb-metrics");
} metricsContainer.innerHTML = "";
}; var metrics = [
container.appendChild(card); { value: overview.nodes || 0, label: "nodes", format: formatCount },
{ value: overview.workstreams || 0, label: "ws", format: formatCount },
{ value: agg.total_tokens || 0, label: "tokens", format: formatTokens },
{ value: agg.total_tool_calls || 0, label: "calls", format: formatCount },
];
metrics.forEach(function (m) {
if (m.value === 0 && m.label !== "nodes" && m.label !== "ws") return;
var el = document.createElement("span");
el.className = "csb-metric";
var valSpan = document.createElement("span");
valSpan.className = "csb-metric-value";
valSpan.textContent = m.format(m.value);
var labelSpan = document.createElement("span");
labelSpan.className = "csb-metric-label";
labelSpan.textContent = m.label;
el.appendChild(valSpan);
el.appendChild(labelSpan);
metricsContainer.appendChild(el);
}); });
} }
function renderAggregateBar(overview) { // --- Node Grouping ---
var agg = overview.aggregate || {}; function extractNodePrefix(nodeId) {
var parts = []; var stripped = nodeId.replace(/[-_][a-z0-9]*\d[a-z0-9]*$/i, "");
if (agg.total_tokens) parts.push(formatTokens(agg.total_tokens) + " tokens"); if (!stripped || stripped === nodeId) {
if (agg.total_tool_calls) stripped = nodeId.replace(/[-_]?\d+$/, "");
parts.push(formatCount(agg.total_tool_calls) + " tool calls"); }
document.getElementById("aggregate-bar").textContent = parts.join(" \u00b7 "); // Clean trailing separators (e.g., FQDN-style "node.prod.01" → "node.prod")
stripped = stripped.replace(/[-_.]$/, "");
return stripped || nodeId;
} }
function renderNodeTable(nodes, total) { function groupNodes(nodes) {
var groupMap = {};
var groupOrder = [];
nodes.forEach(function (node) {
var prefix = extractNodePrefix(node.node_id);
if (!groupMap[prefix]) {
groupMap[prefix] = {
prefix: prefix,
nodes: [],
ws_total: 0,
ws_running: 0,
ws_thinking: 0,
ws_attention: 0,
ws_error: 0,
ws_idle: 0,
total_tokens: 0,
all_reachable: true,
};
groupOrder.push(prefix);
}
var g = groupMap[prefix];
g.nodes.push(node);
g.ws_total += node.ws_total || 0;
g.ws_running += node.ws_running || 0;
g.ws_thinking += node.ws_thinking || 0;
g.ws_attention += node.ws_attention || 0;
g.ws_error += node.ws_error || 0;
g.ws_idle += node.ws_idle || 0;
g.total_tokens += node.total_tokens || 0;
if (!node.reachable) g.all_reachable = false;
});
groupOrder.forEach(function (prefix) {
groupMap[prefix].nodes.sort(function (a, b) {
return b.ws_running + b.ws_attention - (a.ws_running + a.ws_attention);
});
});
var groups = groupOrder.map(function (p) {
return groupMap[p];
});
groups.sort(function (a, b) {
var aAct = a.ws_running + a.ws_attention;
var bAct = b.ws_running + b.ws_attention;
if (bAct !== aAct) return bAct - aAct;
return a.prefix.localeCompare(b.prefix);
});
return groups;
}
function buildNodeRow(node) {
var row = document.createElement("div");
row.className = "node-row";
if (node.ws_attention > 0) row.classList.add("has-attention");
else if (node.ws_running > 0) row.classList.add("has-running");
else if (node.ws_thinking > 0) row.classList.add("has-thinking");
else if (node.ws_error > 0) row.classList.add("has-error");
row.setAttribute("role", "button");
row.setAttribute("tabindex", "0");
row.setAttribute(
"aria-label",
node.node_id +
": " +
node.ws_total +
" workstreams, " +
node.ws_running +
" running, " +
node.ws_attention +
" attention, " +
formatTokens(node.total_tokens) +
" tokens",
);
var dotClass = node.reachable ? "node-dot" : "node-dot unreachable";
var displayTokens = node.total_tokens || node.ws_tokens || 0;
var maxWs = node.max_ws || 10;
var healthPct =
maxWs > 0 ? Math.min(Math.round((node.ws_total / maxWs) * 100), 100) : 0;
var healthFillClass =
healthPct < 50 ? "low" : healthPct < 80 ? "mid" : "high";
var healthFillHtml =
healthPct > 0
? '<span class="health-bar-fill ' +
healthFillClass +
'" style="width:' +
healthPct +
'%"></span>'
: "";
row.innerHTML =
'<span class="node-cell node-cell-name"><span class="' +
dotClass +
'"></span>' +
escapeHtml(node.node_id) +
"</span>" +
'<span class="node-cell node-cell-num' +
(node.ws_total > 0 ? " has-value" : "") +
'">' +
node.ws_total +
"</span>" +
'<span class="node-cell node-cell-num' +
(node.ws_running > 0 ? " has-value" : "") +
'">' +
node.ws_running +
"</span>" +
'<span class="node-cell node-cell-num' +
(node.ws_attention > 0 ? " has-value" : "") +
'">' +
node.ws_attention +
"</span>" +
'<span class="node-cell node-cell-num">' +
formatTokens(displayTokens) +
"</span>" +
'<span class="node-cell node-cell-health"><span class="health-bar">' +
healthFillHtml +
"</span> " +
healthPct +
"%</span>";
row.onclick = function () {
drillDownToNode(node.node_id, node.server_url);
};
row.onkeydown = function (e) {
if (e.key === "Enter" || e.key === " ") {
e.preventDefault();
drillDownToNode(node.node_id, node.server_url);
}
};
return row;
}
function toggleGroup(prefix) {
expandedGroups[prefix] = !expandedGroups[prefix];
var body = document.querySelector(
'.node-group-body[data-prefix="' + prefix.replace(/"/g, '\\"') + '"]',
);
if (!body) return;
var isExpanded = expandedGroups[prefix];
if (isExpanded) body.classList.remove("collapsed");
else body.classList.add("collapsed");
var groupEl = body.parentElement;
if (groupEl) groupEl.setAttribute("aria-expanded", String(isExpanded));
var chevron = groupEl ? groupEl.querySelector(".node-group-chevron") : null;
if (chevron) {
if (isExpanded) chevron.classList.add("expanded");
else chevron.classList.remove("expanded");
}
}
function renderNodeGroups(nodes, total) {
var json = JSON.stringify(nodes);
if (json === _lastNodesJson) return;
_lastNodesJson = json;
var table = document.getElementById("node-table"); var table = document.getElementById("node-table");
table.innerHTML = ""; table.innerHTML = "";
if (!nodes.length) { if (!nodes.length) {
table.innerHTML = '<div class="dashboard-empty">No nodes discovered</div>'; table.innerHTML = '<div class="dashboard-empty">No nodes discovered</div>';
return; return;
} }
nodes.forEach(function (node) {
var row = document.createElement("div"); var topHeaders = document.createElement("div");
row.className = "node-row"; topHeaders.className = "node-colheaders";
if (node.ws_attention > 0) row.classList.add("has-attention"); topHeaders.setAttribute("aria-hidden", "true");
else if (node.ws_running > 0) row.classList.add("has-running"); topHeaders.innerHTML =
else if (node.ws_thinking > 0) row.classList.add("has-thinking"); '<span class="ncol ncol-node">NODE</span>' +
else if (node.ws_error > 0) row.classList.add("has-error"); '<span class="ncol ncol-ws">WS</span>' +
row.setAttribute("role", "button"); '<span class="ncol ncol-run">RUN</span>' +
row.setAttribute("tabindex", "0"); '<span class="ncol ncol-attn">ATTN</span>' +
row.setAttribute( '<span class="ncol ncol-tokens">TOKENS</span>' +
'<span class="ncol ncol-health">LOAD</span>';
table.appendChild(topHeaders);
var groups = groupNodes(nodes);
groups.forEach(function (group) {
// Single-node group — render as plain row
if (group.nodes.length === 1) {
var wrapper = document.createElement("div");
wrapper.className = "node-group node-group-single";
wrapper.appendChild(buildNodeRow(group.nodes[0]));
table.appendChild(wrapper);
return;
}
var groupEl = document.createElement("div");
groupEl.className = "node-group";
var isExpanded = !!expandedGroups[group.prefix];
groupEl.setAttribute("role", "listitem");
groupEl.setAttribute("aria-expanded", String(isExpanded));
// Group header
var header = document.createElement("div");
header.className = "node-group-header";
if (group.ws_attention > 0) header.classList.add("has-attention");
else if (group.ws_running > 0) header.classList.add("has-running");
else if (group.ws_thinking > 0) header.classList.add("has-thinking");
else if (group.ws_error > 0) header.classList.add("has-error");
header.setAttribute("role", "button");
header.setAttribute("tabindex", "0");
header.setAttribute(
"aria-label", "aria-label",
node.node_id + group.prefix +
": " + " group: " +
node.ws_total + group.nodes.length +
" nodes, " +
group.ws_total +
" workstreams, " + " workstreams, " +
node.ws_running + group.ws_running +
" running, " + " running, " +
node.ws_attention + group.ws_attention +
" attention, " + " attention, " +
formatTokens(node.total_tokens) + formatTokens(group.total_tokens) +
" tokens", " tokens",
); );
var dotClass = node.reachable ? "node-dot" : "node-dot unreachable"; var chevronClass = "node-group-chevron" + (isExpanded ? " expanded" : "");
var totalMaxWs = 0;
// Use aggregate tokens, fall back to summed workstream tokens group.nodes.forEach(function (n) {
var displayTokens = node.total_tokens || node.ws_tokens || 0; totalMaxWs += n.max_ws || 10;
});
// Load = workstream count / max capacity var healthPct =
var maxWs = node.max_ws || 10; totalMaxWs > 0
var healthPct = Math.round((node.ws_total / maxWs) * 100); ? Math.min(Math.round((group.ws_total / totalMaxWs) * 100), 100)
: 0;
var healthFillClass = var healthFillClass =
healthPct < 50 ? "low" : healthPct < 80 ? "mid" : "high"; healthPct < 50 ? "low" : healthPct < 80 ? "mid" : "high";
var healthFillHtml = var healthFillHtml =
@@ -281,57 +507,76 @@ function renderNodeTable(nodes, total) {
'%"></span>' '%"></span>'
: ""; : "";
row.innerHTML = header.innerHTML =
'<span class="node-cell node-cell-name"><span class="' + '<span class="node-group-name">' +
dotClass + '<span class="' +
'"></span>' + chevronClass +
escapeHtml(node.node_id) + '" aria-hidden="true">&#x25b8;</span>' +
escapeHtml(group.prefix) +
'<span class="node-group-badge">' +
group.nodes.length +
" nodes</span>" +
"</span>" + "</span>" +
'<span class="node-cell node-cell-num' + '<span class="node-group-cell num' +
(node.ws_total > 0 ? " has-value" : "") + (group.ws_total > 0 ? " has-value" : "") +
'">' + '">' +
node.ws_total + group.ws_total +
"</span>" + "</span>" +
'<span class="node-cell node-cell-num' + '<span class="node-group-cell num' +
(node.ws_running > 0 ? " has-value" : "") + (group.ws_running > 0 ? " has-value" : "") +
'">' + '">' +
node.ws_running + group.ws_running +
"</span>" + "</span>" +
'<span class="node-cell node-cell-num' + '<span class="node-group-cell num' +
(node.ws_attention > 0 ? " has-value" : "") + (group.ws_attention > 0 ? " has-value" : "") +
'">' + '">' +
node.ws_attention + group.ws_attention +
"</span>" + "</span>" +
'<span class="node-cell node-cell-num">' + '<span class="node-group-cell num">' +
formatTokens(displayTokens) + formatTokens(group.total_tokens) +
"</span>" + "</span>" +
'<span class="node-cell node-cell-health">' + '<span class="node-group-cell node-cell-health"><span class="health-bar">' +
'<span class="health-bar">' +
healthFillHtml + healthFillHtml +
"</span>" + "</span> " +
" " +
healthPct + healthPct +
"%" + "%</span>";
"</span>";
row.onclick = function () { var prefix = group.prefix;
drillDownToNode(node.node_id, node.server_url); header.onclick = function () {
toggleGroup(prefix);
}; };
row.onkeydown = function (e) { header.onkeydown = function (e) {
if (e.key === "Enter" || e.key === " ") { if (e.key === "Enter" || e.key === " ") {
e.preventDefault(); e.preventDefault();
drillDownToNode(node.node_id, node.server_url); toggleGroup(prefix);
} }
}; };
table.appendChild(row); groupEl.appendChild(header);
});
// Pagination hint // Group body
var pag = document.getElementById("node-pagination"); var body = document.createElement("div");
pag.innerHTML = ""; body.className = "node-group-body" + (isExpanded ? "" : " collapsed");
if (total > nodes.length) { body.dataset.prefix = group.prefix;
pag.textContent = "Showing " + nodes.length + " of " + total + " nodes";
} var colHeaders = document.createElement("div");
colHeaders.className = "node-colheaders";
colHeaders.setAttribute("aria-hidden", "true");
colHeaders.innerHTML =
'<span class="ncol ncol-node">NODE</span>' +
'<span class="ncol ncol-ws">WS</span>' +
'<span class="ncol ncol-run">RUN</span>' +
'<span class="ncol ncol-attn">ATTN</span>' +
'<span class="ncol ncol-tokens">TOKENS</span>' +
'<span class="ncol ncol-health">LOAD</span>';
body.appendChild(colHeaders);
group.nodes.forEach(function (node) {
body.appendChild(buildNodeRow(node));
});
groupEl.appendChild(body);
table.appendChild(groupEl);
});
} }
// --- Drill-down: Node --- // --- Drill-down: Node ---
@@ -349,30 +594,38 @@ function drillDownToNode(nodeId, serverUrl) {
link.style.display = ""; link.style.display = "";
} }
document.getElementById("main").scrollTop = 0; document.getElementById("main").scrollTop = 0;
document.getElementById("node-ws-table").innerHTML =
'<div class="dashboard-empty">Loading workstreams...</div>';
loadNodeDetail(nodeId); loadNodeDetail(nodeId);
document.getElementById("breadcrumb-home").focus(); document.getElementById("breadcrumb-home").focus();
history.pushState({ view: "node", nodeId: nodeId, serverUrl: serverUrl }, ""); history.pushState({ view: "node", nodeId: nodeId, serverUrl: serverUrl }, "");
} }
function loadNodeDetail(nodeId) { function loadNodeDetail(nodeId) {
authFetch("/api/cluster/node/" + encodeURIComponent(nodeId)) var detailP = authFetch(
.then(function (r) { "/api/cluster/node/" + encodeURIComponent(nodeId),
return r.json(); ).then(function (r) {
}) return r.json();
.then(function (data) { });
if (data.error) { var overviewP = authFetch("/api/cluster/overview").then(function (r) {
document.getElementById("node-ws-table").innerHTML = return r.json();
'<div class="dashboard-empty">' + escapeHtml(data.error) + "</div>"; });
return; Promise.all([detailP, overviewP]).then(function (res) {
} var data = res[0];
var ws = data.workstreams || []; renderStatusBar(res[1]);
var active = ws.filter(function (w) { if (data.error) {
return w.state !== "idle"; document.getElementById("node-ws-table").innerHTML =
}).length; '<div class="dashboard-empty">' + escapeHtml(data.error) + "</div>";
document.getElementById("node-ws-summary").textContent = return;
active + " active \u00b7 " + ws.length + " total"; }
renderWsTable(document.getElementById("node-ws-table"), ws); var ws = data.workstreams || [];
}); var active = ws.filter(function (w) {
return w.state !== "idle";
}).length;
document.getElementById("node-ws-summary").textContent =
active + " active \u00b7 " + ws.length + " total";
renderWsTable(document.getElementById("node-ws-table"), ws);
});
} }
// --- Drill-down: Filtered --- // --- Drill-down: Filtered ---
@@ -417,11 +670,16 @@ function loadFilteredWorkstreams() {
params += "&state=" + encodeURIComponent(currentFilter.state); params += "&state=" + encodeURIComponent(currentFilter.state);
if (currentFilter.node) if (currentFilter.node)
params += "&node=" + encodeURIComponent(currentFilter.node); params += "&node=" + encodeURIComponent(currentFilter.node);
authFetch("/api/cluster/workstreams?" + params) var wsP = authFetch("/api/cluster/workstreams?" + params).then(function (r) {
.then(function (r) { return r.json();
return r.json(); });
}) var overviewP = authFetch("/api/cluster/overview").then(function (r) {
.then(function (data) { return r.json();
});
Promise.all([wsP, overviewP])
.then(function (res) {
var data = res[0];
renderStatusBar(res[1]);
document.getElementById("main").scrollTop = 0; document.getElementById("main").scrollTop = 0;
document.getElementById("filtered-summary").textContent = document.getElementById("filtered-summary").textContent =
"Page " + "Page " +
+15 -17
View File
@@ -4,41 +4,33 @@
<meta charset="utf-8"> <meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1"> <meta name="viewport" content="width=device-width, initial-scale=1">
<title>turnstone console</title> <title>turnstone console</title>
<link rel="preconnect" href="https://fonts.googleapis.com">
<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin>
<link href="https://fonts.googleapis.com/css2?family=IBM+Plex+Mono:wght@400;500;600&family=Outfit:wght@400;500;600;700&display=swap" rel="stylesheet">
<link rel="stylesheet" href="/static/style.css"> <link rel="stylesheet" href="/static/style.css">
</head> </head>
<body> <body>
<div id="header"> <div id="header">
<h1>turnstone console</h1> <h1>turnstone <span class="header-dim">console</span></h1>
<span id="cluster-summary" aria-live="polite"></span> <span id="cluster-summary" aria-live="polite"></span>
<span id="status-bar" role="status" aria-live="assertive"></span> <span id="status-bar" role="status" aria-live="polite"></span>
<button id="logout-btn" onclick="logout()" style="display:none; background:none; border:1px solid var(--border); color:var(--fg-dim); border-radius:var(--radius); padding:4px 8px; cursor:pointer; font:inherit; font-size:12px;">logout</button> <button id="logout-btn" class="header-btn" onclick="logout()" style="display:none">logout</button>
<button id="theme-toggle" onclick="toggleTheme()" aria-label="Toggle light/dark theme" style="margin-left:auto; background:none; border:1px solid var(--border); color:var(--fg); border-radius:var(--radius); padding:4px 8px; cursor:pointer; font:inherit; font-size:12px;">&#9790;</button> <button id="theme-toggle" class="header-btn" onclick="toggleTheme()" aria-label="Toggle light/dark theme">&#9790;</button>
</div> </div>
<nav id="breadcrumb" class="breadcrumb" style="display:none" aria-label="Breadcrumb"> <nav id="breadcrumb" class="breadcrumb" style="display:none" aria-label="Breadcrumb">
<a href="#" id="breadcrumb-home" onclick="showOverview(); return false">Cluster</a> <a href="#" id="breadcrumb-home" onclick="showOverview(); return false">Cluster</a>
<span class="breadcrumb-sep" aria-hidden="true">&gt;</span> <span class="breadcrumb-sep" aria-hidden="true">/</span>
<span id="breadcrumb-label" aria-current="page"></span> <span id="breadcrumb-label" aria-current="page"></span>
</nav> </nav>
<div id="main"> <div id="main">
<!-- CLUSTER OVERVIEW --> <!-- CLUSTER OVERVIEW -->
<div id="view-overview"> <div id="view-overview">
<div class="state-cards" id="state-cards"></div>
<div class="aggregate-bar" id="aggregate-bar"></div>
<div class="section-header">NODES</div> <div class="section-header">NODES</div>
<div class="node-colheaders" aria-hidden="true"> <div id="node-table" role="list" aria-label="Nodes" aria-live="polite">
<span class="ncol ncol-node">NODE</span>
<span class="ncol ncol-ws">WS</span>
<span class="ncol ncol-run">RUN</span>
<span class="ncol ncol-attn">ATTN</span>
<span class="ncol ncol-tokens">TOKENS</span>
<span class="ncol ncol-health">LOAD</span>
</div>
<div id="node-table" role="group" aria-label="Nodes" aria-live="polite">
<div class="dashboard-empty">Loading cluster data...</div> <div class="dashboard-empty">Loading cluster data...</div>
</div> </div>
<div id="node-pagination" class="pagination"></div>
</div> </div>
<!-- NODE DRILL-DOWN --> <!-- NODE DRILL-DOWN -->
@@ -78,6 +70,12 @@
</div> </div>
</div> </div>
<div id="cluster-status-bar" role="region" aria-label="Cluster status">
<div class="csb-states" id="csb-states"><span class="csb-loading">Loading...</span></div>
<div class="csb-divider" aria-hidden="true"></div>
<div class="csb-metrics" id="csb-metrics"></div>
</div>
<script src="/static/app.js"></script> <script src="/static/app.js"></script>
</body> </body>
</html> </html>
File diff suppressed because it is too large Load Diff
+2 -7
View File
@@ -109,9 +109,7 @@ def load_auth_config() -> AuthConfig:
tokens[env_token] = "full" tokens[env_token] = "full"
if enabled and not tokens: if enabled and not tokens:
log.warning( log.warning("Auth enabled but no tokens configured — all API requests will be rejected")
"Auth enabled but no tokens configured — all API requests will be rejected"
)
return AuthConfig(enabled=enabled, tokens=tokens) return AuthConfig(enabled=enabled, tokens=tokens)
@@ -125,10 +123,7 @@ def is_public_path(path: str) -> bool:
"""Return *True* if the path should be accessible without authentication.""" """Return *True* if the path should be accessible without authentication."""
if path in PUBLIC_PATHS: if path in PUBLIC_PATHS:
return True return True
for prefix in PUBLIC_PREFIXES: return any(path.startswith(prefix) for prefix in PUBLIC_PREFIXES)
if path.startswith(prefix):
return True
return False
def required_role(method: str, path: str) -> str: def required_role(method: str, path: str) -> str:
+9 -7
View File
@@ -6,11 +6,13 @@ Precedence: CLI args > env vars > config file > hardcoded defaults.
from __future__ import annotations from __future__ import annotations
import argparse
import logging import logging
from pathlib import Path
import tomllib import tomllib
from pathlib import Path
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
import argparse
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@@ -18,10 +20,10 @@ CONFIG_DIR = Path("~/.config/turnstone").expanduser()
CONFIG_PATH = CONFIG_DIR / "config.toml" CONFIG_PATH = CONFIG_DIR / "config.toml"
# Cache: None = not loaded yet, {} = loaded but empty/missing # Cache: None = not loaded yet, {} = loaded but empty/missing
_cache: dict | None = None _cache: dict[str, Any] | None = None
def load_config(section: str | None = None) -> dict: def load_config(section: str | None = None) -> dict[str, Any]:
"""Load config.toml and return the full dict or a specific section. """Load config.toml and return the full dict or a specific section.
Returns empty dict if file doesn't exist or can't be parsed. Returns empty dict if file doesn't exist or can't be parsed.
@@ -36,7 +38,8 @@ def load_config(section: str | None = None) -> dict:
except Exception as exc: except Exception as exc:
log.warning("Failed to parse %s: %s", CONFIG_PATH, exc) log.warning("Failed to parse %s: %s", CONFIG_PATH, exc)
if section: if section:
return _cache.get(section, {}) result = _cache.get(section, {})
return result if isinstance(result, dict) else {}
return _cache return _cache
@@ -57,7 +60,6 @@ _CONFIG_MAP: dict[str, dict[str, str]] = {
"context_window": "context_window", "context_window": "context_window",
}, },
"session": { "session": {
"persona": "persona",
"instructions": "instructions", "instructions": "instructions",
"retention_days": "session_retention_days", "retention_days": "session_retention_days",
"compact_max_tokens": "compact_max_tokens", "compact_max_tokens": "compact_max_tokens",
+26 -30
View File
@@ -1,8 +1,14 @@
"""SQLite database for persistent memories and conversation history.""" """SQLite database for persistent memories and conversation history."""
from __future__ import annotations
import os import os
import sqlite3 import sqlite3
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from collections.abc import Callable
TURNSTONE_DB = os.path.join(os.getcwd(), ".turnstone.db") TURNSTONE_DB = os.path.join(os.getcwd(), ".turnstone.db")
db_override: str | None = None db_override: str | None = None
@@ -52,9 +58,7 @@ def open_db() -> sqlite3.Connection:
"role TEXT NOT NULL, content TEXT, " "role TEXT NOT NULL, content TEXT, "
"tool_name TEXT, tool_args TEXT)" "tool_name TEXT, tool_args TEXT)"
) )
conn.execute( conn.execute("CREATE INDEX IF NOT EXISTS idx_conv_session ON conversations(session_id)")
"CREATE INDEX IF NOT EXISTS idx_conv_session ON conversations(session_id)"
)
# Migration: add tool_call_id column if missing (for session resume) # Migration: add tool_call_id column if missing (for session resume)
try: try:
conn.execute("SELECT tool_call_id FROM conversations LIMIT 0") conn.execute("SELECT tool_call_id FROM conversations LIMIT 0")
@@ -68,23 +72,18 @@ def open_db() -> sqlite3.Connection:
"title TEXT, created TEXT NOT NULL, updated TEXT NOT NULL)" "title TEXT, created TEXT NOT NULL, updated TEXT NOT NULL)"
) )
conn.execute("CREATE INDEX IF NOT EXISTS idx_sessions_alias ON sessions(alias)") conn.execute("CREATE INDEX IF NOT EXISTS idx_sessions_alias ON sessions(alias)")
conn.execute( conn.execute("CREATE INDEX IF NOT EXISTS idx_sessions_updated ON sessions(updated)")
"CREATE INDEX IF NOT EXISTS idx_sessions_updated ON sessions(updated)"
)
try: try:
# Check if FTS table already exists # Check if FTS table already exists
fts_exists = conn.execute( fts_exists = conn.execute(
"SELECT 1 FROM sqlite_master " "SELECT 1 FROM sqlite_master WHERE type='table' AND name='conversations_fts'"
"WHERE type='table' AND name='conversations_fts'"
).fetchone() ).fetchone()
if not fts_exists: if not fts_exists:
conn.execute( conn.execute(
"CREATE VIRTUAL TABLE conversations_fts " "CREATE VIRTUAL TABLE conversations_fts "
"USING fts5(content, content=conversations, content_rowid=id)" "USING fts5(content, content=conversations, content_rowid=id)"
) )
conn.execute( conn.execute("INSERT INTO conversations_fts(conversations_fts) VALUES('rebuild')")
"INSERT INTO conversations_fts(conversations_fts) VALUES('rebuild')"
)
conn.commit() conn.commit()
_fts5_available = True _fts5_available = True
except Exception: except Exception:
@@ -103,9 +102,7 @@ def load_memories() -> list[tuple[str, str]]:
try: try:
conn = open_db() conn = open_db()
try: try:
return conn.execute( return conn.execute("SELECT key, value FROM memories ORDER BY key").fetchall()
"SELECT key, value FROM memories ORDER BY key"
).fetchall()
finally: finally:
conn.close() conn.close()
except Exception: except Exception:
@@ -172,7 +169,7 @@ def fts5_query(query: str) -> str:
return " ".join(safe) return " ".join(safe)
def search_history(query: str, limit: int = 20) -> list[tuple]: def search_history(query: str, limit: int = 20) -> list[tuple[Any, ...]]:
"""Search conversation history. Returns (timestamp, session_id, role, content, tool_name).""" """Search conversation history. Returns (timestamp, session_id, role, content, tool_name)."""
if not query or not query.strip(): if not query or not query.strip():
return [] return []
@@ -200,7 +197,7 @@ def search_history(query: str, limit: int = 20) -> list[tuple]:
return [] return []
def search_history_recent(limit: int = 20) -> list[tuple]: def search_history_recent(limit: int = 20) -> list[tuple[Any, ...]]:
"""Return most recent conversation messages.""" """Return most recent conversation messages."""
try: try:
conn = open_db() conn = open_db()
@@ -285,7 +282,8 @@ def get_session_name(session_id: str) -> str | None:
(session_id,), (session_id,),
).fetchone() ).fetchone()
if row: if row:
return row[0] or row[1] or None value = row[0] or row[1]
return str(value) if value is not None else None
finally: finally:
conn.close() conn.close()
except Exception: except Exception:
@@ -304,25 +302,24 @@ def resolve_session(alias_or_id: str) -> str | None:
(alias_or_id,), (alias_or_id,),
).fetchone() ).fetchone()
if row: if row:
return row[0] return str(row[0])
# 2. Exact session_id match # 2. Exact session_id match
row = conn.execute( row = conn.execute(
"SELECT session_id FROM sessions WHERE session_id = ?", "SELECT session_id FROM sessions WHERE session_id = ?",
(alias_or_id,), (alias_or_id,),
).fetchone() ).fetchone()
if row: if row:
return row[0] return str(row[0])
# 3. Session_id prefix match # 3. Session_id prefix match
rows = conn.execute( rows = conn.execute(
"SELECT session_id FROM sessions WHERE session_id LIKE ?", "SELECT session_id FROM sessions WHERE session_id LIKE ?",
(alias_or_id + "%",), (alias_or_id + "%",),
).fetchall() ).fetchall()
if len(rows) == 1: if len(rows) == 1:
return rows[0][0] return str(rows[0][0])
# 4. Fallback: check conversations table for legacy sessions # 4. Fallback: check conversations table for legacy sessions
row = conn.execute( row = conn.execute(
"SELECT DISTINCT session_id FROM conversations " "SELECT DISTINCT session_id FROM conversations WHERE session_id = ? LIMIT 1",
"WHERE session_id = ? LIMIT 1",
(alias_or_id,), (alias_or_id,),
).fetchone() ).fetchone()
if row: if row:
@@ -336,7 +333,7 @@ def resolve_session(alias_or_id: str) -> str | None:
(row[0], row[0], row[0]), (row[0], row[0], row[0]),
) )
conn.commit() conn.commit()
return row[0] return str(row[0])
return None return None
finally: finally:
conn.close() conn.close()
@@ -346,7 +343,7 @@ def resolve_session(alias_or_id: str) -> str | None:
def prune_sessions( def prune_sessions(
retention_days: int = 90, retention_days: int = 90,
log_fn=None, log_fn: Callable[[str], None] | None = None,
) -> tuple[int, int]: ) -> tuple[int, int]:
"""Prune orphaned and stale sessions. """Prune orphaned and stale sessions.
@@ -395,15 +392,14 @@ def prune_sessions(
parts.append(f"{orphans} empty session{'s' if orphans != 1 else ''}") parts.append(f"{orphans} empty session{'s' if orphans != 1 else ''}")
if stale: if stale:
parts.append( parts.append(
f"{stale} session{'s' if stale != 1 else ''} " f"{stale} session{'s' if stale != 1 else ''} older than {retention_days} days"
f"older than {retention_days} days"
) )
log_fn(f"[turnstone] Session cleanup: removed {', '.join(parts)}.") log_fn(f"[turnstone] Session cleanup: removed {', '.join(parts)}.")
return (orphans, stale) return (orphans, stale)
def list_sessions(limit: int = 20) -> list[tuple]: def list_sessions(limit: int = 20) -> list[tuple[Any, ...]]:
"""List recent sessions. """List recent sessions.
Returns (session_id, alias, title, created, updated, msg_count) Returns (session_id, alias, title, created, updated, msg_count)
@@ -428,7 +424,7 @@ def list_sessions(limit: int = 20) -> list[tuple]:
return [] return []
def load_session_messages(session_id: str) -> list[dict]: def load_session_messages(session_id: str) -> list[dict[str, Any]]:
"""Load messages for a session and reconstruct OpenAI message format. """Load messages for a session and reconstruct OpenAI message format.
Handles tool_call / tool_result rows by grouping consecutive tool_call Handles tool_call / tool_result rows by grouping consecutive tool_call
@@ -448,7 +444,7 @@ def load_session_messages(session_id: str) -> list[dict]:
except Exception: except Exception:
return [] return []
messages: list[dict] = [] messages: list[dict[str, Any]] = []
i = 0 i = 0
while i < len(rows): while i < len(rows):
role, content, tool_name, tool_args, tc_id = rows[i] role, content, tool_name, tool_args, tc_id = rows[i]
@@ -465,7 +461,7 @@ def load_session_messages(session_id: str) -> list[dict]:
# Collect consecutive tool_call rows into one assistant message. # Collect consecutive tool_call rows into one assistant message.
# If the previous message was an assistant with content (text + # If the previous message was an assistant with content (text +
# tool calls in the same turn), merge tool_calls into it. # tool calls in the same turn), merge tool_calls into it.
assistant_msg: dict = { assistant_msg: dict[str, Any] = {
"role": "assistant", "role": "assistant",
"content": None, "content": None,
"tool_calls": [], "tool_calls": [],
+33 -26
View File
@@ -1,8 +1,11 @@
"""Thread-safe Prometheus-compatible metrics collector for the turnstone web server.""" """Thread-safe Prometheus-compatible metrics collector for the turnstone web server."""
from __future__ import annotations
import threading import threading
import time import time
from collections import defaultdict from collections import defaultdict
from typing import Any
class MetricsCollector: class MetricsCollector:
@@ -10,22 +13,22 @@ class MetricsCollector:
BUCKETS = [0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0] BUCKETS = [0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0]
def __init__(self): def __init__(self) -> None:
self._lock = threading.Lock() self._lock = threading.Lock()
self.start_time = time.monotonic() self.start_time = time.monotonic()
self.model: str = "" self.model: str = ""
# counters # counters
self._req_total: dict = defaultdict(int) # (method, endpoint, status) -> int self._req_total: dict[tuple[str, str, str], int] = defaultdict(int)
self._tokens: dict = defaultdict(int) # ("prompt"|"completion") -> int self._tokens: dict[str, int] = defaultdict(int) # "prompt"|"completion" -> int
self._messages: int = 0 self._messages: int = 0
self._tool_calls: dict = defaultdict(int) # tool_name -> int self._tool_calls: dict[str, int] = defaultdict(int) # tool_name -> int
self._errors: int = 0 self._errors: int = 0
# histograms: (method, endpoint) -> {buckets: [count…], sum: float, count: int} # histograms: (method, endpoint) -> {buckets: [count…], sum: float, count: int}
self._req_duration: dict = {} self._req_duration: dict[tuple[str, str], dict[str, Any]] = {}
# gauge # gauge
self._context_ratio: float = 0.0 self._context_ratio: float = 0.0
def record_request(self, method: str, endpoint: str, status: int, duration: float): def record_request(self, method: str, endpoint: str, status: int, duration: float) -> None:
with self._lock: with self._lock:
self._req_total[(method, endpoint, str(status))] += 1 self._req_total[(method, endpoint, str(status))] += 1
key = (method, endpoint) key = (method, endpoint)
@@ -42,43 +45,53 @@ class MetricsCollector:
h["sum"] += duration h["sum"] += duration
h["count"] += 1 h["count"] += 1
def record_tokens(self, prompt: int, completion: int): def record_tokens(self, prompt: int, completion: int) -> None:
with self._lock: with self._lock:
self._tokens["prompt"] += prompt self._tokens["prompt"] += prompt
self._tokens["completion"] += completion self._tokens["completion"] += completion
def record_tool_call(self, tool_name: str): def record_tool_call(self, tool_name: str) -> None:
with self._lock: with self._lock:
self._tool_calls[tool_name] += 1 self._tool_calls[tool_name] += 1
def record_error(self): def record_error(self) -> None:
with self._lock: with self._lock:
self._errors += 1 self._errors += 1
def record_message_sent(self): def record_message_sent(self) -> None:
with self._lock: with self._lock:
self._messages += 1 self._messages += 1
def record_context_ratio(self, ratio: float): def record_context_ratio(self, ratio: float) -> None:
with self._lock: with self._lock:
self._context_ratio = ratio self._context_ratio = ratio
def generate_text( def generate_text(
self, self,
workstream_states: dict, workstream_states: dict[str, int],
total_workstreams: int, total_workstreams: int,
workstream_metrics: list[dict] | None = None, workstream_metrics: list[dict[str, Any]] | None = None,
) -> str: ) -> str:
"""Return Prometheus text exposition format (v0.0.4).""" """Return Prometheus text exposition format (v0.0.4)."""
lines: list[str] = [] lines: list[str] = []
def gauge(name, help_text, value, labels=None): def gauge(
name: str,
help_text: str,
value: float | int,
labels: dict[str, str] | None = None,
) -> None:
lstr = _fmt_labels(labels) lstr = _fmt_labels(labels)
lines.append(f"# HELP {name} {help_text}") lines.append(f"# HELP {name} {help_text}")
lines.append(f"# TYPE {name} gauge") lines.append(f"# TYPE {name} gauge")
lines.append(f"{name}{lstr} {_fmt_value(value)}") lines.append(f"{name}{lstr} {_fmt_value(value)}")
def counter(name, help_text, value, labels=None): def counter(
name: str,
help_text: str,
value: float | int,
labels: dict[str, str] | None = None,
) -> None:
lstr = _fmt_labels(labels) lstr = _fmt_labels(labels)
lines.append(f"# HELP {name} {help_text}") lines.append(f"# HELP {name} {help_text}")
lines.append(f"# TYPE {name} counter") lines.append(f"# TYPE {name} counter")
@@ -132,8 +145,7 @@ class MetricsCollector:
lines.append("# TYPE turnstone_http_request_duration_seconds histogram") lines.append("# TYPE turnstone_http_request_duration_seconds histogram")
for (method, endpoint), h in sorted(req_duration.items()): for (method, endpoint), h in sorted(req_duration.items()):
prefix = ( prefix = (
f'turnstone_http_request_duration_seconds{{method="{method}",' f'turnstone_http_request_duration_seconds{{method="{method}",endpoint="{endpoint}"'
f'endpoint="{endpoint}"'
) )
for i, b in enumerate(self.BUCKETS): for i, b in enumerate(self.BUCKETS):
lines.append(f'{prefix},le="{b}"}} {h["buckets"][i]}') lines.append(f'{prefix},le="{b}"}} {h["buckets"][i]}')
@@ -154,9 +166,7 @@ class MetricsCollector:
lines.append("# HELP turnstone_tokens_total Total tokens consumed") lines.append("# HELP turnstone_tokens_total Total tokens consumed")
lines.append("# TYPE turnstone_tokens_total counter") lines.append("# TYPE turnstone_tokens_total counter")
for tok_type in ("prompt", "completion"): for tok_type in ("prompt", "completion"):
lines.append( lines.append(f'turnstone_tokens_total{{type="{tok_type}"}} {tokens.get(tok_type, 0)}')
f'turnstone_tokens_total{{type="{tok_type}"}} {tokens.get(tok_type, 0)}'
)
# turnstone_tool_calls_total # turnstone_tool_calls_total
lines.append("# HELP turnstone_tool_calls_total Total tool executions by name") lines.append("# HELP turnstone_tool_calls_total Total tool executions by name")
@@ -212,8 +222,7 @@ class MetricsCollector:
for wm in workstream_metrics: for wm in workstream_metrics:
lstr = _fmt_labels({"ws_id": wm["ws_id"], "name": wm["name"]}) lstr = _fmt_labels({"ws_id": wm["ws_id"], "name": wm["name"]})
lines.append( lines.append(
f"turnstone_workstream_completion_tokens_total{lstr}" f"turnstone_workstream_completion_tokens_total{lstr} {wm['completion_tokens']}"
f" {wm['completion_tokens']}"
) )
lines.append( lines.append(
@@ -232,9 +241,7 @@ class MetricsCollector:
lines.append("# TYPE turnstone_workstream_tool_calls_total counter") lines.append("# TYPE turnstone_workstream_tool_calls_total counter")
for wm in workstream_metrics: for wm in workstream_metrics:
for tool, cnt in sorted(wm["tool_calls"].items()): for tool, cnt in sorted(wm["tool_calls"].items()):
lstr = _fmt_labels( lstr = _fmt_labels({"ws_id": wm["ws_id"], "name": wm["name"], "tool": tool})
{"ws_id": wm["ws_id"], "name": wm["name"], "tool": tool}
)
lines.append(f"turnstone_workstream_tool_calls_total{lstr} {cnt}") lines.append(f"turnstone_workstream_tool_calls_total{lstr} {cnt}")
lines.append( lines.append(
@@ -252,7 +259,7 @@ class MetricsCollector:
return "\n".join(lines) return "\n".join(lines)
def _fmt_labels(labels: dict | None) -> str: def _fmt_labels(labels: dict[str, str] | None) -> str:
if not labels: if not labels:
return "" return ""
parts = [f'{k}="{v}"' for k, v in labels.items()] parts = [f'{k}="{v}"' for k, v in labels.items()]
+36 -26
View File
@@ -1,9 +1,12 @@
"""Sandboxed Python executor for the math tool.""" """Sandboxed Python executor for the math tool."""
from __future__ import annotations
import ast import ast
import multiprocessing import multiprocessing
import re import re
import traceback import traceback
from typing import Any
_MATH_BLOCKED_BUILTINS = { _MATH_BLOCKED_BUILTINS = {
"open", "open",
@@ -56,29 +59,32 @@ _MATH_BLOCKED_MODULES = {
class _ASTValidator(ast.NodeVisitor): class _ASTValidator(ast.NodeVisitor):
"""Validates AST for dangerous constructs.""" """Validates AST for dangerous constructs."""
def __init__(self): def __init__(self) -> None:
self.errors: list[str] = [] self.errors: list[str] = []
def visit_Import(self, node): def visit_Import(self, node: ast.Import) -> None:
for alias in node.names: for alias in node.names:
if alias.name.split(".")[0] in _MATH_BLOCKED_MODULES: if alias.name.split(".")[0] in _MATH_BLOCKED_MODULES:
self.errors.append(f"Import of '{alias.name}' is not allowed") self.errors.append(f"Import of '{alias.name}' is not allowed")
self.generic_visit(node) self.generic_visit(node)
def visit_ImportFrom(self, node): def visit_ImportFrom(self, node: ast.ImportFrom) -> None:
if node.module and node.module.split(".")[0] in _MATH_BLOCKED_MODULES: if node.module and node.module.split(".")[0] in _MATH_BLOCKED_MODULES:
self.errors.append(f"Import from '{node.module}' is not allowed") self.errors.append(f"Import from '{node.module}' is not allowed")
self.generic_visit(node) self.generic_visit(node)
def visit_Call(self, node): def visit_Call(self, node: ast.Call) -> None:
if isinstance(node.func, ast.Name) and node.func.id in _MATH_BLOCKED_BUILTINS: if isinstance(node.func, ast.Name) and node.func.id in _MATH_BLOCKED_BUILTINS:
self.errors.append(f"Call to '{node.func.id}' is not allowed") self.errors.append(f"Call to '{node.func.id}' is not allowed")
self.generic_visit(node) self.generic_visit(node)
def visit_Attribute(self, node): def visit_Attribute(self, node: ast.Attribute) -> None:
if node.attr.startswith("__") and node.attr.endswith("__"): if (
if node.attr not in {"__name__", "__doc__", "__class__"}: node.attr.startswith("__")
self.errors.append(f"Access to '{node.attr}' is not allowed") and node.attr.endswith("__")
and node.attr not in {"__name__", "__doc__", "__class__"}
):
self.errors.append(f"Access to '{node.attr}' is not allowed")
self.generic_visit(node) self.generic_visit(node)
@@ -101,7 +107,7 @@ def validate_math_code(code: str) -> list[str]:
return v.errors return v.errors
def _math_exec_in_process(code: str, result_queue: multiprocessing.Queue): def _math_exec_in_process(code: str, result_queue: multiprocessing.Queue[tuple[str, str]]) -> None:
"""Execute code in a subprocess, put (status, output) in queue.""" """Execute code in a subprocess, put (status, output) in queue."""
import signal as _signal import signal as _signal
import sys as _sys import sys as _sys
@@ -115,7 +121,7 @@ def _math_exec_in_process(code: str, result_queue: multiprocessing.Queue):
captured = StringIO() captured = StringIO()
_sys.stdout = captured _sys.stdout = captured
def _safe_import(name, *args, **kwargs): def _safe_import(name: str, *args: Any, **kwargs: Any) -> Any:
if name.split(".")[0] in _MATH_BLOCKED_MODULES: if name.split(".")[0] in _MATH_BLOCKED_MODULES:
raise ImportError(f"Import of '{name}' is blocked") raise ImportError(f"Import of '{name}' is blocked")
return original_import(name, *args, **kwargs) return original_import(name, *args, **kwargs)
@@ -137,10 +143,18 @@ def _math_exec_in_process(code: str, result_queue: multiprocessing.Queue):
safe_builtins["__import__"] = _safe_import safe_builtins["__import__"] = _safe_import
# Pre-import safe modules # Pre-import safe modules
import math, fractions, itertools, functools, operator import collections
import collections, decimal, random, re, string import decimal
import fractions
import functools
import itertools
import math
import operator
import random
import re
import string
ns: dict = { ns: dict[str, Any] = {
"__builtins__": safe_builtins, "__builtins__": safe_builtins,
"math": math, "math": math,
"fractions": fractions, "fractions": fractions,
@@ -212,7 +226,11 @@ def _math_exec_in_process(code: str, result_queue: multiprocessing.Queue):
pass pass
try: try:
import scipy, scipy.special, scipy.optimize, scipy.integrate, scipy.linalg import scipy # type: ignore[import-untyped]
import scipy.integrate # type: ignore[import-untyped]
import scipy.linalg # type: ignore[import-untyped]
import scipy.optimize # type: ignore[import-untyped]
import scipy.special # type: ignore[import-untyped]
ns["scipy"] = scipy ns["scipy"] = scipy
ns["special"] = scipy.special ns["special"] = scipy.special
@@ -230,11 +248,7 @@ def _math_exec_in_process(code: str, result_queue: multiprocessing.Queue):
printed = captured.getvalue() printed = captured.getvalue()
result_var = ns.get("result") result_var = ns.get("result")
if result_var is not None: if result_var is not None:
out = ( out = f"{printed.rstrip()}\nresult = {result_var}" if printed else str(result_var)
f"{printed.rstrip()}\nresult = {result_var}"
if printed
else str(result_var)
)
elif printed: elif printed:
out = printed.rstrip() out = printed.rstrip()
else: else:
@@ -243,9 +257,7 @@ def _math_exec_in_process(code: str, result_queue: multiprocessing.Queue):
except Exception as e: except Exception as e:
_sys.stdout = _sys.__stdout__ _sys.stdout = _sys.__stdout__
result_queue.put( result_queue.put(("error", f"{type(e).__name__}: {e}\n{traceback.format_exc()}"))
("error", f"{type(e).__name__}: {e}\n{traceback.format_exc()}")
)
def auto_print_wrap(code: str) -> str: def auto_print_wrap(code: str) -> str:
@@ -280,10 +292,8 @@ def execute_math_sandboxed(code: str, timeout: float = 30.0) -> tuple[str, bool]
if errors: if errors:
return "Validation errors:\n" + "\n".join(f"- {e}" for e in errors), True return "Validation errors:\n" + "\n".join(f"- {e}" for e in errors), True
result_queue: multiprocessing.Queue = multiprocessing.Queue() result_queue: multiprocessing.Queue[tuple[str, str]] = multiprocessing.Queue()
proc = multiprocessing.Process( proc = multiprocessing.Process(target=_math_exec_in_process, args=(code, result_queue))
target=_math_exec_in_process, args=(code, result_queue)
)
proc.start() proc.start()
proc.join(timeout=timeout) proc.join(timeout=timeout)
+145 -225
View File
File diff suppressed because it is too large Load Diff
+4 -1
View File
@@ -1,13 +1,16 @@
"""Tool definitions — auto-loaded from turnstone/tools/*.json.""" """Tool definitions — auto-loaded from turnstone/tools/*.json."""
from __future__ import annotations
import json import json
from pathlib import Path from pathlib import Path
from typing import Any
_TOOLS_DIR = Path(__file__).resolve().parent.parent / "tools" _TOOLS_DIR = Path(__file__).resolve().parent.parent / "tools"
_META_KEYS = {"agent", "task_agent", "auto_approve", "primary_key"} _META_KEYS = {"agent", "task_agent", "auto_approve", "primary_key"}
def _load_tools() -> tuple[list[dict], dict]: def _load_tools() -> tuple[list[dict[str, Any]], dict[str, Any]]:
"""Load all .json files from the tools directory. """Load all .json files from the tools directory.
Returns (tool_defs, metadata) where: Returns (tool_defs, metadata) where:
+12 -17
View File
@@ -12,9 +12,11 @@ import threading
import time import time
import uuid import uuid
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Callable from typing import TYPE_CHECKING
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Callable
from turnstone.core.session import ChatSession, SessionUI from turnstone.core.session import ChatSession, SessionUI
@@ -48,7 +50,7 @@ class Workstream:
last_active: float = field(default_factory=time.monotonic, repr=False) last_active: float = field(default_factory=time.monotonic, repr=False)
_lock: threading.Lock = field(default_factory=threading.Lock, repr=False) _lock: threading.Lock = field(default_factory=threading.Lock, repr=False)
def __post_init__(self): def __post_init__(self) -> None:
if not self.name: if not self.name:
self.name = f"ws-{self.id[:4]}" self.name = f"ws-{self.id[:4]}"
@@ -65,7 +67,7 @@ class WorkstreamManager:
def __init__( def __init__(
self, self,
session_factory: Callable[[SessionUI], ChatSession], session_factory: Callable[[SessionUI | None], ChatSession],
): ):
""" """
Args: Args:
@@ -73,7 +75,7 @@ class WorkstreamManager:
config (client, model, temperature, ) so the manager can config (client, model, temperature, ) so the manager can
create sessions without knowing those details. create sessions without knowing those details.
""" """
self._session_factory = session_factory self._session_factory: Callable[[SessionUI | None], ChatSession] = session_factory
self._workstreams: dict[str, Workstream] = {} self._workstreams: dict[str, Workstream] = {}
self._order: list[str] = [] # creation order self._order: list[str] = [] # creation order
self._active_id: str | None = None self._active_id: str | None = None
@@ -94,9 +96,7 @@ class WorkstreamManager:
ws.session = self._session_factory(ws.ui) ws.session = self._session_factory(ws.ui)
with self._lock: with self._lock:
if len(self._workstreams) >= self.MAX_WORKSTREAMS: if len(self._workstreams) >= self.MAX_WORKSTREAMS:
raise RuntimeError( raise RuntimeError(f"Maximum of {self.MAX_WORKSTREAMS} workstreams reached")
f"Maximum of {self.MAX_WORKSTREAMS} workstreams reached"
)
self._workstreams[ws.id] = ws self._workstreams[ws.id] = ws
self._order.append(ws.id) self._order.append(ws.id)
if self._active_id is None: if self._active_id is None:
@@ -117,10 +117,10 @@ class WorkstreamManager:
# Unblock any waiting approval/plan events so worker thread can exit # Unblock any waiting approval/plan events so worker thread can exit
if ws.ui: if ws.ui:
if hasattr(ws.ui, "_approval_event"): if hasattr(ws.ui, "_approval_event"):
ws.ui._approval_result = (False, None) ws.ui._approval_result = False, None # type: ignore[attr-defined]
ws.ui._approval_event.set() ws.ui._approval_event.set()
if hasattr(ws.ui, "_plan_event"): if hasattr(ws.ui, "_plan_event"):
ws.ui._plan_result = "reject" ws.ui._plan_result = "reject" # type: ignore[attr-defined]
ws.ui._plan_event.set() ws.ui._plan_event.set()
if hasattr(ws.ui, "_fg_event"): if hasattr(ws.ui, "_fg_event"):
ws.ui._fg_event.set() ws.ui._fg_event.set()
@@ -143,11 +143,7 @@ class WorkstreamManager:
def list_all(self) -> list[Workstream]: def list_all(self) -> list[Workstream]:
"""Return workstreams in creation order.""" """Return workstreams in creation order."""
with self._lock: with self._lock:
return [ return [self._workstreams[wid] for wid in self._order if wid in self._workstreams]
self._workstreams[wid]
for wid in self._order
if wid in self._workstreams
]
def index_of(self, ws_id: str) -> int: def index_of(self, ws_id: str) -> int:
"""1-based index of a workstream, or 0 if not found.""" """1-based index of a workstream, or 0 if not found."""
@@ -183,7 +179,7 @@ class WorkstreamManager:
# -- state management --------------------------------------------------- # -- state management ---------------------------------------------------
def set_state(self, ws_id: str, state: WorkstreamState, error_msg: str = ""): def set_state(self, ws_id: str, state: WorkstreamState, error_msg: str = "") -> None:
"""Update a workstream's state. Called by UI adapters.""" """Update a workstream's state. Called by UI adapters."""
ws = self._workstreams.get(ws_id) ws = self._workstreams.get(ws_id)
if ws: if ws:
@@ -207,8 +203,7 @@ class WorkstreamManager:
[ [
ws ws
for ws in snapshot for ws in snapshot
if ws.state == WorkstreamState.IDLE if ws.state == WorkstreamState.IDLE and (now - ws.last_active) > max_age_seconds
and (now - ws.last_active) > max_age_seconds
], ],
key=lambda ws: ws.last_active, # oldest first key=lambda ws: ws.last_active, # oldest first
) )
+112 -99
View File
@@ -23,13 +23,15 @@ import sys
import tempfile import tempfile
import textwrap import textwrap
import time import time
from collections.abc import Iterator
from datetime import datetime from datetime import datetime
from typing import Any
from openai import OpenAI from openai import OpenAI, Stream
from turnstone.core.session import ChatSession
from turnstone.core.tools import TOOLS, PRIMARY_KEY_MAP
import turnstone.core.memory as _memory_module import turnstone.core.memory as _memory_module
from turnstone.core.session import ChatSession
from turnstone.core.tools import PRIMARY_KEY_MAP, TOOLS
# ─── ANSI & logging helpers ─────────────────────────────────────────────────── # ─── ANSI & logging helpers ───────────────────────────────────────────────────
@@ -45,44 +47,47 @@ BOLD = "\033[1m"
class NullUI: class NullUI:
"""UI adapter that discards all output. Used by HeadlessSession.""" """UI adapter that discards all output. Used by HeadlessSession."""
def on_thinking_start(self): def on_thinking_start(self) -> None:
pass pass
def on_thinking_stop(self): def on_thinking_stop(self) -> None:
pass pass
def on_reasoning_token(self, text): def on_reasoning_token(self, text: str) -> None:
pass pass
def on_content_token(self, text): def on_content_token(self, text: str) -> None:
pass pass
def on_stream_end(self): def on_stream_end(self) -> None:
pass pass
def approve_tools(self, items): def approve_tools(self, items: list[dict[str, Any]]) -> tuple[bool, str | None]:
return True, None return True, None
def on_tool_result(self, name, output): def on_tool_result(self, name: str, output: str) -> None:
pass pass
def on_status(self, usage, context_window, effort): def on_status(self, usage: dict[str, Any], context_window: int, effort: str) -> None:
pass pass
def on_plan_review(self, content): def on_plan_review(self, content: str) -> str:
return "" return ""
def on_info(self, message): def on_info(self, message: str) -> None:
pass pass
def on_error(self, message): def on_error(self, message: str) -> None:
pass pass
def on_state_change(self, state): def on_state_change(self, state: str) -> None:
pass
def on_rename(self, name: str) -> None:
pass pass
def _log(msg: str, dim: bool = False): def _log(msg: str, dim: bool = False) -> None:
"""Print a log line with optional dim styling.""" """Print a log line with optional dim styling."""
if dim: if dim:
sys.stderr.write(f"{DIM}{msg}{RESET}\n") sys.stderr.write(f"{DIM}{msg}{RESET}\n")
@@ -91,7 +96,7 @@ def _log(msg: str, dim: bool = False):
sys.stderr.flush() sys.stderr.flush()
def _fmt_args(args: dict, max_len: int = 80) -> str: def _fmt_args(args: dict[str, Any], max_len: int = 80) -> str:
"""Format tool args as a compact one-line summary.""" """Format tool args as a compact one-line summary."""
parts = [] parts = []
for k, v in args.items(): for k, v in args.items():
@@ -109,7 +114,7 @@ def _fmt_args(args: dict, max_len: int = 80) -> str:
@contextlib.contextmanager @contextlib.contextmanager
def _suppress_stdout(): def _suppress_stdout() -> Iterator[None]:
"""Redirect stdout to devnull temporarily.""" """Redirect stdout to devnull temporarily."""
old = sys.stdout old = sys.stdout
sys.stdout = io.StringIO() sys.stdout = io.StringIO()
@@ -132,15 +137,43 @@ class HeadlessSession(ChatSession):
- send_headless() uses non-streaming API - send_headless() uses non-streaming API
""" """
def __init__(self, client, model, system_prompt_override=None, **kwargs): def __init__(
kwargs.setdefault("ui", NullUI()) self,
super().__init__(client=client, model=model, **kwargs) client: OpenAI,
self.tool_call_log: list[dict] = [] model: str,
system_prompt_override: str | None = None,
instructions: str | None = None,
temperature: float = 0.7,
max_tokens: int = 32768,
tool_timeout: int = 30,
reasoning_effort: str = "medium",
context_window: int = 131072,
compact_max_tokens: int = 32768,
auto_compact_pct: float = 0.8,
agent_max_turns: int = -1,
tool_truncation: int = 0,
) -> None:
super().__init__(
client=client,
model=model,
ui=NullUI(),
instructions=instructions,
temperature=temperature,
max_tokens=max_tokens,
tool_timeout=tool_timeout,
reasoning_effort=reasoning_effort,
context_window=context_window,
compact_max_tokens=compact_max_tokens,
auto_compact_pct=auto_compact_pct,
agent_max_turns=agent_max_turns,
tool_truncation=tool_truncation,
)
self.tool_call_log: list[dict[str, Any]] = []
self.auto_approve = True self.auto_approve = True
if system_prompt_override is not None: if system_prompt_override is not None:
self._override_system_prompt(system_prompt_override) self._override_system_prompt(system_prompt_override)
def _override_system_prompt(self, content: str): def _override_system_prompt(self, content: str) -> None:
"""Replace the developer message content with a custom prompt.""" """Replace the developer message content with a custom prompt."""
for i, msg in enumerate(self.system_messages): for i, msg in enumerate(self.system_messages):
if msg["role"] == "developer": if msg["role"] == "developer":
@@ -154,7 +187,7 @@ class HeadlessSession(ChatSession):
max_turns: int = 10, max_turns: int = 10,
verbose: bool = False, verbose: bool = False,
log_prefix: str = "", log_prefix: str = "",
) -> list[dict]: ) -> list[dict[str, Any]]:
"""Run a complete conversation turn headlessly. """Run a complete conversation turn headlessly.
Uses non-streaming API calls. Captures all tool calls into Uses non-streaming API calls. Captures all tool calls into
@@ -176,8 +209,8 @@ class HeadlessSession(ChatSession):
response = self.client.chat.completions.create( response = self.client.chat.completions.create(
model=self.model, model=self.model,
messages=msgs, messages=msgs, # type: ignore[arg-type]
tools=TOOLS, tools=TOOLS, # type: ignore[arg-type]
max_completion_tokens=self.max_tokens, max_completion_tokens=self.max_tokens,
temperature=self.temperature, temperature=self.temperature,
stream=False, stream=False,
@@ -189,8 +222,9 @@ class HeadlessSession(ChatSession):
) )
elapsed = time.monotonic() - t0 elapsed = time.monotonic() - t0
assert not isinstance(response, Stream)
choice = response.choices[0] choice = response.choices[0]
assistant_msg: dict = { assistant_msg: dict[str, Any] = {
"role": "assistant", "role": "assistant",
"content": choice.message.content or None, "content": choice.message.content or None,
} }
@@ -203,8 +237,8 @@ class HeadlessSession(ChatSession):
"id": tc.id, "id": tc.id,
"type": "function", "type": "function",
"function": { "function": {
"name": tc.function.name, "name": tc.function.name, # type: ignore[union-attr]
"arguments": tc.function.arguments, "arguments": tc.function.arguments, # type: ignore[union-attr]
}, },
} }
for tc in calls for tc in calls
@@ -237,15 +271,16 @@ class HeadlessSession(ChatSession):
# Log tool calls # Log tool calls
if verbose: if verbose:
names = [tc.function.name for tc in choice.message.tool_calls] names = [tc.function.name for tc in choice.message.tool_calls] # type: ignore[union-attr]
_log(f"{log_prefix} turn {turn}: tools -> {names}") _log(f"{log_prefix} turn {turn}: tools -> {names}")
# Execute tools with stdout suppressed # Execute tools with stdout suppressed
with _suppress_stdout(): with _suppress_stdout():
results, _ = self._execute_tools(assistant_msg["tool_calls"]) results, _ = self._execute_tools(assistant_msg["tool_calls"])
for tc, (tc_id, output) in zip(assistant_msg["tool_calls"], results): for tc, (tc_id, output) in zip(assistant_msg["tool_calls"], results, strict=False):
func_name = tc["function"]["name"] func_name = tc["function"]["name"]
args: dict[str, Any]
try: try:
args = json.loads(tc["function"]["arguments"]) args = json.loads(tc["function"]["arguments"])
except json.JSONDecodeError: except json.JSONDecodeError:
@@ -281,9 +316,7 @@ class HeadlessSession(ChatSession):
"content": output, "content": output,
} }
self.messages.append(tool_msg) self.messages.append(tool_msg)
self._msg_tokens.append( self._msg_tokens.append(max(1, int(len(output) / self._chars_per_token)))
max(1, int(len(output) / self._chars_per_token))
)
return self.tool_call_log return self.tool_call_log
@@ -295,14 +328,14 @@ def _run_single_test(
client: OpenAI, client: OpenAI,
model: str, model: str,
system_prompt: str, system_prompt: str,
case: dict, case: dict[str, Any],
temperature: float, temperature: float,
max_tokens: int, max_tokens: int,
reasoning_effort: str, reasoning_effort: str,
context_window: int, context_window: int,
verbose: bool = False, verbose: bool = False,
log_prefix: str = "", log_prefix: str = "",
) -> dict: ) -> dict[str, Any]:
"""Run a single test case once in an isolated temp directory. """Run a single test case once in an isolated temp directory.
Must be called serially uses os.chdir which is process-global. Must be called serially uses os.chdir which is process-global.
@@ -325,9 +358,7 @@ def _run_single_test(
f.write(content) f.write(content)
if verbose and setup_files: if verbose and setup_files:
_log( _log(f"{log_prefix} setup: created {[p for p, _ in setup_files]}", dim=True)
f"{log_prefix} setup: created {[p for p, _ in setup_files]}", dim=True
)
os.chdir(workdir) os.chdir(workdir)
@@ -335,7 +366,6 @@ def _run_single_test(
client=client, client=client,
model=model, model=model,
system_prompt_override=system_prompt, system_prompt_override=system_prompt,
persona=None,
instructions=None, instructions=None,
temperature=temperature, temperature=temperature,
max_tokens=max_tokens, max_tokens=max_tokens,
@@ -346,7 +376,8 @@ def _run_single_test(
max_turns = case.get("max_turns", 10) max_turns = case.get("max_turns", 10)
# Retry on transient API errors to avoid poisoning eval scores # Retry on transient API errors to avoid poisoning eval scores
_last_err = None tool_log: list[dict[str, Any]] = []
_last_err: Exception | None = None
for _attempt in range(3): for _attempt in range(3):
try: try:
tool_log = session.send_headless( tool_log = session.send_headless(
@@ -363,7 +394,7 @@ def _run_single_test(
_time.sleep(2**_attempt) _time.sleep(2**_attempt)
else: else:
raise _last_err raise _last_err or RuntimeError("send_headless failed after 3 attempts")
final_content = "" final_content = ""
for msg in reversed(session.messages): for msg in reversed(session.messages):
@@ -388,7 +419,7 @@ def _run_single_test(
# ─── Scoring ───────────────────────────────────────────────────────────────── # ─── Scoring ─────────────────────────────────────────────────────────────────
def _match_action(actual: dict, expected: dict) -> bool: def _match_action(actual: dict[str, Any], expected: dict[str, Any]) -> bool:
"""Check if a single actual tool call matches an expected action spec.""" """Check if a single actual tool call matches an expected action spec."""
if actual["tool"] != expected["tool"]: if actual["tool"] != expected["tool"]:
return False return False
@@ -419,10 +450,10 @@ def _match_action(actual: dict, expected: dict) -> bool:
def score_run( def score_run(
tool_log: list[dict], tool_log: list[dict[str, Any]],
expected_actions: list[dict], expected_actions: list[dict[str, Any]],
match_mode: str = "ordered_subset", match_mode: str = "ordered_subset",
) -> dict: ) -> dict[str, Any]:
"""Score a single run's tool log against expected actions. """Score a single run's tool log against expected actions.
Returns dict with: pass, score, matched, unmatched, extra_tools, detail. Returns dict with: pass, score, matched, unmatched, extra_tools, detail.
@@ -441,7 +472,7 @@ def score_run(
if match_mode == "exact": if match_mode == "exact":
matched = [] matched = []
for i, (actual, expected) in enumerate(zip(tool_log, expected_actions)): for i, (actual, expected) in enumerate(zip(tool_log, expected_actions, strict=False)):
if _match_action(actual, expected): if _match_action(actual, expected):
matched.append(i) matched.append(i)
score = len(matched) / n_expected score = len(matched) / n_expected
@@ -537,25 +568,23 @@ def _run_iteration(
client: OpenAI, client: OpenAI,
model: str, model: str,
system_prompt: str, system_prompt: str,
cases: list[dict], cases: list[dict[str, Any]],
n_runs: int, n_runs: int,
temperature: float, temperature: float,
max_tokens: int, max_tokens: int,
reasoning_effort: str, reasoning_effort: str,
context_window: int, context_window: int,
verbose: bool = False, verbose: bool = False,
) -> dict: ) -> dict[str, Any]:
"""Run all test cases n_runs times and score them.""" """Run all test cases n_runs times and score them."""
case_results = {} case_results: dict[str, Any] = {}
for ci, case in enumerate(cases): for ci, case in enumerate(cases):
case_id = case["id"] case_id = case["id"]
case_n = case.get("n_runs", n_runs) case_n = case.get("n_runs", n_runs)
runs = [] runs: list[dict[str, Any]] = []
print( print(f"\n {CYAN}[{ci + 1}/{len(cases)}]{RESET} {BOLD}{case_id}{RESET} ({case_n} runs)")
f"\n {CYAN}[{ci + 1}/{len(cases)}]{RESET} {BOLD}{case_id}{RESET} ({case_n} runs)"
)
if verbose: if verbose:
_log(f" prompt: {case['user_prompt']}", dim=True) _log(f" prompt: {case['user_prompt']}", dim=True)
@@ -582,12 +611,8 @@ def _run_iteration(
match_mode=case.get("match_mode", "ordered_subset"), match_mode=case.get("match_mode", "ordered_subset"),
) )
score_result["tool_sequence"] = [ score_result["tool_sequence"] = [t["tool"] for t in run_result["tool_log"]]
t["tool"] for t in run_result["tool_log"] score_result["tool_args"] = [{t["tool"]: t["args"]} for t in run_result["tool_log"]]
]
score_result["tool_args"] = [
{t["tool"]: t["args"]} for t in run_result["tool_log"]
]
score_result["elapsed"] = run_result.get("elapsed", 0) score_result["elapsed"] = run_result.get("elapsed", 0)
# Detect JSON dumped into final channel (tool call not made) # Detect JSON dumped into final channel (tool call not made)
@@ -619,9 +644,7 @@ def _run_iteration(
status_label = "PASS" if passed else "FAIL" status_label = "PASS" if passed else "FAIL"
tools = score_result.get("tool_sequence", []) tools = score_result.get("tool_sequence", [])
elapsed = score_result.get("elapsed", 0) elapsed = score_result.get("elapsed", 0)
json_flag = ( json_flag = f" {YELLOW}[JSON_DUMP]{RESET}" if score_result.get("json_dump") else ""
f" {YELLOW}[JSON_DUMP]{RESET}" if score_result.get("json_dump") else ""
)
print( print(
f" Run {run_idx + 1}: " f" Run {run_idx + 1}: "
f"{status_color}[{status_label}]{RESET} " f"{status_color}[{status_label}]{RESET} "
@@ -642,9 +665,7 @@ def _run_iteration(
# Aggregate # Aggregate
total_runs = sum(len(cr["runs"]) for cr in case_results.values()) total_runs = sum(len(cr["runs"]) for cr in case_results.values())
total_passes = sum( total_passes = sum(sum(1 for r in cr["runs"] if r["pass"]) for cr in case_results.values())
sum(1 for r in cr["runs"] if r["pass"]) for cr in case_results.values()
)
total_json_dumps = sum( total_json_dumps = sum(
sum(1 for r in cr["runs"] if r.get("json_dump")) for cr in case_results.values() sum(1 for r in cr["runs"] if r.get("json_dump")) for cr in case_results.values()
) )
@@ -661,9 +682,7 @@ def _run_iteration(
if case_results if case_results
else 0 else 0
), ),
"per_case_pass_rates": { "per_case_pass_rates": {cid: cr["pass_rate"] for cid, cr in case_results.items()},
cid: cr["pass_rate"] for cid, cr in case_results.items()
},
}, },
} }
@@ -723,10 +742,10 @@ def _observe_and_update_optimizer(
client: OpenAI, client: OpenAI,
model: str, model: str,
optimizer_system: str, optimizer_system: str,
iterations: list[dict], iterations: list[dict[str, Any]],
) -> str: ) -> str:
"""Analyze optimizer behavior and return a modified OPTIMIZER_SYSTEM.""" """Analyze optimizer behavior and return a modified OPTIMIZER_SYSTEM."""
parts = [] parts: list[str] = []
for i in range(1, len(iterations)): for i in range(1, len(iterations)):
prev, curr = iterations[i - 1], iterations[i] prev, curr = iterations[i - 1], iterations[i]
prev_agg = prev.get("aggregate", {}) prev_agg = prev.get("aggregate", {})
@@ -744,8 +763,8 @@ def _observe_and_update_optimizer(
prev_rates = prev_agg.get("per_case_pass_rates", {}) prev_rates = prev_agg.get("per_case_pass_rates", {})
curr_rates = curr_agg.get("per_case_pass_rates", {}) curr_rates = curr_agg.get("per_case_pass_rates", {})
improved = [] improved: list[str] = []
regressed = [] regressed: list[str] = []
for case_id in set(prev_rates) | set(curr_rates): for case_id in set(prev_rates) | set(curr_rates):
p = prev_rates.get(case_id, 0) p = prev_rates.get(case_id, 0)
c = curr_rates.get(case_id, 0) c = curr_rates.get(case_id, 0)
@@ -764,7 +783,7 @@ def _observe_and_update_optimizer(
# Summarize what the optimizer's output looked like (without showing # Summarize what the optimizer's output looked like (without showing
# full developer messages, which cause the observer to mimic them) # full developer messages, which cause the observer to mimic them)
behavior_notes = [] behavior_notes: list[str] = []
for it in iterations[-3:]: for it in iterations[-3:]:
idx = it.get("iteration", "?") idx = it.get("iteration", "?")
prompt = it.get("prompt", "") prompt = it.get("prompt", "")
@@ -772,7 +791,7 @@ def _observe_and_update_optimizer(
has_bullets = "- " in prompt or "* " in prompt has_bullets = "- " in prompt or "* " in prompt
has_numbers = bool(re.search(r"^\d+\.", prompt, re.MULTILINE)) has_numbers = bool(re.search(r"^\d+\.", prompt, re.MULTILINE))
has_headers = "**" in prompt or "##" in prompt has_headers = "**" in prompt or "##" in prompt
notes = [] notes: list[str] = []
if has_bullets or has_numbers: if has_bullets or has_numbers:
notes.append("used bullet/numbered lists") notes.append("used bullet/numbered lists")
if has_headers: if has_headers:
@@ -791,7 +810,7 @@ def _observe_and_update_optimizer(
f"```\n{optimizer_system}\n```\n\n" f"```\n{optimizer_system}\n```\n\n"
f"## What the Rewriter Produced (do NOT mimic this)\n" f"## What the Rewriter Produced (do NOT mimic this)\n"
+ "\n".join(behavior_notes) + "\n".join(behavior_notes)
+ f"\n\n## Iteration History\n" + "\n\n## Iteration History\n"
+ "\n".join(parts) + "\n".join(parts)
) )
@@ -835,14 +854,14 @@ def _propose_prompt_modification(
client: OpenAI, client: OpenAI,
model: str, model: str,
current_prompt: str, current_prompt: str,
test_cases: list[dict], test_cases: list[dict[str, Any]],
iteration_result: dict, iteration_result: dict[str, Any],
history: list[dict], history: list[dict[str, Any]],
optimizer_system: str = OPTIMIZER_SYSTEM, optimizer_system: str = OPTIMIZER_SYSTEM,
) -> str: ) -> str:
"""Use the model to propose a new prompt based on evaluation results.""" """Use the model to propose a new prompt based on evaluation results."""
# Build summary of results # Build summary of results
summary_parts = [] summary_parts: list[str] = []
for case_id, case_result in iteration_result["cases"].items(): for case_id, case_result in iteration_result["cases"].items():
case_def = next((c for c in test_cases if c["id"] == case_id), None) case_def = next((c for c in test_cases if c["id"] == case_id), None)
if not case_def: if not case_def:
@@ -858,7 +877,7 @@ def _propose_prompt_modification(
) )
# Build history summary (last 3 iterations) # Build history summary (last 3 iterations)
history_parts = [] history_parts: list[str] = []
for h in history[-3:]: for h in history[-3:]:
agg = h.get("aggregate", {}) agg = h.get("aggregate", {})
history_parts.append( history_parts.append(
@@ -926,7 +945,7 @@ def run_optimization(
model: str | None, model: str | None,
test_file: str, test_file: str,
initial_prompt: str | None = None, initial_prompt: str | None = None,
n_runs: int = 3, n_runs: int | None = 3,
max_iterations: int = 5, max_iterations: int = 5,
temperature: float = 0.7, temperature: float = 0.7,
max_tokens: int = 32768, max_tokens: int = 32768,
@@ -934,7 +953,7 @@ def run_optimization(
output_file: str = "eval_results.json", output_file: str = "eval_results.json",
context_window: int = 131072, context_window: int = 131072,
verbose: bool = False, verbose: bool = False,
): ) -> dict[str, Any]:
"""Main optimization loop.""" """Main optimization loop."""
client = OpenAI( client = OpenAI(
base_url=base_url, base_url=base_url,
@@ -946,9 +965,9 @@ def run_optimization(
# Load test cases # Load test cases
with open(test_file) as f: with open(test_file) as f:
suite = json.load(f) suite: dict[str, Any] = json.load(f)
cases = suite["cases"] cases: list[dict[str, Any]] = suite["cases"]
for i, case in enumerate(cases): for i, case in enumerate(cases):
if "id" not in case: if "id" not in case:
raise SystemExit(f"Test case {i} missing required 'id' field") raise SystemExit(f"Test case {i} missing required 'id' field")
@@ -956,8 +975,7 @@ def run_optimization(
raise SystemExit(f"Test case '{case.get('id', i)}' missing 'user_prompt'") raise SystemExit(f"Test case '{case.get('id', i)}' missing 'user_prompt'")
defaults = suite.get("defaults", {}) defaults = suite.get("defaults", {})
# Precedence: CLI arg (non-None) > tests.json defaults > code default (3) # Precedence: CLI arg (non-None) > tests.json defaults > code default (3)
if n_runs is None: resolved_n_runs: int = n_runs if n_runs is not None else int(defaults.get("n_runs", 3))
n_runs = defaults.get("n_runs", 3)
# Get initial prompt # Get initial prompt
if initial_prompt is None: if initial_prompt is None:
@@ -965,7 +983,7 @@ def run_optimization(
tmp = ChatSession( tmp = ChatSession(
client=client, client=client,
model=model, model=model,
persona=None, ui=NullUI(),
instructions=None, instructions=None,
temperature=temperature, temperature=temperature,
max_tokens=max_tokens, max_tokens=max_tokens,
@@ -973,9 +991,7 @@ def run_optimization(
reasoning_effort=reasoning_effort, reasoning_effort=reasoning_effort,
context_window=context_window, context_window=context_window,
) )
initial_prompt = next( initial_prompt = next(m["content"] for m in tmp.system_messages if m["role"] == "developer")
m["content"] for m in tmp.system_messages if m["role"] == "developer"
)
# Strip memory reminder — it's a runtime artifact, not part of the prompt # Strip memory reminder — it's a runtime artifact, not part of the prompt
initial_prompt = re.sub( initial_prompt = re.sub(
r"\n*REMINDER: You currently have \d+ memories stored\..*$", r"\n*REMINDER: You currently have \d+ memories stored\..*$",
@@ -984,13 +1000,13 @@ def run_optimization(
).strip() ).strip()
current_prompt = initial_prompt current_prompt = initial_prompt
results = { results: dict[str, Any] = {
"meta": { "meta": {
"model": model, "model": model,
"base_url": base_url, "base_url": base_url,
"started": datetime.now().isoformat(), "started": datetime.now().isoformat(),
"test_suite": test_file, "test_suite": test_file,
"n_runs_default": n_runs, "n_runs_default": resolved_n_runs,
}, },
"iterations": [], "iterations": [],
} }
@@ -1007,7 +1023,7 @@ def run_optimization(
model=model, model=model,
system_prompt=current_prompt, system_prompt=current_prompt,
cases=cases, cases=cases,
n_runs=n_runs, n_runs=resolved_n_runs,
temperature=temperature, temperature=temperature,
max_tokens=max_tokens, max_tokens=max_tokens,
reasoning_effort=reasoning_effort, reasoning_effort=reasoning_effort,
@@ -1092,10 +1108,7 @@ def run_optimization(
if new_prompt != current_prompt: if new_prompt != current_prompt:
diff = _simple_diff(current_prompt, new_prompt) diff = _simple_diff(current_prompt, new_prompt)
print( print(f"Prompt modified ({len(current_prompt)} -> {len(new_prompt)} chars)")
f"Prompt modified "
f"({len(current_prompt)} -> {len(new_prompt)} chars)"
)
if diff: if diff:
print(diff) print(diff)
iter_result["prompt_diff"] = diff iter_result["prompt_diff"] = diff
@@ -1126,7 +1139,7 @@ def _detect_model(client: OpenAI) -> str:
# ─── CLI ───────────────────────────────────────────────────────────────────── # ─── CLI ─────────────────────────────────────────────────────────────────────
def main(): def main() -> None:
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
description="Prompt optimization and evaluation for turnstone", description="Prompt optimization and evaluation for turnstone",
formatter_class=argparse.RawDescriptionHelpFormatter, formatter_class=argparse.RawDescriptionHelpFormatter,
+1 -1
View File
@@ -6,6 +6,6 @@ commands and subscribe to progress.
""" """
from turnstone.mq.broker import MessageBroker, RedisBroker from turnstone.mq.broker import MessageBroker, RedisBroker
from turnstone.mq.client import TurnstoneClient, TurnResult from turnstone.mq.client import TurnResult, TurnstoneClient
__all__ = ["MessageBroker", "RedisBroker", "TurnstoneClient", "TurnResult"] __all__ = ["MessageBroker", "RedisBroker", "TurnstoneClient", "TurnResult"]
+33 -64
View File
@@ -9,6 +9,7 @@ Run as: ``turnstone-bridge --server-url http://localhost:8080``
from __future__ import annotations from __future__ import annotations
import contextlib
import json import json
import logging import logging
import os import os
@@ -16,7 +17,7 @@ import socket
import threading import threading
import time import time
import uuid import uuid
from collections.abc import Callable, Iterator from typing import TYPE_CHECKING, Any
import httpx import httpx
@@ -46,12 +47,13 @@ from turnstone.mq.protocol import (
WorkstreamRenameEvent, WorkstreamRenameEvent,
) )
if TYPE_CHECKING:
from collections.abc import Callable, Iterator
log = logging.getLogger("turnstone.mq.bridge") log = logging.getLogger("turnstone.mq.bridge")
# Server's default safe tools (auto-approved without user confirmation) # Server's default safe tools (auto-approved without user confirmation)
DEFAULT_SAFE_TOOLS = frozenset( DEFAULT_SAFE_TOOLS = frozenset(["read_file", "search", "man", "remember", "recall", "forget"])
["read_file", "search", "man", "remember", "recall", "forget"]
)
def _default_node_id() -> str: def _default_node_id() -> str:
@@ -85,7 +87,7 @@ class Bridge:
node_id: str = "", node_id: str = "",
heartbeat_ttl: int = 60, heartbeat_ttl: int = 60,
auth_token: str = "", auth_token: str = "",
): ) -> None:
self._server_url = server_url.rstrip("/") self._server_url = server_url.rstrip("/")
self._broker = broker or RedisBroker() self._broker = broker or RedisBroker()
self._approval_timeout = approval_timeout self._approval_timeout = approval_timeout
@@ -99,9 +101,7 @@ class Bridge:
headers: dict[str, str] = {} headers: dict[str, str] = {}
if auth_token: if auth_token:
headers["Authorization"] = f"Bearer {auth_token}" headers["Authorization"] = f"Bearer {auth_token}"
self._http = httpx.Client( self._http = httpx.Client(base_url=self._server_url, timeout=30, headers=headers)
base_url=self._server_url, timeout=30, headers=headers
)
# Protected by _lock — accessed from main, global SSE, and per-ws SSE threads # Protected by _lock — accessed from main, global SSE, and per-ws SSE threads
self._lock = threading.Lock() self._lock = threading.Lock()
@@ -160,9 +160,7 @@ class Bridge:
self._dispatch(msg) self._dispatch(msg)
except Exception as exc: except Exception as exc:
log.error("Failed to process inbound message: %s", exc) log.error("Failed to process inbound message: %s", exc)
self._publish_global( self._publish_global(ErrorEvent(message=f"Failed to process message: {exc}"))
ErrorEvent(message=f"Failed to process message: {exc}")
)
def _dispatch(self, msg: InboundMessage) -> None: def _dispatch(self, msg: InboundMessage) -> None:
# Messages that need routing (have ws_id or target_node) # Messages that need routing (have ws_id or target_node)
@@ -272,9 +270,7 @@ class Bridge:
def _handle_command(self, msg: InboundMessage) -> None: def _handle_command(self, msg: InboundMessage) -> None:
ws_id = getattr(msg, "ws_id", "") ws_id = getattr(msg, "ws_id", "")
command = getattr(msg, "command", "") command = getattr(msg, "command", "")
resp = self._http.post( resp = self._http.post("/api/command", json={"command": command, "ws_id": ws_id})
"/api/command", json={"command": command, "ws_id": ws_id}
)
data = resp.json() data = resp.json()
self._publish_ws( self._publish_ws(
ws_id, ws_id,
@@ -356,7 +352,7 @@ class Bridge:
) )
) )
return "" return ""
ws_id = data["ws_id"] ws_id: str = data["ws_id"]
ws_name = data.get("name", "") ws_name = data.get("name", "")
self._broker.set_ws_owner(ws_id, self._node_id) self._broker.set_ws_owner(ws_id, self._node_id)
@@ -425,21 +421,15 @@ class Bridge:
log.debug("WS SSE reconnecting (%s): %s", ws_id, exc) log.debug("WS SSE reconnecting (%s): %s", ws_id, exc)
time.sleep(2) time.sleep(2)
def _handle_ws_event(self, ws_id: str, data: dict) -> None: def _handle_ws_event(self, ws_id: str, data: dict[str, Any]) -> None:
etype = data.get("type", "") etype = data.get("type", "")
if etype == "content": if etype == "content":
self._publish_ws( self._publish_ws(ws_id, ContentEvent(ws_id=ws_id, text=data.get("text", "")))
ws_id, ContentEvent(ws_id=ws_id, text=data.get("text", ""))
)
elif etype == "reasoning": elif etype == "reasoning":
self._publish_ws( self._publish_ws(ws_id, ReasoningEvent(ws_id=ws_id, text=data.get("text", "")))
ws_id, ReasoningEvent(ws_id=ws_id, text=data.get("text", ""))
)
elif etype == "tool_info": elif etype == "tool_info":
self._publish_ws( self._publish_ws(ws_id, ToolInfoEvent(ws_id=ws_id, items=data.get("items", [])))
ws_id, ToolInfoEvent(ws_id=ws_id, items=data.get("items", []))
)
elif etype == "approve_request": elif etype == "approve_request":
self._handle_approval(ws_id, data) self._handle_approval(ws_id, data)
elif etype == "plan_review": elif etype == "plan_review":
@@ -467,17 +457,13 @@ class Bridge:
), ),
) )
elif etype == "error": elif etype == "error":
self._publish_ws( self._publish_ws(ws_id, ErrorEvent(ws_id=ws_id, message=data.get("message", "")))
ws_id, ErrorEvent(ws_id=ws_id, message=data.get("message", ""))
)
elif etype == "info": elif etype == "info":
self._publish_ws( self._publish_ws(ws_id, InfoEvent(ws_id=ws_id, message=data.get("message", "")))
ws_id, InfoEvent(ws_id=ws_id, message=data.get("message", ""))
)
elif etype == "stream_end": elif etype == "stream_end":
self._publish_ws(ws_id, StreamEndEvent(ws_id=ws_id)) self._publish_ws(ws_id, StreamEndEvent(ws_id=ws_id))
def _handle_approval(self, ws_id: str, data: dict) -> None: def _handle_approval(self, ws_id: str, data: dict[str, Any]) -> None:
"""Handle an approval request — auto-approve or forward to client.""" """Handle an approval request — auto-approve or forward to client."""
items = data.get("items", []) items = data.get("items", [])
@@ -506,9 +492,7 @@ class Bridge:
) )
def _wait_approval() -> None: def _wait_approval() -> None:
raw_resp = self._broker.pop_response( raw_resp = self._broker.pop_response(request_id, timeout=self._approval_timeout)
request_id, timeout=self._approval_timeout
)
if raw_resp: if raw_resp:
resp_msg = InboundMessage.from_json(raw_resp) resp_msg = InboundMessage.from_json(raw_resp)
approved = getattr(resp_msg, "approved", False) approved = getattr(resp_msg, "approved", False)
@@ -524,7 +508,7 @@ class Bridge:
threading.Thread(target=_wait_approval, daemon=True).start() threading.Thread(target=_wait_approval, daemon=True).start()
def _handle_plan_review(self, ws_id: str, data: dict) -> None: def _handle_plan_review(self, ws_id: str, data: dict[str, Any]) -> None:
"""Handle a plan review request — auto-approve or forward to client.""" """Handle a plan review request — auto-approve or forward to client."""
with self._lock: with self._lock:
if self._ws_auto_approve.get(ws_id): if self._ws_auto_approve.get(ws_id):
@@ -542,20 +526,14 @@ class Bridge:
) )
def _wait_plan() -> None: def _wait_plan() -> None:
raw_resp = self._broker.pop_response( raw_resp = self._broker.pop_response(request_id, timeout=self._approval_timeout)
request_id, timeout=self._approval_timeout
)
if raw_resp: if raw_resp:
resp_msg = InboundMessage.from_json(raw_resp) resp_msg = InboundMessage.from_json(raw_resp)
feedback = getattr(resp_msg, "feedback", "") feedback = getattr(resp_msg, "feedback", "")
self._http.post( self._http.post("/api/plan", json={"feedback": feedback, "ws_id": ws_id})
"/api/plan", json={"feedback": feedback, "ws_id": ws_id}
)
else: else:
log.warning("Plan review timeout for ws %s — rejecting", ws_id) log.warning("Plan review timeout for ws %s — rejecting", ws_id)
self._http.post( self._http.post("/api/plan", json={"feedback": "reject", "ws_id": ws_id})
"/api/plan", json={"feedback": "reject", "ws_id": ws_id}
)
threading.Thread(target=_wait_plan, daemon=True).start() threading.Thread(target=_wait_plan, daemon=True).start()
@@ -565,7 +543,7 @@ class Bridge:
approved: bool, approved: bool,
feedback: str | None = None, feedback: str | None = None,
) -> None: ) -> None:
body: dict = {"approved": approved, "ws_id": ws_id} body: dict[str, Any] = {"approved": approved, "ws_id": ws_id}
if feedback: if feedback:
body["feedback"] = feedback body["feedback"] = feedback
self._http.post("/api/approve", json=body) self._http.post("/api/approve", json=body)
@@ -592,7 +570,7 @@ class Bridge:
log.debug("Global SSE reconnecting: %s", exc) log.debug("Global SSE reconnecting: %s", exc)
time.sleep(2) time.sleep(2)
def _handle_global_event(self, data: dict) -> None: def _handle_global_event(self, data: dict[str, Any]) -> None:
etype = data.get("type", "") etype = data.get("type", "")
ws_id = data.get("ws_id", "") ws_id = data.get("ws_id", "")
@@ -617,17 +595,11 @@ class Bridge:
with self._lock: with self._lock:
cid = self._active_sends.pop(ws_id, None) cid = self._active_sends.pop(ws_id, None)
if cid: if cid:
self._publish_ws( self._publish_ws(ws_id, TurnCompleteEvent(ws_id=ws_id, correlation_id=cid))
ws_id, TurnCompleteEvent(ws_id=ws_id, correlation_id=cid)
)
elif etype == "ws_rename": elif etype == "ws_rename":
self._publish_global( self._publish_global(WorkstreamRenameEvent(ws_id=ws_id, name=data.get("name", "")))
WorkstreamRenameEvent(ws_id=ws_id, name=data.get("name", "")) self._publish_cluster(WorkstreamRenameEvent(ws_id=ws_id, name=data.get("name", "")))
)
self._publish_cluster(
WorkstreamRenameEvent(ws_id=ws_id, name=data.get("name", ""))
)
elif etype == "ws_closed": elif etype == "ws_closed":
self._publish_global(WorkstreamClosedEvent(ws_id=ws_id)) self._publish_global(WorkstreamClosedEvent(ws_id=ws_id))
@@ -655,9 +627,7 @@ class Bridge:
def _handle_list_nodes(self, msg: InboundMessage) -> None: def _handle_list_nodes(self, msg: InboundMessage) -> None:
nodes = self._broker.list_nodes() nodes = self._broker.list_nodes()
self._publish_global( self._publish_global(NodeListEvent(correlation_id=msg.correlation_id, nodes=nodes))
NodeListEvent(correlation_id=msg.correlation_id, nodes=nodes)
)
# -- publish helpers ----------------------------------------------------- # -- publish helpers -----------------------------------------------------
@@ -677,14 +647,13 @@ class Bridge:
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def _iter_sse_data(resp: httpx.Response) -> Iterator[dict]: def _iter_sse_data(resp: httpx.Response) -> Iterator[dict[str, Any]]:
"""Yield parsed JSON dicts from an SSE stream.""" """Yield parsed JSON dicts from an SSE stream."""
for line in resp.iter_lines(): for line in resp.iter_lines():
if line.startswith("data: "): if line.startswith("data: "):
try: with contextlib.suppress(json.JSONDecodeError):
yield json.loads(line[6:]) data: dict[str, Any] = json.loads(line[6:])
except json.JSONDecodeError: yield data
pass
# SSE keepalive comments (lines starting with ':') are ignored # SSE keepalive comments (lines starting with ':') are ignored
+26 -19
View File
@@ -7,9 +7,14 @@ RedisBroker is the default provider.
from __future__ import annotations from __future__ import annotations
import contextlib
import json import json
import threading from typing import TYPE_CHECKING, Any, Protocol, cast
from typing import Callable, Protocol
if TYPE_CHECKING:
from collections.abc import Callable
import redis as _redis_t
class MessageBroker(Protocol): class MessageBroker(Protocol):
@@ -74,11 +79,11 @@ class MessageBroker(Protocol):
"""Remove workstream ownership (on close).""" """Remove workstream ownership (on close)."""
... ...
def register_node(self, node_id: str, metadata: dict, ttl: int = 60) -> None: def register_node(self, node_id: str, metadata: dict[str, Any], ttl: int = 60) -> None:
"""Register or refresh a node's heartbeat with metadata.""" """Register or refresh a node's heartbeat with metadata."""
... ...
def list_nodes(self) -> list[dict]: def list_nodes(self) -> list[dict[str, Any]]:
"""List all active nodes (those with unexpired heartbeats).""" """List all active nodes (those with unexpired heartbeats)."""
... ...
@@ -117,12 +122,12 @@ class RedisBroker:
prefix: str = "turnstone", prefix: str = "turnstone",
password: str | None = None, password: str | None = None,
response_ttl: int = 600, response_ttl: int = 600,
): ) -> None:
import redis import redis
self._prefix = prefix self._prefix = prefix
self._response_ttl = response_ttl self._response_ttl = response_ttl
self._pool = redis.ConnectionPool( self._pool: _redis_t.ConnectionPool = redis.ConnectionPool(
host=host, host=host,
port=port, port=port,
db=db, db=db,
@@ -130,9 +135,12 @@ class RedisBroker:
decode_responses=True, decode_responses=True,
retry_on_timeout=True, retry_on_timeout=True,
) )
self._redis = redis.Redis(connection_pool=self._pool) self._redis: _redis_t.Redis[str] = cast(
"_redis_t.Redis[str]",
redis.Redis(connection_pool=self._pool),
)
self._pubsub = self._redis.pubsub(ignore_subscribe_messages=True) self._pubsub = self._redis.pubsub(ignore_subscribe_messages=True)
self._listener_thread: threading.Thread | None = None self._listener_thread: Any = None
self._running = True self._running = True
# -- inbound queue ------------------------------------------------------- # -- inbound queue -------------------------------------------------------
@@ -161,11 +169,12 @@ class RedisBroker:
self._redis.publish(channel, event) self._redis.publish(channel, event)
def subscribe_outbound(self, channel: str, callback: Callable[[str], None]) -> None: def subscribe_outbound(self, channel: str, callback: Callable[[str], None]) -> None:
self._pubsub.subscribe(**{channel: lambda msg: callback(msg["data"])}) def _handler(msg: dict[str, Any]) -> None:
callback(msg["data"])
self._pubsub.subscribe(**{channel: _handler})
if self._listener_thread is None or not self._listener_thread.is_alive(): if self._listener_thread is None or not self._listener_thread.is_alive():
self._listener_thread = self._pubsub.run_in_thread( self._listener_thread = self._pubsub.run_in_thread(sleep_time=0.1, daemon=True)
sleep_time=0.1, daemon=True
)
def unsubscribe_outbound(self, channel: str) -> None: def unsubscribe_outbound(self, channel: str) -> None:
self._pubsub.unsubscribe(channel) self._pubsub.unsubscribe(channel)
@@ -197,19 +206,19 @@ class RedisBroker:
def del_ws_owner(self, ws_id: str) -> None: def del_ws_owner(self, ws_id: str) -> None:
self._redis.delete(f"{self._prefix}:ws:{ws_id}") self._redis.delete(f"{self._prefix}:ws:{ws_id}")
def register_node(self, node_id: str, metadata: dict, ttl: int = 60) -> None: def register_node(self, node_id: str, metadata: dict[str, Any], ttl: int = 60) -> None:
key = f"{self._prefix}:node:{node_id}" key = f"{self._prefix}:node:{node_id}"
self._redis.set(key, json.dumps(metadata), ex=ttl) self._redis.set(key, json.dumps(metadata), ex=ttl)
def list_nodes(self) -> list[dict]: def list_nodes(self) -> list[dict[str, Any]]:
pattern = f"{self._prefix}:node:*" pattern = f"{self._prefix}:node:*"
prefix_len = len(f"{self._prefix}:node:") prefix_len = len(f"{self._prefix}:node:")
nodes = [] nodes: list[dict[str, Any]] = []
for key in self._redis.scan_iter(match=pattern, count=100): for key in self._redis.scan_iter(match=pattern, count=100):
raw = self._redis.get(key) raw = self._redis.get(key)
if raw: if raw:
try: try:
meta = json.loads(raw) meta: dict[str, Any] = json.loads(raw)
except json.JSONDecodeError: except json.JSONDecodeError:
meta = {} meta = {}
meta["node_id"] = key[prefix_len:] meta["node_id"] = key[prefix_len:]
@@ -230,8 +239,6 @@ class RedisBroker:
if self._listener_thread is not None: if self._listener_thread is not None:
self._listener_thread.stop() self._listener_thread.stop()
self._listener_thread = None self._listener_thread = None
try: with contextlib.suppress(Exception):
self._pubsub.close() self._pubsub.close()
except Exception:
pass
self._pool.disconnect() self._pool.disconnect()
+13 -23
View File
@@ -17,7 +17,7 @@ from __future__ import annotations
import threading import threading
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Callable from typing import TYPE_CHECKING, Any
from turnstone.mq.broker import MessageBroker, RedisBroker from turnstone.mq.broker import MessageBroker, RedisBroker
from turnstone.mq.protocol import ( from turnstone.mq.protocol import (
@@ -38,6 +38,9 @@ from turnstone.mq.protocol import (
WorkstreamCreatedEvent, WorkstreamCreatedEvent,
) )
if TYPE_CHECKING:
from collections.abc import Callable
@dataclass @dataclass
class TurnResult: class TurnResult:
@@ -76,7 +79,7 @@ class TurnstoneClient:
broker: MessageBroker | None = None, broker: MessageBroker | None = None,
prefix: str = "turnstone", prefix: str = "turnstone",
**redis_kwargs: object, **redis_kwargs: object,
): ) -> None:
"""Create a client. """Create a client.
Pass ``broker`` for a custom broker, or provide Redis kwargs Pass ``broker`` for a custom broker, or provide Redis kwargs
@@ -112,7 +115,7 @@ class TurnstoneClient:
auto_approve_tools=auto_approve_tools or [], auto_approve_tools=auto_approve_tools or [],
target_node=target_node, target_node=target_node,
) )
node = target_node or (self._broker.get_ws_owner(ws_id) if ws_id else "") node = target_node or (self._broker.get_ws_owner(ws_id) if ws_id else "") or ""
self._broker.push_inbound(msg.to_json(), node_id=node) self._broker.push_inbound(msg.to_json(), node_id=node)
return msg.correlation_id return msg.correlation_id
@@ -157,7 +160,7 @@ class TurnstoneClient:
self._broker.push_inbound(msg.to_json()) self._broker.push_inbound(msg.to_json())
return msg.correlation_id return msg.correlation_id
def list_nodes(self) -> list[dict]: def list_nodes(self) -> list[dict[str, Any]]:
"""List active bridge nodes (reads directly from broker).""" """List active bridge nodes (reads directly from broker)."""
return self._broker.list_nodes() return self._broker.list_nodes()
@@ -234,15 +237,10 @@ class TurnstoneClient:
if on_event: if on_event:
on_event(event) on_event(event)
if ( if isinstance(event, WorkstreamCreatedEvent) and event.correlation_id == cid:
isinstance(event, WorkstreamCreatedEvent)
and event.correlation_id == cid
):
actual_ws_id = event.ws_id actual_ws_id = event.ws_id
result.ws_id = event.ws_id result.ws_id = event.ws_id
self._broker.subscribe_outbound( self._broker.subscribe_outbound(f"{self._prefix}:events:{actual_ws_id}", _on_ws)
f"{self._prefix}:events:{actual_ws_id}", _on_ws
)
def _on_ws(raw: str) -> None: def _on_ws(raw: str) -> None:
event = OutboundEvent.from_json(raw) event = OutboundEvent.from_json(raw)
@@ -263,12 +261,10 @@ class TurnstoneClient:
# Subscribe BEFORE pushing — ensures we don't miss early events # Subscribe BEFORE pushing — ensures we don't miss early events
self._broker.subscribe_outbound(f"{self._prefix}:events:global", _on_global) self._broker.subscribe_outbound(f"{self._prefix}:events:global", _on_global)
if actual_ws_id: if actual_ws_id:
self._broker.subscribe_outbound( self._broker.subscribe_outbound(f"{self._prefix}:events:{actual_ws_id}", _on_ws)
f"{self._prefix}:events:{actual_ws_id}", _on_ws
)
# Now push the message (route to target node or ws owner if known) # Now push the message (route to target node or ws owner if known)
node = target_node or (self._broker.get_ws_owner(ws_id) if ws_id else "") node = target_node or (self._broker.get_ws_owner(ws_id) if ws_id else "") or ""
self._broker.push_inbound(msg.to_json(), node_id=node) self._broker.push_inbound(msg.to_json(), node_id=node)
done.wait(timeout=timeout) done.wait(timeout=timeout)
@@ -290,10 +286,7 @@ class TurnstoneClient:
ws_id: str = "", ws_id: str = "",
) -> None: ) -> None:
"""Subscribe to events for a specific workstream or global events.""" """Subscribe to events for a specific workstream or global events."""
if ws_id: channel = f"{self._prefix}:events:{ws_id}" if ws_id else f"{self._prefix}:events:global"
channel = f"{self._prefix}:events:{ws_id}"
else:
channel = f"{self._prefix}:events:global"
def _cb(raw: str) -> None: def _cb(raw: str) -> None:
event = OutboundEvent.from_json(raw) event = OutboundEvent.from_json(raw)
@@ -303,10 +296,7 @@ class TurnstoneClient:
def unsubscribe(self, ws_id: str = "") -> None: def unsubscribe(self, ws_id: str = "") -> None:
"""Unsubscribe from a workstream or global channel.""" """Unsubscribe from a workstream or global channel."""
if ws_id: channel = f"{self._prefix}:events:{ws_id}" if ws_id else f"{self._prefix}:events:global"
channel = f"{self._prefix}:events:{ws_id}"
else:
channel = f"{self._prefix}:events:global"
self._broker.unsubscribe_outbound(channel) self._broker.unsubscribe_outbound(channel)
# -- lifecycle ----------------------------------------------------------- # -- lifecycle -----------------------------------------------------------
+20 -11
View File
@@ -10,8 +10,8 @@ from __future__ import annotations
import json import json
import time import time
import uuid import uuid
from dataclasses import asdict, dataclass, field from dataclasses import asdict, dataclass, field, fields
from typing import Any
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Inbound messages (client → bridge) # Inbound messages (client → bridge)
@@ -36,7 +36,7 @@ class InboundMessage:
klass = _INBOUND_REGISTRY.get(msg_type) klass = _INBOUND_REGISTRY.get(msg_type)
if klass is None: if klass is None:
raise ValueError(f"Unknown inbound message type: {msg_type!r}") raise ValueError(f"Unknown inbound message type: {msg_type!r}")
valid = {f for f in klass.__dataclass_fields__} valid = {f.name for f in fields(klass)}
return klass(**{k: v for k, v in data.items() if k in valid}) return klass(**{k: v for k, v in data.items() if k in valid})
@@ -146,7 +146,7 @@ class OutboundEvent:
data = json.loads(raw) data = json.loads(raw)
msg_type = data.get("type", "") msg_type = data.get("type", "")
klass = _OUTBOUND_REGISTRY.get(msg_type, OutboundEvent) klass = _OUTBOUND_REGISTRY.get(msg_type, OutboundEvent)
valid = {f for f in klass.__dataclass_fields__} valid = {f.name for f in fields(klass)}
return klass(**{k: v for k, v in data.items() if k in valid}) return klass(**{k: v for k, v in data.items() if k in valid})
@@ -180,7 +180,7 @@ class ToolInfoEvent(OutboundEvent):
"""Tool call info (auto-approved tools).""" """Tool call info (auto-approved tools)."""
type: str = "tool_info" type: str = "tool_info"
items: list = field(default_factory=list) items: list[dict[str, Any]] = field(default_factory=list)
@dataclass @dataclass
@@ -192,7 +192,7 @@ class ApprovalRequestEvent(OutboundEvent):
""" """
type: str = "approval_request" type: str = "approval_request"
items: list = field(default_factory=list) items: list[dict[str, Any]] = field(default_factory=list)
@dataclass @dataclass
@@ -275,7 +275,7 @@ class WorkstreamListEvent(OutboundEvent):
"""Workstream list response.""" """Workstream list response."""
type: str = "ws_list" type: str = "ws_list"
workstreams: list = field(default_factory=list) workstreams: list[dict[str, Any]] = field(default_factory=list)
@dataclass @dataclass
@@ -291,7 +291,7 @@ class HealthResponseEvent(OutboundEvent):
"""Health status response.""" """Health status response."""
type: str = "health_response" type: str = "health_response"
data: dict = field(default_factory=dict) data: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -315,7 +315,7 @@ class NodeListEvent(OutboundEvent):
"""List of active bridge nodes.""" """List of active bridge nodes."""
type: str = "node_list" type: str = "node_list"
nodes: list = field(default_factory=list) nodes: list[dict[str, Any]] = field(default_factory=list)
@dataclass @dataclass
@@ -336,8 +336,17 @@ class ClusterStateEvent(OutboundEvent):
# Type registries (built after all classes are defined) # Type registries (built after all classes are defined)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def _type_default(cls: type[Any]) -> str:
"""Return the default value of the 'type' field for a dataclass."""
for f in fields(cls):
if f.name == "type":
return f.default # type: ignore[return-value]
return ""
_INBOUND_REGISTRY: dict[str, type[InboundMessage]] = { _INBOUND_REGISTRY: dict[str, type[InboundMessage]] = {
cls.__dataclass_fields__["type"].default: cls _type_default(cls): cls
for cls in [ for cls in [
SendMessage, SendMessage,
ApproveMessage, ApproveMessage,
@@ -352,7 +361,7 @@ _INBOUND_REGISTRY: dict[str, type[InboundMessage]] = {
} }
_OUTBOUND_REGISTRY: dict[str, type[OutboundEvent]] = { _OUTBOUND_REGISTRY: dict[str, type[OutboundEvent]] = {
cls.__dataclass_fields__["type"].default: cls _type_default(cls): cls
for cls in [ for cls in [
AckEvent, AckEvent,
ContentEvent, ContentEvent,
+96 -113
View File
@@ -9,7 +9,10 @@ Supports multiple concurrent workstreams (tabs), each with independent
ChatSession and event streams. ChatSession and event streams.
""" """
from __future__ import annotations
import argparse import argparse
import contextlib
import json import json
import os import os
import queue import queue
@@ -17,17 +20,18 @@ import sys
import textwrap import textwrap
import threading import threading
import time import time
from http.server import HTTPServer, BaseHTTPRequestHandler from http.server import BaseHTTPRequestHandler, HTTPServer
from pathlib import Path from pathlib import Path
from socketserver import ThreadingMixIn from socketserver import ThreadingMixIn
from urllib.parse import urlparse, parse_qs from typing import Any
from urllib.parse import ParseResult, parse_qs, urlparse
from openai import OpenAI from openai import OpenAI
from turnstone.core.metrics import metrics as _metrics from turnstone.core.metrics import metrics as _metrics
from turnstone.core.session import ChatSession, SessionUI # noqa: F401 from turnstone.core.session import ChatSession, SessionUI # noqa: F401
from turnstone.core.tools import TOOLS # noqa: F401 — available for introspection from turnstone.core.tools import TOOLS # noqa: F401 — available for introspection
from turnstone.core.workstream import WorkstreamManager, WorkstreamState from turnstone.core.workstream import Workstream, WorkstreamManager
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Static assets — loaded once at startup from turnstone/ui/static/ # Static assets — loaded once at startup from turnstone/ui/static/
@@ -53,15 +57,15 @@ class WebUI:
# Shared global event queue for state-change broadcasts across all # Shared global event queue for state-change broadcasts across all
# workstreams. Set by main() before any WebUI instances are created. # workstreams. Set by main() before any WebUI instances are created.
_global_queue: queue.Queue | None = None _global_queue: queue.Queue[dict[str, Any]] | None = None
def __init__(self, ws_id: str = ""): def __init__(self, ws_id: str = "") -> None:
self.ws_id = ws_id self.ws_id = ws_id
self._event_queue: queue.Queue = queue.Queue() self._event_queue: queue.Queue[dict[str, Any]] = queue.Queue()
self._sse_generation = 0 # incremented on each new SSE connection self._sse_generation = 0 # incremented on each new SSE connection
self._approval_event = threading.Event() self._approval_event = threading.Event()
self._approval_result: tuple[bool, str | None] = (False, None) self._approval_result: tuple[bool, str | None] = (False, None)
self._pending_approval: dict | None = None # re-sent on SSE reconnect self._pending_approval: dict[str, Any] | None = None # re-sent on SSE reconnect
self._plan_event = threading.Event() self._plan_event = threading.Event()
self._plan_result: str = "" self._plan_result: str = ""
self.auto_approve = False self.auto_approve = False
@@ -76,10 +80,10 @@ class WebUI:
self._ws_current_activity: str = "" self._ws_current_activity: str = ""
self._ws_activity_state: str = "" # "tool" | "approval" | "thinking" | "" self._ws_activity_state: str = "" # "tool" | "approval" | "thinking" | ""
def _enqueue(self, data: dict): def _enqueue(self, data: dict[str, Any]) -> None:
self._event_queue.put(data) self._event_queue.put(data)
def _broadcast_state(self, state: str): def _broadcast_state(self, state: str) -> None:
"""Send a state-change event to the global SSE channel.""" """Send a state-change event to the global SSE channel."""
if WebUI._global_queue is not None: if WebUI._global_queue is not None:
with self._ws_lock: with self._ws_lock:
@@ -99,7 +103,7 @@ class WebUI:
} }
) )
def _broadcast_activity(self): def _broadcast_activity(self) -> None:
"""Send an activity-change event to the global SSE channel.""" """Send an activity-change event to the global SSE channel."""
if WebUI._global_queue is not None: if WebUI._global_queue is not None:
with self._ws_lock: with self._ws_lock:
@@ -116,33 +120,31 @@ class WebUI:
# --- SessionUI protocol --- # --- SessionUI protocol ---
def on_thinking_start(self): def on_thinking_start(self) -> None:
with self._ws_lock: with self._ws_lock:
self._ws_current_activity = "Thinking\u2026" self._ws_current_activity = "Thinking\u2026"
self._ws_activity_state = "thinking" self._ws_activity_state = "thinking"
self._broadcast_activity() self._broadcast_activity()
self._enqueue({"type": "thinking_start"}) self._enqueue({"type": "thinking_start"})
def on_thinking_stop(self): def on_thinking_stop(self) -> None:
self._enqueue({"type": "thinking_stop"}) self._enqueue({"type": "thinking_stop"})
def on_reasoning_token(self, text: str): def on_reasoning_token(self, text: str) -> None:
self._enqueue({"type": "reasoning", "text": text}) self._enqueue({"type": "reasoning", "text": text})
def on_content_token(self, text: str): def on_content_token(self, text: str) -> None:
self._enqueue({"type": "content", "text": text}) self._enqueue({"type": "content", "text": text})
def on_stream_end(self): def on_stream_end(self) -> None:
with self._ws_lock: with self._ws_lock:
self._ws_current_activity = "" self._ws_current_activity = ""
self._ws_activity_state = "" self._ws_activity_state = ""
self._broadcast_activity() self._broadcast_activity()
self._enqueue({"type": "stream_end"}) self._enqueue({"type": "stream_end"})
def approve_tools(self, items: list[dict]) -> tuple[bool, str | None]: def approve_tools(self, items: list[dict[str, Any]]) -> tuple[bool, str | None]:
pending = [ pending = [it for it in items if it.get("needs_approval") and not it.get("error")]
it for it in items if it.get("needs_approval") and not it.get("error")
]
# Always send tool info to the browser # Always send tool info to the browser
serialized = [] serialized = []
@@ -152,9 +154,7 @@ class WebUI:
"header": item.get("header", ""), "header": item.get("header", ""),
"preview": item.get("preview", ""), "preview": item.get("preview", ""),
"func_name": item.get("func_name", ""), "func_name": item.get("func_name", ""),
"approval_label": item.get( "approval_label": item.get("approval_label", item.get("func_name", "")),
"approval_label", item.get("func_name", "")
),
"needs_approval": item.get("needs_approval", False), "needs_approval": item.get("needs_approval", False),
"error": item.get("error"), "error": item.get("error"),
} }
@@ -166,9 +166,7 @@ class WebUI:
label = first.get("func_name", "") label = first.get("func_name", "")
preview = first.get("preview", "")[:80] preview = first.get("preview", "")[:80]
with self._ws_lock: with self._ws_lock:
self._ws_current_activity = ( self._ws_current_activity = f"\u2699 {label}: {preview}" if label else ""
f"\u2699 {label}: {preview}" if label else ""
)
self._ws_activity_state = "tool" if label else "" self._ws_activity_state = "tool" if label else ""
self._broadcast_activity() self._broadcast_activity()
self._enqueue({"type": "tool_info", "items": serialized}) self._enqueue({"type": "tool_info", "items": serialized})
@@ -179,9 +177,7 @@ class WebUI:
label = first_pending.get("func_name", "") label = first_pending.get("func_name", "")
preview = first_pending.get("preview", "")[:60] preview = first_pending.get("preview", "")[:60]
with self._ws_lock: with self._ws_lock:
self._ws_current_activity = ( self._ws_current_activity = f"\u23f3 Awaiting approval: {label} \u2014 {preview}"
f"\u23f3 Awaiting approval: {label} \u2014 {preview}"
)
self._ws_activity_state = "approval" self._ws_activity_state = "approval"
self._broadcast_activity() self._broadcast_activity()
@@ -203,7 +199,7 @@ class WebUI:
return approved, feedback return approved, feedback
def on_tool_result(self, name: str, output: str): def on_tool_result(self, name: str, output: str) -> None:
_metrics.record_tool_call(name) _metrics.record_tool_call(name)
with self._ws_lock: with self._ws_lock:
self._ws_tool_calls[name] = self._ws_tool_calls.get(name, 0) + 1 self._ws_tool_calls[name] = self._ws_tool_calls.get(name, 0) + 1
@@ -212,19 +208,15 @@ class WebUI:
self._broadcast_activity() self._broadcast_activity()
self._enqueue({"type": "tool_result", "name": name, "output": output}) self._enqueue({"type": "tool_result", "name": name, "output": output})
def on_status(self, usage: dict, context_window: int, effort: str): def on_status(self, usage: dict[str, Any], context_window: int, effort: str) -> None:
total_tok = usage["prompt_tokens"] + usage["completion_tokens"] total_tok = usage["prompt_tokens"] + usage["completion_tokens"]
pct = total_tok / context_window * 100 if context_window > 0 else 0 pct = total_tok / context_window * 100 if context_window > 0 else 0
_metrics.record_tokens(usage["prompt_tokens"], usage["completion_tokens"]) _metrics.record_tokens(usage["prompt_tokens"], usage["completion_tokens"])
_metrics.record_context_ratio( _metrics.record_context_ratio(total_tok / context_window if context_window > 0 else 0.0)
total_tok / context_window if context_window > 0 else 0.0
)
with self._ws_lock: with self._ws_lock:
self._ws_prompt_tokens += usage["prompt_tokens"] self._ws_prompt_tokens += usage["prompt_tokens"]
self._ws_completion_tokens += usage["completion_tokens"] self._ws_completion_tokens += usage["completion_tokens"]
self._ws_context_ratio = ( self._ws_context_ratio = total_tok / context_window if context_window > 0 else 0.0
total_tok / context_window if context_window > 0 else 0.0
)
self._enqueue( self._enqueue(
{ {
"type": "status", "type": "status",
@@ -243,29 +235,27 @@ class WebUI:
self._plan_event.wait() self._plan_event.wait()
return self._plan_result return self._plan_result
def on_info(self, message: str): def on_info(self, message: str) -> None:
self._enqueue({"type": "info", "message": message}) self._enqueue({"type": "info", "message": message})
def on_error(self, message: str): def on_error(self, message: str) -> None:
_metrics.record_error() _metrics.record_error()
self._enqueue({"type": "error", "message": message}) self._enqueue({"type": "error", "message": message})
def on_state_change(self, state: str): def on_state_change(self, state: str) -> None:
self._broadcast_state(state) self._broadcast_state(state)
def on_rename(self, name: str): def on_rename(self, name: str) -> None:
"""Update the workstream's display name and broadcast to all clients.""" """Update the workstream's display name and broadcast to all clients."""
if WebUI._global_queue is not None: if WebUI._global_queue is not None:
WebUI._global_queue.put( WebUI._global_queue.put({"type": "ws_rename", "ws_id": self.ws_id, "name": name})
{"type": "ws_rename", "ws_id": self.ws_id, "name": name}
)
def resolve_approval(self, approved: bool, feedback: str | None = None): def resolve_approval(self, approved: bool, feedback: str | None = None) -> None:
"""Called by the HTTP handler when the user approves/denies.""" """Called by the HTTP handler when the user approves/denies."""
self._approval_result = (approved, feedback) self._approval_result = (approved, feedback)
self._approval_event.set() self._approval_event.set()
def resolve_plan(self, feedback: str): def resolve_plan(self, feedback: str) -> None:
"""Called by the HTTP handler when the user responds to a plan.""" """Called by the HTTP handler when the user responds to a plan."""
self._plan_result = feedback self._plan_result = feedback
self._plan_event.set() self._plan_event.set()
@@ -276,7 +266,9 @@ class WebUI:
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def _build_history(session, has_pending_approval: bool = False) -> list[dict]: def _build_history(
session: ChatSession, has_pending_approval: bool = False
) -> list[dict[str, Any]]:
"""Build a history replay list from session messages. """Build a history replay list from session messages.
When ``has_pending_approval`` is True, the last assistant entry's When ``has_pending_approval`` is True, the last assistant entry's
@@ -312,10 +304,10 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
""" """
# Suppress default logging to stderr # Suppress default logging to stderr
def log_message(self, format, *args): def log_message(self, fmt: str, *args: Any) -> None: # noqa: N802
pass pass
def _set_headers(self, status=200, content_type="application/json"): def _set_headers(self, status: int = 200, content_type: str = "application/json") -> None:
self._response_status = status self._response_status = status
self.send_response(status) self.send_response(status)
self.send_header("Content-Type", content_type) self.send_header("Content-Type", content_type)
@@ -323,28 +315,30 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
self.send_header("Access-Control-Allow-Origin", "*") self.send_header("Access-Control-Allow-Origin", "*")
self.end_headers() self.end_headers()
def _read_body(self) -> dict: def _read_body(self) -> dict[str, Any]:
length = int(self.headers.get("Content-Length", 0)) length = int(self.headers.get("Content-Length", 0))
if length == 0: if length == 0:
return {} return {}
raw = self.rfile.read(length) raw = self.rfile.read(length)
try: try:
return json.loads(raw.decode("utf-8")) result: dict[str, Any] = json.loads(raw.decode("utf-8"))
return result
except (json.JSONDecodeError, UnicodeDecodeError, ValueError): except (json.JSONDecodeError, UnicodeDecodeError, ValueError):
return {} return {}
def _send_json(self, data: dict, status=200): def _send_json(self, data: dict[str, Any], status: int = 200) -> None:
self._set_headers(status, "application/json") self._set_headers(status, "application/json")
self.wfile.write(json.dumps(data).encode("utf-8")) self.wfile.write(json.dumps(data).encode("utf-8"))
def _get_ws(self, ws_id: str | None): def _get_ws(self, ws_id: str | None) -> tuple[Workstream, WebUI] | tuple[None, None]:
"""Look up workstream by id. Returns (Workstream, WebUI) or (None, None).""" """Look up workstream by id. Returns (Workstream, WebUI) or (None, None)."""
if not ws_id: if not ws_id:
return None, None return None, None
mgr: WorkstreamManager = self.server.workstreams # type: ignore[attr-defined] mgr: WorkstreamManager = self.server.workstreams # type: ignore[attr-defined]
ws = mgr.get(ws_id) ws = mgr.get(ws_id)
if ws: if ws and ws.ui:
return ws, ws.ui ui: WebUI = ws.ui # type: ignore[assignment]
return ws, ui
return None, None return None, None
def _check_auth(self, method: str, path: str) -> bool: def _check_auth(self, method: str, path: str) -> bool:
@@ -354,14 +348,12 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
auth_config = self.server.auth_config # type: ignore[attr-defined] auth_config = self.server.auth_config # type: ignore[attr-defined]
auth_header = self.headers.get("Authorization") auth_header = self.headers.get("Authorization")
cookie_header = self.headers.get("Cookie") cookie_header = self.headers.get("Cookie")
allowed, status, msg = check_request( allowed, status, msg = check_request(auth_config, method, path, auth_header, cookie_header)
auth_config, method, path, auth_header, cookie_header
)
if not allowed: if not allowed:
self._send_json({"error": msg}, status) self._send_json({"error": msg}, status)
return allowed return allowed
def do_GET(self): def do_GET(self) -> None:
_t0 = time.monotonic() _t0 = time.monotonic()
self._response_status = 200 self._response_status = 200
parsed = urlparse(self.path) parsed = urlparse(self.path)
@@ -374,7 +366,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
"GET", parsed.path, self._response_status, time.monotonic() - _t0 "GET", parsed.path, self._response_status, time.monotonic() - _t0
) )
def _do_GET(self, parsed): def _do_GET(self, parsed: ParseResult) -> None: # noqa: N802
if parsed.path == "/": if parsed.path == "/":
self._set_headers(200, "text/html; charset=utf-8") self._set_headers(200, "text/html; charset=utf-8")
self.wfile.write(_HTML.encode("utf-8")) self.wfile.write(_HTML.encode("utf-8"))
@@ -420,6 +412,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
break break
# Send connected event with model info # Send connected event with model info
assert ws.session is not None
session: ChatSession = ws.session session: ChatSession = ws.session
connected_data = json.dumps( connected_data = json.dumps(
{ {
@@ -428,22 +421,20 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
"skip_permissions": ui.auto_approve, "skip_permissions": ui.auto_approve,
} }
) )
self.wfile.write(f"data: {connected_data}\n\n".encode("utf-8")) self.wfile.write(f"data: {connected_data}\n\n".encode())
self.wfile.flush() self.wfile.flush()
# Send conversation history for replay # Send conversation history for replay
history = _build_history( history = _build_history(session, has_pending_approval=ui._pending_approval is not None)
session, has_pending_approval=ui._pending_approval is not None
)
if history: if history:
history_data = json.dumps({"type": "history", "messages": history}) history_data = json.dumps({"type": "history", "messages": history})
self.wfile.write(f"data: {history_data}\n\n".encode("utf-8")) self.wfile.write(f"data: {history_data}\n\n".encode())
self.wfile.flush() self.wfile.flush()
# Re-inject a pending approval request if one was interrupted by a tab switch. # Re-inject a pending approval request if one was interrupted by a tab switch.
if ui._pending_approval is not None: if ui._pending_approval is not None:
pa_data = json.dumps(ui._pending_approval) pa_data = json.dumps(ui._pending_approval)
self.wfile.write(f"data: {pa_data}\n\n".encode("utf-8")) self.wfile.write(f"data: {pa_data}\n\n".encode())
self.wfile.flush() self.wfile.flush()
# Long-running SSE loop # Long-running SSE loop
@@ -452,7 +443,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
try: try:
event = ui._event_queue.get(timeout=5) event = ui._event_queue.get(timeout=5)
data = json.dumps(event) data = json.dumps(event)
self.wfile.write(f"data: {data}\n\n".encode("utf-8")) self.wfile.write(f"data: {data}\n\n".encode())
self.wfile.flush() self.wfile.flush()
except queue.Empty: except queue.Empty:
# Send keepalive comment to prevent timeout # Send keepalive comment to prevent timeout
@@ -469,12 +460,10 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
self.send_header("Access-Control-Allow-Origin", "*") self.send_header("Access-Control-Allow-Origin", "*")
self.end_headers() self.end_headers()
gq: queue.Queue = self.server.global_queue # type: ignore[attr-defined]
# Each global SSE client gets its own consumer queue # Each global SSE client gets its own consumer queue
# (since queue.Queue is single-consumer, we fan out via a listener list) # (since queue.Queue is single-consumer, we fan out via a listener list)
client_queue: queue.Queue = queue.Queue(maxsize=500) client_queue: queue.Queue[dict[str, Any]] = queue.Queue(maxsize=500)
listeners: list = self.server.global_listeners # type: ignore[attr-defined] listeners: list[queue.Queue[dict[str, Any]]] = self.server.global_listeners # type: ignore[attr-defined]
listeners_lock: threading.Lock = self.server.global_listeners_lock # type: ignore[attr-defined] listeners_lock: threading.Lock = self.server.global_listeners_lock # type: ignore[attr-defined]
with listeners_lock: with listeners_lock:
listeners.append(client_queue) listeners.append(client_queue)
@@ -484,7 +473,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
try: try:
event = client_queue.get(timeout=5) event = client_queue.get(timeout=5)
data = json.dumps(event) data = json.dumps(event)
self.wfile.write(f"data: {data}\n\n".encode("utf-8")) self.wfile.write(f"data: {data}\n\n".encode())
self.wfile.flush() self.wfile.flush()
except queue.Empty: except queue.Empty:
self.wfile.write(b": keepalive\n\n") self.wfile.write(b": keepalive\n\n")
@@ -540,7 +529,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
self._set_headers(404, "text/plain") self._set_headers(404, "text/plain")
self.wfile.write(b"Not found") self.wfile.write(b"Not found")
def do_POST(self): def do_POST(self) -> None:
_t0 = time.monotonic() _t0 = time.monotonic()
self._response_status = 200 self._response_status = 200
try: try:
@@ -552,7 +541,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
"POST", self.path, self._response_status, time.monotonic() - _t0 "POST", self.path, self._response_status, time.monotonic() - _t0
) )
def _do_POST(self): def _do_POST(self) -> None: # noqa: N802
if self.path == "/api/send": if self.path == "/api/send":
body = self._read_body() body = self._read_body()
message = body.get("message", "").strip() message = body.get("message", "").strip()
@@ -577,9 +566,13 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
self._send_json({"status": "busy"}) self._send_json({"status": "busy"})
return return
def run(): session = ws.session
assert session is not None
def run() -> None:
assert ui is not None
try: try:
ws.session.send(message) session.send(message)
except Exception as e: except Exception as e:
ui.on_error(f"Error: {e}") ui.on_error(f"Error: {e}")
ui._enqueue({"type": "stream_end"}) ui._enqueue({"type": "stream_end"})
@@ -633,6 +626,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
if not ws or not ui: if not ws or not ui:
self._send_json({"error": "Unknown workstream"}, 404) self._send_json({"error": "Unknown workstream"}, 404)
return return
assert ws.session is not None
try: try:
should_exit = ws.session.handle_command(command) should_exit = ws.session.handle_command(command)
@@ -669,6 +663,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
name=body.get("name", ""), name=body.get("name", ""),
ui_factory=lambda wid: WebUI(ws_id=wid), ui_factory=lambda wid: WebUI(ws_id=wid),
) )
assert isinstance(ws.ui, WebUI)
if skip or body.get("auto_approve", False): if skip or body.get("auto_approve", False):
ws.ui.auto_approve = True ws.ui.auto_approve = True
self._send_json({"ws_id": ws.id, "name": ws.name}) self._send_json({"ws_id": ws.id, "name": ws.name})
@@ -677,8 +672,8 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
elif self.path == "/api/workstreams/close": elif self.path == "/api/workstreams/close":
body = self._read_body() body = self._read_body()
ws_id = body.get("ws_id") ws_id = str(body.get("ws_id", ""))
mgr: WorkstreamManager = self.server.workstreams # type: ignore[attr-defined] mgr = self.server.workstreams # type: ignore[attr-defined]
if mgr.close(ws_id): if mgr.close(ws_id):
self._send_json({"status": "ok"}) self._send_json({"status": "ok"})
else: else:
@@ -697,9 +692,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
self.send_header("Set-Cookie", make_set_cookie(token)) self.send_header("Set-Cookie", make_set_cookie(token))
self.send_header("Cache-Control", "no-cache") self.send_header("Cache-Control", "no-cache")
self.end_headers() self.end_headers()
self.wfile.write( self.wfile.write(json.dumps({"status": "ok", "role": role}).encode("utf-8"))
json.dumps({"status": "ok", "role": role}).encode("utf-8")
)
else: else:
self._send_json({"error": "Invalid token"}, 401) self._send_json({"error": "Invalid token"}, 401)
@@ -717,11 +710,11 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
self._set_headers(404, "text/plain") self._set_headers(404, "text/plain")
self.wfile.write(b"Not found") self.wfile.write(b"Not found")
def _handle_health(self): def _handle_health(self) -> None:
"""Return server health status as JSON.""" """Return server health status as JSON."""
mgr: WorkstreamManager = self.server.workstreams # type: ignore[attr-defined] mgr: WorkstreamManager = self.server.workstreams # type: ignore[attr-defined]
wss = mgr.list_all() wss = mgr.list_all()
states: dict = { states: dict[str, int] = {
"idle": 0, "idle": 0,
"thinking": 0, "thinking": 0,
"running": 0, "running": 0,
@@ -740,7 +733,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
} }
self._send_json(data) self._send_json(data)
def _handle_dashboard(self): def _handle_dashboard(self) -> None:
"""Return enriched workstream data + aggregate stats for the dashboard.""" """Return enriched workstream data + aggregate stats for the dashboard."""
from turnstone.core.memory import get_session_name from turnstone.core.memory import get_session_name
@@ -795,11 +788,11 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
} }
) )
def _handle_metrics(self): def _handle_metrics(self) -> None:
"""Return Prometheus text exposition format metrics.""" """Return Prometheus text exposition format metrics."""
mgr: WorkstreamManager = self.server.workstreams # type: ignore[attr-defined] mgr: WorkstreamManager = self.server.workstreams # type: ignore[attr-defined]
wss = mgr.list_all() wss = mgr.list_all()
states: dict = { states: dict[str, int] = {
"idle": 0, "idle": 0,
"thinking": 0, "thinking": 0,
"running": 0, "running": 0,
@@ -832,7 +825,7 @@ class TurnstoneHTTPHandler(BaseHTTPRequestHandler):
self._set_headers(200, "text/plain; version=0.0.4; charset=utf-8") self._set_headers(200, "text/plain; version=0.0.4; charset=utf-8")
self.wfile.write(content.encode("utf-8")) self.wfile.write(content.encode("utf-8"))
def do_OPTIONS(self): def do_OPTIONS(self) -> None:
"""Handle CORS preflight.""" """Handle CORS preflight."""
self.send_response(200) self.send_response(200)
self.send_header("Access-Control-Allow-Origin", "*") self.send_header("Access-Control-Allow-Origin", "*")
@@ -886,23 +879,23 @@ def detect_model(client: OpenAI) -> str:
def _idle_cleanup_thread( def _idle_cleanup_thread(
mgr: WorkstreamManager, timeout_sec: float, global_queue: queue.Queue mgr: WorkstreamManager, timeout_sec: float, global_queue: queue.Queue[dict[str, Any]]
): ) -> None:
"""Periodically close IDLE workstreams that have been inactive too long.""" """Periodically close IDLE workstreams that have been inactive too long."""
check_every = min(300.0, timeout_sec / 4) # check at ¼ of timeout, max 5 min check_every = min(300.0, timeout_sec / 4) # check at ¼ of timeout, max 5 min
while True: while True:
time.sleep(check_every) time.sleep(check_every)
closed = mgr.close_idle(timeout_sec) closed = mgr.close_idle(timeout_sec)
for ws_id in closed: for ws_id in closed:
try: with contextlib.suppress(queue.Full):
global_queue.put_nowait({"type": "ws_closed", "ws_id": ws_id}) global_queue.put_nowait({"type": "ws_closed", "ws_id": ws_id})
except queue.Full:
pass
def _global_fanout_thread( def _global_fanout_thread(
source_queue: queue.Queue, listeners: list, lock: threading.Lock source_queue: queue.Queue[dict[str, Any]],
): listeners: list[queue.Queue[dict[str, Any]]],
lock: threading.Lock,
) -> None:
"""Reads events from the source queue and copies them to all listener queues.""" """Reads events from the source queue and copies them to all listener queues."""
while True: while True:
try: try:
@@ -910,10 +903,8 @@ def _global_fanout_thread(
with lock: with lock:
snapshot = list(listeners) snapshot = list(listeners)
for lq in snapshot: for lq in snapshot:
try: with contextlib.suppress(queue.Full):
lq.put_nowait(event) lq.put_nowait(event) # drop if a listener is backed up
except queue.Full:
pass # drop if a listener is backed up
except Exception: except Exception:
pass pass
@@ -923,7 +914,7 @@ def _global_fanout_thread(
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def main(): def main() -> None:
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
description="turnstone web server — browser-based chat UI.", description="turnstone web server — browser-based chat UI.",
formatter_class=argparse.RawDescriptionHelpFormatter, formatter_class=argparse.RawDescriptionHelpFormatter,
@@ -945,11 +936,6 @@ def main():
default=None, default=None,
help="Model name (default: auto-detect from server)", help="Model name (default: auto-detect from server)",
) )
parser.add_argument(
"--persona",
default=None,
help="Persona name injected as system message",
)
parser.add_argument( parser.add_argument(
"--instructions", "--instructions",
default=None, default=None,
@@ -1068,24 +1054,21 @@ def main():
) )
# Detect or use provided model # Detect or use provided model
if args.model: model = args.model or detect_model(client)
model = args.model
else:
model = detect_model(client)
# Set up global event queue for state-change broadcasts # Set up global event queue for state-change broadcasts
global_queue: queue.Queue = queue.Queue() global_queue: queue.Queue[dict[str, Any]] = queue.Queue()
global_listeners: list = [] global_listeners: list[queue.Queue[dict[str, Any]]] = []
global_listeners_lock = threading.Lock() global_listeners_lock = threading.Lock()
WebUI._global_queue = global_queue WebUI._global_queue = global_queue
# Session factory — captures shared config # Session factory — captures shared config
def session_factory(ui): def session_factory(ui: SessionUI | None) -> ChatSession:
assert ui is not None
return ChatSession( return ChatSession(
client=client, client=client,
model=model, model=model,
ui=ui, ui=ui,
persona=args.persona,
instructions=args.instructions, instructions=args.instructions,
temperature=args.temperature, temperature=args.temperature,
max_tokens=args.max_tokens, max_tokens=args.max_tokens,
@@ -1104,10 +1087,12 @@ def main():
name="default", name="default",
ui_factory=lambda wid: WebUI(ws_id=wid), ui_factory=lambda wid: WebUI(ws_id=wid),
) )
assert isinstance(ws.ui, WebUI)
if args.skip_permissions: if args.skip_permissions:
ws.ui.auto_approve = True ws.ui.auto_approve = True
# Handle --resume # Handle --resume
assert ws.session is not None
if args.resume: if args.resume:
from turnstone.core.memory import resolve_session from turnstone.core.memory import resolve_session
@@ -1157,8 +1142,6 @@ def main():
print(f"turnstone web server running on http://{args.host}:{args.port}") print(f"turnstone web server running on http://{args.host}:{args.port}")
print(f"Model: {model}") print(f"Model: {model}")
if args.persona:
print(f"Persona: {args.persona}")
print("Press Ctrl+C to stop.") print("Press Ctrl+C to stop.")
try: try:
+3 -4
View File
@@ -7,6 +7,7 @@ import asyncio
import json import json
import logging import logging
import sys import sys
from typing import Any
from turnstone.sim.cluster import SimCluster from turnstone.sim.cluster import SimCluster
from turnstone.sim.config import SimConfig from turnstone.sim.config import SimConfig
@@ -85,9 +86,7 @@ def main() -> None:
parser.add_argument("--redis-password", default=None) parser.add_argument("--redis-password", default=None)
parser.add_argument("--redis-db", type=int, default=0) parser.add_argument("--redis-db", type=int, default=0)
parser.add_argument("--prefix", default="turnstone") parser.add_argument("--prefix", default="turnstone")
parser.add_argument( parser.add_argument("--seed", type=int, default=None, help="Random seed for reproducibility")
"--seed", type=int, default=None, help="Random seed for reproducibility"
)
parser.add_argument("--metrics-file", default="", help="Write JSON metrics to file") parser.add_argument("--metrics-file", default="", help="Write JSON metrics to file")
parser.add_argument( parser.add_argument(
"--log-level", "--log-level",
@@ -152,7 +151,7 @@ async def _run(config: SimConfig) -> None:
await cluster.stop() await cluster.stop()
def _print_report(report: dict, config: SimConfig) -> None: def _print_report(report: dict[str, Any], config: SimConfig) -> None:
lat = report.get("latency", {}) lat = report.get("latency", {})
tp = report.get("throughput", {}) tp = report.get("throughput", {})
util = report.get("utilization", {}) util = report.get("utilization", {})
+27 -25
View File
@@ -7,14 +7,17 @@ import logging
import math import math
import time import time
from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ThreadPoolExecutor
from typing import TYPE_CHECKING, Any, cast
import redis import redis
from turnstone.mq.broker import RedisBroker from turnstone.mq.broker import RedisBroker
from turnstone.sim.config import SimConfig
from turnstone.sim.metrics import MetricsCollector from turnstone.sim.metrics import MetricsCollector
from turnstone.sim.node import SimNode from turnstone.sim.node import SimNode
if TYPE_CHECKING:
from turnstone.sim.config import SimConfig
log = logging.getLogger("turnstone.sim.cluster") log = logging.getLogger("turnstone.sim.cluster")
# How many node queues a single dispatcher watches via one BLPOP call. # How many node queues a single dispatcher watches via one BLPOP call.
@@ -29,16 +32,15 @@ class PooledBroker(RedisBroker):
pool: redis.ConnectionPool, pool: redis.ConnectionPool,
prefix: str = "turnstone", prefix: str = "turnstone",
response_ttl: int = 600, response_ttl: int = 600,
): ) -> None:
# Bypass RedisBroker.__init__ — set up manually with the shared pool. # Bypass RedisBroker.__init__ — set up manually with the shared pool.
import threading
self._prefix = prefix self._prefix = prefix
self._response_ttl = response_ttl self._response_ttl = response_ttl
self._pool = pool self._pool: redis.ConnectionPool = pool
self._redis = redis.Redis(connection_pool=pool) self._redis: redis.Redis[str] = cast("redis.Redis[str]", redis.Redis(connection_pool=pool))
self._pubsub = self._redis.pubsub(ignore_subscribe_messages=True) self._pubsub = self._redis.pubsub(ignore_subscribe_messages=True)
self._listener_thread: threading.Thread | None = None self._listener_thread: Any = None
self._running = True self._running = True
def close(self) -> None: def close(self) -> None:
@@ -56,11 +58,11 @@ class InboundDispatcher:
def __init__( def __init__(
self, self,
redis_client: redis.Redis, redis_client: redis.Redis[str],
node_ids: list[str], node_ids: list[str],
nodes: dict[str, SimNode], nodes: dict[str, SimNode],
prefix: str, prefix: str,
): ) -> None:
self._redis = redis_client self._redis = redis_client
self._node_ids = node_ids self._node_ids = node_ids
self._nodes = nodes self._nodes = nodes
@@ -72,9 +74,7 @@ class InboundDispatcher:
self._keys.append(f"{prefix}:inbound") self._keys.append(f"{prefix}:inbound")
# Pre-compute key → node_id mapping # Pre-compute key → node_id mapping
self._key_to_node: dict[str, str] = { self._key_to_node: dict[str, str] = {f"{prefix}:inbound:{nid}": nid for nid in node_ids}
f"{prefix}:inbound:{nid}": nid for nid in node_ids
}
async def run(self) -> None: async def run(self) -> None:
while self._running: while self._running:
@@ -146,15 +146,16 @@ class SimCluster:
await cluster.stop() await cluster.stop()
""" """
def __init__(self, config: SimConfig): def __init__(self, config: SimConfig) -> None:
self._config = config self._config = config
self._metrics = MetricsCollector() self._metrics = MetricsCollector()
self._nodes: dict[str, SimNode] = {} self._nodes: dict[str, SimNode] = {}
self._node_order: list[str] = [] self._node_order: list[str] = []
self._dispatchers: list[InboundDispatcher] = [] self._dispatchers: list[InboundDispatcher] = []
self._tasks: list[asyncio.Task] = [] self._tasks: list[asyncio.Task[None]] = []
self._pool: redis.ConnectionPool | None = None self._pool: redis.ConnectionPool | None = None
self._redis_client: redis.Redis | None = None self._redis_client: redis.Redis[str] | None = None
self._executor: ThreadPoolExecutor | None = None
self._running = True self._running = True
@property @property
@@ -174,7 +175,7 @@ class SimCluster:
self._executor = ThreadPoolExecutor(max_workers=64) self._executor = ThreadPoolExecutor(max_workers=64)
# Shared Redis pool # Shared Redis pool
self._pool = redis.ConnectionPool( pool: redis.ConnectionPool = redis.ConnectionPool(
host=self._config.redis_host, host=self._config.redis_host,
port=self._config.redis_port, port=self._config.redis_port,
db=self._config.redis_db, db=self._config.redis_db,
@@ -183,13 +184,14 @@ class SimCluster:
retry_on_timeout=True, retry_on_timeout=True,
max_connections=64, max_connections=64,
) )
self._redis_client = redis.Redis(connection_pool=self._pool) self._pool = pool
self._redis_client = cast("redis.Redis[str]", redis.Redis(connection_pool=pool))
# Create nodes # Create nodes
for i in range(self._config.num_nodes): for i in range(self._config.num_nodes):
node_id = f"sim-{i:04d}" node_id = f"sim-{i:04d}"
broker = PooledBroker( broker = PooledBroker(
self._pool, pool,
prefix=self._config.prefix, prefix=self._config.prefix,
) )
node = SimNode(node_id, broker, self._config, self._metrics) node = SimNode(node_id, broker, self._config, self._metrics)
@@ -203,7 +205,7 @@ class SimCluster:
start = i * NODES_PER_DISPATCHER start = i * NODES_PER_DISPATCHER
batch_ids = all_ids[start : start + NODES_PER_DISPATCHER] batch_ids = all_ids[start : start + NODES_PER_DISPATCHER]
# Each dispatcher gets its own Redis client from the shared pool # Each dispatcher gets its own Redis client from the shared pool
client = redis.Redis(connection_pool=self._pool) client: redis.Redis[str] = cast("redis.Redis[str]", redis.Redis(connection_pool=pool))
dispatcher = InboundDispatcher( dispatcher = InboundDispatcher(
client, client,
batch_ids, batch_ids,
@@ -246,9 +248,7 @@ class SimCluster:
while self._running: while self._running:
await asyncio.sleep(self._config.metrics_interval) await asyncio.sleep(self._config.metrics_interval)
counts = { counts = {
nid: node.workstream_count nid: node.workstream_count for nid, node in self._nodes.items() if node._running
for nid, node in self._nodes.items()
if node._running
} }
self._metrics.snapshot_utilization(counts) self._metrics.snapshot_utilization(counts)
@@ -261,12 +261,14 @@ class SimCluster:
] ]
await asyncio.gather(*tasks, return_exceptions=True) await asyncio.gather(*tasks, return_exceptions=True)
assert self._redis_client is not None
redis_client = self._redis_client
registered = 0 registered = 0
deadline = time.monotonic() + 30 deadline = time.monotonic() + 30
while time.monotonic() < deadline: while time.monotonic() < deadline:
keys = await loop.run_in_executor( keys = await loop.run_in_executor(
self._executor, self._executor,
self._redis_client.keys, redis_client.keys,
f"{self._config.prefix}:node:sim-*", f"{self._config.prefix}:node:sim-*",
) )
registered = len(keys) registered = len(keys)
@@ -298,7 +300,7 @@ class SimCluster:
d.remove_node(node_id) d.remove_node(node_id)
log.info("Killed node %s", node_id) log.info("Killed node %s", node_id)
def report(self) -> dict: def report(self) -> dict[str, Any]:
"""Generate final metrics report.""" """Generate final metrics report."""
return self._metrics.summary() return self._metrics.summary()
@@ -312,8 +314,8 @@ class SimCluster:
for task in self._tasks: for task in self._tasks:
task.cancel() task.cancel()
await asyncio.gather(*self._tasks, return_exceptions=True) await asyncio.gather(*self._tasks, return_exceptions=True)
if hasattr(self, "_executor"): if self._executor is not None:
self._executor.shutdown(wait=False) self._executor.shutdown(wait=False)
if self._pool: if self._pool is not None:
self._pool.disconnect() self._pool.disconnect()
log.info("Cluster stopped") log.info("Cluster stopped")
+4 -2
View File
@@ -4,8 +4,10 @@ from __future__ import annotations
import asyncio import asyncio
import random import random
from typing import TYPE_CHECKING, Any
from turnstone.sim.config import SimConfig if TYPE_CHECKING:
from turnstone.sim.config import SimConfig
_WORD_POOL = [ _WORD_POOL = [
"the", "the",
@@ -66,7 +68,7 @@ class SimEngine:
async def simulate_llm_response( async def simulate_llm_response(
self, first_round: bool, turn_number: int self, first_round: bool, turn_number: int
) -> tuple[str, list[dict]]: ) -> tuple[str, list[dict[str, Any]]]:
"""Simulate an LLM response. """Simulate an LLM response.
Returns ``(content_text, tool_calls)`` where *tool_calls* may be Returns ``(content_text, tool_calls)`` where *tool_calls* may be
+3 -2
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import threading import threading
import time import time
from collections import defaultdict from collections import defaultdict
from typing import Any
class MetricsCollector: class MetricsCollector:
@@ -48,7 +49,7 @@ class MetricsCollector:
with self._lock: with self._lock:
self._ws_counts.append(dict(ws_counts)) self._ws_counts.append(dict(ws_counts))
def summary(self) -> dict: def summary(self) -> dict[str, Any]:
"""Generate final metrics report with percentiles and aggregates.""" """Generate final metrics report with percentiles and aggregates."""
with self._lock: with self._lock:
latencies = sorted(self._turn_latencies) latencies = sorted(self._turn_latencies)
@@ -60,7 +61,7 @@ class MetricsCollector:
duration = 0.0 duration = 0.0
# Utilization from latest snapshot # Utilization from latest snapshot
util: dict = {} util: dict[str, Any] = {}
if self._ws_counts: if self._ws_counts:
last = self._ws_counts[-1] last = self._ws_counts[-1]
counts = list(last.values()) counts = list(last.values())
+11 -6
View File
@@ -13,6 +13,7 @@ import json
import logging import logging
import time import time
import uuid import uuid
from typing import TYPE_CHECKING
from turnstone.mq.protocol import ( from turnstone.mq.protocol import (
AckEvent, AckEvent,
@@ -22,6 +23,7 @@ from turnstone.mq.protocol import (
HealthResponseEvent, HealthResponseEvent,
InboundMessage, InboundMessage,
NodeListEvent, NodeListEvent,
OutboundEvent,
StateChangeEvent, StateChangeEvent,
StatusEvent, StatusEvent,
StreamEndEvent, StreamEndEvent,
@@ -31,9 +33,12 @@ from turnstone.mq.protocol import (
WorkstreamCreatedEvent, WorkstreamCreatedEvent,
WorkstreamListEvent, WorkstreamListEvent,
) )
from turnstone.sim.config import SimConfig
from turnstone.sim.engine import SimEngine, ToolSimulationError from turnstone.sim.engine import SimEngine, ToolSimulationError
from turnstone.sim.metrics import MetricsCollector
if TYPE_CHECKING:
from turnstone.mq.broker import RedisBroker
from turnstone.sim.config import SimConfig
from turnstone.sim.metrics import MetricsCollector
log = logging.getLogger("turnstone.sim.node") log = logging.getLogger("turnstone.sim.node")
@@ -212,7 +217,7 @@ class SimNode:
def __init__( def __init__(
self, self,
node_id: str, node_id: str,
broker: object, broker: RedisBroker,
config: SimConfig, config: SimConfig,
metrics: MetricsCollector, metrics: MetricsCollector,
): ):
@@ -420,19 +425,19 @@ class SimNode:
# -- event publishing helpers -------------------------------------------- # -- event publishing helpers --------------------------------------------
def _publish_global(self, event: object) -> None: def _publish_global(self, event: OutboundEvent) -> None:
self._broker.publish_outbound( self._broker.publish_outbound(
f"{self._prefix}:events:global", f"{self._prefix}:events:global",
event.to_json(), event.to_json(),
) )
def _publish_ws(self, ws_id: str, event: object) -> None: def _publish_ws(self, ws_id: str, event: OutboundEvent) -> None:
self._broker.publish_outbound( self._broker.publish_outbound(
f"{self._prefix}:events:{ws_id}", f"{self._prefix}:events:{ws_id}",
event.to_json(), event.to_json(),
) )
def _publish_cluster(self, event: object) -> None: def _publish_cluster(self, event: OutboundEvent) -> None:
self._broker.publish_outbound( self._broker.publish_outbound(
f"{self._prefix}:events:cluster", f"{self._prefix}:events:cluster",
event.to_json(), event.to_json(),
+6 -5
View File
@@ -5,15 +5,15 @@ from __future__ import annotations
import asyncio import asyncio
import logging import logging
import time import time
from typing import TYPE_CHECKING, Protocol from typing import TYPE_CHECKING, Any, Protocol
from turnstone.mq.broker import RedisBroker from turnstone.mq.broker import RedisBroker
from turnstone.mq.protocol import SendMessage from turnstone.mq.protocol import SendMessage
from turnstone.sim.config import SimConfig
from turnstone.sim.metrics import MetricsCollector
if TYPE_CHECKING: if TYPE_CHECKING:
from turnstone.sim.cluster import SimCluster from turnstone.sim.cluster import SimCluster
from turnstone.sim.config import SimConfig
from turnstone.sim.metrics import MetricsCollector
log = logging.getLogger("turnstone.sim.scenario") log = logging.getLogger("turnstone.sim.scenario")
@@ -157,6 +157,7 @@ class LifecycleScenario:
from turnstone.mq.protocol import ( from turnstone.mq.protocol import (
CloseWorkstreamMessage, CloseWorkstreamMessage,
CreateWorkstreamMessage, CreateWorkstreamMessage,
InboundMessage,
) )
broker = _make_broker(config) broker = _make_broker(config)
@@ -166,7 +167,7 @@ class LifecycleScenario:
# Phase 1: Create workstreams # Phase 1: Create workstreams
create_count = min(50, config.num_nodes * 2) create_count = min(50, config.num_nodes * 2)
for i in range(create_count): for i in range(create_count):
msg = CreateWorkstreamMessage( msg: InboundMessage = CreateWorkstreamMessage(
name=f"lifecycle-ws-{i}", name=f"lifecycle-ws-{i}",
auto_approve=True, auto_approve=True,
) )
@@ -225,7 +226,7 @@ def _make_broker(config: SimConfig) -> RedisBroker:
) )
SCENARIOS: dict[str, type] = { SCENARIOS: dict[str, type[Any]] = {
"steady": SteadyStateScenario, "steady": SteadyStateScenario,
"burst": BurstScenario, "burst": BurstScenario,
"node_failure": NodeFailureScenario, "node_failure": NodeFailureScenario,
+6 -6
View File
@@ -22,25 +22,25 @@ CYAN = "\033[36m" if _use_color else ""
GRAY = "\033[90m" if _use_color else "" GRAY = "\033[90m" if _use_color else ""
def red(s): def red(s: str) -> str:
return f"{RED}{s}{RESET}" return f"{RED}{s}{RESET}"
def yellow(s): def yellow(s: str) -> str:
return f"{YELLOW}{s}{RESET}" return f"{YELLOW}{s}{RESET}"
def dim(s): def dim(s: str) -> str:
return f"{DIM}{s}{RESET}" return f"{DIM}{s}{RESET}"
def bold(s): def bold(s: str) -> str:
return f"{BOLD}{s}{RESET}" return f"{BOLD}{s}{RESET}"
def cyan(s): def cyan(s: str) -> str:
return f"{CYAN}{s}{RESET}" return f"{CYAN}{s}{RESET}"
def green(s): def green(s: str) -> str:
return f"{GREEN}{s}{RESET}" return f"{GREEN}{s}{RESET}"
+2 -4
View File
@@ -13,7 +13,7 @@ class MarkdownRenderer:
(fenced code blocks) track state across lines. (fenced code blocks) track state across lines.
""" """
def __init__(self): def __init__(self) -> None:
self.in_code_block = False self.in_code_block = False
self._buf = "" self._buf = ""
@@ -52,9 +52,7 @@ class MarkdownRenderer:
# Inline formatting (order matters: bold before italic) # Inline formatting (order matters: bold before italic)
line = re.sub(r"\*\*(.+?)\*\*", f"{BOLD}\\1{RESET}", line) line = re.sub(r"\*\*(.+?)\*\*", f"{BOLD}\\1{RESET}", line)
line = re.sub(r"__(.+?)__", f"{BOLD}\\1{RESET}", line) line = re.sub(r"__(.+?)__", f"{BOLD}\\1{RESET}", line)
line = re.sub( line = re.sub(r"(?<!\*)\*(?!\*)(.+?)(?<!\*)\*(?!\*)", f"{ITALIC}\\1{RESET}", line)
r"(?<!\*)\*(?!\*)(.+?)(?<!\*)\*(?!\*)", f"{ITALIC}\\1{RESET}", line
)
line = re.sub(r"`(.+?)`", f"{CYAN}\\1{RESET}", line) line = re.sub(r"`(.+?)`", f"{CYAN}\\1{RESET}", line)
# Bullet lists — cyan bullet # Bullet lists — cyan bullet
+5 -5
View File
@@ -16,12 +16,12 @@ class Spinner:
self._stop_event = threading.Event() self._stop_event = threading.Event()
self._thread: threading.Thread | None = None self._thread: threading.Thread | None = None
def start(self): def start(self) -> None:
self._stop_event.clear() self._stop_event.clear()
self._thread = threading.Thread(target=self._spin, daemon=True) self._thread = threading.Thread(target=self._spin, daemon=True)
self._thread.start() self._thread.start()
def _spin(self): def _spin(self) -> None:
i = 0 i = 0
while not self._stop_event.wait(0.08): while not self._stop_event.wait(0.08):
frame = self._FRAMES[i % len(self._FRAMES)] frame = self._FRAMES[i % len(self._FRAMES)]
@@ -29,7 +29,7 @@ class Spinner:
sys.stderr.flush() sys.stderr.flush()
i += 1 i += 1
def stop(self): def stop(self) -> None:
if self._stop_event.is_set(): if self._stop_event.is_set():
return return
self._stop_event.set() self._stop_event.set()
@@ -39,9 +39,9 @@ class Spinner:
sys.stderr.write("\r\033[2K") sys.stderr.write("\r\033[2K")
sys.stderr.flush() sys.stderr.flush()
def __enter__(self): def __enter__(self) -> "Spinner":
self.start() self.start()
return self return self
def __exit__(self, *_): def __exit__(self, *_: object) -> None:
self.stop() self.stop()
+5 -5
View File
@@ -243,7 +243,7 @@ function toggleTheme() {
var current = document.documentElement.dataset.theme; var current = document.documentElement.dataset.theme;
var next = current === "light" ? "" : "light"; var next = current === "light" ? "" : "light";
document.documentElement.dataset.theme = next; document.documentElement.dataset.theme = next;
localStorage.setItem("pcode-theme", next || "dark"); localStorage.setItem("turnstone-theme", next || "dark");
updateThemeMenuItem(); updateThemeMenuItem();
} }
function updateThemeMenuItem() { function updateThemeMenuItem() {
@@ -265,7 +265,7 @@ function updateThemeMenuItem() {
); );
} }
(function () { (function () {
if (localStorage.getItem("pcode-theme") === "light") if (localStorage.getItem("turnstone-theme") === "light")
document.documentElement.dataset.theme = "light"; document.documentElement.dataset.theme = "light";
updateThemeMenuItem(); updateThemeMenuItem();
})(); })();
@@ -613,7 +613,7 @@ function switchTab(wsId) {
// Push history entry so back button can retrace tab navigation. // Push history entry so back button can retrace tab navigation.
if (!_historyNavigation) { if (!_historyNavigation) {
history.pushState({ pcode: "workstream", wsId: wsId }, ""); history.pushState({ turnstone: "workstream", wsId: wsId }, "");
} }
} }
@@ -1829,7 +1829,7 @@ authFetch("/api/workstreams")
} }
connectGlobalSSE(); connectGlobalSSE();
// Seed the history stack so back-from-workstream returns here. // Seed the history stack so back-from-workstream returns here.
history.replaceState({ pcode: "dashboard" }, ""); history.replaceState({ turnstone: "dashboard" }, "");
showDashboard(); showDashboard();
}); });
@@ -1837,7 +1837,7 @@ authFetch("/api/workstreams")
window.addEventListener("popstate", function (e) { window.addEventListener("popstate", function (e) {
_historyNavigation = true; _historyNavigation = true;
try { try {
if (e.state && e.state.pcode === "workstream") { if (e.state && e.state.turnstone === "workstream") {
// Navigating to a workstream state (forward, or back between tabs). // Navigating to a workstream state (forward, or back between tabs).
if (dashboardVisible) hideDashboard(); if (dashboardVisible) hideDashboard();
if (e.state.wsId && workstreams[e.state.wsId]) switchTab(e.state.wsId); if (e.state.wsId && workstreams[e.state.wsId]) switchTab(e.state.wsId);
+5 -2
View File
@@ -4,6 +4,9 @@
<meta charset="utf-8"> <meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1"> <meta name="viewport" content="width=device-width, initial-scale=1">
<title>turnstone</title> <title>turnstone</title>
<link rel="preconnect" href="https://fonts.googleapis.com">
<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin>
<link href="https://fonts.googleapis.com/css2?family=IBM+Plex+Mono:wght@400;500;600&family=Outfit:wght@400;500;600;700&display=swap" rel="stylesheet">
<link rel="stylesheet" href="/static/style.css"> <link rel="stylesheet" href="/static/style.css">
</head> </head>
<body> <body>
@@ -22,10 +25,10 @@
</button> </button>
</div> </div>
</div> </div>
<h1>pcode</h1> <h1>turnstone</h1>
<span id="model-name"></span> <span id="model-name"></span>
<span id="status-bar"></span> <span id="status-bar"></span>
<button id="logout-btn" onclick="logout()" style="display:none; background:none; border:1px solid var(--border); color:var(--fg-dim); border-radius:var(--radius); padding:4px 8px; cursor:pointer; font:inherit; font-size:12px;">logout</button> <button id="logout-btn" class="header-btn" onclick="logout()" style="display:none">logout</button>
</div> </div>
<div id="tab-bar" role="tablist"> <div id="tab-bar" role="tablist">
File diff suppressed because it is too large Load Diff