Skip to content

Commit 4853d41

Browse files
committed
fix(codex): offload project resolution
1 parent d3d6f8b commit 4853d41

4 files changed

Lines changed: 161 additions & 19 deletions

File tree

headroom/providers/codex/project_context.py

Lines changed: 66 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -2,13 +2,15 @@
22

33
from __future__ import annotations
44

5+
import asyncio
56
import json
67
import logging
78
import os
89
import sqlite3
910
import time
1011
from collections.abc import Mapping
1112
from dataclasses import dataclass
13+
from functools import lru_cache
1214
from pathlib import Path
1315
from typing import Any, Literal, cast
1416

@@ -26,6 +28,7 @@
2628
_MAX_METADATA_BYTES = 16 * 1024
2729
_SQLITE_TIMEOUT_SECONDS = 0.1
2830
_SQLITE_ATTEMPTS = 2
31+
_ROLLOUT_CACHE_MAX_ENTRIES = 256
2932

3033

3134
@dataclass(frozen=True, slots=True)
@@ -121,6 +124,27 @@ def resolve(
121124
metadata_reason,
122125
)
123126

127+
async def resolve_async(
128+
self,
129+
*,
130+
headers: Mapping[str, str],
131+
body: Mapping[str, Any],
132+
pinned_cwd: Path | None = None,
133+
project_root_override: str | None = None,
134+
) -> CodexResolvedProject:
135+
"""Resolve optional project context without blocking async model traffic."""
136+
try:
137+
return await asyncio.to_thread(
138+
self.resolve,
139+
headers=headers,
140+
body=body,
141+
pinned_cwd=pinned_cwd,
142+
project_root_override=project_root_override,
143+
)
144+
except Exception:
145+
logger.warning("event=codex_project_resolution_failed", exc_info=True)
146+
return self._skip("resolver_failed")
147+
124148
@staticmethod
125149
def _header(headers: Mapping[str, str], name: str) -> str | None:
126150
lowered = name.lower()
@@ -150,6 +174,10 @@ def _turn_identity(
150174
metadata = container.get("client_metadata")
151175
identity = self._identity_from_metadata(metadata)
152176
if identity is not None:
177+
if sum(len(value.encode("utf-8")) for value in identity if value) > (
178+
_MAX_METADATA_BYTES
179+
):
180+
return None, "codex-client-metadata", "metadata_too_large"
153181
return identity, "codex-client-metadata", "resolved"
154182
return None, "unresolved", "metadata_missing"
155183

@@ -359,7 +387,30 @@ def _rollout_for_thread(state_path: Path, thread_id: str) -> tuple[Path | None,
359387
return None, "state_locked"
360388

361389
def _cwd_from_rollout(self, rollout: Path, turn_id: str) -> tuple[Path | None, str]:
362-
matches: set[Path] = set()
390+
try:
391+
metadata = rollout.stat()
392+
except OSError:
393+
return None, "rollout_stale"
394+
fingerprint = (
395+
metadata.st_dev,
396+
metadata.st_ino,
397+
metadata.st_size,
398+
metadata.st_mtime_ns,
399+
metadata.st_ctime_ns,
400+
)
401+
raw_cwds, reason = _cached_raw_cwds_from_rollout(rollout, fingerprint, turn_id)
402+
if reason != "resolved":
403+
return None, reason
404+
matches = {cwd for raw_cwd in raw_cwds if (cwd := self._canonical_cwd(raw_cwd))}
405+
if len(matches) > 1:
406+
return None, "turn_ambiguous"
407+
if not matches:
408+
return None, "turn_context_missing"
409+
return next(iter(matches)), "resolved"
410+
411+
@staticmethod
412+
def _read_raw_cwds_from_rollout(rollout: Path, turn_id: str) -> tuple[tuple[str, ...], str]:
413+
matches: set[str] = set()
363414
truncated = False
364415
try:
365416
with rollout.open(encoding="utf-8") as handle:
@@ -378,18 +429,14 @@ def _cwd_from_rollout(self, rollout: Path, turn_id: str) -> tuple[Path | None, s
378429
continue
379430
cwd = payload.get("cwd")
380431
if isinstance(cwd, str):
381-
canonical = self._canonical_cwd(cwd)
382-
if canonical is not None:
383-
matches.add(canonical)
432+
matches.add(cwd)
384433
except OSError:
385-
return None, "rollout_stale"
434+
return (), "rollout_stale"
386435
if truncated:
387-
return None, "rollout_truncated"
388-
if len(matches) > 1:
389-
return None, "turn_ambiguous"
436+
return (), "rollout_truncated"
390437
if matches:
391-
return next(iter(matches)), "resolved"
392-
return None, "turn_context_missing"
438+
return tuple(sorted(matches)), "resolved"
439+
return (), "turn_context_missing"
393440

394441
@staticmethod
395442
def _skip(reason: str) -> CodexResolvedProject:
@@ -402,4 +449,13 @@ def _skip(reason: str) -> CodexResolvedProject:
402449
)
403450

404451

452+
@lru_cache(maxsize=_ROLLOUT_CACHE_MAX_ENTRIES)
453+
def _cached_raw_cwds_from_rollout(
454+
rollout: Path,
455+
_fingerprint: tuple[int, int, int, int, int],
456+
turn_id: str,
457+
) -> tuple[tuple[str, ...], str]:
458+
return CodexProjectContextResolver._read_raw_cwds_from_rollout(rollout, turn_id)
459+
460+
405461
__all__ = ["CodexProjectContextResolver", "CodexResolvedProject"]

headroom/proxy/handlers/openai.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5163,7 +5163,7 @@ async def handle_openai_responses(
51635163
or None
51645164
)
51655165
if client == "codex":
5166-
codex_project = CodexProjectContextResolver().resolve(
5166+
codex_project = await CodexProjectContextResolver().resolve_async(
51675167
headers=dict(request.headers),
51685168
body=body,
51695169
project_root_override=codex_project_root_override,
@@ -7057,7 +7057,7 @@ async def _prepare_memory_frame(frame_body: dict[str, Any], frame_raw: str) -> s
70577057

70587058
resolved_project_root_override = ws_project_root_override
70597059
if client == "codex":
7060-
resolved_project = ws_project_resolver.resolve(
7060+
resolved_project = await ws_project_resolver.resolve_async(
70617061
headers=ws_headers,
70627062
body=frame_body,
70637063
pinned_cwd=ws_pinned_project_cwd,

tests/test_codex_project_context.py

Lines changed: 83 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
import json
55
import sqlite3
66
import sys
7+
import threading
78
from pathlib import Path
89
from types import SimpleNamespace
910
from unittest.mock import patch
@@ -129,6 +130,42 @@ def test_codex_http_projects_are_isolated_and_ws_mismatch_fails_closed(
129130
assert mismatch.cwd is None
130131
assert mismatch.reason == "project_mismatch"
131132

133+
_append_turn(rollout_a, "turn-a", project_b)
134+
changed_rollout = resolver.resolve(
135+
headers={},
136+
body={"client_metadata": {"thread_id": "thread-a", "turn_id": "turn-a"}},
137+
)
138+
assert changed_rollout.cwd is None
139+
assert changed_rollout.reason == "turn_ambiguous"
140+
141+
142+
def test_cached_rollout_recanonicalizes_symlink_before_pinned_check(
143+
monkeypatch, tmp_path: Path
144+
) -> None:
145+
codex_home = tmp_path / "codex"
146+
codex_home.mkdir()
147+
project_a = tmp_path / "a"
148+
project_b = tmp_path / "b"
149+
project_a.mkdir()
150+
project_b.mkdir()
151+
project_link = tmp_path / "current-project"
152+
project_link.symlink_to(project_a, target_is_directory=True)
153+
rollout = codex_home / "rollout.jsonl"
154+
_seed_rollout(rollout, "turn", project_link)
155+
_seed_thread(codex_home, "thread", rollout)
156+
monkeypatch.setenv("CODEX_HOME", str(codex_home))
157+
resolver = CodexProjectContextResolver()
158+
body = {"client_metadata": {"thread_id": "thread", "turn_id": "turn"}}
159+
160+
first = resolver.resolve(headers={}, body=body)
161+
project_link.unlink()
162+
project_link.symlink_to(project_b, target_is_directory=True)
163+
second = resolver.resolve(headers={}, body=body, pinned_cwd=first.cwd)
164+
165+
assert first.cwd == project_a.resolve()
166+
assert second.cwd is None
167+
assert second.reason == "project_mismatch"
168+
132169

133170
def test_turn_specific_resume_and_fork_use_exact_context(monkeypatch, tmp_path: Path) -> None:
134171
codex_home = tmp_path / "codex"
@@ -284,7 +321,7 @@ def test_explicit_overrides_precede_state_and_body_cwd(monkeypatch, tmp_path: Pa
284321
)
285322

286323
assert project_id.source == "x-headroom-project-id"
287-
assert project_id.project_key == "chosen"
324+
assert project_id.project_key.startswith("chosen-")
288325
assert explicit_cwd.cwd == project_a.resolve()
289326
assert explicit_cwd.source == "x-headroom-cwd"
290327

@@ -465,6 +502,49 @@ async def _record_request_outcome(self, _outcome):
465502
assert unresolved_handler.observed == []
466503

467504

505+
@pytest.mark.asyncio
506+
async def test_http_project_resolution_does_not_block_event_loop(monkeypatch) -> None:
507+
resolver_entered = threading.Event()
508+
resolver_release = threading.Event()
509+
510+
def blocked_resolve(self, **kwargs):
511+
resolver_entered.set()
512+
assert resolver_release.wait(timeout=3)
513+
raise RuntimeError("blocked resolver failed")
514+
515+
monkeypatch.setattr(CodexProjectContextResolver, "resolve", blocked_resolve)
516+
monkeypatch.setattr("headroom.tokenizers.get_tokenizer", lambda model: _DummyTokenizer())
517+
request = _build_request(
518+
{"model": "gpt-5.4", "input": "hello"},
519+
{"Authorization": "Bearer test", "X-Client": "codex"},
520+
)
521+
handler = _HTTPHandler()
522+
progressed_before_release = False
523+
524+
async def unrelated_work() -> None:
525+
nonlocal progressed_before_release
526+
while not resolver_entered.is_set():
527+
await asyncio.sleep(0)
528+
progressed_before_release = not resolver_release.is_set()
529+
resolver_release.set()
530+
531+
release_timer = threading.Timer(2, resolver_release.set)
532+
release_timer.start()
533+
try:
534+
response, _ = await asyncio.gather(
535+
handler.handle_openai_responses(request),
536+
unrelated_work(),
537+
)
538+
finally:
539+
resolver_release.set()
540+
release_timer.cancel()
541+
release_timer.join()
542+
543+
assert resolver_entered.is_set()
544+
assert progressed_before_release
545+
assert response.status_code == 200
546+
547+
468548
@pytest.mark.asyncio
469549
async def test_ws_mismatch_skips_project_memory_but_forwards_main_traffic(
470550
monkeypatch, tmp_path: Path
@@ -521,7 +601,8 @@ def compute_memory_tool_definitions(self, _provider):
521601
[
522602
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
523603
json.dumps({"type": "response.completed", "response": {"id": "r_1"}}),
524-
]
604+
],
605+
hold_after_events=True,
525606
)
526607
websocket = _FakeWebSocket(frames=[first, mismatch])
527608
handler = _WSHandler()

tests/test_codex_ws_per_frame_memory.py

Lines changed: 10 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -147,7 +147,8 @@ async def test_memory_lookup_runs_for_each_issue_artifact_frame_and_preserves_no
147147
[
148148
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
149149
json.dumps({"type": "response.completed", "response": {"id": "r_1"}}),
150-
]
150+
],
151+
hold_after_events=True,
151152
)
152153
first_turn, later_turn = _issue_2059_turns()
153154
first_input, later_input = _issue_2059_inputs()
@@ -183,7 +184,8 @@ async def test_memory_lookup_skips_input_bearing_non_create_first_frame():
183184
[
184185
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
185186
json.dumps({"type": "response.completed", "response": {"id": "r_1"}}),
186-
]
187+
],
188+
hold_after_events=True,
187189
)
188190
_first_input, later_input = _issue_2059_inputs()
189191
cancel_frame = json.dumps(
@@ -214,7 +216,8 @@ async def test_memory_lookup_skips_bypassed_frames():
214216
[
215217
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
216218
json.dumps({"type": "response.completed", "response": {"id": "r_1"}}),
217-
]
219+
],
220+
hold_after_events=True,
218221
)
219222
first, later = _issue_2059_turns()
220223
client_ws = _FakeWebSocket(
@@ -238,7 +241,8 @@ async def test_memory_lookup_keeps_legacy_direct_first_frame():
238241
[
239242
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
240243
json.dumps({"type": "response.completed", "response": {"id": "r_1"}}),
241-
]
244+
],
245+
hold_after_events=True,
242246
)
243247
first_input, later_input = _issue_2059_inputs()
244248
first = _direct_turn(first_input)
@@ -265,7 +269,8 @@ async def test_memory_lookup_skips_disabled_memory(monkeypatch):
265269
[
266270
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
267271
json.dumps({"type": "response.completed", "response": {"id": "r_1"}}),
268-
]
272+
],
273+
hold_after_events=True,
269274
)
270275
first, later = _issue_2059_turns()
271276
client_ws = _FakeWebSocket(frames=[first, later])

0 commit comments

Comments
 (0)