Normalize Store FastAPI (#241)

This commit is contained in:
Yuge Zhang
2025-10-29 21:43:37 +08:00
committed by GitHub
parent f8c45b6ca8
commit 5f67bfe137
2 changed files with 240 additions and 160 deletions
+181 -147
View File
@@ -9,14 +9,14 @@ import threading
import time
import traceback
from contextlib import suppress
from typing import Any, Awaitable, Callable, Dict, List, Literal, Optional, Sequence, Union
from typing import Any, Awaitable, Callable, Dict, List, Literal, Optional, Sequence
import aiohttp
import uvicorn
from fastapi import FastAPI, Request, Response
from fastapi import Body, FastAPI, HTTPException, Request, Response
from fastapi.responses import JSONResponse
from opentelemetry.sdk.trace import ReadableSpan
from pydantic import BaseModel, Field
from pydantic import BaseModel, TypeAdapter
from agentlightning.types import (
Attempt,
@@ -35,9 +35,7 @@ from .base import UNSET, LightningStore, Unset
logger = logging.getLogger(__name__)
class PydanticUnset(BaseModel):
_type: Literal["UNSET"] = "UNSET"
AGL_API_V1_PREFIX = "/agl/v1"
class RolloutRequest(BaseModel):
@@ -58,31 +56,29 @@ class WaitForRolloutsRequest(BaseModel):
timeout: Optional[float] = None
class RolloutId(BaseModel):
class NextSequenceIdRequest(BaseModel):
rollout_id: str
attempt_id: str
class AddResourcesRequest(BaseModel):
resources: NamedResources
class NextSequenceIdResponse(BaseModel):
sequence_id: int
class UpdateRolloutRequest(BaseModel):
rollout_id: str
input: Union[TaskInput, PydanticUnset] = Field(default_factory=PydanticUnset)
mode: Union[Optional[Literal["train", "val", "test"]], PydanticUnset] = Field(default_factory=PydanticUnset)
resources_id: Union[Optional[str], PydanticUnset] = Field(default_factory=PydanticUnset)
status: Union[RolloutStatus, PydanticUnset] = Field(default_factory=PydanticUnset)
config: Union[RolloutConfig, PydanticUnset] = Field(default_factory=PydanticUnset)
metadata: Union[Dict[str, Any], PydanticUnset] = Field(default_factory=PydanticUnset)
input: Optional[TaskInput] = None
mode: Optional[Literal["train", "val", "test"]] = None
resources_id: Optional[str] = None
status: Optional[RolloutStatus] = None
config: Optional[RolloutConfig] = None
metadata: Optional[Dict[str, Any]] = None
class UpdateAttemptRequest(BaseModel):
rollout_id: str
attempt_id: Union[str, Literal["latest"]]
status: Union[AttemptStatus, PydanticUnset] = Field(default_factory=PydanticUnset)
worker_id: Union[str, PydanticUnset] = Field(default_factory=PydanticUnset)
last_heartbeat_time: Union[float, PydanticUnset] = Field(default_factory=PydanticUnset)
metadata: Union[Dict[str, Any], PydanticUnset] = Field(default_factory=PydanticUnset)
status: Optional[AttemptStatus] = None
worker_id: Optional[str] = None
last_heartbeat_time: Optional[float] = None
metadata: Optional[Dict[str, Any]] = None
class LightningStoreServer(LightningStore):
@@ -213,7 +209,7 @@ class LightningStoreServer(LightningStore):
while time.time() - current_time < 10:
async with aiohttp.ClientSession() as session:
with suppress(Exception):
async with session.get(f"{self.endpoint}/health") as response:
async with session.get(f"{self.endpoint}{AGL_API_V1_PREFIX}/health") as response:
if response.status == 200:
return True
await asyncio.sleep(0.1)
@@ -307,30 +303,43 @@ class LightningStoreServer(LightningStore):
"""Set up FastAPI routes for all store operations."""
assert self.app is not None
@self.app.exception_handler(Exception)
async def _app_exception_handler(request: Request, exc: Exception): # pyright: ignore[reportUnusedFunction]
@self.app.middleware("http")
async def _app_exception_handler( # pyright: ignore[reportUnusedFunction]
request: Request, call_next: Callable[[Request], Awaitable[Response]]
) -> Response:
"""
Convert unhandled application exceptions into 400 responses.
Convert unhandled application exceptions into 500 responses.
Only covers /agl/v1 requests.
- Client needs a reliable signal to distinguish "app bug / bad request"
from transport/session failures.
- 400 here means "do not retry"; network issues will surface as aiohttp
- 400 means "do not retry"; network issues will surface as aiohttp
exceptions or 5xx and will be retried by the client shield.
"""
logger.exception("Unhandled application error", exc_info=exc)
return JSONResponse(
status_code=400,
content={
"detail": str(exc),
"error_type": type(exc).__name__,
"traceback": traceback.format_exc(),
},
)
try:
return await call_next(request)
except Exception as exc:
# decide whether to convert this into your 400 JSONResponse
if request.url.path.startswith(AGL_API_V1_PREFIX):
logger.exception("Unhandled application error", exc_info=exc)
payload = {
"detail": "Internal server error",
"error_type": type(exc).__name__,
"traceback": traceback.format_exc(),
}
# 500 so clients can decide to retry
return JSONResponse(status_code=500, content=payload)
# otherwise re-raise and let FastAPI/Starlette handle it (500 or other handlers)
raise
@self.app.middleware("http")
async def _log_time( # pyright: ignore[reportUnusedFunction]
request: Request, call_next: Callable[[Request], Awaitable[Response]]
):
if not request.url.path.startswith("/agl/v1/"):
return await call_next(request)
start = time.perf_counter()
response = await call_next(request)
duration = (time.perf_counter() - start) * 1000
@@ -346,21 +355,11 @@ class LightningStoreServer(LightningStore):
)
return response
@self.app.get("/health")
@self.app.get(AGL_API_V1_PREFIX + "/health")
async def health(): # pyright: ignore[reportUnusedFunction]
return {"status": "ok"}
@self.app.post("/start_rollout", response_model=AttemptedRollout)
async def start_rollout(request: RolloutRequest): # pyright: ignore[reportUnusedFunction]
return await self.start_rollout(
input=request.input,
mode=request.mode,
resources_id=request.resources_id,
config=request.config,
metadata=request.metadata,
)
@self.app.post("/enqueue_rollout", response_model=Rollout)
@self.app.post(AGL_API_V1_PREFIX + "/queues/rollouts/enqueue", status_code=201, response_model=Rollout)
async def enqueue_rollout(request: RolloutRequest): # pyright: ignore[reportUnusedFunction]
return await self.enqueue_rollout(
input=request.input,
@@ -370,88 +369,121 @@ class LightningStoreServer(LightningStore):
metadata=request.metadata,
)
@self.app.get("/dequeue_rollout", response_model=Optional[AttemptedRollout])
@self.app.post(AGL_API_V1_PREFIX + "/queues/rollouts/dequeue", response_model=Optional[AttemptedRollout])
async def dequeue_rollout(): # pyright: ignore[reportUnusedFunction]
return await self.dequeue_rollout()
@self.app.post("/start_attempt", response_model=AttemptedRollout)
async def start_attempt(request: RolloutId): # pyright: ignore[reportUnusedFunction]
return await self.start_attempt(request.rollout_id)
@self.app.post(AGL_API_V1_PREFIX + "/rollouts", status_code=201, response_model=AttemptedRollout)
async def start_rollout(request: RolloutRequest): # pyright: ignore[reportUnusedFunction]
return await self.start_rollout(
input=request.input,
mode=request.mode,
resources_id=request.resources_id,
config=request.config,
metadata=request.metadata,
)
@self.app.post("/query_rollouts", response_model=List[Rollout])
async def query_rollouts(request: QueryRolloutsRequest): # pyright: ignore[reportUnusedFunction]
@self.app.get(AGL_API_V1_PREFIX + "/rollouts", response_model=List[Rollout])
async def query_rollouts(): # pyright: ignore[reportUnusedFunction]
return await self.query_rollouts()
@self.app.post(AGL_API_V1_PREFIX + "/rollouts/search", response_model=List[Rollout])
async def search_rollouts(request: QueryRolloutsRequest): # pyright: ignore[reportUnusedFunction]
return await self.query_rollouts(status=request.status, rollout_ids=request.rollout_ids)
@self.app.get("/query_attempts/{rollout_id}", response_model=List[Attempt])
async def query_attempts(rollout_id: str): # pyright: ignore[reportUnusedFunction]
return await self.query_attempts(rollout_id)
@self.app.get("/get_latest_attempt/{rollout_id}", response_model=Optional[Attempt])
async def get_latest_attempt(rollout_id: str): # pyright: ignore[reportUnusedFunction]
return await self.get_latest_attempt(rollout_id)
@self.app.get("/get_rollout_by_id/{rollout_id}", response_model=Optional[Rollout])
@self.app.get(AGL_API_V1_PREFIX + "/rollouts/{rollout_id}", response_model=Rollout)
async def get_rollout_by_id(rollout_id: str): # pyright: ignore[reportUnusedFunction]
return await self.get_rollout_by_id(rollout_id)
@self.app.post("/add_resources", response_model=ResourcesUpdate)
async def add_resources(resources: AddResourcesRequest): # pyright: ignore[reportUnusedFunction]
return await self.add_resources(resources.resources)
def _get_mandatory_field_or_unset(request: BaseModel, 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 field in request.model_fields_set:
value = getattr(request, field)
if value is None:
raise HTTPException(status_code=400, detail=f"{field} is invalid; it cannot be a null value.")
return value
else:
return UNSET
@self.app.post("/update_resources", response_model=ResourcesUpdate)
async def update_resources(update: ResourcesUpdate): # pyright: ignore[reportUnusedFunction]
return await self.update_resources(update.resources_id, update.resources)
@self.app.post(AGL_API_V1_PREFIX + "/rollouts/{rollout_id}", response_model=Rollout)
async def update_rollout( # pyright: ignore[reportUnusedFunction]
rollout_id: str, request: UpdateRolloutRequest = Body(...)
):
return await self.update_rollout(
rollout_id=rollout_id,
input=request.input if "input" in request.model_fields_set else UNSET,
mode=request.mode if "mode" in request.model_fields_set else UNSET,
resources_id=request.resources_id if "resources_id" in request.model_fields_set else UNSET,
status=_get_mandatory_field_or_unset(request, "status"),
config=_get_mandatory_field_or_unset(request, "config"),
metadata=request.metadata if "metadata" in request.model_fields_set else UNSET,
)
@self.app.get("/get_resources_by_id/{resources_id}", response_model=Optional[ResourcesUpdate])
async def get_resources_by_id(resources_id: str): # pyright: ignore[reportUnusedFunction]
return await self.get_resources_by_id(resources_id)
@self.app.post(
AGL_API_V1_PREFIX + "/rollouts/{rollout_id}/attempts", status_code=201, response_model=AttemptedRollout
)
async def start_attempt(rollout_id: str): # pyright: ignore[reportUnusedFunction]
return await self.start_attempt(rollout_id)
@self.app.get("/get_latest_resources", response_model=Optional[ResourcesUpdate])
@self.app.post(AGL_API_V1_PREFIX + "/rollouts/{rollout_id}/attempts/{attempt_id}", response_model=Attempt)
async def update_attempt( # pyright: ignore[reportUnusedFunction]
rollout_id: str, attempt_id: str, request: UpdateAttemptRequest = Body(...)
):
return await self.update_attempt(
rollout_id=rollout_id,
attempt_id=attempt_id,
status=_get_mandatory_field_or_unset(request, "status"),
worker_id=_get_mandatory_field_or_unset(request, "worker_id"),
last_heartbeat_time=_get_mandatory_field_or_unset(request, "last_heartbeat_time"),
metadata=_get_mandatory_field_or_unset(request, "metadata"),
)
@self.app.get(AGL_API_V1_PREFIX + "/rollouts/{rollout_id}/attempts", response_model=List[Attempt])
async def query_attempts(rollout_id: str): # pyright: ignore[reportUnusedFunction]
return await self.query_attempts(rollout_id)
@self.app.get(AGL_API_V1_PREFIX + "/rollouts/{rollout_id}/attempts/latest", response_model=Optional[Attempt])
async def get_latest_attempt(rollout_id: str): # pyright: ignore[reportUnusedFunction]
return await self.get_latest_attempt(rollout_id)
@self.app.post(AGL_API_V1_PREFIX + "/resources", status_code=201, response_model=ResourcesUpdate)
async def add_resources(resources: NamedResources): # pyright: ignore[reportUnusedFunction]
return await self.add_resources(resources)
@self.app.get(AGL_API_V1_PREFIX + "/resources/latest", response_model=Optional[ResourcesUpdate])
async def get_latest_resources(): # pyright: ignore[reportUnusedFunction]
return await self.get_latest_resources()
@self.app.post("/add_span", response_model=Span)
@self.app.post(AGL_API_V1_PREFIX + "/resources/{resources_id}", response_model=ResourcesUpdate)
async def update_resources( # pyright: ignore[reportUnusedFunction]
resources_id: str, resources: NamedResources
):
return await self.update_resources(resources_id, resources)
@self.app.get(AGL_API_V1_PREFIX + "/resources/{resources_id}", response_model=Optional[ResourcesUpdate])
async def get_resources_by_id(resources_id: str): # pyright: ignore[reportUnusedFunction]
return await self.get_resources_by_id(resources_id)
@self.app.post(AGL_API_V1_PREFIX + "/spans", status_code=201, response_model=Span)
async def add_span(span: Span): # pyright: ignore[reportUnusedFunction]
return await self.add_span(span)
@self.app.get("/get_next_span_sequence_id/{rollout_id}/{attempt_id}", response_model=int)
async def get_next_span_sequence_id(rollout_id: str, attempt_id: str): # pyright: ignore[reportUnusedFunction]
return await self.get_next_span_sequence_id(rollout_id, attempt_id)
@self.app.post("/wait_for_rollouts", response_model=List[Rollout])
async def wait_for_rollouts(request: WaitForRolloutsRequest): # pyright: ignore[reportUnusedFunction]
return await self.wait_for_rollouts(rollout_ids=request.rollout_ids, timeout=request.timeout)
@self.app.get("/query_spans/{rollout_id}", response_model=List[Span])
@self.app.get(AGL_API_V1_PREFIX + "/spans", response_model=List[Span])
async def query_spans( # pyright: ignore[reportUnusedFunction]
rollout_id: str, attempt_id: Optional[str] = None
rollout_id: str,
attempt_id: Optional[str] = None,
):
return await self.query_spans(rollout_id, attempt_id)
@self.app.post("/update_rollout", response_model=Rollout)
async def update_rollout(request: UpdateRolloutRequest): # pyright: ignore[reportUnusedFunction]
return await self.update_rollout(
rollout_id=request.rollout_id,
input=request.input if not isinstance(request.input, PydanticUnset) else UNSET,
mode=request.mode if not isinstance(request.mode, PydanticUnset) else UNSET,
resources_id=request.resources_id if not isinstance(request.resources_id, PydanticUnset) else UNSET,
status=request.status if not isinstance(request.status, PydanticUnset) else UNSET,
config=request.config if not isinstance(request.config, PydanticUnset) else UNSET,
metadata=request.metadata if not isinstance(request.metadata, PydanticUnset) else UNSET,
)
@self.app.post(AGL_API_V1_PREFIX + "/spans/next", response_model=NextSequenceIdResponse)
async def get_next_span_sequence_id(request: NextSequenceIdRequest): # pyright: ignore[reportUnusedFunction]
sequence_id = await self.get_next_span_sequence_id(request.rollout_id, request.attempt_id)
return NextSequenceIdResponse(sequence_id=sequence_id)
@self.app.post("/update_attempt", response_model=Attempt)
async def update_attempt(request: UpdateAttemptRequest): # pyright: ignore[reportUnusedFunction]
return await self.update_attempt(
rollout_id=request.rollout_id,
attempt_id=request.attempt_id,
status=request.status if not isinstance(request.status, PydanticUnset) else UNSET,
worker_id=request.worker_id if not isinstance(request.worker_id, PydanticUnset) else UNSET,
last_heartbeat_time=(
request.last_heartbeat_time if not isinstance(request.last_heartbeat_time, PydanticUnset) else UNSET
),
metadata=request.metadata if not isinstance(request.metadata, PydanticUnset) else UNSET,
)
@self.app.post(AGL_API_V1_PREFIX + "/waits/rollouts", response_model=List[Rollout])
async def wait_for_rollouts(request: WaitForRolloutsRequest): # pyright: ignore[reportUnusedFunction]
return await self.wait_for_rollouts(rollout_ids=request.rollout_ids, timeout=request.timeout)
# Delegate methods
async def _call_store_method(self, method_name: str, *args: Any, **kwargs: Any) -> Any:
@@ -623,7 +655,7 @@ class LightningStoreClient(LightningStore):
retry_delays: Sequence[float] = (1.0, 2.0, 5.0),
health_retry_delays: Sequence[float] = (0.1, 0.2, 0.5),
):
self.server_address = server_address.rstrip("/")
self.server_address = server_address.rstrip("/") + AGL_API_V1_PREFIX
self._sessions: Dict[int, aiohttp.ClientSession] = {} # id(loop) -> ClientSession
self._lock = threading.RLock()
@@ -717,6 +749,7 @@ class LightningStoreClient(LightningStore):
path: str,
*,
json: Any | None = None,
params: Dict[str, Any] | None = None,
) -> Any:
"""
Make an HTTP request with:
@@ -742,11 +775,11 @@ class LightningStoreClient(LightningStore):
await asyncio.sleep(delay)
try:
http_call = getattr(session, method)
async with http_call(url, json=json) as resp:
async with http_call(url, json=json, params=params) as resp:
resp.raise_for_status()
return await resp.json()
except aiohttp.ClientResponseError as cre:
# Respect app-level 4xx as final (server marks app faults as 400)
# Respect app-level 4xx as final
# 4xx => application issue; do not retry (except 408 which is transient)
logger.debug(f"ClientResponseError: {cre.status} {cre.message}", exc_info=True)
if 400 <= cre.status < 500 and cre.status != 408:
@@ -805,7 +838,7 @@ class LightningStoreClient(LightningStore):
) -> AttemptedRollout:
data = await self._request_json(
"post",
"/start_rollout",
"/rollouts",
json=RolloutRequest(
input=input,
mode=mode,
@@ -826,7 +859,7 @@ class LightningStoreClient(LightningStore):
) -> Rollout:
data = await self._request_json(
"post",
"/enqueue_rollout",
"/queues/rollouts/enqueue",
json=RolloutRequest(
input=input,
mode=mode,
@@ -849,9 +882,9 @@ class LightningStoreClient(LightningStore):
server error, etc.), it logs the error and returns None immediately.
"""
session = await self._get_session()
url = f"{self.server_address}/dequeue_rollout"
url = f"{self.server_address}/queues/rollouts/dequeue"
try:
async with session.get(url) as resp:
async with session.post(url) as resp:
resp.raise_for_status()
data = await resp.json()
self._dequeue_was_successful = True
@@ -868,26 +901,25 @@ class LightningStoreClient(LightningStore):
async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
data = await self._request_json(
"post",
"/start_attempt",
json=RolloutId(rollout_id=rollout_id).model_dump(),
f"/rollouts/{rollout_id}/attempts",
)
return AttemptedRollout.model_validate(data)
async def query_rollouts(
self, *, status: Optional[Sequence[RolloutStatus]] = None, rollout_ids: Optional[Sequence[str]] = None
) -> List[Rollout]:
data = await self._request_json(
"post",
"/query_rollouts",
json=QueryRolloutsRequest(
if status or rollout_ids:
payload = QueryRolloutsRequest(
status=list(status) if status else None,
rollout_ids=list(rollout_ids) if rollout_ids else None,
).model_dump(),
)
).model_dump(exclude_none=True)
data = await self._request_json("post", "/rollouts/search", json=payload)
else:
data = await self._request_json("get", "/rollouts")
return [Rollout.model_validate(item) for item in data]
async def query_attempts(self, rollout_id: str) -> List[Attempt]:
data = await self._request_json("get", f"/query_attempts/{rollout_id}")
data = await self._request_json("get", f"/rollouts/{rollout_id}/attempts")
return [Attempt.model_validate(item) for item in data]
async def get_latest_attempt(self, rollout_id: str) -> Optional[Attempt]:
@@ -905,7 +937,7 @@ class LightningStoreClient(LightningStore):
If all retries fail, it logs the error and returns None instead of raising an exception.
"""
try:
data = await self._request_json("get", f"/get_latest_attempt/{rollout_id}")
data = await self._request_json("get", f"/rollouts/{rollout_id}/attempts/latest")
return Attempt.model_validate(data) if data else None
except Exception as e:
logger.error(f"get_latest_attempt failed after all retries for rollout_id={rollout_id}: {e}", exc_info=True)
@@ -926,22 +958,19 @@ class LightningStoreClient(LightningStore):
If all retries fail, it logs the error and returns None instead of raising an exception.
"""
try:
data = await self._request_json("get", f"/get_rollout_by_id/{rollout_id}")
data = await self._request_json("get", f"/rollouts/{rollout_id}")
return Rollout.model_validate(data) if data else None
except Exception as e:
logger.error(f"get_rollout_by_id failed after all retries for rollout_id={rollout_id}: {e}", exc_info=True)
return None
async def add_resources(self, resources: NamedResources) -> ResourcesUpdate:
request = AddResourcesRequest(resources=resources)
data = await self._request_json("post", "/add_resources", json=request.model_dump())
data = await self._request_json("post", "/resources", json=TypeAdapter(NamedResources).dump_python(resources))
return ResourcesUpdate.model_validate(data)
async def update_resources(self, resources_id: str, resources: NamedResources) -> ResourcesUpdate:
data = await self._request_json(
"post",
"/update_resources",
json=ResourcesUpdate(resources_id=resources_id, resources=resources).model_dump(),
"post", f"/resources/{resources_id}", json=TypeAdapter(NamedResources).dump_python(resources)
)
return ResourcesUpdate.model_validate(data)
@@ -960,7 +989,7 @@ class LightningStoreClient(LightningStore):
If all retries fail, it logs the error and returns None instead of raising an exception.
"""
try:
data = await self._request_json("get", f"/get_resources_by_id/{resources_id}")
data = await self._request_json("get", f"/resources/{resources_id}")
return ResourcesUpdate.model_validate(data) if data else None
except Exception as e:
logger.error(
@@ -980,20 +1009,24 @@ class LightningStoreClient(LightningStore):
If all retries fail, it logs the error and returns None instead of raising an exception.
"""
try:
data = await self._request_json("get", "/get_latest_resources")
data = await self._request_json("get", "/resources/latest")
return ResourcesUpdate.model_validate(data) if data else None
except Exception as e:
logger.error(f"get_latest_resources failed after all retries: {e}", exc_info=True)
return None
async def add_span(self, span: Span) -> Span:
data = await self._request_json("post", "/add_span", json=span.model_dump(mode="json"))
data = await self._request_json("post", "/spans", json=span.model_dump(mode="json"))
return Span.model_validate(data)
async def get_next_span_sequence_id(self, rollout_id: str, attempt_id: str) -> int:
data = await self._request_json("get", f"/get_next_span_sequence_id/{rollout_id}/{attempt_id}")
# endpoint returns a plain JSON number
return int(data)
data = await self._request_json(
"post",
"/spans/next",
json=NextSequenceIdRequest(rollout_id=rollout_id, attempt_id=attempt_id).model_dump(),
)
response = NextSequenceIdResponse.model_validate(data)
return response.sequence_id
async def add_otel_span(
self,
@@ -1030,7 +1063,7 @@ class LightningStoreClient(LightningStore):
)
data = await self._request_json(
"post",
"/wait_for_rollouts",
"/waits/rollouts",
json=WaitForRolloutsRequest(rollout_ids=rollout_ids, timeout=timeout).model_dump(),
)
return [Rollout.model_validate(item) for item in data]
@@ -1040,10 +1073,10 @@ class LightningStoreClient(LightningStore):
rollout_id: str,
attempt_id: str | Literal["latest"] | None = None,
) -> List[Span]:
path = f"/query_spans/{rollout_id}"
params: Dict[str, str] = {"rollout_id": rollout_id}
if attempt_id is not None:
path += f"?attempt_id={attempt_id}"
data = await self._request_json("get", path)
params["attempt_id"] = attempt_id
data = await self._request_json("get", "/spans", params=params)
return [Span.model_validate(item) for item in data]
async def update_rollout(
@@ -1056,7 +1089,7 @@ class LightningStoreClient(LightningStore):
config: RolloutConfig | Unset = UNSET,
metadata: Optional[Dict[str, Any]] | Unset = UNSET,
) -> Rollout:
payload: Dict[str, Any] = {"rollout_id": rollout_id}
payload: Dict[str, Any] = {}
if not isinstance(input, Unset):
payload["input"] = input
if not isinstance(mode, Unset):
@@ -1070,7 +1103,7 @@ class LightningStoreClient(LightningStore):
if not isinstance(metadata, Unset):
payload["metadata"] = metadata
data = await self._request_json("post", "/update_rollout", json=payload)
data = await self._request_json("post", f"/rollouts/{rollout_id}", json=payload)
return Rollout.model_validate(data)
async def update_attempt(
@@ -1082,10 +1115,7 @@ class LightningStoreClient(LightningStore):
last_heartbeat_time: float | Unset = UNSET,
metadata: Optional[Dict[str, Any]] | Unset = UNSET,
) -> Attempt:
payload: Dict[str, Any] = {
"rollout_id": rollout_id,
"attempt_id": attempt_id,
}
payload: Dict[str, Any] = {}
if not isinstance(status, Unset):
payload["status"] = status
if not isinstance(worker_id, Unset):
@@ -1095,5 +1125,9 @@ class LightningStoreClient(LightningStore):
if not isinstance(metadata, Unset):
payload["metadata"] = metadata
data = await self._request_json("post", "/update_attempt", json=payload)
data = await self._request_json(
"post",
f"/rollouts/{rollout_id}/attempts/{attempt_id}",
json=payload,
)
return Attempt.model_validate(data)
+59 -13
View File
@@ -5,7 +5,7 @@ import contextlib
import multiprocessing
import socket
import sys
from typing import Any, AsyncGenerator, Tuple, cast
from typing import Any, AsyncGenerator, Dict, Tuple, cast
from unittest.mock import patch
import aiohttp
@@ -137,6 +137,9 @@ async def test_add_resources_via_server(server_client: Tuple[LightningStoreServe
assert isinstance(retrieved.resources["main_llm"], LLM)
assert retrieved.resources["main_llm"].model == "test-model"
assert isinstance(retrieved.resources["greeting"], PromptTemplate)
assert retrieved.resources["greeting"].template == "Hello {name}!"
# Verify it's set as latest
latest = await server.get_latest_resources()
assert latest is not None
@@ -326,6 +329,28 @@ async def test_update_rollout_none_vs_unset(server_client: Tuple[LightningStoreS
assert preserved.status == "running"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"bad_payload",
[
{"status": None},
{"config": None},
],
)
async def test_update_rollout_rejects_none_values(
server_client: Tuple[LightningStoreServer, LightningStoreClient],
bad_payload: Dict[str, Any],
) -> None:
_, client = server_client
attempted = await client.start_rollout(input={"payload": "bad-none"})
with pytest.raises((ClientResponseError, AttributeError)) as exc_info:
await client.update_rollout(attempted.rollout_id, **bad_payload)
if isinstance(exc_info.value, ClientResponseError):
assert exc_info.value.status == 400
@pytest.mark.asyncio
async def test_update_attempt_none_vs_unset(server_client: Tuple[LightningStoreServer, LightningStoreClient]) -> None:
_, client = server_client
@@ -355,6 +380,27 @@ async def test_update_attempt_none_vs_unset(server_client: Tuple[LightningStoreS
assert preserved.status == "running"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"bad_payload",
[
{"last_heartbeat_time": None},
{"metadata": None},
],
)
async def test_update_attempt_rejects_none_values(
server_client: Tuple[LightningStoreServer, LightningStoreClient],
bad_payload: Dict[str, Any],
) -> None:
_, client = server_client
attempted = await client.start_rollout(input={"payload": "bad-none-attempt"})
with pytest.raises(ClientResponseError) as exc_info:
await client.update_attempt(attempted.rollout_id, attempted.attempt.attempt_id, **bad_payload)
assert exc_info.value.status == 400
@pytest.mark.asyncio
async def test_concurrent_add_otel_span_sequence_ids_unique(
server_client: Tuple[LightningStoreServer, LightningStoreClient], mock_readable_span: ReadableSpan
@@ -521,11 +567,11 @@ async def test_subprocess_client_operations_work_but_direct_store_access_fails()
@pytest.mark.parametrize(
"status,endpoint,make_app_error",
[
(400, "/enqueue_rollout", True), # server-marked app error -> 400 -> no retry
(404, "/update_rollout", False), # non-408 4xx -> no retry
(400, "/queues/rollouts/enqueue", True), # server-marked app error -> 500 -> retry
(404, "/rollouts/nonexistent", False), # non-408 4xx -> no retry
],
)
async def test_no_retry_on_4xx_application_and_non408(
async def test_retry_on_4xx_application_and_non408(
server_client: Tuple[LightningStoreServer, LightningStoreClient],
monkeypatch: MonkeyPatch,
status: int,
@@ -547,12 +593,12 @@ async def test_no_retry_on_4xx_application_and_non408(
with pytest.raises(ClientResponseError) as ei:
await client.enqueue_rollout(input={"origin": "should-fail"})
assert ei.value.status == 400
assert call_count["n"] == 1
assert ei.value.status == 500
assert call_count["n"] == 4
monkeypatch.setattr(server.store, "enqueue_rollout", original, raising=True)
else:
# Raise 404 once for /update_rollout; client must not retry.
# Raise 404 once for /rollouts/nonexistent; client must not retry.
original_post = aiohttp.ClientSession.post
calls = {"n": 0}
@@ -584,7 +630,7 @@ async def test_retry_on_transient_network_errors_then_success(
counters = {"post_calls": 0}
def flaky_post(self: aiohttp.ClientSession, url: Any, *args: Any, **kwargs: Any) -> MockResponse:
if str(url).endswith("/start_rollout"):
if str(url).endswith("/rollouts"):
if counters["post_calls"] == 0:
counters["post_calls"] += 1
# raise chosen transient network exception
@@ -609,7 +655,7 @@ async def test_retry_on_transient_http_status_then_success(
fired = {"once": False}
def post_then_ok(self: aiohttp.ClientSession, url: Any, *args: Any, **kwargs: Any) -> MockResponse:
if str(url).endswith("/enqueue_rollout") and not fired["once"]:
if str(url).endswith("/queues/rollouts/enqueue") and not fired["once"]:
fired["once"] = True
req_info = aiohttp.RequestInfo(
url=URL(str(url)), method="POST", headers=cast(Any, {}), real_url=URL(str(url))
@@ -635,7 +681,7 @@ async def test_unhealthy_health_probe_stops_retries(
post_calls = {"n": 0}
def failing_post(self: aiohttp.ClientSession, url: Any, *args: Any, **kwargs: Any) -> MockResponse:
if str(url).endswith("/start_rollout"):
if str(url).endswith("/rollouts"):
post_calls["n"] += 1
raise ServerDisconnectedError("synthetic disconnect")
return MockResponse(original_post(self, url, *args, **kwargs))
@@ -694,7 +740,7 @@ async def test_retry_mechanism_with_custom_delays_and_health_recovery(
timestamps: list[float] = []
def monitored_post(self: aiohttp.ClientSession, url: Any, *args: Any, **kwargs: Any) -> MockResponse:
if str(url).endswith("/start_rollout"):
if str(url).endswith("/rollouts"):
import time
counters["post_attempts"] += 1
@@ -755,7 +801,7 @@ async def test_client_response_error_with_different_status_codes(
# Test 403 Forbidden - should NOT retry
def post_403(self: aiohttp.ClientSession, url: Any, *args: Any, **kwargs: Any) -> MockResponse:
if str(url).endswith("/start_rollout"):
if str(url).endswith("/rollouts"):
req_info = aiohttp.RequestInfo(
url=URL(str(url)), method="POST", headers=cast(Any, {}), real_url=URL(str(url))
)
@@ -772,7 +818,7 @@ async def test_client_response_error_with_different_status_codes(
call_count = {"n": 0}
def post_503_then_ok(self: aiohttp.ClientSession, url: Any, *args: Any, **kwargs: Any) -> MockResponse:
if str(url).endswith("/enqueue_rollout"):
if str(url).endswith("/queues/rollouts/enqueue"):
call_count["n"] += 1
if call_count["n"] == 1:
req_info = aiohttp.RequestInfo(