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
46 changes: 46 additions & 0 deletions src/polar/config/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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` |
Expand Down Expand Up @@ -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 `<gateway-node-id>:<session-id>` 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
Expand Down
1 change: 1 addition & 0 deletions src/polar/config/topology.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
27 changes: 24 additions & 3 deletions src/polar/gateway/node.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
from contextlib import suppress
from pathlib import Path
from tempfile import mkdtemp
from typing import Awaitable, Callable

import httpx

Expand Down Expand Up @@ -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("/")
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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
Expand Down
183 changes: 162 additions & 21 deletions src/polar/gateway/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Loading