Edit on GitHub

agent_search_gateway.concurrency

Provider quotas, singleflight, and keyed serialization primitives.

  1"""Provider quotas, singleflight, and keyed serialization primitives."""
  2
  3import asyncio
  4import logging
  5import time
  6from collections.abc import Awaitable, Callable, Hashable, Mapping
  7from dataclasses import dataclass
  8from typing import Generic, TypeVar
  9
 10from .observability import log_event
 11
 12T = TypeVar("T")
 13K = TypeVar("K", bound=Hashable)
 14
 15
 16class CapacityLease:
 17    def __init__(self, gate: "CapacityGate", *, acquired: bool = False) -> None:
 18        self._gate = gate
 19        self._acquired = acquired
 20        self._released = False
 21
 22    async def __aenter__(self) -> "CapacityLease":
 23        if not self._acquired:
 24            await self._gate._acquire()
 25            self._acquired = True
 26        return self
 27
 28    async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
 29        await self.release()
 30
 31    async def release(self) -> None:
 32        if not self._acquired or self._released:
 33            return
 34        self._released = True
 35        await self._gate._release()
 36
 37
 38class CapacityGate:
 39    def __init__(
 40        self,
 41        limit: int,
 42        *,
 43        provider: str | None = None,
 44        quota_kind: str | None = None,
 45        logger: logging.Logger | None = None,
 46        monotonic: Callable[[], float] = time.monotonic,
 47    ) -> None:
 48        if limit <= 0:
 49            raise ValueError("capacity limit must be positive")
 50        self.limit = limit
 51        self.in_use = 0
 52        self.max_observed_in_use = 0
 53        self._provider = provider
 54        self._quota_kind = quota_kind
 55        self._logger = logger
 56        self._monotonic = monotonic
 57        self._condition = asyncio.Condition()
 58
 59    def lease(self) -> CapacityLease:
 60        return CapacityLease(self)
 61
 62    async def try_lease(self) -> CapacityLease | None:
 63        async with self._condition:
 64            if self.in_use >= self.limit:
 65                self._log("quota_waiting")
 66                return None
 67            self._claim(waited_ms=0)
 68            return CapacityLease(self, acquired=True)
 69
 70    async def wait_until_available(self) -> None:
 71        async with self._condition:
 72            if self.in_use >= self.limit:
 73                self._log("quota_waiting")
 74            await self._condition.wait_for(lambda: self.in_use < self.limit)
 75
 76    async def _acquire(self) -> None:
 77        async with self._condition:
 78            started = self._monotonic()
 79            waiting = self.in_use >= self.limit
 80            if waiting:
 81                self._log("quota_waiting")
 82            await self._condition.wait_for(lambda: self.in_use < self.limit)
 83            waited_ms = max(0, int((self._monotonic() - started) * 1000)) if waiting else 0
 84            self._claim(waited_ms=waited_ms)
 85
 86    def _claim(self, *, waited_ms: int) -> None:
 87        self.in_use += 1
 88        self.max_observed_in_use = max(self.max_observed_in_use, self.in_use)
 89        self._log("quota_acquired", waited_ms=waited_ms)
 90
 91    async def _release(self) -> None:
 92        async with self._condition:
 93            if self.in_use <= 0:
 94                raise RuntimeError("capacity lease released more than once")
 95            self.in_use -= 1
 96            self._log("quota_released")
 97            self._condition.notify_all()
 98
 99    def _log(self, event: str, *, waited_ms: int | None = None) -> None:
