Skip to content
Merged
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
136 changes: 86 additions & 50 deletions agentlightning/store/client_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,7 @@ class LightningStoreServer(LightningStore):
def __init__(self, store: LightningStore, host: str, port: int):
super().__init__()
self.store = store
self._lock = threading.Lock()
self.host = host
self.port = port
self.app: FastAPI | None = FastAPI(title="LightningStore Server")
Expand Down Expand Up @@ -149,6 +150,7 @@ def __setstate__(self, state: Dict[str, Any]):
self.port = state["port"]
self._owner_pid = state["_owner_pid"]
self._client = None
self._lock = threading.Lock()
# Do NOT reconstruct app, _uvicorn_config, _uvicorn_server
# to avoid transferring server state to subprocess

Expand Down Expand Up @@ -280,7 +282,7 @@ async def health(): # pyright: ignore[reportUnusedFunction]

@self.app.post("/start_rollout", response_model=AttemptedRollout)
async def start_rollout(request: RolloutRequest): # pyright: ignore[reportUnusedFunction]
return await self.store.start_rollout(
return await self.start_rollout(
input=request.input,
mode=request.mode,
resources_id=request.resources_id,
Expand All @@ -290,7 +292,7 @@ async def start_rollout(request: RolloutRequest): # pyright: ignore[reportUnuse

@self.app.post("/enqueue_rollout", response_model=Rollout)
async def enqueue_rollout(request: RolloutRequest): # pyright: ignore[reportUnusedFunction]
return await self.store.enqueue_rollout(
return await self.enqueue_rollout(
input=request.input,
mode=request.mode,
resources_id=request.resources_id,
Expand All @@ -300,65 +302,65 @@ async def enqueue_rollout(request: RolloutRequest): # pyright: ignore[reportUnu

@self.app.get("/dequeue_rollout", response_model=Optional[AttemptedRollout])
async def dequeue_rollout(): # pyright: ignore[reportUnusedFunction]
return await self.store.dequeue_rollout()
return await self.dequeue_rollout()

@self.app.post("/start_attempt", response_model=AttemptedRollout)
async def start_attempt(request: RolloutId): # pyright: ignore[reportUnusedFunction]
return await self.store.start_attempt(request.rollout_id)
return await self.start_attempt(request.rollout_id)

@self.app.post("/query_rollouts", response_model=List[Rollout])
async def query_rollouts(request: QueryRolloutsRequest): # pyright: ignore[reportUnusedFunction]
return await self.store.query_rollouts(status=request.status)
return await self.query_rollouts(status=request.status, rollout_ids=request.rollout_ids)

@self.app.get("/query_attempts/{rollout_id}", response_model=List[Attempt])
async def query_attempts(rollout_id: str): # pyright: ignore[reportUnusedFunction]
return await self.store.query_attempts(rollout_id)
return await self.query_attempts(rollout_id)

@self.app.get("/get_latest_attempt/{rollout_id}", response_model=Optional[Attempt])
async def get_latest_attempt(rollout_id: str): # pyright: ignore[reportUnusedFunction]
return await self.store.get_latest_attempt(rollout_id)
return await self.get_latest_attempt(rollout_id)

@self.app.get("/get_rollout_by_id/{rollout_id}", response_model=Optional[Rollout])
async def get_rollout_by_id(rollout_id: str): # pyright: ignore[reportUnusedFunction]
return await self.store.get_rollout_by_id(rollout_id)
return await self.get_rollout_by_id(rollout_id)

@self.app.post("/add_resources", response_model=ResourcesUpdate)
async def add_resources(resources: AddResourcesRequest): # pyright: ignore[reportUnusedFunction]
return await self.store.add_resources(resources.resources)
return await self.add_resources(resources.resources)

@self.app.post("/update_resources", response_model=ResourcesUpdate)
async def update_resources(update: ResourcesUpdate): # pyright: ignore[reportUnusedFunction]
return await self.store.update_resources(update.resources_id, update.resources)
return await self.update_resources(update.resources_id, update.resources)

@self.app.get("/get_resources_by_id/{resources_id}", response_model=Optional[ResourcesUpdate])
async def get_resources_by_id(resources_id: str): # pyright: ignore[reportUnusedFunction]
return await self.store.get_resources_by_id(resources_id)
return await self.get_resources_by_id(resources_id)

@self.app.get("/get_latest_resources", response_model=Optional[ResourcesUpdate])
async def get_latest_resources(): # pyright: ignore[reportUnusedFunction]
return await self.store.get_latest_resources()
return await self.get_latest_resources()

@self.app.post("/add_span", response_model=Span)
async def add_span(span: Span): # pyright: ignore[reportUnusedFunction]
return await self.store.add_span(span)
return await self.add_span(span)

@self.app.get("/get_next_span_sequence_id/{rollout_id}/{attempt_id}", response_model=int)
async def get_next_span_sequence_id(rollout_id: str, attempt_id: str): # pyright: ignore[reportUnusedFunction]
return await self.store.get_next_span_sequence_id(rollout_id, attempt_id)
return await self.get_next_span_sequence_id(rollout_id, attempt_id)

@self.app.post("/wait_for_rollouts", response_model=List[Rollout])
async def wait_for_rollouts(request: WaitForRolloutsRequest): # pyright: ignore[reportUnusedFunction]
return await self.store.wait_for_rollouts(rollout_ids=request.rollout_ids, timeout=request.timeout)
return await self.wait_for_rollouts(rollout_ids=request.rollout_ids, timeout=request.timeout)

@self.app.get("/query_spans/{rollout_id}", response_model=List[Span])
async def query_spans( # pyright: ignore[reportUnusedFunction]
rollout_id: str, attempt_id: Optional[str] = None
):
return await self.store.query_spans(rollout_id, attempt_id)
return await self.query_spans(rollout_id, attempt_id)

@self.app.post("/update_rollout", response_model=Rollout)
async def update_rollout(request: UpdateRolloutRequest): # pyright: ignore[reportUnusedFunction]
return await self.store.update_rollout(
return await self.update_rollout(
rollout_id=request.rollout_id,
input=request.input if not isinstance(request.input, PydanticUnset) else UNSET,
mode=request.mode if not isinstance(request.mode, PydanticUnset) else UNSET,
Expand All @@ -370,7 +372,7 @@ async def update_rollout(request: UpdateRolloutRequest): # pyright: ignore[repo

@self.app.post("/update_attempt", response_model=Attempt)
async def update_attempt(request: UpdateAttemptRequest): # pyright: ignore[reportUnusedFunction]
return await self.store.update_attempt(
return await self.update_attempt(
rollout_id=request.rollout_id,
attempt_id=request.attempt_id,
status=request.status if not isinstance(request.status, PydanticUnset) else UNSET,
Expand All @@ -382,6 +384,18 @@ async def update_attempt(request: UpdateAttemptRequest): # pyright: ignore[repo
)

# Delegate methods
async def _call_store_method(self, method_name: str, *args: Any, **kwargs: Any) -> Any:
backend = self._backend()
method = getattr(backend, method_name)
if backend is self.store:
if method_name == "wait_for_rollouts":
# wait_for_rollouts can block for a long time; avoid holding the lock
# so other requests can make progress while we wait.
return await method(*args, **kwargs)
with self._lock:

Copilot AI Oct 20, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Using a blocking threading.Lock inside an async code path will block the entire event loop when contention occurs (another request waiting to acquire the lock will hard-block the loop), defeating concurrency and risking stalls. Replace threading.Lock with an asyncio.Lock (initialized as self._lock = asyncio.Lock()) and use async with self._lock: so awaiting the method does not block other tasks from running while they wait for the lock.

Suggested change
with self._lock:
async with self._lock:

Copilot uses AI. Check for mistakes.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This change is intentional to make _call_store_method thread-safe.

return await method(*args, **kwargs)
return await method(*args, **kwargs)

async def start_rollout(
self,
input: TaskInput,
Expand All @@ -390,7 +404,14 @@ async def start_rollout(
config: RolloutConfig | None = None,
metadata: Dict[str, Any] | None = None,
) -> AttemptedRollout:
return await self._backend().start_rollout(input, mode, resources_id, config, metadata)
return await self._call_store_method(
"start_rollout",
input,
mode,
resources_id,
config,
metadata,
)

async def enqueue_rollout(
self,
Expand All @@ -400,45 +421,52 @@ async def enqueue_rollout(
config: RolloutConfig | None = None,
metadata: Dict[str, Any] | None = None,
) -> Rollout:
return await self._backend().enqueue_rollout(input, mode, resources_id, config, metadata)
return await self._call_store_method(
"enqueue_rollout",
input,
mode,
resources_id,
config,
metadata,
)

async def dequeue_rollout(self) -> Optional[AttemptedRollout]:
return await self._backend().dequeue_rollout()
return await self._call_store_method("dequeue_rollout")

async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
return await self._backend().start_attempt(rollout_id)
return await self._call_store_method("start_attempt", rollout_id)

async def query_rollouts(
self, *, status: Optional[Sequence[RolloutStatus]] = None, rollout_ids: Optional[Sequence[str]] = None
) -> List[Rollout]:
return await self._backend().query_rollouts(status=status, rollout_ids=rollout_ids)
return await self._call_store_method("query_rollouts", status=status, rollout_ids=rollout_ids)

async def query_attempts(self, rollout_id: str) -> List[Attempt]:
return await self._backend().query_attempts(rollout_id)
return await self._call_store_method("query_attempts", rollout_id)

async def get_latest_attempt(self, rollout_id: str) -> Optional[Attempt]:
return await self._backend().get_latest_attempt(rollout_id)
return await self._call_store_method("get_latest_attempt", rollout_id)

async def get_rollout_by_id(self, rollout_id: str) -> Optional[Rollout]:
return await self._backend().get_rollout_by_id(rollout_id)
return await self._call_store_method("get_rollout_by_id", rollout_id)

async def add_resources(self, resources: NamedResources) -> ResourcesUpdate:
return await self._backend().add_resources(resources)
return await self._call_store_method("add_resources", resources)

async def update_resources(self, resources_id: str, resources: NamedResources) -> ResourcesUpdate:
return await self._backend().update_resources(resources_id, resources)
return await self._call_store_method("update_resources", resources_id, resources)

async def get_resources_by_id(self, resources_id: str) -> Optional[ResourcesUpdate]:
return await self._backend().get_resources_by_id(resources_id)
return await self._call_store_method("get_resources_by_id", resources_id)

async def get_latest_resources(self) -> Optional[ResourcesUpdate]:
return await self._backend().get_latest_resources()
return await self._call_store_method("get_latest_resources")

async def add_span(self, span: Span) -> Span:
return await self._backend().add_span(span)
return await self._call_store_method("add_span", span)

async def get_next_span_sequence_id(self, rollout_id: str, attempt_id: str) -> int:
return await self._backend().get_next_span_sequence_id(rollout_id, attempt_id)
return await self._call_store_method("get_next_span_sequence_id", rollout_id, attempt_id)

async def add_otel_span(
self,
Expand All @@ -447,17 +475,23 @@ async def add_otel_span(
readable_span: ReadableSpan,
sequence_id: int | None = None,
) -> Span:
return await self._backend().add_otel_span(rollout_id, attempt_id, readable_span, sequence_id)
return await self._call_store_method(
"add_otel_span",
rollout_id,
attempt_id,
readable_span,
sequence_id,
)

async def wait_for_rollouts(self, *, rollout_ids: List[str], timeout: Optional[float] = None) -> List[Rollout]:
return await self._backend().wait_for_rollouts(rollout_ids=rollout_ids, timeout=timeout)
return await self._call_store_method("wait_for_rollouts", rollout_ids=rollout_ids, timeout=timeout)

async def query_spans(
self,
rollout_id: str,
attempt_id: str | Literal["latest"] | None = None,
) -> List[Span]:
return await self._backend().query_spans(rollout_id, attempt_id)
return await self._call_store_method("query_spans", rollout_id, attempt_id)

async def update_rollout(
self,
Expand All @@ -469,14 +503,15 @@ async def update_rollout(
config: RolloutConfig | Unset = UNSET,
metadata: Optional[Dict[str, Any]] | Unset = UNSET,
) -> Rollout:
return await self._backend().update_rollout(
rollout_id=rollout_id,
input=input,
mode=mode,
resources_id=resources_id,
status=status,
config=config,
metadata=metadata,
return await self._call_store_method(
"update_rollout",
rollout_id,
input,
mode,
resources_id,
status,
config,
metadata,
)

async def update_attempt(
Expand All @@ -488,13 +523,14 @@ async def update_attempt(
last_heartbeat_time: float | Unset = UNSET,
metadata: Optional[Dict[str, Any]] | Unset = UNSET,
) -> Attempt:
return await self._backend().update_attempt(
rollout_id=rollout_id,
attempt_id=attempt_id,
status=status,
worker_id=worker_id,
last_heartbeat_time=last_heartbeat_time,
metadata=metadata,
return await self._call_store_method(
"update_attempt",
rollout_id,
attempt_id,
status,
worker_id,
last_heartbeat_time,
metadata,
)


Expand Down
43 changes: 42 additions & 1 deletion agentlightning/store/memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
import threading
import time
import uuid
import weakref
from collections import deque
from collections.abc import Iterable
from collections.abc import Mapping as MappingABC
Expand Down Expand Up @@ -52,6 +53,46 @@
logger = logging.getLogger(__name__)


class _LoopAwareAsyncLock:
"""Async lock that transparently rebinds to the current event loop.

The lock intentionally remains *thread-unsafe*: callers must only use it from
one thread at a time. If multiple threads interact with the store, each
thread gets its own event loop specific lock.
"""

def __init__(self) -> None:
self._locks: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, asyncio.Lock] = weakref.WeakKeyDictionary()

# When serializing and deserializing, we don't need to serialize the locks.
# Because another process will have its own set of event loops and its own lock.
def __getstate__(self) -> dict[str, Any]:
return {}

def __setstate__(self, state: dict[str, Any]) -> None:
self._locks = weakref.WeakKeyDictionary()

def _get_lock_for_current_loop(self) -> asyncio.Lock:
loop = asyncio.get_running_loop()
lock = self._locks.get(loop)
if lock is None:
lock = asyncio.Lock()
self._locks[loop] = lock
return lock

async def __aenter__(self) -> asyncio.Lock:
lock = self._get_lock_for_current_loop()
await lock.acquire()
return lock

async def __aexit__(self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: Any) -> None:
loop = asyncio.get_running_loop()
lock = self._locks.get(loop)
if lock is None or not lock.locked():
raise RuntimeError("Lock released without being acquired")
lock.release()


def estimate_model_size(obj: Any) -> int:
"""Rough recursive size estimate for Pydantic BaseModel instances."""

Expand Down Expand Up @@ -151,7 +192,7 @@ def __init__(
safe_memory_threshold: float | int | None = None,
span_size_estimator: Callable[[Span], int] | None = None,
):
self._lock = asyncio.Lock()
self._lock = _LoopAwareAsyncLock()

# Task queue and rollouts storage
self._task_queue: deque[Rollout] = deque()
Expand Down
21 changes: 21 additions & 0 deletions tests/store/test_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -182,6 +182,27 @@ async def test_get_rollout_by_id(inmemory_store: InMemoryLightningStore) -> None
assert updated.status == "running"


@pytest.mark.asyncio
async def test_store_lock_rebinds_to_new_event_loop(
inmemory_store: InMemoryLightningStore,
) -> None:
"""The in-memory store can be reused after switching to a new event loop."""

rollout = await inmemory_store.enqueue_rollout(input={"foo": "bar"})

def run_in_new_loop() -> Optional[Rollout]:
loop = asyncio.new_event_loop()
try:
return loop.run_until_complete(inmemory_store.get_rollout_by_id(rollout.rollout_id))
finally:
loop.close()

retrieved = await asyncio.to_thread(run_in_new_loop)

assert retrieved is not None
assert retrieved.rollout_id == rollout.rollout_id


@pytest.mark.asyncio
async def test_query_rollouts_by_rollout_ids(inmemory_store: InMemoryLightningStore) -> None:
"""Test querying rollouts filtered by rollout IDs."""
Expand Down
Loading