diff --git a/src/ai/backend/appproxy/coordinator/types.py b/src/ai/backend/appproxy/coordinator/types.py index d85f1387497..c5612df4a53 100644 --- a/src/ai/backend/appproxy/coordinator/types.py +++ b/src/ai/backend/appproxy/coordinator/types.py @@ -3,6 +3,7 @@ import asyncio import itertools import logging +import weakref from collections import defaultdict from collections.abc import AsyncIterator, Callable, Sequence from contextlib import AbstractAsyncContextManager, AsyncExitStack @@ -105,19 +106,19 @@ class CircuitManager: event_producer: EventProducer traefik_etcd: TraefikEtcd | None local_config: ServerConfig - _circuit_locks: dict[UUID, asyncio.Lock] = field(default_factory=dict) - - def _get_lock(self, circuit_id: UUID) -> asyncio.Lock: - if circuit_id not in self._circuit_locks: - self._circuit_locks[circuit_id] = asyncio.Lock() - return self._circuit_locks[circuit_id] - - def _release_circuit_lock(self, circuit_id: UUID) -> None: - self._circuit_locks.pop(circuit_id, None) + # Weak values: an entry vanishes with its last user, so no explicit release. + _circuit_locks: weakref.WeakValueDictionary[UUID, asyncio.Lock] = field( + default_factory=weakref.WeakValueDictionary + ) @actxmgr async def circuit_lock(self, circuit_id: UUID) -> AsyncIterator[None]: - async with self._get_lock(circuit_id): + # No await between lookup and store — callers converge on one lock. + lock = self._circuit_locks.get(circuit_id) + if lock is None: + lock = asyncio.Lock() + self._circuit_locks[circuit_id] = lock + async with lock: yield async def initialize_circuits(self, circuits: Sequence[Circuit]) -> None: @@ -264,8 +265,6 @@ async def unload_circuits(self, circuits: Sequence[Circuit]) -> None: await self.unload_legacy_circuit(circuit) except Exception: log.exception("Failed to unload circuit {}", circuit.id) - finally: - self._release_circuit_lock(circuit.id) async def unload_traefik_circuit(self, circuit: Circuit) -> None: log.debug("unload_traefik_circuit(): start") diff --git a/tests/unit/appproxy/coordinator/test_circuit_locking.py b/tests/unit/appproxy/coordinator/test_circuit_locking.py index 6ef046e158b..96c759b5483 100644 --- a/tests/unit/appproxy/coordinator/test_circuit_locking.py +++ b/tests/unit/appproxy/coordinator/test_circuit_locking.py @@ -1,13 +1,11 @@ from __future__ import annotations import asyncio -from collections.abc import AsyncIterator -from contextlib import asynccontextmanager from dataclasses import dataclass, field from types import SimpleNamespace from typing import Any, cast from unittest.mock import AsyncMock, MagicMock -from uuid import UUID, uuid4 +from uuid import uuid4 import pytest @@ -15,54 +13,6 @@ from ai.backend.appproxy.coordinator.types import CircuitManager, CircuitRouteUpdateItem -class _ReadonlySessionContext: - def __init__(self, order: list[str]) -> None: - self._order = order - - async def __aenter__(self) -> object: - self._order.append("db_enter") - return object() - - async def __aexit__( - self, - exc_type: type[BaseException] | None, - exc: BaseException | None, - tb: Any, - ) -> None: - self._order.append("db_exit") - - -class _FakeDB: - def __init__(self, order: list[str]) -> None: - self._order = order - - def begin_readonly_session(self) -> _ReadonlySessionContext: - return _ReadonlySessionContext(self._order) - - -class _FakeCircuitManager: - def __init__(self, order: list[str]) -> None: - self._order = order - - @asynccontextmanager - async def circuit_lock(self, _circuit_id: UUID) -> AsyncIterator[None]: - self._order.append("lock_enter") - try: - yield - finally: - self._order.append("lock_exit") - - def release_circuit_lock(self, _circuit_id: UUID) -> None: - self._order.append("release_lock") - - async def _update_circuit_routes_unlocked( - self, - _circuit: object, - _old_routes: list[object], - ) -> None: - self._order.append("update") - - @dataclass class UpdateControl: """Control handle for a two-invocation fake update side effect.""" @@ -187,27 +137,44 @@ async def test_queued_updates_before_unload_are_serialized( # Lock entry is cleaned up after unload assert circuit.id not in circuit_manager._circuit_locks - @pytest.fixture - def patch_unload( + async def test_circuit_lock_entry_lives_only_while_used( self, circuit_manager: CircuitManager, - monkeypatch: pytest.MonkeyPatch, + circuit: Circuit, ) -> None: - monkeypatch.setattr(circuit_manager, "unload_traefik_circuit", AsyncMock(return_value=None)) + async with circuit_manager.circuit_lock(circuit.id): + assert circuit.id in circuit_manager._circuit_locks + assert circuit.id not in circuit_manager._circuit_locks - async def test_unload_removes_circuit_lock( + async def test_waiters_keep_lock_entry_alive( self, circuit_manager: CircuitManager, circuit: Circuit, - patch_unload: None, ) -> None: - # Populate the lock entry - async with circuit_manager.circuit_lock(circuit.id): - pass - assert circuit.id in circuit_manager._circuit_locks + entered = asyncio.Event() + release = asyncio.Event() - # Act - await circuit_manager.unload_circuits([circuit]) + async def _holder() -> None: + async with circuit_manager.circuit_lock(circuit.id): + entered.set() + await release.wait() + + async def _waiter() -> None: + async with circuit_manager.circuit_lock(circuit.id): + pass + + holder_task = asyncio.create_task(_holder()) + await entered.wait() + held_lock = circuit_manager._circuit_locks.get(circuit.id) + assert held_lock is not None + + waiter_task = asyncio.create_task(_waiter()) + await asyncio.sleep(0) + + assert circuit_manager._circuit_locks.get(circuit.id) is held_lock - # Assert - lock entry should be cleaned up + release.set() + await asyncio.gather(holder_task, waiter_task) + # Drop the test's own strong ref so the entry can be collected. + del held_lock assert circuit.id not in circuit_manager._circuit_locks