feat: add hash ring tables and storage protocol (migration 030)

Add three tables for the hash ring routing system:
- hash_ring_buckets: bucket-to-node assignments (65536 rows, rebalancer-managed)
- bucket_stats: per-bucket workstream counts (server-managed lifecycle counters)
- workstream_overrides: per-workstream routing pins (targeted/admin/pinned)

Add 10 storage protocol methods with SQLite and PostgreSQL implementations.
Wire bucket_stats lifecycle hooks into WorkstreamManager create/close/set_state.

Tables start empty — the rebalancer (Phase 3) seeds hash_ring_buckets on
first run. bucket_stats rows are upserted lazily on workstream lifecycle.
This commit is contained in:
Patrick Buckley
2026-03-30 08:52:31 -07:00
committed by Patrick Buckley
parent 2bb55590bf
commit 262a6a9918
8 changed files with 578 additions and 3 deletions
+135
View File
@@ -0,0 +1,135 @@
"""Tests for the hash ring routing storage methods."""
from __future__ import annotations
class TestHashRingBuckets:
def test_list_empty(self, storage):
assert storage.list_ring_buckets() == []
def test_seed_and_list(self, storage):
storage.seed_ring_buckets([(0, "node-a"), (1, "node-b"), (2, "node-a")])
rows = storage.list_ring_buckets()
assert len(rows) == 3
assert rows[0] == {"bucket": 0, "node_id": "node-a"}
assert rows[1] == {"bucket": 1, "node_id": "node-b"}
assert rows[2] == {"bucket": 2, "node_id": "node-a"}
def test_seed_idempotent(self, storage):
storage.seed_ring_buckets([(0, "node-a"), (1, "node-b")])
# Re-seed with conflicting assignment: should keep original
storage.seed_ring_buckets([(0, "node-x"), (2, "node-c")])
rows = storage.list_ring_buckets()
by_bucket = {r["bucket"]: r["node_id"] for r in rows}
assert by_bucket[0] == "node-a" # original preserved
assert by_bucket[1] == "node-b"
assert by_bucket[2] == "node-c" # new bucket added
def test_assign_buckets(self, storage):
storage.seed_ring_buckets([(0, "node-a"), (1, "node-a"), (2, "node-b")])
storage.assign_buckets([0, 1], "node-c")
rows = storage.list_ring_buckets()
by_bucket = {r["bucket"]: r["node_id"] for r in rows}
assert by_bucket[0] == "node-c"
assert by_bucket[1] == "node-c"
assert by_bucket[2] == "node-b"
def test_assign_returns_count(self, storage):
storage.seed_ring_buckets([(0, "node-a"), (1, "node-a")])
count = storage.assign_buckets([0, 1], "node-b")
assert count == 2
# Empty list returns 0
assert storage.assign_buckets([], "node-x") == 0
class TestBucketStats:
def test_increment_creates_row(self, storage):
storage.increment_bucket_count(42)
stats = storage.list_bucket_stats()
assert len(stats) == 1
assert stats[0]["bucket"] == 42
assert stats[0]["ws_count"] == 1
assert stats[0]["active_count"] == 0
def test_increment_active(self, storage):
storage.increment_bucket_count(10, active=True)
stats = storage.list_bucket_stats()
assert stats[0]["ws_count"] == 1
assert stats[0]["active_count"] == 1
# Increment again without active
storage.increment_bucket_count(10)
stats = storage.list_bucket_stats()
assert stats[0]["ws_count"] == 2
assert stats[0]["active_count"] == 1
def test_decrement(self, storage):
storage.increment_bucket_count(5, active=True)
storage.increment_bucket_count(5, active=True)
storage.decrement_bucket_count(5, active=True)
stats = storage.list_bucket_stats()
assert stats[0]["ws_count"] == 1
assert stats[0]["active_count"] == 1
def test_decrement_clamps_at_zero(self, storage):
storage.increment_bucket_count(7)
storage.decrement_bucket_count(7)
storage.decrement_bucket_count(7) # already at 0
stats = storage.list_bucket_stats()
# ws_count is 0, so should not appear (filter ws_count > 0)
assert len(stats) == 0
def test_adjust_active_only(self, storage):
storage.increment_bucket_count(20, active=True)
storage.increment_bucket_count(20, active=True)
# Decrease active without changing ws_count
storage.adjust_bucket_active(20, -1)
stats = storage.list_bucket_stats()
assert stats[0]["ws_count"] == 2
assert stats[0]["active_count"] == 1
# Clamp at zero
storage.adjust_bucket_active(20, -5)
stats = storage.list_bucket_stats()
assert stats[0]["active_count"] == 0
def test_list_sparse(self, storage):
storage.increment_bucket_count(100)
storage.increment_bucket_count(200)
storage.increment_bucket_count(300)
# Decrement 200 to zero
storage.decrement_bucket_count(200)
stats = storage.list_bucket_stats()
buckets = [s["bucket"] for s in stats]
assert 100 in buckets
assert 200 not in buckets
assert 300 in buckets
class TestWorkstreamOverrides:
def test_set_and_list(self, storage):
storage.set_workstream_override("ws-001", "node-a", reason="affinity")
overrides = storage.list_workstream_overrides()
assert len(overrides) == 1
assert overrides[0]["ws_id"] == "ws-001"
assert overrides[0]["node_id"] == "node-a"
assert overrides[0]["reason"] == "affinity"
def test_upsert(self, storage):
storage.set_workstream_override("ws-002", "node-a")
storage.set_workstream_override("ws-002", "node-b", reason="migration")
overrides = storage.list_workstream_overrides()
assert len(overrides) == 1
assert overrides[0]["node_id"] == "node-b"
assert overrides[0]["reason"] == "migration"
def test_delete(self, storage):
storage.set_workstream_override("ws-003", "node-a")
result = storage.delete_workstream_override("ws-003")
assert result is True
assert storage.list_workstream_overrides() == []
def test_delete_nonexistent(self, storage):
result = storage.delete_workstream_override("ws-nope")
assert result is False
def test_list_empty(self, storage):
assert storage.list_workstream_overrides() == []
+40
View File
@@ -110,6 +110,46 @@ def update_workstream_state(ws_id: str, state: str) -> None:
log.warning("Failed to update workstream state ws=%s state=%s", ws_id, state, exc_info=True)
# -- Hash ring bucket counts --------------------------------------------------
def _bucket_of(ws_id: str) -> int:
"""Extract bucket from ws_id prefix. First 4 hex chars -> 0-65535."""
return int(ws_id[:4], 16)
def increment_bucket_count(ws_id: str, active: bool = False) -> None:
"""Fire-and-forget bucket count increment."""
try:
get_storage().increment_bucket_count(_bucket_of(ws_id), active)
except Exception:
log.warning("bucket count increment failed for %s", ws_id[:8], exc_info=True)
def decrement_bucket_count(ws_id: str, active: bool = False) -> None:
"""Fire-and-forget bucket count decrement."""
try:
get_storage().decrement_bucket_count(_bucket_of(ws_id), active)
except Exception:
log.warning("bucket count decrement failed for %s", ws_id[:8], exc_info=True)
def adjust_bucket_active(ws_id: str, delta: int) -> None:
"""Fire-and-forget active count adjustment."""
try:
get_storage().adjust_bucket_active(_bucket_of(ws_id), delta)
except Exception:
log.warning("bucket active adjust failed for %s", ws_id[:8], exc_info=True)
def delete_workstream_override(ws_id: str) -> None:
"""Fire-and-forget override deletion."""
try:
get_storage().delete_workstream_override(ws_id)
except Exception:
log.warning("override delete failed for %s", ws_id[:8], exc_info=True)
def update_workstream_name(ws_id: str, name: str) -> None:
"""Update a workstream's display name."""
try:
+128
View File
@@ -11,9 +11,11 @@ from turnstone.core.log import get_logger
from turnstone.core.storage._schema import (
api_tokens,
audit_events,
bucket_stats,
channel_routes,
channel_users,
conversations,
hash_ring_buckets,
intent_verdicts,
mcp_servers,
metadata,
@@ -40,6 +42,7 @@ from turnstone.core.storage._schema import (
users,
watches,
workstream_config,
workstream_overrides,
workstreams,
)
from turnstone.core.storage._utils import (
@@ -1216,6 +1219,131 @@ class PostgreSQLBackend:
conn.commit()
return result.rowcount > 0
# -- Hash ring routing -----------------------------------------------------
def list_ring_buckets(self) -> list[dict[str, Any]]:
with self._engine.connect() as conn:
rows = conn.execute(
sa.select(hash_ring_buckets).order_by(hash_ring_buckets.c.bucket)
).fetchall()
return [dict(r._mapping) for r in rows]
def seed_ring_buckets(self, assignments: list[tuple[int, str]]) -> None:
from sqlalchemy.dialects.postgresql import insert as pg_insert
chunk_size = 500
with self._engine.connect() as conn:
for i in range(0, len(assignments), chunk_size):
chunk = assignments[i : i + chunk_size]
stmt = pg_insert(hash_ring_buckets).values(
[{"bucket": b, "node_id": n} for b, n in chunk]
)
stmt = stmt.on_conflict_do_nothing(index_elements=[hash_ring_buckets.c.bucket])
conn.execute(stmt)
conn.commit()
def assign_buckets(self, buckets: list[int], node_id: str) -> int:
if not buckets:
return 0
with self._engine.connect() as conn:
result = conn.execute(
sa.update(hash_ring_buckets)
.where(hash_ring_buckets.c.bucket.in_(buckets))
.values(node_id=node_id)
)
conn.commit()
return result.rowcount
def increment_bucket_count(self, bucket: int, active: bool = False) -> None:
from sqlalchemy.dialects.postgresql import insert as pg_insert
with self._engine.connect() as conn:
stmt = pg_insert(bucket_stats).values(
bucket=bucket,
ws_count=1,
active_count=1 if active else 0,
)
set_: dict[str, Any] = {"ws_count": bucket_stats.c.ws_count + 1}
if active:
set_["active_count"] = bucket_stats.c.active_count + 1
stmt = stmt.on_conflict_do_update(index_elements=[bucket_stats.c.bucket], set_=set_)
conn.execute(stmt)
conn.commit()
def decrement_bucket_count(self, bucket: int, active: bool = False) -> None:
vals: dict[str, Any] = {
"ws_count": sa.case(
(bucket_stats.c.ws_count > 0, bucket_stats.c.ws_count - 1),
else_=0,
)
}
if active:
vals["active_count"] = sa.case(
(bucket_stats.c.active_count > 0, bucket_stats.c.active_count - 1),
else_=0,
)
with self._engine.connect() as conn:
conn.execute(
sa.update(bucket_stats).where(bucket_stats.c.bucket == bucket).values(**vals)
)
conn.commit()
def adjust_bucket_active(self, bucket: int, delta: int) -> None:
with self._engine.connect() as conn:
conn.execute(
sa.update(bucket_stats)
.where(bucket_stats.c.bucket == bucket)
.values(
active_count=sa.case(
(
bucket_stats.c.active_count + sa.literal(delta) >= 0,
bucket_stats.c.active_count + sa.literal(delta),
),
else_=0,
)
)
)
conn.commit()
def list_bucket_stats(self) -> list[dict[str, Any]]:
with self._engine.connect() as conn:
rows = conn.execute(
sa.select(bucket_stats)
.where(bucket_stats.c.ws_count > 0)
.order_by(bucket_stats.c.bucket)
).fetchall()
return [dict(r._mapping) for r in rows]
def set_workstream_override(self, ws_id: str, node_id: str, reason: str = "targeted") -> None:
from sqlalchemy.dialects.postgresql import insert as pg_insert
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
stmt = pg_insert(workstream_overrides).values(
ws_id=ws_id, node_id=node_id, reason=reason, created=now
)
stmt = stmt.on_conflict_do_update(
index_elements=[workstream_overrides.c.ws_id],
set_={"node_id": node_id, "reason": reason},
)
with self._engine.connect() as conn:
conn.execute(stmt)
conn.commit()
def delete_workstream_override(self, ws_id: str) -> bool:
with self._engine.connect() as conn:
result = conn.execute(
sa.delete(workstream_overrides).where(workstream_overrides.c.ws_id == ws_id)
)
conn.commit()
return result.rowcount > 0
def list_workstream_overrides(self) -> list[dict[str, str]]:
with self._engine.connect() as conn:
rows = conn.execute(
sa.select(workstream_overrides).order_by(workstream_overrides.c.ws_id)
).fetchall()
return [dict(r._mapping) for r in rows]
# -- Roles -----------------------------------------------------------------
def create_role(
+42
View File
@@ -464,6 +464,48 @@ class StorageBackend(Protocol):
"""Remove a service registration. Returns True if existed."""
...
# -- Hash ring routing ---
def list_ring_buckets(self) -> list[dict[str, Any]]:
"""Return all rows from hash_ring_buckets. Empty if not seeded."""
...
def seed_ring_buckets(self, assignments: list[tuple[int, str]]) -> None:
"""Insert (bucket, node_id) rows. Idempotent (ON CONFLICT DO NOTHING)."""
...
def assign_buckets(self, buckets: list[int], node_id: str) -> int:
"""Reassign buckets to a node. Returns rows updated."""
...
def increment_bucket_count(self, bucket: int, active: bool = False) -> None:
"""Increment ws_count (and active_count if active) in bucket_stats. Upserts."""
...
def decrement_bucket_count(self, bucket: int, active: bool = False) -> None:
"""Decrement ws_count (and active_count if active). Clamps at zero."""
...
def adjust_bucket_active(self, bucket: int, delta: int) -> None:
"""Adjust active_count only (not ws_count). For state transitions."""
...
def list_bucket_stats(self) -> list[dict[str, Any]]:
"""Return all bucket_stats rows with ws_count > 0."""
...
def set_workstream_override(self, ws_id: str, node_id: str, reason: str = "targeted") -> None:
"""Pin a workstream to a specific node. Upserts."""
...
def delete_workstream_override(self, ws_id: str) -> bool:
"""Remove a pin. Returns True if one existed."""
...
def list_workstream_overrides(self) -> list[dict[str, str]]:
"""Return all overrides."""
...
# -- Roles (RBAC) ----------------------------------------------------------
def create_role(
+32
View File
@@ -232,6 +232,38 @@ services = sa.Table(
sa.Index("idx_services_type_heartbeat", services.c.service_type, services.c.last_heartbeat)
# ---------------------------------------------------------------------------
# Hash ring routing tables
# ---------------------------------------------------------------------------
hash_ring_buckets = sa.Table(
"hash_ring_buckets",
metadata,
sa.Column("bucket", sa.Integer, primary_key=True),
sa.Column("node_id", sa.Text, nullable=False),
)
sa.Index("idx_ring_buckets_node", hash_ring_buckets.c.node_id)
bucket_stats = sa.Table(
"bucket_stats",
metadata,
sa.Column("bucket", sa.Integer, primary_key=True),
sa.Column("ws_count", sa.Integer, nullable=False, server_default="0"),
sa.Column("active_count", sa.Integer, nullable=False, server_default="0"),
)
workstream_overrides = sa.Table(
"workstream_overrides",
metadata,
sa.Column("ws_id", sa.Text, primary_key=True),
sa.Column("node_id", sa.Text, nullable=False),
sa.Column("reason", sa.Text, nullable=False, server_default="targeted"),
sa.Column("created", sa.Text, nullable=False),
)
sa.Index("idx_ws_overrides_node", workstream_overrides.c.node_id)
# ---------------------------------------------------------------------------
# Governance tables — RBAC, orgs, policies, skills, usage, audit
# ---------------------------------------------------------------------------
+128
View File
@@ -11,9 +11,11 @@ from turnstone.core.log import get_logger
from turnstone.core.storage._schema import (
api_tokens,
audit_events,
bucket_stats,
channel_routes,
channel_users,
conversations,
hash_ring_buckets,
intent_verdicts,
mcp_servers,
metadata,
@@ -40,6 +42,7 @@ from turnstone.core.storage._schema import (
users,
watches,
workstream_config,
workstream_overrides,
workstreams,
)
from turnstone.core.storage._utils import (
@@ -1291,6 +1294,131 @@ class SQLiteBackend:
conn.commit()
return result.rowcount > 0
# -- Hash ring routing -----------------------------------------------------
def list_ring_buckets(self) -> list[dict[str, Any]]:
with self._engine.connect() as conn:
rows = conn.execute(
sa.select(hash_ring_buckets).order_by(hash_ring_buckets.c.bucket)
).fetchall()
return [dict(r._mapping) for r in rows]
def seed_ring_buckets(self, assignments: list[tuple[int, str]]) -> None:
from sqlalchemy.dialects.sqlite import insert as sqlite_insert
chunk_size = 500
with self._engine.connect() as conn:
for i in range(0, len(assignments), chunk_size):
chunk = assignments[i : i + chunk_size]
stmt = sqlite_insert(hash_ring_buckets).values(
[{"bucket": b, "node_id": n} for b, n in chunk]
)
stmt = stmt.on_conflict_do_nothing(index_elements=["bucket"])
conn.execute(stmt)
conn.commit()
def assign_buckets(self, buckets: list[int], node_id: str) -> int:
if not buckets:
return 0
with self._engine.connect() as conn:
result = conn.execute(
sa.update(hash_ring_buckets)
.where(hash_ring_buckets.c.bucket.in_(buckets))
.values(node_id=node_id)
)
conn.commit()
return result.rowcount
def increment_bucket_count(self, bucket: int, active: bool = False) -> None:
from sqlalchemy.dialects.sqlite import insert as sqlite_insert
with self._engine.connect() as conn:
stmt = sqlite_insert(bucket_stats).values(
bucket=bucket,
ws_count=1,
active_count=1 if active else 0,
)
set_ = {"ws_count": bucket_stats.c.ws_count + 1}
if active:
set_["active_count"] = bucket_stats.c.active_count + 1
stmt = stmt.on_conflict_do_update(index_elements=["bucket"], set_=set_)
conn.execute(stmt)
conn.commit()
def decrement_bucket_count(self, bucket: int, active: bool = False) -> None:
vals: dict[str, Any] = {
"ws_count": sa.case(
(bucket_stats.c.ws_count > 0, bucket_stats.c.ws_count - 1),
else_=0,
)
}
if active:
vals["active_count"] = sa.case(
(bucket_stats.c.active_count > 0, bucket_stats.c.active_count - 1),
else_=0,
)
with self._engine.connect() as conn:
conn.execute(
sa.update(bucket_stats).where(bucket_stats.c.bucket == bucket).values(**vals)
)
conn.commit()
def adjust_bucket_active(self, bucket: int, delta: int) -> None:
with self._engine.connect() as conn:
conn.execute(
sa.update(bucket_stats)
.where(bucket_stats.c.bucket == bucket)
.values(
active_count=sa.case(
(
bucket_stats.c.active_count + sa.literal(delta) >= 0,
bucket_stats.c.active_count + sa.literal(delta),
),
else_=0,
)
)
)
conn.commit()
def list_bucket_stats(self) -> list[dict[str, Any]]:
with self._engine.connect() as conn:
rows = conn.execute(
sa.select(bucket_stats)
.where(bucket_stats.c.ws_count > 0)
.order_by(bucket_stats.c.bucket)
).fetchall()
return [dict(r._mapping) for r in rows]
def set_workstream_override(self, ws_id: str, node_id: str, reason: str = "targeted") -> None:
from sqlalchemy.dialects.sqlite import insert as sqlite_insert
now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S")
stmt = sqlite_insert(workstream_overrides).values(
ws_id=ws_id, node_id=node_id, reason=reason, created=now
)
stmt = stmt.on_conflict_do_update(
index_elements=["ws_id"],
set_={"node_id": node_id, "reason": reason},
)
with self._engine.connect() as conn:
conn.execute(stmt)
conn.commit()
def delete_workstream_override(self, ws_id: str) -> bool:
with self._engine.connect() as conn:
result = conn.execute(
sa.delete(workstream_overrides).where(workstream_overrides.c.ws_id == ws_id)
)
conn.commit()
return result.rowcount > 0
def list_workstream_overrides(self) -> list[dict[str, str]]:
with self._engine.connect() as conn:
rows = conn.execute(
sa.select(workstream_overrides).order_by(workstream_overrides.c.ws_id)
).fetchall()
return [dict(r._mapping) for r in rows]
# -- Roles -----------------------------------------------------------------
def create_role(
@@ -0,0 +1,45 @@
"""Create hash ring routing tables.
Revision ID: 030
Revises: 029
Create Date: 2026-03-30
"""
import sqlalchemy as sa
from alembic import op
revision = "030"
down_revision = "029"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"hash_ring_buckets",
sa.Column("bucket", sa.Integer, primary_key=True),
sa.Column("node_id", sa.Text, nullable=False),
)
op.create_index("idx_ring_buckets_node", "hash_ring_buckets", ["node_id"])
op.create_table(
"bucket_stats",
sa.Column("bucket", sa.Integer, primary_key=True),
sa.Column("ws_count", sa.Integer, nullable=False, server_default="0"),
sa.Column("active_count", sa.Integer, nullable=False, server_default="0"),
)
op.create_table(
"workstream_overrides",
sa.Column("ws_id", sa.Text, primary_key=True),
sa.Column("node_id", sa.Text, nullable=False),
sa.Column("reason", sa.Text, nullable=False, server_default="targeted"),
sa.Column("created", sa.Text, nullable=False),
)
op.create_index("idx_ws_overrides_node", "workstream_overrides", ["node_id"])
def downgrade() -> None:
op.drop_table("workstream_overrides")
op.drop_table("bucket_stats")
op.drop_table("hash_ring_buckets")
+28 -3
View File
@@ -44,6 +44,9 @@ class WorkstreamState(enum.Enum):
ERROR = "error" # last operation failed
_ACTIVE_STATES = {"running", "thinking", "attention"}
# ---------------------------------------------------------------------------
# Workstream dataclass
# ---------------------------------------------------------------------------
@@ -161,8 +164,12 @@ class WorkstreamManager:
if first_evicted is not None:
self._cleanup_ui(first_evicted)
self._last_evicted = first_evicted
from turnstone.core.memory import decrement_bucket_count as _dbc1
from turnstone.core.memory import delete_workstream_override as _dwo1
from turnstone.core.metrics import metrics as _m1
_dbc1(first_evicted.id, active=False) # evicted ws is always idle
_dwo1(first_evicted.id)
_m1.record_eviction()
# Create workstream and ChatSession outside the lock (construction is
@@ -186,7 +193,7 @@ class WorkstreamManager:
self._active_id = ws.id
# Persist to storage only after successful insertion
from turnstone.core.memory import register_workstream
from turnstone.core.memory import increment_bucket_count, register_workstream
register_workstream(
ws.id,
@@ -195,13 +202,18 @@ class WorkstreamManager:
skill_id=skill_id,
skill_version=skill_version,
)
increment_bucket_count(ws.id)
# Cleanup second-phase eviction outside the lock.
if second_evicted is not None:
self._cleanup_ui(second_evicted)
self._last_evicted = second_evicted
from turnstone.core.memory import decrement_bucket_count as _dbc2
from turnstone.core.memory import delete_workstream_override as _dwo2
from turnstone.core.metrics import metrics as _m2
_dbc2(second_evicted.id, active=False)
_dwo2(second_evicted.id)
_m2.record_eviction()
return ws
@@ -272,14 +284,21 @@ class WorkstreamManager:
ws = self._workstreams.pop(ws_id, None)
if ws is None:
return False
was_active = ws.state.value in _ACTIVE_STATES
self._order.remove(ws_id)
if self._active_id == ws_id:
self._active_id = self._order[0]
# Unblock any waiting approval/plan events so worker thread can exit
self._cleanup_ui(ws)
from turnstone.core.memory import update_workstream_state
from turnstone.core.memory import (
decrement_bucket_count,
delete_workstream_override,
update_workstream_state,
)
update_workstream_state(ws_id, "closed")
decrement_bucket_count(ws_id, active=was_active)
delete_workstream_override(ws_id)
return True
# -- lookup -------------------------------------------------------------
@@ -340,12 +359,18 @@ class WorkstreamManager:
ws = self._workstreams.get(ws_id)
if ws:
with ws._lock:
old_active = ws.state.value in _ACTIVE_STATES
ws.state = state
ws.last_active = time.monotonic()
ws.error_message = error_msg
from turnstone.core.memory import update_workstream_state
from turnstone.core.memory import adjust_bucket_active, update_workstream_state
update_workstream_state(ws_id, state.value)
new_active = state.value in _ACTIVE_STATES
if old_active and not new_active:
adjust_bucket_active(ws_id, -1)
elif new_active and not old_active:
adjust_bucket_active(ws_id, 1)
if self._on_state_change:
self._on_state_change(ws_id, state)