# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 import argparse import json import httpx import uvicorn from fastapi import FastAPI, Request from fastapi.responses import JSONResponse from starlette.responses import Response from slime.utils.misc import load_function def run_router(args): """ Run the Slime router with the specified configuration. """ # Initialize the router with tokenizer and lazy worker initialization slime_router = SlimeRouter(args, verbose=False) # Start the server uvicorn.run(slime_router.app, host=args.sglang_router_ip, port=args.sglang_router_port, log_level="info") class SlimeRouter: def __init__(self, args, verbose=False): """Initialize the slime-router with SGLang router address""" self.args = args self.verbose = verbose self.app = FastAPI() # Worker information self.worker_urls: dict[str, int] = {} self.max_weight_version = None max_connections = getattr(args, "slime_router_max_connections", None) if max_connections is None: max_connections = ( args.sglang_server_concurrency * args.rollout_num_gpus // args.rollout_num_gpus_per_engine ) timeout = getattr(args, "slime_router_timeout", None) self.client = httpx.AsyncClient( limits=httpx.Limits(max_connections=max_connections), timeout=httpx.Timeout(timeout), ) self._setup_routes() for middleware_path in args.slime_router_middleware_paths or []: if self.verbose: print(f"[slime-router] Loading middleware from: {middleware_path}") middleware = load_function(middleware_path) self.app.add_middleware(middleware, router=self) def _setup_routes(self): """Setup all the HTTP routes""" # sglang-router api self.app.post("/add_worker")(self.add_worker) self.app.get("/list_workers")(self.list_workers) self.app.post("/retrieve_from_text")(self.retrieve_from_text) # Catch-all route for proxying to SGLang - must be registered LAST self.app.api_route("/{path:path}", methods=["GET", "POST", "PUT", "DELETE"])(self.proxy) async def health_check(self, request: Request): # TODO: do health check in background pass async def proxy(self, request: Request, path: str): """Proxy all other requests to the SGLang router""" # Forward all other paths to SGLang router worker_url = self._use_url() url = f"{worker_url}/{path}" # Get request body and headers body = await request.body() headers = dict(request.headers) try: response = await self.client.request(request.method, url, content=body, headers=headers) # Eagerly read content so we can return JSON (not streaming) content = await response.aread() content_type = response.headers.get("content-type", "") try: # Prefer parsing JSON if possible data = json.loads(content) return JSONResponse( content=data, status_code=response.status_code, headers=dict(response.headers), ) except Exception: # Fall back to raw body with original content type return Response( content=content, status_code=response.status_code, headers=dict(response.headers), media_type=content_type or None, ) finally: self._finish_url(worker_url) async def add_worker(self, request: Request): """Add a new worker to the router. Supports providing the URL via query string or JSON body. Examples: - POST /add_worker?url=http://127.0.0.1:10090 - POST /add_worker with body {"url": "http://127.0.0.1:10090"} """ # 1) Prefer query param worker_url = request.query_params.get("url") or request.query_params.get("worker_url") # 2) Fallback to JSON body if not worker_url: body = await request.body() payload = json.loads(body) if body else {} worker_url = payload.get("url") or payload.get("worker_url") if not worker_url: return JSONResponse( status_code=400, content={"error": "worker_url is required (use query ?url=... or JSON body)"} ) # Add if new, keep a simple request count per worker if worker_url not in self.worker_urls: self.worker_urls[worker_url] = 0 if self.verbose: print(f"[slime-router] Added new worker: {worker_url}") return {"status": "success", "worker_urls": self.worker_urls} async def list_workers(self, request: Request): """List all registered workers""" return {"urls": list(self.worker_urls.keys())} async def retrieve_from_text(self, request: Request): """Get token information from text input""" body = await request.body() payload = json.loads(body) if body else {} text = payload.get("text", "") # Use radix tree's retrieve_from_text method (no need to fetch weight version here) token_ids, logp, loss_mask = self.radix_tree.retrieve_from_text(text, return_logprob=True) # Handle the result based on whether logp was requested result = { "tokens": token_ids, # token IDs "response": text, # The input text "loss_mask": loss_mask, # Loss mask for the tokens "token_length": len(token_ids), "loss_mask_length": len(loss_mask), "rollout_logp": logp, } return result def _use_url(self): """Select a worker URL using round-robin strategy""" assert len(self.worker_urls) > 0, "No workers available" # get the url with mininal count url = min(self.worker_urls, key=self.worker_urls.get) self.worker_urls[url] += 1 return url def _finish_url(self, url): """Mark the request to the given URL as finished""" assert url in self.worker_urls, f"URL {url} not recognized" self.worker_urls[url] -= 1 assert self.worker_urls[url] >= 0, f"URL {url} count went negative" if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--host", type=str, default="0.0.0.0") parser.add_argument("--port", type=int, default=30000) parser.add_argument("--sglang-host", type=str, required=True) parser.add_argument("--sglang-port", type=int, required=True) parser.add_argument("--tokenizer-name", type=str, help="Name of the tokenizer to use for tokenization") parser.add_argument("--verbose", action="store_true", help="Enable verbose output") args = parser.parse_args() # Run the router run_router(args)