Compare commits
22 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 0a050533b1 | |||
| 2ff6fb58c4 | |||
| 21f7a1826b | |||
| 761912a6e8 | |||
| 1af34e8e71 | |||
| 3ca278bfea | |||
| e9c069fca8 | |||
| eb05ec5cf8 | |||
| 179e6397e6 | |||
| 986aa389b5 | |||
| 1b69fab51d | |||
| 0094bb2eb0 | |||
| b554be2634 | |||
| b9ee261b4c | |||
| 5ad0f3616c | |||
| 737edebec0 | |||
| 4c3e59ee2e | |||
| 5fe113fc9a | |||
| 06652d5b6b | |||
| 59b5bc48a8 | |||
| 6e2637c6fb | |||
| f2810c2b72 |
+134
-31
@@ -11,8 +11,21 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any, List, Literal, Optional, Sequence, TypeVar, cast
|
||||
from contextlib import suppress
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Awaitable,
|
||||
Callable,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Sequence,
|
||||
TypeVar,
|
||||
cast,
|
||||
)
|
||||
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
|
||||
@@ -30,6 +43,7 @@ from agentlightning.types import (
|
||||
RolloutRawResult,
|
||||
Span,
|
||||
)
|
||||
from agentlightning.utils.system_snapshot import system_snapshot
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agentlightning.execution.events import ExecutionEvent
|
||||
@@ -52,7 +66,14 @@ class LitAgentRunner(Runner[T_task]):
|
||||
worker_id: Identifier for the active worker process, if any.
|
||||
"""
|
||||
|
||||
def __init__(self, tracer: Tracer, max_rollouts: Optional[int] = None, poll_interval: float = 5.0) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
tracer: Tracer,
|
||||
max_rollouts: Optional[int] = None,
|
||||
poll_interval: float = 5.0,
|
||||
heartbeat_interval: float = 10.0,
|
||||
heartbeat_launch_mode: Literal["asyncio", "thread"] = "asyncio",
|
||||
) -> None:
|
||||
"""Initialize the agent runner.
|
||||
|
||||
Args:
|
||||
@@ -60,11 +81,16 @@ class LitAgentRunner(Runner[T_task]):
|
||||
max_rollouts: Optional cap on iterations processed by
|
||||
[`iter`][agentlightning.LitAgentRunner.iter].
|
||||
poll_interval: Seconds to wait between store polls when no work is available.
|
||||
heartbeat_interval: Seconds to wait between sending heartbeats to the store.
|
||||
heartbeat_launch_mode: Launch mode for the heartbeat loop. Can be "asyncio" or "thread".
|
||||
"asyncio" is the default and recommended mode. Use "thread" if you are experiencing blocking coroutines.
|
||||
"""
|
||||
super().__init__()
|
||||
self._tracer = tracer
|
||||
self._max_rollouts = max_rollouts
|
||||
self._poll_interval = poll_interval
|
||||
self._heartbeat_interval = heartbeat_interval
|
||||
self._heartbeat_launch_mode = heartbeat_launch_mode
|
||||
|
||||
# Set later
|
||||
self._agent: Optional[LitAgent[T_task]] = None
|
||||
@@ -304,6 +330,67 @@ class LitAgentRunner(Runner[T_task]):
|
||||
|
||||
return trace_spans
|
||||
|
||||
async def _emit_heartbeat(self, store: LightningStore) -> None:
|
||||
"""Send a heartbeat tick to the store."""
|
||||
worker_id = self.get_worker_id()
|
||||
|
||||
try:
|
||||
await store.update_worker(worker_id, system_snapshot())
|
||||
except asyncio.CancelledError:
|
||||
# bypass the exception
|
||||
raise
|
||||
except Exception:
|
||||
logger.exception("%s Unable to update worker heartbeat.", self._log_prefix())
|
||||
|
||||
def _start_heartbeat_loop(self, store: LightningStore) -> Optional[Callable[[], Awaitable[None]]]:
|
||||
"""Start a background heartbeat loop and return an async stopper."""
|
||||
|
||||
if self._heartbeat_interval <= 0:
|
||||
return None
|
||||
|
||||
if self.worker_id is None:
|
||||
logger.warning("%s Cannot start heartbeat loop without worker_id.", self._log_prefix())
|
||||
return None
|
||||
|
||||
if self._heartbeat_launch_mode == "asyncio":
|
||||
stop_event = asyncio.Event()
|
||||
|
||||
async def heartbeat_loop() -> None:
|
||||
while not stop_event.is_set():
|
||||
await self._emit_heartbeat(store)
|
||||
with suppress(asyncio.TimeoutError):
|
||||
await asyncio.wait_for(stop_event.wait(), timeout=self._heartbeat_interval)
|
||||
|
||||
task = asyncio.create_task(heartbeat_loop(), name=f"{self.get_worker_id()}-heartbeat")
|
||||
|
||||
async def stop() -> None:
|
||||
stop_event.set()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
return stop
|
||||
|
||||
if self._heartbeat_launch_mode == "thread":
|
||||
stop_evt = threading.Event()
|
||||
|
||||
def thread_worker() -> None:
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
while not stop_evt.is_set():
|
||||
loop.run_until_complete(self._emit_heartbeat(store))
|
||||
stop_evt.wait(self._heartbeat_interval)
|
||||
|
||||
thread = threading.Thread(target=thread_worker, name=f"{self.get_worker_id()}-heartbeat", daemon=True)
|
||||
thread.start()
|
||||
|
||||
async def stop() -> None:
|
||||
stop_evt.set()
|
||||
await asyncio.to_thread(thread.join)
|
||||
|
||||
return stop
|
||||
|
||||
raise ValueError(f"Unsupported heartbeat launch mode: {self._heartbeat_launch_mode}")
|
||||
|
||||
async def _sleep_until_next_poll(self, event: Optional[ExecutionEvent] = None) -> None:
|
||||
"""Sleep until the next poll interval, with optional event-based interruption.
|
||||
|
||||
@@ -450,39 +537,49 @@ class LitAgentRunner(Runner[T_task]):
|
||||
logger.info(f"{self._log_prefix()} Started async rollouts (max: {self._max_rollouts or 'unlimited'}).")
|
||||
store = self.get_store()
|
||||
|
||||
while not (event is not None and event.is_set()) and (
|
||||
self._max_rollouts is None or num_tasks_processed < self._max_rollouts
|
||||
):
|
||||
# Retrieve the next rollout
|
||||
next_rollout: Optional[Rollout] = None
|
||||
while not (event is not None and event.is_set()):
|
||||
logger.debug(f"{self._log_prefix()} Try to poll for next rollout.")
|
||||
next_rollout = await store.dequeue_rollout()
|
||||
stop_heartbeat = self._start_heartbeat_loop(store)
|
||||
|
||||
try:
|
||||
while not (event is not None and event.is_set()) and (
|
||||
self._max_rollouts is None or num_tasks_processed < self._max_rollouts
|
||||
):
|
||||
# Retrieve the next rollout
|
||||
next_rollout: Optional[Rollout] = None
|
||||
while not (event is not None and event.is_set()):
|
||||
logger.debug(f"{self._log_prefix()} Try to poll for next rollout.")
|
||||
next_rollout = await store.dequeue_rollout(worker_id=self.get_worker_id())
|
||||
if next_rollout is None:
|
||||
logger.debug(
|
||||
f"{self._log_prefix()} No rollout to poll. Waiting for {self._poll_interval} seconds."
|
||||
)
|
||||
await self._sleep_until_next_poll(event)
|
||||
else:
|
||||
break
|
||||
|
||||
if next_rollout is None:
|
||||
logger.debug(f"{self._log_prefix()} No rollout to poll. Waiting for {self._poll_interval} seconds.")
|
||||
await self._sleep_until_next_poll(event)
|
||||
else:
|
||||
break
|
||||
return
|
||||
|
||||
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
|
||||
|
||||
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)
|
||||
|
||||
# Execute the step
|
||||
await self._step_impl(next_rollout)
|
||||
|
||||
num_tasks_processed += 1
|
||||
if num_tasks_processed % 10 == 0 or num_tasks_processed == 1:
|
||||
logger.info(f"{self._log_prefix()} Progress: {num_tasks_processed}/{self._max_rollouts or 'unlimited'}")
|
||||
num_tasks_processed += 1
|
||||
if num_tasks_processed % 10 == 0 or num_tasks_processed == 1:
|
||||
logger.info(
|
||||
f"{self._log_prefix()} Progress: {num_tasks_processed}/{self._max_rollouts or 'unlimited'}"
|
||||
)
|
||||
finally:
|
||||
if stop_heartbeat is not None:
|
||||
await stop_heartbeat()
|
||||
|
||||
logger.info(f"{self._log_prefix()} Finished async rollouts. Processed {num_tasks_processed} tasks.")
|
||||
|
||||
@@ -526,6 +623,12 @@ class LitAgentRunner(Runner[T_task]):
|
||||
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(),
|
||||
)
|
||||
rollout_id = await self._step_impl(attempted_rollout, raise_on_exception=True)
|
||||
|
||||
completed_rollout = await store.get_rollout_by_id(rollout_id)
|
||||
|
||||
@@ -17,6 +17,7 @@ from agentlightning.types import (
|
||||
RolloutStatus,
|
||||
Span,
|
||||
TaskInput,
|
||||
Worker,
|
||||
)
|
||||
|
||||
|
||||
@@ -167,7 +168,7 @@ class LightningStore:
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def dequeue_rollout(self) -> Optional[AttemptedRollout]:
|
||||
async def dequeue_rollout(self, worker_id: Optional[str] = None) -> Optional[AttemptedRollout]:
|
||||
"""Claim the oldest queued rollout and transition it to `preparing`.
|
||||
|
||||
This function do not block.
|
||||
@@ -180,6 +181,8 @@ class LightningStore:
|
||||
the number of attempts already registered for the rollout plus one.
|
||||
* Return an [`AttemptedRollout`][agentlightning.AttemptedRollout] snapshot so the
|
||||
runner knows both rollout metadata and the attempt identifier.
|
||||
* Optionally refresh the caller's [`Worker`][agentlightning.Worker] telemetry
|
||||
(e.g., `last_dequeue_time`) when `worker_id` is provided.
|
||||
|
||||
Returns:
|
||||
The next attempt to execute, or `None` when no eligible rollouts are queued.
|
||||
@@ -527,6 +530,12 @@ class LightningStore:
|
||||
Similar to [`update_rollout()`][agentlightning.LightningStore.update_rollout],
|
||||
parameters also default to the sentinel [`UNSET`][agentlightning.store.base.UNSET].
|
||||
|
||||
If `worker_id` is present, the worker status will be updated following the rules:
|
||||
|
||||
1. If attempt status is "succeeded" or "failed", the corresponding worker status will be set to "idle".
|
||||
2. If attempt status is "unresponsive" or "timeout", the corresponding worker status will be set to "unknown".
|
||||
3. Otherwise, the worker status will be set to "busy".
|
||||
|
||||
Args:
|
||||
rollout_id: Identifier of the rollout whose attempt will be updated.
|
||||
attempt_id: Attempt identifier or `"latest"` as a convenience.
|
||||
@@ -543,3 +552,46 @@ class LightningStore:
|
||||
ValueError: Implementations must raise when the rollout or attempt is unknown.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def query_workers(
|
||||
self,
|
||||
) -> List[Worker]:
|
||||
"""Query all workers in the system.
|
||||
|
||||
Returns:
|
||||
A list of all workers.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def get_worker_by_id(self, worker_id: str) -> Optional[Worker]:
|
||||
"""Retrieve a single worker by identifier.
|
||||
|
||||
Args:
|
||||
worker_id: Identifier of the worker.
|
||||
|
||||
Returns:
|
||||
The worker record if it exists, otherwise `None`.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Subclasses must implement lookup semantics.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def update_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
heartbeat_stats: Dict[str, Any] | Unset = UNSET,
|
||||
) -> Worker:
|
||||
"""Record a heartbeat for `worker_id` and refresh telemetry.
|
||||
|
||||
Implementations must treat this API as heartbeat-only: it should snapshot
|
||||
the latest stats when provided, stamp `last_heartbeat_time` with the
|
||||
current wall clock, and rely on other store mutations (`dequeue_rollout`,
|
||||
`update_attempt`, etc.) to drive the worker's busy/idle status,
|
||||
assignment, and activity timestamps.
|
||||
|
||||
Args:
|
||||
worker_id: Identifier of the worker to update.
|
||||
heartbeat_stats: Replacement worker heartbeat statistics (non-null when provided).
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
@@ -34,6 +34,8 @@ from agentlightning.types import (
|
||||
RolloutStatus,
|
||||
Span,
|
||||
TaskInput,
|
||||
Worker,
|
||||
WorkerStatus,
|
||||
)
|
||||
|
||||
from .base import UNSET, LightningStore, LightningStoreCapabilities, Unset
|
||||
@@ -62,6 +64,10 @@ class RolloutRequest(BaseModel):
|
||||
metadata: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class DequeueRolloutRequest(BaseModel):
|
||||
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))
|
||||
@@ -106,6 +112,10 @@ class UpdateAttemptRequest(BaseModel):
|
||||
metadata: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class UpdateWorkerRequest(BaseModel):
|
||||
heartbeat_stats: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class QueryAttemptsRequest(BaseModel):
|
||||
# Pagination
|
||||
limit: int = -1
|
||||
@@ -148,6 +158,19 @@ class QuerySpansRequest(BaseModel):
|
||||
sort_order: Literal["asc", "desc"] = "asc"
|
||||
|
||||
|
||||
class QueryWorkersRequest(BaseModel):
|
||||
status_in: Optional[List[WorkerStatus]] = Field(FastAPIQuery(default=None))
|
||||
worker_id_contains: Optional[str] = None
|
||||
# Pagination
|
||||
limit: int = -1
|
||||
offset: int = 0
|
||||
# Sorting
|
||||
sort_by: Optional[str] = None
|
||||
sort_order: Literal["asc", "desc"] = "asc"
|
||||
# Filtering logic
|
||||
filter_logic: Literal["and", "or"] = "and"
|
||||
|
||||
|
||||
def _apply_filters_sort_paginate(
|
||||
items: List[T],
|
||||
filters: Dict[str, Any],
|
||||
@@ -600,8 +623,11 @@ class LightningStoreServer(LightningStore):
|
||||
)
|
||||
|
||||
@api.post(API_AGL_PREFIX + "/queues/rollouts/dequeue", response_model=Optional[AttemptedRollout])
|
||||
async def dequeue_rollout(): # pyright: ignore[reportUnusedFunction]
|
||||
return await self.dequeue_rollout()
|
||||
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 + "/rollouts", status_code=201, response_model=AttemptedRollout)
|
||||
async def start_rollout(request: RolloutRequest): # pyright: ignore[reportUnusedFunction]
|
||||
@@ -641,9 +667,11 @@ class LightningStoreServer(LightningStore):
|
||||
async def get_rollout_by_id(rollout_id: str): # pyright: ignore[reportUnusedFunction]
|
||||
return await self.get_rollout_by_id(rollout_id)
|
||||
|
||||
def _get_mandatory_field_or_unset(request: BaseModel, field: str) -> Any:
|
||||
def _get_mandatory_field_or_unset(request: BaseModel | None, field: str) -> Any:
|
||||
# If some fields are mandatory by the underlying store, but optional in the FastAPI,
|
||||
# we make sure it's set to non-null value or UNSET via this function.
|
||||
if request is None:
|
||||
return UNSET
|
||||
if field in request.model_fields_set:
|
||||
value = getattr(request, field)
|
||||
if value is None:
|
||||
@@ -683,6 +711,39 @@ class LightningStoreServer(LightningStore):
|
||||
metadata=_get_mandatory_field_or_unset(request, "metadata"),
|
||||
)
|
||||
|
||||
@api.get(API_AGL_PREFIX + "/workers", response_model=PaginatedResponse[Worker])
|
||||
async def query_workers(params: QueryWorkersRequest = Depends()): # pyright: ignore[reportUnusedFunction]
|
||||
all_workers = await self.query_workers()
|
||||
|
||||
filters: Dict[str, Any] = {}
|
||||
if params.status_in:
|
||||
filters["status_in"] = params.status_in
|
||||
if params.worker_id_contains is not None:
|
||||
filters["worker_id_contains"] = params.worker_id_contains
|
||||
|
||||
return _apply_filters_sort_paginate(
|
||||
all_workers,
|
||||
filters,
|
||||
params.filter_logic,
|
||||
params.sort_by,
|
||||
params.sort_order,
|
||||
params.limit,
|
||||
params.offset,
|
||||
)
|
||||
|
||||
@api.get(API_AGL_PREFIX + "/workers/{worker_id}", response_model=Optional[Worker])
|
||||
async def get_worker(worker_id: str): # pyright: ignore[reportUnusedFunction]
|
||||
return await self.get_worker_by_id(worker_id)
|
||||
|
||||
@api.post(API_AGL_PREFIX + "/workers/{worker_id}", response_model=Worker)
|
||||
async def update_worker( # pyright: ignore[reportUnusedFunction]
|
||||
worker_id: str, request: UpdateWorkerRequest | None = Body(None)
|
||||
):
|
||||
return await self.update_worker(
|
||||
worker_id=worker_id,
|
||||
heartbeat_stats=_get_mandatory_field_or_unset(request, "heartbeat_stats"),
|
||||
)
|
||||
|
||||
@api.get(API_AGL_PREFIX + "/rollouts/{rollout_id}/attempts", response_model=PaginatedResponse[Attempt])
|
||||
async def query_attempts( # pyright: ignore[reportUnusedFunction]
|
||||
rollout_id: str, params: QueryAttemptsRequest = Depends()
|
||||
@@ -893,8 +954,8 @@ class LightningStoreServer(LightningStore):
|
||||
metadata,
|
||||
)
|
||||
|
||||
async def dequeue_rollout(self) -> Optional[AttemptedRollout]:
|
||||
return await self._call_store_method("dequeue_rollout")
|
||||
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)
|
||||
@@ -999,6 +1060,23 @@ class LightningStoreServer(LightningStore):
|
||||
metadata,
|
||||
)
|
||||
|
||||
async def query_workers(self) -> List[Worker]:
|
||||
return await self._call_store_method("query_workers")
|
||||
|
||||
async def get_worker_by_id(self, worker_id: str) -> Optional[Worker]:
|
||||
return await self._call_store_method("get_worker_by_id", worker_id)
|
||||
|
||||
async def update_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
heartbeat_stats: Dict[str, Any] | Unset = UNSET,
|
||||
) -> Worker:
|
||||
return await self._call_store_method(
|
||||
"update_worker",
|
||||
worker_id,
|
||||
heartbeat_stats,
|
||||
)
|
||||
|
||||
|
||||
class LightningStoreClient(LightningStore):
|
||||
"""HTTP client that talks to a remote LightningStoreServer.
|
||||
@@ -1242,7 +1320,7 @@ class LightningStoreClient(LightningStore):
|
||||
)
|
||||
return Rollout.model_validate(data)
|
||||
|
||||
async def dequeue_rollout(self) -> Optional[AttemptedRollout]:
|
||||
async def dequeue_rollout(self, worker_id: Optional[str] = None) -> Optional[AttemptedRollout]:
|
||||
"""
|
||||
Dequeue a rollout from the server queue.
|
||||
|
||||
@@ -1255,8 +1333,11 @@ class LightningStoreClient(LightningStore):
|
||||
"""
|
||||
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) as resp:
|
||||
async with session.post(url, **request_kwargs) as resp:
|
||||
resp.raise_for_status()
|
||||
data = await resp.json()
|
||||
self._dequeue_was_successful = True
|
||||
@@ -1525,3 +1606,27 @@ class LightningStoreClient(LightningStore):
|
||||
json=payload,
|
||||
)
|
||||
return Attempt.model_validate(data)
|
||||
|
||||
async def query_workers(self) -> List[Worker]:
|
||||
data = await self._request_json("get", "/workers")
|
||||
items = data.get("items", [])
|
||||
return [Worker.model_validate(item) for item in items]
|
||||
|
||||
async def get_worker_by_id(self, worker_id: str) -> Optional[Worker]:
|
||||
data = await self._request_json("get", f"/workers/{worker_id}")
|
||||
if data is None:
|
||||
return None
|
||||
return Worker.model_validate(data)
|
||||
|
||||
async def update_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
heartbeat_stats: Dict[str, Any] | Unset = UNSET,
|
||||
) -> Worker:
|
||||
payload: Dict[str, Any] = {}
|
||||
if not isinstance(heartbeat_stats, Unset):
|
||||
payload["heartbeat_stats"] = heartbeat_stats
|
||||
json_payload = payload if payload else None
|
||||
|
||||
data = await self._request_json("post", f"/workers/{worker_id}", json=json_payload)
|
||||
return Worker.model_validate(data)
|
||||
|
||||
@@ -44,6 +44,7 @@ from agentlightning.types import (
|
||||
RolloutStatus,
|
||||
Span,
|
||||
TaskInput,
|
||||
Worker,
|
||||
)
|
||||
|
||||
from .base import UNSET, LightningStore, LightningStoreCapabilities, Unset, is_finished, is_queuing
|
||||
@@ -242,6 +243,48 @@ class InMemoryLightningStore(LightningStore):
|
||||
|
||||
# Completion tracking for wait_for_rollouts (cross-loop safe)
|
||||
self._completion_events: Dict[str, threading.Event] = {}
|
||||
# Worker tracking
|
||||
self._workers: Dict[str, Worker] = {}
|
||||
|
||||
# Running rollouts cache, including preparing and running rollouts
|
||||
self._running_rollout_ids: Set[str] = set()
|
||||
|
||||
def _get_or_create_worker(self, worker_id: str) -> Worker:
|
||||
worker = self._workers.get(worker_id)
|
||||
if worker is None:
|
||||
worker = Worker(worker_id=worker_id)
|
||||
self._workers[worker_id] = worker
|
||||
return worker
|
||||
|
||||
def _sync_worker_with_attempt(self, attempt: Attempt) -> None:
|
||||
worker_id = attempt.worker_id
|
||||
if not worker_id:
|
||||
return
|
||||
|
||||
worker = self._get_or_create_worker(worker_id)
|
||||
now = time.time()
|
||||
|
||||
if attempt.status in ("succeeded", "failed"):
|
||||
if worker.status != "idle":
|
||||
worker.last_idle_time = now
|
||||
worker.status = "idle"
|
||||
worker.current_rollout_id = None
|
||||
worker.current_attempt_id = None
|
||||
elif attempt.status in ("timeout", "unresponsive"):
|
||||
if worker.status != "unknown":
|
||||
worker.last_idle_time = now
|
||||
worker.status = "unknown"
|
||||
worker.current_rollout_id = None
|
||||
worker.current_attempt_id = None
|
||||
else:
|
||||
transitioned = worker.status != "busy" or worker.current_attempt_id != attempt.attempt_id
|
||||
if transitioned:
|
||||
worker.last_busy_time = now
|
||||
worker.status = "busy"
|
||||
worker.current_rollout_id = attempt.rollout_id
|
||||
worker.current_attempt_id = attempt.attempt_id
|
||||
|
||||
Worker.model_validate(worker.model_dump())
|
||||
|
||||
def capabilities(self) -> LightningStoreCapabilities:
|
||||
"""Return the capabilities of the store."""
|
||||
@@ -281,6 +324,7 @@ class InMemoryLightningStore(LightningStore):
|
||||
config=rollout_config,
|
||||
metadata=rollout_metadata,
|
||||
)
|
||||
self._running_rollout_ids.add(rollout.rollout_id)
|
||||
|
||||
# Create the initial attempt
|
||||
attempt_id = _generate_attempt_id()
|
||||
@@ -338,7 +382,7 @@ class InMemoryLightningStore(LightningStore):
|
||||
return rollout
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def dequeue_rollout(self) -> Optional[AttemptedRollout]:
|
||||
async def dequeue_rollout(self, worker_id: Optional[str] = None) -> Optional[AttemptedRollout]:
|
||||
"""Retrieves the next task from the queue without blocking.
|
||||
Returns `None` if the queue is empty.
|
||||
|
||||
@@ -347,6 +391,11 @@ class InMemoryLightningStore(LightningStore):
|
||||
See [`LightningStore.dequeue_rollout()`][agentlightning.LightningStore.dequeue_rollout] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
if worker_id is not None:
|
||||
worker = self._get_or_create_worker(worker_id)
|
||||
worker.last_dequeue_time = time.time()
|
||||
worker.status = "idle"
|
||||
|
||||
# Keep looking until we find a rollout that's still in queuing status
|
||||
# or the queue is empty
|
||||
while self._task_queue:
|
||||
@@ -357,6 +406,7 @@ class InMemoryLightningStore(LightningStore):
|
||||
if is_queuing(rollout):
|
||||
# Update status to preparing
|
||||
rollout.status = "preparing"
|
||||
self._running_rollout_ids.add(rollout.rollout_id)
|
||||
|
||||
# Create a new attempt (could be first attempt or retry)
|
||||
attempt_id = _generate_attempt_id()
|
||||
@@ -655,6 +705,7 @@ class InMemoryLightningStore(LightningStore):
|
||||
if current_attempt == latest_attempt:
|
||||
if rollout.status == "preparing":
|
||||
rollout.status = "running"
|
||||
self._running_rollout_ids.add(rollout.rollout_id)
|
||||
elif rollout.status in ["queuing", "requeuing"]:
|
||||
try:
|
||||
self._task_queue.remove(rollout)
|
||||
@@ -663,6 +714,7 @@ class InMemoryLightningStore(LightningStore):
|
||||
f"Trying to remove rollout {rollout.rollout_id} from the queue but it's not in the queue."
|
||||
)
|
||||
rollout.status = "running"
|
||||
self._running_rollout_ids.add(rollout.rollout_id)
|
||||
|
||||
return span
|
||||
|
||||
@@ -908,6 +960,12 @@ class InMemoryLightningStore(LightningStore):
|
||||
elif is_queuing(rollout) and rollout not in self._task_queue:
|
||||
self._task_queue.append(rollout)
|
||||
|
||||
# Updating running rollouts cache
|
||||
if rollout.status in ["preparing", "running"]:
|
||||
self._running_rollout_ids.add(rollout.rollout_id)
|
||||
else:
|
||||
self._running_rollout_ids.discard(rollout.rollout_id)
|
||||
|
||||
# If the rollout is no longer in a queueing state, remove it from the queue.
|
||||
if not isinstance(status, Unset) and not is_queuing(rollout) and rollout in self._task_queue:
|
||||
try:
|
||||
@@ -951,19 +1009,26 @@ class InMemoryLightningStore(LightningStore):
|
||||
if not attempt:
|
||||
raise ValueError(f"Attempt {attempt_id} not found for rollout {rollout_id}")
|
||||
|
||||
worker_sync_required = False
|
||||
|
||||
# Update fields if they are not UNSET
|
||||
if not isinstance(worker_id, Unset):
|
||||
attempt.worker_id = worker_id
|
||||
worker_sync_required = worker_sync_required or bool(worker_id)
|
||||
if not isinstance(status, Unset):
|
||||
attempt.status = status
|
||||
# Also update end_time if the status indicates completion
|
||||
if status in ["failed", "succeeded"]:
|
||||
attempt.end_time = time.time()
|
||||
if not isinstance(worker_id, Unset):
|
||||
attempt.worker_id = worker_id
|
||||
worker_sync_required = worker_sync_required or bool(attempt.worker_id)
|
||||
if not isinstance(last_heartbeat_time, Unset):
|
||||
attempt.last_heartbeat_time = last_heartbeat_time
|
||||
if not isinstance(metadata, Unset):
|
||||
attempt.metadata = metadata
|
||||
|
||||
if worker_sync_required and attempt.worker_id:
|
||||
self._sync_worker_with_attempt(attempt)
|
||||
|
||||
# Re-validate the attempt to ensure legality
|
||||
Attempt.model_validate(attempt.model_dump())
|
||||
|
||||
@@ -981,12 +1046,40 @@ class InMemoryLightningStore(LightningStore):
|
||||
|
||||
return attempt
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def query_workers(self) -> List[Worker]:
|
||||
"""Return the current snapshot of all workers."""
|
||||
async with self._lock:
|
||||
return list(self._workers.values())
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def get_worker_by_id(self, worker_id: str) -> Optional[Worker]:
|
||||
async with self._lock:
|
||||
return self._workers.get(worker_id)
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def update_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
heartbeat_stats: Dict[str, Any] | Unset = UNSET,
|
||||
) -> Worker:
|
||||
"""Create or update a worker entry."""
|
||||
async with self._lock:
|
||||
worker = self._get_or_create_worker(worker_id)
|
||||
if not isinstance(heartbeat_stats, Unset):
|
||||
worker.heartbeat_stats = dict(heartbeat_stats)
|
||||
worker.last_heartbeat_time = time.time()
|
||||
|
||||
Worker.model_validate(worker.model_dump())
|
||||
return worker
|
||||
|
||||
async def _healthcheck(self) -> None:
|
||||
"""Perform healthcheck against all running rollouts in the store."""
|
||||
async with self._lock:
|
||||
running_rollouts: List[AttemptedRollout] = []
|
||||
for rollout in self._rollouts.values():
|
||||
if rollout.status in ["preparing", "running"]:
|
||||
for rollout_id in self._running_rollout_ids:
|
||||
rollout = self._rollouts.get(rollout_id)
|
||||
if rollout is not None and rollout.status in ["preparing", "running"]:
|
||||
all_attempts = self._attempts.get(rollout.rollout_id, [])
|
||||
if not all_attempts:
|
||||
# The rollout is running but has no attempts, this should not happen
|
||||
|
||||
@@ -18,6 +18,7 @@ from agentlightning.types import (
|
||||
RolloutStatus,
|
||||
Span,
|
||||
TaskInput,
|
||||
Worker,
|
||||
)
|
||||
|
||||
from .base import UNSET, LightningStore, LightningStoreCapabilities, Unset
|
||||
@@ -66,9 +67,9 @@ class LightningStoreThreaded(LightningStore):
|
||||
with self._lock:
|
||||
return await self.store.enqueue_rollout(input, mode, resources_id, config, metadata)
|
||||
|
||||
async def dequeue_rollout(self) -> Optional[AttemptedRollout]:
|
||||
async def dequeue_rollout(self, worker_id: Optional[str] = None) -> Optional[AttemptedRollout]:
|
||||
with self._lock:
|
||||
return await self.store.dequeue_rollout()
|
||||
return await self.store.dequeue_rollout(worker_id=worker_id)
|
||||
|
||||
async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
|
||||
with self._lock:
|
||||
@@ -180,3 +181,22 @@ class LightningStoreThreaded(LightningStore):
|
||||
last_heartbeat_time=last_heartbeat_time,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
async def query_workers(self) -> List[Worker]:
|
||||
with self._lock:
|
||||
return await self.store.query_workers()
|
||||
|
||||
async def get_worker_by_id(self, worker_id: str) -> Optional[Worker]:
|
||||
with self._lock:
|
||||
return await self.store.get_worker_by_id(worker_id)
|
||||
|
||||
async def update_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
heartbeat_stats: Dict[str, Any] | Unset = UNSET,
|
||||
) -> Worker:
|
||||
with self._lock:
|
||||
return await self.store.update_worker(
|
||||
worker_id=worker_id,
|
||||
heartbeat_stats=heartbeat_stats,
|
||||
)
|
||||
|
||||
@@ -49,6 +49,8 @@ __all__ = [
|
||||
"Attempt",
|
||||
"AttemptedRollout",
|
||||
"Hook",
|
||||
"Worker",
|
||||
"WorkerStatus",
|
||||
]
|
||||
|
||||
T_co = TypeVar("T_co", covariant=True)
|
||||
@@ -200,6 +202,32 @@ class AttemptedRollout(Rollout):
|
||||
return self
|
||||
|
||||
|
||||
WorkerStatus = Literal["idle", "busy", "unknown"]
|
||||
|
||||
|
||||
class Worker(BaseModel):
|
||||
"""Worker information. This is actually the same as Runner info."""
|
||||
|
||||
worker_id: str
|
||||
"""The ID of the worker."""
|
||||
status: WorkerStatus = "unknown"
|
||||
"""The status of the worker."""
|
||||
heartbeat_stats: Optional[Dict[str, Any]] = None
|
||||
"""Statistics about the worker's heartbeat."""
|
||||
last_heartbeat_time: Optional[float] = None
|
||||
"""The last time when the worker has reported the stats."""
|
||||
last_dequeue_time: Optional[float] = None
|
||||
"""The last time when the worker has tried to dequeue a rollout."""
|
||||
last_busy_time: Optional[float] = None
|
||||
"""The last time when the worker has started an attempt and became busy."""
|
||||
last_idle_time: Optional[float] = None
|
||||
"""The last time when the worker has triggered the end of an attempt and became idle."""
|
||||
current_rollout_id: Optional[str] = None
|
||||
"""The ID of the current rollout that the worker is processing."""
|
||||
current_attempt_id: Optional[str] = None
|
||||
"""The ID of the current attempt that the worker is processing."""
|
||||
|
||||
|
||||
TaskInput = Any
|
||||
"""Task input type. Accepts arbitrary payloads."""
|
||||
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import platform
|
||||
import socket
|
||||
from contextlib import suppress
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
import psutil
|
||||
from gpustat import GPUStat, GPUStatCollection
|
||||
|
||||
|
||||
def system_snapshot(include_gpu: bool = False) -> Dict[str, Any]:
|
||||
# CPU
|
||||
cpu = {
|
||||
"cpu_name": platform.processor(),
|
||||
"cpu_cores": psutil.cpu_count(logical=False),
|
||||
"cpu_threads": psutil.cpu_count(logical=True),
|
||||
"cpu_usage_pct": psutil.cpu_percent(0.05),
|
||||
}
|
||||
|
||||
# Memory
|
||||
vm = psutil.virtual_memory()
|
||||
mem = {
|
||||
"mem_used_gb": round(vm.used / (2**30), 2),
|
||||
"mem_total_gb": round(vm.total / (2**30), 2),
|
||||
"mem_pct": vm.percent,
|
||||
}
|
||||
|
||||
# Disk
|
||||
du = psutil.disk_usage("/")
|
||||
disk = {
|
||||
"disk_used_gb": round(du.used / (2**30), 2),
|
||||
"disk_total_gb": round(du.total / (2**30), 2),
|
||||
"disk_pct": du.percent,
|
||||
}
|
||||
|
||||
# GPU
|
||||
gpus: List[Dict[str, Any]] = []
|
||||
with suppress(Exception):
|
||||
for g in GPUStatCollection.new_query().gpus: # type: ignore
|
||||
g = cast(GPUStat, g)
|
||||
gpus.append(
|
||||
{
|
||||
"gpu": g.name, # type: ignore
|
||||
"util_pct": g.utilization,
|
||||
"mem_used_mb": g.memory_used,
|
||||
"mem_total_mb": g.memory_total,
|
||||
"temp_c": g.temperature,
|
||||
}
|
||||
)
|
||||
|
||||
# Network
|
||||
net = psutil.net_io_counters()
|
||||
netinfo = {
|
||||
"bytes_sent_mb": round(net.bytes_sent / (2**20), 2),
|
||||
"bytes_recv_mb": round(net.bytes_recv / (2**20), 2),
|
||||
}
|
||||
|
||||
# OS / meta
|
||||
return {
|
||||
"timestamp": datetime.now().isoformat(timespec="seconds"),
|
||||
"host": socket.gethostname(),
|
||||
"os": platform.platform(),
|
||||
**cpu,
|
||||
**mem,
|
||||
**disk,
|
||||
**netinfo,
|
||||
**({"gpus": gpus} if include_gpu else {}),
|
||||
}
|
||||
@@ -6,6 +6,7 @@ import { ResourcesPage } from './pages/Resources.page';
|
||||
import { RolloutsPage } from './pages/Rollouts.page';
|
||||
import { SettingsPage } from './pages/Settings.page';
|
||||
import { TracesPage } from './pages/Traces.page';
|
||||
import { WorkersPage } from './pages/Workers.page';
|
||||
|
||||
const router = createBrowserRouter([
|
||||
{
|
||||
@@ -28,6 +29,10 @@ const router = createBrowserRouter([
|
||||
path: 'traces',
|
||||
element: <TracesPage />,
|
||||
},
|
||||
{
|
||||
path: 'runners',
|
||||
element: <WorkersPage />,
|
||||
},
|
||||
{
|
||||
path: 'settings',
|
||||
element: <SettingsPage />,
|
||||
|
||||
@@ -21,7 +21,7 @@ import {
|
||||
import { useGetSpansQuery } from '@/features/rollouts';
|
||||
import { closeDrawer, openDrawer, selectDrawerContent, selectDrawerIsOpen } from '@/features/ui/drawer';
|
||||
import { useAppDispatch, useAppSelector } from '@/store/hooks';
|
||||
import type { Attempt, AttemptStatus, Rollout, RolloutStatus, Span } from '@/types';
|
||||
import type { Attempt, AttemptStatus, Rollout, RolloutStatus, Span, Worker } from '@/types';
|
||||
import { formatStatusLabel } from '@/utils/format';
|
||||
import { TracesTable, type TracesTableRecord } from './TracesTable.component';
|
||||
|
||||
@@ -50,6 +50,12 @@ const SPAN_STATUS_COLORS: Record<Span['status']['status_code'], string> = {
|
||||
ERROR: 'red',
|
||||
};
|
||||
|
||||
const WORKER_STATUS_COLORS: Record<Worker['status'], string> = {
|
||||
busy: 'orange',
|
||||
idle: 'teal',
|
||||
unknown: 'gray',
|
||||
};
|
||||
|
||||
const TRACES_SORT_FIELD_MAP: Record<string, string> = {
|
||||
name: 'name',
|
||||
traceId: 'trace_id',
|
||||
@@ -408,6 +414,60 @@ function RolloutTracesDrawerBody({ rollout, attempt, onShowRollout, onShowSpanDe
|
||||
);
|
||||
}
|
||||
|
||||
type WorkerDrawerTitleProps = {
|
||||
worker: Worker;
|
||||
};
|
||||
|
||||
function WorkerDrawerTitle({ worker }: WorkerDrawerTitleProps) {
|
||||
const badgeColor = WORKER_STATUS_COLORS[worker.status] ?? 'gray';
|
||||
return (
|
||||
<Stack gap={3}>
|
||||
<Group gap={6} align='center'>
|
||||
<Text fw={600}>{worker.workerId}</Text>
|
||||
<CopyButton value={worker.workerId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy worker ID ${worker.workerId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
<Badge size='sm' variant='light' color={badgeColor}>
|
||||
{formatStatusLabel(worker.status)}
|
||||
</Badge>
|
||||
</Group>
|
||||
<Group gap='xl'>
|
||||
<Group gap={4}>
|
||||
<Text size='sm' c='dimmed' fw={500}>
|
||||
Rollout
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed'>
|
||||
{worker.currentRolloutId ?? '—'}
|
||||
</Text>
|
||||
</Group>
|
||||
<Group gap={4}>
|
||||
<Text size='sm' c='dimmed' fw={500}>
|
||||
Attempt
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed'>
|
||||
{worker.currentAttemptId ?? '—'}
|
||||
</Text>
|
||||
</Group>
|
||||
</Group>
|
||||
</Stack>
|
||||
);
|
||||
}
|
||||
|
||||
export function AppDrawerContainer() {
|
||||
const dispatch = useAppDispatch();
|
||||
const isOpen = useAppSelector(selectDrawerIsOpen);
|
||||
@@ -428,6 +488,13 @@ export function AppDrawerContainer() {
|
||||
return null;
|
||||
}
|
||||
|
||||
if (content.type === 'worker-detail') {
|
||||
const { worker } = content;
|
||||
const title = <WorkerDrawerTitle worker={worker} />;
|
||||
const body = <JsonEditor value={worker} />;
|
||||
return { title, body };
|
||||
}
|
||||
|
||||
if (content.type === 'trace-detail') {
|
||||
const { span } = content;
|
||||
const title = <TraceDrawerTitle span={span} />;
|
||||
|
||||
@@ -22,6 +22,7 @@ const DEFAULT_RECORDS_PER_PAGE_OPTIONS = [50, 100, 200, 500];
|
||||
|
||||
const COLUMN_VISIBILITY: Record<string, ColumnVisibilityConfig> = {
|
||||
name: { minWidth: 12.5, priority: 0 },
|
||||
sequenceId: { fixedWidth: 6, priority: 1 },
|
||||
spanId: { fixedWidth: 14, priority: 1 },
|
||||
traceId: { fixedWidth: 24, priority: 3 },
|
||||
parentId: { fixedWidth: 12, priority: 2 },
|
||||
@@ -86,6 +87,12 @@ function createTracesColumns({
|
||||
</Text>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'sequenceId',
|
||||
title: 'Seq.',
|
||||
sortable: true,
|
||||
render: ({ sequenceId }) => <Text size='sm'>{sequenceId}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'traceId',
|
||||
title: 'Trace ID',
|
||||
|
||||
@@ -0,0 +1,362 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useCallback, useEffect, useMemo } from 'react';
|
||||
import { IconCheck, IconCopy, IconInfoCircle, IconRefresh } from '@tabler/icons-react';
|
||||
import { DataTable, type DataTableColumn, type DataTableSortStatus } from 'mantine-datatable';
|
||||
import { ActionIcon, Badge, Box, Button, CopyButton, Group, Stack, Text, Tooltip } from '@mantine/core';
|
||||
import { useElementSize, useViewportSize } from '@mantine/hooks';
|
||||
import { getLayoutAwareWidth } from '@/layouts/helper';
|
||||
import type { Worker } from '@/types';
|
||||
import { getErrorDescriptor } from '@/utils/error';
|
||||
import { formatDateTime, formatRelativeTime, formatStatusLabel } from '@/utils/format';
|
||||
import { createResponsiveColumns, type ColumnVisibilityConfig } from '@/utils/table';
|
||||
|
||||
const DEFAULT_RECORDS_PER_PAGE_OPTIONS = [50, 100, 200, 500];
|
||||
|
||||
const COLUMN_VISIBILITY: Record<string, ColumnVisibilityConfig> = {
|
||||
workerId: { fixedWidth: 12, priority: 0 },
|
||||
status: { fixedWidth: 6, priority: 1 },
|
||||
currentRolloutId: { fixedWidth: 14, priority: 3 },
|
||||
currentAttemptId: { fixedWidth: 14, priority: 3 },
|
||||
lastHeartbeatTime: { fixedWidth: 10, priority: 2 },
|
||||
lastBusyTime: { fixedWidth: 10, priority: 3 },
|
||||
lastIdleTime: { fixedWidth: 10, priority: 3 },
|
||||
lastDequeueTime: { fixedWidth: 10, priority: 1 },
|
||||
actions: { fixedWidth: 5, priority: 0 },
|
||||
};
|
||||
|
||||
export type WorkersTableRecord = Worker & {
|
||||
timestamps: Record<
|
||||
'lastHeartbeatTime' | 'lastBusyTime' | 'lastIdleTime' | 'lastDequeueTime',
|
||||
{ absolute: string; relative: string }
|
||||
>;
|
||||
};
|
||||
|
||||
const buildTimestampMeta = (value: Worker['lastHeartbeatTime']) => ({
|
||||
absolute: formatDateTime(value),
|
||||
relative: formatRelativeTime(value),
|
||||
});
|
||||
|
||||
function buildWorkerRecord(worker: Worker): WorkersTableRecord {
|
||||
return {
|
||||
...worker,
|
||||
timestamps: {
|
||||
lastHeartbeatTime: buildTimestampMeta(worker.lastHeartbeatTime),
|
||||
lastBusyTime: buildTimestampMeta(worker.lastBusyTime),
|
||||
lastIdleTime: buildTimestampMeta(worker.lastIdleTime),
|
||||
lastDequeueTime: buildTimestampMeta(worker.lastDequeueTime),
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
type WorkersColumnsOptions = {
|
||||
onShowDetails: (worker: Worker) => void;
|
||||
};
|
||||
|
||||
const STATUS_COLORS: Record<Worker['status'], string> = {
|
||||
busy: 'orange',
|
||||
idle: 'teal',
|
||||
unknown: 'gray',
|
||||
};
|
||||
|
||||
function createWorkersColumns({ onShowDetails }: WorkersColumnsOptions): DataTableColumn<WorkersTableRecord>[] {
|
||||
return [
|
||||
{
|
||||
accessor: 'workerId',
|
||||
title: 'Runner ID',
|
||||
sortable: true,
|
||||
render: ({ workerId }) => (
|
||||
<Group gap={2} wrap='nowrap'>
|
||||
<Text fw={500} size='sm'>
|
||||
{workerId}
|
||||
</Text>
|
||||
<CopyButton value={workerId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy worker ID ${workerId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
</Group>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'status',
|
||||
title: 'Status',
|
||||
sortable: true,
|
||||
render: ({ status }) => {
|
||||
const color = STATUS_COLORS[status] ?? 'gray';
|
||||
return (
|
||||
<Badge size='sm' variant='light' color={color} radius='sm'>
|
||||
{formatStatusLabel(status)}
|
||||
</Badge>
|
||||
);
|
||||
},
|
||||
},
|
||||
{
|
||||
accessor: 'currentRolloutId',
|
||||
title: 'Current Rollout',
|
||||
sortable: true,
|
||||
render: ({ currentRolloutId }) => <Text size='sm'>{currentRolloutId ?? '—'}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'currentAttemptId',
|
||||
title: 'Current Attempt',
|
||||
sortable: true,
|
||||
render: ({ currentAttemptId }) => <Text size='sm'>{currentAttemptId ?? '—'}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'lastHeartbeatTime',
|
||||
title: 'Heartbeat',
|
||||
sortable: true,
|
||||
render: ({ timestamps }) => (
|
||||
<Stack gap={0} justify='center'>
|
||||
<Text size='sm'>{timestamps.lastHeartbeatTime.relative}</Text>
|
||||
{timestamps.lastHeartbeatTime.absolute !== '—' && (
|
||||
<Text size='xs' c='dimmed'>
|
||||
{timestamps.lastHeartbeatTime.absolute}
|
||||
</Text>
|
||||
)}
|
||||
</Stack>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'lastBusyTime',
|
||||
title: 'Last Busy',
|
||||
sortable: true,
|
||||
render: ({ timestamps }) => (
|
||||
<Stack gap={0} justify='center'>
|
||||
<Text size='sm'>{timestamps.lastBusyTime.relative}</Text>
|
||||
{timestamps.lastBusyTime.absolute !== '—' && (
|
||||
<Text size='xs' c='dimmed'>
|
||||
{timestamps.lastBusyTime.absolute}
|
||||
</Text>
|
||||
)}
|
||||
</Stack>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'lastIdleTime',
|
||||
title: 'Last Idle',
|
||||
sortable: true,
|
||||
render: ({ timestamps }) => (
|
||||
<Stack gap={0} justify='center'>
|
||||
<Text size='sm'>{timestamps.lastIdleTime.relative}</Text>
|
||||
{timestamps.lastIdleTime.absolute !== '—' && (
|
||||
<Text size='xs' c='dimmed'>
|
||||
{timestamps.lastIdleTime.absolute}
|
||||
</Text>
|
||||
)}
|
||||
</Stack>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'lastDequeueTime',
|
||||
title: 'Last Dequeue',
|
||||
sortable: true,
|
||||
render: ({ timestamps }) => (
|
||||
<Stack gap={0} justify='center'>
|
||||
<Text size='sm'>{timestamps.lastDequeueTime.relative}</Text>
|
||||
{timestamps.lastDequeueTime.absolute !== '—' && (
|
||||
<Text size='xs' c='dimmed'>
|
||||
{timestamps.lastDequeueTime.absolute}
|
||||
</Text>
|
||||
)}
|
||||
</Stack>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'actions',
|
||||
title: 'Actions',
|
||||
textAlign: 'left',
|
||||
render: (record) => (
|
||||
<Tooltip label='Show runner detail' withArrow disabled={!onShowDetails}>
|
||||
<ActionIcon
|
||||
aria-label='Show runner detail'
|
||||
variant='subtle'
|
||||
color='gray'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
onShowDetails(record);
|
||||
}}
|
||||
>
|
||||
<IconInfoCircle size={16} />
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
),
|
||||
},
|
||||
];
|
||||
}
|
||||
|
||||
export type WorkersTableProps = {
|
||||
workers: Worker[] | undefined;
|
||||
totalRecords: number;
|
||||
isFetching: boolean;
|
||||
isError: boolean;
|
||||
error: unknown;
|
||||
searchTerm: string;
|
||||
sort: { column: string; direction: 'asc' | 'desc' };
|
||||
page: number;
|
||||
recordsPerPage: number;
|
||||
onSortStatusChange: (status: DataTableSortStatus<WorkersTableRecord>) => void;
|
||||
onPageChange: (page: number) => void;
|
||||
onRecordsPerPageChange: (value: number) => void;
|
||||
onResetFilters: () => void;
|
||||
onRefetch: () => void;
|
||||
onShowDetails: (worker: Worker) => void;
|
||||
recordsPerPageOptions?: number[];
|
||||
};
|
||||
|
||||
export function WorkersTable({
|
||||
workers,
|
||||
totalRecords,
|
||||
isFetching,
|
||||
isError,
|
||||
error,
|
||||
searchTerm,
|
||||
sort,
|
||||
page,
|
||||
recordsPerPage,
|
||||
onSortStatusChange,
|
||||
onPageChange,
|
||||
onRecordsPerPageChange,
|
||||
onResetFilters,
|
||||
onRefetch,
|
||||
onShowDetails,
|
||||
recordsPerPageOptions = DEFAULT_RECORDS_PER_PAGE_OPTIONS,
|
||||
}: WorkersTableProps) {
|
||||
const { ref: tableContainerRef, width: containerWidth } = useElementSize();
|
||||
const { width: viewportWidth } = useViewportSize();
|
||||
|
||||
const layoutAwareContainerWidth = useMemo(
|
||||
() => getLayoutAwareWidth(containerWidth, viewportWidth),
|
||||
[containerWidth, viewportWidth],
|
||||
);
|
||||
|
||||
const workerRecords = useMemo<WorkersTableRecord[]>(() => {
|
||||
if (!workers) {
|
||||
return [];
|
||||
}
|
||||
return workers.map((worker) => buildWorkerRecord(worker));
|
||||
}, [workers]);
|
||||
|
||||
const columns = useMemo(() => createWorkersColumns({ onShowDetails }), [onShowDetails]);
|
||||
const responsiveColumns = useMemo(
|
||||
() => createResponsiveColumns(columns, layoutAwareContainerWidth, COLUMN_VISIBILITY),
|
||||
[columns, layoutAwareContainerWidth],
|
||||
);
|
||||
|
||||
const totalPages = useMemo(
|
||||
() => Math.max(1, Math.ceil(Math.max(0, totalRecords) / Math.max(1, recordsPerPage))),
|
||||
[recordsPerPage, totalRecords],
|
||||
);
|
||||
|
||||
useEffect(() => {
|
||||
if (page > totalPages) {
|
||||
onPageChange(totalPages);
|
||||
}
|
||||
}, [onPageChange, page, totalPages]);
|
||||
|
||||
const hasActiveFilters = searchTerm.trim().length > 0;
|
||||
|
||||
const sortStatus: DataTableSortStatus<WorkersTableRecord> = {
|
||||
columnAccessor: sort.column,
|
||||
direction: sort.direction,
|
||||
};
|
||||
|
||||
const handleSortStatusChange = useCallback(
|
||||
(status: DataTableSortStatus<WorkersTableRecord>) => {
|
||||
onSortStatusChange(status);
|
||||
},
|
||||
[onSortStatusChange],
|
||||
);
|
||||
|
||||
const errorDescriptor = isError ? getErrorDescriptor(error) : null;
|
||||
const errorMessage = isError
|
||||
? `Workers are temporarily unavailable${errorDescriptor ? ` (${errorDescriptor})` : ''}.`
|
||||
: 'Workers are temporarily unavailable.';
|
||||
|
||||
const emptyState = (
|
||||
<Stack gap='sm' align='center' py='lg'>
|
||||
{isError ? (
|
||||
<>
|
||||
<Text fw={600} size='sm'>
|
||||
{errorMessage}
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
Use the retry button to try again, or adjust the search to broaden the results.
|
||||
</Text>
|
||||
<Group gap='xs'>
|
||||
<Button size='xs' variant='light' color='gray' leftSection={<IconRefresh size={14} />} onClick={onRefetch}>
|
||||
Retry
|
||||
</Button>
|
||||
{hasActiveFilters ? (
|
||||
<Button size='xs' variant='subtle' onClick={onResetFilters}>
|
||||
Clear filters
|
||||
</Button>
|
||||
) : null}
|
||||
</Group>
|
||||
</>
|
||||
) : (
|
||||
<>
|
||||
<Text fw={600} size='sm'>
|
||||
No workers found
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
{hasActiveFilters
|
||||
? 'Try adjusting the search to see more results.'
|
||||
: 'Try refreshing to fetch the latest worker status.'}
|
||||
</Text>
|
||||
<Group gap='xs'>
|
||||
<Button size='xs' variant='light' leftSection={<IconRefresh size={14} />} onClick={onRefetch}>
|
||||
Refresh
|
||||
</Button>
|
||||
{hasActiveFilters ? (
|
||||
<Button size='xs' variant='subtle' onClick={onResetFilters}>
|
||||
Clear filters
|
||||
</Button>
|
||||
) : null}
|
||||
</Group>
|
||||
</>
|
||||
)}
|
||||
</Stack>
|
||||
);
|
||||
|
||||
return (
|
||||
<Box ref={tableContainerRef}>
|
||||
<DataTable<WorkersTableRecord>
|
||||
classNames={{ root: 'workers-table' }}
|
||||
withTableBorder
|
||||
withColumnBorders
|
||||
highlightOnHover
|
||||
verticalAlign='center'
|
||||
minHeight={workerRecords.length === 0 ? 400 : undefined}
|
||||
idAccessor='workerId'
|
||||
records={workerRecords}
|
||||
columns={responsiveColumns}
|
||||
totalRecords={totalRecords}
|
||||
recordsPerPage={recordsPerPage}
|
||||
page={page}
|
||||
onPageChange={onPageChange}
|
||||
onRecordsPerPageChange={onRecordsPerPageChange}
|
||||
recordsPerPageOptions={recordsPerPageOptions}
|
||||
sortStatus={sortStatus}
|
||||
onSortStatusChange={handleSortStatusChange}
|
||||
fetching={isFetching}
|
||||
loaderSize='sm'
|
||||
emptyState={workerRecords.length === 0 ? emptyState : undefined}
|
||||
/>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,154 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useMemo, useState } from 'react';
|
||||
import type { Meta, StoryObj } from '@storybook/react';
|
||||
import { IconSearch } from '@tabler/icons-react';
|
||||
import { Box, Stack, TextInput, Title } from '@mantine/core';
|
||||
import type { Worker } from '@/types';
|
||||
import { WorkersTable } from './WorkersTable.component';
|
||||
|
||||
const meta: Meta<typeof WorkersTable> = {
|
||||
title: 'Components/WorkersTable',
|
||||
component: WorkersTable,
|
||||
parameters: {
|
||||
layout: 'fullscreen',
|
||||
},
|
||||
};
|
||||
|
||||
export default meta;
|
||||
|
||||
type Story = StoryObj<typeof WorkersTable>;
|
||||
|
||||
const now = Math.floor(Date.now() / 1000);
|
||||
|
||||
const sampleWorkers: Worker[] = [
|
||||
{
|
||||
workerId: 'worker-east',
|
||||
status: 'busy',
|
||||
heartbeatStats: { queueDepth: 2, gpuUtilization: 0.82 },
|
||||
lastHeartbeatTime: now - 20,
|
||||
lastDequeueTime: now - 60,
|
||||
lastBusyTime: now - 120,
|
||||
lastIdleTime: now - 600,
|
||||
currentRolloutId: 'ro-story-001',
|
||||
currentAttemptId: 'at-story-010',
|
||||
},
|
||||
{
|
||||
workerId: 'worker-west',
|
||||
status: 'busy',
|
||||
heartbeatStats: { queueDepth: 1 },
|
||||
lastHeartbeatTime: now - 45,
|
||||
lastDequeueTime: now - 300,
|
||||
lastBusyTime: now - 200,
|
||||
lastIdleTime: now - 4800,
|
||||
currentRolloutId: 'ro-story-003',
|
||||
currentAttemptId: 'at-story-033',
|
||||
},
|
||||
{
|
||||
workerId: 'worker-north',
|
||||
status: 'idle',
|
||||
heartbeatStats: { queueDepth: 0 },
|
||||
lastHeartbeatTime: now - 90,
|
||||
lastDequeueTime: now - 3600,
|
||||
lastBusyTime: now - 5400,
|
||||
lastIdleTime: now - 5400,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
{
|
||||
workerId: 'worker-south',
|
||||
status: 'idle',
|
||||
heartbeatStats: null,
|
||||
lastHeartbeatTime: now - 900,
|
||||
lastDequeueTime: now - 7200,
|
||||
lastBusyTime: now - 8600,
|
||||
lastIdleTime: now - 8600,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
{
|
||||
workerId: 'worker-standby',
|
||||
status: 'unknown',
|
||||
heartbeatStats: { queueDepth: 0 },
|
||||
lastHeartbeatTime: now - 15,
|
||||
lastDequeueTime: now - 4000,
|
||||
lastBusyTime: null,
|
||||
lastIdleTime: null,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
];
|
||||
|
||||
type WorkersTableStoryWrapperProps = {
|
||||
maxWidth: number;
|
||||
initialSort?: { column: string; direction: 'asc' | 'desc' };
|
||||
};
|
||||
|
||||
function WorkersTableStoryWrapper({ maxWidth, initialSort }: WorkersTableStoryWrapperProps) {
|
||||
const [searchTerm, setSearchTerm] = useState('');
|
||||
const [page, setPage] = useState(1);
|
||||
const [recordsPerPage, setRecordsPerPage] = useState(5);
|
||||
const [sort, setSort] = useState<{ column: string; direction: 'asc' | 'desc' }>(
|
||||
() => initialSort ?? { column: 'lastHeartbeatTime', direction: 'desc' },
|
||||
);
|
||||
|
||||
const filteredWorkers = useMemo(() => {
|
||||
const normalized = searchTerm.trim().toLowerCase();
|
||||
if (!normalized) {
|
||||
return sampleWorkers;
|
||||
}
|
||||
return sampleWorkers.filter((worker) => worker.workerId.toLowerCase().includes(normalized));
|
||||
}, [searchTerm]);
|
||||
|
||||
return (
|
||||
<Stack gap='md' p='lg'>
|
||||
<Title order={2}>Workers ({maxWidth}px max width)</Title>
|
||||
<TextInput
|
||||
placeholder='Search'
|
||||
leftSection={<IconSearch size={16} />}
|
||||
value={searchTerm}
|
||||
onChange={(event) => setSearchTerm(event.currentTarget.value)}
|
||||
w='100%'
|
||||
style={{ maxWidth: 360 }}
|
||||
/>
|
||||
<Box style={{ maxWidth }}>
|
||||
<WorkersTable
|
||||
workers={filteredWorkers}
|
||||
totalRecords={filteredWorkers.length}
|
||||
isFetching={false}
|
||||
isError={false}
|
||||
error={null}
|
||||
searchTerm={searchTerm}
|
||||
sort={sort}
|
||||
page={page}
|
||||
recordsPerPage={recordsPerPage}
|
||||
onSortStatusChange={(status) => {
|
||||
return setSort({ column: status.columnAccessor as string, direction: status.direction });
|
||||
}}
|
||||
onPageChange={setPage}
|
||||
onRecordsPerPageChange={setRecordsPerPage}
|
||||
onResetFilters={() => {
|
||||
setSearchTerm('');
|
||||
setPage(1);
|
||||
}}
|
||||
onRefetch={() => {}}
|
||||
onShowDetails={() => {}}
|
||||
/>
|
||||
</Box>
|
||||
</Stack>
|
||||
);
|
||||
}
|
||||
|
||||
export const Wide: Story = {
|
||||
render: () => <WorkersTableStoryWrapper maxWidth={1600} />,
|
||||
};
|
||||
|
||||
export const Narrow: Story = {
|
||||
render: () => <WorkersTableStoryWrapper maxWidth={780} />,
|
||||
};
|
||||
|
||||
export const SortedByCurrentRollout: Story = {
|
||||
render: () => (
|
||||
<WorkersTableStoryWrapper maxWidth={1200} initialSort={{ column: 'currentRolloutId', direction: 'asc' }} />
|
||||
),
|
||||
};
|
||||
@@ -13,6 +13,8 @@ import type {
|
||||
RolloutStatus,
|
||||
Span,
|
||||
Timestamp,
|
||||
Worker,
|
||||
WorkerStatus,
|
||||
} from '../../types';
|
||||
|
||||
const rawBaseQuery = fetchBaseQuery({
|
||||
@@ -122,6 +124,21 @@ const normalizeResources = (value: unknown): Resources => {
|
||||
};
|
||||
};
|
||||
|
||||
const normalizeWorker = (value: unknown): Worker => {
|
||||
const camelized = camelCaseKeys(value) as Worker;
|
||||
return {
|
||||
workerId: camelized.workerId,
|
||||
status: camelized.status,
|
||||
heartbeatStats: camelized.heartbeatStats ?? null,
|
||||
lastHeartbeatTime: camelized.lastHeartbeatTime ?? null,
|
||||
lastDequeueTime: camelized.lastDequeueTime ?? null,
|
||||
lastBusyTime: camelized.lastBusyTime ?? null,
|
||||
lastIdleTime: camelized.lastIdleTime ?? null,
|
||||
currentRolloutId: camelized.currentRolloutId ?? null,
|
||||
currentAttemptId: camelized.currentAttemptId ?? null,
|
||||
};
|
||||
};
|
||||
|
||||
const normalizePaginatedResponse = <T>(value: unknown, normalizer: (item: unknown) => T): PaginatedResponse<T> => {
|
||||
if (!value || typeof value !== 'object') {
|
||||
throw new Error('Expected paginated response payload');
|
||||
@@ -183,6 +200,15 @@ export type GetResourcesQueryArgs = {
|
||||
resourcesIdContains?: string | null;
|
||||
};
|
||||
|
||||
export type GetWorkersQueryArgs = {
|
||||
limit: number;
|
||||
offset: number;
|
||||
sortBy?: string | null;
|
||||
sortOrder?: 'asc' | 'desc';
|
||||
workerIdContains?: string | null;
|
||||
statusIn?: WorkerStatus[];
|
||||
};
|
||||
|
||||
export type GetRolloutAttemptsQueryArgs = {
|
||||
rolloutId: string;
|
||||
limit?: number;
|
||||
@@ -208,7 +234,7 @@ export type GetSpansQueryArgs = {
|
||||
export const rolloutsApi = createApi({
|
||||
reducerPath: 'rolloutsApi',
|
||||
baseQuery: dynamicBaseQuery,
|
||||
tagTypes: ['Rollout', 'Span', 'Resources'],
|
||||
tagTypes: ['Rollout', 'Span', 'Resources', 'Worker'],
|
||||
endpoints: (builder) => ({
|
||||
getResources: builder.query<PaginatedResponse<Resources>, GetResourcesQueryArgs>({
|
||||
query: ({ limit, offset, sortBy, sortOrder, resourcesIdContains }) => {
|
||||
@@ -238,6 +264,37 @@ export const rolloutsApi = createApi({
|
||||
]
|
||||
: [{ type: 'Resources' as const, id: 'LIST' }],
|
||||
}),
|
||||
getWorkers: builder.query<PaginatedResponse<Worker>, GetWorkersQueryArgs>({
|
||||
query: ({ limit, offset, sortBy, sortOrder, workerIdContains, statusIn }) => {
|
||||
const searchParams = new URLSearchParams();
|
||||
searchParams.set('limit', String(typeof limit === 'number' ? limit : -1));
|
||||
searchParams.set('offset', String(typeof offset === 'number' ? offset : 0));
|
||||
if (sortBy) {
|
||||
searchParams.set('sort_by', sortBy);
|
||||
}
|
||||
if (sortOrder) {
|
||||
searchParams.set('sort_order', sortOrder);
|
||||
}
|
||||
if (workerIdContains && workerIdContains.trim().length > 0) {
|
||||
searchParams.set('worker_id_contains', workerIdContains.trim());
|
||||
}
|
||||
if (statusIn && statusIn.length > 0) {
|
||||
statusIn.forEach((status) => searchParams.append('status_in', status));
|
||||
}
|
||||
|
||||
const queryString = searchParams.toString();
|
||||
const url = queryString.length > 0 ? `v1/agl/workers?${queryString}` : 'v1/agl/workers';
|
||||
return { url, method: 'GET' };
|
||||
},
|
||||
transformResponse: (response: unknown) => normalizePaginatedResponse(response, normalizeWorker),
|
||||
providesTags: (result) =>
|
||||
result
|
||||
? [
|
||||
{ type: 'Worker' as const, id: 'LIST' },
|
||||
...result.items.map((worker) => ({ type: 'Worker' as const, id: worker.workerId })),
|
||||
]
|
||||
: [{ type: 'Worker' as const, id: 'LIST' }],
|
||||
}),
|
||||
getRollouts: builder.query<PaginatedResponse<Rollout>, GetRolloutsQueryArgs>({
|
||||
query: ({ limit, offset, sortBy, sortOrder, statusIn, rolloutIdContains, modeIn }) => {
|
||||
const searchParams = new URLSearchParams();
|
||||
@@ -343,4 +400,10 @@ export const rolloutsApi = createApi({
|
||||
}),
|
||||
});
|
||||
|
||||
export const { useGetResourcesQuery, useGetRolloutsQuery, useGetRolloutAttemptsQuery, useGetSpansQuery } = rolloutsApi;
|
||||
export const {
|
||||
useGetResourcesQuery,
|
||||
useGetWorkersQuery,
|
||||
useGetRolloutsQuery,
|
||||
useGetRolloutAttemptsQuery,
|
||||
useGetSpansQuery,
|
||||
} = rolloutsApi;
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { createSlice, type PayloadAction } from '@reduxjs/toolkit';
|
||||
import type { Attempt, Rollout, Span } from '@/types';
|
||||
import type { Attempt, Rollout, Span, Worker } from '@/types';
|
||||
|
||||
export type DrawerType = 'rollout-json' | 'rollout-traces' | 'trace-detail';
|
||||
export type DrawerType = 'rollout-json' | 'rollout-traces' | 'trace-detail' | 'worker-detail';
|
||||
|
||||
export type DrawerContent =
|
||||
| {
|
||||
@@ -17,6 +17,10 @@ export type DrawerContent =
|
||||
span: Span;
|
||||
rollout: Rollout | null;
|
||||
attempt: Attempt | null;
|
||||
}
|
||||
| {
|
||||
type: 'worker-detail';
|
||||
worker: Worker;
|
||||
};
|
||||
|
||||
export type DrawerState = {
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
export * from './slice';
|
||||
export * from './selectors';
|
||||
export { useGetWorkersQuery } from '../rollouts';
|
||||
@@ -0,0 +1,45 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { createSelector } from '@reduxjs/toolkit';
|
||||
import type { GetWorkersQueryArgs } from '@/features/rollouts';
|
||||
import type { RootState } from '@/store';
|
||||
import type { WorkersSortState } from './slice';
|
||||
|
||||
const WORKERS_SORT_FIELD_MAP: Record<string, string> = {
|
||||
workerId: 'worker_id',
|
||||
status: 'status',
|
||||
currentRolloutId: 'current_rollout_id',
|
||||
currentAttemptId: 'current_attempt_id',
|
||||
lastHeartbeatTime: 'last_heartbeat_time',
|
||||
lastDequeueTime: 'last_dequeue_time',
|
||||
lastBusyTime: 'last_busy_time',
|
||||
lastIdleTime: 'last_idle_time',
|
||||
};
|
||||
|
||||
const resolveWorkersSortField = (sort: WorkersSortState): string =>
|
||||
WORKERS_SORT_FIELD_MAP[sort.column] ?? 'last_heartbeat_time';
|
||||
|
||||
export const selectWorkersUiState = (state: RootState) => state.workers;
|
||||
|
||||
export const selectWorkersSearchTerm = (state: RootState) => selectWorkersUiState(state).searchTerm;
|
||||
export const selectWorkersPage = (state: RootState) => selectWorkersUiState(state).page;
|
||||
export const selectWorkersRecordsPerPage = (state: RootState) => selectWorkersUiState(state).recordsPerPage;
|
||||
export const selectWorkersSort = (state: RootState) => selectWorkersUiState(state).sort;
|
||||
|
||||
export const selectWorkersQueryArgs = createSelector(
|
||||
[selectWorkersSearchTerm, selectWorkersPage, selectWorkersRecordsPerPage, selectWorkersSort],
|
||||
(searchTerm, page, recordsPerPage, sort): GetWorkersQueryArgs => {
|
||||
const normalizedSearch = searchTerm.trim();
|
||||
const limit = Math.max(1, recordsPerPage);
|
||||
const offset = Math.max(0, (page - 1) * limit);
|
||||
const sortBy = resolveWorkersSortField(sort);
|
||||
|
||||
return {
|
||||
limit,
|
||||
offset,
|
||||
sortBy,
|
||||
sortOrder: sort.direction,
|
||||
workerIdContains: normalizedSearch.length > 0 ? normalizedSearch : undefined,
|
||||
};
|
||||
},
|
||||
);
|
||||
@@ -0,0 +1,59 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { createSlice, type PayloadAction } from '@reduxjs/toolkit';
|
||||
|
||||
export type SortDirection = 'asc' | 'desc';
|
||||
|
||||
export type WorkersSortState = {
|
||||
column: string;
|
||||
direction: SortDirection;
|
||||
};
|
||||
|
||||
export type WorkersUiState = {
|
||||
searchTerm: string;
|
||||
page: number;
|
||||
recordsPerPage: number;
|
||||
sort: WorkersSortState;
|
||||
};
|
||||
|
||||
export const initialWorkersUiState: WorkersUiState = {
|
||||
searchTerm: '',
|
||||
page: 1,
|
||||
recordsPerPage: 50,
|
||||
sort: {
|
||||
column: 'lastHeartbeatTime',
|
||||
direction: 'desc',
|
||||
},
|
||||
};
|
||||
|
||||
const workersSlice = createSlice({
|
||||
name: 'workers',
|
||||
initialState: initialWorkersUiState,
|
||||
reducers: {
|
||||
setWorkersSearchTerm(state, action: PayloadAction<string>) {
|
||||
state.searchTerm = action.payload;
|
||||
state.page = 1;
|
||||
},
|
||||
setWorkersPage(state, action: PayloadAction<number>) {
|
||||
state.page = action.payload;
|
||||
},
|
||||
setWorkersRecordsPerPage(state, action: PayloadAction<number>) {
|
||||
state.recordsPerPage = action.payload;
|
||||
state.page = 1;
|
||||
},
|
||||
setWorkersSort(state, action: PayloadAction<WorkersSortState>) {
|
||||
state.sort = action.payload;
|
||||
},
|
||||
resetWorkersFilters(state) {
|
||||
state.searchTerm = initialWorkersUiState.searchTerm;
|
||||
state.page = initialWorkersUiState.page;
|
||||
state.recordsPerPage = initialWorkersUiState.recordsPerPage;
|
||||
state.sort = initialWorkersUiState.sort;
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
export const { setWorkersSearchTerm, setWorkersPage, setWorkersRecordsPerPage, setWorkersSort, resetWorkersFilters } =
|
||||
workersSlice.actions;
|
||||
|
||||
export const workersReducer = workersSlice.reducer;
|
||||
@@ -0,0 +1,102 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { createServerBackedStore } from '@test-utils';
|
||||
import { describe, expect, it } from 'vitest';
|
||||
import { rolloutsApi } from '@/features/rollouts';
|
||||
import type { Worker } from '@/types';
|
||||
import { selectWorkersQueryArgs } from './selectors';
|
||||
import {
|
||||
resetWorkersFilters,
|
||||
setWorkersPage,
|
||||
setWorkersRecordsPerPage,
|
||||
setWorkersSearchTerm,
|
||||
setWorkersSort,
|
||||
} from './slice';
|
||||
|
||||
const extractWorkerIds = (workers: Worker[]): string[] => workers.map((worker) => worker.workerId);
|
||||
|
||||
describe('workers feature integration', () => {
|
||||
it('builds default query arguments from the UI state', () => {
|
||||
const store = createServerBackedStore();
|
||||
const queryArgs = selectWorkersQueryArgs(store.getState());
|
||||
|
||||
expect(queryArgs).toMatchObject({
|
||||
limit: 50,
|
||||
offset: 0,
|
||||
sortBy: 'last_heartbeat_time',
|
||||
sortOrder: 'desc',
|
||||
workerIdContains: undefined,
|
||||
});
|
||||
});
|
||||
|
||||
it('fetches workers from the Python LightningStore server', async () => {
|
||||
const store = createServerBackedStore();
|
||||
const queryArgs = selectWorkersQueryArgs(store.getState());
|
||||
|
||||
const subscription = store.dispatch(rolloutsApi.endpoints.getWorkers.initiate(queryArgs));
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
expect(data.total).toBeGreaterThanOrEqual(4);
|
||||
expect(data.items).toHaveLength(Math.min(queryArgs.limit, data.total));
|
||||
|
||||
const workerIds = extractWorkerIds(data.items);
|
||||
expect(workerIds).toEqual(expect.arrayContaining(['worker-east', 'worker-west']));
|
||||
|
||||
const heartbeatTimes = data.items.map((worker) => worker.lastHeartbeatTime ?? 0);
|
||||
const sortedHeartbeatTimes = [...heartbeatTimes].sort((a, b) => b - a);
|
||||
expect(heartbeatTimes).toEqual(sortedHeartbeatTimes);
|
||||
});
|
||||
|
||||
it('paginates worker results based on UI state', async () => {
|
||||
const store = createServerBackedStore();
|
||||
store.dispatch(setWorkersRecordsPerPage(2));
|
||||
store.dispatch(setWorkersPage(2));
|
||||
|
||||
const queryArgs = selectWorkersQueryArgs(store.getState());
|
||||
expect(queryArgs).toMatchObject({ limit: 2, offset: 2 });
|
||||
|
||||
const subscription = store.dispatch(rolloutsApi.endpoints.getWorkers.initiate(queryArgs));
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
expect(data.items).toHaveLength(2);
|
||||
expect(data.total).toBeGreaterThanOrEqual(4);
|
||||
});
|
||||
|
||||
it('applies search and sorting preferences', async () => {
|
||||
const store = createServerBackedStore();
|
||||
store.dispatch(resetWorkersFilters());
|
||||
store.dispatch(setWorkersSearchTerm('worker-west'));
|
||||
store.dispatch(setWorkersSort({ column: 'workerId', direction: 'asc' }));
|
||||
|
||||
const queryArgs = selectWorkersQueryArgs(store.getState());
|
||||
expect(queryArgs).toMatchObject({
|
||||
limit: 50,
|
||||
offset: 0,
|
||||
sortBy: 'worker_id',
|
||||
sortOrder: 'asc',
|
||||
workerIdContains: 'worker-west',
|
||||
});
|
||||
|
||||
const subscription = store.dispatch(rolloutsApi.endpoints.getWorkers.initiate(queryArgs));
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
expect(data.items).toHaveLength(1);
|
||||
expect(data.items[0].workerId).toBe('worker-west');
|
||||
});
|
||||
|
||||
it('maps current rollout/attempt sorting to backend fields', () => {
|
||||
const store = createServerBackedStore();
|
||||
store.dispatch(setWorkersSort({ column: 'currentRolloutId', direction: 'desc' }));
|
||||
let queryArgs = selectWorkersQueryArgs(store.getState());
|
||||
expect(queryArgs.sortBy).toBe('current_rollout_id');
|
||||
expect(queryArgs.sortOrder).toBe('desc');
|
||||
|
||||
store.dispatch(setWorkersSort({ column: 'currentAttemptId', direction: 'asc' }));
|
||||
queryArgs = selectWorkersQueryArgs(store.getState());
|
||||
expect(queryArgs.sortBy).toBe('current_attempt_id');
|
||||
expect(queryArgs.sortOrder).toBe('asc');
|
||||
});
|
||||
});
|
||||
@@ -39,6 +39,10 @@ const ROUTES = [
|
||||
path: 'traces',
|
||||
element: <Placeholder title='Traces' description='Browse telemetry spans across attempts.' />,
|
||||
},
|
||||
{
|
||||
path: 'runners',
|
||||
element: <Placeholder title='Runners' description='Monitor runner activity and status.' />,
|
||||
},
|
||||
{
|
||||
path: 'settings',
|
||||
element: (
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useEffect, useMemo, useState, type ReactNode } from 'react';
|
||||
import { IconCpu, IconRouteSquare, IconSettings, IconTimeline } from '@tabler/icons-react';
|
||||
import { IconCpu, IconRouteSquare, IconRun, IconSettings, IconTimeline } from '@tabler/icons-react';
|
||||
import { Outlet, NavLink as RouterNavLink, useLocation, useNavigate } from 'react-router-dom';
|
||||
import { AppShell, Badge, Group, Image, NavLink as MantineNavLink, Stack, Text, UnstyledButton } from '@mantine/core';
|
||||
import { AppAlertBanner } from '@/components/AppAlertBanner';
|
||||
@@ -22,6 +22,7 @@ const NAV_ITEMS: NavItem[] = [
|
||||
{ label: 'Rollouts', to: '/rollouts', icon: <IconRouteSquare size={16} /> },
|
||||
{ label: 'Resources', to: '/resources', icon: <IconCpu size={16} /> },
|
||||
{ label: 'Traces', to: '/traces', icon: <IconTimeline size={16} /> },
|
||||
{ label: 'Runners', to: '/runners', icon: <IconRun size={16} /> },
|
||||
{ label: 'Settings', to: '/settings', icon: <IconSettings size={16} /> },
|
||||
];
|
||||
|
||||
|
||||
@@ -0,0 +1,270 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import type { Meta, StoryObj } from '@storybook/react';
|
||||
import { waitFor, within } from '@testing-library/dom';
|
||||
import userEvent from '@testing-library/user-event';
|
||||
import { Provider } from 'react-redux';
|
||||
import { createMemoryRouter, RouterProvider } from 'react-router-dom';
|
||||
import { AppAlertBanner } from '@/components/AppAlertBanner';
|
||||
import { AppDrawerContainer } from '@/components/AppDrawer.component';
|
||||
import { initialConfigState } from '@/features/config/slice';
|
||||
import { initialWorkersUiState } from '@/features/workers/slice';
|
||||
import { AppLayout } from '@/layouts/AppLayout';
|
||||
import { createAppStore } from '@/store';
|
||||
import type { Worker } from '@/types';
|
||||
import { createWorkersHandlers } from '@/utils/mock';
|
||||
import { STORY_BASE_URL, STORY_DATE_NOW_SECONDS } from '../../.storybook/constants';
|
||||
import { allModes } from '../../.storybook/modes';
|
||||
import { WorkersPage } from './Workers.page';
|
||||
|
||||
const meta: Meta<typeof WorkersPage> = {
|
||||
title: 'Pages/WorkersPage',
|
||||
component: WorkersPage,
|
||||
parameters: {
|
||||
layout: 'fullscreen',
|
||||
chromatic: {
|
||||
modes: allModes,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export default meta;
|
||||
|
||||
type Story = StoryObj<typeof WorkersPage>;
|
||||
|
||||
const now = STORY_DATE_NOW_SECONDS;
|
||||
|
||||
const sampleWorkers: Worker[] = [
|
||||
{
|
||||
workerId: 'worker-east',
|
||||
status: 'busy',
|
||||
heartbeatStats: { queueDepth: 2, gpuUtilization: 0.82 },
|
||||
lastHeartbeatTime: now - 20,
|
||||
lastDequeueTime: now - 120,
|
||||
lastBusyTime: now - 60,
|
||||
lastIdleTime: now - 600,
|
||||
currentRolloutId: 'ro-story-001',
|
||||
currentAttemptId: 'at-story-010',
|
||||
},
|
||||
{
|
||||
workerId: 'worker-west',
|
||||
status: 'busy',
|
||||
heartbeatStats: { queueDepth: 1 },
|
||||
lastHeartbeatTime: now - 45,
|
||||
lastDequeueTime: now - 300,
|
||||
lastBusyTime: now - 120,
|
||||
lastIdleTime: now - 4800,
|
||||
currentRolloutId: 'ro-story-003',
|
||||
currentAttemptId: 'at-story-033',
|
||||
},
|
||||
{
|
||||
workerId: 'worker-north',
|
||||
status: 'idle',
|
||||
heartbeatStats: { queueDepth: 0 },
|
||||
lastHeartbeatTime: now - 120,
|
||||
lastDequeueTime: now - 3600,
|
||||
lastBusyTime: now - 5400,
|
||||
lastIdleTime: now - 180,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
{
|
||||
workerId: 'worker-south',
|
||||
status: 'idle',
|
||||
heartbeatStats: null,
|
||||
lastHeartbeatTime: now - 900,
|
||||
lastDequeueTime: now - 7200,
|
||||
lastBusyTime: now - 8600,
|
||||
lastIdleTime: now - 8600,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
{
|
||||
workerId: 'worker-central',
|
||||
status: 'busy',
|
||||
heartbeatStats: { queueDepth: 3, cpuUtilization: 0.55 },
|
||||
lastHeartbeatTime: now - 8,
|
||||
lastDequeueTime: now - 45,
|
||||
lastBusyTime: now - 10,
|
||||
lastIdleTime: now - 900,
|
||||
currentRolloutId: 'ro-story-005',
|
||||
currentAttemptId: 'at-story-013',
|
||||
},
|
||||
{
|
||||
workerId: 'worker-standby',
|
||||
status: 'idle',
|
||||
heartbeatStats: { queueDepth: 0, threads: 32 },
|
||||
lastHeartbeatTime: now - 300,
|
||||
lastDequeueTime: now - 10800,
|
||||
lastBusyTime: now - 14400,
|
||||
lastIdleTime: now - 200,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
{
|
||||
workerId: 'worker-observer',
|
||||
status: 'unknown',
|
||||
heartbeatStats: { queueDepth: 0 },
|
||||
lastHeartbeatTime: now - 30,
|
||||
lastDequeueTime: now - 6400,
|
||||
lastBusyTime: null,
|
||||
lastIdleTime: null,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
];
|
||||
|
||||
const defaultHandlers = createWorkersHandlers(sampleWorkers);
|
||||
|
||||
function createStoryStore(configOverrides?: Partial<typeof initialConfigState>) {
|
||||
return createAppStore({
|
||||
config: {
|
||||
...initialConfigState,
|
||||
baseUrl: STORY_BASE_URL,
|
||||
autoRefreshMs: 0,
|
||||
...configOverrides,
|
||||
},
|
||||
workers: initialWorkersUiState,
|
||||
});
|
||||
}
|
||||
|
||||
function renderWithStore(configOverrides?: Partial<typeof initialConfigState>) {
|
||||
const store = createStoryStore(configOverrides);
|
||||
|
||||
return (
|
||||
<Provider store={store}>
|
||||
<>
|
||||
<WorkersPage />
|
||||
<AppAlertBanner />
|
||||
<AppDrawerContainer />
|
||||
</>
|
||||
</Provider>
|
||||
);
|
||||
}
|
||||
|
||||
function renderWithinAppLayout(configOverrides?: Partial<typeof initialConfigState>) {
|
||||
const store = createStoryStore(configOverrides);
|
||||
const router = createMemoryRouter(
|
||||
[
|
||||
{
|
||||
path: '/',
|
||||
element: (
|
||||
<AppLayout
|
||||
config={{
|
||||
baseUrl: store.getState().config.baseUrl,
|
||||
autoRefreshMs: store.getState().config.autoRefreshMs,
|
||||
}}
|
||||
/>
|
||||
),
|
||||
children: [
|
||||
{
|
||||
path: '/runners',
|
||||
element: <WorkersPage />,
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
{ initialEntries: ['/runners'] },
|
||||
);
|
||||
|
||||
return (
|
||||
<Provider store={store}>
|
||||
<>
|
||||
<RouterProvider router={router} />
|
||||
<AppDrawerContainer />
|
||||
</>
|
||||
</Provider>
|
||||
);
|
||||
}
|
||||
|
||||
const manyWorkers = Array.from({ length: 80 }, (_, index) => {
|
||||
const suffix = (index + 1).toString().padStart(3, '0');
|
||||
const busy = index % 2 === 0;
|
||||
return {
|
||||
workerId: `worker-batch-${suffix}`,
|
||||
status: busy ? 'busy' : 'idle',
|
||||
heartbeatStats: busy ? { queueDepth: (index % 5) + 1 } : { queueDepth: 0 },
|
||||
lastHeartbeatTime: now - (index * 5 + 15),
|
||||
lastDequeueTime: now - (index * 20 + 60),
|
||||
lastBusyTime: busy ? now - (index * 10 + 30) : null,
|
||||
lastIdleTime: busy ? null : now - (index * 10 + 45),
|
||||
currentRolloutId: busy ? `ro-many-${suffix}` : null,
|
||||
currentAttemptId: busy ? `at-many-${suffix}` : null,
|
||||
} satisfies Worker;
|
||||
});
|
||||
|
||||
export const Default: Story = {
|
||||
render: () => renderWithinAppLayout(),
|
||||
parameters: {
|
||||
msw: {
|
||||
handlers: defaultHandlers,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export const Search: Story = {
|
||||
render: () => renderWithStore(),
|
||||
parameters: {
|
||||
msw: {
|
||||
handlers: defaultHandlers,
|
||||
},
|
||||
},
|
||||
play: async ({ canvasElement }) => {
|
||||
const canvas = within(canvasElement);
|
||||
await canvas.findByText('worker-east');
|
||||
|
||||
const searchInput = canvas.getByPlaceholderText('Search by Runner ID');
|
||||
await userEvent.type(searchInput, 'worker-west');
|
||||
|
||||
await waitFor(() => {
|
||||
if (canvas.queryByText('worker-east')) {
|
||||
throw new Error('Expected filtered table to hide worker-east');
|
||||
}
|
||||
if (!canvas.queryByText('worker-west')) {
|
||||
throw new Error('Expected worker-west to remain visible');
|
||||
}
|
||||
});
|
||||
},
|
||||
};
|
||||
|
||||
export const DrawerOpen: Story = {
|
||||
render: () => renderWithStore(),
|
||||
parameters: {
|
||||
msw: {
|
||||
handlers: defaultHandlers,
|
||||
},
|
||||
},
|
||||
play: async ({ canvasElement }) => {
|
||||
const canvas = within(canvasElement);
|
||||
await canvas.findByText('worker-east');
|
||||
|
||||
const detailsButtons = await canvas.findAllByRole('button', { name: /detail/i });
|
||||
await userEvent.click(detailsButtons[0]);
|
||||
|
||||
const body = within(document.body);
|
||||
await waitFor(() => {
|
||||
if (!body.queryByTestId('json-editor-container')) {
|
||||
throw new Error('Expected worker detail drawer with JSON view');
|
||||
}
|
||||
});
|
||||
},
|
||||
};
|
||||
|
||||
export const ManyWorkers: Story = {
|
||||
render: () => renderWithStore(),
|
||||
parameters: {
|
||||
msw: {
|
||||
handlers: createWorkersHandlers(manyWorkers),
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export const DarkTheme: Story = {
|
||||
render: () => renderWithStore({ theme: 'dark' }),
|
||||
parameters: {
|
||||
theme: 'dark',
|
||||
msw: {
|
||||
handlers: defaultHandlers,
|
||||
},
|
||||
},
|
||||
};
|
||||
@@ -0,0 +1,159 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useCallback, useEffect } from 'react';
|
||||
import { IconSearch } from '@tabler/icons-react';
|
||||
import type { DataTableSortStatus } from 'mantine-datatable';
|
||||
import { Skeleton, Stack, TextInput, Title } from '@mantine/core';
|
||||
import { WorkersTable, type WorkersTableRecord } from '@/components/WorkersTable.component';
|
||||
import { selectAutoRefreshMs } from '@/features/config';
|
||||
import { hideAlert, showAlert } from '@/features/ui/alert';
|
||||
import { openDrawer } from '@/features/ui/drawer';
|
||||
import {
|
||||
resetWorkersFilters,
|
||||
selectWorkersPage,
|
||||
selectWorkersQueryArgs,
|
||||
selectWorkersRecordsPerPage,
|
||||
selectWorkersSearchTerm,
|
||||
selectWorkersSort,
|
||||
setWorkersPage,
|
||||
setWorkersRecordsPerPage,
|
||||
setWorkersSearchTerm,
|
||||
setWorkersSort,
|
||||
useGetWorkersQuery,
|
||||
} from '@/features/workers';
|
||||
import { useAppDispatch, useAppSelector } from '@/store/hooks';
|
||||
import type { PaginatedResponse, Worker } from '@/types';
|
||||
import { getErrorDescriptor } from '@/utils/error';
|
||||
|
||||
export function WorkersPage() {
|
||||
const dispatch = useAppDispatch();
|
||||
const autoRefreshMs = useAppSelector(selectAutoRefreshMs);
|
||||
const searchTerm = useAppSelector(selectWorkersSearchTerm);
|
||||
const page = useAppSelector(selectWorkersPage);
|
||||
const recordsPerPage = useAppSelector(selectWorkersRecordsPerPage);
|
||||
const sort = useAppSelector(selectWorkersSort);
|
||||
const queryArgs = useAppSelector(selectWorkersQueryArgs);
|
||||
|
||||
const workersQueryResult = useGetWorkersQuery(queryArgs, {
|
||||
pollingInterval: autoRefreshMs > 0 ? autoRefreshMs : undefined,
|
||||
});
|
||||
|
||||
const workersData = workersQueryResult.data as PaginatedResponse<Worker> | undefined;
|
||||
const { isLoading, isFetching, isError, error, refetch } = workersQueryResult;
|
||||
|
||||
const handleSearchTermChange = useCallback(
|
||||
(value: string) => {
|
||||
dispatch(setWorkersSearchTerm(value));
|
||||
},
|
||||
[dispatch],
|
||||
);
|
||||
|
||||
const handleSortStatusChange = useCallback(
|
||||
(status: DataTableSortStatus<WorkersTableRecord>) => {
|
||||
dispatch(
|
||||
setWorkersSort({
|
||||
column: status.columnAccessor,
|
||||
direction: status.direction,
|
||||
}),
|
||||
);
|
||||
},
|
||||
[dispatch],
|
||||
);
|
||||
|
||||
const handlePageChange = useCallback(
|
||||
(nextPage: number) => {
|
||||
dispatch(setWorkersPage(nextPage));
|
||||
},
|
||||
[dispatch],
|
||||
);
|
||||
|
||||
const handleRecordsPerPageChange = useCallback(
|
||||
(value: number) => {
|
||||
dispatch(setWorkersRecordsPerPage(value));
|
||||
},
|
||||
[dispatch],
|
||||
);
|
||||
|
||||
const handleResetFilters = useCallback(() => {
|
||||
dispatch(resetWorkersFilters());
|
||||
}, [dispatch]);
|
||||
|
||||
const handleShowWorkerDetails = useCallback(
|
||||
(worker: Worker) => {
|
||||
dispatch(
|
||||
openDrawer({
|
||||
type: 'worker-detail',
|
||||
worker,
|
||||
}),
|
||||
);
|
||||
},
|
||||
[dispatch],
|
||||
);
|
||||
|
||||
const hasWorkers = Array.isArray(workersData?.items) && workersData.items.length > 0;
|
||||
const showSkeleton = isLoading && !hasWorkers;
|
||||
|
||||
useEffect(() => {
|
||||
if (isError) {
|
||||
const descriptor = getErrorDescriptor(error);
|
||||
const suffix = descriptor ? ` (${descriptor})` : '';
|
||||
dispatch(
|
||||
showAlert({
|
||||
id: 'workers-fetch',
|
||||
message: `Unable to refresh workers${suffix}. The table may be out of date until the connection recovers.`,
|
||||
tone: 'error',
|
||||
}),
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
if (!isLoading && !isFetching) {
|
||||
dispatch(hideAlert({ id: 'workers-fetch' }));
|
||||
}
|
||||
}, [dispatch, error, isError, isFetching, isLoading]);
|
||||
|
||||
useEffect(
|
||||
() => () => {
|
||||
dispatch(hideAlert({ id: 'workers-fetch' }));
|
||||
},
|
||||
[dispatch],
|
||||
);
|
||||
|
||||
return (
|
||||
<Stack gap='md'>
|
||||
<Title order={1}>Runners</Title>
|
||||
|
||||
<TextInput
|
||||
placeholder='Search by Runner ID'
|
||||
value={searchTerm}
|
||||
onChange={(event) => handleSearchTermChange(event.currentTarget.value)}
|
||||
leftSection={<IconSearch size={16} />}
|
||||
data-testid='workers-search-input'
|
||||
w='100%'
|
||||
style={{ maxWidth: 360 }}
|
||||
/>
|
||||
|
||||
{showSkeleton ? (
|
||||
<Skeleton height={360} radius='md' />
|
||||
) : (
|
||||
<WorkersTable
|
||||
workers={workersData?.items}
|
||||
totalRecords={workersData?.total ?? 0}
|
||||
isFetching={isFetching}
|
||||
isError={isError}
|
||||
error={error}
|
||||
searchTerm={searchTerm}
|
||||
sort={sort}
|
||||
page={page}
|
||||
recordsPerPage={recordsPerPage}
|
||||
onSortStatusChange={handleSortStatusChange}
|
||||
onPageChange={handlePageChange}
|
||||
onRecordsPerPageChange={handleRecordsPerPageChange}
|
||||
onResetFilters={handleResetFilters}
|
||||
onRefetch={refetch}
|
||||
onShowDetails={handleShowWorkerDetails}
|
||||
/>
|
||||
)}
|
||||
</Stack>
|
||||
);
|
||||
}
|
||||
@@ -7,6 +7,7 @@ import { rolloutsApi, rolloutsReducer } from '../features/rollouts';
|
||||
import { tracesReducer } from '../features/traces';
|
||||
import { alertReducer } from '../features/ui/alert';
|
||||
import { drawerReducer } from '../features/ui/drawer';
|
||||
import { workersReducer } from '../features/workers';
|
||||
|
||||
const rootReducer = combineReducers({
|
||||
config: configReducer,
|
||||
@@ -14,6 +15,7 @@ const rootReducer = combineReducers({
|
||||
alert: alertReducer,
|
||||
rollouts: rolloutsReducer,
|
||||
resources: resourcesReducer,
|
||||
workers: workersReducer,
|
||||
traces: tracesReducer,
|
||||
[rolloutsApi.reducerPath]: rolloutsApi.reducer,
|
||||
});
|
||||
|
||||
@@ -28,6 +28,24 @@ export type Attempt = {
|
||||
metadata: Record<string, any> | null;
|
||||
};
|
||||
|
||||
export type WorkerStatus = 'idle' | 'busy' | 'unknown';
|
||||
|
||||
/**
|
||||
* Synced with agentlightning.types.core.Worker
|
||||
* with camel case and snake case conversions
|
||||
*/
|
||||
export type Worker = {
|
||||
workerId: string;
|
||||
status: WorkerStatus;
|
||||
heartbeatStats: Record<string, any> | null;
|
||||
lastHeartbeatTime: Timestamp | null;
|
||||
lastDequeueTime: Timestamp | null;
|
||||
lastBusyTime: Timestamp | null;
|
||||
lastIdleTime: Timestamp | null;
|
||||
currentRolloutId: string | null;
|
||||
currentAttemptId: string | null;
|
||||
};
|
||||
|
||||
/**
|
||||
* Synced with agentlightning.types.core.Rollout
|
||||
* with camel case and snake case conversions
|
||||
|
||||
@@ -8,27 +8,32 @@
|
||||
*/
|
||||
|
||||
import { describe, expect, it } from 'vitest';
|
||||
import type { Attempt, Resources, Rollout, Span } from '@/types';
|
||||
import type { Attempt, Resources, Rollout, Span, Worker } from '@/types';
|
||||
import {
|
||||
buildAttemptsResponse,
|
||||
buildResourcesResponse,
|
||||
buildRolloutsResponse,
|
||||
buildSpansResponse,
|
||||
buildWorkersResponse,
|
||||
createMockHandlers,
|
||||
createResourcesHandlers,
|
||||
createRolloutsHandlers,
|
||||
createSpansHandlers,
|
||||
createWorkersHandlers,
|
||||
filterResourcesForParams,
|
||||
filterRolloutsForParams,
|
||||
filterSpansForParams,
|
||||
filterWorkersForParams,
|
||||
getResourcesSortValue,
|
||||
getRolloutSortValue,
|
||||
getSpanSortValue,
|
||||
getWorkerSortValue,
|
||||
parseNumberParam,
|
||||
sortAttemptsForParams,
|
||||
sortResourcesForParams,
|
||||
sortRolloutsForParams,
|
||||
sortSpansForParams,
|
||||
sortWorkersForParams,
|
||||
} from './mock';
|
||||
|
||||
const now = Math.floor(Date.now() / 1000);
|
||||
@@ -219,6 +224,53 @@ const sampleResources: Resources[] = [
|
||||
},
|
||||
];
|
||||
|
||||
const sampleWorkers: Worker[] = [
|
||||
{
|
||||
workerId: 'worker-alpha',
|
||||
status: 'busy',
|
||||
heartbeatStats: { queueDepth: 2 },
|
||||
lastHeartbeatTime: now - 30,
|
||||
lastDequeueTime: now - 300,
|
||||
lastBusyTime: now - 60,
|
||||
lastIdleTime: now - 600,
|
||||
currentRolloutId: 'ro-001',
|
||||
currentAttemptId: 'at-001',
|
||||
},
|
||||
{
|
||||
workerId: 'worker-beta',
|
||||
status: 'idle',
|
||||
heartbeatStats: { queueDepth: 0 },
|
||||
lastHeartbeatTime: now - 120,
|
||||
lastDequeueTime: now - 1200,
|
||||
lastBusyTime: now - 3600,
|
||||
lastIdleTime: now - 180,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
{
|
||||
workerId: 'worker-gamma',
|
||||
status: 'busy',
|
||||
heartbeatStats: null,
|
||||
lastHeartbeatTime: now - 10,
|
||||
lastDequeueTime: now - 60,
|
||||
lastBusyTime: now - 20,
|
||||
lastIdleTime: now - 4000,
|
||||
currentRolloutId: 'ro-003',
|
||||
currentAttemptId: 'at-003',
|
||||
},
|
||||
{
|
||||
workerId: 'worker-delta',
|
||||
status: 'unknown',
|
||||
heartbeatStats: { queueDepth: 0 },
|
||||
lastHeartbeatTime: now - 5,
|
||||
lastDequeueTime: now - 80,
|
||||
lastBusyTime: null,
|
||||
lastIdleTime: null,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
];
|
||||
|
||||
describe('parseNumberParam', () => {
|
||||
it('returns default value when param is missing', () => {
|
||||
const params = new URLSearchParams();
|
||||
@@ -725,6 +777,115 @@ describe('createResourcesHandlers', () => {
|
||||
});
|
||||
});
|
||||
|
||||
describe('filterWorkersForParams', () => {
|
||||
it('returns all workers without filters', () => {
|
||||
const params = new URLSearchParams();
|
||||
const result = filterWorkersForParams(sampleWorkers, params);
|
||||
expect(result).toHaveLength(4);
|
||||
});
|
||||
|
||||
it('filters by status and worker ID substring using AND logic', () => {
|
||||
const params = new URLSearchParams('status_in=busy&worker_id_contains=gamma');
|
||||
const result = filterWorkersForParams(sampleWorkers, params);
|
||||
expect(result).toHaveLength(1);
|
||||
expect(result[0].workerId).toBe('worker-gamma');
|
||||
});
|
||||
|
||||
it('supports filter_logic=or', () => {
|
||||
const params = new URLSearchParams('status_in=idle&worker_id_contains=gamma&filter_logic=or');
|
||||
const result = filterWorkersForParams(sampleWorkers, params);
|
||||
expect(result).toHaveLength(2);
|
||||
});
|
||||
|
||||
it('filters by unknown status', () => {
|
||||
const params = new URLSearchParams('status_in=unknown');
|
||||
const result = filterWorkersForParams(sampleWorkers, params);
|
||||
expect(result).toHaveLength(1);
|
||||
expect(result[0].workerId).toBe('worker-delta');
|
||||
});
|
||||
});
|
||||
|
||||
describe('getWorkerSortValue', () => {
|
||||
const worker = sampleWorkers[0];
|
||||
|
||||
it('returns worker_id', () => {
|
||||
expect(getWorkerSortValue(worker, 'worker_id')).toBe('worker-alpha');
|
||||
});
|
||||
|
||||
it('returns status', () => {
|
||||
expect(getWorkerSortValue(worker, 'status')).toBe('busy');
|
||||
});
|
||||
|
||||
it('returns timestamp fields', () => {
|
||||
expect(getWorkerSortValue(worker, 'last_busy_time')).toBe(worker.lastBusyTime);
|
||||
expect(getWorkerSortValue(worker, 'last_idle_time')).toBe(worker.lastIdleTime);
|
||||
expect(getWorkerSortValue(worker, 'last_dequeue_time')).toBe(worker.lastDequeueTime);
|
||||
});
|
||||
|
||||
it('returns rollout and attempt identifiers', () => {
|
||||
expect(getWorkerSortValue(worker, 'current_rollout_id')).toBe(worker.currentRolloutId);
|
||||
expect(getWorkerSortValue(worker, 'current_attempt_id')).toBe(worker.currentAttemptId);
|
||||
});
|
||||
|
||||
it('falls back to last_heartbeat_time', () => {
|
||||
expect(getWorkerSortValue(worker, 'unknown')).toBe(worker.lastHeartbeatTime);
|
||||
});
|
||||
});
|
||||
|
||||
describe('sortWorkersForParams', () => {
|
||||
it('sorts by last heartbeat ascending by default', () => {
|
||||
const result = sortWorkersForParams(sampleWorkers, null, 'asc');
|
||||
expect(result.map((worker) => worker.workerId)).toEqual([
|
||||
'worker-beta',
|
||||
'worker-alpha',
|
||||
'worker-gamma',
|
||||
'worker-delta',
|
||||
]);
|
||||
});
|
||||
|
||||
it('sorts descending by worker_id when requested', () => {
|
||||
const result = sortWorkersForParams(sampleWorkers, 'worker_id', 'desc');
|
||||
expect(result.map((worker) => worker.workerId)).toEqual([
|
||||
'worker-gamma',
|
||||
'worker-delta',
|
||||
'worker-beta',
|
||||
'worker-alpha',
|
||||
]);
|
||||
});
|
||||
|
||||
it('sorts by current_rollout_id', () => {
|
||||
const result = sortWorkersForParams(sampleWorkers, 'current_rollout_id', 'asc');
|
||||
expect(result.map((worker) => worker.currentRolloutId)).toEqual([null, null, 'ro-001', 'ro-003']);
|
||||
});
|
||||
});
|
||||
|
||||
describe('buildWorkersResponse', () => {
|
||||
it('applies filters before pagination', () => {
|
||||
const request = new Request('http://localhost/v1/agl/workers?worker_id_contains=beta&limit=5');
|
||||
const response = buildWorkersResponse(sampleWorkers, request);
|
||||
expect(response.items).toHaveLength(1);
|
||||
const items = response.items as Array<Record<string, unknown>>;
|
||||
expect(items[0].worker_id).toBe('worker-beta');
|
||||
});
|
||||
|
||||
it('applies sort and pagination parameters', () => {
|
||||
const request = new Request('http://localhost/v1/agl/workers?sort_by=worker_id&limit=2&offset=1');
|
||||
const response = buildWorkersResponse(sampleWorkers, request);
|
||||
expect(response.items).toHaveLength(2);
|
||||
const items = response.items as Array<Record<string, unknown>>;
|
||||
expect(items[0].worker_id).toBe('worker-beta');
|
||||
expect(response.total).toBe(4);
|
||||
});
|
||||
});
|
||||
|
||||
describe('createWorkersHandlers', () => {
|
||||
it('creates handler for workers endpoint', () => {
|
||||
const handlers = createWorkersHandlers(sampleWorkers);
|
||||
expect(handlers).toHaveLength(1);
|
||||
expect(handlers[0].info.header).toContain('GET');
|
||||
});
|
||||
});
|
||||
|
||||
describe('createRolloutsHandlers', () => {
|
||||
it('creates handlers that return correct rollout data', async () => {
|
||||
const attemptsByRollout = { 'ro-001': sampleAttempts };
|
||||
|
||||
+112
-1
@@ -14,7 +14,7 @@
|
||||
*/
|
||||
|
||||
import { delay, http, HttpResponse } from 'msw';
|
||||
import type { Attempt, Resources, Rollout, Span } from '@/types';
|
||||
import type { Attempt, Resources, Rollout, Span, Worker } from '@/types';
|
||||
import { snakeCaseKeys } from './format';
|
||||
|
||||
/**
|
||||
@@ -434,6 +434,117 @@ export function buildResourcesResponse(resources: Resources[], request: Request)
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* Filter workers based on query parameters.
|
||||
* Supports: status_in, worker_id_contains
|
||||
*/
|
||||
export function filterWorkersForParams(workers: Worker[], params: URLSearchParams): Worker[] {
|
||||
const statusFilters = params.getAll('status_in');
|
||||
const workerIdContains = params.get('worker_id_contains');
|
||||
const filterLogic = params.get('filter_logic') === 'or' ? 'or' : 'and';
|
||||
|
||||
return workers.filter((worker) => {
|
||||
const checks: boolean[] = [];
|
||||
if (statusFilters.length > 0) {
|
||||
checks.push(statusFilters.includes(worker.status));
|
||||
}
|
||||
if (workerIdContains) {
|
||||
checks.push(worker.workerId.toLowerCase().includes(workerIdContains.toLowerCase()));
|
||||
}
|
||||
if (checks.length === 0) {
|
||||
return true;
|
||||
}
|
||||
return filterLogic === 'or' ? checks.some(Boolean) : checks.every(Boolean);
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* Resolve a worker sort value for the given column.
|
||||
*/
|
||||
export function getWorkerSortValue(worker: Worker, sortBy: string): string | number | null {
|
||||
switch (sortBy) {
|
||||
case 'worker_id':
|
||||
return worker.workerId;
|
||||
case 'status':
|
||||
return worker.status;
|
||||
case 'current_rollout_id':
|
||||
return worker.currentRolloutId ?? '';
|
||||
case 'current_attempt_id':
|
||||
return worker.currentAttemptId ?? '';
|
||||
case 'last_busy_time':
|
||||
return worker.lastBusyTime ?? null;
|
||||
case 'last_idle_time':
|
||||
return worker.lastIdleTime ?? null;
|
||||
case 'last_dequeue_time':
|
||||
return worker.lastDequeueTime ?? null;
|
||||
case 'last_heartbeat_time':
|
||||
default:
|
||||
return worker.lastHeartbeatTime ?? null;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Sort workers based on query parameters.
|
||||
* Default sort_by is 'last_heartbeat_time'.
|
||||
*/
|
||||
export function sortWorkersForParams(workers: Worker[], sortBy: string | null, sortOrder: 'asc' | 'desc'): Worker[] {
|
||||
const resolvedSortBy = sortBy ?? 'last_heartbeat_time';
|
||||
const sorted = [...workers].sort((a, b) => {
|
||||
const aValue = getWorkerSortValue(a, resolvedSortBy);
|
||||
const bValue = getWorkerSortValue(b, resolvedSortBy);
|
||||
if (aValue === bValue) {
|
||||
return 0;
|
||||
}
|
||||
if (aValue == null) {
|
||||
return -1;
|
||||
}
|
||||
if (bValue == null) {
|
||||
return 1;
|
||||
}
|
||||
if (typeof aValue === 'number' && typeof bValue === 'number') {
|
||||
return aValue - bValue;
|
||||
}
|
||||
return String(aValue).localeCompare(String(bValue));
|
||||
});
|
||||
|
||||
if (sortOrder === 'desc') {
|
||||
sorted.reverse();
|
||||
}
|
||||
|
||||
return sorted;
|
||||
}
|
||||
|
||||
/**
|
||||
* Build a paginated workers response matching the Python server's format.
|
||||
*/
|
||||
export function buildWorkersResponse(workers: Worker[], request: Request): Record<string, unknown> {
|
||||
const url = new URL(request.url);
|
||||
const params = url.searchParams;
|
||||
const filtered = filterWorkersForParams(workers, params);
|
||||
const sortBy = params.get('sort_by');
|
||||
const sortOrder = params.get('sort_order') === 'desc' ? 'desc' : 'asc';
|
||||
const sorted = sortWorkersForParams(filtered, sortBy, sortOrder);
|
||||
const limitParam = parseNumberParam(params, 'limit', sorted.length);
|
||||
const offsetParam = parseNumberParam(params, 'offset', 0);
|
||||
const effectiveLimit = limitParam < 0 ? sorted.length : limitParam;
|
||||
const offset = offsetParam < 0 ? 0 : offsetParam;
|
||||
const paginated = effectiveLimit >= 0 ? sorted.slice(offset, offset + effectiveLimit) : [...sorted];
|
||||
|
||||
return snakeCaseKeys({
|
||||
items: paginated,
|
||||
limit: effectiveLimit,
|
||||
offset,
|
||||
total: filtered.length,
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* Create MSW handlers for workers endpoints.
|
||||
*/
|
||||
export function createWorkersHandlers(workers: Worker[]) {
|
||||
return [http.get('*/v1/agl/workers', ({ request }) => HttpResponse.json(buildWorkersResponse(workers, request)))];
|
||||
}
|
||||
|
||||
/**
|
||||
* Create MSW handlers for resources endpoints.
|
||||
*
|
||||
|
||||
@@ -33,6 +33,7 @@ from agentlightning.types import (
|
||||
RolloutConfig,
|
||||
Span,
|
||||
TraceStatus,
|
||||
Worker,
|
||||
)
|
||||
|
||||
|
||||
@@ -634,6 +635,68 @@ def inject_mock_data(store: InMemoryLightningStore, now: float | None = None) ->
|
||||
store._resources["rs-story-005"] = resource5
|
||||
store._latest_resources_id = "rs-story-005"
|
||||
|
||||
# Register workers with diverse states and activity windows.
|
||||
workers = [
|
||||
Worker(
|
||||
worker_id="worker-east",
|
||||
status="busy",
|
||||
heartbeat_stats={"queue_depth": 2, "gpu_utilization": 0.82},
|
||||
last_heartbeat_time=now - 20,
|
||||
last_dequeue_time=now - 60,
|
||||
last_busy_time=now - 120,
|
||||
last_idle_time=now - 600,
|
||||
current_rollout_id="ro-story-001",
|
||||
current_attempt_id="at-story-010",
|
||||
),
|
||||
Worker(
|
||||
worker_id="worker-north",
|
||||
status="idle",
|
||||
heartbeat_stats={"queue_depth": 0, "gpu_utilization": 0.15},
|
||||
last_heartbeat_time=now - 90,
|
||||
last_dequeue_time=now - 3600,
|
||||
last_busy_time=now - 5400,
|
||||
last_idle_time=now - 5400,
|
||||
current_rollout_id=None,
|
||||
current_attempt_id=None,
|
||||
),
|
||||
Worker(
|
||||
worker_id="worker-west",
|
||||
status="busy",
|
||||
heartbeat_stats={"queue_depth": 1, "gpu_utilization": 0.41},
|
||||
last_heartbeat_time=now - 45,
|
||||
last_dequeue_time=now - 300,
|
||||
last_busy_time=now - 200,
|
||||
last_idle_time=now - 4800,
|
||||
current_rollout_id="ro-story-003",
|
||||
current_attempt_id="at-story-033",
|
||||
),
|
||||
Worker(
|
||||
worker_id="worker-south",
|
||||
status="idle",
|
||||
heartbeat_stats={"queue_depth": 0},
|
||||
last_heartbeat_time=now - 900,
|
||||
last_dequeue_time=now - 7200,
|
||||
last_busy_time=now - 8600,
|
||||
last_idle_time=now - 8600,
|
||||
current_rollout_id=None,
|
||||
current_attempt_id=None,
|
||||
),
|
||||
Worker(
|
||||
worker_id="worker-observer",
|
||||
status="unknown",
|
||||
heartbeat_stats={"queue_depth": 0},
|
||||
last_heartbeat_time=now - 15,
|
||||
last_dequeue_time=now - 4000,
|
||||
last_busy_time=None,
|
||||
last_idle_time=None,
|
||||
current_rollout_id=None,
|
||||
current_attempt_id=None,
|
||||
),
|
||||
]
|
||||
|
||||
for worker in workers:
|
||||
store._workers[worker.worker_id] = worker
|
||||
|
||||
|
||||
async def main():
|
||||
parser = argparse.ArgumentParser(description="Run a Python server for the LightningStore")
|
||||
|
||||
@@ -72,6 +72,15 @@ Each attempt begins in **preparing**, created either when a rollout is dequeued
|
||||
|
||||
This simple model allows the system to distinguish between normal termination, abnormal stalling, and recoverable interruption without additional state flags.
|
||||
|
||||
## Worker Telemetry
|
||||
|
||||
Workers track runner-level activity timestamps (`last_heartbeat_time`, `last_dequeue_time`, `last_busy_time`, `last_idle_time`) plus their current rollout assignment. Those fields are now derived automatically:
|
||||
|
||||
- [`dequeue_rollout(worker_id=...)`][agentlightning.LightningStore.dequeue_rollout] records which worker polled the queue and refreshes `last_dequeue_time`.
|
||||
- [`update_attempt(..., worker_id=...)`][agentlightning.LightningStore.update_attempt] drives the worker status machine. Assigning an attempt marks the worker **busy** and stamps `last_busy_time`; finishing with `status in {"succeeded","failed"}` switches to **idle**, while watchdog transitions such as `timeout`/`unresponsive` make the worker **unknown** and clear `current_rollout_id` / `current_attempt_id`.
|
||||
- [`update_worker(...)`][agentlightning.LightningStore.update_worker] is reserved for heartbeats. It snapshots optional `heartbeat_stats` and always updates `last_heartbeat_time`.
|
||||
|
||||
Because every transition flows through these APIs, worker status is derived automatically from rollout execution and heartbeats. Note, however, that calling `update_worker` with a new `worker_id` will create a new worker record with status "unknown" if one does not exist. Thus, while manual status changes are not allowed, new worker records can be created externally via heartbeats.
|
||||
|
||||
## Rollout Transition Map
|
||||
|
||||
|
||||
@@ -26,6 +26,10 @@
|
||||
|
||||
::: agentlightning.AttemptedRollout
|
||||
|
||||
::: agentlightning.Worker
|
||||
|
||||
::: agentlightning.WorkerStatus
|
||||
|
||||
::: agentlightning.Hook
|
||||
|
||||
## Resources
|
||||
|
||||
@@ -7,6 +7,7 @@ requires-python = ">=3.10"
|
||||
dependencies = [
|
||||
"graphviz",
|
||||
"psutil",
|
||||
"gpustat",
|
||||
"setproctitle",
|
||||
"flask",
|
||||
"agentops>=0.4.13",
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
import asyncio
|
||||
import random
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any, AsyncGenerator, Dict, List, Optional, Sequence, cast
|
||||
from typing import Any, AsyncGenerator, Dict, List, Literal, Optional, Sequence, Tuple, cast
|
||||
|
||||
import pytest
|
||||
from opentelemetry import trace as trace_api
|
||||
@@ -17,10 +17,10 @@ from agentlightning.litagent import LitAgent
|
||||
from agentlightning.reward import emit_reward, find_final_reward
|
||||
from agentlightning.runner import LitAgentRunner
|
||||
from agentlightning.runner.base import Runner
|
||||
from agentlightning.store.base import LightningStore
|
||||
from agentlightning.store.base import UNSET, LightningStore, Unset
|
||||
from agentlightning.store.memory import InMemoryLightningStore
|
||||
from agentlightning.tracer.base import Tracer
|
||||
from agentlightning.types import LLM, Hook, NamedResources, PromptTemplate, Rollout, Span, SpanNames
|
||||
from agentlightning.types import LLM, Hook, NamedResources, PromptTemplate, Rollout, Span, SpanNames, Worker
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
@@ -114,6 +114,49 @@ class DummyTracer(Tracer):
|
||||
return span
|
||||
|
||||
|
||||
class RecordingStore(InMemoryLightningStore):
|
||||
"""In-memory store that records worker heartbeat updates for inspection in tests."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.worker_updates: List[Tuple[str, Optional[Dict[str, Any]]]] = []
|
||||
|
||||
async def update_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
heartbeat_stats: Dict[str, Any] | Unset = UNSET,
|
||||
) -> Worker:
|
||||
payload = None if isinstance(heartbeat_stats, Unset) else heartbeat_stats
|
||||
self.worker_updates.append((worker_id, payload))
|
||||
return await super().update_worker(worker_id, heartbeat_stats=heartbeat_stats)
|
||||
|
||||
|
||||
class HeartbeatAgent(LitAgent[Dict[str, Any]]):
|
||||
"""Minimal agent used for heartbeat-only runner tests."""
|
||||
|
||||
def validation_rollout(self, task: Dict[str, Any], resources: Dict[str, Any], rollout: Any) -> float:
|
||||
return 0.0
|
||||
|
||||
|
||||
async def setup_heartbeat_runner(
|
||||
*,
|
||||
heartbeat_interval: float = 0.05,
|
||||
heartbeat_launch_mode: Literal["asyncio", "thread"] = "asyncio",
|
||||
) -> tuple[LitAgentRunner[Any], RecordingStore]:
|
||||
"""Create a runner wired to a RecordingStore for heartbeat tests."""
|
||||
|
||||
store = RecordingStore()
|
||||
runner = LitAgentRunner[Any](
|
||||
tracer=DummyTracer(),
|
||||
heartbeat_interval=heartbeat_interval,
|
||||
heartbeat_launch_mode=heartbeat_launch_mode,
|
||||
)
|
||||
agent = HeartbeatAgent()
|
||||
runner.init(agent)
|
||||
runner.init_worker(worker_id=0, store=store)
|
||||
return runner, store
|
||||
|
||||
|
||||
async def setup_runner(
|
||||
agent: LitAgent[Any],
|
||||
*,
|
||||
@@ -658,3 +701,46 @@ async def test_step_with_custom_resources_returns_rollout() -> None:
|
||||
|
||||
# Verify the rollout has the correct resources_id
|
||||
assert result.resources_id is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_emit_heartbeat_updates_worker_snapshot(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
snapshot = {"cpu_pct": 42.0, "mem_pct": 10.5}
|
||||
monkeypatch.setattr("agentlightning.runner.agent.system_snapshot", lambda: snapshot)
|
||||
|
||||
runner, store = await setup_heartbeat_runner(heartbeat_interval=0.1)
|
||||
worker_label = runner.get_worker_id()
|
||||
try:
|
||||
await runner._emit_heartbeat(store) # pyright: ignore[reportPrivateUsage]
|
||||
finally:
|
||||
teardown_runner(runner)
|
||||
|
||||
assert store.worker_updates == [(worker_label, snapshot)]
|
||||
worker = await store.get_worker_by_id(worker_label)
|
||||
assert worker is not None
|
||||
assert worker.heartbeat_stats == snapshot
|
||||
assert worker.last_heartbeat_time is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_heartbeat_loop_runs_until_stopped(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
snapshot = {"timestamp": 1234567890}
|
||||
monkeypatch.setattr("agentlightning.runner.agent.system_snapshot", lambda: snapshot)
|
||||
|
||||
runner, store = await setup_heartbeat_runner(heartbeat_interval=0.05)
|
||||
stop_heartbeat = runner._start_heartbeat_loop(store) # pyright: ignore[reportPrivateUsage]
|
||||
assert stop_heartbeat is not None
|
||||
|
||||
try:
|
||||
await asyncio.sleep(0.12)
|
||||
finally:
|
||||
await stop_heartbeat()
|
||||
|
||||
update_count = len(store.worker_updates)
|
||||
assert update_count >= 1
|
||||
assert all(stats == snapshot for _, stats in store.worker_updates if stats is not None)
|
||||
|
||||
await asyncio.sleep(0.06)
|
||||
assert len(store.worker_updates) == update_count
|
||||
|
||||
teardown_runner(runner)
|
||||
|
||||
@@ -16,6 +16,7 @@ from agentlightning.types import (
|
||||
RolloutStatus,
|
||||
Span,
|
||||
TaskInput,
|
||||
Worker,
|
||||
)
|
||||
|
||||
|
||||
@@ -47,8 +48,8 @@ class DummyLightningStore(LightningStore):
|
||||
self.calls.append(("enqueue_rollout", (input, mode, resources_id, config, metadata), {}))
|
||||
return self.return_values["enqueue_rollout"]
|
||||
|
||||
async def dequeue_rollout(self) -> Optional[AttemptedRollout]:
|
||||
self.calls.append(("dequeue_rollout", (), {}))
|
||||
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:
|
||||
@@ -156,6 +157,31 @@ class DummyLightningStore(LightningStore):
|
||||
)
|
||||
return self.return_values["update_attempt"]
|
||||
|
||||
async def query_workers(self) -> List[Worker]:
|
||||
self.calls.append(("query_workers", (), {}))
|
||||
return self.return_values["query_workers"]
|
||||
|
||||
async def get_worker_by_id(self, worker_id: str) -> Optional[Worker]:
|
||||
self.calls.append(("get_worker_by_id", (worker_id,), {}))
|
||||
return self.return_values["get_worker_by_id"]
|
||||
|
||||
async def update_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
heartbeat_stats: Dict[str, Any] | Any = UNSET,
|
||||
) -> Worker:
|
||||
self.calls.append(
|
||||
(
|
||||
"update_worker",
|
||||
(
|
||||
worker_id,
|
||||
heartbeat_stats,
|
||||
),
|
||||
{},
|
||||
)
|
||||
)
|
||||
return self.return_values["update_worker"]
|
||||
|
||||
|
||||
def minimal_dummy_store() -> DummyLightningStore:
|
||||
# Provide minimal return values
|
||||
@@ -180,5 +206,8 @@ def minimal_dummy_store() -> DummyLightningStore:
|
||||
"query_spans": [],
|
||||
"update_rollout": None,
|
||||
"update_attempt": None,
|
||||
"query_workers": [],
|
||||
"get_worker_by_id": None,
|
||||
"update_worker": Worker(worker_id="worker-0"),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -229,7 +229,13 @@ async def test_client_server_end_to_end(
|
||||
server_queue_config = RolloutConfig(unresponsive_seconds=4.2, max_attempts=2)
|
||||
queued_rollout = await server.enqueue_rollout(input={"origin": "server-queue"}, config=server_queue_config)
|
||||
assert queued_rollout.config.unresponsive_seconds == 4.2
|
||||
dequeued = await server.dequeue_rollout()
|
||||
server_worker_id = "server-worker"
|
||||
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.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)
|
||||
|
||||
await server.query_rollouts()
|
||||
@@ -260,10 +266,25 @@ async def test_client_server_end_to_end(
|
||||
queued_rollout.rollout_id,
|
||||
started_attempt.attempt.attempt_id,
|
||||
status="running",
|
||||
worker_id="server-worker",
|
||||
worker_id=server_worker_id,
|
||||
metadata={"phase": "warmup"},
|
||||
)
|
||||
server_worker_busy = await server.get_worker_by_id(server_worker_id)
|
||||
assert server_worker_busy is not None
|
||||
assert server_worker_busy.status == "busy"
|
||||
assert server_worker_busy.current_rollout_id == queued_rollout.rollout_id
|
||||
assert server_worker_busy.current_attempt_id == started_attempt.attempt.attempt_id
|
||||
assert server_worker_busy.last_busy_time is not None
|
||||
assert server_worker_busy.last_busy_time >= dequeue_time
|
||||
|
||||
await server.update_attempt(queued_rollout.rollout_id, "latest", status="succeeded")
|
||||
server_worker_idle = await server.get_worker_by_id(server_worker_id)
|
||||
assert server_worker_idle is not None
|
||||
assert server_worker_idle.status == "idle"
|
||||
assert server_worker_idle.current_rollout_id is None
|
||||
assert server_worker_idle.current_attempt_id is None
|
||||
assert server_worker_idle.last_idle_time is not None
|
||||
assert server_worker_idle.last_idle_time >= server_worker_busy.last_busy_time
|
||||
completed = await server.wait_for_rollouts(rollout_ids=[queued_rollout.rollout_id], timeout=0.1)
|
||||
assert completed and completed[0].status in {"succeeded", "failed", "cancelled"}
|
||||
|
||||
@@ -285,8 +306,14 @@ async def test_client_server_end_to_end(
|
||||
client_queue_config = RolloutConfig(unresponsive_seconds=6.0)
|
||||
enqueued = await client.enqueue_rollout(input={"origin": "client-queue"}, config=client_queue_config)
|
||||
assert enqueued.config.unresponsive_seconds == 6.0
|
||||
dequeued_client = await client.dequeue_rollout()
|
||||
client_worker_id = "client-worker"
|
||||
dequeued_client = await client.dequeue_rollout(worker_id=client_worker_id)
|
||||
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.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)
|
||||
|
||||
all_rollouts = await client.query_rollouts()
|
||||
@@ -324,11 +351,26 @@ async def test_client_server_end_to_end(
|
||||
await client.update_attempt(
|
||||
dequeued_client.rollout_id,
|
||||
started_client_attempt.attempt.attempt_id,
|
||||
worker_id="client-worker",
|
||||
worker_id=client_worker_id,
|
||||
metadata={"info": "started"},
|
||||
)
|
||||
client_worker_busy = await client.get_worker_by_id(client_worker_id)
|
||||
assert client_worker_busy is not None
|
||||
assert client_worker_busy.status == "busy"
|
||||
assert client_worker_busy.current_rollout_id == dequeued_client.rollout_id
|
||||
assert client_worker_busy.current_attempt_id == started_client_attempt.attempt.attempt_id
|
||||
assert client_worker_busy.last_busy_time is not None
|
||||
assert client_worker_busy.last_busy_time >= client_dequeue_time
|
||||
|
||||
await client.update_attempt(dequeued_client.rollout_id, "latest", status="succeeded")
|
||||
await client.update_rollout(dequeued_client.rollout_id, status="succeeded")
|
||||
client_worker_idle = await client.get_worker_by_id(client_worker_id)
|
||||
assert client_worker_idle is not None
|
||||
assert client_worker_idle.status == "idle"
|
||||
assert client_worker_idle.current_rollout_id is None
|
||||
assert client_worker_idle.current_attempt_id is None
|
||||
assert client_worker_idle.last_idle_time is not None
|
||||
assert client_worker_idle.last_idle_time >= client_worker_busy.last_busy_time
|
||||
|
||||
wait_result = await client.wait_for_rollouts(rollout_ids=[dequeued_client.rollout_id], timeout=0.05)
|
||||
assert wait_result and wait_result[0].status == "succeeded"
|
||||
@@ -415,6 +457,77 @@ async def test_update_attempt_none_vs_unset(server_client: Tuple[LightningStoreS
|
||||
assert preserved.status == "running"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_worker_records_heartbeat(
|
||||
server_client: Tuple[LightningStoreServer, LightningStoreClient],
|
||||
) -> None:
|
||||
_, client = server_client
|
||||
|
||||
first = await client.update_worker("runner-1", heartbeat_stats={"cpu": 0.4})
|
||||
assert first.status == "unknown"
|
||||
assert first.heartbeat_stats == {"cpu": 0.4}
|
||||
assert first.last_heartbeat_time is not None
|
||||
|
||||
second = await client.update_worker("runner-1")
|
||||
assert second.last_heartbeat_time is not None
|
||||
assert second.last_heartbeat_time >= first.last_heartbeat_time
|
||||
assert second.heartbeat_stats == {"cpu": 0.4}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_worker_rejects_none_stats(
|
||||
server_client: Tuple[LightningStoreServer, LightningStoreClient],
|
||||
) -> None:
|
||||
_, client = server_client
|
||||
with pytest.raises(ClientResponseError) as exc_info:
|
||||
await client.update_worker("runner-err", heartbeat_stats=cast(Any, None))
|
||||
assert exc_info.value.status == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_worker_status_transitions_via_attempts(
|
||||
server_client: Tuple[LightningStoreServer, LightningStoreClient],
|
||||
) -> None:
|
||||
_, client = server_client
|
||||
|
||||
await client.enqueue_rollout(input={"payload": "worker"})
|
||||
claimed = await client.dequeue_rollout(worker_id="runner-auto")
|
||||
assert claimed is not None
|
||||
|
||||
await client.update_attempt(claimed.rollout_id, claimed.attempt.attempt_id, worker_id="runner-auto")
|
||||
busy = await client.get_worker_by_id("runner-auto")
|
||||
assert busy is not None
|
||||
assert busy.status == "busy"
|
||||
assert busy.current_rollout_id == claimed.rollout_id
|
||||
assert busy.current_attempt_id == claimed.attempt.attempt_id
|
||||
assert busy.last_dequeue_time is not None
|
||||
assert busy.last_busy_time is not None
|
||||
|
||||
await client.update_attempt(claimed.rollout_id, claimed.attempt.attempt_id, status="succeeded")
|
||||
idle = await client.get_worker_by_id("runner-auto")
|
||||
assert idle is not None
|
||||
assert idle.status == "idle"
|
||||
assert idle.current_rollout_id is None
|
||||
assert idle.current_attempt_id is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_worker_by_id(server_client: Tuple[LightningStoreServer, LightningStoreClient]) -> None:
|
||||
server, client = server_client
|
||||
|
||||
await server.update_worker("runner-lookup", heartbeat_stats={"cpu": 0.3})
|
||||
|
||||
server_worker = await server.get_worker_by_id("runner-lookup")
|
||||
assert server_worker is not None
|
||||
assert server_worker.worker_id == "runner-lookup"
|
||||
assert await server.get_worker_by_id("missing") is None
|
||||
|
||||
client_worker = await client.get_worker_by_id("runner-lookup")
|
||||
assert client_worker is not None
|
||||
assert client_worker.worker_id == "runner-lookup"
|
||||
assert await client.get_worker_by_id("missing") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"bad_payload",
|
||||
|
||||
@@ -446,6 +446,49 @@ async def test_requeue_mechanism(store_fixture: LightningStore) -> None:
|
||||
assert latest_attempt.sequence_id == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_and_query_workers(store_fixture: LightningStore) -> None:
|
||||
"""Workers can be created, heartbeats recorded, and telemetry auto-updated."""
|
||||
first = await store_fixture.update_worker("worker-1", heartbeat_stats={"cpu": 0.5})
|
||||
assert first.worker_id == "worker-1"
|
||||
assert first.heartbeat_stats == {"cpu": 0.5}
|
||||
assert isinstance(first.last_heartbeat_time, float)
|
||||
assert first.status == "unknown"
|
||||
|
||||
rollout = await store_fixture.enqueue_rollout(input={"task": "work"})
|
||||
claimed = await store_fixture.dequeue_rollout(worker_id="worker-1")
|
||||
assert claimed is not None
|
||||
assert claimed.rollout_id == rollout.rollout_id
|
||||
|
||||
await store_fixture.update_attempt(claimed.rollout_id, claimed.attempt.attempt_id, worker_id="worker-1")
|
||||
busy = await store_fixture.get_worker_by_id("worker-1")
|
||||
assert busy is not None
|
||||
assert busy.status == "busy"
|
||||
assert busy.current_rollout_id == claimed.rollout_id
|
||||
assert busy.current_attempt_id == claimed.attempt.attempt_id
|
||||
assert isinstance(busy.last_dequeue_time, float)
|
||||
assert isinstance(busy.last_busy_time, float)
|
||||
|
||||
heartbeat = await store_fixture.update_worker("worker-1")
|
||||
assert heartbeat.last_heartbeat_time is not None
|
||||
assert heartbeat.last_heartbeat_time >= first.last_heartbeat_time
|
||||
|
||||
await store_fixture.update_attempt(claimed.rollout_id, claimed.attempt.attempt_id, status="succeeded")
|
||||
idle = await store_fixture.get_worker_by_id("worker-1")
|
||||
assert idle is not None
|
||||
assert idle.status == "idle"
|
||||
assert idle.current_rollout_id is None
|
||||
assert idle.current_attempt_id is None
|
||||
assert isinstance(idle.last_idle_time, float)
|
||||
|
||||
workers = await store_fixture.query_workers()
|
||||
assert any(w.worker_id == "worker-1" for w in workers)
|
||||
assert await store_fixture.get_worker_by_id("missing") is None
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
await store_fixture.update_worker("worker-1", heartbeat_stats=None) # type: ignore[arg-type]
|
||||
|
||||
|
||||
# Resource Management Tests
|
||||
|
||||
|
||||
|
||||
@@ -1012,3 +1012,123 @@ async def test_client_query_with_filters(
|
||||
rollouts = await client.query_rollouts(rollout_ids=[r2.rollout_id])
|
||||
assert len(rollouts) == 1
|
||||
assert rollouts[0].rollout_id == r2.rollout_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_workers_endpoint_supports_updates(
|
||||
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
|
||||
) -> None:
|
||||
_server, _client, session, api_endpoint = server_client
|
||||
|
||||
async with session.post(
|
||||
f"{api_endpoint}/workers/worker-1",
|
||||
json={"heartbeat_stats": {"cpu": 0.7}},
|
||||
) as resp:
|
||||
assert resp.status == 200
|
||||
created = await resp.json()
|
||||
assert created["worker_id"] == "worker-1"
|
||||
assert created["status"] == "unknown"
|
||||
assert created["heartbeat_stats"] == {"cpu": 0.7}
|
||||
first_heartbeat = created["last_heartbeat_time"]
|
||||
|
||||
async with session.get(f"{api_endpoint}/workers") as resp:
|
||||
assert resp.status == 200
|
||||
data = await resp.json()
|
||||
workers = data["items"]
|
||||
assert len(workers) == 1
|
||||
assert workers[0]["worker_id"] == "worker-1"
|
||||
|
||||
async with session.post(
|
||||
f"{api_endpoint}/workers/worker-1",
|
||||
json={"heartbeat_stats": {"cpu": 0.8}},
|
||||
) as resp:
|
||||
assert resp.status == 200
|
||||
updated = await resp.json()
|
||||
assert updated["last_heartbeat_time"] >= first_heartbeat
|
||||
|
||||
async with session.get(f"{api_endpoint}/workers") as resp:
|
||||
assert resp.status == 200
|
||||
data = await resp.json()
|
||||
workers = data["items"]
|
||||
assert workers[0]["status"] == "unknown"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_workers_endpoint_rejects_none_stats(
|
||||
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
|
||||
) -> None:
|
||||
_server, _client, session, api_endpoint = server_client
|
||||
|
||||
async with session.post(
|
||||
f"{api_endpoint}/workers/worker-err",
|
||||
json={"heartbeat_stats": None},
|
||||
) as resp:
|
||||
assert resp.status == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_worker_by_id_restful(
|
||||
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
|
||||
) -> None:
|
||||
server, _client, session, api_endpoint = server_client
|
||||
|
||||
await server.update_worker("worker-fetch", heartbeat_stats={"cpu": 0.4})
|
||||
|
||||
async with session.get(f"{api_endpoint}/workers/worker-fetch") as resp:
|
||||
assert resp.status == 200
|
||||
data = await resp.json()
|
||||
assert data["worker_id"] == "worker-fetch"
|
||||
|
||||
async with session.get(f"{api_endpoint}/workers/missing") as resp:
|
||||
assert resp.status == 200
|
||||
data = await resp.json()
|
||||
assert data is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_workers_endpoint_filter_and_sort(
|
||||
server_client: Tuple[LightningStoreServer, LightningStoreClient, aiohttp.ClientSession, str],
|
||||
) -> None:
|
||||
server, _client, session, api_endpoint = server_client
|
||||
|
||||
# Worker A: finishes an attempt and becomes idle.
|
||||
await server.update_worker("worker-a", heartbeat_stats={"cpu": 0.1})
|
||||
await server.enqueue_rollout(input={"worker": "a"})
|
||||
claimed_a = await server.dequeue_rollout(worker_id="worker-a")
|
||||
assert claimed_a is not None
|
||||
await server.update_attempt(
|
||||
claimed_a.rollout_id, claimed_a.attempt.attempt_id, worker_id="worker-a", status="succeeded"
|
||||
)
|
||||
|
||||
# Worker B: currently busy on an attempt.
|
||||
await server.update_worker("worker-b", heartbeat_stats={"cpu": 0.9})
|
||||
await server.enqueue_rollout(input={"worker": "b"})
|
||||
claimed_b = await server.dequeue_rollout(worker_id="worker-b")
|
||||
assert claimed_b is not None
|
||||
await server.update_attempt(claimed_b.rollout_id, claimed_b.attempt.attempt_id, worker_id="worker-b")
|
||||
|
||||
# Worker C: also busy.
|
||||
await server.update_worker("worker-c", heartbeat_stats={"cpu": 0.2})
|
||||
await server.enqueue_rollout(input={"worker": "c"})
|
||||
claimed_c = await server.dequeue_rollout(worker_id="worker-c")
|
||||
assert claimed_c is not None
|
||||
await server.update_attempt(claimed_c.rollout_id, claimed_c.attempt.attempt_id, worker_id="worker-c")
|
||||
|
||||
async with session.get(
|
||||
f"{api_endpoint}/workers",
|
||||
params={"status_in": ["busy"], "worker_id_contains": "worker", "sort_by": "worker_id", "sort_order": "desc"},
|
||||
) as resp:
|
||||
assert resp.status == 200
|
||||
data = await resp.json()
|
||||
items = data["items"]
|
||||
assert [w["worker_id"] for w in items] == ["worker-c", "worker-b"]
|
||||
|
||||
async with session.get(
|
||||
f"{api_endpoint}/workers",
|
||||
params={"limit": 1, "offset": 1, "sort_by": "worker_id", "sort_order": "asc"},
|
||||
) as resp:
|
||||
assert resp.status == 200
|
||||
data = await resp.json()
|
||||
assert data["limit"] == 1
|
||||
assert data["offset"] == 1
|
||||
assert len(data["items"]) == 1
|
||||
|
||||
@@ -25,6 +25,7 @@ from agentlightning.types import (
|
||||
SpanContext,
|
||||
TaskInput,
|
||||
TraceStatus,
|
||||
Worker,
|
||||
)
|
||||
|
||||
from .dummy_store import DummyLightningStore
|
||||
@@ -160,6 +161,8 @@ async def test_threaded_store_delegates_all_methods() -> None:
|
||||
last_heartbeat_time=1.5,
|
||||
metadata={"idx": 0},
|
||||
)
|
||||
worker_list = [Worker(worker_id="worker-1", status="busy")]
|
||||
updated_worker = Worker(worker_id="worker-1", status="idle")
|
||||
|
||||
return_values = {
|
||||
"start_rollout": attempted_rollout,
|
||||
@@ -180,6 +183,9 @@ async def test_threaded_store_delegates_all_methods() -> None:
|
||||
"query_spans": [span],
|
||||
"update_rollout": updated_rollout,
|
||||
"update_attempt": updated_attempt,
|
||||
"query_workers": worker_list,
|
||||
"get_worker_by_id": worker_list[0],
|
||||
"update_worker": updated_worker,
|
||||
}
|
||||
|
||||
dummy_store = DummyLightningStore(return_values)
|
||||
@@ -231,6 +237,9 @@ async def test_threaded_store_delegates_all_methods() -> None:
|
||||
)
|
||||
== updated_attempt
|
||||
)
|
||||
assert await threaded_store.query_workers() == worker_list
|
||||
assert await threaded_store.get_worker_by_id("worker-1") == worker_list[0]
|
||||
assert await threaded_store.update_worker("worker-1", heartbeat_stats={"cpu": 0.5}) == updated_worker
|
||||
|
||||
expected_order = [
|
||||
"start_rollout",
|
||||
@@ -251,6 +260,9 @@ async def test_threaded_store_delegates_all_methods() -> None:
|
||||
"query_spans",
|
||||
"update_rollout",
|
||||
"update_attempt",
|
||||
"query_workers",
|
||||
"get_worker_by_id",
|
||||
"update_worker",
|
||||
]
|
||||
assert [name for name, *_ in dummy_store.calls] == expected_order
|
||||
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Optional
|
||||
|
||||
import pytest
|
||||
|
||||
from agentlightning.utils import system_snapshot
|
||||
|
||||
try:
|
||||
import torch # type: ignore
|
||||
|
||||
GPU_AVAILABLE = torch.cuda.is_available()
|
||||
except Exception:
|
||||
GPU_AVAILABLE = False # type: ignore
|
||||
|
||||
|
||||
def _patch_system_snapshot(monkeypatch: pytest.MonkeyPatch, include_gpu: bool = False) -> Optional[SimpleNamespace]:
|
||||
monkeypatch.setattr(system_snapshot.platform, "processor", lambda: "test-cpu")
|
||||
monkeypatch.setattr(system_snapshot.platform, "platform", lambda: "test-platform")
|
||||
monkeypatch.setattr(system_snapshot.socket, "gethostname", lambda: "test-host")
|
||||
|
||||
def fake_cpu_count(logical: bool = True) -> int:
|
||||
return 4 if logical else 2
|
||||
|
||||
monkeypatch.setattr(system_snapshot.psutil, "cpu_count", fake_cpu_count)
|
||||
monkeypatch.setattr(system_snapshot.psutil, "cpu_percent", lambda _: 33.3) # type: ignore
|
||||
|
||||
vm = SimpleNamespace(used=5 * (2**30), total=10 * (2**30), percent=50.0)
|
||||
monkeypatch.setattr(system_snapshot.psutil, "virtual_memory", lambda: vm)
|
||||
|
||||
du = SimpleNamespace(used=2 * (2**30), total=8 * (2**30), percent=25.0)
|
||||
monkeypatch.setattr(system_snapshot.psutil, "disk_usage", lambda _: du) # type: ignore
|
||||
|
||||
net = SimpleNamespace(bytes_sent=4 * (2**20), bytes_recv=6 * (2**20))
|
||||
monkeypatch.setattr(system_snapshot.psutil, "net_io_counters", lambda: net)
|
||||
|
||||
if include_gpu:
|
||||
dummy_gpu = SimpleNamespace(
|
||||
name="Test GPU",
|
||||
utilization=37.5,
|
||||
memory_used=1024,
|
||||
memory_total=4096,
|
||||
temperature=68,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
system_snapshot.GPUStatCollection,
|
||||
"new_query",
|
||||
lambda *args, **kwargs: SimpleNamespace(gpus=[dummy_gpu]), # type: ignore
|
||||
)
|
||||
|
||||
return dummy_gpu
|
||||
return None
|
||||
|
||||
|
||||
def test_system_snapshot_excludes_gpu_by_default(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_patch_system_snapshot(monkeypatch)
|
||||
|
||||
snapshot = system_snapshot.system_snapshot()
|
||||
|
||||
assert snapshot["cpu_name"] == "test-cpu"
|
||||
assert snapshot["cpu_cores"] == 2
|
||||
assert snapshot["cpu_threads"] == 4
|
||||
assert snapshot["cpu_usage_pct"] == 33.3
|
||||
assert snapshot["mem_used_gb"] == 5.0
|
||||
assert snapshot["mem_total_gb"] == 10.0
|
||||
assert snapshot["mem_pct"] == 50.0
|
||||
assert snapshot["disk_used_gb"] == 2.0
|
||||
assert snapshot["disk_total_gb"] == 8.0
|
||||
assert snapshot["disk_pct"] == 25.0
|
||||
assert snapshot["bytes_sent_mb"] == 4.0
|
||||
assert snapshot["bytes_recv_mb"] == 6.0
|
||||
assert snapshot["host"] == "test-host"
|
||||
assert snapshot["os"] == "test-platform"
|
||||
assert "gpus" not in snapshot
|
||||
|
||||
|
||||
def test_system_snapshot_includes_gpus_when_requested(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
dummy_gpu = _patch_system_snapshot(monkeypatch, include_gpu=True)
|
||||
assert dummy_gpu is not None
|
||||
|
||||
snapshot = system_snapshot.system_snapshot(include_gpu=True)
|
||||
|
||||
assert snapshot["gpus"] == [
|
||||
{
|
||||
"gpu": dummy_gpu.name,
|
||||
"util_pct": dummy_gpu.utilization,
|
||||
"mem_used_mb": dummy_gpu.memory_used,
|
||||
"mem_total_mb": dummy_gpu.memory_total,
|
||||
"temp_c": dummy_gpu.temperature,
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_sanity_check() -> None:
|
||||
snapshot = system_snapshot.system_snapshot()
|
||||
assert snapshot is not None
|
||||
|
||||
snapshot = system_snapshot.system_snapshot(include_gpu=True)
|
||||
assert snapshot is not None
|
||||
|
||||
if GPU_AVAILABLE:
|
||||
assert snapshot["gpus"] is not None
|
||||
assert len(snapshot["gpus"]) > 0
|
||||
@@ -130,6 +130,7 @@ dependencies = [
|
||||
{ name = "aiohttp", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
{ name = "fastapi", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
{ name = "flask", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
{ name = "gpustat", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
{ name = "graphviz", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
{ name = "gunicorn", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
{ name = "httpdbg", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
@@ -348,6 +349,7 @@ requires-dist = [
|
||||
{ name = "aiohttp" },
|
||||
{ name = "fastapi" },
|
||||
{ name = "flask" },
|
||||
{ name = "gpustat" },
|
||||
{ name = "graphviz" },
|
||||
{ name = "gunicorn" },
|
||||
{ name = "httpdbg" },
|
||||
@@ -1482,6 +1484,18 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/b0/5c/dbd00727a3dd165d7e0e8af40e630cd7e45d77b525a3218afaff8a87358e/blake3-1.0.8-cp314-cp314t-win_amd64.whl", hash = "sha256:421b99cdf1ff2d1bf703bc56c454f4b286fce68454dd8711abbcb5a0df90c19a", size = 215133, upload-time = "2025-10-14T06:47:16.069Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "blessed"
|
||||
version = "1.23.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "wcwidth", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/c6/70/057c35a79a2015a6e35e45101d710b7a84b92af0c1104399bb83d33cd0d4/blessed-1.23.0.tar.gz", hash = "sha256:56591a32966f704f6131f1400af4151d9e8f5f4144133a5ca034019763dee77b", size = 6745236, upload-time = "2025-11-03T02:50:12.633Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/bc/f4/668909d1273be078ce5fb9a6d75bcfd3bf0832f238b4b09e0eb020d0eec9/blessed-1.23.0-py3-none-any.whl", hash = "sha256:4c432dcde0d45112372d1d096b2c4c0a6a5db1b94d546124872d0c3e64b5ea26", size = 95330, upload-time = "2025-11-03T02:50:10.64Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "blinker"
|
||||
version = "1.9.0"
|
||||
@@ -3178,6 +3192,17 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/86/f1/62a193f0227cf15a920390abe675f386dec35f7ae3ffe6da582d3ade42c7/googleapis_common_protos-1.70.0-py3-none-any.whl", hash = "sha256:b8bfcca8c25a2bb253e0e0b0adaf8c00773e5e6af6fd92397576680b807e0fd8", size = 294530, upload-time = "2025-04-14T10:17:01.271Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "gpustat"
|
||||
version = "1.1.1"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "blessed", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
{ name = "nvidia-ml-py", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
{ name = "psutil", marker = "sys_platform == 'linux' or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-core-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-tinker') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-core-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-core-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-tinker' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-cu128') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-legacy') or (extra == 'group-14-agentlightning-torch-cpu' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-gpu-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-gpu-legacy' and extra == 'group-14-agentlightning-trl') or (extra == 'group-14-agentlightning-torch-gpu-stable' and extra == 'group-14-agentlightning-torch-legacy') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-torch-stable') or (extra == 'group-14-agentlightning-torch-legacy' and extra == 'group-14-agentlightning-trl')" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/79/c4/46d005aec3bf911cb030467d91e062a5386ff4a03e51874424cacc0f60c1/gpustat-1.1.1.tar.gz", hash = "sha256:c18d3ed5518fc16300c42d694debc70aebb3be55cae91f1db64d63b5fa8af9d8", size = 98052, upload-time = "2023-08-22T19:39:06.062Z" }
|
||||
|
||||
[[package]]
|
||||
name = "graphviz"
|
||||
version = "0.21"
|
||||
@@ -6695,6 +6720,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/2f/d8/a6b0d0d0c2435e9310f3e2bb0d9c9dd4c33daef86aa5f30b3681defd37ea/nvidia_cusparselt_cu12-0.7.1-py3-none-win_amd64.whl", hash = "sha256:f67fbb5831940ec829c9117b7f33807db9f9678dc2a617fbe781cac17b4e1075", size = 271020911, upload-time = "2025-02-26T00:14:47.204Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "nvidia-ml-py"
|
||||
version = "13.580.82"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/dd/6c/4a533f2c0185027c465adb6063086bc3728301e95f483665bfa9ebafb2d3/nvidia_ml_py-13.580.82.tar.gz", hash = "sha256:0c028805dc53a0e2a6985ea801888197765ac2ef8f1c9e29a7bf0d3616a5efc7", size = 47999, upload-time = "2025-09-11T16:44:56.267Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/7f/96/d6d25a4c307d6645f4a9b91d620c0151c544ad38b5e371313a87d2761004/nvidia_ml_py-13.580.82-py3-none-any.whl", hash = "sha256:4361db337b0c551e2d101936dae2e9a60f957af26818e8c0c3a1f32b8db8d0a7", size = 49008, upload-time = "2025-09-11T16:44:54.915Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "nvidia-nccl-cu12"
|
||||
version = "2.26.2"
|
||||
|
||||
Reference in New Issue
Block a user