From eaa8f21b1fc4514ee54c7f6260222eac126b5348 Mon Sep 17 00:00:00 2001 From: Blue <3067670134@qq.com> Date: Sat, 3 Oct 2026 22:26:36 +0800 Subject: [PATCH 1/2] fix(controller): fail locally invalid job templates --- agentlightning/controller/k8s_reconciler.py | 11 +++ .../controller/test_k8s_invalid_templates.py | 77 +++++++++++++++++++ 2 files changed, 88 insertions(+) create mode 100644 tests/controller/test_k8s_invalid_templates.py diff --git a/agentlightning/controller/k8s_reconciler.py b/agentlightning/controller/k8s_reconciler.py index 7cf5f7c07..12bb7fd0e 100644 --- a/agentlightning/controller/k8s_reconciler.py +++ b/agentlightning/controller/k8s_reconciler.py @@ -258,6 +258,17 @@ async def _create_job(self, rollout: Rollout) -> None: try: manifest = build_job_spec(rollout, self._config) attempt_id = manifest["metadata"]["labels"]["agentlightning/attempt-id"] + except Exception 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}", + ) + return + + try: api = await self._get_k8s_api() job = k8s_objects.Job(manifest, api=api) await job.async_create() diff --git a/tests/controller/test_k8s_invalid_templates.py b/tests/controller/test_k8s_invalid_templates.py new file mode 100644 index 000000000..0fb812caf --- /dev/null +++ b/tests/controller/test_k8s_invalid_templates.py @@ -0,0 +1,77 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Regression tests for local K8s Job template validation.""" + +from unittest.mock import AsyncMock + +import httpx +import pytest +from omegaconf import OmegaConf + +from agentlightning.client import AgentLightningAsyncClient +from agentlightning.controller.k8s_reconciler import K8sReconciler +from agentlightning.schemas import Rollout, RolloutConfig, RolloutK8sConfig, RolloutLifecycleStatus + + +def _reconciler() -> tuple[K8sReconciler, AsyncMock]: + api = AsyncMock(spec=AgentLightningAsyncClient) + api.patch.return_value = httpx.Response(200, request=httpx.Request("PATCH", "http://store")) + config = OmegaConf.create( + { + "agl_server": {"url": "http://store", "key": ""}, + "k8s_runner": { + "namespace": "default", + "ttl_after_finished": 600, + "max_jobs_per_minute": 100, + }, + } + ) + return K8sReconciler(api, config), api + + +def _rollout(template: str) -> Rollout: + return Rollout( + rollout_id="invalid-template", + input={}, + config=RolloutConfig(k8s=RolloutK8sConfig(job_template=template)), + status=RolloutLifecycleStatus(created_at=1.0, updated_at=1.0), + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "template", + [ + "kind: Job\nspec: [\n", + "kind: Job\nmetadata:\n name: {{ job_name\n", + "apiVersion: v1\nkind: Pod\nmetadata: {}\n", + ], + ids=["malformed-yaml", "malformed-jinja", "invalid-kind"], +) +async def test_invalid_local_template_fails_without_cluster_call(template: str) -> None: + reconciler, api = _reconciler() + get_k8s_api = AsyncMock() + reconciler._get_k8s_api = get_k8s_api + + await reconciler._create_job(_rollout(template)) + + get_k8s_api.assert_not_awaited() + api.patch.assert_awaited_once() + request = api.patch.await_args + assert request.args == ("/api/rollouts/invalid-template",) + assert request.kwargs["json"]["status"]["state"] == "failed" + assert request.kwargs["json"]["status"]["error_message"].startswith("Invalid Job spec: ") + assert not reconciler._job_creation_timestamps + + +@pytest.mark.asyncio +async def test_transient_cluster_error_for_valid_template_is_retried() -> None: + reconciler, api = _reconciler() + get_k8s_api = AsyncMock(side_effect=RuntimeError("cluster temporarily unavailable")) + reconciler._get_k8s_api = get_k8s_api + + await reconciler._create_job(_rollout("apiVersion: batch/v1\nkind: Job\nspec: {}\n")) + + get_k8s_api.assert_awaited_once() + api.patch.assert_not_awaited() + assert not reconciler._job_creation_timestamps From 2e8e88563718993b3d73dda3504881ea22c2defe Mon Sep 17 00:00:00 2001 From: Blue <3067670134@qq.com> Date: Sun, 4 Oct 2026 10:26:49 +0800 Subject: [PATCH 2/2] fix(controller): validate queued jobs before cluster gates --- agentlightning/controller/k8s_reconciler.py | 57 +++-- .../controller/test_k8s_invalid_templates.py | 237 +++++++++++++++++- 2 files changed, 264 insertions(+), 30 deletions(-) diff --git a/agentlightning/controller/k8s_reconciler.py b/agentlightning/controller/k8s_reconciler.py index 12bb7fd0e..877ada806 100644 --- a/agentlightning/controller/k8s_reconciler.py +++ b/agentlightning/controller/k8s_reconciler.py @@ -161,6 +161,17 @@ async def _periodic_reconcile_loop(self) -> None: async def _reconcile_once(self) -> None: """One reconcile cycle: align queuing/running rollouts with K8s Jobs.""" rollouts = await self._query_rollouts(state_in=[RolloutState.QUEUING, RolloutState.RUNNING], limit=500) + prepared_rollouts: list[tuple[Rollout, dict[str, Any] | None]] = [] + for rollout in rollouts: + manifest = None + if rollout.status.state == RolloutState.QUEUING: + manifest = await self._prepare_job_manifest(rollout) + if manifest is None: + continue + prepared_rollouts.append((rollout, manifest)) + if not prepared_rollouts: + return + api = await self._get_k8s_api() jobs = [ cast(k8s_objects.Job, job).raw @@ -172,13 +183,13 @@ async def _reconcile_once(self) -> None: ] jobs_by_name = {job.get("metadata", {}).get("name", ""): job for job in jobs} - for rollout in rollouts: + for rollout, manifest in prepared_rollouts: job_name = rollout.status.k8s_job_name or build_job_name(rollout.rollout_id) job = jobs_by_name.get(job_name) if job is None: if rollout.status.state == RolloutState.QUEUING: - await self._create_job(rollout) + await self._create_job(rollout, manifest=manifest) continue log.warning("Orphaned running rollout — Job gone", rollout_id=rollout.rollout_id, job_name=job_name) await self._patch_status(rollout.rollout_id, state=RolloutState.FAILED, error_message="Job disappeared") @@ -239,8 +250,33 @@ async def _reconcile_once(self) -> None: error_message=error_message, ) - async def _create_job(self, rollout: Rollout) -> None: - """Create a K8s Job for a queuing rollout without changing rollout state.""" + async def _prepare_job_manifest(self, rollout: Rollout) -> dict[str, Any] | None: + """Validate locally and report invalid templates as FAILED before cluster or quota checks.""" + try: + manifest = build_job_spec(rollout, self._config) + _ = manifest["metadata"]["labels"]["agentlightning/attempt-id"] + return manifest + except Exception 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}", + ) + return None + + async def _create_job(self, rollout: Rollout, *, manifest: dict[str, Any] | None = None) -> None: + """Submit a queuing rollout's Job, marking invalid specs FAILED. + + Valid rollouts remain QUEUING until Job observations update their state. + Rate limits and transient cluster failures defer submission for retry. + """ + if manifest is None: + manifest = await self._prepare_job_manifest(rollout) + if manifest is None: + return + attempt_id = manifest["metadata"]["labels"]["agentlightning/attempt-id"] job_name = build_job_name(rollout.rollout_id) now = time.monotonic() window_start = now - JOB_CREATION_WINDOW_SECONDS @@ -255,19 +291,6 @@ async def _create_job(self, rollout: Rollout) -> None: ) return - try: - manifest = build_job_spec(rollout, self._config) - attempt_id = manifest["metadata"]["labels"]["agentlightning/attempt-id"] - except Exception 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}", - ) - return - try: api = await self._get_k8s_api() job = k8s_objects.Job(manifest, api=api) diff --git a/tests/controller/test_k8s_invalid_templates.py b/tests/controller/test_k8s_invalid_templates.py index 0fb812caf..3619a9fb8 100644 --- a/tests/controller/test_k8s_invalid_templates.py +++ b/tests/controller/test_k8s_invalid_templates.py @@ -2,15 +2,27 @@ """Regression tests for local K8s Job template validation.""" -from unittest.mock import AsyncMock +import time +from collections.abc import AsyncIterator +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock, Mock import httpx import pytest from omegaconf import OmegaConf from agentlightning.client import AgentLightningAsyncClient +from agentlightning.controller import k8s_reconciler from agentlightning.controller.k8s_reconciler import K8sReconciler -from agentlightning.schemas import Rollout, RolloutConfig, RolloutK8sConfig, RolloutLifecycleStatus +from agentlightning.schemas import Rollout, RolloutConfig, RolloutK8sConfig, RolloutLifecycleStatus, RolloutState + +INVALID_TEMPLATES = [ + "kind: Job\nspec: [\n", + "kind: Job\nmetadata:\n name: {{ job_name\n", + "apiVersion: v1\nkind: Pod\nmetadata: {}\n", +] +VALID_TEMPLATE = "apiVersion: batch/v1\nkind: Job\nspec: {}\n" def _reconciler() -> tuple[K8sReconciler, AsyncMock]: @@ -29,29 +41,30 @@ def _reconciler() -> tuple[K8sReconciler, AsyncMock]: return K8sReconciler(api, config), api -def _rollout(template: str) -> Rollout: +def _rollout( + template: str, *, rollout_id: str = "invalid-template", state: RolloutState = RolloutState.QUEUING +) -> Rollout: return Rollout( - rollout_id="invalid-template", + rollout_id=rollout_id, input={}, config=RolloutConfig(k8s=RolloutK8sConfig(job_template=template)), - status=RolloutLifecycleStatus(created_at=1.0, updated_at=1.0), + status=RolloutLifecycleStatus(created_at=1.0, updated_at=1.0, state=state), ) @pytest.mark.asyncio @pytest.mark.parametrize( "template", - [ - "kind: Job\nspec: [\n", - "kind: Job\nmetadata:\n name: {{ job_name\n", - "apiVersion: v1\nkind: Pod\nmetadata: {}\n", - ], + INVALID_TEMPLATES, ids=["malformed-yaml", "malformed-jinja", "invalid-kind"], ) -async def test_invalid_local_template_fails_without_cluster_call(template: str) -> None: +@pytest.mark.parametrize("quota_full", [False, True]) +async def test_invalid_local_template_fails_without_cluster_call(template: str, quota_full: bool) -> None: reconciler, api = _reconciler() get_k8s_api = AsyncMock() reconciler._get_k8s_api = get_k8s_api + if quota_full: + reconciler._job_creation_timestamps.extend([time.monotonic()] * 100) await reconciler._create_job(_rollout(template)) @@ -61,7 +74,7 @@ async def test_invalid_local_template_fails_without_cluster_call(template: str) assert request.args == ("/api/rollouts/invalid-template",) assert request.kwargs["json"]["status"]["state"] == "failed" assert request.kwargs["json"]["status"]["error_message"].startswith("Invalid Job spec: ") - assert not reconciler._job_creation_timestamps + assert len(reconciler._job_creation_timestamps) == (100 if quota_full else 0) @pytest.mark.asyncio @@ -70,8 +83,206 @@ async def test_transient_cluster_error_for_valid_template_is_retried() -> None: get_k8s_api = AsyncMock(side_effect=RuntimeError("cluster temporarily unavailable")) reconciler._get_k8s_api = get_k8s_api - await reconciler._create_job(_rollout("apiVersion: batch/v1\nkind: Job\nspec: {}\n")) + await reconciler._create_job(_rollout(VALID_TEMPLATE)) get_k8s_api.assert_awaited_once() api.patch.assert_not_awaited() assert not reconciler._job_creation_timestamps + + +def _query_response(api: AsyncMock, *rollouts: Rollout) -> None: + api.get.return_value = httpx.Response( + 200, + request=httpx.Request("GET", "http://store/api/rollouts"), + json=[rollout.model_dump(mode="json") for rollout in rollouts], + ) + + +async def _job_listing(jobs: list[dict[str, Any]], error: Exception | None = None) -> AsyncIterator[Any]: + if error is not None: + raise error + for job in jobs: + yield SimpleNamespace(raw=job) + + +def _mock_job_type(monkeypatch: pytest.MonkeyPatch, *listings: AsyncIterator[Any]) -> Mock: + job = Mock() + job.async_create = AsyncMock() + job_type = Mock(return_value=job) + job_type.async_list = Mock(side_effect=listings) + monkeypatch.setattr(k8s_reconciler.k8s_objects, "Job", job_type) + return job_type + + +@pytest.mark.asyncio +@pytest.mark.parametrize("template", INVALID_TEMPLATES, ids=["malformed-yaml", "malformed-jinja", "invalid-kind"]) +@pytest.mark.parametrize("gate", ["api-unavailable", "listing-unavailable", "quota-full"]) +async def test_reconcile_rejects_invalid_templates_before_cluster_and_quota( + monkeypatch: pytest.MonkeyPatch, template: str, gate: str +) -> None: + reconciler, api = _reconciler() + _query_response(api, _rollout(template)) + get_k8s_api = AsyncMock(return_value=object()) + if gate == "api-unavailable": + get_k8s_api.side_effect = RuntimeError("cluster temporarily unavailable") + reconciler._get_k8s_api = get_k8s_api + listing_error = RuntimeError("list temporarily unavailable") if gate == "listing-unavailable" else None + job_type = _mock_job_type(monkeypatch, _job_listing([], listing_error)) + if gate == "quota-full": + reconciler._job_creation_timestamps.extend([time.monotonic()] * 100) + + await reconciler._reconcile_once() + + get_k8s_api.assert_not_awaited() + job_type.async_list.assert_not_called() + job_type.assert_not_called() + api.patch.assert_awaited_once() + assert api.patch.await_args.kwargs["json"]["status"]["state"] == "failed" + assert api.patch.await_args.kwargs["json"]["status"]["error_message"].startswith("Invalid Job spec: ") + assert len(reconciler._job_creation_timestamps) == (100 if gate == "quota-full" else 0) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("gate", ["api", "listing"]) +async def test_reconcile_valid_job_retries_cluster_failure(monkeypatch: pytest.MonkeyPatch, gate: str) -> None: + reconciler, api = _reconciler() + _query_response(api, _rollout(VALID_TEMPLATE)) + get_k8s_api = AsyncMock(return_value=object()) + failure = RuntimeError("cluster temporarily unavailable") + if gate == "api": + get_k8s_api.side_effect = [failure, object(), object()] + listings = [_job_listing([])] + else: + listings = [_job_listing([], failure), _job_listing([])] + reconciler._get_k8s_api = get_k8s_api + job_type = _mock_job_type(monkeypatch, *listings) + + with pytest.raises(RuntimeError, match="cluster temporarily unavailable"): + await reconciler._reconcile_once() + api.patch.assert_not_awaited() + job_type.return_value.async_create.assert_not_awaited() + assert not reconciler._job_creation_timestamps + + await reconciler._reconcile_once() + job_type.return_value.async_create.assert_awaited_once() + api.patch.assert_not_awaited() + assert len(reconciler._job_creation_timestamps) == 1 + + +@pytest.mark.asyncio +async def test_reconcile_valid_job_waits_for_quota_and_reuses_manifest(monkeypatch: pytest.MonkeyPatch) -> None: + reconciler, api = _reconciler() + _query_response(api, _rollout(VALID_TEMPLATE)) + reconciler._get_k8s_api = AsyncMock(return_value=object()) + job_type = _mock_job_type(monkeypatch, _job_listing([]), _job_listing([])) + build = Mock(wraps=k8s_reconciler.build_job_spec) + monkeypatch.setattr(k8s_reconciler, "build_job_spec", build) + reconciler._job_creation_timestamps.extend([time.monotonic()] * 100) + + await reconciler._reconcile_once() + job_type.return_value.async_create.assert_not_awaited() + api.patch.assert_not_awaited() + assert len(reconciler._job_creation_timestamps) == 100 + + # Let the same occupied slots age out of the real rate-limit window. + reconciler._job_creation_timestamps.clear() + reconciler._job_creation_timestamps.extend( + [time.monotonic() - k8s_reconciler.JOB_CREATION_WINDOW_SECONDS - 1] * 100 + ) + await reconciler._reconcile_once() + job_type.return_value.async_create.assert_awaited_once() + api.patch.assert_not_awaited() + assert len(reconciler._job_creation_timestamps) == 1 + assert build.call_count == 2 # One render per cycle, reused for submission. + + +@pytest.mark.asyncio +async def test_reconcile_keeps_observing_running_jobs_without_rendering_template( + monkeypatch: pytest.MonkeyPatch, +) -> None: + reconciler, api = _reconciler() + _query_response(api, _rollout(INVALID_TEMPLATES[0], state=RolloutState.RUNNING)) + reconciler._get_k8s_api = AsyncMock(return_value=object()) + job_type = _mock_job_type( + monkeypatch, + _job_listing([{"metadata": {"name": "agl-rollout-invalid-template"}, "status": {"succeeded": 1}}]), + ) + + await reconciler._reconcile_once() + + job_type.assert_not_called() + api.patch.assert_awaited_once() + assert api.patch.await_args.kwargs["json"]["status"]["state"] == "succeeded" + + +@pytest.mark.asyncio +async def test_reconcile_retries_failed_invalid_template_status_patch(monkeypatch: pytest.MonkeyPatch) -> None: + reconciler, api = _reconciler() + _query_response(api, _rollout(INVALID_TEMPLATES[0])) + api.patch.return_value = httpx.Response(503, request=httpx.Request("PATCH", "http://store")) + get_k8s_api = AsyncMock(side_effect=RuntimeError("cluster temporarily unavailable")) + reconciler._get_k8s_api = get_k8s_api + job_type = _mock_job_type(monkeypatch) + + for _ in range(2): + await reconciler._reconcile_once() + + assert api.patch.await_count == 2 + assert all(call.kwargs["json"]["status"]["state"] == "failed" for call in api.patch.await_args_list) + get_k8s_api.assert_not_awaited() + job_type.async_list.assert_not_called() + + +@pytest.mark.asyncio +async def test_reconcile_valid_first_does_not_hide_invalid_template_during_cluster_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + reconciler, api = _reconciler() + _query_response(api, _rollout(VALID_TEMPLATE, rollout_id="valid"), _rollout(INVALID_TEMPLATES[0])) + + async def unavailable() -> None: + # Every queued rollout must be validated before the first cluster call. + api.patch.assert_awaited_once() + assert api.patch.await_args.args == ("/api/rollouts/invalid-template",) + raise RuntimeError("cluster temporarily unavailable") + + reconciler._get_k8s_api = AsyncMock(side_effect=unavailable) + job_type = _mock_job_type(monkeypatch) + with pytest.raises(RuntimeError, match="cluster temporarily unavailable"): + await reconciler._reconcile_once() + + api.patch.assert_awaited_once() + assert api.patch.await_args.kwargs["json"]["status"]["state"] == "failed" + job_type.assert_not_called() + assert not reconciler._job_creation_timestamps + + +@pytest.mark.asyncio +async def test_reconcile_invalid_first_keeps_creating_valid_jobs_and_observing_running_jobs( + monkeypatch: pytest.MonkeyPatch, +) -> None: + reconciler, api = _reconciler() + _query_response( + api, + _rollout(INVALID_TEMPLATES[0]), + _rollout(VALID_TEMPLATE, rollout_id="valid"), + _rollout(INVALID_TEMPLATES[0], rollout_id="running", state=RolloutState.RUNNING), + ) + reconciler._get_k8s_api = AsyncMock(return_value=object()) + job_type = _mock_job_type( + monkeypatch, + _job_listing([{"metadata": {"name": "agl-rollout-running"}, "status": {"succeeded": 1}}]), + ) + build = Mock(wraps=k8s_reconciler.build_job_spec) + monkeypatch.setattr(k8s_reconciler, "build_job_spec", build) + + await reconciler._reconcile_once() + + job_type.return_value.async_create.assert_awaited_once() + assert job_type.call_args.args[0]["metadata"]["name"] == "agl-rollout-valid" + assert [(call.args[0], call.kwargs["json"]["status"]["state"]) for call in api.patch.await_args_list] == [ + ("/api/rollouts/invalid-template", "failed"), + ("/api/rollouts/running", "succeeded"), + ] + assert [call.args[0].rollout_id for call in build.call_args_list] == ["invalid-template", "valid"] + assert len(reconciler._job_creation_timestamps) == 1