From 3de299804e952bea1f2a789d264d300800c94499 Mon Sep 17 00:00:00 2001 From: Yuge Zhang Date: Sun, 19 Oct 2025 19:33:04 -0700 Subject: [PATCH 1/2] Ensure store server serializes access per thread --- agentlightning/store/client_server.py | 136 ++++++++++++++++---------- agentlightning/store/memory.py | 35 ++++++- tests/store/test_memory.py | 21 ++++ 3 files changed, 141 insertions(+), 51 deletions(-) diff --git a/agentlightning/store/client_server.py b/agentlightning/store/client_server.py index 06c6daf77..9bd5281a3 100644 --- a/agentlightning/store/client_server.py +++ b/agentlightning/store/client_server.py @@ -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") @@ -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 @@ -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, @@ -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, @@ -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, @@ -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, @@ -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: + return await method(*args, **kwargs) + return await method(*args, **kwargs) + async def start_rollout( self, input: TaskInput, @@ -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, @@ -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, @@ -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, @@ -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( @@ -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, ) diff --git a/agentlightning/store/memory.py b/agentlightning/store/memory.py index b8f720176..5b62a93e7 100644 --- a/agentlightning/store/memory.py +++ b/agentlightning/store/memory.py @@ -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 @@ -52,6 +53,38 @@ 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() + + 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.""" @@ -151,7 +184,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() diff --git a/tests/store/test_memory.py b/tests/store/test_memory.py index 886c330a3..20a3ea732 100644 --- a/tests/store/test_memory.py +++ b/tests/store/test_memory.py @@ -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.""" From 9388f1066f0b2ae07fdeacd9c300dc29602525b2 Mon Sep 17 00:00:00 2001 From: Yuge Zhang Date: Mon, 20 Oct 2025 11:14:41 +0800 Subject: [PATCH 2/2] fix getstate and setstate --- agentlightning/store/memory.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/agentlightning/store/memory.py b/agentlightning/store/memory.py index 5b62a93e7..00f793a3e 100644 --- a/agentlightning/store/memory.py +++ b/agentlightning/store/memory.py @@ -62,7 +62,15 @@ class _LoopAwareAsyncLock: """ def __init__(self) -> None: - self._locks: "weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, asyncio.Lock]" = weakref.WeakKeyDictionary() + 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()