100        if self._logger is None or self._provider is None or self._quota_kind is None:
101            return
102        if waited_ms is None:
103            log_event(
104                self._logger,
105                logging.DEBUG,
106                event,
107                provider=self._provider,
108                quota_kind=self._quota_kind,
109                in_use=self.in_use,
110                limit=self.limit,
111            )
112            return
113        log_event(
114            self._logger,
115            logging.DEBUG,
116            event,
117            provider=self._provider,
118            quota_kind=self._quota_kind,
119            in_use=self.in_use,
120            limit=self.limit,
121            waited_ms=waited_ms,
122        )
123
124
125class ProviderQuotaManager:
126    def __init__(
127        self,
128        *,
129        web_limits: Mapping[str, int],
130        llm_limits: Mapping[str, int],
131        academic_limits: Mapping[str, int] | None = None,
132        logger: logging.Logger | None = None,
133        monotonic: Callable[[], float] = time.monotonic,
134    ) -> None:
135        event_logger = logger or logging.getLogger(__name__)
136        self._web = {
137            name: CapacityGate(
138                limit,
139                provider=name,
140                quota_kind="web",
141                logger=event_logger,
142                monotonic=monotonic,
143            )
144            for name, limit in web_limits.items()
145        }
146        self._llm = {
147            name: CapacityGate(
148                limit,
149                provider=name,
150                quota_kind="llm",
151                logger=event_logger,
152                monotonic=monotonic,
153            )
154            for name, limit in llm_limits.items()
155        }
156        self._academic = {
157            name: CapacityGate(
158                limit,
159                provider=name,
160                quota_kind="academic",
161                logger=event_logger,
162                monotonic=monotonic,
163            )
164            for name, limit in (academic_limits or {}).items()
165        }
166
167    def get_web(self, name: str) -> CapacityGate:
168        return self._web[name]
169
170    def get_llm(self, name: str) -> CapacityGate:
171        return self._llm[name]
172
173    def get_academic(self, name: str) -> CapacityGate:
174        return self._academic[name]
175
176    async def wait_until_any_web_available(self, candidate_names: tuple[str, ...]) -> None:
177        if not candidate_names:
178            return
179        tasks = [
180            asyncio.create_task(self.get_web(name).wait_until_available())
181            for name in candidate_names
182        ]
183        try:
184            done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
185            for task in done:
186                task.result()
187            for task in pending:
188                task.cancel()
189            if pending:
190                await asyncio.gather(*pending, return_exceptions=True)
191        finally:
192            for task in tasks:
193                if not task.done():
194                    task.cancel()
195
196
197@dataclass(slots=True)
198class _LockEntry:
199    lock: asyncio.Lock
200    references: int = 0
201
202
203class _KeyedLockLease(Generic[K]):
204    def __init__(self, pool: "PerKeyLockPool[K]", key: K) -> None:
205        self._pool = pool
206        self._key = key
207        self._entry: _LockEntry | None = None
208
209    async def __aenter__(self) -> None:
210        self._entry = await self._pool._reserve(self._key)
211        try:
212            await self._entry.lock.acquire()
213        except BaseException:
214            await self._pool._unreserve(self._key, self._entry)
215            self._entry = None
216            raise
217
218    async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
219        entry = self._entry
220        if entry is None:
221            return
222        entry.lock.release()
223        await self._pool._unreserve(self._key, entry)
224        self._entry = None
225
226
227class PerKeyLockPool(Generic[K]):
228    def __init__(self) -> None:
229        self._entries: dict[K, _LockEntry] = {}
230        self._guard = asyncio.Lock()
231
232    def acquire(self, key: K) -> _KeyedLockLease[K]:
233        return _KeyedLockLease(self, key)
234
235    async def _reserve(self, key: K) -> _LockEntry:
236        async with self._guard:
237            entry = self._entries.get(key)
238            if entry is None:
239                entry = _LockEntry(asyncio.Lock())
240                self._entries[key] = entry
241            entry.references += 1
242            return entry
243
244    async def _unreserve(self, key: K, entry: _LockEntry) -> None:
245        async with self._guard:
246            entry.references -= 1
247            if entry.references == 0:
248                del self._entries[key]
249
250
251class SingleflightGroup(Generic[K, T]):
252    def __init__(self) -> None:
253        self._guard = asyncio.Lock()
254        self._inflight: dict[K, asyncio.Future[T]] = {}
255
256    async def do(
257        self,
258        key: K,
259        factory: Callable[[], Awaitable[T]],
260        *,
261        on_leader: Callable[[], None] | None = None,
262        on_follower: Callable[[], None] | None = None,
263    ) -> T:
264        async with self._guard:
265            future = self._inflight.get(key)
266            if future is None:
267                future = asyncio.get_running_loop().create_future()
268                self._inflight[key] = future
269                leader = True
270            else:
271                leader = False
272
273        if not leader:
274            self._run_role_callback(on_follower)
275            return await asyncio.shield(future)
276        self._run_role_callback(on_leader)
277
278        try:
279            result = await factory()
280        except BaseException as exc:
281            if isinstance(exc, asyncio.CancelledError):
282                future.cancel()
283            else:
284                future.set_exception(exc)
285                future.exception()
286            raise
287        else:
288            future.set_result(result)
289            return result
290        finally:
291            await self._cleanup(key, future)
292
293    @staticmethod
294    def _run_role_callback(callback: Callable[[], None] | None) -> None:
295        if callback is None:
296            return
297        try:
298            callback()
299        except Exception:
300            return
301
302    async def _cleanup(self, key: K, future: asyncio.Future[T]) -> None:
303        async with self._guard:
304            if self._inflight.get(key) is future:
305                del self._inflight[key]
class CapacityLease:
17class CapacityLease:
18    def __init__(self, gate: "CapacityGate", *, acquired: bool = False) -> None:
19        self._gate = gate
20        self._acquired = acquired
21        self._released = False
22
23    async def __aenter__(self) -> "CapacityLease":
24        if not self._acquired:
25            await self._gate._acquire()
26            self._acquired = True
27        return self
28
29    async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
30        await self.release()
31
32    async def release(self) -> None:
33        if not self._acquired or self._released:
34            return
35        self._released = True
36        await self._gate._release()
CapacityLease( gate: CapacityGate, *, acquired: bool = False)
18    def __init__(self, gate: "CapacityGate", *, acquired: bool = False) -> None:
19        self._gate = gate
20        self._acquired = acquired
21        self._released = False
async def release(self) -> None:
32    async def release(self) -> None:
33        if not self._acquired or self._released:
34            return
35        self._released = True
36        await self._gate._release()
class CapacityGate:
 39class CapacityGate:
 40    def __init__(
 41        self,
 42        limit: int,
 43        *,
 44        provider: str | None = None,
 45        quota_kind: str | None = None,
 46        logger: logging.Logger | None = None,
 47        monotonic: Callable[[], float] = time.monotonic,
 48    ) -> None:
 49        if limit <= 0:
 50            raise ValueError("capacity limit must be positive")
 51        self.limit = limit
 52        self.in_use = 0
 53        self.max_observed_in_use = 0
 54        self._provider = provider
 55        self._quota_kind = quota_kind
 56        self._logger = logger
 57        self._monotonic = monotonic
 58        self._condition = asyncio.Condition()
 59
 60    def lease(self) -> CapacityLease:
 61        return CapacityLease(self)
 62
 63    async def try_lease(self) -> CapacityLease | None:
 64        async with self._condition:
 65            if self.in_use >= self.limit:
 66                self._log("quota_waiting")
 67                return None
 68            self._claim(waited_ms=0)
 69            return CapacityLease(self, acquired=True)
 70
 71    async def wait_until_available(self) -> None:
 72        async with self._condition:
 73            if self.in_use >= self.limit:
 74                self._log("quota_waiting")
 75            await self._condition.wait_for(lambda: self.in_use < self.limit)
 76
 77    async def _acquire(self) -> None:
 78        async with self._condition:
 79            started = self._monotonic()
 80            waiting = self.in_use >= self.limit
 81            if waiting:
 82                self._log("quota_waiting")
 83            await self._condition.wait_for(lambda: self.in_use < self.limit)
 84            waited_ms = max(0, int((self._monotonic() - started) * 1000)) if waiting else 0
 85            self._claim(waited_ms=waited_ms)
 86
 87    def _claim(self, *, waited_ms: int) -> None:
 88        self.in_use += 1
 89        self.max_observed_in_use = max(self.max_observed_in_use, self.in_use)
 90        self._log("quota_acquired", waited_ms=waited_ms)
 91
 92    async def _release(self) -> None:
 93        async with self._condition:
 94            if self.in_use <= 0:
 95                raise RuntimeError("capacity lease released more than once")
 96            self.in_use -= 1
 97            self._log("quota_released")
 98            self._condition.notify_all()
 99
