From 6cf8be6bfa12c645850040d33beccdffab0a0b37 Mon Sep 17 00:00:00 2001 From: HaoKang-Timmy Date: Thu, 13 Aug 2026 23:26:00 +0000 Subject: [PATCH] Add opt-in ThunderAgent scheduler support --- src/polar/config/README.md | 46 ++ src/polar/config/topology.py | 1 + src/polar/gateway/node.py | 27 +- src/polar/gateway/proxy.py | 183 +++++- src/polar/gateway/server.py | 34 +- tests/config/test_topology.py | 30 +- tests/gateway/test_thunderagent.py | 633 +++++++++++++++++++ tests/gateway/test_thunderagent_lifecycle.py | 136 ++++ 8 files changed, 1056 insertions(+), 34 deletions(-) create mode 100644 tests/gateway/test_thunderagent.py create mode 100644 tests/gateway/test_thunderagent_lifecycle.py diff --git a/src/polar/config/README.md b/src/polar/config/README.md index f1f1d5707..05adab05f 100644 --- a/src/polar/config/README.md +++ b/src/polar/config/README.md @@ -60,6 +60,7 @@ falls back to `rollout.public_url` when omitted. | `public_url` | str | derived from `host:port` | | `model_served` | str | `""` | | `inference.engine` | `sglang` \| `vllm` | `sglang` | +| `inference.scheduler` | `none` \| `thunderagent` | `none` | | `inference.base_url` | str | `http://127.0.0.1:8000` | | `max_init_workers` | int | `4` | | `max_run_workers` | int | `2` | @@ -88,9 +89,54 @@ gateway: max_postrun_workers: 4 inference: engine: sglang # or vllm + scheduler: none # or thunderagent; opt-in base_url: http://127.0.0.1:8000 ``` +## ThunderAgent scheduler + +ThunderAgent runs as an external OpenAI-compatible proxy; Polar does not start +or configure it. `inference.engine` remains `sglang` or `vllm` (matching the +backend behind ThunderAgent), while `inference.base_url` points to the +ThunderAgent listener. The default `scheduler: none` keeps the direct inference +path unchanged. Polar uses `:` as the stable +ThunderAgent program identity, so gateways sharing one proxy remain isolated. + +Install the external ThunderAgent checkout in editable mode if needed: + +```bash +python -m pip install -e /path/to/ThunderAgent +``` + +For example, with a vLLM server on port 8000, start ThunderAgent with the fixed +rollout settings: + +```bash +thunderagent \ + --backend-type vllm \ + --backends http://127.0.0.1:8000 \ + --port 9000 \ + --router tr \ + --metrics \ + --acting-token-weight 0.0 +``` + +Then opt the gateway node in and point it at that proxy: + +```yaml +inference: + engine: vllm + scheduler: thunderagent + base_url: http://127.0.0.1:9000 +``` + +When gateways share a ThunderAgent instance, coordinate each weight update as +one shared operation: pause and drain all gateways, call +`POST /weight_sync/begin` once per ThunderAgent instance, update the weights, +call `POST /weight_sync/end` once per instance, and then resume all gateways. +Do not map each gateway's pause/resume pair to its own ThunderAgent begin/end +pair, which can release the shared barrier too early. + ## Reachable URLs and multi-node `public_url`s must be reachable by whoever calls them: the rollout server calls diff --git a/src/polar/config/topology.py b/src/polar/config/topology.py index ffa3f0826..d79b728c8 100644 --- a/src/polar/config/topology.py +++ b/src/polar/config/topology.py @@ -21,6 +21,7 @@ class _StrictModel(BaseModel): class _InferenceConfig(_StrictModel): engine: Literal["sglang", "vllm"] = "sglang" + scheduler: Literal["none", "thunderagent"] = "none" base_url: str = "http://127.0.0.1:8000" @field_validator("base_url") diff --git a/src/polar/gateway/node.py b/src/polar/gateway/node.py index 5dc72cd5d..59c3dfca2 100644 --- a/src/polar/gateway/node.py +++ b/src/polar/gateway/node.py @@ -8,6 +8,7 @@ from contextlib import suppress from pathlib import Path from tempfile import mkdtemp +from typing import Awaitable, Callable import httpx @@ -63,6 +64,7 @@ def __init__( session_base_dir: str | None = None, rollout_server_url: str | None = None, heartbeat_interval_seconds: int = 30, + program_releaser: Callable[[str], Awaitable[bool]] | None = None, ) -> None: self.node_id = node_id self.gateway_url = gateway_url.rstrip("/") @@ -90,6 +92,7 @@ def __init__( self._heartbeat_interval_seconds = heartbeat_interval_seconds self._control_client: httpx.AsyncClient | None = None self._heartbeat_task: asyncio.Task[None] | None = None + self._program_releaser = program_releaser async def start(self) -> None: await self._dispatcher.start() @@ -201,6 +204,21 @@ async def dispatch(self, request: SessionDispatchRequest) -> None: async def cancel(self, session_id: str) -> bool: return await self._dispatcher.cancel(session_id) + async def release_program(self, session_id: str) -> None: + """Best-effort release of external scheduler state for one session.""" + if self._program_releaser is None: + return + try: + await self._program_releaser(session_id) + except asyncio.CancelledError: + raise + except Exception: + logger.warning( + "Failed to release scheduler program for session %s", + session_id, + exc_info=True, + ) + async def active_sessions(self) -> int: return await self._dispatcher.active_count() @@ -552,9 +570,12 @@ async def _handle_postrun(self, managed: ManagedSession) -> None: # status/task_id visible for debugging via the polling endpoint. self.session_registry.clear_result_payload(request.session_id) finally: - await self._remove_session_dir_best_effort( - managed.session_dir, request.session_id - ) + try: + await self.release_program(request.session_id) + finally: + await self._remove_session_dir_best_effort( + managed.session_dir, request.session_id + ) async def _build_session_result(self, managed: ManagedSession) -> SessionResult: request = managed.request diff --git a/src/polar/gateway/proxy.py b/src/polar/gateway/proxy.py index 7b7ab3331..faa182b0b 100644 --- a/src/polar/gateway/proxy.py +++ b/src/polar/gateway/proxy.py @@ -10,7 +10,8 @@ import asyncio import json import logging -from typing import Any +from typing import Any, Literal +from urllib.parse import quote import httpx @@ -56,7 +57,7 @@ class UpstreamTransportError(UpstreamError): class InferenceClient: - """Direct httpx client to an inference server's OpenAI-compatible API. + """HTTP client to an inference server or opted-in scheduler proxy. Per-call bound comes from the session's remaining-timeout budget (`_await_with_budget` at the gateway node). The internal httpx timeout @@ -67,13 +68,28 @@ class InferenceClient: _LIVENESS_TIMEOUT_SECONDS = 900.0 - def __init__(self, base_url: str, engine: InferenceEngine): + def __init__( + self, + base_url: str, + engine: InferenceEngine, + *, + scheduler: Literal["none", "thunderagent"] = "none", + program_namespace: str | None = None, + ) -> None: self.base_url = base_url.rstrip("/") self.engine = engine + self.scheduler = scheduler + self.program_namespace = program_namespace self._client: httpx.AsyncClient | None = None self._generation_paused = False self._inflight_generations = 0 self._generation_condition = asyncio.Condition() + self._program_condition = asyncio.Condition() + self._program_inflight: dict[str, int] = {} + self._programs_may_exist: set[str] = set() + self._terminal_programs: set[str] = set() + self._program_release_tasks: dict[str, asyncio.Task[None]] = {} + self._closing = False async def _get_client(self) -> httpx.AsyncClient: if self._client is None or self._client.is_closed: @@ -111,29 +127,134 @@ def _translate_transport_error(exc: httpx.RequestError) -> UpstreamError: return UpstreamTimeoutError("Upstream request timed out") return UpstreamTransportError(f"Upstream request failed: {exc}") - async def completion(self, request: dict[str, Any]) -> dict[str, Any]: + async def completion( + self, + request: dict[str, Any], + *, + session_id: str | None = None, + ) -> dict[str, Any]: """Non-streaming chat completion. Returns the full JSON response.""" - await self._acquire_generation_slot() + if self.scheduler == "thunderagent" and not session_id: + raise ValueError("ThunderAgent completions require a session_id") + program_registered = False + program_id: str | None = None + slot_acquired = False + if self.scheduler == "thunderagent": + assert session_id is not None + program_id = await self._begin_program_request(session_id) + program_registered = True + try: + await self._acquire_generation_slot() + slot_acquired = True + client = await self._get_client() + from copy import deepcopy + + request_copy = deepcopy(request) + request_copy.pop("stream", None) + request_copy["stream"] = False + request_copy = self.engine.prepare_request(request_copy) + headers = {"Content-Type": "application/json"} + if program_id is not None: + # ThunderAgent checks body fields before X-Session-ID. Reserve + # the program identity for Polar so callers cannot split one session. + request_copy.pop("program_id", None) + extra_body = request_copy.get("extra_body") + if isinstance(extra_body, dict): + extra_body.pop("program_id", None) + headers["X-Session-ID"] = program_id + try: + resp = await client.post( + "/v1/chat/completions", + json=request_copy, + headers=headers, + ) + except httpx.RequestError as exc: + raise self._translate_transport_error(exc) from exc + + await self._raise_for_status(resp) + return self.engine.normalize_response(resp.json()) + finally: + try: + if slot_acquired: + await self._release_generation_slot() + finally: + if program_registered: + assert session_id is not None + await self._end_program_request(session_id) + + async def release_program(self, session_id: str) -> bool: + """Release one program from an opted-in ThunderAgent scheduler.""" + if self.scheduler != "thunderagent": + return False + async with self._program_condition: + self._terminal_programs.add(session_id) + await self._program_condition.wait_for( + lambda: self._program_inflight.get(session_id, 0) == 0 + ) + if session_id not in self._programs_may_exist: + return False + release_task = self._program_release_tasks.get(session_id) + owns_task = release_task is None + if release_task is None: + release_task = asyncio.create_task( + self._release_program_upstream(self._program_id(session_id)) + ) + self._program_release_tasks[session_id] = release_task + try: + await asyncio.shield(release_task) + except asyncio.CancelledError: + raise + except Exception: + async with self._program_condition: + if self._program_release_tasks.get(session_id) is release_task: + self._program_release_tasks.pop(session_id, None) + if not owns_task: + return await self.release_program(session_id) + raise + async with self._program_condition: + if self._program_release_tasks.get(session_id) is release_task: + self._program_release_tasks.pop(session_id, None) + self._programs_may_exist.discard(session_id) + return owns_task + + def _program_id(self, session_id: str) -> str: + if self.program_namespace: + return f"{quote(self.program_namespace, safe='')}:{session_id}" + return session_id + + async def _begin_program_request(self, session_id: str) -> str: + async with self._program_condition: + if self._closing or session_id in self._terminal_programs: + raise UpstreamHTTPError( + 409, + {"error": {"message": f"Session {session_id} has terminated"}}, + ) + self._program_inflight[session_id] = self._program_inflight.get(session_id, 0) + 1 + # Mark before any await to make release safe even if the proxy POST + # is cancelled after ThunderAgent has accepted it. + self._programs_may_exist.add(session_id) + return self._program_id(session_id) + + async def _end_program_request(self, session_id: str) -> None: + async with self._program_condition: + remaining = self._program_inflight.get(session_id, 0) - 1 + if remaining > 0: + self._program_inflight[session_id] = remaining + else: + self._program_inflight.pop(session_id, None) + self._program_condition.notify_all() + + async def _release_program_upstream(self, program_id: str) -> None: client = await self._get_client() - from copy import deepcopy - - request_copy = deepcopy(request) - request_copy.pop("stream", None) - request_copy["stream"] = False - request_copy = self.engine.prepare_request(request_copy) try: - resp = await client.post( - "/v1/chat/completions", - json=request_copy, - headers={"Content-Type": "application/json"}, + response = await client.post( + "/programs/release", + json={"program_id": program_id}, + timeout=5.0, ) except httpx.RequestError as exc: raise self._translate_transport_error(exc) from exc - finally: - await self._release_generation_slot() - - await self._raise_for_status(resp) - return self.engine.normalize_response(resp.json()) + await self._raise_for_status(response) async def _acquire_generation_slot(self) -> None: async with self._generation_condition: @@ -201,6 +322,26 @@ async def health(self) -> dict[str, Any]: except json.JSONDecodeError: return {"status": "ok", "body": text} - async def close(self): + async def close(self) -> None: + program_ids: list[str] = [] + if self.scheduler == "thunderagent": + async with self._program_condition: + self._closing = True + program_ids = list(self._programs_may_exist) + async with self._generation_condition: + self._generation_paused = False + self._generation_condition.notify_all() + if program_ids: + results = await asyncio.gather( + *(self.release_program(program_id) for program_id in program_ids), + return_exceptions=True, + ) + for program_id, result in zip(program_ids, results): + if isinstance(result, BaseException): + logger.warning( + "Failed to release ThunderAgent program %s during shutdown: %s", + program_id, + result, + ) if self._client and not self._client.is_closed: await self._client.aclose() diff --git a/src/polar/gateway/server.py b/src/polar/gateway/server.py index 1bc9466ea..dbbd44442 100644 --- a/src/polar/gateway/server.py +++ b/src/polar/gateway/server.py @@ -5,7 +5,6 @@ import asyncio from contextlib import asynccontextmanager from dataclasses import dataclass -import hashlib import json import logging import os @@ -42,7 +41,6 @@ from polar.gateway.transform.base import BaseTransformer from polar.platform.events import SSE_HEADERS, EventBus from polar.rollout.models import SessionDispatchRequest, SessionDispatchResponse, SessionStatus -from polar.runtime.models import RuntimeSpec from polar.trajectory.registry import default_builder_registry, default_evaluator_registry logging.basicConfig( @@ -81,7 +79,12 @@ def configure_server(topology_path: str = "topology.yaml", *, node_id: str | Non def _build_state(topology: TopologyConfig, node_id: str | None) -> GatewayState: node = topology.select_gateway_node(node_id) - inference = InferenceClient(node.inference_base_url, get_engine(node.engine)) + inference = InferenceClient( + node.inference_base_url, + get_engine(node.engine), + scheduler=node.inference.scheduler, + program_namespace=node.id, + ) persistence_config = topology.gateway.completion_persistence save_dir = topology.rollout.save_dir completion_writer = CompletionWriter( @@ -111,6 +114,9 @@ def _build_state(topology: TopologyConfig, node_id: str | None) -> GatewayState: default_runtime=node.default_runtime, rollout_server_url=topology.gateway.rollout_server_url or None, heartbeat_interval_seconds=topology.gateway.heartbeat_interval_seconds, + program_releaser=( + inference.release_program if node.inference.scheduler == "thunderagent" else None + ), ) return GatewayState( topology=topology, @@ -187,10 +193,16 @@ async def _lifespan(_: FastAPI): try: yield finally: - await state.node_manager.close() - await state.inference.close() - state.storage.close() - await state.completion_writer.close() + try: + await state.node_manager.close() + finally: + try: + await state.inference.close() + finally: + try: + state.storage.close() + finally: + await state.completion_writer.close() app = FastAPI(title="Polar Gateway", version="0.1.0", lifespan=_lifespan) @@ -598,6 +610,10 @@ async def delete_session(session_id: str): if info is None and deleted_count == 0: raise HTTPException(status_code=404, detail="Session not found") + await state.node_manager.release_program(safe_session_id) + # A completion that was already in flight may have persisted after the + # first delete but before scheduler draining finished. + deleted_count += state.storage.delete_session(safe_session_id) state.session_registry.remove(safe_session_id) return SessionDeleteResponse( session_id=safe_session_id, @@ -680,7 +696,7 @@ async def _handle_non_streaming( ) -> JSONResponse: state = get_state() try: - response = await state.inference.completion(openai_request) + response = await state.inference.completion(openai_request, session_id=session_id) except UpstreamError as exc: logger.warning("Non-streaming upstream error for session %s: %s", session_id, exc) return _upstream_error_response(api_type, exc) @@ -715,7 +731,7 @@ async def _handle_streaming( non_stream_request = {k: v for k, v in openai_request.items() if k != "stream_options"} non_stream_request["stream"] = False try: - response = await state.inference.completion(non_stream_request) + response = await state.inference.completion(non_stream_request, session_id=session_id) except UpstreamError as exc: logger.warning("Upstream error for streaming session %s: %s", session_id, exc) return _upstream_error_response(api_type, exc) diff --git a/tests/config/test_topology.py b/tests/config/test_topology.py index 6065ee990..21f9e7ae2 100644 --- a/tests/config/test_topology.py +++ b/tests/config/test_topology.py @@ -89,7 +89,11 @@ def test_inference_block_selects_engine_and_base_url(tmp_path: Path) -> None: { "id": "node-a", "public_url": "http://127.0.0.1:8100", - "inference": {"engine": "vllm", "base_url": "http://127.0.0.1:8000"}, + "inference": { + "engine": "vllm", + "scheduler": "thunderagent", + "base_url": "http://127.0.0.1:8000", + }, } ], }, @@ -97,6 +101,7 @@ def test_inference_block_selects_engine_and_base_url(tmp_path: Path) -> None: ) node = TopologyConfig.load(path).gateway.nodes[0] assert node.engine == "vllm" + assert node.inference.scheduler == "thunderagent" assert node.inference_base_url == "http://127.0.0.1:8000" @@ -117,6 +122,7 @@ def test_inference_engine_defaults_to_sglang(tmp_path: Path) -> None: ) node = TopologyConfig.load(path).gateway.nodes[0] assert node.engine == "sglang" + assert node.inference.scheduler == "none" def test_inference_defaults_when_block_omitted(tmp_path: Path) -> None: @@ -148,6 +154,28 @@ def test_invalid_inference_engine_is_rejected(tmp_path: Path) -> None: TopologyConfig.load(path) +def test_invalid_inference_scheduler_is_rejected(tmp_path: Path) -> None: + path = _write_yaml( + tmp_path / "topology.yaml", + { + "gateway": { + "nodes": [ + { + "id": "node-a", + "public_url": "http://127.0.0.1:8100", + "inference": { + "scheduler": "other", + "base_url": "http://127.0.0.1:8000", + }, + } + ], + }, + }, + ) + with pytest.raises(ValueError, match="scheduler"): + TopologyConfig.load(path) + + def test_invalid_inference_base_url_is_rejected(tmp_path: Path) -> None: path = _write_yaml( tmp_path / "topology.yaml", diff --git a/tests/gateway/test_thunderagent.py b/tests/gateway/test_thunderagent.py new file mode 100644 index 000000000..082dddc13 --- /dev/null +++ b/tests/gateway/test_thunderagent.py @@ -0,0 +1,633 @@ +from __future__ import annotations + +import asyncio +import json +from collections import Counter +from collections.abc import Callable +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock + +import httpx +import pytest +import yaml + +from polar.config import TopologyConfig +from polar.gateway import server +from polar.gateway.detection import APIType +from polar.gateway.engine import SGLangEngine, VLLMEngine +from polar.gateway.proxy import InferenceClient, UpstreamHTTPError +from polar.gateway.transform.openai_chat import OpenAIChatTransformer + + +def _completion_response() -> httpx.Response: + return httpx.Response( + 200, + json={ + "id": "chatcmpl-test", + "model": "test-model", + "prompt_token_ids": [1, 2], + "choices": [ + { + "message": { + "role": "assistant", + "content": "ok", + "reasoning": "because", + }, + "finish_reason": "stop", + "token_ids": [10, 11], + "logprobs": { + "content": [ + {"token": "o", "logprob": -0.1}, + {"token": "k", "logprob": -0.2}, + ] + }, + } + ], + }, + ) + + +def _make_client( + handler: Callable[[httpx.Request], httpx.Response], + *, + scheduler: str = "thunderagent", +) -> tuple[InferenceClient, httpx.AsyncClient]: + client = InferenceClient( + "http://inference:9000", + VLLMEngine(), + scheduler=scheduler, + program_namespace="node-a", + ) + transport_client = httpx.AsyncClient( + base_url="http://inference:9000", + transport=httpx.MockTransport(handler), + ) + client._client = transport_client + return client, transport_client + + +@pytest.mark.asyncio +async def test_thunderagent_requires_session_identity() -> None: + client, _ = _make_client(lambda _: _completion_response()) + try: + with pytest.raises(ValueError, match="require a session_id"): + await client.completion({"messages": []}) + finally: + await client.close() + + +@pytest.mark.asyncio +async def test_release_before_first_completion_fences_session() -> None: + client, _ = _make_client(lambda _: _completion_response()) + assert await client.release_program("session-a") is False + with pytest.raises(UpstreamHTTPError) as exc_info: + await client.completion({"messages": []}, session_id="session-a") + assert exc_info.value.status_code == 409 + await client.close() + + +@pytest.mark.asyncio +async def test_default_scheduler_sends_no_session_header_and_release_is_noop() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + assert request.url.path == "/v1/chat/completions" + body = json.loads(request.content) + assert body["program_id"] == "caller-controlled" + assert body["extra_body"]["program_id"] == "caller-controlled" + return _completion_response() + + client, _ = _make_client(handler, scheduler="none") + try: + await client.completion( + { + "messages": [], + "program_id": "caller-controlled", + "extra_body": {"program_id": "caller-controlled"}, + }, + session_id="session-a", + ) + assert await client.release_program("session-a") is False + finally: + await client.close() + + assert [request.url.path for request in requests] == ["/v1/chat/completions"] + assert "x-session-id" not in requests[0].headers + + +@pytest.mark.asyncio +async def test_thunderagent_session_identity_and_vllm_training_fields_are_preserved() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + if request.url.path == "/v1/chat/completions": + return _completion_response() + if request.url.path == "/programs/release": + return httpx.Response(200, json={"released": True}) + raise AssertionError(f"unexpected request: {request.method} {request.url}") + + client, _ = _make_client(handler) + responses: list[dict] = [] + try: + for session_id in ("session-a", "session-a", "session-b"): + responses.append( + await client.completion( + { + "messages": [], + "program_id": "caller-controlled", + "extra_body": {"program_id": "caller-controlled"}, + }, + session_id=session_id, + ) + ) + assert await client.release_program("session-a") is True + assert await client.release_program("session-b") is True + finally: + await client.close() + + completions = [request for request in requests if request.url.path == "/v1/chat/completions"] + assert [request.headers["x-session-id"] for request in completions] == [ + "node-a:session-a", + "node-a:session-a", + "node-a:session-b", + ] + for request in completions: + body = json.loads(request.content) + assert body["stream"] is False + assert body["logprobs"] is True + assert body["return_token_ids"] is True + assert body["top_logprobs"] == 0 + assert "program_id" not in body + assert "program_id" not in body["extra_body"] + + for response in responses: + choice = response["choices"][0] + assert response["prompt_token_ids"] == [1, 2] + assert choice["token_ids"] == [10, 11] + assert [entry["token_id"] for entry in choice["logprobs"]["content"]] == [10, 11] + assert choice["message"]["reasoning_content"] == "because" + + +@pytest.mark.asyncio +async def test_gateway_namespace_isolates_equal_session_ids() -> None: + program_ids: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/v1/chat/completions": + program_ids.append(request.headers["x-session-id"]) + return _completion_response() + program_ids.append(json.loads(request.content)["program_id"]) + return httpx.Response(200, json={"released": True}) + + for namespace in ("node-a", "node-b"): + client = InferenceClient( + "http://inference:9000", + VLLMEngine(), + scheduler="thunderagent", + program_namespace=namespace, + ) + client._client = httpx.AsyncClient( + base_url="http://inference:9000", + transport=httpx.MockTransport(handler), + ) + await client.completion({"messages": []}, session_id="same-session") + await client.release_program("same-session") + await client.close() + + assert program_ids == [ + "node-a:same-session", + "node-a:same-session", + "node-b:same-session", + "node-b:same-session", + ] + + +@pytest.mark.asyncio +async def test_thunderagent_preserves_sglang_training_contract() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + if request.url.path == "/programs/release": + return httpx.Response(200, json={"released": True}) + return httpx.Response( + 200, + json={ + "choices": [ + { + "prompt_token_ids": [1, 2], + "message": {"role": "assistant", "content": "ok"}, + "logprobs": { + "content": [ + {"token": "o", "logprob": -0.1}, + {"token": "k", "logprob": -0.2}, + ] + }, + "meta_info": { + "output_token_logprobs": [ + [-0.1, 10, "o"], + [-0.2, 11, "k"], + ] + }, + } + ] + }, + ) + + client = InferenceClient( + "http://inference:9000", + SGLangEngine(), + scheduler="thunderagent", + program_namespace="node-a", + ) + client._client = httpx.AsyncClient( + base_url="http://inference:9000", + transport=httpx.MockTransport(handler), + ) + try: + response = await client.completion({"messages": []}, session_id="session-a") + assert await client.release_program("session-a") is True + finally: + await client.close() + + completion = requests[0] + body = json.loads(completion.content) + assert completion.headers["x-session-id"] == "node-a:session-a" + assert body["return_prompt_token_ids"] is True + assert body["return_meta_info"] is True + choice = response["choices"][0] + assert choice["input_token_ids"] == [1, 2] + assert choice["token_ids"] == [10, 11] + assert [entry["token_id"] for entry in choice["logprobs"]["content"]] == [10, 11] + + +@pytest.mark.asyncio +async def test_concurrent_release_of_one_program_posts_once() -> None: + release_requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/v1/chat/completions": + return _completion_response() + if request.url.path == "/programs/release": + release_requests.append(request) + return httpx.Response(200, json={"released": True}) + raise AssertionError(f"unexpected request: {request.method} {request.url}") + + client, _ = _make_client(handler) + try: + await client.completion({"messages": []}, session_id="session-a") + results = await asyncio.gather( + client.release_program("session-a"), + client.release_program("session-a"), + ) + finally: + await client.close() + + assert sorted(results) == [False, True] + assert len(release_requests) == 1 + assert json.loads(release_requests[0].content) == {"program_id": "node-a:session-a"} + + +@pytest.mark.asyncio +async def test_release_waits_for_paused_completion_and_fences_late_requests() -> None: + paths: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + paths.append(request.url.path) + if request.url.path == "/v1/chat/completions": + return _completion_response() + return httpx.Response(200, json={"released": True}) + + client, _ = _make_client(handler) + await client.pause_generation() + completion_task = asyncio.create_task( + client.completion({"messages": []}, session_id="session-a") + ) + await asyncio.sleep(0) + assert client._program_inflight == {"session-a": 1} + + release_task = asyncio.create_task(client.release_program("session-a")) + await asyncio.sleep(0) + assert not release_task.done() + + await client.resume_generation() + await completion_task + assert await release_task is True + + assert paths == ["/v1/chat/completions", "/programs/release"] + with pytest.raises(UpstreamHTTPError) as exc_info: + await client.completion({"messages": []}, session_id="session-a") + assert exc_info.value.status_code == 409 + await client.close() + + +@pytest.mark.asyncio +async def test_failed_release_can_be_retried() -> None: + release_attempts = 0 + + def handler(request: httpx.Request) -> httpx.Response: + nonlocal release_attempts + if request.url.path == "/v1/chat/completions": + return _completion_response() + if request.url.path == "/programs/release": + release_attempts += 1 + if release_attempts == 1: + return httpx.Response(500, json={"error": {"message": "release failed"}}) + return httpx.Response(200, json={"released": True}) + raise AssertionError(f"unexpected request: {request.method} {request.url}") + + client, _ = _make_client(handler) + try: + await client.completion({"messages": []}, session_id="session-a") + with pytest.raises(UpstreamHTTPError, match="release failed"): + await client.release_program("session-a") + assert await client.release_program("session-a") is True + finally: + await client.close() + + assert release_attempts == 2 + + +@pytest.mark.asyncio +async def test_concurrent_waiter_retries_failed_release() -> None: + first_release_started = asyncio.Event() + finish_first_release = asyncio.Event() + release_attempts = 0 + + async def handler(request: httpx.Request) -> httpx.Response: + nonlocal release_attempts + if request.url.path == "/v1/chat/completions": + return _completion_response() + release_attempts += 1 + if release_attempts == 1: + first_release_started.set() + await finish_first_release.wait() + return httpx.Response(500, json={"error": {"message": "release failed"}}) + return httpx.Response(200, json={"released": True}) + + client = InferenceClient( + "http://inference:9000", + VLLMEngine(), + scheduler="thunderagent", + program_namespace="node-a", + ) + client._client = httpx.AsyncClient( + base_url="http://inference:9000", + transport=httpx.MockTransport(handler), + ) + await client.completion({"messages": []}, session_id="session-a") + first = asyncio.create_task(client.release_program("session-a")) + await first_release_started.wait() + second = asyncio.create_task(client.release_program("session-a")) + finish_first_release.set() + results = await asyncio.gather(first, second, return_exceptions=True) + + assert release_attempts == 2 + assert sum(result is True for result in results) == 1 + assert sum(isinstance(result, UpstreamHTTPError) for result in results) == 1 + await client.close() + + +@pytest.mark.asyncio +async def test_close_releases_each_active_program_once() -> None: + released_programs: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/v1/chat/completions": + return _completion_response() + if request.url.path == "/programs/release": + released_programs.append(json.loads(request.content)["program_id"]) + return httpx.Response(200, json={"released": True}) + raise AssertionError(f"unexpected request: {request.method} {request.url}") + + client, transport_client = _make_client(handler) + await client.completion({"messages": []}, session_id="session-a") + await client.completion({"messages": []}, session_id="session-a") + await client.completion({"messages": []}, session_id="session-b") + + await client.close() + await client.close() + + assert Counter(released_programs) == Counter({"node-a:session-a": 1, "node-a:session-b": 1}) + assert transport_client.is_closed + + +@pytest.mark.asyncio +async def test_close_releases_program_while_completion_is_in_flight() -> None: + completion_started = asyncio.Event() + allow_completion = asyncio.Event() + released_programs: list[str] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/v1/chat/completions": + completion_started.set() + await allow_completion.wait() + return _completion_response() + if request.url.path == "/programs/release": + released_programs.append(json.loads(request.content)["program_id"]) + return httpx.Response(200, json={"released": True}) + raise AssertionError(f"unexpected request: {request.method} {request.url}") + + client = InferenceClient( + "http://inference:9000", + VLLMEngine(), + scheduler="thunderagent", + program_namespace="node-a", + ) + client._client = httpx.AsyncClient( + base_url="http://inference:9000", + transport=httpx.MockTransport(handler), + ) + completion_task = asyncio.create_task( + client.completion({"messages": []}, session_id="session-a") + ) + await completion_started.wait() + completion_task.cancel() + with pytest.raises(asyncio.CancelledError): + await completion_task + + await client.close() + allow_completion.set() + + assert released_programs == ["node-a:session-a"] + + +@pytest.mark.asyncio +async def test_close_drains_completion_waiting_at_pause_gate() -> None: + paths: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + paths.append(request.url.path) + if request.url.path == "/v1/chat/completions": + return _completion_response() + return httpx.Response(200, json={"released": True}) + + client, _ = _make_client(handler) + await client.pause_generation() + completion_task = asyncio.create_task( + client.completion({"messages": []}, session_id="session-a") + ) + await asyncio.sleep(0) + assert client._program_inflight == {"session-a": 1} + + await asyncio.gather(completion_task, client.close()) + + assert paths == ["/v1/chat/completions", "/programs/release"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +async def test_gateway_handlers_forward_session_identity( + monkeypatch: pytest.MonkeyPatch, + streaming: bool, +) -> None: + response = _completion_response().json() + inference = SimpleNamespace(completion=AsyncMock(return_value=response)) + storage = SimpleNamespace(save_message=Mock()) + monkeypatch.setattr( + server, + "get_state", + lambda: SimpleNamespace( + inference=inference, + storage=storage, + ), + ) + openai_request = { + "model": "served-model", + "messages": [], + "stream": streaming, + } + handler = server._handle_streaming if streaming else server._handle_non_streaming + + await handler( + APIType.OPENAI_CHAT, + OpenAIChatTransformer(), + openai_request, + {"model": "requested-model", "stream": streaming}, + "session-a", + original_model="requested-model", + session_info=None, + ) + + expected_request = {**openai_request, "stream": False} if streaming else openai_request + inference.completion.assert_awaited_once_with(expected_request, session_id="session-a") + + +@pytest.mark.parametrize( + ("scheduler", "expects_releaser"), + [("none", False), ("thunderagent", True)], +) +@pytest.mark.asyncio +async def test_build_state_wires_scheduler_opt_in( + tmp_path, + scheduler: str, + expects_releaser: bool, +) -> None: + topology_path = tmp_path / "topology.yaml" + topology_path.write_text( + yaml.safe_dump( + { + "gateway": { + "nodes": [ + { + "id": "node-a", + "public_url": "http://gateway:8100", + "inference": { + "engine": "vllm", + "scheduler": scheduler, + "base_url": "http://inference:9000", + }, + } + ] + } + } + ) + ) + + state = server._build_state(TopologyConfig.load(topology_path), "node-a") + try: + assert state.inference.scheduler == scheduler + assert state.inference.program_namespace == "node-a" + assert (state.node_manager._program_releaser is not None) is expects_releaser + finally: + await state.node_manager.close() + await state.inference.close() + state.storage.close() + await state.completion_writer.close() + + +@pytest.mark.asyncio +async def test_lifespan_releases_inference_if_node_shutdown_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + node_manager = SimpleNamespace( + start=AsyncMock(), + close=AsyncMock(side_effect=RuntimeError("node shutdown failed")), + ) + inference = SimpleNamespace(close=AsyncMock()) + storage = SimpleNamespace(close=Mock()) + completion_writer = SimpleNamespace(start=AsyncMock(), close=AsyncMock()) + monkeypatch.setattr( + server, + "get_state", + lambda: SimpleNamespace( + node_manager=node_manager, + inference=inference, + storage=storage, + completion_writer=completion_writer, + ), + ) + + with pytest.raises(RuntimeError, match="node shutdown failed"): + async with server._lifespan(None): + pass + + inference.close.assert_awaited_once() + storage.close.assert_called_once() + completion_writer.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_delete_removes_completion_persisted_while_release_drains( + monkeypatch: pytest.MonkeyPatch, +) -> None: + release_started = asyncio.Event() + finish_release = asyncio.Event() + storage = server.SessionStore() + registry = server.SessionRegistry() + registry.register("session-a") + storage.ensure_session("session-a", None, None, None) + + async def release_program(session_id: str) -> None: + release_started.set() + await finish_release.wait() + + node_manager = SimpleNamespace( + cancel=AsyncMock(return_value=True), + release_program=release_program, + ) + monkeypatch.setattr( + server, + "get_state", + lambda: SimpleNamespace( + node_manager=node_manager, + session_registry=registry, + storage=storage, + ), + ) + + delete_task = asyncio.create_task(server.delete_session("session-a")) + await release_started.wait() + storage.save_message( + "session-a", + {"model": "test"}, + {"choices": []}, + ) + finish_release.set() + response = await delete_task + + assert response.messages_deleted == 1 + assert registry.get("session-a") is None + assert storage.get_session_metadata("session-a") is None diff --git a/tests/gateway/test_thunderagent_lifecycle.py b/tests/gateway/test_thunderagent_lifecycle.py new file mode 100644 index 000000000..fb72d0e76 --- /dev/null +++ b/tests/gateway/test_thunderagent_lifecycle.py @@ -0,0 +1,136 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from pathlib import Path +from unittest.mock import AsyncMock + +import pytest + +from polar.agent.models import AgentSpec +from polar.gateway.dispatcher import ManagedSession +from polar.gateway.node import GatewayNodeManager +from polar.gateway.session import SessionRegistry +from polar.gateway.storage import SessionStore +from polar.rollout.models import SessionDispatchRequest, SessionResult, SessionStatus +from polar.rollout.timer import StageTimer +from polar.trajectory.models import Trajectory +from polar.trajectory.registry import StrategyRegistry + + +def _manager( + program_releaser: Callable[[str], Awaitable[bool]], +) -> tuple[GatewayNodeManager, SessionRegistry, SessionStore]: + registry = SessionRegistry() + storage = SessionStore() + manager = GatewayNodeManager( + node_id="node-1", + gateway_url="http://gateway.test", + max_init_workers=1, + max_run_workers=1, + max_postrun_workers=1, + storage=storage, + session_registry=registry, + builders=StrategyRegistry(object), + evaluators=StrategyRegistry(object), + program_releaser=program_releaser, + ) + return manager, registry, storage + + +def _managed_session( + tmp_path: Path, + *, + session_id: str, + status: SessionStatus | None, +) -> ManagedSession: + request = SessionDispatchRequest( + session_id=session_id, + task_id="task-1", + instruction="test", + remaining_timeout_seconds=10, + agent=AgentSpec(harness="codex"), + ) + session_dir = tmp_path / session_id + artifacts_dir = session_dir / "artifacts" + artifacts_dir.mkdir(parents=True) + result = None + if status is not None: + error = None if status == SessionStatus.COMPLETED else "terminal error" + result = SessionResult( + session_id=session_id, + task_id=request.task_id, + status=status, + trajectory=Trajectory(status=status.value, error=error), + error=error, + ) + return ManagedSession( + request=request, + timer=StageTimer(), + session_dir=session_dir, + artifacts_dir=artifacts_dir, + final_result=result, + cancel_requested=status is None, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("status", "expected_status"), + [ + (SessionStatus.COMPLETED, SessionStatus.COMPLETED), + (SessionStatus.ERROR, SessionStatus.ERROR), + (SessionStatus.TIMEOUT, SessionStatus.TIMEOUT), + (None, SessionStatus.ERROR), + ], + ids=["completed", "error", "timeout", "cancelled"], +) +async def test_postrun_releases_program_for_every_terminal_state( + tmp_path: Path, + status: SessionStatus | None, + expected_status: SessionStatus, +) -> None: + releaser = AsyncMock(return_value=True) + manager, registry, storage = _manager(releaser) + managed = _managed_session(tmp_path, session_id=f"session-{expected_status}", status=status) + registry.register(managed.session_id, task_id=managed.request.task_id) + storage.ensure_session(managed.session_id, None, None, None) + + try: + await manager._handle_postrun(managed) + + info = registry.get(managed.session_id) + assert info is not None + assert info.status == expected_status + releaser.assert_awaited_once_with(managed.session_id) + assert storage.get_session_metadata(managed.session_id) is None + assert not managed.session_dir.exists() + finally: + await manager.close() + storage.close() + + +@pytest.mark.asyncio +async def test_releaser_failure_preserves_terminal_result_and_cleanup(tmp_path: Path) -> None: + releaser = AsyncMock(side_effect=RuntimeError("release failed")) + manager, registry, storage = _manager(releaser) + managed = _managed_session( + tmp_path, + session_id="session-release-failure", + status=SessionStatus.TIMEOUT, + ) + registry.register(managed.session_id, task_id=managed.request.task_id) + storage.ensure_session(managed.session_id, None, None, None) + + try: + await manager._handle_postrun(managed) + + info = registry.get(managed.session_id) + assert info is not None + assert info.result is not None + assert info.result.status == SessionStatus.TIMEOUT + releaser.assert_awaited_once_with(managed.session_id) + assert storage.get_session_metadata(managed.session_id) is None + assert not managed.session_dir.exists() + finally: + await manager.close() + storage.close()