fix(dashboard): coalesce expensive reads before worker admission

This commit is contained in:
Gianpietro Dal Zio
2026-09-07 18:02:31 +02:00
committed by kshitij
parent 1eaeb73839
commit 653418d842
5 changed files with 559 additions and 4 deletions
+67
View File
@@ -0,0 +1,67 @@
"""Coalesce concurrent synchronous HTTP reads before threadpool admission."""
import asyncio
import inspect
from functools import partial, wraps
from weakref import WeakKeyDictionary
from typing import get_type_hints
from anyio.to_thread import run_sync
from hermes_constants import get_hermes_home
def _freeze(value):
if isinstance(value, dict):
return (type(value), frozenset((_freeze(k), _freeze(v)) for k, v in value.items()))
if isinstance(value, (list, tuple)):
return (type(value), tuple(map(_freeze, value)))
if isinstance(value, (set, frozenset)):
return (type(value), frozenset(map(_freeze, value)))
return (type(value), value)
def coalesced_read(func, *, thread_runner=run_sync):
"""Return an async, in-flight-only wrapper around a synchronous read.
One worker per decorated function/event loop; identical home + bound
arguments share its result. Arguments must be hashable or built-in
containers of hashable values and must not be mutated during the read.
A custom async thread_runner must await the supplied zero-argument callable
through a context-preserving threadpool, not abandon running work.
"""
signature = inspect.signature(func)
hints = get_type_hints(func, include_extras=True)
signature = signature.replace(
parameters=[p.replace(annotation=hints.get(p.name, p.annotation))
for p in signature.parameters.values()],
return_annotation=hints.get("return", signature.return_annotation),
)
states = WeakKeyDictionary()
@wraps(func)
async def wrapped(*args, **kwargs):
loop = asyncio.get_running_loop()
bound = signature.bind(*args, **kwargs)
bound.apply_defaults()
key = (str(get_hermes_home()), _freeze(bound.arguments))
hash(key)
pending, admission = states.setdefault(loop, ({}, asyncio.Semaphore(1)))
task = pending.get(key)
if task is None:
async def execute():
try:
async with admission:
return await thread_runner(partial(func, *args, **kwargs))
finally:
pending.pop(key, None)
if not pending:
states.pop(loop, None)
task = loop.create_task(execute())
pending[key] = task
# Observe failures even if every HTTP caller has disconnected.
task.add_done_callback(lambda done: None if done.cancelled() else done.exception())
# A disconnected caller must not release admission for a live worker.
return await asyncio.shield(task)
wrapped.__signature__ = signature
wrapped.__annotations__ = hints
return wrapped
+9 -3
View File
@@ -12,6 +12,7 @@ so a test's ``monkeypatch.setattr(<owning module>, "_helper", ...)`` keeps worki
import contextlib
import copy
import functools
from hermes_cli.web_read_coalescing import coalesced_read
import inspect
import json
import logging
@@ -632,17 +633,22 @@ def post_profiles_sessions_pull_requests(body: SessionPrScanBody):
return {"pull_requests": found, "scanned": wanted}
@router.get("/api/profiles")
async def list_profiles_endpoint():
@functools.partial(coalesced_read, thread_runner=lambda func: run_in_threadpool(func))
def _read_profiles():
from hermes_cli import profiles as profiles_mod
try:
profiles = await run_in_threadpool(profiles_mod.list_profiles)
profiles = profiles_mod.list_profiles()
return {"profiles": [_profile_to_dict(p) for p in profiles]}
except Exception:
_log.exception("GET /api/profiles failed; falling back to profile directory scan")
return {"profiles": _fallback_profile_dicts(profiles_mod)}
@router.get("/api/profiles")
async def list_profiles_endpoint():
return await _read_profiles()
@router.post("/api/profiles")
async def create_profile_endpoint(body: ProfileCreate):
from hermes_cli import profiles as profiles_mod
+22 -1
View File
@@ -29,6 +29,7 @@ from fastapi.responses import FileResponse
from pydantic import BaseModel, Field
from hermes_cli import kanban_db
from hermes_cli.web_read_coalescing import coalesced_read
from hermes_cli import kanban_db_connect as kbc
from hermes_cli import kanban_db_notify as kbn
from hermes_cli import kanban_db_dispatch as kbd
@@ -263,7 +264,6 @@ def _links_for(conn: sqlite3.Connection, task_id: str) -> dict[str, list[str]]:
# --- GET /board -------------------------------------------------------------
@router.get("/board")
def get_board(
tenant: Optional[str] = Query(None, description="Filter to a single tenant"),
include_archived: bool = Query(False),
@@ -315,6 +315,27 @@ def get_board(
"assignees": assignees, "latest_event_id": int(latest_event_id), "now": int(time.time())}
_read_board = coalesced_read(get_board)
@router.get("/board")
async def get_board_endpoint(
tenant: Optional[str] = Query(None, description="Filter to a single tenant"),
include_archived: bool = Query(False),
board: Optional[str] = _BOARD_Q,
workflow_template_id: Optional[str] = Query(None, description="Restrict to tasks using this workflow template id"),
current_step_key: Optional[str] = Query(None, description="Restrict to tasks at this workflow step key"),
):
# Resolve selection before keying so a board switch cannot join an older read.
return await _read_board(
tenant=tenant,
include_archived=include_archived,
board=board or kanban_db.get_current_board(),
workflow_template_id=workflow_template_id,
current_step_key=current_step_key,
)
# --- GET /tasks/:id ---------------------------------------------------------
@router.get("/tasks/{task_id}")
@@ -0,0 +1,374 @@
"""Behavioral coverage for pre-threadpool HTTP read admission."""
import asyncio
import threading
from functools import wraps
import pytest
async def wait_until(predicate):
async with asyncio.timeout(3):
while not predicate():
await asyncio.sleep(0.005)
def async_test(fn):
@wraps(fn)
def run(*args, **kwargs):
return asyncio.run(fn(*args, **kwargs))
return run
def decorator():
from hermes_cli import web_read_coalescing
return web_read_coalescing.coalesced_read
@async_test
async def test_identical_inflight_calls_share_one_read():
release = threading.Event()
calls = []
@decorator()
def read(key=1):
calls.append(key)
assert release.wait(3)
return key
tasks = [asyncio.create_task(read()) for _ in range(12)]
try:
await wait_until(lambda: calls)
await asyncio.sleep(0.05)
assert calls == [1]
finally:
release.set()
results = await asyncio.gather(*tasks)
assert results == [1] * 12
assert await read() == 1
assert calls == [1, 1]
@async_test
async def test_differing_keys_wait_before_threadpool_admission():
import anyio.to_thread
release = threading.Event()
calls = []
@decorator()
def read(key):
calls.append(key)
assert release.wait(3)
return key
limiter = anyio.to_thread.current_default_thread_limiter()
tasks = [asyncio.create_task(read(key)) for key in range(8)]
try:
await wait_until(lambda: calls)
await asyncio.sleep(0.05)
assert len(calls) == 1
assert limiter.borrowed_tokens == 1
finally:
release.set()
results = await asyncio.gather(*tasks)
assert results == list(range(8))
@async_test
async def test_caller_cancellation_keeps_worker_and_admission():
release = threading.Event()
calls = []
@decorator()
def read(key):
calls.append(key)
assert release.wait(3)
return key
first = asyncio.create_task(read(1))
tasks = []
try:
await wait_until(lambda: calls)
first.cancel()
with pytest.raises(asyncio.CancelledError):
await first
tasks = [asyncio.create_task(read(1)), asyncio.create_task(read(2))]
await asyncio.sleep(0.05)
assert calls == [1]
finally:
release.set()
results = await asyncio.gather(*tasks)
assert results == [1, 2]
assert calls == [1, 2]
@async_test
async def test_canonical_arguments_coalesce():
release = threading.Event()
calls = []
@decorator()
def read(key=1, *, limit=2):
calls.append(key)
assert release.wait(3)
return (key, limit)
tasks = [asyncio.create_task(read()), asyncio.create_task(read(1)),
asyncio.create_task(read(limit=2, key=1))]
try:
await wait_until(lambda: calls)
await asyncio.sleep(0.05)
finally:
release.set()
results = await asyncio.gather(*tasks)
assert results == [(1, 2)] * 3
assert calls == [1]
@async_test
async def test_home_isolation_preserves_worker_context(monkeypatch):
from contextvars import ContextVar
from pathlib import Path
from hermes_cli import web_read_coalescing
home = ContextVar("test_read_home", default=Path("/first"))
monkeypatch.setattr(web_read_coalescing, "get_hermes_home", home.get, raising=False)
release = threading.Event()
calls = []
@decorator()
def read():
calls.append(str(home.get()))
assert release.wait(3)
return str(home.get())
first = asyncio.create_task(read())
home.set(Path("/second"))
second = asyncio.create_task(read())
try:
await wait_until(lambda: calls)
await asyncio.sleep(0.05)
finally:
release.set()
results = await asyncio.gather(first, second)
assert results == ["/first", "/second"]
assert calls == results
def test_fastapi_can_resolve_postponed_annotations():
import inspect
from pathlib import Path
from fastapi import FastAPI
namespace = {"Path": Path}
exec("from __future__ import annotations\ndef read(key: Path) -> dict:\n return {'key': str(key)}", namespace)
read = decorator()(namespace["read"])
assert inspect.iscoroutinefunction(read)
assert inspect.signature(read).parameters["key"].annotation is Path
app = FastAPI()
app.get("/read")(read)
assert app.openapi()["paths"]["/read"]["get"]["parameters"][0]["name"] == "key"
@async_test
async def test_failed_shared_read_is_retried():
release = threading.Event()
calls = []
@decorator()
def read():
calls.append(1)
assert release.wait(3)
if len(calls) == 1:
raise ValueError("read failed")
return "recovered"
tasks = [asyncio.create_task(read()) for _ in range(5)]
await wait_until(lambda: calls)
await asyncio.sleep(0.05)
release.set()
errors = await asyncio.gather(*tasks, return_exceptions=True)
assert all(isinstance(error, ValueError) for error in errors)
assert calls == [1]
assert await read() == "recovered"
assert calls == [1, 1]
@async_test
async def test_abandoned_failure_is_observed_and_retryable():
import gc
release = threading.Event()
calls = []
errors = []
loop = asyncio.get_running_loop()
loop.set_exception_handler(lambda loop, context: errors.append(context))
@decorator()
def read():
calls.append(1)
assert release.wait(3)
raise ValueError("abandoned read")
caller = asyncio.create_task(read())
try:
await wait_until(lambda: calls)
caller.cancel()
with pytest.raises(asyncio.CancelledError):
await caller
finally:
release.set()
await asyncio.sleep(0.05)
gc.collect()
assert errors == []
with pytest.raises(ValueError, match="abandoned read"):
await read()
assert calls == [1, 1]
@async_test
async def test_profiles_burst_leaves_threadpool_status_responsive(monkeypatch):
import anyio.to_thread
import httpx
from fastapi import FastAPI
from hermes_cli import profiles as profiles_mod
from hermes_cli.web_routers import profiles
from starlette.concurrency import run_in_threadpool
release = threading.Event()
calls = []
def slow_profiles():
calls.append(1)
assert release.wait(3)
return ["example"]
monkeypatch.setattr(profiles_mod, "list_profiles", slow_profiles)
monkeypatch.setattr(profiles, "_profile_to_dict", lambda p: {"name": p})
monkeypatch.setattr(profiles, "run_in_threadpool", run_in_threadpool)
app = FastAPI()
app.include_router(profiles.router)
@app.get("/api/status")
def status():
return {"ok": True}
limiter = anyio.to_thread.current_default_thread_limiter()
previous = limiter.total_tokens
limiter.total_tokens = 2
try:
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client:
tasks = [asyncio.create_task(client.get("/api/profiles")) for _ in range(12)]
try:
await wait_until(lambda: calls)
await asyncio.sleep(0.05)
response = await asyncio.wait_for(client.get("/api/status"), 0.5)
assert response.status_code == 200
assert response.json() == {"ok": True}
assert calls == [1]
assert limiter.borrowed_tokens == 1
finally:
release.set()
responses = await asyncio.gather(*tasks)
assert all(r.status_code == 200 for r in responses)
assert all(r.json() == {"profiles": [{"name": "example"}]} for r in responses)
finally:
limiter.total_tokens = previous
@async_test
async def test_profiles_retains_threadpool_monkeypatch_seam(monkeypatch):
from hermes_cli import profiles as profiles_mod
from hermes_cli.web_routers import profiles
calls = []
async def runner(func):
calls.append("runner")
return func()
monkeypatch.setattr(profiles, "run_in_threadpool", runner)
monkeypatch.setattr(profiles_mod, "list_profiles", lambda: ["example"])
monkeypatch.setattr(profiles, "_profile_to_dict", lambda p: {"name": p})
assert await profiles.list_profiles_endpoint() == {"profiles": [{"name": "example"}]}
assert calls == ["runner"]
@async_test
async def test_container_arguments_are_canonical_and_type_safe():
release = threading.Event()
calls = []
@decorator()
def read(options):
calls.append(options)
assert release.wait(3)
return options
values = [{"a": [1], "b": {2, 3}}, {"b": {3, 2}, "a": [1]}, True, 1]
tasks = [asyncio.create_task(read(value)) for value in values]
try:
await asyncio.sleep(0.05)
finally:
release.set()
results = await asyncio.gather(*tasks)
assert results == values
assert len(calls) == 3
assert type(results[2]) is bool
assert type(results[3]) is int
def test_decorated_function_is_isolated_between_live_event_loops():
from concurrent.futures import ThreadPoolExecutor
rendezvous = threading.Barrier(2)
@decorator()
def read():
rendezvous.wait(timeout=3)
return threading.get_ident()
with ThreadPoolExecutor(max_workers=2) as executor:
futures = [executor.submit(asyncio.run, read()) for _ in range(2)]
results = [future.result(timeout=5) for future in futures]
assert len(set(results)) == 2
@async_test
async def test_profiles_failure_preserves_fallback_and_retries(monkeypatch):
from hermes_cli import profiles as profiles_mod
from hermes_cli.web_routers import profiles
from starlette.concurrency import run_in_threadpool
calls = []
def read():
calls.append(1)
if len(calls) == 1:
raise OSError("unavailable")
return ["recovered"]
monkeypatch.setattr(profiles, "run_in_threadpool", run_in_threadpool)
monkeypatch.setattr(profiles_mod, "list_profiles", read)
monkeypatch.setattr(profiles, "_profile_to_dict", lambda p: {"name": p})
monkeypatch.setattr(profiles, "_fallback_profile_dicts", lambda mod: [{"name": "fallback"}])
assert await profiles.list_profiles_endpoint() == {"profiles": [{"name": "fallback"}]}
assert await profiles.list_profiles_endpoint() == {"profiles": [{"name": "recovered"}]}
assert calls == [1, 1]
@async_test
async def test_profiles_fallback_is_coalesced_off_event_loop(monkeypatch):
from hermes_cli import profiles as profiles_mod
from hermes_cli.web_routers import profiles
from starlette.concurrency import run_in_threadpool
event_loop_thread = threading.get_ident()
fallback_threads = []
expected = {"profiles": [{"name": "fallback"}]}
def unavailable():
raise OSError("unavailable")
def fallback(_):
fallback_threads.append(threading.get_ident())
return expected["profiles"]
monkeypatch.setattr(profiles, "run_in_threadpool", run_in_threadpool)
monkeypatch.setattr(profiles_mod, "list_profiles", unavailable)
monkeypatch.setattr(profiles, "_fallback_profile_dicts", fallback)
results = await asyncio.gather(*(profiles.list_profiles_endpoint() for _ in range(8)))
assert all(result == expected for result in results)
assert event_loop_thread not in fallback_threads
assert len(fallback_threads) == 1
@@ -0,0 +1,87 @@
"""HTTP admission regression: board bursts must not consume every worker."""
import asyncio
import importlib.util
import os
import sys
import threading
from pathlib import Path
import anyio
import httpx
import pytest
from fastapi import FastAPI
from starlette.concurrency import run_in_threadpool
from hermes_cli import kanban_db as kb
OK = 200
BOARD_PATH = "/api/plugins/kanban/board"
STATUS_PATH = "/probe"
BURST_SIZE = 8
@pytest.fixture
def anyio_backend():
return "asyncio"
@pytest.mark.anyio
async def test_board_burst_preserves_http_worker_capacity(tmp_path, monkeypatch):
home = tmp_path / "home"
home.mkdir()
monkeypatch.setenv("HERMES_HOME", str(home))
monkeypatch.setattr(Path, "home", lambda: tmp_path)
kb.init_db()
plugin_path = Path(os.environ.get(
"HERMES_TEST_KANBAN_PLUGIN",
str(Path(__file__).resolve().parents[2] / "plugins/kanban/dashboard/plugin_api.py"),
))
spec = importlib.util.spec_from_file_location("kanban_admission_test", plugin_path)
plugin = importlib.util.module_from_spec(spec)
monkeypatch.setitem(sys.modules, spec.name, plugin)
spec.loader.exec_module(plugin)
app = FastAPI()
app.include_router(plugin.router, prefix="/api/plugins/kanban")
@app.get(STATUS_PATH)
async def probe():
return await run_in_threadpool(lambda: {"ready": True})
entered = threading.Event()
release = threading.Event()
count_lock = threading.Lock()
calls = 0
original = kb.list_tasks
def blocked_read(*args, **kwargs):
nonlocal calls
with count_lock:
calls += 1
entered.set()
assert release.wait(10), "test did not release board read"
return original(*args, **kwargs)
monkeypatch.setattr(kb, "list_tasks", blocked_read)
limiter = anyio.to_thread.current_default_thread_limiter()
old_tokens = limiter.total_tokens
limiter.total_tokens = 2
requests = []
try:
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client:
requests = [asyncio.create_task(client.get(BOARD_PATH)) for _ in range(BURST_SIZE)]
for _ in range(200):
if entered.is_set():
break
await asyncio.sleep(0.01)
assert entered.is_set()
response = await asyncio.wait_for(client.get(STATUS_PATH), timeout=2)
assert response.status_code == OK
release.set()
responses = await asyncio.gather(*requests)
assert all(r.status_code == OK for r in responses)
assert calls == 1
finally:
release.set()
if requests:
await asyncio.gather(*requests, return_exceptions=True)
limiter.total_tokens = old_tokens