Add reference counting to ModelInstance for parallel test safety (#16672)

This commit is contained in:
Simo Lin
2026-01-07 08:28:15 -08:00
committed by GitHub
parent b5a94f8a8e
commit 7385834c8d
3 changed files with 100 additions and 6 deletions
+24 -2
View File
@@ -151,6 +151,7 @@ def model_client(request: pytest.FixtureRequest, model_pool: "ModelPool"):
def test_chat(model_client): def test_chat(model_client):
response = model_client.chat.completions.create(...) response = model_client.chat.completions.create(...)
""" """
import openai
from infra import PARAM_MODEL from infra import PARAM_MODEL
marker = request.node.get_closest_marker(PARAM_MODEL) marker = request.node.get_closest_marker(PARAM_MODEL)
@@ -163,10 +164,23 @@ def model_client(request: pytest.FixtureRequest, model_pool: "ModelPool"):
model_id = marker.args[0] model_id = marker.args[0]
try: try:
return model_pool.get_client(model_id) instance = model_pool.get(model_id)
except KeyError: except KeyError:
pytest.skip(f"Model {model_id} not available in model pool") pytest.skip(f"Model {model_id} not available in model pool")
# Acquire reference to prevent eviction during test
instance.acquire()
client = openai.OpenAI(
base_url=f"{instance.base_url}/v1",
api_key="not-used",
)
yield client
# Release reference to allow eviction
instance.release()
@pytest.fixture @pytest.fixture
def model_base_url(request: pytest.FixtureRequest, model_pool: "ModelPool") -> str: def model_base_url(request: pytest.FixtureRequest, model_pool: "ModelPool") -> str:
@@ -189,6 +203,14 @@ def model_base_url(request: pytest.FixtureRequest, model_pool: "ModelPool") -> s
model_id = marker.args[0] model_id = marker.args[0]
try: try:
return model_pool.get_base_url(model_id) instance = model_pool.get(model_id)
except KeyError: except KeyError:
pytest.skip(f"Model {model_id} not available in model pool") pytest.skip(f"Model {model_id} not available in model pool")
# Acquire reference to prevent eviction during test
instance.acquire()
yield instance.base_url
# Release reference to allow eviction
instance.release()
@@ -212,6 +212,11 @@ def _setup_pd_backend(
prefills = existing_prefills + new_prefills prefills = existing_prefills + new_prefills
decodes = existing_decodes + new_decodes decodes = existing_decodes + new_decodes
# Acquire references to prevent eviction during test
all_workers = prefills + decodes
for worker in all_workers:
worker.acquire()
model_path = prefills[0].model_path if prefills else None model_path = prefills[0].model_path if prefills else None
# Launch PD gateway # Launch PD gateway
@@ -244,6 +249,9 @@ def _setup_pd_backend(
finally: finally:
logger.info("Tearing down PD gateway") logger.info("Tearing down PD gateway")
gateway.shutdown() gateway.shutdown()
# Release references to allow eviction
for worker in all_workers:
worker.release()
def _setup_local_backend( def _setup_local_backend(
@@ -260,6 +268,7 @@ def _setup_local_backend(
from infra import Gateway, WorkerIdentity, WorkerType from infra import Gateway, WorkerIdentity, WorkerType
num_workers = workers_config.get("count") or 1 num_workers = workers_config.get("count") or 1
instances: list = [] # Track instances for reference counting
try: try:
if num_workers > 1: if num_workers > 1:
@@ -290,8 +299,13 @@ def _setup_local_backend(
model_path = instances[0].model_path model_path = instances[0].model_path
else: else:
instance = model_pool.get(model_id, connection_mode) instance = model_pool.get(model_id, connection_mode)
instances = [instance]
worker_urls = [instance.worker_url] worker_urls = [instance.worker_url]
model_path = instance.model_path model_path = instance.model_path
# Acquire references to prevent eviction during test
for inst in instances:
inst.acquire()
except RuntimeError as e: except RuntimeError as e:
pytest.fail(str(e)) pytest.fail(str(e))
@@ -324,6 +338,9 @@ def _setup_local_backend(
finally: finally:
logger.info("Tearing down gateway for %s backend", backend_name) logger.info("Tearing down gateway for %s backend", backend_name)
gateway.shutdown() gateway.shutdown()
# Release references to allow eviction
for inst in instances:
inst.release()
def _setup_cloud_backend(backend_name: str): def _setup_cloud_backend(backend_name: str):
@@ -382,6 +399,9 @@ def backend_router(request: pytest.FixtureRequest, model_pool: "ModelPool"):
except RuntimeError as e: except RuntimeError as e:
pytest.fail(str(e)) pytest.fail(str(e))
# Acquire reference to prevent eviction during test
instance.acquire()
gateway = Gateway() gateway = Gateway()
gateway.start( gateway.start(
worker_urls=[instance.worker_url], worker_urls=[instance.worker_url],
@@ -392,3 +412,5 @@ def backend_router(request: pytest.FixtureRequest, model_pool: "ModelPool"):
yield gateway yield gateway
finally: finally:
gateway.shutdown() gateway.shutdown()
# Release reference to allow eviction
instance.release()
+54 -4
View File
@@ -5,8 +5,9 @@ from __future__ import annotations
import logging import logging
import os import os
import subprocess import subprocess
import threading
import time import time
from dataclasses import dataclass from dataclasses import dataclass, field
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
import httpx import httpx
@@ -102,6 +103,10 @@ class ModelInstance:
last_used: float = 0.0 # Timestamp for MRU eviction last_used: float = 0.0 # Timestamp for MRU eviction
_healthy: bool = False # Track if initial health check passed _healthy: bool = False # Track if initial health check passed
# Reference counting for safe parallel test execution
_ref_count: int = 0
_ref_lock: threading.Lock = field(default_factory=threading.Lock)
@property @property
def identity(self) -> WorkerIdentity: def identity(self) -> WorkerIdentity:
"""Get the identity (model_id, mode, worker_type) of this instance.""" """Get the identity (model_id, mode, worker_type) of this instance."""
@@ -111,6 +116,44 @@ class ModelInstance:
worker_type=self.worker_type, worker_type=self.worker_type,
) )
@property
def is_in_use(self) -> bool:
"""Check if this instance has active references (tests using it)."""
with self._ref_lock:
return self._ref_count > 0
def acquire(self) -> None:
"""Acquire a reference to this instance.
Call this before using the instance in a test to prevent eviction.
Must be paired with a release() call when done.
Also updates last_used timestamp atomically with ref count.
"""
with self._ref_lock:
self._ref_count += 1
self.last_used = time.time()
logger.debug(
"Acquired reference to %s (ref_count=%d)", self.key, self._ref_count
)
def release(self) -> None:
"""Release a reference to this instance.
Call this when done using the instance in a test.
"""
with self._ref_lock:
if self._ref_count > 0:
self._ref_count -= 1
logger.debug(
"Released reference to %s (ref_count=%d)",
self.key,
self._ref_count,
)
else:
logger.warning(
"Attempted to release reference to %s with ref_count=0", self.key
)
@property @property
def worker_url(self) -> str: def worker_url(self) -> str:
"""URL to use when connecting router to this worker.""" """URL to use when connecting router to this worker."""
@@ -629,8 +672,8 @@ class ModelPool:
instance = self.instances[key] instance = self.instances[key]
# Update last_used timestamp # Note: last_used is updated in acquire() which should be called by fixtures
instance.last_used = time.time() # to prevent eviction during test execution
# Verify worker is still alive and healthy # Verify worker is still alive and healthy
if not instance.is_alive(): if not instance.is_alive():
@@ -671,8 +714,15 @@ class ModelPool:
# Sort by last_used descending (MRU eviction) - evict most recently used first # Sort by last_used descending (MRU eviction) - evict most recently used first
# Store (dict_key, instance) tuples to preserve the actual key for eviction # Store (dict_key, instance) tuples to preserve the actual key for eviction
# Note: Make a copy of items to avoid RuntimeError if dict is modified during iteration
evictable: list[tuple[str, ModelInstance]] = [] evictable: list[tuple[str, ModelInstance]] = []
for dict_key, inst in self.instances.items(): for dict_key, inst in list(self.instances.items()):
# Skip instances with active references (tests using them)
if inst.is_in_use:
logger.debug(
"Skipping eviction of %s - has active references", dict_key
)
continue
if exclude_worker_types is not None: if exclude_worker_types is not None:
# Precise matching with worker types # Precise matching with worker types
# Must match model_id AND worker_type, mode is optional # Must match model_id AND worker_type, mode is optional