Compare commits
14 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 601f6d9fe0 | |||
| 4aceeaa9aa | |||
| 76bc122f52 | |||
| 64918057aa | |||
| 3ad19c033d | |||
| 7eac44e375 | |||
| 0a4a4f6c8d | |||
| 1ce6d5a295 | |||
| add4a1396f | |||
| 9fbec629af | |||
| d4f6dfeb07 | |||
| 28b7f939f7 | |||
| 85454b88ea | |||
| aa1da5b2f9 |
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.",
|
||||
)
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -22,6 +22,8 @@
|
||||
|
||||
::: agentlightning.Rollout
|
||||
|
||||
::: agentlightning.EnqueueRolloutRequest
|
||||
|
||||
::: agentlightning.Attempt
|
||||
|
||||
::: agentlightning.AttemptedRollout
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user