Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions agentlightning/config/controller.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,10 @@ k8s_runner:
ttl_after_finished: 1200
max_jobs_per_minute: 100
poll_interval: 5
image_readiness:
enabled: true
heartbeat_seconds: 5
lease_seconds: 30

local_runner:
maximum_size: 50
Expand Down
178 changes: 147 additions & 31 deletions agentlightning/controller/k8s_reconciler.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,22 +12,28 @@
from __future__ import annotations

import asyncio
import json
import time
from collections import deque
from collections.abc import Mapping, Sequence
from typing import Any, cast

import httpx
import kr8s
import kr8s.asyncio
import structlog
import yaml
from jinja2 import Environment
from kr8s.asyncio import objects as k8s_objects
from omegaconf import DictConfig

from agentlightning.client import AgentLightningAsyncClient
from agentlightning.schemas import DEFAULT_ATTEMPT_ID, Rollout, RolloutPatch, RolloutState, RolloutStatusPatch
from agentlightning.k8s import extract_pod_images, normalize_image_reference, render_job_template
from agentlightning.schemas import (
DEFAULT_ATTEMPT_ID,
K8sImageReadinessReport,
Rollout,
RolloutPatch,
RolloutState,
RolloutStatusPatch,
)

log = structlog.get_logger()

Expand All @@ -40,25 +46,49 @@ def build_job_name(rollout_id: str) -> str:
return f"agl-rollout-{rollout_id}"


def images_available_on_all_ready_nodes(
nodes: Sequence[Mapping[str, Any]],
) -> tuple[frozenset[str], int]:
"""Return the image-name intersection for Ready, schedulable nodes."""
eligible: list[Mapping[str, Any]] = []
for node in nodes:
if node.get("spec", {}).get("unschedulable", False):
continue
conditions = node.get("status", {}).get("conditions", [])
if not any(condition.get("type") == "Ready" and condition.get("status") == "True" for condition in conditions):
continue
eligible.append(node)

if not eligible:
raise RuntimeError("no Ready schedulable Kubernetes nodes found")

per_node: list[set[str]] = []
for node in eligible:
names = {
normalize_image_reference(name)
for image in node.get("status", {}).get("images", [])
for name in image.get("names", [])
if isinstance(name, str) and name.strip()
}
per_node.append(names)

common = set(per_node[0])
for names in per_node[1:]:
common.intersection_update(names)
return frozenset(common), len(eligible)


def build_job_spec(rollout: Rollout, controller_config: DictConfig) -> dict[str, Any]:
"""Build a K8s Job manifest from the rollout's complete Jinja2 Job template."""
template = rollout.config.k8s.job_template if rollout.config.k8s else None
if not template:
raise ValueError("invalid rollout config: missing config.k8s.job_template")

