22
33from __future__ import annotations
44
5+ import asyncio
56import json
67import logging
78import os
89import sqlite3
910import time
1011from collections .abc import Mapping
1112from dataclasses import dataclass
13+ from functools import lru_cache
1214from pathlib import Path
1315from typing import Any , Literal , cast
1416
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" ]
0 commit comments