100    def _log(self, event: str, *, waited_ms: int | None = None) -> None:
101        if self._logger is None or self._provider is None or self._quota_kind is None:
102            return
103        if waited_ms is None:
104            log_event(
105                self._logger,
106                logging.DEBUG,
107                event,
108                provider=self._provider,
109                quota_kind=self._quota_kind,
110                in_use=self.in_use,
111                limit=self.limit,
112            )
113            return
114        log_event(
115            self._logger,
116            logging.DEBUG,
117            event,
118            provider=self._provider,
119            quota_kind=self._quota_kind,
120            in_use=self.in_use,
121            limit=self.limit,
122            waited_ms=waited_ms,
123        )
CapacityGate( limit: int, *, provider: str | None = None, quota_kind: str | None = None, logger: logging.Logger | None = None, monotonic: Callable[[], float] = <built-in function monotonic>)
40    def __init__(
41        self,
42        limit: int,
43        *,
44        provider: str | None = None,
45        quota_kind: str | None = None,
46        logger: logging.Logger | None = None,
47        monotonic: Callable[[], float] = time.monotonic,
48    ) -> None:
49        if limit <= 0:
50            raise ValueError("capacity limit must be positive")
51        self.limit = limit
52        self.in_use = 0
53        self.max_observed_in_use = 0
54        self._provider = provider
55        self._quota_kind = quota_kind
56        self._logger = logger
57        self._monotonic = monotonic
58        self._condition = asyncio.Condition()
limit
in_use
max_observed_in_use
def lease(self) -> CapacityLease:
60    def lease(self) -> CapacityLease:
61        return CapacityLease(self)
async def try_lease(self) -> CapacityLease | None:
63    async def try_lease(self) -> CapacityLease | None:
64        async with self._condition:
65            if self.in_use >= self.limit:
66                self._log("quota_waiting")
67                return None
68            self._claim(waited_ms=0)
69            return CapacityLease(self, acquired=True)
async def wait_until_available(self) -> None:
71    async def wait_until_available(self) -> None:
72        async with self._condition:
73            if self.in_use >= self.limit:
74                self._log("quota_waiting")
75            await self._condition.wait_for(lambda: self.in_use < self.limit)
class ProviderQuotaManager:
126class ProviderQuotaManager:
127    def __init__(
128        self,
129        *,
130        web_limits: Mapping[str, int],
131        llm_limits: Mapping[str, int],
132        academic_limits: Mapping[str, int] | None = None,
133        logger: logging.Logger | None = None,
134        monotonic: Callable[[], float] = time.monotonic,
135    ) -> None:
136        event_logger = logger or logging.getLogger(__name__)
137        self._web = {
138            name: CapacityGate(
139                limit,
140                provider=name,
141                quota_kind="web",
142                logger=event_logger,
143                monotonic=monotonic,
144            )
145            for name, limit in web_limits.items()
146        }
147        self._llm = {
148            name: CapacityGate(
149                limit,
150                provider=name,
151                quota_kind="llm",
152                logger=event_logger,
153                monotonic=monotonic,
154            )
155            for name, limit in llm_limits.items()
156        }
157        self._academic = {
158            name: CapacityGate(
159                limit,
160                provider=name,
161                quota_kind="academic",
162                logger=event_logger,
163                monotonic=monotonic,
164            )
165            for name, limit in (academic_limits or {}).items()
166        }
167
168    def get_web(self, name: str) -> CapacityGate:
169        return self._web[name]
170
171    def get_llm(self, name: str) -> CapacityGate:
172        return self._llm[name]
173
174    def get_academic(self, name: str) -> CapacityGate:
175        return self._academic[name]
176
177    async def wait_until_any_web_available(self, candidate_names: tuple[str, ...]) -> None:
178        if not candidate_names:
179            return
180        tasks = [
181            asyncio.create_task(self.get_web(name).wait_until_available())
182            for name in candidate_names
183        ]
184        try:
185            done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
186            for task in done:
187                task.result()
188            for task in pending:
189                task.cancel()
190            if pending:
191                await asyncio.gather(*pending, return_exceptions=True)
192        finally:
193            for task in tasks:
194                if not task.done():
195                    task.cancel()
ProviderQuotaManager( *, web_limits: Mapping[str, int], llm_limits: Mapping[str, int], academic_limits: Mapping[str, int] | None = None, logger: logging.Logger | None = None, monotonic: Callable[[], float] = <built-in function monotonic>)
127    def __init__(
128        self,
129        *,
130        web_limits: Mapping[str, int],
131        llm_limits: Mapping[str, int],
132        academic_limits: Mapping[str, int] | None = None,
133        logger: logging.Logger | None = None,
134        monotonic: Callable[[], float] = time.monotonic,
135    ) -> None:
136        event_logger = logger or logging.getLogger(__name__)
137        self._web = {
138            name: CapacityGate(
139                limit,
140                provider=name,
141                quota_kind="web",
142                logger=event_logger,
143                monotonic=monotonic,
144            )
145            for name, limit in web_limits.items()
146        }
147        self._llm = {
148            name: CapacityGate(
149                limit,
150                provider=name,
151                quota_kind="llm",
152                logger=event_logger,
153                monotonic=monotonic,
154            )
155            for name, limit in llm_limits.items()
156        }
157        self._academic = {
158            name: CapacityGate(
159                limit,
160                provider=name,
161                quota_kind="academic",
162                logger=event_logger,
163                monotonic=monotonic,
164            )
165            for name, limit in (academic_limits or {}).items()
166        }
def get_web(self, name: str) -> CapacityGate:
168    def get_web(self, name: str) -> CapacityGate:
169        return self._web[name]
def get_llm(self, name: str) -> CapacityGate:
171    def get_llm(self, name: str) -> CapacityGate:
172        return self._llm[name]
def get_academic(self, name: str) -> CapacityGate:
174    def get_academic(self, name: str) -> CapacityGate:
175        return self._academic[name]
async def wait_until_any_web_available(self, candidate_names: tuple[str, ...]) -> None:
177    async def wait_until_any_web_available(self, candidate_names: tuple[str, ...]) -> None:
178        if not candidate_names:
179            return
180        tasks = [
181            asyncio.create_task(self.get_web(name).wait_until_available())
182            for name in candidate_names
183        ]
184        try:
185            done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
186            for task in done:
187                task.result()
188            for task in pending:
189                task.cancel()
190            if pending:
191                await asyncio.gather(*pending, return_exceptions=True)
192        finally:
193            for task in tasks:
194                if not task.done():
195                    task.cancel()
class PerKeyLockPool(typing.Generic[~K]):
228class PerKeyLockPool(Generic[K]):
229    def __init__(self) -> None:
230        self._entries: dict[K, _LockEntry] = {}
231        self._guard = asyncio.Lock()
232
233    def acquire(self, key: K) -> _KeyedLockLease[K]:
234        return _KeyedLockLease(self, key)
235
236    async def _reserve(self, key: K) -> _LockEntry:
237        async with self._guard:
238            entry = self._entries.get(key)
239            if entry is None:
240                entry = _LockEntry(asyncio.Lock())
241                self._entries[key] = entry
242            entry.references += 1
243            return entry
244
245    async def _unreserve(self, key: K, entry: _LockEntry) -> None:
246        async with self._guard:
247            entry.references -= 1
248            if entry.references == 0:
249                del self._entries[key]

