Normalize Store FastAPI (#241)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user