Compare commits

...

14 Commits

Author SHA1 Message Date
Yuge Zhang 601f6d9fe0 second change 2025-12-03 23:47:40 +08:00
Yuge Zhang 4aceeaa9aa first change 2025-12-03 23:46:40 +08:00
Yuge Zhang 76bc122f52 resolve comments 2025-12-03 18:52:22 +08:00
Yuge Zhang 64918057aa . 2025-12-03 18:19:22 +08:00
Yuge Zhang 3ad19c033d . 2025-12-03 18:06:53 +08:00
Yuge Zhang 7eac44e375 update trainer comment 2025-12-03 17:55:44 +08:00
Yuge Zhang 0a4a4f6c8d optimize verl 2025-12-03 17:26:54 +08:00
Yuge Zhang 1ce6d5a295 add tests for concurrent enqueue 2025-12-03 17:22:37 +08:00
Yuge Zhang add4a1396f update tests for enqueue_rollout and dequeue_rollout 2025-12-03 16:58:31 +08:00
Yuge Zhang 9fbec629af enqueue many and dequeue many 2025-12-03 14:54:48 +08:00
Yuge Zhang d4f6dfeb07 add tests 2025-12-03 14:11:40 +08:00
Yuge Zhang 28b7f939f7 add worker_id to several methods 2025-12-03 13:53:08 +08:00
Yuge Zhang 85454b88ea add worker id to signatures 2025-12-03 12:45:11 +08:00
Yuge Zhang aa1da5b2f9 turn on thread safe by default in some cases 2025-12-03 11:40:18 +08:00
18 changed files with 986 additions and 197 deletions
+3 -1
View File
@@ -64,7 +64,9 @@ def main(argv: Iterable[str] | None = None) -> int:
setup_logging(args.log_level)
if args.backend == "memory":
store = InMemoryLightningStore(prometheus=args.prometheus)
store = InMemoryLightningStore(
prometheus=args.prometheus, thread_safe=True
) # Using thread_safe store for server
elif args.backend == "mongo":
from agentlightning.store.mongo import MongoLightningStore
+4 -18
View File
@@ -72,8 +72,8 @@ class LitAgentRunner(Runner[T_task]):
tracer: Tracer,
max_rollouts: Optional[int] = None,
poll_interval: float = 5.0,
heartbeat_interval: float = 10.0,
interval_jitter: float = 0.1,
heartbeat_interval: float = 0.0,
interval_jitter: float = 0.5,
heartbeat_launch_mode: Literal["asyncio", "thread"] = "asyncio",
) -> None:
"""Initialize the agent runner.
@@ -577,16 +577,6 @@ class LitAgentRunner(Runner[T_task]):
if next_rollout is None:
return
try:
# Claim the rollout but updating the current worker id
await store.update_attempt(
next_rollout.rollout_id, next_rollout.attempt.attempt_id, worker_id=self.get_worker_id()
)
except Exception:
# This exception could happen if the rollout is dequeued and the other end died for some reason
logger.exception(f"{self._log_prefix()} Exception during update_attempt, giving up the rollout.")
continue
# Execute the step
await self._step_impl(next_rollout)
@@ -640,12 +630,8 @@ class LitAgentRunner(Runner[T_task]):
else:
resources_id = None
attempted_rollout = await self.get_store().start_rollout(input=input, mode=mode, resources_id=resources_id)
# Register the attempt as running by the current worker
await self.get_store().update_attempt(
attempted_rollout.rollout_id,
attempted_rollout.attempt.attempt_id,
worker_id=self.get_worker_id(),
attempted_rollout = await self.get_store().start_rollout(
input=input, mode=mode, resources_id=resources_id, worker_id=self.get_worker_id()
)
rollout_id = await self._step_impl(attempted_rollout, raise_on_exception=True)
+11 -17
View File
@@ -10,6 +10,7 @@ from agentlightning.types import (
Attempt,
AttemptedRollout,
AttemptStatus,
EnqueueRolloutRequest,
NamedResources,
ResourcesUpdate,
Rollout,
@@ -100,19 +101,6 @@ class LightningStoreStatistics(TypedDict, total=False):
"""Memory capacity of the store in bytes."""
class _EnqueueRolloutRequestRequired(TypedDict):
input: TaskInput
class EnqueueRolloutRequest(_EnqueueRolloutRequestRequired, total=False):
"""Payload describing a rollout to be queued via `enqueue_rollout`."""
mode: Optional[RolloutMode]
resources_id: Optional[str]
config: Optional[RolloutConfig]
metadata: Optional[Dict[str, Any]]
class LightningStore:
"""Contract for the persistent control-plane that coordinates training rollouts.
@@ -174,6 +162,7 @@ class LightningStore:
resources_id: str | None = None,
config: RolloutConfig | None = None,
metadata: Dict[str, Any] | None = None,
worker_id: str | None = None,
) -> AttemptedRollout:
"""Register a rollout and immediately create its first attempt.
@@ -196,6 +185,7 @@ class LightningStore:
resources_id: Concrete resource snapshot to execute against; defaults to the latest stored snapshot.
config: Rollout retry/timeout policy. Should default to a fresh [`RolloutConfig`][agentlightning.RolloutConfig].
metadata: Free-form metadata persisted verbatim with the rollout.
worker_id: Optional worker identifier to associate the new attempt with.
Returns:
The fully-populated [`AttemptedRollout`][agentlightning.AttemptedRollout] including
@@ -241,7 +231,7 @@ class LightningStore:
"""
raise NotImplementedError()
async def enqueue_many_rollouts(self, inputs: Sequence[EnqueueRolloutRequest]) -> Sequence[Rollout]:
async def enqueue_many_rollouts(self, rollouts: Sequence[EnqueueRolloutRequest]) -> Sequence[Rollout]:
"""Persist multiple rollouts in `queuing` state.
The implementation can delegate to [`enqueue_rollout()`][agentlightning.LightningStore.enqueue_rollout]
@@ -249,11 +239,11 @@ class LightningStore:
more efficient bulk enqueue semantics.
Args:
inputs: Rollout submission payloads mirroring [`enqueue_rollout()`][agentlightning.LightningStore.enqueue_rollout]'s
rollouts: Rollout submission payloads mirroring [`enqueue_rollout()`][agentlightning.LightningStore.enqueue_rollout]'s
parameters. Each entry requires `input` and can optionally include other fields.
Returns:
Rollouts enqueued in the same order as `inputs`.
Rollouts enqueued in the same order as `rollouts`.
"""
raise NotImplementedError()
@@ -273,6 +263,9 @@ class LightningStore:
* Optionally refresh the caller's [`Worker`][agentlightning.Worker] telemetry
(e.g., `last_dequeue_time`) when `worker_id` is provided.
Args:
worker_id: Optional worker identifier to associate the claimed attempt with.
Returns:
The next attempt to execute, or `None` when no eligible rollouts are queued.
@@ -304,7 +297,7 @@ class LightningStore:
"""
raise NotImplementedError()
async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
async def start_attempt(self, rollout_id: str, worker_id: Optional[str] = None) -> AttemptedRollout:
"""Create a manual retry attempt for an existing rollout.
This is typically invoked by runners that wish to retry outside of the
@@ -315,6 +308,7 @@ class LightningStore:
Args:
rollout_id: Unique identifier of the rollout receiving a new attempt.
worker_id: Optional worker identifier to associate the new attempt with.
Returns:
The rollout paired with its newly-created attempt.
+136 -47
View File
@@ -47,6 +47,7 @@ from agentlightning.types import (
Attempt,
AttemptedRollout,
AttemptStatus,
EnqueueRolloutRequest,
NamedResources,
PaginatedResult,
ResourcesUpdate,
@@ -81,12 +82,26 @@ class RolloutRequest(BaseModel):
resources_id: Optional[str] = None
config: Optional[RolloutConfig] = None
metadata: Optional[Dict[str, Any]] = None
worker_id: Optional[str] = None
class DequeueRolloutRequest(BaseModel):
worker_id: Optional[str] = None
class StartAttemptRequest(BaseModel):
worker_id: Optional[str] = None
class EnqueueManyRolloutsRequest(BaseModel):
rollouts: List[EnqueueRolloutRequest]
class DequeueManyRolloutsRequest(BaseModel):
limit: int = 1
worker_id: Optional[str] = None
class QueryRolloutsRequest(BaseModel):
status_in: Optional[List[RolloutStatus]] = Field(FastAPIQuery(default=None))
rollout_id_in: Optional[List[str]] = Field(FastAPIQuery(default=None))
@@ -522,22 +537,38 @@ class LightningStoreServer(LightningStore):
async def health(): # pyright: ignore[reportUnusedFunction]
return {"status": "ok"}
@api.post(API_AGL_PREFIX + "/queues/rollouts/enqueue", status_code=201, response_model=Rollout)
async def enqueue_rollout(request: RolloutRequest): # pyright: ignore[reportUnusedFunction]
return await self.enqueue_rollout(
input=request.input,
mode=request.mode,
resources_id=request.resources_id,
config=request.config,
metadata=request.metadata,
)
@api.post(API_AGL_PREFIX + "/queues/rollouts/enqueue", status_code=201, response_model=List[Rollout])
async def enqueue_rollouts( # pyright: ignore[reportUnusedFunction]
request: EnqueueManyRolloutsRequest,
) -> List[Rollout]:
enqueue_requests = request.rollouts
if not enqueue_requests:
return []
if len(enqueue_requests) == 1:
single = enqueue_requests[0]
rollout = await self.enqueue_rollout(
input=single.input,
mode=single.mode,
resources_id=single.resources_id,
config=single.config,
metadata=single.metadata,
)
return [rollout]
rollouts = await self.enqueue_many_rollouts(enqueue_requests)
return list(rollouts)
@api.post(API_AGL_PREFIX + "/queues/rollouts/dequeue", response_model=Optional[AttemptedRollout])
async def dequeue_rollout( # pyright: ignore[reportUnusedFunction]
request: DequeueRolloutRequest | None = Body(None),
):
worker_id = request.worker_id if request else None
return await self.dequeue_rollout(worker_id=worker_id)
@api.post(API_AGL_PREFIX + "/queues/rollouts/dequeue", response_model=List[AttemptedRollout])
async def dequeue_rollouts( # pyright: ignore[reportUnusedFunction]
request: DequeueManyRolloutsRequest | None = Body(None),
) -> List[AttemptedRollout]:
payload = request or DequeueManyRolloutsRequest()
if payload.limit <= 0:
return []
if payload.limit == 1:
single = await self.dequeue_rollout(worker_id=payload.worker_id)
return [single] if single else []
rollouts = await self.dequeue_many_rollouts(limit=payload.limit, worker_id=payload.worker_id)
return list(rollouts)
@api.post(API_AGL_PREFIX + "/rollouts", status_code=201, response_model=AttemptedRollout)
async def start_rollout(request: RolloutRequest): # pyright: ignore[reportUnusedFunction]
@@ -547,6 +578,7 @@ class LightningStoreServer(LightningStore):
resources_id=request.resources_id,
config=request.config,
metadata=request.metadata,
worker_id=request.worker_id,
)
@api.get(API_AGL_PREFIX + "/rollouts", response_model=PaginatedResult[Union[AttemptedRollout, Rollout]])
@@ -615,8 +647,11 @@ class LightningStoreServer(LightningStore):
)
@api.post(API_AGL_PREFIX + "/rollouts/{rollout_id}/attempts", status_code=201, response_model=AttemptedRollout)
async def start_attempt(rollout_id: str): # pyright: ignore[reportUnusedFunction]
return await self.start_attempt(rollout_id)
async def start_attempt( # pyright: ignore[reportUnusedFunction]
rollout_id: str, request: StartAttemptRequest | None = Body(None)
):
worker_id = request.worker_id if request else None
return await self.start_attempt(rollout_id, worker_id=worker_id)
@api.post(API_AGL_PREFIX + "/rollouts/{rollout_id}/attempts/search", response_model=PaginatedResult[Attempt])
async def search_attempts( # pyright: ignore[reportUnusedFunction]
@@ -1007,6 +1042,7 @@ class LightningStoreServer(LightningStore):
resources_id: str | None = None,
config: RolloutConfig | None = None,
metadata: Dict[str, Any] | None = None,
worker_id: Optional[str] = None,
) -> AttemptedRollout:
return await self._call_store_method(
"start_rollout",
@@ -1015,6 +1051,7 @@ class LightningStoreServer(LightningStore):
resources_id,
config,
metadata,
worker_id,
)
async def enqueue_rollout(
@@ -1034,11 +1071,22 @@ class LightningStoreServer(LightningStore):
metadata,
)
async def enqueue_many_rollouts(self, rollouts: Sequence[EnqueueRolloutRequest]) -> Sequence[Rollout]:
return await self._call_store_method("enqueue_many_rollouts", rollouts)
async def dequeue_rollout(self, worker_id: Optional[str] = None) -> Optional[AttemptedRollout]:
return await self._call_store_method("dequeue_rollout", worker_id)
async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
return await self._call_store_method("start_attempt", rollout_id)
async def dequeue_many_rollouts(
self,
*,
limit: int = 1,
worker_id: Optional[str] = None,
) -> Sequence[AttemptedRollout]:
return await self._call_store_method("dequeue_many_rollouts", limit=limit, worker_id=worker_id)
async def start_attempt(self, rollout_id: str, worker_id: Optional[str] = None) -> AttemptedRollout:
return await self._call_store_method("start_attempt", rollout_id, worker_id)
async def query_rollouts(
self,
@@ -1512,6 +1560,7 @@ class LightningStoreClient(LightningStore):
resources_id: str | None = None,
config: RolloutConfig | None = None,
metadata: Dict[str, Any] | None = None,
worker_id: Optional[str] = None,
) -> AttemptedRollout:
data = await self._request_json(
"post",
@@ -1522,6 +1571,7 @@ class LightningStoreClient(LightningStore):
resources_id=resources_id,
config=config,
metadata=metadata,
worker_id=worker_id,
).model_dump(exclude_none=False),
)
return AttemptedRollout.model_validate(data)
@@ -1534,18 +1584,64 @@ class LightningStoreClient(LightningStore):
config: RolloutConfig | None = None,
metadata: Dict[str, Any] | None = None,
) -> Rollout:
request_body = EnqueueManyRolloutsRequest(
rollouts=[
EnqueueRolloutRequest(
input=input,
mode=mode,
resources_id=resources_id,
config=config,
metadata=metadata,
)
]
).model_dump(exclude_none=False)
data = await self._request_json(
"post",
"/queues/rollouts/enqueue",
json=RolloutRequest(
input=input,
mode=mode,
resources_id=resources_id,
config=config,
metadata=metadata,
).model_dump(exclude_none=False),
json=request_body,
)
return Rollout.model_validate(data)
if not data:
raise RuntimeError("enqueue_rollout returned no rollouts")
return Rollout.model_validate(data[0])
async def enqueue_many_rollouts(self, rollouts: Sequence[EnqueueRolloutRequest]) -> Sequence[Rollout]:
if not rollouts:
return []
request_body = EnqueueManyRolloutsRequest(rollouts=list(rollouts)).model_dump(exclude_none=False)
data = await self._request_json(
"post",
"/queues/rollouts/enqueue",
json=request_body,
)
return [Rollout.model_validate(entry) for entry in data]
async def _dequeue_batch(
self,
*,
limit: int,
worker_id: Optional[str],
) -> List[AttemptedRollout]:
if limit <= 0:
return []
session = await self._get_session()
url = f"{self.server_address}/queues/rollouts/dequeue"
payload: Dict[str, Any] = {"limit": limit}
if worker_id is not None:
payload["worker_id"] = worker_id
try:
async with session.post(url, json=payload) as resp:
resp.raise_for_status()
data = await resp.json()
self._dequeue_was_successful = True
return [AttemptedRollout.model_validate(item) for item in data]
except Exception as e:
if self._dequeue_was_successful:
if self._dequeue_first_unsuccessful:
client_logger.warning(f"dequeue_rollout failed with exception: {e}")
self._dequeue_first_unsuccessful = False
client_logger.debug("dequeue_rollout failed with exception. Details:", exc_info=True)
# Else ignore the exception because the server is not ready yet
return []
async def dequeue_rollout(self, worker_id: Optional[str] = None) -> Optional[AttemptedRollout]:
"""
@@ -1558,30 +1654,23 @@ class LightningStoreClient(LightningStore):
This method does NOT retry on failures. If any exception occurs (network error,
server error, etc.), it logs the error and returns None immediately.
"""
session = await self._get_session()
url = f"{self.server_address}/queues/rollouts/dequeue"
request_kwargs: Dict[str, Any] = {}
if worker_id is not None:
request_kwargs["json"] = {"worker_id": worker_id}
try:
async with session.post(url, **request_kwargs) as resp:
resp.raise_for_status()
data = await resp.json()
self._dequeue_was_successful = True
return AttemptedRollout.model_validate(data) if data else None
except Exception as e:
if self._dequeue_was_successful:
if self._dequeue_first_unsuccessful:
client_logger.warning(f"dequeue_rollout failed with exception: {e}")
self._dequeue_first_unsuccessful = False
client_logger.debug("dequeue_rollout failed with exception. Details:", exc_info=True)
# Else ignore the exception because the server is not ready yet
return None
attempts = await self._dequeue_batch(limit=1, worker_id=worker_id)
return attempts[0] if attempts else None
async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
async def dequeue_many_rollouts(
self,
*,
limit: int = 1,
worker_id: Optional[str] = None,
) -> Sequence[AttemptedRollout]:
return await self._dequeue_batch(limit=limit, worker_id=worker_id)
async def start_attempt(self, rollout_id: str, worker_id: Optional[str] = None) -> AttemptedRollout:
payload = {"worker_id": worker_id} if worker_id is not None else None
data = await self._request_json(
"post",
f"/rollouts/{rollout_id}/attempts",
json=payload,
)
return AttemptedRollout.model_validate(data)
@@ -853,6 +853,9 @@ class _ThreadSafeAsyncLock:
async def __aenter__(self):
# We run the blocking .acquire() in a thread pool so we don't block the event loop
loop = asyncio.get_running_loop()
# NOTE: If this fails to acquire, it will block the executor thread that
# is running it. That thread will not auto-terminate when asyncio is cancelled.
# Therefore, zombie thread is possible if the lock is held for a long time.
await loop.run_in_executor(None, self._lock.acquire)
return self
+1 -1
View File
@@ -1404,7 +1404,7 @@ class MongoLightningCollections(LightningCollections):
"""Perform a atomic operation on the collections."""
if commit:
raise ValueError("Commit should be used with execute() instead.")
with self._prometheus_tracker.track("atomic", self._database_name, self._collection_name):
with self._prometheus_tracker.track(f"atomic__{mode}", self._database_name, self._collection_name):
# First step: ensure all collections exist before going into the atomic block
if not self._collection_ensured:
await self._ensure_collections()
+200 -75
View File
@@ -46,6 +46,7 @@ from agentlightning.types import (
Attempt,
AttemptedRollout,
AttemptStatus,
EnqueueRolloutRequest,
FilterField,
NamedResources,
PaginatedResult,
@@ -277,7 +278,9 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
return updated_workers[0]
@tracked("_unlocked_sync_worker_with_attempt")
async def _unlocked_sync_worker_with_attempt(self, collections: T_collections, attempt: Attempt) -> None:
async def _unlocked_sync_worker_with_attempt(
self, collections: T_collections, attempt: Attempt, dequeue: bool
) -> None:
"""Update the worker's status. This can be done in a separate session."""
worker_id = attempt.worker_id
if not worker_id:
@@ -288,6 +291,10 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
worker = Worker(worker_id=worker_id)
now = time.time()
# This is called from dequeue_rollout
if dequeue:
worker.last_dequeue_time = now
if attempt.status in ("succeeded", "failed"):
if worker.status != "idle":
worker.last_idle_time = now
@@ -322,10 +329,30 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
@tracked("_sync_workers_with_attempts")
@_with_collections_execute(labels=["workers", "attempts"])
async def _sync_workers_with_attempts(self, collections: T_collections, attempts: Sequence[Attempt]) -> None:
"""Update the worker's status. Locked bulk version of `_unlocked_sync_workers_with_attempts`."""
async def _sync_workers_with_attempts(
self, collections: T_collections, attempts: Sequence[Attempt], dequeue: bool = False
) -> None:
"""Update the worker's status. Locked bulk version of `_unlocked_sync_workers_with_attempts`.
Use `dequeue = True` if `last_dequeue_time` should be updated.
"""
for attempt in attempts:
await self._unlocked_sync_worker_with_attempt(collections, attempt)
await self._unlocked_sync_worker_with_attempt(collections, attempt, dequeue)
@tracked("_dequeue_mark_worker_idle")
async def _dequeue_mark_worker_idle(self, worker_id: str) -> None:
"""Dequeue fails and mark the worker as idle."""
async with self.collections.atomic(mode="r", snapshot=self._read_snapshot, labels=["workers"]) as collections:
worker = await collections.workers.get({"worker_id": {"exact": worker_id}})
now = time.time()
if not worker or worker.status != "idle":
# should mark the worker as idle
worker = Worker(worker_id=worker_id, status="idle", last_idle_time=now, last_dequeue_time=now)
await self._update_or_insert_worker(worker, update_fields=["status", "last_idle_time", "last_dequeue_time"])
else:
# only update last_dequeue_time
worker = Worker(worker_id=worker_id, last_dequeue_time=now)
await self._update_or_insert_worker(worker, update_fields=["last_dequeue_time"])
@tracked("start_rollout")
@healthcheck_before
@@ -336,6 +363,7 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
resources_id: str | None = None,
config: RolloutConfig | None = None,
metadata: Dict[str, Any] | None = None,
worker_id: str | None = None,
) -> AttemptedRollout:
"""Notify the store that I'm about to run a rollout.
@@ -370,6 +398,7 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
sequence_id=1,
start_time=current_time,
status="preparing",
worker_id=worker_id,
)
async def _insert_rollout_and_attempt(collections: T_collections) -> None:
@@ -387,9 +416,50 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
all_fields = list(rollout.__class__.model_fields.keys())
await self._post_update_rollout([(rollout, all_fields)])
if worker_id is not None:
await self._sync_workers_with_attempts([attempt])
# Return a rollout with attempt attached.
return AttemptedRollout(**rollout.model_dump(), attempt=attempt)
@tracked("_enqueue_many_rollouts")
@_with_collections_execute(labels=["rollouts", "rollout_queue"])
async def _enqueue_many_rollouts(self, collections: T_collections, rollouts: Sequence[Rollout]) -> None:
"""Enqueue many rollouts into the rollout queue. Locked bulk version."""
rollout_ids = [rollout.rollout_id for rollout in rollouts]
await collections.rollout_queue.enqueue(rollout_ids)
await collections.rollouts.insert(rollouts)
@tracked("_prepare_single_rollout")
async def _prepare_single_rollout(
self,
input: TaskInput,
mode: Literal["train", "val", "test"] | None = None,
resources_id: str | None = None,
config: RolloutConfig | None = None,
metadata: Dict[str, Any] | None = None,
) -> Rollout:
"""Prepare a single rollout object without enqueuing it.
Expects resources_id to have been resolved.
"""
rollout_id = _generate_rollout_id()
current_time = time.time()
rollout_config = config.model_copy(deep=True) if config is not None else RolloutConfig()
rollout_metadata = dict(metadata) if metadata is not None else {}
return Rollout(
rollout_id=rollout_id,
input=input,
mode=mode,
resources_id=resources_id,
start_time=current_time,
status="queuing", # should be queuing
config=rollout_config,
metadata=rollout_metadata,
)
@tracked("enqueue_rollout")
@healthcheck_before
async def enqueue_rollout(
@@ -404,90 +474,103 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
See [`LightningStore.enqueue_rollout()`][agentlightning.LightningStore.enqueue_rollout] for semantics.
"""
rollout_id = _generate_rollout_id()
current_time = time.time()
rollout_config = config.model_copy(deep=True) if config is not None else RolloutConfig()
rollout_metadata = dict(metadata) if metadata is not None else {}
if resources_id is None:
latest_resources = await self._get_latest_resources()
resources_id = latest_resources.resources_id if latest_resources is not None else None
rollout = Rollout(
rollout_id=rollout_id,
rollout = await self._prepare_single_rollout(
input=input,
mode=mode,
resources_id=resources_id,
start_time=current_time,
status="queuing", # should be queuing
config=rollout_config,
metadata=rollout_metadata,
mode=mode,
config=config,
metadata=metadata,
)
async def _insert_rollout_and_enqueue(collections: T_collections) -> None:
await collections.rollouts.insert([rollout])
await collections.rollout_queue.enqueue([rollout.rollout_id]) # add it to the end of the queue
await self.collections.execute(
_insert_rollout_and_enqueue,
mode="rw",
snapshot=self._read_snapshot,
commit=True,
labels=["rollouts", "rollout_queue"],
)
await self._enqueue_many_rollouts([rollout])
# Notify the subclass that the rollout status has changed.
all_fields = list(rollout.__class__.model_fields.keys())
all_fields = list(Rollout.model_fields.keys())
await self._post_update_rollout([(rollout, all_fields)])
# Return the rollout with no attempt attached.
return rollout
@tracked("_post_dequeue_rollout")
@tracked("enqueue_many_rollouts")
@healthcheck_before
async def enqueue_many_rollouts(self, rollouts: Sequence[EnqueueRolloutRequest]) -> Sequence[Rollout]:
"""Adds many rollouts in a batch."""
prepared_rollouts: List[Rollout] = []
latest_resources = await self._get_latest_resources()
for request in rollouts:
resources_id = request.resources_id
if resources_id is None:
resources_id = latest_resources.resources_id if latest_resources is not None else None
rollout = await self._prepare_single_rollout(
input=request.input,
resources_id=resources_id,
mode=request.mode,
config=request.config,
metadata=request.metadata,
)
prepared_rollouts.append(rollout)
await self._enqueue_many_rollouts(prepared_rollouts)
all_fields = list(Rollout.model_fields.keys())
rollout_updates = [(rollout, all_fields) for rollout in prepared_rollouts]
await self._post_update_rollout(rollout_updates)
return prepared_rollouts
@tracked("_post_dequeue_rollouts")
@_with_collections_execute(labels=["rollouts", "attempts"])
async def _post_dequeue_rollout(
self, collections: T_collections, rollout_id: str
) -> Optional[Tuple[AttemptedRollout, Sequence[str]]]:
async def _post_dequeue_rollouts(
self, collections: T_collections, rollout_ids: Sequence[str], worker_id: Optional[str]
) -> Sequence[Tuple[AttemptedRollout, Sequence[str]]]:
"""Post-dequeue logic for the rollout. Returns the rollout and the update fields (for post-update logic)."""
rollout = await collections.rollouts.get({"rollout_id": {"exact": rollout_id}})
if not rollout:
logger.warning(f"Rollout {rollout_id} not found, skipping dequeuing")
return None
rollouts = await collections.rollouts.query({"rollout_id": {"within": rollout_ids}})
if not rollouts:
logger.warning(f"No rollout found for rollout IDs: {rollout_ids}, skipping dequeuing")
return []
# Check if rollout is still in a queuing state
# (it might have been updated to a different status while in queue)
if is_queuing(rollout):
# Create a new attempt (could be first attempt or retry)
attempt_id = _generate_attempt_id()
current_time = time.time()
dequeue_results: List[Tuple[AttemptedRollout, Sequence[str]]] = []
for rollout in rollouts:
# Check if rollout is still in a queuing state
# (it might have been updated to a different status while in queue)
if is_queuing(rollout):
# Create a new attempt (could be first attempt or retry)
attempt_id = _generate_attempt_id()
current_time = time.time()
# Get existing attempts to determine sequence number
existing_attempts = await self._unlocked_query_attempts_for_rollout(collections, rollout.rollout_id)
sequence_id = len(existing_attempts) + 1
# Get existing attempts to determine sequence number
existing_attempts = await self._unlocked_query_attempts_for_rollout(collections, rollout.rollout_id)
sequence_id = len(existing_attempts) + 1
attempt = Attempt(
rollout_id=rollout.rollout_id,
attempt_id=attempt_id,
sequence_id=sequence_id,
start_time=current_time,
status="preparing",
)
attempt = Attempt(
rollout_id=rollout.rollout_id,
attempt_id=attempt_id,
sequence_id=sequence_id,
start_time=current_time,
status="preparing",
worker_id=worker_id,
)
await collections.attempts.insert([attempt])
await collections.attempts.insert([attempt])
# Sync attempt status to rollout
rollout, update_fields = await self._unlocked_update_rollout_only(
collections, rollout.rollout_id, status="preparing"
)
return AttemptedRollout(**rollout.model_dump(), attempt=attempt), update_fields
# Sync attempt status to rollout
rollout, update_fields = await self._unlocked_update_rollout_only(
collections, rollout.rollout_id, status="preparing"
)
dequeue_results.append((AttemptedRollout(**rollout.model_dump(), attempt=attempt), update_fields))
else:
# If not in queuing state, skip this rollout and continue
# (it was updated externally and should not be processed)
logger.warning(
f"Rollout {rollout.rollout_id} is not in queuing state: {rollout.status}, skipping dequeuing"
)
return None
else:
# If not in queuing state, skip this rollout and continue
# (it was updated externally and should not be processed)
logger.warning(
f"Rollout {rollout.rollout_id} is not in queuing state: {rollout.status}, skipping dequeuing"
)
return dequeue_results
@tracked("dequeue_rollout")
@healthcheck_before
@@ -499,10 +582,6 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
See [`LightningStore.dequeue_rollout()`][agentlightning.LightningStore.dequeue_rollout] for semantics.
"""
if worker_id is not None:
new_worker = Worker(worker_id=worker_id, status="idle", last_dequeue_time=time.time())
await self._update_or_insert_worker(new_worker, update_fields=["status", "last_dequeue_time"])
# Keep looking until we find a rollout that's still in queuing status
# or the queue is empty
while True:
@@ -514,20 +593,62 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
break
rollout_id = dequeued[0]
post_dequeue_result = await self._post_dequeue_rollout(rollout_id)
if post_dequeue_result is not None:
attempted_rollout, update_fields = post_dequeue_result
await self._post_update_rollout([(attempted_rollout, update_fields)])
post_dequeue_result = await self._post_dequeue_rollouts([rollout_id], worker_id)
if post_dequeue_result:
await self._post_update_rollout(post_dequeue_result)
attempted_rollout, _ = post_dequeue_result[0]
if worker_id is not None:
await self._sync_workers_with_attempts([attempted_rollout.attempt], dequeue=True)
return attempted_rollout
# else continue the loop
# No valid rollouts found
# if worker_id is not None:
# # Mark the current worker as idle
# await self._dequeue_mark_worker_idle(worker_id)
return None
@tracked("dequeue_many_rollouts")
@healthcheck_before
async def dequeue_many_rollouts(
self, *, limit: int = 1, worker_id: Optional[str] = None
) -> Sequence[AttemptedRollout]:
"""Retrieves up to `limit` tasks from the queue without blocking."""
dequeued_rollouts: List[AttemptedRollout] = []
# Keep looking until we find a rollout that's still in queuing status
# or the queue is empty
while len(dequeued_rollouts) < limit:
rest_limit = limit - len(dequeued_rollouts)
async with self.collections.atomic(
mode="rw", snapshot=self._read_snapshot, labels=["rollout_queue"]
) as collections:
dequeued = await collections.rollout_queue.dequeue(rest_limit)
if not dequeued:
# have no more rollouts in the queue; break.
break
post_dequeue_result = await self._post_dequeue_rollouts(dequeued, worker_id)
if post_dequeue_result:
await self._post_update_rollout(post_dequeue_result)
dequeued_rollouts.extend([item for item, _ in post_dequeue_result])
# else continue the loop
# Final cleanup and worker status update
if worker_id is not None:
if dequeued_rollouts:
# NOTE: One worker can currently only associated with one attempt.
# Assuming the worker is working on the last dequeued rollout.
await self._sync_workers_with_attempts([dequeued_rollouts[-1].attempt], dequeue=True)
else:
# Mark the current worker as idle
await self._dequeue_mark_worker_idle(worker_id)
return dequeued_rollouts
@tracked("start_attempt")
@healthcheck_before
async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
async def start_attempt(self, rollout_id: str, worker_id: Optional[str] = None) -> AttemptedRollout:
"""Creates a new attempt for a given rollout ID and return the attempt details.
See [`LightningStore.start_attempt()`][agentlightning.LightningStore.start_attempt] for semantics.
@@ -556,6 +677,7 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
sequence_id=sequence_id,
start_time=current_time,
status="preparing",
worker_id=worker_id,
)
# Add attempt to storage
@@ -572,6 +694,9 @@ class CollectionBasedLightningStore(LightningStore, Generic[T_collections]):
)
await self._post_update_rollout([(rollout, update_fields)])
if worker_id is not None:
await self._sync_workers_with_attempts([attempt])
# Return the rollout with the new attempt attached.
return AttemptedRollout(**rollout.model_dump(), attempt=attempt)
+25 -3
View File
@@ -11,6 +11,7 @@ from agentlightning.types import (
Attempt,
AttemptedRollout,
AttemptStatus,
EnqueueRolloutRequest,
NamedResources,
ResourcesUpdate,
Rollout,
@@ -59,9 +60,17 @@ class LightningStoreThreaded(LightningStore):
resources_id: str | None = None,
config: RolloutConfig | None = None,
metadata: Dict[str, Any] | None = None,
worker_id: Optional[str] = None,
) -> AttemptedRollout:
with self._lock:
return await self.store.start_rollout(input, mode, resources_id, config, metadata)
return await self.store.start_rollout(
input,
mode,
resources_id,
config,
metadata,
worker_id,
)
async def enqueue_rollout(
self,
@@ -74,13 +83,26 @@ class LightningStoreThreaded(LightningStore):
with self._lock:
return await self.store.enqueue_rollout(input, mode, resources_id, config, metadata)
async def enqueue_many_rollouts(self, rollouts: Sequence[EnqueueRolloutRequest]) -> Sequence[Rollout]:
with self._lock:
return await self.store.enqueue_many_rollouts(rollouts)
async def dequeue_rollout(self, worker_id: Optional[str] = None) -> Optional[AttemptedRollout]:
with self._lock:
return await self.store.dequeue_rollout(worker_id=worker_id)
async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
async def dequeue_many_rollouts(
self,
*,
limit: int = 1,
worker_id: Optional[str] = None,
) -> Sequence[AttemptedRollout]:
with self._lock:
return await self.store.start_attempt(rollout_id)
return await self.store.dequeue_many_rollouts(limit=limit, worker_id=worker_id)
async def start_attempt(self, rollout_id: str, worker_id: Optional[str] = None) -> AttemptedRollout:
with self._lock:
return await self.store.start_attempt(rollout_id, worker_id)
async def query_rollouts(
self,
+14 -7
View File
@@ -213,10 +213,6 @@ class Trainer(TrainerLegacy):
# We might be able to support a list of resources in future.
self.initial_resources = initial_resources
# The active store for the current execution context
self.store = self._make_store(store)
self.runner = self._make_runner(runner)
self.port = port
self.strategy = self._make_strategy(
@@ -224,6 +220,11 @@ class Trainer(TrainerLegacy):
n_runners=self.n_runners,
port=port,
)
# The active store for the current execution context
self.store = self._make_store(store, self.strategy)
self.runner = self._make_runner(runner)
if hasattr(self.strategy, "n_runners"):
strategy_runners = getattr(self.strategy, "n_runners")
if isinstance(strategy_runners, int) and strategy_runners > 0:
@@ -282,13 +283,19 @@ class Trainer(TrainerLegacy):
type_error_fmt="Adapter factory returned {type_name}, which is not a TraceAdapter subclass.",
)
def _make_store(self, store: ComponentSpec[LightningStore]) -> LightningStore:
"""Resolve the store implementation backing rollouts, attempts, spans, and resources."""
def _make_store(self, store: ComponentSpec[LightningStore], strategy: ExecutionStrategy) -> LightningStore:
"""Resolve the store implementation backing rollouts, attempts, spans, and resources.
By default, it's always a in-memory store. If using a client/server execution strategy,
the in-memory store will be initialized in a thread-safe manner.
"""
is_client_server = isinstance(strategy, ClientServerExecutionStrategy)
default_store_factory = lambda: InMemoryLightningStore(thread_safe=is_client_server)
return build_component(
store,
expected_type=LightningStore,
spec_name="store",
default_factory=InMemoryLightningStore,
default_factory=default_store_factory,
invalid_spec_error_fmt="Invalid store type: {actual_type}. Expected LightningStore, str, dict, or None.",
type_error_fmt="Store factory returned {type_name}, which is not a LightningStore subclass.",
)
+19
View File
@@ -53,6 +53,7 @@ __all__ = [
"Rollout",
"Attempt",
"AttemptedRollout",
"EnqueueRolloutRequest",
"Hook",
"Worker",
"WorkerStatus",
@@ -211,6 +212,24 @@ class AttemptedRollout(Rollout):
return self
class EnqueueRolloutRequest(BaseModel):
"""Payload describing a rollout to be queued via [`enqueue_rollout`][agentlightning.LightningStore.enqueue_rollout].
A subset of fields from [`Rollout`][agentlightning.Rollout] used for queuing new rollouts.
"""
input: TaskInput
"""Task input used to generate the rollout."""
mode: Optional[RolloutMode] = None
"""Execution mode such as `"train"`, `"val"` or `"test"`. See [`RolloutMode`][agentlightning.RolloutMode]."""
resources_id: Optional[str] = None
"""Identifier of the resources required to execute the rollout."""
config: Optional[RolloutConfig] = None
"""Retry and timeout configuration associated with the rollout."""
metadata: Optional[Dict[str, Any]] = None
"""Additional metadata attached to the rollout."""
WorkerStatus = Literal["idle", "busy", "unknown"]
+37 -22
View File
@@ -9,7 +9,7 @@ import time
import uuid
from collections import defaultdict
from collections.abc import Mapping
from typing import Any, Dict, List, Literal, Optional, Tuple
from typing import Any, Dict, List, Literal, Optional, Tuple, cast
import numpy as np
import requests
@@ -22,7 +22,7 @@ from agentlightning import LLM, AgentLightningServer, NamedResources, RolloutLeg
from agentlightning.adapter.triplet import TracerTraceToTriplet, TraceToTripletBase
from agentlightning.llm_proxy import LLMProxy, ModelConfig
from agentlightning.store.base import LightningStore
from agentlightning.types import Rollout, RolloutConfig, Task
from agentlightning.types import EnqueueRolloutRequest, Rollout, RolloutConfig, Task
__all__ = [
"AgentModeDaemon",
@@ -377,42 +377,57 @@ class AgentModeDaemon:
num_samples = len(data[keys[0]])
rollouts_per_sample = self.train_rollout_n if is_train else 1
enqueue_rollout_requests: List[EnqueueRolloutRequest] = []
data_id_to_original_sample: Dict[str, Dict[str, Any]] = {}
for i in range(num_samples):
data_id = str(uuid.uuid4())
original_sample = {key: data[key][i] for key in keys}
original_sample["data_id"] = data_id
data_id_to_original_sample[data_id] = original_sample
# For training, each sample is rolled out multiple times
# Data ID is different from Rollout ID, as one data can have multiple rollouts.
for _ in range(rollouts_per_sample):
task_metadata = {"data_id": data_id, "is_train": is_train}
# Data ID is different from Rollout ID, as one data can have multiple rollouts.
if self.mode == "v0":
# Queue immediately
rollout_id = await self.server.queue_task(
sample=_to_native(original_sample),
mode="train" if is_train else "val",
resources_id=resources_id,
metadata=task_metadata,
)
else:
rollout = await self.store.enqueue_rollout(
input=_to_native(original_sample),
mode="train" if is_train else "val",
resources_id=resources_id,
metadata=task_metadata,
)
await self.store.update_rollout(
rollout_id=rollout.rollout_id,
config=RolloutConfig(
unresponsive_seconds=self.llm_timeout_seconds,
timeout_seconds=self.llm_timeout_seconds,
),
)
rollout_id = rollout.rollout_id
# Store original sample data to reconstruct batch information later
self._task_id_to_original_sample[rollout_id] = original_sample
self._total_tasks_queued += 1
# Store original sample data to reconstruct batch information later
self._task_id_to_original_sample[rollout_id] = original_sample
self._total_tasks_queued += 1
else:
# Collect tasks to enqueue in batch and queue them later
enqueue_rollout_requests.append(
EnqueueRolloutRequest(
input=_to_native(original_sample),
mode="train" if is_train else "val",
resources_id=resources_id,
config=RolloutConfig(
unresponsive_seconds=self.llm_timeout_seconds,
timeout_seconds=self.llm_timeout_seconds,
),
metadata=task_metadata,
)
)
if self.mode == "v1":
# Enqueue all the tasks in a single batch
rollouts = await self.store.enqueue_many_rollouts(enqueue_rollout_requests)
self._task_id_to_original_sample.update(
{
# Recover the original data and store it for later use.
rollout.rollout_id: data_id_to_original_sample[cast(Dict[str, Any], rollout.metadata)["data_id"]]
for rollout in rollouts
}
)
self._total_tasks_queued += len(rollouts)
def set_up_data_and_server(self, data: Dict[str, Any], server_addresses: List[str], is_train: bool = True):
"""Synchronous wrapper for setting up data and server resources."""
+2
View File
@@ -22,6 +22,8 @@
::: agentlightning.Rollout
::: agentlightning.EnqueueRolloutRequest
::: agentlightning.Attempt
::: agentlightning.AttemptedRollout
+66
View File
@@ -706,6 +706,72 @@ async def test_step_with_custom_resources_returns_rollout() -> None:
assert result.resources_id is not None
@pytest.mark.asyncio
async def test_step_registers_worker_id_on_start_rollout(monkeypatch: pytest.MonkeyPatch) -> None:
"""runner.step should pass the formatted worker ID down to the store."""
class WorkerAwareAgent(LitAgent[Dict[str, Any]]):
def validation_rollout(self, task: Dict[str, Any], resources: Dict[str, Any], rollout: Any) -> float:
return 1.0
agent = WorkerAwareAgent()
runner, store, _ = await setup_runner(agent)
expected_worker_label = runner.get_worker_id()
captured: Dict[str, Optional[str]] = {}
original_start_rollout = store.start_rollout
async def wrapped_start_rollout(*args: Any, **kwargs: Any):
captured["worker_id"] = kwargs.get("worker_id")
return await original_start_rollout(*args, **kwargs)
monkeypatch.setattr(store, "start_rollout", wrapped_start_rollout)
try:
await runner.step({"task": "worker-aware"})
finally:
teardown_runner(runner)
assert captured["worker_id"] == expected_worker_label
@pytest.mark.asyncio
async def test_iter_passes_worker_id_to_dequeue(monkeypatch: pytest.MonkeyPatch) -> None:
"""iter() should poll the store with the formatted worker identifier."""
class IdleAgent(LitAgent[Dict[str, Any]]):
def validation_rollout(
self, task: Dict[str, Any], resources: Dict[str, Any], rollout: Any
) -> float: # pragma: no cover - not invoked
return 0.0
agent = IdleAgent()
runner, store, _ = await setup_runner(agent, poll_interval=0.01)
expected_worker_label = runner.get_worker_id()
captured: Dict[str, Optional[str]] = {}
event = ThreadingEvent()
async def fake_dequeue(*, worker_id: Optional[str] = None):
captured["worker_id"] = worker_id
event.set()
return None
async def fast_sleep(self: LitAgentRunner[Any], event: Optional[ExecutionEvent] = None) -> None:
if event is not None:
event.set()
monkeypatch.setattr(store, "dequeue_rollout", fake_dequeue)
monkeypatch.setattr(LitAgentRunner, "_sleep_until_next_poll", fast_sleep)
try:
await runner.iter(event=event)
finally:
teardown_runner(runner)
assert captured["worker_id"] == expected_worker_label
@pytest.mark.asyncio
async def test_emit_heartbeat_updates_worker_snapshot(monkeypatch: pytest.MonkeyPatch) -> None:
snapshot = {"cpu_pct": 42.0, "mem_pct": 10.5}
+18 -3
View File
@@ -10,6 +10,7 @@ from agentlightning.types import (
Attempt,
AttemptedRollout,
AttemptStatus,
EnqueueRolloutRequest,
NamedResources,
ResourcesUpdate,
Rollout,
@@ -42,8 +43,9 @@ class DummyLightningStore(LightningStore):
resources_id: Optional[str] = None,
config: Optional[RolloutConfig] = None,
metadata: Optional[Dict[str, Any]] = None,
worker_id: Optional[str] = None,
) -> AttemptedRollout:
self.calls.append(("start_rollout", (input, mode, resources_id, config, metadata), {}))
self.calls.append(("start_rollout", (input, mode, resources_id, config, metadata, worker_id), {}))
return self.return_values["start_rollout"]
async def enqueue_rollout(
@@ -57,12 +59,25 @@ class DummyLightningStore(LightningStore):
self.calls.append(("enqueue_rollout", (input, mode, resources_id, config, metadata), {}))
return self.return_values["enqueue_rollout"]
async def enqueue_many_rollouts(self, rollouts: Sequence[EnqueueRolloutRequest]) -> Sequence[Rollout]:
self.calls.append(("enqueue_many_rollouts", (rollouts,), {}))
return self.return_values["enqueue_many_rollouts"]
async def dequeue_rollout(self, worker_id: Optional[str] = None) -> Optional[AttemptedRollout]:
self.calls.append(("dequeue_rollout", (worker_id,), {}))
return self.return_values["dequeue_rollout"]
async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
self.calls.append(("start_attempt", (rollout_id,), {}))
async def dequeue_many_rollouts(
self,
*,
limit: int = 1,
worker_id: Optional[str] = None,
) -> Sequence[AttemptedRollout]:
self.calls.append(("dequeue_many_rollouts", (), {"limit": limit, "worker_id": worker_id}))
return self.return_values["dequeue_many_rollouts"]
async def start_attempt(self, rollout_id: str, worker_id: Optional[str] = None) -> AttemptedRollout:
self.calls.append(("start_attempt", (rollout_id, worker_id), {}))
return self.return_values["start_attempt"]
async def query_rollouts(self, *args: Any, **kwargs: Any) -> List[Rollout]:
+194 -3
View File
@@ -18,7 +18,16 @@ from yarl import URL
from agentlightning.store.base import UNSET, LightningStore
from agentlightning.store.client_server import LightningStoreClient, LightningStoreServer
from agentlightning.store.memory import InMemoryLightningStore
from agentlightning.types import LLM, OtelResource, PaginatedResult, PromptTemplate, RolloutConfig, Span, TraceStatus
from agentlightning.types import (
LLM,
EnqueueRolloutRequest,
OtelResource,
PaginatedResult,
PromptTemplate,
RolloutConfig,
Span,
TraceStatus,
)
from agentlightning.utils.server_launcher import LaunchMode, PythonServerLauncherArgs
@@ -159,6 +168,188 @@ async def test_server_client_statistics_match(server_client: Tuple[LightningStor
assert server_stats["total_rollouts"] >= 1 # type: ignore
@pytest.mark.asyncio
async def test_client_start_rollout_propagates_worker_id(
server_client: Tuple[LightningStoreServer, LightningStoreClient],
) -> None:
server, client = server_client
attempt = await client.start_rollout(input={"source": "remote-worker"}, worker_id="client-worker-start")
assert attempt.attempt.worker_id == "client-worker-start"
worker = await server.get_worker_by_id("client-worker-start")
assert worker is not None
assert worker.status == "busy"
assert worker.current_rollout_id == attempt.rollout_id
assert worker.current_attempt_id == attempt.attempt.attempt_id
@pytest.mark.asyncio
async def test_client_start_attempt_propagates_worker_id(
server_client: Tuple[LightningStoreServer, LightningStoreClient],
) -> None:
server, client = server_client
initial = await client.start_rollout(input={"source": "retry-worker"})
retry = await client.start_attempt(initial.rollout_id, worker_id="client-worker-retry")
assert retry.attempt.sequence_id == 2
assert retry.attempt.worker_id == "client-worker-retry"
worker = await server.get_worker_by_id("client-worker-retry")
assert worker is not None
assert worker.status == "busy"
assert worker.current_rollout_id == retry.rollout_id
assert worker.current_attempt_id == retry.attempt.attempt_id
@pytest.mark.asyncio
async def test_client_enqueue_many_rollouts_uses_batch_payload(monkeypatch: MonkeyPatch) -> None:
client = LightningStoreClient("http://localhost:9000")
captured: Dict[str, Any] = {}
async def fake_request_json(_, method: str, path: str, *, json: Any = None, params: Any = None) -> Any:
captured.update({"method": method, "path": path, "json": json})
count = len(json["rollouts"]) if json and "rollouts" in json else 0 # type: ignore[index]
return [{"rollout_id": f"bulk-{idx}", "input": {"idx": idx}, "start_time": float(idx)} for idx in range(count)]
monkeypatch.setattr(LightningStoreClient, "_request_json", fake_request_json, raising=False) # type: ignore
requests = [
EnqueueRolloutRequest(input={"idx": 0}, mode="train", metadata={"batch": "left"}),
EnqueueRolloutRequest(input={"idx": 1}, resources_id="resources-1"),
]
rollouts = await client.enqueue_many_rollouts(requests)
assert captured["method"] == "post"
assert captured["path"] == "/queues/rollouts/enqueue"
assert len(captured["json"]["rollouts"]) == 2 # type: ignore[index]
assert captured["json"]["rollouts"][0]["mode"] == "train" # type: ignore[index]
assert captured["json"]["rollouts"][1]["resources_id"] == "resources-1" # type: ignore[index]
assert len(rollouts) == 2
await client.close()
@pytest.mark.asyncio
async def test_client_dequeue_methods_share_batch_logic(monkeypatch: MonkeyPatch) -> None:
client = LightningStoreClient("http://localhost:9001")
def attempt_payload(idx: int) -> Dict[str, Any]:
attempt_id = f"attempt-{idx}"
rollout_id = f"rollout-{idx}"
return {
"rollout_id": rollout_id,
"input": {"idx": idx},
"start_time": float(idx),
"status": "preparing",
"attempt": {
"rollout_id": rollout_id,
"attempt_id": attempt_id,
"sequence_id": 1,
"start_time": float(idx),
"status": "preparing",
"worker_id": "batch-worker",
},
}
payload_queue = [
[attempt_payload(0), attempt_payload(1)],
[attempt_payload(0)],
]
class FakeResponse:
def __init__(self, body: Any):
self._body = body
self.status = 200
async def __aenter__(self) -> "FakeResponse":
return self
async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None:
return None
def raise_for_status(self) -> None:
return None
async def json(self) -> Any:
return self._body
class RecordingSession:
def __init__(self) -> None:
self.calls: list[Dict[str, Any]] = []
def post(self, url: str, json: Dict[str, Any]) -> FakeResponse:
self.calls.append({"url": url, "json": json})
body = payload_queue.pop(0)
return FakeResponse(body)
session = RecordingSession()
async def fake_get_session() -> RecordingSession:
return session
monkeypatch.setattr(client, "_get_session", fake_get_session)
batch = await client.dequeue_many_rollouts(limit=2, worker_id="batch-worker")
assert len(batch) == 2
single = await client.dequeue_rollout(worker_id="batch-worker")
assert single is not None
assert [call["json"]["limit"] for call in session.calls] == [2, 1]
await client.close()
@pytest.mark.asyncio
async def test_client_dequeue_many_rollouts_skips_network_for_non_positive_limit(monkeypatch: MonkeyPatch) -> None:
client = LightningStoreClient("http://localhost:9002")
async def fail_get_session() -> None:
pytest.fail("Client should not request a session when limit <= 0")
monkeypatch.setattr(client, "_get_session", fail_get_session)
assert await client.dequeue_many_rollouts(limit=0, worker_id="idle") == []
assert await client.dequeue_many_rollouts(limit=-5) == []
await client.close()
@pytest.mark.asyncio
async def test_client_concurrent_enqueue_many_rollouts(
server_client: Tuple[LightningStoreServer, LightningStoreClient],
) -> None:
_, client = server_client
async def enqueue_batch(batch_idx: int) -> list[str]:
requests = [EnqueueRolloutRequest(input={"batch": batch_idx, "idx": item}) for item in range(3)]
rollouts = await client.enqueue_many_rollouts(requests)
return [rollout.rollout_id for rollout in rollouts]
batches = await asyncio.gather(*(enqueue_batch(batch_idx) for batch_idx in range(5)))
all_ids = {rollout_id for batch in batches for rollout_id in batch}
assert len(all_ids) == 15
queried = await client.query_rollouts(limit=-1)
assert isinstance(queried, PaginatedResult)
assert queried.total >= 15
@pytest.mark.asyncio
async def test_client_concurrent_dequeue_many_rollouts(
server_client: Tuple[LightningStoreServer, LightningStoreClient],
) -> None:
server, client = server_client
requests = [EnqueueRolloutRequest(input={"idx": idx}) for idx in range(6)]
# Seed queue from the server to avoid races with background processing
await asyncio.gather(*(server.enqueue_rollout(**req.model_dump()) for req in requests))
async def consume(limit: int, worker: str):
return await client.dequeue_many_rollouts(limit=limit, worker_id=worker)
batches = await asyncio.gather(
consume(3, "worker-a"),
consume(3, "worker-b"),
)
claimed_ids = {attempt.rollout_id for batch in batches for attempt in batch}
assert len(claimed_ids) == 6
assert await client.dequeue_many_rollouts(limit=1) == []
@pytest.mark.asyncio
async def test_add_resources_via_server(server_client: Tuple[LightningStoreServer, LightningStoreClient]) -> None:
"""Test that add_resources works correctly via server."""
@@ -313,7 +504,7 @@ async def test_client_server_end_to_end(
dequeued = await server.dequeue_rollout(worker_id=server_worker_id)
server_worker_after_dequeue = await server.get_worker_by_id(server_worker_id)
assert server_worker_after_dequeue is not None
assert server_worker_after_dequeue.status == "idle"
assert server_worker_after_dequeue.status == "busy" # should be busy after dequeue
assert server_worker_after_dequeue.last_dequeue_time is not None
dequeue_time = server_worker_after_dequeue.last_dequeue_time
started_attempt = await server.start_attempt(queued_rollout.rollout_id)
@@ -391,7 +582,7 @@ async def test_client_server_end_to_end(
assert dequeued_client is not None
client_worker_after_dequeue = await client.get_worker_by_id(client_worker_id)
assert client_worker_after_dequeue is not None
assert client_worker_after_dequeue.status == "idle"
assert client_worker_after_dequeue.status == "busy" # should be busy after dequeue
assert client_worker_after_dequeue.last_dequeue_time is not None
client_dequeue_time = client_worker_after_dequeue.last_dequeue_time
started_client_attempt = await client.start_attempt(dequeued_client.rollout_id)
+87
View File
@@ -31,6 +31,7 @@ from agentlightning.types import (
LLM,
Attempt,
AttemptedRollout,
EnqueueRolloutRequest,
Event,
Link,
OtelResource,
@@ -675,6 +676,49 @@ async def test_requeue_mechanism(store_fixture: LightningStore) -> None:
assert latest_attempt.sequence_id == 2
@pytest.mark.asyncio
async def test_enqueue_many_rollouts_preserves_order(store_fixture: LightningStore) -> None:
"""enqueue_many_rollouts should enqueue tasks in the provided order with matching metadata."""
requests = [
EnqueueRolloutRequest(input={"idx": 0}, metadata={"batch": "a"}),
EnqueueRolloutRequest(input={"idx": 1}, mode="train"),
EnqueueRolloutRequest(input={"idx": 2}, config=RolloutConfig(timeout_seconds=3.5)),
]
rollouts = await store_fixture.enqueue_many_rollouts(requests)
assert [rollout.input["idx"] for rollout in rollouts] == [0, 1, 2]
assert all(rollout.status == "queuing" for rollout in rollouts)
assert rollouts[0].metadata == {"batch": "a"}
assert rollouts[1].mode == "train"
assert rollouts[2].config.timeout_seconds == 3.5
for expected_idx in range(3):
dequeued = await store_fixture.dequeue_rollout()
assert dequeued is not None
assert dequeued.input["idx"] == expected_idx
@pytest.mark.asyncio
async def test_dequeue_many_rollouts_with_limit(store_fixture: LightningStore) -> None:
"""dequeue_many_rollouts should honor limit and propagate worker IDs to attempts."""
requests = [EnqueueRolloutRequest(input={"idx": idx}) for idx in range(4)]
await store_fixture.enqueue_many_rollouts(requests)
first_batch = await store_fixture.dequeue_many_rollouts(limit=2, worker_id="bulk-worker")
assert len(first_batch) == 2
assert [attempt.input["idx"] for attempt in first_batch] == [0, 1]
assert all(attempt.attempt.worker_id == "bulk-worker" for attempt in first_batch)
second_batch = await store_fixture.dequeue_many_rollouts(limit=5, worker_id="bulk-worker")
assert len(second_batch) == 2
assert [attempt.input["idx"] for attempt in second_batch] == [2, 3]
assert await store_fixture.dequeue_many_rollouts(limit=1) == []
@pytest.mark.asyncio
async def test_update_and_query_workers(store_fixture: LightningStore) -> None:
"""Workers can be created, heartbeats recorded, and telemetry auto-updated."""
@@ -718,6 +762,49 @@ async def test_update_and_query_workers(store_fixture: LightningStore) -> None:
await store_fixture.update_worker("worker-1", heartbeat_stats=None) # type: ignore[arg-type]
@pytest.mark.asyncio
async def test_start_rollout_assigns_worker(store_fixture: LightningStore) -> None:
"""start_rollout should immediately associate attempts with the provided worker."""
attempted = await store_fixture.start_rollout(input={"task": "direct"}, worker_id="worker-direct")
assert attempted.attempt.worker_id == "worker-direct"
worker = await store_fixture.get_worker_by_id("worker-direct")
assert worker is not None
assert worker.status == "busy"
assert worker.current_rollout_id == attempted.rollout_id
assert worker.current_attempt_id == attempted.attempt.attempt_id
@pytest.mark.asyncio
async def test_dequeue_rollout_assigns_worker(store_fixture: LightningStore) -> None:
"""dequeue_rollout should stamp attempts and worker telemetry with worker_id."""
await store_fixture.enqueue_rollout(input={"task": "queued"})
dequeued = await store_fixture.dequeue_rollout(worker_id="worker-dequeue")
assert dequeued is not None
assert dequeued.attempt.worker_id == "worker-dequeue"
worker = await store_fixture.get_worker_by_id("worker-dequeue")
assert worker is not None
assert worker.status == "busy"
assert worker.current_rollout_id == dequeued.rollout_id
assert worker.current_attempt_id == dequeued.attempt.attempt_id
@pytest.mark.asyncio
async def test_start_attempt_assigns_worker(store_fixture: LightningStore) -> None:
"""Manual retries should also update worker state when worker_id is provided."""
initial = await store_fixture.start_rollout(input={"task": "retry-seed"})
retry = await store_fixture.start_attempt(initial.rollout_id, worker_id="worker-retry")
assert retry.attempt.sequence_id == 2
assert retry.attempt.worker_id == "worker-retry"
worker = await store_fixture.get_worker_by_id("worker-retry")
assert worker is not None
assert worker.status == "busy"
assert worker.current_rollout_id == retry.rollout_id
assert worker.current_attempt_id == retry.attempt.attempt_id
@pytest.mark.asyncio
async def test_query_workers_supports_filters(store_fixture: LightningStore) -> None:
"""Worker queries should support filtering, sorting, and pagination."""
+129
View File
@@ -208,6 +208,89 @@ async def test_cors_allows_wildcard_origin() -> None:
assert allow_credentials == "true"
@pytest.mark.asyncio
async def test_enqueue_endpoint_batches_payloads(
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
) -> None:
_server, _client, session, api_endpoint = server_client
single_payload = {"rollouts": [{"input": {"task": "single"}}]}
async with session.post(f"{api_endpoint}/queues/rollouts/enqueue", json=single_payload) as resp:
assert resp.status == 201
body = await resp.json()
assert isinstance(body, list)
assert len(body) == 1 # type: ignore
assert body[0]["input"] == {"task": "single"}
batch_payload = {
"rollouts": [
{"input": {"task": "batch-1"}, "metadata": {"batch": 1}},
{"input": {"task": "batch-2"}},
]
}
async with session.post(f"{api_endpoint}/queues/rollouts/enqueue", json=batch_payload) as resp:
assert resp.status == 201
body = await resp.json()
assert len(body) == 2
assert [item["input"]["task"] for item in body] == ["batch-1", "batch-2"]
assert body[0]["metadata"] == {"batch": 1}
@pytest.mark.asyncio
async def test_dequeue_endpoint_returns_batches(
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
) -> None:
server, _client, session, api_endpoint = server_client
for idx in range(3):
await server.enqueue_rollout(input={"idx": idx})
async with session.post(
f"{api_endpoint}/queues/rollouts/dequeue", json={"limit": 2, "worker_id": "rest-worker"}
) as resp:
assert resp.status == 200
body = await resp.json()
assert len(body) == 2
assert all(item["attempt"]["worker_id"] == "rest-worker" for item in body)
async with session.post(f"{api_endpoint}/queues/rollouts/dequeue") as resp:
assert resp.status == 200
body = await resp.json()
assert len(body) == 1
async with session.post(f"{api_endpoint}/queues/rollouts/dequeue") as resp:
assert resp.status == 200
body = await resp.json()
assert body == []
@pytest.mark.asyncio
async def test_enqueue_endpoint_requires_rollouts_field(
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
) -> None:
_server, _client, session, api_endpoint = server_client
async with session.post(f"{api_endpoint}/queues/rollouts/enqueue", json={}) as resp:
assert resp.status == 422
@pytest.mark.asyncio
async def test_dequeue_endpoint_zero_limit_returns_empty(
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
) -> None:
server, _client, session, api_endpoint = server_client
await server.enqueue_rollout(input={"idx": 0})
async with session.post(f"{api_endpoint}/queues/rollouts/dequeue", json={"limit": 0}) as resp:
assert resp.status == 200
assert await resp.json() == []
async with session.post(f"{api_endpoint}/queues/rollouts/dequeue", json={"limit": 1}) as resp:
assert resp.status == 200
body = await resp.json()
assert len(body) == 1
@pytest.mark.asyncio
async def test_statistics_endpoint_returns_counts(
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
@@ -225,6 +308,52 @@ async def test_statistics_endpoint_returns_counts(
assert payload["total_rollouts"] >= 1
@pytest.mark.asyncio
async def test_rest_start_rollout_propagates_worker_id(
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
) -> None:
server, _client, session, api_endpoint = server_client
payload = {"input": {"source": "rest-worker"}, "worker_id": "rest-start-worker"}
async with session.post(f"{api_endpoint}/rollouts", json=payload) as resp:
assert resp.status == 201
data = await resp.json()
assert data["attempt"]["worker_id"] == "rest-start-worker"
worker = await server.get_worker_by_id("rest-start-worker")
assert worker is not None
assert worker.status == "busy"
assert worker.current_rollout_id == data["rollout_id"]
assert worker.current_attempt_id == data["attempt"]["attempt_id"]
@pytest.mark.asyncio
async def test_rest_start_attempt_propagates_worker_id(
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
) -> None:
server, _client, session, api_endpoint = server_client
async with session.post(f"{api_endpoint}/rollouts", json={"input": {"source": "rest-retry"}}) as resp:
assert resp.status == 201
base_rollout = await resp.json()
attempt_worker = "rest-attempt-worker"
async with session.post(
f"{api_endpoint}/rollouts/{base_rollout['rollout_id']}/attempts",
json={"worker_id": attempt_worker},
) as resp:
assert resp.status == 201
retry_payload = await resp.json()
assert retry_payload["attempt"]["sequence_id"] == 2
assert retry_payload["attempt"]["worker_id"] == attempt_worker
worker = await server.get_worker_by_id(attempt_worker)
assert worker is not None
assert worker.status == "busy"
assert worker.current_rollout_id == retry_payload["rollout_id"]
assert worker.current_attempt_id == retry_payload["attempt"]["attempt_id"]
# Rollouts Pagination, Sorting, and Filtering Tests
+37
View File
@@ -17,6 +17,7 @@ from agentlightning.types import (
Attempt,
AttemptedRollout,
AttemptStatus,
EnqueueRolloutRequest,
NamedResources,
OtelResource,
ResourcesUpdate,
@@ -273,6 +274,42 @@ async def test_threaded_store_delegates_all_methods() -> None:
assert [name for name, *_ in dummy_store.calls] == expected_order
@pytest.mark.asyncio
async def test_threaded_store_enqueue_many_rollouts_delegates() -> None:
requests = [
EnqueueRolloutRequest(input={"idx": 0}, mode="train", metadata={"batch": "left"}),
EnqueueRolloutRequest(input={"idx": 1}, mode=None, resources_id="resources-1"),
]
rollouts = [
Rollout(rollout_id="bulk-0", input={"idx": 0}, start_time=0.0),
Rollout(rollout_id="bulk-1", input={"idx": 1}, start_time=1.0),
]
dummy_store = DummyLightningStore({"enqueue_many_rollouts": rollouts})
threaded_store = LightningStoreThreaded(dummy_store)
result = await threaded_store.enqueue_many_rollouts(requests)
assert result == rollouts
assert dummy_store.calls[-1][0] == "enqueue_many_rollouts"
assert dummy_store.calls[-1][1][0] == requests
@pytest.mark.asyncio
async def test_threaded_store_dequeue_many_rollouts_delegates() -> None:
attempt_a = Attempt(rollout_id="bulk-0", attempt_id="attempt-0", sequence_id=1, start_time=0.0)
attempt_b = Attempt(rollout_id="bulk-1", attempt_id="attempt-1", sequence_id=1, start_time=0.0)
attempts = [
AttemptedRollout(rollout_id="bulk-0", input={"idx": 0}, start_time=0.0, attempt=attempt_a),
AttemptedRollout(rollout_id="bulk-1", input={"idx": 1}, start_time=0.0, attempt=attempt_b),
]
dummy_store = DummyLightningStore({"dequeue_many_rollouts": attempts})
threaded_store = LightningStoreThreaded(dummy_store)
result = await threaded_store.dequeue_many_rollouts(limit=2, worker_id="thread-worker")
assert result == attempts
assert dummy_store.calls[-1][0] == "dequeue_many_rollouts"
assert dummy_store.calls[-1][2] == {"limit": 2, "worker_id": "thread-worker"}
def test_threaded_store_serializes_update_attempt_calls() -> None:
store = SlowAttemptStore()
threaded_store = LightningStoreThreaded(store)