Abstract base class for generic types.

On Python 3.12 and newer, generic classes implicitly inherit from Generic when they declare a parameter list after the class's name::

class Mapping[KT, VT]:
    def __getitem__(self, key: KT) -> VT:
        ...
    # Etc.

On older versions of Python, however, generic classes have to explicitly inherit from Generic.

After a class has been declared to be generic, it can then be used as follows::

def lookup_name[KT, VT](mapping: Mapping[KT, VT], key: KT, default: VT) -> VT:
    try:
        return mapping[key]
    except KeyError:
        return default
def acquire(self, key: ~K) -> agent_search_gateway.concurrency._KeyedLockLease[~K]:
233    def acquire(self, key: K) -> _KeyedLockLease[K]:
234        return _KeyedLockLease(self, key)
class SingleflightGroup(typing.Generic[~K, ~T]):
252class SingleflightGroup(Generic[K, T]):
253    def __init__(self) -> None:
254        self._guard = asyncio.Lock()
255        self._inflight: dict[K, asyncio.Future[T]] = {}
256
257    async def do(
258        self,
259        key: K,
260        factory: Callable[[], Awaitable[T]],
261        *,
262        on_leader: Callable[[], None] | None = None,
263        on_follower: Callable[[], None] | None = None,
264    ) -> T:
265        async with self._guard:
266            future = self._inflight.get(key)
267            if future is None:
268                future = asyncio.get_running_loop().create_future()
269                self._inflight[key] = future
270                leader = True
271            else:
272                leader = False
273
274        if not leader:
275            self._run_role_callback(on_follower)
276            return await asyncio.shield(future)
277        self._run_role_callback(on_leader)
278
279        try:
280            result = await factory()
281        except BaseException as exc:
282            if isinstance(exc, asyncio.CancelledError):
283                future.cancel()
284            else:
285                future.set_exception(exc)
286                future.exception()
287            raise
288        else:
289            future.set_result(result)
290            return result
291        finally:
292            await self._cleanup(key, future)
293
294    @staticmethod
295    def _run_role_callback(callback: Callable[[], None] | None) -> None:
296        if callback is None:
297            return
298        try:
299            callback()
300        except Exception:
301            return
302
303    async def _cleanup(self, key: K, future: asyncio.Future[T]) -> None:
304        async with self._guard:
305            if self._inflight.get(key) is future:
306                del self._inflight[key]

