40 lines
1.1 KiB
Python
40 lines
1.1 KiB
Python
"""Small in-memory gate used by legacy model-concurrency tests."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
|
|
|
|
class ModelConcurrencyGate:
|
|
"""Track per-model active counts with atomic try-acquire semantics."""
|
|
|
|
def __init__(self) -> None:
|
|
self._lock = asyncio.Lock()
|
|
self._counts: dict[str, int] = {}
|
|
|
|
async def try_acquire(self, model_name: str, limit: int) -> bool:
|
|
key = str(model_name).strip()
|
|
if not key:
|
|
return True
|
|
async with self._lock:
|
|
current = self._counts.get(key, 0)
|
|
if current >= max(1, int(limit)):
|
|
return False
|
|
self._counts[key] = current + 1
|
|
return True
|
|
|
|
async def release(self, model_name: str) -> None:
|
|
key = str(model_name).strip()
|
|
if not key:
|
|
return
|
|
async with self._lock:
|
|
current = self._counts.get(key, 0)
|
|
if current <= 1:
|
|
self._counts.pop(key, None)
|
|
else:
|
|
self._counts[key] = current - 1
|
|
|
|
async def snapshot(self) -> dict[str, int]:
|
|
async with self._lock:
|
|
return dict(self._counts)
|