Edit on GitHub

agent_search_gateway.request_ids

Request identifiers used for workflow correlation and search result names.

 1"""Request identifiers used for workflow correlation and search result names."""
 2
 3import re
 4import secrets
 5from collections.abc import Callable, Iterator
 6from contextlib import contextmanager
 7from contextvars import ContextVar
 8from pathlib import Path
 9from typing import Literal
10
11RequestIdFactory = Callable[[], str]
12ResultKind = Literal["keyword", "llm", "paper"]
13
14_RESULT_KINDS: tuple[ResultKind, ...] = ("keyword", "llm", "paper")
15_REQUEST_ID_PATTERN = re.compile(r"[0-9a-f]{8}")
16_REQUEST_ID: ContextVar[str | None] = ContextVar("agent_search_gateway_request_id", default=None)
17
18
19def validate_request_id(value: str) -> str:
20    if _REQUEST_ID_PATTERN.fullmatch(value) is None:
21        raise ValueError("request ID must be exactly 8 lowercase hexadecimal characters")
22    return value
23
24
25def generate_request_id() -> str:
26    return secrets.token_hex(4)
27
28
29def current_request_id() -> str | None:
30    return _REQUEST_ID.get()
31
32
33@contextmanager
34def bind_request_id(request_id: str) -> Iterator[None]:
35    token = _REQUEST_ID.set(validate_request_id(request_id))
36    try:
37        yield
38    finally:
39        _REQUEST_ID.reset(token)
40
41
42def result_filename(kind: ResultKind, request_id: str) -> str:
43    if kind not in _RESULT_KINDS:
44        raise ValueError(f"invalid result kind: {kind}")
45    return f"{kind}-{validate_request_id(request_id)}.jsonl"
46
47
48class RequestIdRegistry:
49    def __init__(
50        self,
51        results_dir: Path,
52        *,
53        factory: RequestIdFactory = generate_request_id,
54        max_attempts: int = 256,
55    ) -> None:
56        if max_attempts <= 0:
57            raise ValueError("max_attempts must be positive")
58        self._results_dir = results_dir
59        self._factory = factory
60        self._max_attempts = max_attempts
61        self._active: set[str] = set()
62
63    @contextmanager
64    def reserve(self, *, may_write_search_result: bool) -> Iterator[str]:
65        request_id = self._select_available(may_write_search_result=may_write_search_result)
66        self._active.add(request_id)
67        try:
68            yield request_id
69        finally:
70            self._active.remove(request_id)
71
72    def _select_available(self, *, may_write_search_result: bool) -> str:
73        for _ in range(self._max_attempts):
74            candidate = self._next_candidate()
75            if candidate in self._active:
76                continue
77            if may_write_search_result and self._has_result_collision(candidate):
78                continue
79            return candidate
80        raise RuntimeError("unable to reserve an available request ID")
81
82    def _next_candidate(self) -> str:
83        try:
84            candidate = self._factory()
85        except Exception as exc:
86            raise RuntimeError("request ID factory failed") from exc
87        try:
88            return validate_request_id(candidate)
89        except (TypeError, ValueError) as exc:
90            raise RuntimeError("invalid request ID returned by factory") from exc
91
92    def _has_result_collision(self, request_id: str) -> bool:
93        return any(
94            (self._results_dir / result_filename(kind, request_id)).exists()
95            for kind in _RESULT_KINDS
96        )
RequestIdFactory = collections.abc.Callable[[], str]
ResultKind = typing.Literal['keyword', 'llm', 'paper']
def validate_request_id(value: str) -> str:
20def validate_request_id(value: str) -> str:
21    if _REQUEST_ID_PATTERN.fullmatch(value) is None:
22        raise ValueError("request ID must be exactly 8 lowercase hexadecimal characters")
23    return value
def generate_request_id() -> str:
26def generate_request_id() -> str:
27    return secrets.token_hex(4)
def current_request_id() -> str | None:
30def current_request_id() -> str | None:
31    return _REQUEST_ID.get()
@contextmanager
def bind_request_id(request_id: str) -> Iterator[None]:
34@contextmanager
35def bind_request_id(request_id: str) -> Iterator[None]:
36    token = _REQUEST_ID.set(validate_request_id(request_id))
37    try:
38        yield
39    finally:
40        _REQUEST_ID.reset(token)
def result_filename(kind: Literal['keyword', 'llm', 'paper'], request_id: str) -> str:
43def result_filename(kind: ResultKind, request_id: str) -> str:
44    if kind not in _RESULT_KINDS:
45        raise ValueError(f"invalid result kind: {kind}")
46    return f"{kind}-{validate_request_id(request_id)}.jsonl"
class RequestIdRegistry:
49class RequestIdRegistry:
50    def __init__(
51        self,
52        results_dir: Path,
53        *,
54        factory: RequestIdFactory = generate_request_id,
55        max_attempts: int = 256,
56    ) -> None:
57        if max_attempts <= 0:
58            raise ValueError("max_attempts must be positive")
59        self._results_dir = results_dir
60        self._factory = factory
61        self._max_attempts = max_attempts
62        self._active: set[str] = set()
63
64    @contextmanager
65    def reserve(self, *, may_write_search_result: bool) -> Iterator[str]:
66        request_id = self._select_available(may_write_search_result=may_write_search_result)
67        self._active.add(request_id)
68        try:
69            yield request_id
70        finally:
71            self._active.remove(request_id)
72
73    def _select_available(self, *, may_write_search_result: bool) -> str:
74        for _ in range(self._max_attempts):
75            candidate = self._next_candidate()
76            if candidate in self._active:
77                continue
78            if may_write_search_result and self._has_result_collision(candidate):
79                continue
80            return candidate
81        raise RuntimeError("unable to reserve an available request ID")
82
83    def _next_candidate(self) -> str:
84        try:
85            candidate = self._factory()
86        except Exception as exc:
87            raise RuntimeError("request ID factory failed") from exc
88        try:
89            return validate_request_id(candidate)
90        except (TypeError, ValueError) as exc:
91            raise RuntimeError("invalid request ID returned by factory") from exc
92
93    def _has_result_collision(self, request_id: str) -> bool:
94        return any(
95            (self._results_dir / result_filename(kind, request_id)).exists()
96            for kind in _RESULT_KINDS
97        )
RequestIdRegistry( results_dir: pathlib.Path, *, factory: Callable[[], str] = <function generate_request_id>, max_attempts: int = 256)
50    def __init__(
51        self,
52        results_dir: Path,
53        *,
54        factory: RequestIdFactory = generate_request_id,
55        max_attempts: int = 256,
56    ) -> None:
57        if max_attempts <= 0:
58            raise ValueError("max_attempts must be positive")
59        self._results_dir = results_dir
60        self._factory = factory
61        self._max_attempts = max_attempts
62        self._active: set[str] = set()
@contextmanager
def reserve(self, *, may_write_search_result: bool) -> Iterator[str]:
64    @contextmanager
65    def reserve(self, *, may_write_search_result: bool) -> Iterator[str]:
66        request_id = self._select_available(may_write_search_result=may_write_search_result)
67        self._active.add(request_id)
68        try:
69            yield request_id
70        finally:
71            self._active.remove(request_id)