Abstract base class for generic types.

On Python 3.12 and newer, generic classes implicitly inherit from Generic when they declare a parameter list after the class's name::

class Mapping[KT, VT]:
    def __getitem__(self, key: KT) -> VT:
        ...
    # Etc.

On older versions of Python, however, generic classes have to explicitly inherit from Generic.

After a class has been declared to be generic, it can then be used as follows::

def lookup_name[KT, VT](mapping: Mapping[KT, VT], key: KT, default: VT) -> VT:
    try:
        return mapping[key]
    except KeyError:
        return default
async def do( self, key: ~K, factory: Callable[[], Awaitable[~T]], *, on_leader: Callable[[], None] | None = None, on_follower: Callable[[], None] | None = None) -> ~T:
257    async def do(
258        self,
259        key: K,
260        factory: Callable[[], Awaitable[T]],
261        *,
262        on_leader: Callable[[], None] | None = None,
263        on_follower: Callable[[], None] | None = None,
264    ) -> T:
265        async with self._guard:
266            future = self._inflight.get(key)
267            if future is None:
268                future = asyncio.get_running_loop().create_future()
269                self._inflight[key] = future
270                leader = True
271            else:
272                leader = False
273
274        if not leader:
275            self._run_role_callback(on_follower)
276            return await asyncio.shield(future)
277        self._run_role_callback(on_leader)
278
279        try:
280            result = await factory()
281        except BaseException as exc:
282            if isinstance(exc, asyncio.CancelledError):
283                future.cancel()
284            else:
285                future.set_exception(exc)
286                future.exception()
287            raise
288        else:
289            future.set_result(result)
290            return result
291        finally:
292            await self._cleanup(key, future)