diff --git a/cron/scheduler.py b/cron/scheduler.py index 7cf688f3734a..511cbda5d384 100644 --- a/cron/scheduler.py +++ b/cron/scheduler.py @@ -3557,7 +3557,14 @@ def _notify_provider_jobs_changed() -> None: logger.debug("on_jobs_changed notify failed: %s", e) -def tick(verbose: bool = True, adapters=None, loop=None, sync: bool = True) -> int: +def tick( + verbose: bool = True, + adapters=None, + loop=None, + sync: bool = True, + *, + can_dispatch=None, +): """ Check and run all due jobs. @@ -3568,7 +3575,9 @@ def tick(verbose: bool = True, adapters=None, loop=None, sync: bool = True) -> i verbose: Whether to print status messages adapters: Optional dict mapping Platform → live adapter (from gateway) loop: Optional asyncio event loop (from gateway) for live adapter sends - + can_dispatch: Optional synchronous gate; false leaves due jobs untouched + for the next allowed tick + Returns: Number of jobs executed (0 if another tick is already running) """ @@ -3590,6 +3599,10 @@ def tick(verbose: bool = True, adapters=None, loop=None, sync: bool = True) -> i return 0 try: + if can_dispatch is not None and not can_dispatch(): + logger.debug("Cron dispatch paused while gateway drains existing work") + return 0 + due_jobs = get_due_jobs() if verbose and not due_jobs: diff --git a/cron/scheduler_provider.py b/cron/scheduler_provider.py index ab3121bfa3ba..5429f79cf362 100644 --- a/cron/scheduler_provider.py +++ b/cron/scheduler_provider.py @@ -156,14 +156,16 @@ class InProcessCronScheduler(CronScheduler): ``start()`` blocks in the tick loop until ``stop_event`` is set, identical to the pre-refactor ``_start_cron_ticker`` core loop. The caller runs it in - a daemon thread. + a daemon thread. ``can_dispatch`` is an optional synchronous gate supplied + by GatewayRunner during external drain; skipped ticks leave due jobs intact + for the next allowed tick. """ @property def name(self) -> str: return "builtin" - def start(self, stop_event, *, adapters=None, loop=None, interval=60): + def start(self, stop_event, *, adapters=None, loop=None, interval=60, can_dispatch=None): import logging from cron.scheduler import tick as cron_tick from cron.jobs import record_ticker_heartbeat @@ -176,7 +178,16 @@ class InProcessCronScheduler(CronScheduler): while not stop_event.is_set(): ok = False try: - cron_tick(verbose=False, adapters=adapters, loop=loop, sync=False) + if can_dispatch is not None and not can_dispatch(): + logger.debug("Cron dispatch paused while gateway drains existing work") + else: + cron_tick( + verbose=False, + adapters=adapters, + loop=loop, + sync=False, + can_dispatch=can_dispatch, + ) ok = True except BaseException as e: # Catch BaseException (not just Exception) so a SystemExit from diff --git a/gateway/platforms/api_server.py b/gateway/platforms/api_server.py index cff62197fa8f..76f69e73e814 100644 --- a/gateway/platforms/api_server.py +++ b/gateway/platforms/api_server.py @@ -35,6 +35,8 @@ import asyncio import hashlib import hmac import json +from contextvars import ContextVar +from functools import wraps import logging import os import socket as _socket @@ -664,6 +666,43 @@ def _openai_error(message: str, err_type: str = "invalid_request_error", param: } +_api_agent_request_reservation: ContextVar[Optional[dict[str, bool]]] = ContextVar( + "api_agent_request_reservation", default=None +) + + +def _admit_api_agent_request(handler): + """Reserve an authenticated API turn before its handler first awaits. + + Gateway shutdown and aiohttp requests share an event loop. Keeping the + drain check and reservation in one non-awaiting block prevents a request + admitted immediately before shutdown from becoming invisible while it is + still parsing its body or resolving session state. The mutable reservation + is intentionally shared with child tasks so agent/task bookkeeping releases + this one slot exactly once. + """ + @wraps(handler) + async def _wrapped(self, request, *args, **kwargs): + auth_err = self._check_auth(request) + if auth_err: + return auth_err + draining = self._draining_response() + if draining is not None: + return draining + reservation = {"active": True} + token = _api_agent_request_reservation.set(reservation) + self._pending_agent_requests += 1 + try: + return await handler(self, request, *args, **kwargs) + finally: + if reservation["active"]: + reservation["active"] = False + self._pending_agent_requests = max(0, self._pending_agent_requests - 1) + _api_agent_request_reservation.reset(token) + + return _wrapped + + if AIOHTTP_AVAILABLE: @web.middleware async def body_limit_middleware(request, handler): @@ -912,6 +951,10 @@ class APIServerAdapter(BasePlatformAdapter): # (the /v1/runs path tracks its own in-flight set via # _active_run_tasks). self._inflight_agent_runs: int = 0 + # Requests admitted before their handler reaches agent bookkeeping. + # Shutdown counts this reservation so the request cannot slip through + # the drain between its first await and _run_agent()/task registration. + self._pending_agent_requests: int = 0 def active_agent_work_count(self) -> int: """Return all live agent work owned by this API adapter. @@ -922,8 +965,10 @@ class APIServerAdapter(BasePlatformAdapter): covers that gap and excludes completed tasks retained until cleanup. """ try: - return int(self._inflight_agent_runs) + sum( - not task.done() for task in self._active_run_tasks.values() + return ( + int(getattr(self, "_pending_agent_requests", 0)) + + int(self._inflight_agent_runs) + + sum(not task.done() for task in self._active_run_tasks.values()) ) except Exception: return 0 @@ -958,6 +1003,13 @@ class APIServerAdapter(BasePlatformAdapter): headers={"Retry-After": "1"}, ) + def _activate_admitted_request(self) -> None: + """Transfer this request's drain reservation to agent bookkeeping.""" + reservation = _api_agent_request_reservation.get() + if reservation and reservation["active"]: + reservation["active"] = False + self._pending_agent_requests = max(0, self._pending_agent_requests - 1) + def _readiness_work_counts(self) -> tuple[int, int, int]: """Return bounded work counts from each subsystem's public state.""" active_api_runs = sum( @@ -1964,14 +2016,9 @@ class APIServerAdapter(BasePlatformAdapter): fork = db.get_session(fork_id) or {"id": fork_id, "parent_session_id": source_id} return web.json_response({"object": "hermes.session", "session": self._session_response(fork)}, status=201) + @_admit_api_agent_request async def _handle_session_chat(self, request: "web.Request") -> "web.Response": """POST /api/sessions/{session_id}/chat — one synchronous agent turn.""" - auth_err = self._check_auth(request) - if auth_err: - return auth_err - draining = self._draining_response() - if draining is not None: - return draining gateway_session_key, key_err = self._parse_session_key_header(request) if key_err is not None: return key_err @@ -2011,14 +2058,9 @@ class APIServerAdapter(BasePlatformAdapter): headers=headers, ) + @_admit_api_agent_request async def _handle_session_chat_stream(self, request: "web.Request") -> "web.StreamResponse": """POST /api/sessions/{session_id}/chat/stream — SSE wrapper over _run_agent.""" - auth_err = self._check_auth(request) - if auth_err: - return auth_err - draining = self._draining_response() - if draining is not None: - return draining gateway_session_key, key_err = self._parse_session_key_header(request) if key_err is not None: return key_err @@ -2155,15 +2197,9 @@ class APIServerAdapter(BasePlatformAdapter): logger.debug("[api_server] session SSE stream error: %s", exc) return response + @_admit_api_agent_request async def _handle_chat_completions(self, request: "web.Request") -> "web.Response": """POST /v1/chat/completions — OpenAI Chat Completions format.""" - auth_err = self._check_auth(request) - if auth_err: - return auth_err - draining = self._draining_response() - if draining is not None: - return draining - # Bound total in-flight agent runs (configurable; #7483). limited = self._concurrency_limited_response() if limited is not None: @@ -2365,8 +2401,8 @@ class APIServerAdapter(BasePlatformAdapter): # ``tool_progress_callback`` is intentionally not wired here: # it would duplicate every emit because ``run_agent`` fires it # side-by-side with ``tool_start_callback``/``tool_complete_callback``. - # The structured callbacks are strictly richer (they carry the - # tool_call id), so they own the chat-completions SSE channel. + # The structured callbacks are strictly richer (they carry + # the tool_call id), so they own the chat-completions SSE channel. agent_ref = [None] agent_task = asyncio.ensure_future(self._run_agent( user_message=user_message, @@ -3294,15 +3330,9 @@ class APIServerAdapter(BasePlatformAdapter): return response + @_admit_api_agent_request async def _handle_responses(self, request: "web.Request") -> "web.Response": """POST /v1/responses — OpenAI Responses API format.""" - auth_err = self._check_auth(request) - if auth_err: - return auth_err - draining = self._draining_response() - if draining is not None: - return draining - # Bound total in-flight agent runs (configurable; #7483). limited = self._concurrency_limited_response() if limited is not None: @@ -3845,6 +3875,9 @@ class APIServerAdapter(BasePlatformAdapter): auth_err = self._check_auth(request) if auth_err: return auth_err + draining = self._draining_response() + if draining is not None: + return draining cron_err = self._check_jobs_available() if cron_err: return cron_err @@ -3890,6 +3923,9 @@ class APIServerAdapter(BasePlatformAdapter): self._request_audit_log_suffix(request), ) return web.json_response({"error": "invalid fire token"}, status=401) + draining = self._draining_response() + if draining is not None: + return draining try: body = await request.json() @@ -4068,19 +4104,22 @@ class APIServerAdapter(BasePlatformAdapter): """Return a 429 response if the concurrent-run cap is reached, else None. The cap bounds total in-flight agent activity across every - agent-serving endpoint: the non-streaming chat/responses paths - (tracked by ``_inflight_agent_runs``) plus the ``/v1/runs`` path - (tracked by live entries in ``_active_run_tasks``). Stream queues are - transport state and may disappear while their underlying run remains - active, so they must not define run concurrency. A configured value of - 0 disables the cap entirely. + agent-serving endpoint. Reuse the same adapter-owned work count that + shutdown draining uses, including an admitted request before it reaches + agent/task bookkeeping. Stream queues are transport state and may + disappear while their underlying run remains active, so they must not + define run concurrency. A configured value of 0 disables the cap. """ limit = self._max_concurrent_runs if limit <= 0: return None - inflight = self._inflight_agent_runs + sum( - not task.done() for task in self._active_run_tasks.values() - ) + inflight = self.active_agent_work_count() + # The current request owns one reservation until it hands off to + # _run_agent() or /v1/runs task registration. It must not consume its + # own last available slot; other admitted requests remain counted. + reservation = _api_agent_request_reservation.get() + if reservation and reservation["active"]: + inflight -= 1 if inflight >= limit: return web.json_response( _openai_error( @@ -4198,6 +4237,7 @@ class APIServerAdapter(BasePlatformAdapter): finally: clear_session_vars(tokens) + self._activate_admitted_request() self._inflight_agent_runs += 1 try: return await loop.run_in_executor(None, _run) @@ -4272,15 +4312,9 @@ class APIServerAdapter(BasePlatformAdapter): return _callback + @_admit_api_agent_request async def _handle_runs(self, request: "web.Request") -> "web.Response": """POST /v1/runs — start an agent run, return run_id immediately.""" - auth_err = self._check_auth(request) - if auth_err: - return auth_err - draining = self._draining_response() - if draining is not None: - return draining - # Long-term memory scope header (see chat_completions for details). gateway_session_key, key_err = self._parse_session_key_header(request) if key_err is not None: @@ -4602,6 +4636,7 @@ class APIServerAdapter(BasePlatformAdapter): self._run_approval_sessions.pop(run_id, None) self._stopping_run_ids.discard(run_id) + self._activate_admitted_request() task = asyncio.create_task(_run_and_close()) self._active_run_tasks[run_id] = task try: diff --git a/gateway/run.py b/gateway/run.py index 3382d5323c2a..d5b3fbff4743 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -4122,6 +4122,14 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew def _running_agent_count(self) -> int: return len(self._running_agents) + def _active_work_count(self) -> int: + """All agent work the gateway must expose and drain as one total.""" + return ( + self._running_agent_count() + + self._active_cron_job_count() + + self._active_api_run_count() + ) + def _active_cron_job_count(self) -> int: """Count of cron jobs currently executing, from the cron scheduler's own in-flight tracking (``cron.scheduler._running_job_ids``). @@ -4543,7 +4551,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew gateway_state=gateway_state, exit_reason=exit_reason, restart_requested=self._restart_requested, - active_agents=self._running_agent_count(), + active_agents=self._active_work_count(), ) except Exception: pass @@ -4566,7 +4574,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew """ try: from gateway.status import write_runtime_status - write_runtime_status(active_agents=self._running_agent_count()) + write_runtime_status(active_agents=self._active_work_count()) except Exception: pass @@ -4593,7 +4601,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew logger.info( "External drain ENGAGED (.drain_request.json present) — refusing " "new turns; %d in-flight turn(s) will finish. Process stays up.", - self._running_agent_count(), + self._active_work_count(), ) # Flip the persisted lifecycle state so /api/status.gateway_busy / # gateway_drainable track the drain. Preserve active_agents (the @@ -4644,6 +4652,10 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew try: if drain_requested(): self._enter_external_drain() + # API and cron work live outside messaging's + # _running_agents map. Refresh the aggregate while an + # external caller polls this reversible drain state. + self._persist_active_agents() else: self._exit_external_drain() except asyncio.CancelledError: @@ -21004,13 +21016,21 @@ async def start_gateway(config: Optional[GatewayConfig] = None, replace: bool = # historical in-process 60s ticker; an external provider (e.g. chronos) # may arm a schedule and return. Pass the event loop so cron delivery can # use live adapters (E2EE support). - from cron.scheduler_provider import resolve_cron_scheduler + from cron.scheduler_provider import InProcessCronScheduler, resolve_cron_scheduler cron_stop = threading.Event() cron_provider = resolve_cron_scheduler() + cron_start_kwargs = {"adapters": runner.adapters, "loop": asyncio.get_running_loop()} + # External cron providers own their remote scheduling contract. Only the + # in-process ticker polls local due jobs, so only it receives the local + # external-drain dispatch gate. + if isinstance(cron_provider, InProcessCronScheduler): + cron_start_kwargs["can_dispatch"] = lambda: not ( + runner._draining or runner._external_drain_active + ) cron_thread = threading.Thread( target=cron_provider.start, args=(cron_stop,), - kwargs={"adapters": runner.adapters, "loop": asyncio.get_running_loop()}, + kwargs=cron_start_kwargs, daemon=True, name="cron-scheduler", ) diff --git a/tests/cron/test_scheduler.py b/tests/cron/test_scheduler.py index 3d262155ba3a..37db53890138 100644 --- a/tests/cron/test_scheduler.py +++ b/tests/cron/test_scheduler.py @@ -1533,6 +1533,25 @@ class TestRunJobSessionPersistence: assert "final fallback report" in output assert "(FAILED)" not in output + def test_tick_skips_due_jobs_while_dispatch_is_paused(self, tmp_path): + """The drain gate runs before advancing a due job's schedule.""" + from cron.scheduler import tick + + job = { + "id": "paused-due-job", + "name": "paused due job", + "schedule": {"kind": "interval", "seconds": 60}, + "next_run_at": "2020-01-01T00:00:00+00:00", + "enabled": True, + } + with patch("cron.scheduler.get_due_jobs", return_value=[job]), patch( + "cron.scheduler.advance_next_run" + ) as advance, patch("cron.scheduler.run_one_job") as run_one: + assert tick(verbose=False, sync=True, can_dispatch=lambda: False) == 0 + + advance.assert_not_called() + run_one.assert_not_called() + def test_tick_marks_empty_response_as_error(self, tmp_path): """When run_job returns success=True but final_response is empty, tick() should mark the job as error so last_status != 'ok'. diff --git a/tests/cron/test_scheduler_provider.py b/tests/cron/test_scheduler_provider.py index 348caa4adff8..2779fdd9f3f2 100644 --- a/tests/cron/test_scheduler_provider.py +++ b/tests/cron/test_scheduler_provider.py @@ -129,13 +129,11 @@ def test_cronscheduler_default_is_available_true(): def test_abc_growth_stays_additive(): - """Forward-compat guard: the ABC's REQUIRED surface is exactly name+start. + """The provider interface stays source-compatible with existing plugins. - Any optional hook added later for the external provider - (on_jobs_changed/fire_due/reconcile) must be NON-abstract (carry a default), - so the built-in keeps satisfying the ABC without overriding them. This test - fails loudly if someone makes a future hook abstract (a breaking change that - would force every provider — including the built-in — to implement it). + ``start`` must be the only required implementation hook: future optional + behavior belongs in non-abstract default methods so custom plugins do not + break on import after an upgrade. """ from cron.scheduler_provider import CronScheduler @@ -174,6 +172,33 @@ def test_inprocess_provider_ticks_and_stops(): assert calls[0].get("sync") is False +def test_inprocess_provider_skips_dispatch_while_draining(): + """A drain pause keeps due work pending until dispatch is re-enabled.""" + from cron.scheduler_provider import InProcessCronScheduler + + calls = [] + stop = threading.Event() + allow_dispatch = threading.Event() + provider = InProcessCronScheduler() + + with patch("cron.scheduler.tick", side_effect=lambda *a, **k: calls.append(k) or 0): + thread = threading.Thread( + target=provider.start, + args=(stop,), + kwargs={"interval": 0.01, "can_dispatch": allow_dispatch.is_set}, + daemon=True, + ) + thread.start() + time.sleep(0.05) + assert calls == [] + allow_dispatch.set() + assert _wait_until(lambda: len(calls) >= 1), "provider never resumed dispatch" + stop.set() + thread.join(timeout=5) + + assert not thread.is_alive() + + def test_inprocess_provider_stop_is_noop(): """The default stop() hook is a safe no-op (the stop_event is the real stop signal for the built-in).""" diff --git a/tests/gateway/restart_test_helpers.py b/tests/gateway/restart_test_helpers.py index 8b1c66c5ba89..c0466838bc8c 100644 --- a/tests/gateway/restart_test_helpers.py +++ b/tests/gateway/restart_test_helpers.py @@ -115,6 +115,18 @@ def make_restart_runner( runner._running_agent_count = GatewayRunner._running_agent_count.__get__( runner, GatewayRunner ) + runner._active_cron_job_count = GatewayRunner._active_cron_job_count.__get__( + runner, GatewayRunner + ) + runner._active_api_run_count = GatewayRunner._active_api_run_count.__get__( + runner, GatewayRunner + ) + runner._active_work_count = GatewayRunner._active_work_count.__get__( + runner, GatewayRunner + ) + runner._persist_active_agents = GatewayRunner._persist_active_agents.__get__( + runner, GatewayRunner + ) runner._snapshot_running_agents = GatewayRunner._snapshot_running_agents.__get__( runner, GatewayRunner ) diff --git a/tests/gateway/test_api_server_active_work_drain.py b/tests/gateway/test_api_server_active_work_drain.py index 780eb88d8906..9ca6c5c4639b 100644 --- a/tests/gateway/test_api_server_active_work_drain.py +++ b/tests/gateway/test_api_server_active_work_drain.py @@ -9,7 +9,7 @@ turns once the gateway starts draining. import asyncio from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from aiohttp import web @@ -37,9 +37,9 @@ def _make_api_adapter(*, inflight: int = 0, queued_ids=()): ) def active_agent_work_count() -> int: - return int(adapter._inflight_agent_runs) + sum( - not task.done() for task in adapter._active_run_tasks.values() - ) + return int(getattr(adapter, "_pending_agent_requests", 0)) + int( + adapter._inflight_agent_runs + ) + sum(not task.done() for task in adapter._active_run_tasks.values()) adapter.active_agent_work_count = active_agent_work_count return adapter @@ -94,6 +94,37 @@ class TestActiveApiRunCount: class TestAPIServerAdapterWorkCount: + def test_concurrency_limit_counts_other_pending_admissions(self): + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + adapter._max_concurrent_runs = 1 + adapter._pending_agent_requests = 1 + + response = adapter._concurrency_limited_response() + + assert response is not None + assert response.status == 429 + + @pytest.mark.asyncio + async def test_concurrency_limit_excludes_current_pending_admission(self): + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + adapter._max_concurrent_runs = 1 + app = _make_admission_app(adapter) + + async with TestClient(TestServer(app)) as client: + with patch.object(adapter, "_run_agent", new=AsyncMock(return_value=({}, {}))): + response = await client.post( + "/api/sessions/s/chat", + json={"message": "hello"}, + ) + + assert response.status == 404 + + def test_counts_pending_admission_before_agent_bookkeeping(self): + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + adapter._pending_agent_requests = 1 + + assert adapter.active_agent_work_count() == 1 + def test_counts_live_run_task_before_agent_creation(self): adapter = APIServerAdapter(PlatformConfig(enabled=True)) adapter._inflight_agent_runs = 2 @@ -226,3 +257,70 @@ class TestDrainAdmission: assert response.status == 503 assert response.headers["Retry-After"] == "1" assert payload["error"]["code"] == "gateway_draining" + + @pytest.mark.asyncio + async def test_external_drain_refuses_every_agent_start_endpoint(self): + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + runner = SimpleNamespace(_draining=False, _external_drain_active=True) + app = _make_admission_app(adapter) + paths = ( + "/api/sessions/missing/chat", + "/api/sessions/missing/chat/stream", + "/v1/chat/completions", + "/v1/responses", + "/v1/runs", + ) + + with patch("gateway.run._gateway_runner_ref", lambda: runner): + async with TestClient(TestServer(app)) as client: + for path in paths: + response = await client.post(path, json={}) + payload = await response.json() + + assert response.status == 503 + assert payload["error"]["code"] == "gateway_draining" + + @pytest.mark.asyncio + async def test_admitted_request_blocks_drain_before_agent_bookkeeping(self): + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + runner, _adapter = make_restart_runner() + runner.adapters = {Platform.API_SERVER: adapter} + app = _make_admission_app(adapter) + body_read_started = asyncio.Event() + allow_body_read = asyncio.Event() + + async def delayed_read_json(_request): + body_read_started.set() + await allow_body_read.wait() + return {"message": "hello"}, None + + with patch.object( + adapter, + "_get_existing_session_or_404", + return_value=({}, None), + ), patch.object( + adapter, + "_read_json_body", + side_effect=delayed_read_json, + ), patch.object( + adapter, + "_run_agent", + new=AsyncMock(return_value=({"final_response": "done"}, {})), + ): + async with TestClient(TestServer(app)) as client: + request_task = asyncio.create_task( + client.post("/api/sessions/missing/chat", json={}) + ) + await body_read_started.wait() + + assert adapter._pending_agent_requests == 1 + drain_task = asyncio.create_task(runner._drain_active_agents(2.0)) + await asyncio.sleep(0.1) + assert not drain_task.done() + + allow_body_read.set() + response = await request_task + assert response.status == 200 + _snapshot, timed_out = await drain_task + + assert timed_out is False diff --git a/tests/gateway/test_api_server_jobs.py b/tests/gateway/test_api_server_jobs.py index 082ab6cf1671..71f35619505d 100644 --- a/tests/gateway/test_api_server_jobs.py +++ b/tests/gateway/test_api_server_jobs.py @@ -11,6 +11,7 @@ Covers: """ import logging +from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest @@ -494,6 +495,19 @@ class TestRunJob: assert data["job"] == triggered_job mock_trigger.assert_called_once_with(VALID_JOB_ID) + @pytest.mark.asyncio + async def test_run_job_refuses_during_gateway_drain(self, adapter): + app = _create_app(adapter) + runner = SimpleNamespace(_draining=False, _external_drain_active=True) + + with patch("gateway.run._gateway_runner_ref", lambda: runner): + async with TestClient(TestServer(app)) as cli: + resp = await cli.post(f"/api/jobs/{VALID_JOB_ID}/run") + payload = await resp.json() + + assert resp.status == 503 + assert payload["error"]["code"] == "gateway_draining" + # --------------------------------------------------------------------------- # 17. test_auth_required diff --git a/tests/gateway/test_cron_fire_webhook.py b/tests/gateway/test_cron_fire_webhook.py index 06eaba671dae..fad9f4de1021 100644 --- a/tests/gateway/test_cron_fire_webhook.py +++ b/tests/gateway/test_cron_fire_webhook.py @@ -8,6 +8,8 @@ test_chronos_verify.py. """ import asyncio +from types import SimpleNamespace +from unittest.mock import patch import pytest from aiohttp import web @@ -108,6 +110,31 @@ async def test_missing_token_401(adapter, monkeypatch): assert spy.fired == [] +@pytest.mark.asyncio +async def test_valid_token_refuses_during_gateway_drain(adapter, monkeypatch): + spy = _SpyProvider() + runner = SimpleNamespace(_draining=False, _external_drain_active=True) + monkeypatch.setattr("cron.scheduler_provider.resolve_cron_scheduler", lambda: spy) + monkeypatch.setattr( + "plugins.cron_providers.chronos.verify.get_fire_verifier", + lambda: (lambda **kw: {"purpose": "cron_fire"}), + ) + + app = _create_app(adapter) + with patch("gateway.run._gateway_runner_ref", lambda: runner): + async with TestClient(TestServer(app)) as cli: + response = await cli.post( + "/api/cron/fire", + headers={"Authorization": "Bearer good"}, + json={"job_id": "abc123"}, + ) + payload = await response.json() + + assert response.status == 503 + assert payload["error"]["code"] == "gateway_draining" + assert spy.fired == [] + + @pytest.mark.asyncio async def test_missing_job_id_400(adapter, monkeypatch): """Valid token but no job_id → 400, no fire.""" diff --git a/tests/gateway/test_external_drain_control.py b/tests/gateway/test_external_drain_control.py index d37b42224359..eb1c42443899 100644 --- a/tests/gateway/test_external_drain_control.py +++ b/tests/gateway/test_external_drain_control.py @@ -19,6 +19,7 @@ import pytest import gateway.drain_control as dc from gateway.run import GatewayRunner +from gateway.config import Platform from gateway.platforms.base import MessageEvent, MessageType from tests.gateway.restart_test_helpers import make_restart_runner, make_restart_source @@ -238,6 +239,16 @@ def _drain_runner(): class TestDrainStateMachine: + def test_active_work_count_includes_api_and_cron_work(self, monkeypatch): + runner, _ = _drain_runner() + runner.adapters = { + Platform.API_SERVER: MagicMock(active_agent_work_count=MagicMock(return_value=2)) + } + runner._running_agents = {"session": MagicMock()} + monkeypatch.setattr("cron.scheduler.get_running_job_ids", lambda: {"job-1"}) + + assert runner._active_work_count() == 4 + def test_enter_sets_flag_and_flips_state(self): runner, _ = _drain_runner() runner._enter_external_drain() @@ -289,6 +300,26 @@ class TestDrainStateMachine: class TestDrainWatcher: + @pytest.mark.asyncio + async def test_watcher_persists_aggregate_work_during_external_drain(self, home, monkeypatch): + runner, _ = _drain_runner() + runner._drain_control_watcher = GatewayRunner._drain_control_watcher.__get__( + runner, GatewayRunner + ) + runner._persist_active_agents = MagicMock() + dc.write_drain_request() + task = asyncio.create_task(runner._drain_control_watcher(interval=0.01)) + await asyncio.sleep(0.03) + runner._running = False + await asyncio.sleep(0.02) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + runner._persist_active_agents.assert_called() + @pytest.mark.asyncio async def test_watcher_enters_then_exits_with_marker(self, home): runner, _ = _drain_runner()