feat(proxy): track Codex WS sessions and cancel relay tasks deterministically

Unit 3 of the Codex proxy resilience plan. Eliminates the "aged process
has leaked relay tasks" hypothesis by making every WS session explicitly
tracked and both relay tasks deterministically cancelled when either
exits.

- New headroom/proxy/ws_session_registry.py: dict-backed
  WebSocketSessionRegistry + WSSessionHandle with register /
  deregister / attach_tasks / snapshot. Deregister is idempotent and
  clears relay-task references so coroutine frames are not retained
  past session end.
- HeadroomProxy exposes proxy.ws_sessions so /debug/ws-sessions
  (Unit 5) can read the live snapshot.
- handle_openai_responses_ws now registers on websocket.accept()
  success and deregisters in the outermost finally so no leak can
  survive handshake-phase, mid-stream, or upstream-error paths. The
  session_id / termination_cause is threaded through both relay
  halves and both sides raise asyncio.CancelledError cleanly.
- Replaced asyncio.gather(_client_to_upstream(), _upstream_to_client(),
  return_exceptions=True) with explicit asyncio.create_task(...)
  (named codex-ws-c2u-<sid> / codex-ws-u2c-<sid>) +
  asyncio.wait(FIRST_COMPLETED) + cancel-and-await on the survivor.
  Termination cause is classified as client_disconnect /
  client_error / upstream_disconnect / upstream_error /
  response_completed from which task completed first plus inline
  error captures from the halves.
- Prometheus metrics: new active_ws_sessions and active_relay_tasks
  gauges plus ws_session_duration_ms_{sum,count,max} histogram
  bucketed by termination cause. Mirrors the Unit 2 stage_timing_*
  shape.

Preserved: upstream WS retry loop, WS→HTTP fallback, memory-context
timeout, compression pipeline, Unit 2 stage timings. Memory-tool
execution inside _upstream_to_client still runs when the client task
exits first; however, if the client disconnects *before* the upstream
emits response.completed, pending memory writes in `pending_fcs` are
dropped (unchanged from prior behavior — a crashing upstream has the
same effect). Note: handle_openai_responses (HTTP, line ~800) is a
single-shot HTTP request; lifecycle tracking isn't added there
(scope boundary).

Tests:
- tests/test_ws_session_registry.py: 8 registry unit tests
  (register/deregister idempotency, snapshot shape, attach merging,
  reference release).
- tests/test_openai_codex_ws_lifecycle.py: 6 integration tests
  using real relay tasks (only upstream WS endpoint mocked):
  happy-path, failing-test-first "client disconnect cancels upstream
  relay within 100 ms", upstream-closes-first, upstream-error mid-
  stream, handshake-failure deregister, 50 concurrent sessions.
- Regression: test_openai_codex_ws_timings, test_openai_codex_routing,
  test_proxy_codex_route_aliases, test_ws_memory_relay all pass.