env = Environment()
env.filters["yaml_escape"] = lambda value: json.dumps(str(value), ensure_ascii=True)
rendered = env.from_string(template).render(
job = render_job_template(
template,
job_name=build_job_name(rollout.rollout_id),
input=rollout.input,
input_data=rollout.input,
)
docs = [doc for doc in yaml.safe_load_all(rendered) if doc is not None]
if len(docs) != 1:
raise ValueError("invalid rollout config: config.k8s.job_template must render exactly one YAML document")

job = docs[0]
if not isinstance(job, dict) or job.get("kind") != "Job":
raise ValueError("invalid rollout config: config.k8s.job_template must render a Kubernetes Job")

metadata = job.setdefault("metadata", {})
metadata["name"] = build_job_name(rollout.rollout_id)
Expand Down Expand Up @@ -117,6 +147,15 @@ def __init__(self, api: AgentLightningAsyncClient, config: DictConfig) -> None:
self._k8s_api: Any | None = None
self._stop = asyncio.Event()
self._job_creation_timestamps: deque[float] = deque()
self._ready_images: frozenset[str] | None = None
self._ready_images_expires_at = 0.0

readiness_config = self._runner_config.image_readiness
if readiness_config.enabled:
heartbeat = float(readiness_config.heartbeat_seconds)
lease = float(readiness_config.lease_seconds)
if not 0 < heartbeat < lease <= 300:
raise ValueError("k8s_runner.image_readiness requires 0 < heartbeat_seconds < lease_seconds <= 300")

async def _get_k8s_api(self) -> Any:
if self._k8s_api is None:
Expand All @@ -131,17 +170,67 @@ async def run(self) -> None:
poll_interval=self._runner_config.poll_interval,
)
try:
await asyncio.gather(
self._periodic_reconcile_loop(),
self._watch_jobs_loop(),
tasks = []
if self._runner_config.image_readiness.enabled:
try:
await self._publish_image_readiness_once()
except Exception:
log.exception("Initial K8s image readiness publish failed")
tasks.append(self._image_readiness_loop())
tasks.extend(
[
self._periodic_reconcile_loop(),
self._watch_jobs_loop(),
]
)
await asyncio.gather(*tasks)
except asyncio.CancelledError:
log.info("Controller stopped")

def stop(self) -> None:
"""Signal the controller to stop."""
self._stop.set()

async def _scan_preloaded_images(self) -> tuple[frozenset[str], int]:
api = await self._get_k8s_api()
nodes = [cast(k8s_objects.Node, node).raw async for node in k8s_objects.Node.async_list(api=api)]
return images_available_on_all_ready_nodes(nodes)

async def _publish_image_readiness_once(self) -> None:
images, node_count = await self._scan_preloaded_images()
lease_seconds = float(self._runner_config.image_readiness.lease_seconds)
scanned_at = time.monotonic()
report = K8sImageReadinessReport(
images=sorted(images),
node_count=node_count,
lease_seconds=lease_seconds,
)
response = await self._api.put(
"/api/runner-readiness/k8s",
json=report.model_dump(mode="json"),
)
response.raise_for_status()
self._ready_images = images
self._ready_images_expires_at = scanned_at + lease_seconds

async def _image_readiness_loop(self) -> None:
heartbeat = float(self._runner_config.image_readiness.heartbeat_seconds)
while not self._stop.is_set():
try:
await asyncio.wait_for(self._stop.wait(), timeout=heartbeat)
return
except TimeoutError:
pass
try:
await self._publish_image_readiness_once()
except Exception:
log.exception("K8s image readiness publish failed")

def _fresh_ready_images(self) -> frozenset[str] | None:
if self._ready_images is None or time.monotonic() >= self._ready_images_expires_at:
return None
return self._ready_images

# --- Periodic reconcile ---

async def _periodic_reconcile_loop(self) -> None:
Expand Down Expand Up @@ -242,27 +331,54 @@ async def _reconcile_once(self) -> None:
async def _create_job(self, rollout: Rollout) -> None:
"""Create a K8s Job for a queuing rollout without changing rollout state."""
job_name = build_job_name(rollout.rollout_id)
now = time.monotonic()
window_start = now - JOB_CREATION_WINDOW_SECONDS
while self._job_creation_timestamps and self._job_creation_timestamps[0] <= window_start:
self._job_creation_timestamps.popleft()
if len(self._job_creation_timestamps) >= self._runner_config.max_jobs_per_minute:
log.info(
"Job creation rate limit reached — deferring queued rollouts",
rollout_id=rollout.rollout_id,
jobs_in_last_minute=len(self._job_creation_timestamps),
max_jobs_per_minute=self._runner_config.max_jobs_per_minute,
)
return

try:
manifest = build_job_spec(rollout, self._config)
attempt_id = manifest["metadata"]["labels"]["agentlightning/attempt-id"]
requires_preloaded = bool(rollout.config.k8s and rollout.config.k8s.require_preloaded_images)
if requires_preloaded:
ready_images = self._fresh_ready_images()
if ready_images is None:
await self._patch_status(
rollout.rollout_id,
state=RolloutState.FAILED,
error_message="Fresh Kubernetes image readiness is unavailable",
)
return
missing_images = sorted(extract_pod_images(manifest) - ready_images)
if missing_images:
await self._patch_status(
rollout.rollout_id,
state=RolloutState.FAILED,
error_message=("Required Kubernetes image(s) are not preloaded: " + ", ".join(missing_images)),
)
return

now = time.monotonic()
window_start = now - JOB_CREATION_WINDOW_SECONDS
while self._job_creation_timestamps and self._job_creation_timestamps[0] <= window_start:
self._job_creation_timestamps.popleft()
if len(self._job_creation_timestamps) >= self._runner_config.max_jobs_per_minute:
log.info(
"Job creation rate limit reached — deferring queued rollouts",
rollout_id=rollout.rollout_id,
jobs_in_last_minute=len(self._job_creation_timestamps),
max_jobs_per_minute=self._runner_config.max_jobs_per_minute,
)
return

api = await self._get_k8s_api()
job = k8s_objects.Job(manifest, api=api)
await job.async_create()
self._job_creation_timestamps.append(time.monotonic())
log.info("Job created", rollout_id=rollout.rollout_id, job_name=job_name, attempt_id=attempt_id)
except ValueError as exc:
error_str = str(exc)
log.error("Invalid Job spec — marking failed", rollout_id=rollout.rollout_id, error=error_str)
await self._patch_status(
rollout.rollout_id,
state=RolloutState.FAILED,
error_message=f"Invalid Job spec: {error_str}",
)
except Exception as exc:
error_str = str(exc)
lower_error = error_str.lower()
Expand Down
86 changes: 86 additions & 0 deletions agentlightning/k8s.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,86 @@
# Copyright (c) Microsoft. All rights reserved.

"""Pure helpers shared by Kubernetes rollout producers and consumers."""

from __future__ import annotations

import json
from collections.abc import Mapping
from functools import lru_cache
from typing import Any

import yaml
from jinja2 import Environment, Template

__all__ = [
"extract_pod_images",
"normalize_image_reference",
"render_job_template",
]


def normalize_image_reference(image: str) -> str:
"""Return a canonical image reference suitable for exact comparison."""
reference = image.strip()
if not reference:
raise ValueError("container image must be a non-empty string")
if "://" in reference:
reference = reference.split("://", 1)[1]

if "/" not in reference:
reference = f"docker.io/library/{reference}"
else:
first, remainder = reference.split("/", 1)
if first in {"docker.io", "index.docker.io"}:
if "/" not in remainder:
remainder = f"library/{remainder}"
reference = f"docker.io/{remainder}"
elif "." not in first and ":" not in first and first != "localhost":
reference = f"docker.io/{reference}"

last_component = reference.rsplit("/", 1)[-1]
if "@" not in reference and ":" not in last_component:
reference = f"{reference}:latest"
return reference


@lru_cache(maxsize=32)
def _compile_job_template(job_template: str) -> Template:
environment = Environment()
environment.filters["yaml_escape"] = lambda value: json.dumps(str(value), ensure_ascii=True)
return environment.from_string(job_template)


def render_job_template(
job_template: str,
*,
job_name: str,
input_data: Any,
) -> dict[str, Any]:
"""Render one Kubernetes Job from the controller-compatible template."""
rendered = _compile_job_template(job_template).render(
job_name=job_name,
input=input_data,
)
documents = [document for document in yaml.safe_load_all(rendered) if document is not None]
if len(documents) != 1:
raise ValueError("job template must render exactly one YAML document")
job = documents[0]
if not isinstance(job, dict) or job.get("kind") != "Job":
raise ValueError("job template must render a Kubernetes Job")
return job


def extract_pod_images(job: Mapping[str, Any]) -> frozenset[str]:
"""Extract normalized images from every container list in a Job Pod spec."""
pod_spec = job.get("spec", {}).get("template", {}).get("spec", {})
images: set[str] = set()
for container_key in ("initContainers", "containers", "ephemeralContainers"):
for container in pod_spec.get(container_key, []) or []:
image = container.get("image")
if not isinstance(image, str) or not image.strip():
raise ValueError(f"{container_key} entry is missing a non-empty image")
images.add(normalize_image_reference(image))
if not images:
raise ValueError("rendered Kubernetes Job contains no container images")
return frozenset(images)
18 changes: 18 additions & 0 deletions agentlightning/schemas.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,23 @@ class RolloutState(StrEnum):
DEFAULT_ATTEMPT_ID = "0"


class K8sImageReadinessReport(BaseModel):
"""Controller report of images cached on every eligible K8s node."""

images: list[str] = Field(default_factory=list)
node_count: int = Field(ge=1)
lease_seconds: float = Field(gt=0, le=300)


class K8sImageReadinessSnapshot(BaseModel):
"""Server-timestamped, leased K8s image inventory."""

images: list[str] = Field(default_factory=list)
node_count: int = Field(ge=1)
observed_at: float
expires_at: float


class RolloutLocalConfig(BaseModel):
"""Local runner config for a rollout."""

Expand All @@ -113,6 +130,7 @@ class RolloutK8sConfig(BaseModel):
"""K8s runner config for a rollout."""

job_template: str | None = None
require_preloaded_images: bool = False


class RolloutConfig(BaseModel):
Expand Down
3 changes: 2 additions & 1 deletion agentlightning/server/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
from omegaconf import DictConfig, OmegaConf

from agentlightning.server.proxy import ProxyPauseState, ProxyRouter
from agentlightning.server.routes import events, models, proxy, rollouts
from agentlightning.server.routes import events, models, proxy, readiness, rollouts

log = structlog.get_logger()

Expand Down Expand Up @@ -83,6 +83,7 @@ async def healthz() -> dict[str, str]:
app.include_router(rollouts.router, prefix="/api", dependencies=[Depends(verify_key)])
app.include_router(events.router, prefix="/api", dependencies=[Depends(verify_key)])
app.include_router(models.router, prefix="/api", dependencies=[Depends(verify_key)])
app.include_router(readiness.router, prefix="/api", dependencies=[Depends(verify_key)])

# Proxy routes (LLM proxy + event ingestion) — require agent-facing auth.
app.include_router(proxy.router, dependencies=[Depends(verify_key)])
Expand Down
Loading