Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 11 additions & 12 deletions src/ai/backend/appproxy/coordinator/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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")
Expand Down
95 changes: 31 additions & 64 deletions tests/unit/appproxy/coordinator/test_circuit_locking.py
Original file line number Diff line number Diff line change
@@ -1,68 +1,18 @@
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

from ai.backend.appproxy.coordinator.models import Circuit
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."""
Expand Down Expand Up @@ -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
Loading