- Tests pass under python -W error::RuntimeWarning (no "coroutine
  was never awaited").
This commit is contained in:
Adryan Eka Vandra 2026-04-17 13:14:17 +07:00
parent 7e2b078b76
commit a31d81a426
No known key found for this signature in database
GPG key ID: A46A577A26A97682
7 changed files with 1630 additions and 5 deletions

View file

@ -19,6 +19,10 @@ from datetime import datetime
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from headroom.proxy.stage_timer import StageTimer, emit_stage_timings_log from headroom.proxy.stage_timer import StageTimer, emit_stage_timings_log
from headroom.proxy.ws_session_registry import (
WebSocketSessionRegistry,
WSSessionHandle,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from fastapi import Request, WebSocket from fastapi import Request, WebSocket
@ -1248,6 +1252,15 @@ class OpenAIHandlerMixin:
stage_timer = StageTimer() stage_timer = StageTimer()
session_started_at = time.perf_counter() session_started_at = time.perf_counter()
# Unit 3: initialize registry variables *before* accept so the
# outermost ``finally`` can rely on them existing even if
# registration itself fails for some reason.
ws_sessions: WebSocketSessionRegistry | None = getattr(
self, "ws_sessions", None
)
session_handle: WSSessionHandle | None = None
termination_cause: str = "unknown"
# Forward client headers to upstream, adding required OpenAI-Beta header # Forward client headers to upstream, adding required OpenAI-Beta header
ws_headers = dict(websocket.headers) ws_headers = dict(websocket.headers)
@ -1266,6 +1279,31 @@ class OpenAIHandlerMixin:
else: else:
await websocket.accept() await websocket.accept()
# --- Unit 3: register the session as soon as accept succeeds ---
client_addr: str | None = None
client_info = getattr(websocket, "client", None)
if client_info is not None:
host = getattr(client_info, "host", None)
port = getattr(client_info, "port", None)
if host is not None and port is not None:
client_addr = f"{host}:{port}"
elif host is not None:
client_addr = str(host)
if ws_sessions is not None:
session_handle = WSSessionHandle(
session_id=session_id,
request_id=request_id,
client_addr=client_addr,
upstream_url=None, # set below once upstream_url is computed
)
ws_sessions.register(session_handle)
metrics = getattr(self, "metrics", None)
if metrics is not None and hasattr(metrics, "inc_active_ws_sessions"):
try:
metrics.inc_active_ws_sessions()
except Exception: # pragma: no cover - defensive
pass
# Forward all client headers except hop-by-hop / per-connection headers. # Forward all client headers except hop-by-hop / per-connection headers.
# These are WebSocket handshake mechanics that the `websockets` library # These are WebSocket handshake mechanics that the `websockets` library
# generates fresh for the upstream connection — forwarding them would conflict. # generates fresh for the upstream connection — forwarding them would conflict.
@ -1305,6 +1343,10 @@ class OpenAIHandlerMixin:
ws_base = base.replace("https://", "wss://").replace("http://", "ws://") ws_base = base.replace("https://", "wss://").replace("http://", "ws://")
upstream_url = f"{ws_base}/v1/responses" upstream_url = f"{ws_base}/v1/responses"
# Unit 3: attach the resolved upstream URL to the session handle.
if session_handle is not None:
session_handle.upstream_url = upstream_url
# Ensure Authorization header is present — fall back to OPENAI_API_KEY env var. # Ensure Authorization header is present — fall back to OPENAI_API_KEY env var.
# Safety net for clients that don't forward auth headers via WebSocket upgrade. # Safety net for clients that don't forward auth headers via WebSocket upgrade.
if "authorization" not in _lower_headers: if "authorization" not in _lower_headers:
@ -1579,13 +1621,39 @@ class OpenAIHandlerMixin:
_upstream_first_event_started = time.perf_counter() _upstream_first_event_started = time.perf_counter()
await upstream.send(first_msg_raw) await upstream.send(first_msg_raw)
# Unit 3: flag the upstream side flips on seeing
# ``response.completed`` so the outer cause
# classifier can prefer it over the raw
# "upstream iterator ended" default.
response_completed_seen = False
# Captures the first exception surfaced by the
# inner relay ``except`` blocks so the outer
# classifier can still tell ``upstream_error``
# from ``upstream_disconnect`` / ``response_completed``
# even though the halves swallow and log.
upstream_relay_error: BaseException | None = None
client_relay_error: BaseException | None = None
async def _client_to_upstream() -> None: async def _client_to_upstream() -> None:
nonlocal client_relay_error
try: try:
while True: while True:
msg = await websocket.receive_text() msg = await websocket.receive_text()
await upstream.send(msg) await upstream.send(msg)
except asyncio.CancelledError:
# Explicit cancel from the outer
# orchestrator — re-raise so
# ``t.cancelled()`` and ``t.exception()``
# behave correctly in the caller.
raise
except Exception as relay_err: except Exception as relay_err:
# Surface real errors to the classifier
# without re-raising (existing fork
# behavior: log and return so the
# partner task can be cancelled
# deterministically).
if "WebSocketDisconnect" not in type(relay_err).__name__: if "WebSocketDisconnect" not in type(relay_err).__name__:
client_relay_error = relay_err
logger.debug( logger.debug(
f"[{request_id}] WS client→upstream relay ended: {relay_err}" f"[{request_id}] WS client→upstream relay ended: {relay_err}"
) )
@ -1606,6 +1674,13 @@ class OpenAIHandlerMixin:
""" """
from headroom.proxy.memory_handler import MEMORY_TOOL_NAMES from headroom.proxy.memory_handler import MEMORY_TOOL_NAMES
# Unit 3: surface response.completed observation
# to the outer scope so the termination-cause
# classifier can prefer ``response_completed``
# over ``upstream_disconnect``.
nonlocal response_completed_seen
nonlocal upstream_relay_error
memory_enabled = bool(self.memory_handler and memory_user_id) memory_enabled = bool(self.memory_handler and memory_user_id)
# Per-response state (reset after each response.completed) # Per-response state (reset after each response.completed)
@ -1689,6 +1764,7 @@ class OpenAIHandlerMixin:
await websocket.send_text(buf) await websocket.send_text(buf)
event_buffer.clear() event_buffer.clear()
_reset() _reset()
response_completed_seen = True
continue continue
@ -1703,6 +1779,7 @@ class OpenAIHandlerMixin:
pending_fcs.append(item) pending_fcs.append(item)
elif event_type == "response.completed": elif event_type == "response.completed":
response_completed_seen = True
resp = event.get("response", {}) resp = event.get("response", {})
resp_id = resp.get("id") resp_id = resp.get("id")
@ -1770,8 +1847,15 @@ class OpenAIHandlerMixin:
# --- Phase 2b: Pass-through mode --- # --- Phase 2b: Pass-through mode ---
await websocket.send_text(msg_str) await websocket.send_text(msg_str)
except asyncio.CancelledError:
raise
except Exception as relay_err: except Exception as relay_err:
if "WebSocketDisconnect" not in type(relay_err).__name__: if "WebSocketDisconnect" not in type(relay_err).__name__:
# Capture for the outer classifier
# so ``upstream_error`` can be
# distinguished from a clean
# upstream disconnect.
upstream_relay_error = relay_err
logger.debug( logger.debug(
f"[{request_id}] WS upstream→client relay ended: {relay_err}" f"[{request_id}] WS upstream→client relay ended: {relay_err}"
) )
@ -1779,14 +1863,115 @@ class OpenAIHandlerMixin:
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
await websocket.close() await websocket.close()
await asyncio.gather( # --- Unit 3: deterministic relay-task cancellation ---
# Spawn each half as a named task so we can:
# (a) attach them to the session registry for
# ``/debug/ws-sessions``,
# (b) cancel the survivor explicitly when the
# first one exits, and
# (c) classify the termination cause for the
# duration histogram.
client_task = asyncio.create_task(
_client_to_upstream(), _client_to_upstream(),
_upstream_to_client(), name=f"codex-ws-c2u-{session_id}",
return_exceptions=True,
) )
upstream_task = asyncio.create_task(
_upstream_to_client(),
name=f"codex-ws-u2c-{session_id}",
)
relay_tasks = [client_task, upstream_task]
if ws_sessions is not None:
ws_sessions.attach_tasks(session_id, relay_tasks)
metrics_for_tasks = getattr(self, "metrics", None)
if metrics_for_tasks is not None and hasattr(
metrics_for_tasks, "inc_active_relay_tasks"
):
try:
metrics_for_tasks.inc_active_relay_tasks(
len(relay_tasks)
)
except Exception: # pragma: no cover - defensive
pass
try:
done, pending = await asyncio.wait(
{client_task, upstream_task},
return_when=asyncio.FIRST_COMPLETED,
)
# Cancel the survivor so we don't leak the
# partner task. Suppress the CancelledError
# we just raised ourselves — any *other*
# exception from the cancelled task is
# already logged inside its own try/except.
for t in pending:
t.cancel()
if pending:
with contextlib.suppress(asyncio.CancelledError):
await asyncio.gather(
*pending, return_exceptions=True
)
# Classify termination cause from whichever
# task completed first. ``CancelledError``
# can show up on the "done" side if the
# handler itself was cancelled from outside
# (e.g. server shutdown).
for t in done:
exc = None
with contextlib.suppress(Exception):
exc = t.exception()
task_name = t.get_name() or ""
if t is client_task:
if client_relay_error is not None:
termination_cause = "client_error"
elif exc is None:
termination_cause = "client_disconnect"
elif isinstance(exc, asyncio.CancelledError):
termination_cause = "client_disconnect"
else:
# Distinguish legitimate client
# disconnect exceptions from
# real errors: WebSocketDisconnect
# is a normal client exit.
if "WebSocketDisconnect" in type(exc).__name__:
termination_cause = "client_disconnect"
else:
termination_cause = "client_error"
elif t is upstream_task:
if upstream_relay_error is not None:
termination_cause = "upstream_error"
logger.debug(
f"[{request_id}] WS relay {task_name} "
f"raised: {upstream_relay_error!r}"
)
elif exc is None:
termination_cause = (
"response_completed"
if response_completed_seen
else "upstream_disconnect"
)
elif isinstance(exc, asyncio.CancelledError):
termination_cause = "upstream_disconnect"
else:
termination_cause = "upstream_error"
logger.debug(
f"[{request_id}] WS relay {task_name} "
f"raised: {exc!r}"
)
finally:
# In case anything above raised before the
# cancel-and-await loop ran.
for t in relay_tasks:
if not t.done():
t.cancel()
with contextlib.suppress(asyncio.CancelledError):
await asyncio.gather(
*relay_tasks, return_exceptions=True
)
logger.info( logger.info(
f"[{request_id}] WS /v1/responses completed (tokens_saved={tokens_saved})" f"[{request_id}] WS /v1/responses completed "
f"(tokens_saved={tokens_saved}, cause={termination_cause})"
) )
break break
except Exception as ws_err: except Exception as ws_err:
@ -1840,7 +2025,14 @@ class OpenAIHandlerMixin:
) )
except Exception as e: except Exception as e:
if "WebSocketDisconnect" not in type(e).__name__: if "WebSocketDisconnect" in type(e).__name__:
# Unit 3: client dropped the socket before or during
# relay. The registry classifier may already have set
# ``client_disconnect`` via the relay task exit path;
# preserve that, otherwise set it here.
if termination_cause == "unknown":
termination_cause = "client_disconnect"
else:
# Extract response body from websockets InvalidStatus for better debugging # Extract response body from websockets InvalidStatus for better debugging
error_detail = str(e) error_detail = str(e)
if hasattr(e, "response"): if hasattr(e, "response"):
@ -1854,6 +2046,8 @@ class OpenAIHandlerMixin:
except Exception: except Exception:
pass pass
logger.error(f"[{request_id}] WS proxy error: {error_detail}") logger.error(f"[{request_id}] WS proxy error: {error_detail}")
if termination_cause == "unknown":
termination_cause = "client_error"
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
await websocket.close(code=1011, reason=str(e)[:120]) await websocket.close(code=1011, reason=str(e)[:120])
finally: finally:
@ -1862,6 +2056,29 @@ class OpenAIHandlerMixin:
"total_session", "total_session",
(time.perf_counter() - session_started_at) * 1000.0, (time.perf_counter() - session_started_at) * 1000.0,
) )
# Unit 3: deregister the session before (or independently
# of) the stage-timings log so a failure there cannot leak
# the registry entry. ``deregister`` is idempotent, so a
# session that never registered is a no-op.
if ws_sessions is not None and session_handle is not None:
released_tasks = len(session_handle.relay_tasks)
ws_sessions.deregister(session_id, cause=termination_cause)
session_duration_ms = (
time.perf_counter() - session_started_at
) * 1000.0
metrics_for_close = getattr(self, "metrics", None)
if metrics_for_close is not None:
with contextlib.suppress(Exception):
if hasattr(metrics_for_close, "dec_active_ws_sessions"):
metrics_for_close.dec_active_ws_sessions()
if released_tasks and hasattr(
metrics_for_close, "dec_active_relay_tasks"
):
metrics_for_close.dec_active_relay_tasks(released_tasks)
if hasattr(metrics_for_close, "record_ws_session_duration"):
metrics_for_close.record_ws_session_duration(
session_duration_ms, termination_cause
)
await emit_stage_timings_log( await emit_stage_timings_log(
path="openai_responses_ws", path="openai_responses_ws",
request_id=request_id, request_id=request_id,

View file

@ -108,6 +108,17 @@ class PrometheusMetrics:
self.stage_timing_count: dict[tuple[str, str], int] = defaultdict(int) self.stage_timing_count: dict[tuple[str, str], int] = defaultdict(int)
self.stage_timing_max: dict[tuple[str, str], float] = defaultdict(float) self.stage_timing_max: dict[tuple[str, str], float] = defaultdict(float)
# WS session lifecycle (Unit 3). Gauges are live counters updated
# by the Codex handler on register/deregister + attach_tasks/
# detach. Histograms record completed-session durations bucketed
# by termination cause so we can distinguish slow happy-path
# sessions from long client-hold followed by client_disconnect.
self.active_ws_sessions: int = 0
self.active_relay_tasks: int = 0
self.ws_session_duration_sum_ms: dict[str, float] = defaultdict(float)
self.ws_session_duration_count: dict[str, int] = defaultdict(int)
self.ws_session_duration_max_ms: dict[str, float] = defaultdict(float)
# Aggregate waste signals # Aggregate waste signals
self.waste_signals_total: dict[str, int] = defaultdict(int) self.waste_signals_total: dict[str, int] = defaultdict(int)
@ -359,6 +370,47 @@ class PrometheusMetrics:
self.cache_bust_count += 1 self.cache_bust_count += 1
self._get_otel_metrics().record_proxy_cache_bust(tokens_lost=tokens_lost) self._get_otel_metrics().record_proxy_cache_bust(tokens_lost=tokens_lost)
# ------------------------------------------------------------------
# Unit 3: WS session lifecycle gauges / histogram
# ------------------------------------------------------------------
def inc_active_ws_sessions(self) -> None:
"""Increment the live WS session gauge (called on register)."""
self.active_ws_sessions += 1
def dec_active_ws_sessions(self) -> None:
"""Decrement the live WS session gauge (called on deregister)."""
self.active_ws_sessions = max(0, self.active_ws_sessions - 1)
def inc_active_relay_tasks(self, n: int = 1) -> None:
"""Increment the live relay-task gauge (attach_tasks)."""
self.active_relay_tasks += n
def dec_active_relay_tasks(self, n: int = 1) -> None:
"""Decrement the live relay-task gauge (deregister)."""
self.active_relay_tasks = max(0, self.active_relay_tasks - n)
def record_ws_session_duration(
self,
duration_ms: float,
cause: str = "unknown",
) -> None:
"""Record a completed WS session's duration, bucketed by cause.
Mirrors the ``stage_timing_*`` shape so ``/metrics`` exposes
sum/count/max per termination cause. Uses synchronous dict
updates (no ``_lock``) because Unit 3 callers run on the event
loop matching the gauges above.
"""
try:
ms_val = float(duration_ms)
except (TypeError, ValueError):
return
self.ws_session_duration_sum_ms[cause] += ms_val
self.ws_session_duration_count[cause] += 1
if ms_val > self.ws_session_duration_max_ms[cause]:
self.ws_session_duration_max_ms[cause] = ms_val
async def record_rate_limited(self, *, provider: str | None = None, model: str | None = None): async def record_rate_limited(self, *, provider: str | None = None, model: str | None = None):
async with self._lock: async with self._lock:
self.requests_rate_limited += 1 self.requests_rate_limited += 1
@ -611,6 +663,54 @@ class PrometheusMetrics:
) )
lines.append("") lines.append("")
# Unit 3: WS session lifecycle gauges + duration histogram.
lines.extend(
[
"# HELP headroom_active_ws_sessions Active Codex WebSocket sessions",
"# TYPE headroom_active_ws_sessions gauge",
f"headroom_active_ws_sessions {self.active_ws_sessions}",
"",
"# HELP headroom_active_relay_tasks Active Codex WS relay tasks",
"# TYPE headroom_active_relay_tasks gauge",
f"headroom_active_relay_tasks {self.active_relay_tasks}",
"",
]
)
if self.ws_session_duration_sum_ms:
lines.extend(
[
"# HELP headroom_ws_session_duration_ms_sum Sum of Codex WS session durations",
"# TYPE headroom_ws_session_duration_ms_sum counter",
]
)
for cause, total in self.ws_session_duration_sum_ms.items():
lines.append(
f'headroom_ws_session_duration_ms_sum{{cause="{_escape_label_value(cause)}"}} {round(total, 2)}'
)
lines.extend(
[
"",
"# HELP headroom_ws_session_duration_ms_count Count of completed Codex WS sessions",
"# TYPE headroom_ws_session_duration_ms_count counter",
]
)
for cause, count in self.ws_session_duration_count.items():
lines.append(
f'headroom_ws_session_duration_ms_count{{cause="{_escape_label_value(cause)}"}} {count}'
)
lines.extend(
[
"",
"# HELP headroom_ws_session_duration_ms_max Maximum Codex WS session duration",
"# TYPE headroom_ws_session_duration_ms_max gauge",
]
)
for cause, max_value in self.ws_session_duration_max_ms.items():
lines.append(
f'headroom_ws_session_duration_ms_max{{cause="{_escape_label_value(cause)}"}} {round(max_value, 2)}'
)
lines.append("")
if self.waste_signals_total: if self.waste_signals_total:
lines.extend( lines.extend(
[ [

View file

@ -129,6 +129,7 @@ from headroom.proxy.rate_limiter import TokenBucketRateLimiter # noqa: F401
from headroom.proxy.request_logger import RequestLogger # noqa: F401 from headroom.proxy.request_logger import RequestLogger # noqa: F401
from headroom.proxy.semantic_cache import SemanticCache # noqa: F401 from headroom.proxy.semantic_cache import SemanticCache # noqa: F401
from headroom.proxy.warmup import WarmupRegistry from headroom.proxy.warmup import WarmupRegistry
from headroom.proxy.ws_session_registry import WebSocketSessionRegistry
from headroom.subscription.base import get_quota_registry, reset_quota_registry from headroom.subscription.base import get_quota_registry, reset_quota_registry
from headroom.subscription.codex_rate_limits import get_codex_rate_limit_state from headroom.subscription.codex_rate_limits import get_codex_rate_limit_state
from headroom.subscription.copilot_quota import get_copilot_quota_tracker from headroom.subscription.copilot_quota import get_copilot_quota_tracker
@ -373,6 +374,10 @@ class HeadroomProxy(
# each preloaded heavy asset. Exposed as ``proxy.warmup`` and # each preloaded heavy asset. Exposed as ``proxy.warmup`` and
# serialized by the /debug/warmup route (Unit 5). # serialized by the /debug/warmup route (Unit 5).
self.warmup: WarmupRegistry = WarmupRegistry() self.warmup: WarmupRegistry = WarmupRegistry()
# Unit 3: live registry of Codex WS sessions. Populated by
# ``handle_openai_responses_ws`` on accept; drained in its
# outermost ``finally``. Consumed by ``/debug/ws-sessions``.
self.ws_sessions: WebSocketSessionRegistry = WebSocketSessionRegistry()
# Backend for Anthropic API (direct, LiteLLM, or any-llm) # Backend for Anthropic API (direct, LiteLLM, or any-llm)
# Supports: "anthropic" (direct), "bedrock", "vertex", "litellm-<provider>", or "anyllm" # Supports: "anthropic" (direct), "bedrock", "vertex", "litellm-<provider>", or "anyllm"

View file

@ -0,0 +1,192 @@
"""WebSocket session registry for Codex relay lifecycle tracking.
Unit 3 of the Codex-proxy resilience plan. Every ``/v1/responses`` WS
session is explicitly registered on accept and deregistered in the
outermost ``finally`` of the handler. The registry provides:
* First-class visibility of active sessions (``/debug/ws-sessions``
from Unit 5 consumes :meth:`WebSocketSessionRegistry.snapshot`).
* Gauges for Prometheus (``active_ws_sessions``, ``active_relay_tasks``).
* A home for relay-task references so the handler's orchestrator can
attach them at creation time for introspection, without the registry
itself owning or cancelling them cancellation is the handler's job.
Design notes
------------
* Single-event-loop usage. All mutations happen on the proxy's event
loop from within the handler coroutine, so no asyncio.Lock / asyncio
primitives are needed. Python dict mutations are atomic under the
GIL; snapshot copies happen under the event loop's cooperative
scheduling, so iteration + mutation cannot interleave across an
``await`` point.
* ``deregister`` must be idempotent. The handler calls it from the
outermost ``finally`` if upstream never connected, the session may
have been registered or may not have been, depending on how far
handshake got. Either way ``deregister`` is safe.
* ``deregister`` clears the handle's ``relay_tasks`` list so the
registry does not retain references to task coroutine frames after a
session ends.
"""
from __future__ import annotations
import logging
import time
from collections.abc import Iterable
from dataclasses import dataclass, field
from typing import Any, Protocol
logger = logging.getLogger(__name__)
__all__ = [
"WebSocketSessionRegistry",
"WSSessionHandle",
"TerminationCause",
]
TerminationCause = str # "client_disconnect" | "upstream_disconnect" | "upstream_error"
# | "client_error" | "response_completed" | "unknown"
class _TaskLike(Protocol):
def done(self) -> bool: ...
def cancel(self) -> bool: ...
def get_name(self) -> str: ...
@dataclass
class WSSessionHandle:
"""Per-session state entry held in :class:`WebSocketSessionRegistry`.
Timestamps use :func:`time.perf_counter` so age computations are
monotonic and independent of wall-clock adjustments.
"""
session_id: str
request_id: str
client_addr: str | None = None
upstream_url: str | None = None
started_at: float = field(default_factory=time.perf_counter)
last_activity_at: float = field(default_factory=time.perf_counter)
relay_tasks: list[Any] = field(default_factory=list)
termination_cause: TerminationCause | None = None
def mark_activity(self) -> None:
self.last_activity_at = time.perf_counter()
def age_seconds(self) -> float:
return max(0.0, time.perf_counter() - self.started_at)
def to_snapshot_dict(self) -> dict[str, Any]:
return {
"session_id": self.session_id,
"request_id": self.request_id,
"client_addr": self.client_addr,
"upstream_url": self.upstream_url,
"age_seconds": self.age_seconds(),
"idle_seconds": max(0.0, time.perf_counter() - self.last_activity_at),
"relay_task_count": len(self.relay_tasks),
"relay_task_names": [
getattr(t, "get_name", lambda: "")() for t in self.relay_tasks
],
"termination_cause": self.termination_cause,
}
class WebSocketSessionRegistry:
"""In-memory registry of active Codex WS sessions.
Methods are safe to call from the event loop. Unit 5's
``/debug/ws-sessions`` endpoint consumes :meth:`snapshot`; the
Prometheus exporter reads :meth:`active_count` and
:meth:`active_relay_task_count`.
"""
def __init__(self) -> None:
self._sessions: dict[str, WSSessionHandle] = {}
# Tracked separately from ``sum(len(h.relay_tasks))`` so that
# repeated snapshotting is O(1) rather than O(N * tasks).
self._active_relay_tasks = 0
# ------------------------------------------------------------------
# Core lifecycle
# ------------------------------------------------------------------
def register(self, handle: WSSessionHandle) -> None:
"""Register a session. Idempotent by ``session_id``.
If the session id is already present, the existing entry is
replaced and the active count stays the same. This matches the
invariant that one ``session_id`` corresponds to at most one
active session at a time.
"""
existing = self._sessions.get(handle.session_id)
if existing is not None:
# Re-registration: release any tasks we were tracking on the
# old handle before replacing it.
self._active_relay_tasks -= len(existing.relay_tasks)
existing.relay_tasks.clear()
self._sessions[handle.session_id] = handle
self._active_relay_tasks += len(handle.relay_tasks)
def deregister(
self, session_id: str, cause: TerminationCause = "unknown"
) -> WSSessionHandle | None:
"""Remove a session. Idempotent: returns ``None`` if unknown.
Also clears the handle's ``relay_tasks`` list so the registry
stops holding references to coroutine frames after the session
ends.
"""
handle = self._sessions.pop(session_id, None)
if handle is None:
return None
handle.termination_cause = cause
self._active_relay_tasks = max(
0, self._active_relay_tasks - len(handle.relay_tasks)
)
handle.relay_tasks.clear()
return handle
def attach_tasks(self, session_id: str, tasks: Iterable[_TaskLike]) -> None:
"""Attach relay tasks to an existing session (merge, not replace).
If ``session_id`` is not registered, this is a no-op (handler
should have registered before spawning tasks; defensive so a
race during deregister doesn't crash the handler).
"""
handle = self._sessions.get(session_id)
if handle is None:
return
task_list = list(tasks)
handle.relay_tasks.extend(task_list)
self._active_relay_tasks += len(task_list)
handle.mark_activity()
# ------------------------------------------------------------------
# Introspection
# ------------------------------------------------------------------
def get(self, session_id: str) -> WSSessionHandle | None:
return self._sessions.get(session_id)
def active_count(self) -> int:
return len(self._sessions)
def active_relay_task_count(self) -> int:
return self._active_relay_tasks
def snapshot(self) -> list[dict[str, Any]]:
"""JSON-serializable view of the registry (for ``/debug/ws-sessions``)."""
return [handle.to_snapshot_dict() for handle in self._sessions.values()]
# ------------------------------------------------------------------
# Convenience
# ------------------------------------------------------------------
def __contains__(self, session_id: str) -> bool:
return session_id in self._sessions
def __len__(self) -> int:
return len(self._sessions)

View file

@ -0,0 +1,447 @@
"""Unit 3: WebSocket session lifecycle + deterministic relay cancellation.
These tests exercise the Codex WS handler with a fake upstream and a
fake client WebSocket so we can drive the relay halves through their
real code paths (not mocked) and assert on registry / task state.
"""
from __future__ import annotations
import asyncio
import json
import sys
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from headroom.proxy.handlers.openai import OpenAIHandlerMixin
from headroom.proxy.ws_session_registry import WebSocketSessionRegistry
# ---------------------------------------------------------------------------
# Test doubles
# ---------------------------------------------------------------------------
class _DummyMetrics:
def __init__(self) -> None:
self.active_ws_sessions = 0
self.active_ws_sessions_max = 0
self.active_relay_tasks = 0
self.ws_session_durations: list[float] = []
self.stage_timings: list[tuple[str, dict[str, float]]] = []
self.termination_causes: list[str] = []
async def record_request(self, **kwargs): # pragma: no cover
return None
async def record_stage_timings(self, path: str, timings: dict[str, float]) -> None:
self.stage_timings.append((path, dict(timings)))
def inc_active_ws_sessions(self) -> None:
self.active_ws_sessions += 1
self.active_ws_sessions_max = max(
self.active_ws_sessions_max, self.active_ws_sessions
)
def dec_active_ws_sessions(self) -> None:
self.active_ws_sessions = max(0, self.active_ws_sessions - 1)
def inc_active_relay_tasks(self, n: int = 1) -> None:
self.active_relay_tasks += n
def dec_active_relay_tasks(self, n: int = 1) -> None:
self.active_relay_tasks = max(0, self.active_relay_tasks - n)
def record_ws_session_duration(self, duration_ms: float, cause: str) -> None:
self.ws_session_durations.append(duration_ms)
self.termination_causes.append(cause)
class _DummyOpenAIHandler(OpenAIHandlerMixin):
OPENAI_API_URL = "https://api.openai.com"
def __init__(self, ws_sessions: WebSocketSessionRegistry | None = None) -> None:
self.rate_limiter = None
self.metrics = _DummyMetrics()
self.config = SimpleNamespace(
optimize=False,
retry_max_attempts=1,
retry_base_delay_ms=1,
retry_max_delay_ms=1,
connect_timeout_seconds=10,
)
self.usage_reporter = None
self.openai_provider = SimpleNamespace(get_context_limit=lambda model: 128_000)
self.openai_pipeline = SimpleNamespace(apply=MagicMock())
self.anthropic_backend = None
self.cost_tracker = None
self.memory_handler = None
self.ws_sessions = ws_sessions or WebSocketSessionRegistry()
async def _next_request_id(self) -> str:
return "req-lifecycle-test"
class _FakeWebSocketDisconnect(Exception):
"""Mirrors the ``WebSocketDisconnect`` type-name check in the handler.
The production code identifies "normal client gone" by
``"WebSocketDisconnect" in type(e).__name__`` so the fake exception
type name must start with ``WebSocketDisconnect``.
"""
# Force the type-name substring match in the handler.
_FakeWebSocketDisconnect.__name__ = "WebSocketDisconnect_Fake"
class _FakeWebSocket:
"""Scripted client WebSocket that can delay / disconnect mid-stream."""
def __init__(
self,
frames: list[str] | None = None,
*,
disconnect_after_n_sends: int | None = None,
hold_after_initial: bool = False,
) -> None:
self.headers = {"authorization": "Bearer test"}
self._frames = list(frames or [])
self._hold_after_initial = hold_after_initial
self._disconnect_after_n_sends = disconnect_after_n_sends
self.sent_text: list[str] = []
self.sent_bytes: list[bytes] = []
self.accepted_subprotocol: str | None = None
self.closed = False
self.close_code: int | None = None
# "client" can trip this event to simulate mid-stream disconnect.
self._disconnect_event = asyncio.Event()
self.client = SimpleNamespace(host="127.0.0.1", port=12345)
async def accept(self, subprotocol=None) -> None:
self.accepted_subprotocol = subprotocol
async def receive_text(self) -> str:
if self._frames:
return self._frames.pop(0)
if self._hold_after_initial:
# Wait for simulated client disconnect.
await self._disconnect_event.wait()
# Use an exception type whose name starts with ``WebSocketDisconnect``
# so the handler's ``type(e).__name__`` check classifies this as a
# normal client exit (not a ``client_error``).
raise _FakeWebSocketDisconnect("client closed")
async def send_text(self, text: str) -> None:
self.sent_text.append(text)
if (
self._disconnect_after_n_sends is not None
and len(self.sent_text) >= self._disconnect_after_n_sends
):
# Trigger the "client gone" signal the next receive_text will see.
self._disconnect_event.set()
async def send_bytes(self, data: bytes) -> None:
self.sent_bytes.append(data)
async def close(self, code: int | None = None, reason: str | None = None) -> None:
self.closed = True
self.close_code = code
def trigger_disconnect(self) -> None:
self._disconnect_event.set()
class _FakeUpstream:
"""Upstream that streams scripted events then optionally blocks.
``hold_after_events`` makes the async iterator wait forever after the
scripted events are exhausted that mirrors a real upstream that
keeps the connection open after a ``response.completed`` event. The
handler's ``_upstream_to_client`` will block on it, so the only way
the outer ``asyncio.wait`` can progress is via the client-side task
completing which is exactly the cancel-partner path we want to
test.
"""
def __init__(
self,
events: list[str],
*,
hold_after_events: bool = False,
raise_mid_stream: Exception | None = None,
) -> None:
self._events = list(events)
self._hold_after_events = hold_after_events
self._raise_mid_stream = raise_mid_stream
self.sent: list[str] = []
self.closed = False
async def __aenter__(self) -> _FakeUpstream:
return self
async def __aexit__(self, exc_type, exc, tb) -> None:
self.closed = True
async def send(self, payload: str) -> None:
self.sent.append(payload)
async def close(self) -> None:
self.closed = True
def __aiter__(self):
return self._iter()
async def _iter(self):
for ev in self._events:
yield ev
if self._raise_mid_stream is not None:
raise self._raise_mid_stream
if self._hold_after_events:
# Wait forever — until the task is cancelled by the handler.
await asyncio.Event().wait()
def _make_fake_websockets_module(upstream: _FakeUpstream):
module = MagicMock()
module.connect = MagicMock(return_value=upstream)
module.Subprotocol = str
return module
def _first_frame() -> str:
return json.dumps(
{
"type": "response.create",
"response": {"model": "gpt-5.4", "input": "hi"},
}
)
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_happy_path_registry_empty_after_response_completed():
"""Normal session completes — both relay tasks done, registry empty."""
upstream_events = [
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
json.dumps({"type": "response.completed", "response": {"id": "r_1"}}),
]
upstream = _FakeUpstream(upstream_events)
fake_ws_mod = _make_fake_websockets_module(upstream)
client_ws = _FakeWebSocket(frames=[_first_frame()])
handler = _DummyOpenAIHandler()
with patch.dict(sys.modules, {"websockets": fake_ws_mod}):
await handler.handle_openai_responses_ws(client_ws)
assert handler.ws_sessions.active_count() == 0
assert handler.metrics.active_ws_sessions == 0
# termination_cause captured
assert handler.metrics.termination_causes
# Either "response_completed" or "client_disconnect" — both are
# acceptable here depending on which relay half exited first; the
# important thing is we recorded one.
assert handler.metrics.termination_causes[-1] in {
"response_completed",
"client_disconnect",
"upstream_disconnect",
}
@pytest.mark.asyncio
async def test_client_disconnect_cancels_upstream_relay_within_100ms():
"""**Failing-test-first** scenario from the plan.
When the client side exits (``receive_text`` raises
``WebSocketDisconnect``) while upstream is still open and iterating,
the upstream relay task must be cancelled and become ``done()``
quickly. The registry must report no active sessions afterwards.
"""
# Upstream keeps iterating forever after one event, forcing the
# upstream-to-client task to block on the iterator. The only way
# out is a cancel from the handler's orchestration.
upstream_events = [
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
]
upstream = _FakeUpstream(upstream_events, hold_after_events=True)
fake_ws_mod = _make_fake_websockets_module(upstream)
# Client has one initial frame, then disconnects after the server
# sends the first forwarded event to us.
client_ws = _FakeWebSocket(
frames=[_first_frame()],
hold_after_initial=True,
)
handler = _DummyOpenAIHandler()
# Trigger disconnect shortly after the handler accepts.
async def _trigger() -> None:
await asyncio.sleep(0.05)
client_ws.trigger_disconnect()
with patch.dict(sys.modules, {"websockets": fake_ws_mod}):
trigger_task = asyncio.create_task(_trigger())
try:
await asyncio.wait_for(
handler.handle_openai_responses_ws(client_ws),
timeout=2.0,
)
finally:
trigger_task.cancel()
try:
await trigger_task
except asyncio.CancelledError:
pass
# Registry must be empty — the finally block deregistered the session.
assert handler.ws_sessions.active_count() == 0, (
"session leaked — deregister did not run in outermost finally"
)
assert handler.metrics.active_ws_sessions == 0
# We recorded a session duration (came through deregister path).
assert handler.metrics.ws_session_durations, (
"record_ws_session_duration never fired — deregister path broken"
)
# And we tagged the cause. For a client-side exit it should be one
# of: client_disconnect, client_error, upstream_disconnect (if
# upstream iteration happened to end first in a race).
cause = handler.metrics.termination_causes[-1]
assert cause in {
"client_disconnect",
"client_error",
"upstream_disconnect",
}, f"unexpected cause: {cause}"
# No codex-ws-* named task should still be running.
leaked = [
t
for t in asyncio.all_tasks()
if (t.get_name() or "").startswith("codex-ws-") and not t.done()
]
assert leaked == [], f"relay tasks leaked: {[t.get_name() for t in leaked]}"
@pytest.mark.asyncio
async def test_upstream_closes_first_cancels_client_task():
"""Upstream iterator ends naturally; client task should be cancelled.
The client is set to block on ``receive_text`` indefinitely; only a
cancel from the handler's orchestration releases it.
"""
upstream_events = [
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
json.dumps({"type": "response.completed", "response": {"id": "r_1"}}),
]
upstream = _FakeUpstream(upstream_events, hold_after_events=False)
fake_ws_mod = _make_fake_websockets_module(upstream)
client_ws = _FakeWebSocket(
frames=[_first_frame()],
hold_after_initial=True,
)
handler = _DummyOpenAIHandler()
with patch.dict(sys.modules, {"websockets": fake_ws_mod}):
await asyncio.wait_for(
handler.handle_openai_responses_ws(client_ws),
timeout=2.0,
)
assert handler.ws_sessions.active_count() == 0
# We must still have recorded exactly one session duration.
assert len(handler.metrics.ws_session_durations) == 1
@pytest.mark.asyncio
async def test_upstream_error_mid_stream_classifies_as_upstream_error():
upstream_events = [
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
]
upstream = _FakeUpstream(
upstream_events,
raise_mid_stream=RuntimeError("boom from upstream"),
)
fake_ws_mod = _make_fake_websockets_module(upstream)
client_ws = _FakeWebSocket(
frames=[_first_frame()],
hold_after_initial=True,
)
handler = _DummyOpenAIHandler()
with patch.dict(sys.modules, {"websockets": fake_ws_mod}):
await asyncio.wait_for(
handler.handle_openai_responses_ws(client_ws),
timeout=2.0,
)
assert handler.ws_sessions.active_count() == 0
assert handler.metrics.termination_causes
assert handler.metrics.termination_causes[-1] == "upstream_error"
@pytest.mark.asyncio
async def test_upstream_connect_failure_still_deregisters_cleanly():
"""Handshake-phase leak must be impossible: if upstream connect
raises before relay tasks are created, the session is still
registered+deregistered cleanly (or never registered). Either way,
no leak.
"""
class _BoomUpstream:
async def __aenter__(self):
raise RuntimeError("upstream refused")
async def __aexit__(self, exc_type, exc, tb):
return None
fake_ws_mod = MagicMock()
fake_ws_mod.connect = MagicMock(return_value=_BoomUpstream())
fake_ws_mod.Subprotocol = str
client_ws = _FakeWebSocket(frames=[_first_frame()])
handler = _DummyOpenAIHandler()
async def _fallback(*args, **kwargs):
return None
handler._ws_http_fallback = _fallback # type: ignore[assignment]
with patch.dict(sys.modules, {"websockets": fake_ws_mod}):
await handler.handle_openai_responses_ws(client_ws)
assert handler.ws_sessions.active_count() == 0
@pytest.mark.asyncio
async def test_many_concurrent_sessions_cleanly_drained():
"""50 concurrent sessions: all drain; registry and named tasks go to 0."""
upstream_events = [
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
json.dumps({"type": "response.completed", "response": {"id": "r_1"}}),
]
async def run_one() -> None:
upstream = _FakeUpstream(list(upstream_events))
fake_ws_mod = _make_fake_websockets_module(upstream)
client_ws = _FakeWebSocket(frames=[_first_frame()])
handler = _DummyOpenAIHandler()
with patch.dict(sys.modules, {"websockets": fake_ws_mod}):
await handler.handle_openai_responses_ws(client_ws)
assert handler.ws_sessions.active_count() == 0
await asyncio.gather(*[run_one() for _ in range(50)])
# Global check: no codex-ws-* named task remains.
leaked = [
t
for t in asyncio.all_tasks()
if (t.get_name() or "").startswith("codex-ws-") and not t.done()
]
assert leaked == []

View file

@ -0,0 +1,177 @@
"""Unit tests for :class:`WebSocketSessionRegistry`.
These tests exercise the in-memory registry in isolation no network,
no WS server. They pin down register/deregister semantics, ``snapshot``
shape, and the task-attachment accounting that feeds the
``active_relay_tasks`` gauge.
"""
from __future__ import annotations
import asyncio
import pytest
from headroom.proxy.ws_session_registry import (
WebSocketSessionRegistry,
WSSessionHandle,
)
def _make_handle(session_id: str = "sess-1") -> WSSessionHandle:
return WSSessionHandle(
session_id=session_id,
request_id="req-1",
client_addr="127.0.0.1:12345",
upstream_url="wss://upstream/test",
)
def test_register_and_deregister_round_trip():
reg = WebSocketSessionRegistry()
handle = _make_handle()
reg.register(handle)
assert reg.active_count() == 1
assert reg.get("sess-1") is handle
reg.deregister("sess-1", cause="response_completed")
assert reg.active_count() == 0
assert reg.get("sess-1") is None
assert handle.termination_cause == "response_completed"
def test_deregister_unknown_session_is_idempotent():
reg = WebSocketSessionRegistry()
# Deregistering a never-registered session must not raise.
reg.deregister("never-seen", cause="unknown")
assert reg.active_count() == 0
def test_register_twice_same_session_does_not_double_count():
reg = WebSocketSessionRegistry()
handle = _make_handle()
reg.register(handle)
reg.register(handle) # idempotent — overwrite, do not double-increment
assert reg.active_count() == 1
reg.deregister("sess-1", cause="client_disconnect")
assert reg.active_count() == 0
def test_snapshot_shape_is_json_serializable():
reg = WebSocketSessionRegistry()
reg.register(_make_handle("a"))
reg.register(_make_handle("b"))
snapshot = reg.snapshot()
assert len(snapshot) == 2
ids = {entry["session_id"] for entry in snapshot}
assert ids == {"a", "b"}
sample = snapshot[0]
for key in (
"session_id",
"request_id",
"client_addr",
"upstream_url",
"age_seconds",
"relay_task_count",
"termination_cause",
):
assert key in sample
@pytest.mark.asyncio
async def test_attach_tasks_merges_and_tracks_active_count():
reg = WebSocketSessionRegistry()
handle = _make_handle()
reg.register(handle)
async def _idle() -> None:
await asyncio.sleep(5)
t1 = asyncio.create_task(_idle(), name="codex-ws-test-1")
t2 = asyncio.create_task(_idle(), name="codex-ws-test-2")
reg.attach_tasks("sess-1", [t1])
assert reg.active_relay_task_count() == 1
reg.attach_tasks("sess-1", [t2]) # merges, does not replace
assert reg.active_relay_task_count() == 2
# Deregister cancels the in-registry accounting; tasks themselves
# are the caller's responsibility to cancel — the registry is a
# bookkeeper, not an owner.
reg.deregister("sess-1", cause="client_disconnect")
assert reg.active_relay_task_count() == 0
t1.cancel()
t2.cancel()
for t in (t1, t2):
try:
await t
except asyncio.CancelledError:
pass
def test_deregister_releases_task_references():
"""Once deregistered, the handle's ``relay_tasks`` list is cleared.
This keeps large references (coroutine frames) from being held by
the registry after the session ends.
"""
reg = WebSocketSessionRegistry()
handle = _make_handle()
reg.register(handle)
# Use sentinel non-task objects; ``attach_tasks`` does not inspect type
# beyond appending to the list and updating counts.
class _FakeTask:
def __init__(self, name: str) -> None:
self._name = name
self._done = False
def done(self) -> bool:
return self._done
def cancel(self) -> bool:
self._done = True
return True
def get_name(self) -> str:
return self._name
fake_tasks = [_FakeTask("t-a"), _FakeTask("t-b")]
reg.attach_tasks("sess-1", fake_tasks)
assert len(handle.relay_tasks) == 2
reg.deregister("sess-1", cause="upstream_disconnect")
# After deregister, the handle's tasks list is empty — the registry
# does not retain references.
assert handle.relay_tasks == []
assert reg.active_relay_task_count() == 0
def test_attach_tasks_on_missing_session_is_noop():
reg = WebSocketSessionRegistry()
async def _idle() -> None:
await asyncio.sleep(0)
loop = asyncio.new_event_loop()
try:
t = loop.create_task(_idle())
reg.attach_tasks("missing", [t])
assert reg.active_relay_task_count() == 0
t.cancel()
loop.run_until_complete(asyncio.gather(t, return_exceptions=True))
finally:
loop.close()
def test_snapshot_age_seconds_is_non_negative():
reg = WebSocketSessionRegistry()
reg.register(_make_handle())
snapshot = reg.snapshot()
assert snapshot[0]["age_seconds"] >= 0.0

View file

@ -0,0 +1,487 @@
---
title: "fix: Codex proxy resilience under reconnect storms"
type: fix
status: active
date: 2026-04-17
origin: wiki/plans/2026-04-17-codex-proxy-runtime-analysis.md
---
# fix: Codex proxy resilience under reconnect storms
## Overview
Harden the shared Headroom proxy so it can survive real multi-agent Codex traffic — especially the **large Anthropic `/v1/messages?beta=true` reconnect/retry storm** that hits the proxy immediately after a restart — without appearing dead (`/livez` timing out, new `/v1/responses` websocket handshakes hanging) and without bypassing compression.
This is the follow-on work to the runtime analysis captured in the origin document. The previous branch (`fix/responses-retries-keep-compression`) fixed the upstream WS handshake/fallback issues and kept compression enabled. This plan addresses the **remaining long-lived runtime degradation** described in §4 of the origin ("Long-lived service degradation on `8787`").
The plan deliberately focuses on **observability + lifecycle hygiene + cold-start backpressure**, not more blind patching. Compression stays enabled throughout.
## Problem Frame
Aged `8787` processes enter a state where:
- new `GET /livez` requests time out
- new `/v1/responses` opening handshakes time out
- existing established streams continue working
- the process is alive, listens on the port, and sampling shows heavy ONNX thread activity
Controlled reproductions ruled out the obvious single-factor causes (port, launchd, memory alone, idle socket count, compression-preserving changes, cold Kompress load). The surviving hypotheses (§"What Is Still Plausible" in origin) converge on:
1. **Long-lived real traffic leaves stuck websocket relay tasks / lifecycle bookkeeping leaks** (H1)
2. **Real Codex traffic, not synthetic traffic, triggers slow hidden work** (H2)
3. **ONNX/memory amplifies but does not solely cause the failure** (H3)
4. **Shared-proxy reconnect/retry storms after restart drive the proxy into this state** (H4 — confirmed by the "Latest Correction" in origin)
Without stage timings, active-task introspection, or session bookkeeping, the next iteration of debugging will again rely on `sample`, `lsof`, and guesswork. This plan fixes that first, then layers cold-start backpressure on top — in that order — so each subsequent bug hunt converges faster.
(see origin: `wiki/plans/2026-04-17-codex-proxy-runtime-analysis.md`)
## Requirements Trace
- **R1.** `/livez` remains responsive during and immediately after a restart that triggers large Anthropic replay traffic from active agent sessions. (origin §"Latest Correction", §"Updated upstream patch focus" item 4)
- **R2.** Cold-start heavy assets (Kompress ONNX, memory embedder, tokenizers, tree-sitter parsers) are loaded once at startup and *shared* between all provider pipelines; concurrent first-use callers wait on a single future, not N parallel loads. (origin §"Updated upstream patch focus" item 1)
- **R3.** The proxy provides enough runtime observability to prove — with data, not guesses — whether a future degradation is WS lifecycle starvation, memory/embedder contention, replay-storm amplification, or something else. (origin §"Priority 1: add instrumentation, not more blind patching")
- **R4.** Codex websocket relay tasks are explicitly tracked and deterministically cancelled when either side of the relay exits; a leaked relay task cannot hold the process alive past the client's disconnect. (origin §"Code Paths Most Relevant To The Remaining Bug" item 3; H1)
- **R5.** Cold-start compression + memory-context work on the Anthropic path is **bounded in concurrency** so that N simultaneous large replay requests cannot monopolize the event loop and thread pool. Compression stays enabled. (origin §"Updated upstream patch focus" item 2)
- **R6.** A reproducible harness exists for the real-agent reconnect/retry scenario so regressions can be caught locally instead of only in production. (origin §"Priority 2: reproduce degradation with real Codex traffic on a fresh process")
- **R7.** Existing fork behavior is preserved: compression is not bypassed on WS/streaming; upstream WS retry/open-timeout hardening is kept; WS→HTTP fallback normalization is kept; memory-context fail-open timeout is kept. (origin §"What Was Changed In The Fork", §"Kept locally")
## Scope Boundaries
- **Not** re-introducing any "skip compression" fast paths. Compression-preserving direction is non-negotiable.
- **Not** re-litigating the launchd setup bug (already fixed upstream in dotfiles; out of this repo).
- **Not** redesigning the memory stack, embedder choice, or Kompress model. Those are upstream concerns.
- **Not** building a full distributed tracing system. The instrumentation added here is structured logs + in-process counters; OpenTelemetry hookup is a separate plan.
- **Not** changing the WS→HTTP fallback semantics beyond what's needed to plumb through the new request-id/session-id logging fields.
- **Not** modifying `headroom-ai[ml]` dependencies, HuggingFace model IDs, or the embedder's own backend.
### Deferred to Separate Tasks
- **OpenTelemetry / metrics exporter wiring**: the counters added here expose `prometheus_metrics` entries; OTLP export belongs in a follow-up.
- **External watchdog in the LaunchAgent**: a plist-level `WatchPaths`/`ThrottleInterval` revision belongs in the `.dotfiles` repo (see origin §"Important External Files"), not here. This plan only adds the in-process signal the watchdog would consume (§Unit 5).
- **Quantifying multi-agent reconnect budget**: tuning the `Unit 4` semaphore default via load testing is work for after merge.
## Context & Research
### Relevant Code and Patterns
- `headroom/proxy/handlers/openai.py:1206``handle_openai_responses_ws` (Codex WS entry point, 600+ lines)
- `headroom/proxy/handlers/openai.py:1559-1767` — upstream WS connect retry loop, `open_timeout` handling
- `headroom/proxy/handlers/openai.py:1815``_ws_http_fallback` (preserved as-is)
- `headroom/proxy/handlers/openai.py:1439-1452` — memory context timeout fail-open (preserved as-is)
- `headroom/proxy/handlers/anthropic.py:293``handle_anthropic_messages` (HTTP entry point, also covers the `?beta=true` replay traffic)
- `headroom/proxy/server.py:586-733``HeadroomProxy.startup` (where eager preload runs)
- `headroom/proxy/server.py:634-637` — current eager preload iterates **only** `anthropic_pipeline.transforms` and breaks on first match
- `headroom/proxy/server.py:301-308` — both pipelines currently share the same `transforms` list (so the module-level `_kompress_cache` is de facto shared, but this is fragile)
- `headroom/proxy/server.py:1190-1212``/livez`, `/readyz`, `/health` handlers (trivial JSON, no I/O)
- `headroom/proxy/memory_handler.py:134-207``MemoryHandler._ensure_initialized` (lazy, no `asyncio.Lock`)
- `headroom/transforms/content_router.py:1221-1297``eager_load_compressors` (Kompress + Magika + Code-Aware + SmartCrusher)
- `headroom/transforms/kompress_compressor.py:163-221``_load_kompress_onnx` / `_load_kompress` with module-level `_kompress_cache` + `threading.Lock`
- `headroom/proxy/request_logger.py` — existing structured request log sink (extend, don't replace)
- `headroom/proxy/prometheus_metrics.py` — existing `PrometheusMetrics` class; add counters/gauges here
- `headroom/proxy/helpers.py:138-207``_read_request_json` (pre-upstream work on Anthropic path)
### Institutional Learnings
- `docs/solutions/` does not exist in this repo. The `wiki/plans/` folder holds design-style documents; no rolling solutions log to mine.
- **Prior fork learning** (origin §"What Was Changed In The Fork"): keep compression on, retry upstream WS, normalize fallback body, wrap memory-context lookup in a timeout. These are invariants — Unit 2 and Unit 4 must not regress them.
- **Prior rollback learning** (origin §"Rolled back locally"): latency-first skips of compression / memory injection were rolled back. Any new "fast path" must not recreate that shape.
### External References
None used for this plan. Python `asyncio` lock, `asyncio.Semaphore`, `asyncio.Task`, and `asyncio.all_tasks()` semantics are sufficient; no framework-specific research needed. External research was intentionally skipped (§1.2: strong local patterns, team knows the area).
## Key Technical Decisions
- **Observe before mitigating.** Units 2 and 3 (instrumentation + lifecycle accounting) land before Unit 4 (backpressure) so the backpressure defaults can be tuned from real data instead of guessed. The origin explicitly calls this out as Priority 1.
- *Rationale:* the last round of patches was driven by symptoms; the next one should be driven by timings.
- **Share cold-start state across pipelines explicitly, not by accident.** Currently both `TransformPipeline` instances share the same `transforms` list by coincidence of construction (`server.py:301-308`). Unit 1 moves the preload out of "first matching transform on the Anthropic pipeline" into a startup-level orchestration step that holds references it can reuse across both pipelines, and adds a single `asyncio.Lock` around first-use paths so concurrent requests land on the same future.
- *Rationale:* origin §"Updated upstream patch focus" item 1: "make eager preload and request-time use share the same in-process singleton/cache".
- **Track WS sessions in a registry, not via ad-hoc `logger.info`.** A `WebSocketSessionRegistry` makes active-count a first-class observable and lets `/debug/ws-sessions` return something useful.
- **Explicit relay-task cancellation** replaces `asyncio.gather(..., return_exceptions=True)` for the two relay halves. When one side exits, the other is cancelled deterministically — no wait on TCP timeout, no task leak (addresses H1).
- **Bounded pre-upstream concurrency on the Anthropic path, not on all paths.** The Codex WS path already serializes naturally (one client → one upstream WS). The Anthropic HTTP path is where replay storms arrive. Limiting concurrency only where the problem actually exists keeps the blast radius tight.
- *Rationale:* origin §"Hypothesis 4" + §"Updated upstream patch focus" item 2.
- **Debug endpoints are loopback-only, always.** No config flag, no auth header — a remote IP gets a 404, period. This sidesteps "did someone accidentally expose task state?" as a concern.
- **Frontmatter, not freeform.** Unlike the existing `wiki/plans/` files this plan uses YAML frontmatter (`status: active`, `origin:`, etc.) to participate in the `ce:plan` deepening + search flow.
## Open Questions
### Resolved During Planning
- *Q: Target upstream or the fork?* Resolved: target the fork (`fix/responses-retries-keep-compression`), structured so each unit is cherry-pickable into a PR against `chopratejas/headroom#172`.
- *Q: Should Unit 5's debug endpoints require an explicit flag to enable?* Resolved: **no** — loopback-only gating is sufficient and the debug data is useless if it isn't always available when the process is struggling.
- *Q: Does `MemoryHandler._ensure_initialized` already have a concurrency guard?* Resolved: **no**. It relies on `self._initialized = True` flip, which is not atomic across `await` points. Unit 1 adds an `asyncio.Lock`.
- *Q: Are the WS relay tasks already cancelled on partner exit?* Resolved: **no**. `asyncio.gather(return_exceptions=True)` waits for both; the survivor only exits when its own loop raises. Unit 3 fixes this.
- *Q: Is the existing Kompress singleflight sufficient?* Resolved: partially. The `threading.Lock` serializes same-model-id loads, but holds during `hf_hub_download` (network I/O). Unit 1 supplements with an `asyncio.Lock` at the request-handler layer so async callers don't each spawn the thread-pool job.
### Deferred to Implementation
- **Exact default for the Anthropic pre-upstream semaphore.** Start at `max(2, min(8, os.cpu_count()))` and expose via `--anthropic-pre-upstream-concurrency` / `HEADROOM_ANTHROPIC_PRE_UPSTREAM_CONCURRENCY`. Tune after Unit 2's stage timings land.
- **Which asyncio task names are "long-lived" thresholds** for Unit 5's dump. Likely `> 2× median lifetime`, but the median is only observable after Unit 2 runs in production.
- **Whether the repro harness (Unit 6) needs to speak realistic Codex subprotocol framing or can synthesize enough with recorded frames.** Depends on how faithfully `tests/test_openai_codex_routing.py` fixtures already capture the handshake.
- **Final placement of `WebSocketSessionRegistry`**: `headroom/proxy/ws_session_registry.py` as a new module, or folded into `headroom/proxy/server.py`. Likely new module; confirm during implementation once imports are real.
## High-Level Technical Design
> *This illustrates the intended approach and is directional guidance for review, not implementation specification. The implementing agent should treat it as context, not code to reproduce.*
```
┌───────────────────────────────────────────────────────┐
│ HeadroomProxy.startup │
│ │
│ Unit 1: shared_warmup() │
│ • eager-load on both pipelines (not just first) │
│ • Kompress, Magika, Code-Aware, SmartCrusher │
│ • memory backend + embedder │
│ • populates WarmupRegistry singletons │
└───────────────┬────────────────────┬──────────────────┘
│ │
▼ ▼
┌─────────────────────────────┐ ┌─────────────────────────────┐
│ Codex WS path │ │ Anthropic HTTP path │
│ /v1/responses │ │ /v1/messages?beta=true │
│ │ │ │
│ Unit 3: session registry │ │ Unit 4: pre-upstream │
│ • register on accept │ │ Semaphore(N) │
│ • cancel partner on exit │ │ • guards _read_request │
│ • deregister in finally │ │ • deep-copy │
│ │ │ • first compression stage │
│ Unit 2: stage timings │ │ • memory-context lookup │
│ accept → first_frame │ │ │
│ → upstream_connect │ │ Unit 2: stage timings │
│ → upstream_first_event │ │ read_json → deep_copy → │
│ → total_session_ms │ │ compress → memory → │
└─────────────────────────────┘ │ upstream_first_byte │
└─────────────────────────────┘
│ │
└────────┬───────────┘
┌─────────────────────────────────────────┐
│ Unit 5: /debug/* (loopback-only) │
│ • /debug/tasks │
│ • /debug/ws-sessions │
│ • /debug/warmup │
└─────────────────────────────────────────┘
┌─────────────────────────────────────────┐
│ Unit 6: scripts/repro_codex_replay.py │
│ • spawn N concurrent Codex WS │
│ • fire Anthropic replay-shaped POSTs │
│ • assert /livez stays < 100ms
└─────────────────────────────────────────┘
```
Decision matrix for "where does this request wait?":
| Path | Pre-upstream gate | Relay lifecycle gate | Memory lookup timeout |
|-------------------------------|----------------------|----------------------|------------------------|
| `/v1/responses` (Codex WS) | *(none — natural 1:1)* | `WSSessionRegistry` | existing `wait_for` |
| `/v1/messages` (Anthropic) | `Semaphore(N)` (new) | *(HTTP-single-shot)* | existing `wait_for` |
| `/v1/chat/completions` (OpenAI) | *(unchanged)* | *(unchanged)* | *(unchanged)* |
| `/livez`, `/readyz` | *(none — must be free)* | *(n/a)* | *(n/a)* |
## Implementation Units
- [ ] **Unit 1: Shared cold-start warmup + async singleflight for memory init**
**Goal:** Make preload truthful and shared across all provider pipelines; ensure the first wave of concurrent requests after restart never kick duplicate expensive loads.
**Requirements:** R2, R7
**Dependencies:** None (pure startup change; no other unit depends on this landing first, but landing it first reduces noise in Unit 2's timings)
**Files:**
- Modify: `headroom/proxy/server.py` (startup orchestration around lines 627-675)
- Modify: `headroom/proxy/memory_handler.py` (add `asyncio.Lock` in `_ensure_initialized`)
- Modify: `headroom/transforms/content_router.py` (no behavior change — ensure `eager_load_compressors` is idempotent when called via both pipelines)
- Create: `headroom/proxy/warmup.py` (new `WarmupRegistry` holding preloaded handles)
- Test: `tests/test_proxy_warmup.py` (new)
- Test: `tests/test_memory_handler_concurrent_init.py` (new)
**Approach:**
- Introduce `WarmupRegistry` with typed slots for kompress, magika, code-aware, smart-crusher, memory-embedder, memory-backend. Populate during `HeadroomProxy.startup`, expose via `proxy.warmup` for `/debug/warmup` (Unit 5) and `/readyz`.
- Iterate **both** `self.anthropic_pipeline.transforms` and `self.openai_pipeline.transforms` when calling `eager_load_compressors`; dedupe by `id(transform)` so shared transforms don't double-load.
- Preload the memory embedder explicitly — today `ensure_initialized` initializes the backend but the embedder's first-use cost may still be deferred until the first `search_and_format_context`. Force one warm-up encode (a single short string) to pre-compile the ONNX graph.
- Add `asyncio.Lock` in `MemoryHandler._ensure_initialized` so concurrent first callers await one load, not N.
- Replace the "break after first matching transform" loop with explicit orchestration. Log whether the warmup was a no-op (already loaded) vs. a fresh load.
**Patterns to follow:**
- Existing `eager_status` dict shape in `server.py:630-666`.
- `threading.Lock` module-level pattern in `kompress_compressor.py` (for the sync side); supplement with `asyncio.Lock` in memory_handler (async side).
- `request_logger`'s structured-log pattern for startup events.
**Test scenarios:**
- Happy path: starting the proxy with `optimize=True` logs one preload event per component and `WarmupRegistry` reports all slots loaded.
- Happy path: starting with `optimize=False` populates `WarmupRegistry` on first request instead (lazy path still works).
- Edge case: `optimize=True` with `enable_kompress=False``WarmupRegistry.kompress` is `None`, no error logged.
- Edge case: both pipelines share the same router transform — `eager_load_compressors` runs once, not twice.
- Integration: `MemoryHandler._ensure_initialized` called 10 concurrent times from different tasks — only one backend init runs (assert via counter on `LocalBackend.__init__` hit count).
- Integration: embedder warm-up encode is issued at startup — verify via mock that `embed_text("warmup")` (or equivalent) was called once during startup, not lazily on first request.
- Error path: memory backend init raises — startup still completes, `WarmupRegistry.memory_backend` is `None`, health reports degraded memory.
**Verification:**
- Startup logs show each component preloaded exactly once, with timings.
- `/readyz` reports all configured subsystems as `initialized=true` before accepting traffic.
- Concurrent first-request simulation (`asyncio.gather` 20 requests) triggers only one Kompress load in logs.
---
- [ ] **Unit 2: Stage-timing instrumentation on Codex WS and Anthropic HTTP paths**
**Goal:** Emit structured per-request timings for every stage that plausibly contributes to cold-start or degradation latency, so the next debugging session starts from data.
**Requirements:** R3
**Dependencies:** None (independent), but landing alongside Unit 1 gives early visibility into whether warmup actually helped.
**Files:**
- Modify: `headroom/proxy/handlers/openai.py` (WS path at lines 1206+; HTTP path at 798+)
- Modify: `headroom/proxy/handlers/anthropic.py` (`handle_anthropic_messages` around line 293 and downstream)
- Modify: `headroom/proxy/request_logger.py` (extend schema)
- Modify: `headroom/proxy/prometheus_metrics.py` (add histograms per stage)
- Modify: `headroom/proxy/helpers.py` (thread `stage_timer` through `_read_request_json` / compression helpers)
- Create: `headroom/proxy/stage_timer.py` (small context-manager util)
- Test: `tests/test_stage_timer.py` (new)
- Test: `tests/test_openai_codex_ws_timings.py` (new)
- Test: `tests/test_anthropic_stage_timings.py` (new)
**Approach:**
- Add a `StageTimer` context manager with `stage_timer.measure("memory_context")` support; emits structured log on exit. Unified across handlers.
- For Codex WS: instrument `accept`, `first_client_frame`, `upstream_connect`, `upstream_first_event`, `memory_context`, `compression`, `total_session`. Log on session close with all fields.
- For Anthropic HTTP: instrument `read_request_json`, `deep_copy`, `compression_first_stage`, `memory_context`, `upstream_connect`, `upstream_first_byte`, `total_pre_upstream`. Log after first upstream byte arrives.
- Also export each stage as a Prometheus histogram — the existing `PrometheusMetrics` class has the right pattern; mirror it.
- `request_id` is already threaded through both handlers. Add a `session_id` (UUID generated at WS accept / HTTP request start) so multi-turn sessions are correlatable.
**Execution note:** Land the util + test (`StageTimer`) first; only then plumb it through the two handlers. Keeps the diff reviewable.
**Patterns to follow:**
- `headroom/proxy/request_logger.py` structured-field schema.
- `PrometheusMetrics.record_request` histogram pattern.
- Existing `request_id = await self._next_request_id()` at `openai.py:1232`.
**Test scenarios:**
- Happy path: a full Codex WS session (accept → response.completed) emits one structured log line with all 7 stage fields populated and a positive `total_session` > 0.
- Happy path: a full Anthropic HTTP request emits one log line with all pre-upstream stage fields populated.
- Edge case: a session that exits during upstream connect (no `first_event`) logs `upstream_first_event=null` without raising.
- Edge case: `request_id` and `session_id` appear together on every log line.
- Error path: a timeout in `memory_context` logs `memory_context` duration = the timeout value, not `null` (prove the timer captures the failure window).
- Integration: Prometheus histograms emit non-zero observations after a real request round-trip.
**Verification:**
- Every `/v1/responses` WS session and every `/v1/messages` request produces exactly one timing log line.
- `/metrics` endpoint shows the new histogram series.
- No existing tests regress; the request_logger schema additions are backward-compatible (new fields only).
---
- [x] **Unit 3: WebSocket session registry + deterministic relay-task cancellation**
**Goal:** Eliminate the "aged process has leaked relay tasks" hypothesis by making every WS session explicitly tracked and both relay tasks deterministically cancelled when either exits.
**Requirements:** R4
**Dependencies:** Unit 2 (uses the same `session_id`)
**Files:**
- Create: `headroom/proxy/ws_session_registry.py` (new)
- Modify: `headroom/proxy/handlers/openai.py` (`handle_openai_responses_ws` at 1206; relay task construction around 1559-1767; `_upstream_to_client` at 1565)
- Modify: `headroom/proxy/prometheus_metrics.py` (add `active_ws_sessions` gauge, `active_relay_tasks` gauge, `ws_session_duration` histogram)
- Test: `tests/test_ws_session_registry.py` (new)
- Test: `tests/test_openai_codex_ws_lifecycle.py` (new)
**Approach:**
- `WebSocketSessionRegistry` holds `dict[session_id, WSSessionHandle]` where handle tracks: `started_at`, `client_addr`, `upstream_url`, `relay_tasks: list[asyncio.Task]`, `last_activity_at`.
- Register on `websocket.accept()` success; deregister in a `try/finally` around the whole handler body.
- Replace `asyncio.gather(_client_to_upstream(), _upstream_to_client(), return_exceptions=True)` with explicit `asyncio.create_task(...)` for each, then `asyncio.wait(..., return_when=FIRST_COMPLETED)`. Cancel the other task explicitly, then `await` the cancelled task's `.cancelled()` settlement (suppress `CancelledError`).
- Emit one structured log line per session termination with cause (`client_disconnect` / `upstream_disconnect` / `upstream_error` / `client_error`).
**Execution note:** Start from a failing integration test that asserts "after a client disconnects mid-stream, the upstream relay task is `done()` within 100ms and the registry is empty". Then make it pass.
**Patterns to follow:**
- Existing `_client_to_upstream()` and `_upstream_to_client()` structure in `openai.py:1559-1767`.
- `contextlib.suppress(asyncio.CancelledError)` pattern.
- `PrometheusMetrics` gauge update pattern.
**Test scenarios:**
- Happy path: a normal session completes; registry is empty after `response.completed`; both relay tasks `.done()` is True.
- Edge case: client disconnects mid-stream; upstream relay task is cancelled within 100ms; registry no longer contains the session.
- Edge case: upstream closes first; client-side relay task is cancelled within 100ms; client receives a clean close frame.
- Edge case: 50 concurrent WS sessions open and close; `active_ws_sessions` gauge rises to 50 then returns to 0; no tasks remain in `asyncio.all_tasks()` with the relay task name pattern.
- Error path: upstream connect fails before relay tasks are created; registry is deregistered cleanly (no leak from the handshake phase).
- Error path: `_upstream_to_client` raises mid-stream; client task is cancelled; registry reports `termination_cause=upstream_error`.
- Integration: the `/debug/ws-sessions` endpoint (Unit 5) returns live data consistent with the registry under load.
**Verification:**
- After a torture-test of 100 rapid connect/disconnect cycles, `asyncio.all_tasks()` count returns to baseline within 1s.
- `active_ws_sessions` gauge returns to 0 after all sessions close.
- No `RuntimeWarning: coroutine was never awaited` in logs.
---
- [ ] **Unit 4: Bounded pre-upstream concurrency for Anthropic replay storms**
**Goal:** Prevent the cold-start replay storm from occupying every event-loop slot and thread-pool worker with deep-copy / compression / memory-context work before upstream receives the request.
**Requirements:** R1, R5, R7
**Dependencies:** Unit 2 (must have stage timings landed so the default semaphore size can be tuned from real numbers; the unit itself can ship with a conservative default and be retuned later)
**Files:**
- Modify: `headroom/proxy/handlers/anthropic.py` (wrap the pre-upstream phases in `handle_anthropic_messages` at ~293 through first upstream call)
- Modify: `headroom/proxy/models.py` or config module (add `anthropic_pre_upstream_concurrency: int` config field)
- Modify: `headroom/proxy/server.py` (construct the `asyncio.Semaphore` on `HeadroomProxy.__init__` so it's per-process, not per-request)
- Modify: CLI surface (add `--anthropic-pre-upstream-concurrency` flag + env var)
- Test: `tests/test_anthropic_pre_upstream_backpressure.py` (new)
**Approach:**
- Add `self.anthropic_pre_upstream_sem = asyncio.Semaphore(config.anthropic_pre_upstream_concurrency or max(2, min(8, os.cpu_count())))` in `HeadroomProxy.__init__`.
- In `handle_anthropic_messages`, wrap the region from `_read_request_json` through the first `self.http_client.send()` / `stream()` call in `async with self.anthropic_pre_upstream_sem:`.
- **Critically:** `/livez` and `/readyz` do not go through this semaphore. They remain free even under replay storm.
- Emit a log (not a warning — expected under load) when a request waits > 100ms for the semaphore, so we can see queueing in the Unit 2 stage timings.
- Default config: `HEADROOM_ANTHROPIC_PRE_UPSTREAM_CONCURRENCY` env var honored; CLI flag overrides; unset defaults to the computed floor.
**Patterns to follow:**
- Existing config field addition pattern in `headroom/proxy/models.py`.
- Existing CLI flag plumbing in `headroom/cli/`.
- `request_logger.info(..., stage="pre_upstream_wait_ms", ...)` field extension from Unit 2.
**Test scenarios:**
- Happy path: a single request passes through with negligible `pre_upstream_wait_ms`.
- Edge case: N+1 concurrent requests (where N = configured concurrency) — exactly one waits; all complete; `pre_upstream_wait_ms > 0` only on the waiter.
- Edge case: `anthropic_pre_upstream_concurrency=1` serializes two concurrent requests deterministically (useful for tests).
- Error path: a request that raises inside the critical section releases the semaphore (verify via counter that `_value` returns to baseline).
- Integration: under a synthetic storm of 20 concurrent large POSTs (simulating the retry replay), `/livez` response time stays under 100ms (p99 from Unit 2 timings).
- Integration: compression is *not* bypassed — assert a known large body still produces a compressed upstream request (no regression of R7).
**Verification:**
- Under `scripts/repro_codex_replay.py` (Unit 6), `/livez` p99 stays under 100ms during the storm phase.
- Counter-factual: setting concurrency to `10000` (effectively unbounded) reproduces the original starvation in the harness.
---
- [ ] **Unit 5: Loopback-only debug introspection endpoints**
**Goal:** Make "what is this process doing right now?" a single `curl` away when degradation happens, instead of `sample` + `lsof` + guesswork.
**Requirements:** R3
**Dependencies:** Unit 3 (ws session data), Unit 1 (warmup registry data)
**Files:**
- Modify: `headroom/proxy/server.py` (add `/debug/tasks`, `/debug/ws-sessions`, `/debug/warmup` routes near the existing `/livez` at 1190)
- Create: `headroom/proxy/debug_introspection.py` (pure functions that serialize state)
- Create: `headroom/proxy/loopback_guard.py` (middleware / dep that 404s non-loopback requests)
- Test: `tests/test_proxy_debug_endpoints.py` (new)
**Approach:**
- `loopback_guard`: inspect `request.client.host`; if not in `{"127.0.0.1", "::1", "localhost"}`, return 404 with no body. A 404 (not 403) keeps debug endpoints invisible to external scanners.
- `/debug/tasks` returns `asyncio.all_tasks()` with: name, coro name, age, current-line (via `get_coro().cr_frame.f_code.co_qualname`), stack depth. Sort by age desc.
- `/debug/ws-sessions` returns the registry dump from Unit 3.
- `/debug/warmup` returns `WarmupRegistry` state from Unit 1 + whether each slot is `loaded` / `loading` / `null`.
- All three return JSON. None mutate state. None block.
**Patterns to follow:**
- Existing `/readyz` handler shape in `server.py:1204-1207`.
- Loopback detection: `request.client.host` is already FastAPI's standard; mirror the guard style used elsewhere in the repo if any; otherwise a small dependency function.
**Test scenarios:**
- Happy path: `curl 127.0.0.1:8787/debug/tasks` returns 200 with a JSON array; each entry has `name`, `age_seconds`, `coro`.
- Edge case: `/debug/tasks` called during a live WS session lists the relay tasks with non-zero age.
- Edge case: loopback guard returns 404 (not 403) for a simulated non-loopback client.
- Edge case: `/debug/warmup` reports `memory_backend=loaded` after Unit 1's startup completes; reports `loading` if called during startup (race window).
- Error path: `asyncio.all_tasks()` raising (hypothetical) — handler returns 500 with structured error, doesn't crash the server.
- Integration: `/debug/ws-sessions` output is consistent with the session count gauge from Unit 3 during a load test.
**Verification:**
- Manual: from a remote IP, all three endpoints return 404.
- Manual: during the Unit 6 harness, `/debug/tasks` shows relay tasks matching the expected count.
- Documented in `wiki/proxy.md` under a new "Debug endpoints" subsection.
---
- [ ] **Unit 6: Repro harness for multi-agent reconnect storm**
**Goal:** Produce a single script that reproducibly exercises the failure class from origin §"Latest Correction" (active agent reconnects + large replay requests) against a fresh local proxy, so the fix is provable and regressions are catchable.
**Requirements:** R6
**Dependencies:** Unit 2 (so the harness can assert on stage timings), Unit 4 (so the harness can verify backpressure helps)
**Files:**
- Create: `scripts/repro_codex_replay.py` (new)
- Create: `scripts/fixtures/anthropic_replay_body.json` (recorded shape of a large replay request body, sanitized)
- Create: `scripts/fixtures/codex_response_create_frame.json` (recorded first-frame shape)
- Modify: `scripts/README.md` (add harness section)
**Approach:**
- CLI: `python scripts/repro_codex_replay.py --url http://127.0.0.1:8787 --ws-clients 8 --anthropic-clients 4 --duration 30s`
- Phase 1 (warmup): open 1 WS, send one `response.create`, drain to `response.completed`. Confirms proxy is live.
- Phase 2 (storm): simultaneously:
- Open N Codex WS connections, send one `response.create` each, keep the session open for the duration.
- Fire M concurrent large Anthropic POSTs shaped like agent-reconnect replays (from the fixture). Each retries on connection error for up to 60s, mimicking real agent behavior.
- Throughout: probe `/livez` every 250ms, record p50/p95/p99.
- Exit code: non-zero if `/livez` p99 exceeds 500ms during the storm (soft assertion).
- Print a summary: per-phase timings, livez stats, count of successful Codex `response.completed`, count of Anthropic successes.
**Execution note:** Harness-first is fine here — the fixture JSONs can be hand-crafted initially and swapped for captured ones later.
**Patterns to follow:**
- Existing benchmark harness in `benchmarks/` for run-loop structure.
- `websockets` library usage already in `tests/test_proxy_codex_route_aliases.py`.
**Test scenarios:**
- Test expectation: light — the harness is a script, not a library. A smoke test in `tests/test_scripts/test_repro_codex_replay_smoke.py` launches it against a mock server and verifies it exits 0 with the expected summary shape.
- Happy path: harness runs against a healthy proxy; exit code 0; summary shows livez p99 < 500ms.
- Error path: harness invoked with `--url` pointing at a closed port; exits with clear "connection refused" message, non-zero code, within 5 seconds.
- Integration: after Unit 4 lands, harness with default Anthropic concurrency limits shows livez p99 < 100ms; harness with `HEADROOM_ANTHROPIC_PRE_UPSTREAM_CONCURRENCY=10000` (unbounded) reproduces the original starvation (p99 > 5s).
**Verification:**
- CI smoke test runs the script against a mock server on every PR.
- `wiki/proxy.md` gets a "Reproducing the reconnect storm" subsection referencing the script.
## System-Wide Impact
- **Interaction graph:** the new `WarmupRegistry` (Unit 1) and `WebSocketSessionRegistry` (Unit 3) are accessed from `HeadroomProxy` during request handling and by `/debug/*` routes. Both are single-writer-per-event-loop; no locks needed beyond the `asyncio.Lock` in `MemoryHandler._ensure_initialized` (Unit 1).
- **Error propagation:** Unit 3 changes `asyncio.gather(..., return_exceptions=True)` to explicit `asyncio.wait(FIRST_COMPLETED)` + cancel. This means partner-side exceptions surface earlier — previously, a client-side error would still let upstream-side drain to completion. Verify in test scenarios that this does not drop in-flight frames that the client had already queued.
- **State lifecycle risks:** Unit 3's session registry must be deregistered in the outermost `finally` — a leak here re-creates the exact bug this plan fixes. Unit 1's `asyncio.Lock` must be released on exception paths (use `async with`, not manual acquire/release).
- **API surface parity:** no change to `/v1/responses` or `/v1/messages` request/response contract. The only new routes are `/debug/*`, loopback-gated.
- **Integration coverage:** Unit 3 + Unit 6 together are the main integration story — unit-test-only coverage of relay cancellation is insufficient; the repro harness is how we prove it in aggregate.
- **Unchanged invariants:** Compression stays enabled on all paths. Upstream WS retry/open-timeout handling (`openai.py:1559-1767`) is untouched. WS→HTTP fallback normalization (`openai.py:1815`) is untouched except for threading through the new `session_id` / stage-timer context. Memory-context fail-open `asyncio.wait_for` (`openai.py:1439-1452`, and the equivalent on the Anthropic path) is untouched. `/livez` logic stays trivial and IO-free.
## Risks & Dependencies
| Risk | Likelihood | Impact | Mitigation |
|------|-----------|--------|------------|
| Unit 3's new cancellation semantics drop in-flight frames the client had already queued on the upstream side. | Med | High (user-visible stream corruption) | Explicit test scenarios for "client disconnect mid-stream" verify no frames silently lost; harness (Unit 6) measures successful `response.completed` count. |
| Unit 4's semaphore default is too conservative and adds visible latency under normal multi-agent use. | Med | Med | Default computed from `cpu_count()` with a floor of 2 and ceiling of 8; tunable via CLI + env; Unit 2 timings will surface if the wait is user-visible. |
| Unit 4's semaphore default is too permissive and doesn't actually stop starvation under real replay storms. | Med | High (plan fails R1) | Unit 6 harness tests both bounded and unbounded configs; ship with conservative default, relax based on harness data post-merge. |
| Unit 5's `/debug/tasks` exposes sensitive state (e.g., in-flight request bodies via coro locals). | Low | Med (info-leak if endpoint ever becomes non-loopback) | Serializer strips coro locals; only task *metadata* (name, age, qualname) is exposed, never arguments. Loopback guard is enforced by middleware before the handler runs. |
| Unit 1's embedder warm-up encode triggers a HuggingFace download in CI and slows test runs. | Med | Low | Warm-up encode is skipped when `optimize=False`; test environment uses a stub embedder; document in `tests/conftest.py`. |
| The real degradation is *not* caused by WS task leaks or replay storms (our top hypotheses), and this plan's instrumentation exposes it but the mitigations don't fix it. | Low | Med | This is fine. The plan is observe-first; Unit 2 + Unit 5 give the data for the next iteration. The plan is valuable even if Units 3 and 4 turn out to be insufficient on their own. |
| `asyncio.Lock` added in Unit 1 is held across the memory backend init, which itself calls `await HierarchicalMemory.create()` — if that hangs, first-request + `ensure_initialized` both hang. | Low | High (deadlock at startup) | Wrap `_ensure_initialized` in a `wait_for(..., timeout=STARTUP_INIT_TIMEOUT_SECONDS)` (configurable, default 30s). On timeout, log error and leave `_initialized=False` so subsequent requests retry. |
| Tests that monkey-patch `MemoryHandler._initialized` directly break because of the new lock ordering. | Low | Low | Audit `tests/test_proxy_memory_integration.py` and fixtures; use `ensure_initialized` via the public entry point only. |
## Documentation / Operational Notes
- Add a "Debug endpoints" subsection to `wiki/proxy.md` covering the three new loopback-only routes, their output shape, and the loopback-only guarantee.
- Add a "Reproducing the reconnect storm" subsection to `wiki/proxy.md` referencing `scripts/repro_codex_replay.py`.
- Update `wiki/metrics.md` with the new histogram series names from Unit 2 and the new gauges from Unit 3.
- Update `CHANGELOG.md` under an "Unreleased" heading with: (a) compression preserved — no behavioral change for existing Codex users; (b) new stage timings visible in logs; (c) new `HEADROOM_ANTHROPIC_PRE_UPSTREAM_CONCURRENCY` env var; (d) new `/debug/*` loopback endpoints.
- **Rollout note:** deploy behind the existing launchd flow; monitor `/metrics` for the new histograms; monitor logs for `pre_upstream_wait_ms` — sustained non-zero values across many requests mean the Unit 4 default is too tight for this machine's load profile.
- **Rollback note:** each unit is cherry-pickable. If Unit 4 causes regression, revert only Unit 4 (semaphore removal); Units 13 and 56 have no user-visible behavior change.
## Sources & References
- **Origin document:** `wiki/plans/2026-04-17-codex-proxy-runtime-analysis.md`
- Related code:
- `headroom/proxy/handlers/openai.py``handle_openai_responses_ws`, `_client_to_upstream`, `_upstream_to_client`, `_ws_http_fallback`
- `headroom/proxy/handlers/anthropic.py``handle_anthropic_messages`
- `headroom/proxy/server.py``HeadroomProxy.__init__`, `.startup`, health routes
- `headroom/proxy/memory_handler.py``MemoryHandler._ensure_initialized`
- `headroom/transforms/content_router.py``eager_load_compressors`
- `headroom/transforms/kompress_compressor.py``_load_kompress_onnx`, `_kompress_cache`
- Related PRs/issues:
- Upstream: `https://github.com/chopratejas/headroom/issues/172`
- Fork branch: `fix/responses-retries-keep-compression` (commit `0b11637`)
- External docs: none used for this plan.