Mercurial
diff mrjunejune/inference/litellm_proxy.py @ 260:1f9877b637e9
Add Copilot-powered cyberpunk JRPG chat
Integrate the production JRPG chat with Seobeo streaming, Deita persistence, and a Bazel-managed Copilot SDK and LiteLLM inference stack.
Co-authored-by: Copilot <[email protected]>
| author | MrJuneJune <mrjunejune@users.noreply.github.com> |
|---|---|
| date | Wed, 05 Aug 2026 09:19:41 -0700 |
| parents | |
| children |
line wrap: on
line diff
--- /dev/null Thu Jan 01 00:00:00 1970 +0000 +++ b/mrjunejune/inference/litellm_proxy.py Wed Aug 05 09:19:41 2026 -0700 @@ -0,0 +1,166 @@ +from __future__ import annotations + +import argparse +import asyncio +import json +import os +import pathlib +from collections.abc import AsyncIterator, Awaitable, Callable +from typing import Any + +import litellm +import uvicorn +import yaml +from fastapi import FastAPI, HTTPException, Request +from fastapi.responses import JSONResponse, StreamingResponse + +Completion = Callable[..., Awaitable[Any]] + + +def resolve_token_directory(explicit: str | None = None) -> pathlib.Path: + if explicit: + return pathlib.Path(explicit).expanduser().resolve() + configured = os.environ.get("GITHUB_COPILOT_TOKEN_DIR") + if configured: + return pathlib.Path(configured).expanduser().resolve() + state_root = os.environ.get("MRJUNEJUNE_INFERENCE_STATE") + if state_root: + return ( + pathlib.Path(state_root).expanduser().resolve() + / "litellm-copilot" + ) + xdg_state = os.environ.get("XDG_STATE_HOME") + root = ( + pathlib.Path(xdg_state).expanduser() + if xdg_state + else pathlib.Path.home() / ".local" / "state" + ) + return (root / "mrjunejune" / "inference" / "litellm-copilot").resolve() + + +def _serialize(value: Any) -> dict[str, Any]: + if hasattr(value, "model_dump"): + return value.model_dump(exclude_none=True) + if hasattr(value, "dict"): + return value.dict() + if isinstance(value, dict): + return value + raise TypeError("LiteLLM returned an unsupported response") + + +def create_app( + completion: Completion = litellm.acompletion, + *, + model_alias: str = "jrpg-copilot", + upstream_model: str = "github_copilot/gpt-4", +) -> FastAPI: + app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None) + + @app.get("/health/liveliness") + async def health() -> dict[str, str]: + return {"status": "ready"} + + @app.get("/health/readiness") + async def readiness() -> dict[str, str]: + from litellm.llms.github_copilot.authenticator import Authenticator + + try: + await asyncio.to_thread(Authenticator().get_api_key) + except Exception as error: + raise HTTPException( + status_code=503, + detail="GitHub Copilot authentication unavailable", + ) from error + return {"status": "ready"} + + @app.post("/v1/chat/completions") + async def chat(request: Request): + master_key = os.environ.get("LITELLM_MASTER_KEY", "") + authorization = request.headers.get("Authorization", "") + if not master_key or authorization != f"Bearer {master_key}": + raise HTTPException(status_code=401, detail="Invalid gateway key") + + payload = await request.json() + if not isinstance(payload, dict) or payload.get("model") != model_alias: + raise HTTPException(status_code=400, detail="Unknown model alias") + payload = dict(payload) + payload["model"] = upstream_model + configured_max = int(os.environ.get("LITELLM_MAX_OUTPUT_TOKENS", "1024")) + requested_max = payload.get("max_tokens") + if not isinstance(requested_max, int) or requested_max > configured_max: + payload["max_tokens"] = configured_max + + try: + response = await completion(**payload) + except Exception as error: + raise HTTPException( + status_code=502, + detail="GitHub Copilot provider request failed", + ) from error + + if payload.get("stream"): + async def events() -> AsyncIterator[str]: + async for chunk in response: + yield ( + "data: " + + json.dumps(_serialize(chunk), separators=(",", ":")) + + "\n\n" + ) + yield "data: [DONE]\n\n" + + return StreamingResponse(events(), media_type="text/event-stream") + return JSONResponse(_serialize(response)) + + return app + + +def _load_model(config_path: str) -> tuple[str, str]: + with open(config_path, encoding="utf-8") as config_file: + config = yaml.safe_load(config_file) + models = config.get("model_list", []) + if len(models) != 1: + raise ValueError("LiteLLM config must contain exactly one model") + model = models[0] + return model["model_name"], model["litellm_params"]["model"] + + +def main() -> None: + parser = argparse.ArgumentParser(description="Minimal LiteLLM gateway") + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, default=4000) + parser.add_argument("--config") + parser.add_argument("--authenticate", action="store_true") + parser.add_argument( + "--token-dir", + help=( + "Persistent GitHub Copilot token directory. Defaults to " + "GITHUB_COPILOT_TOKEN_DIR, MRJUNEJUNE_INFERENCE_STATE, or " + "~/.local/state/mrjunejune/inference/litellm-copilot." + ), + ) + args = parser.parse_args() + if args.authenticate: + from litellm.llms.github_copilot.authenticator import Authenticator + + token_directory = resolve_token_directory(args.token_dir) + token_directory.mkdir(parents=True, exist_ok=True, mode=0o700) + token_directory.chmod(0o700) + os.environ["GITHUB_COPILOT_TOKEN_DIR"] = str(token_directory) + authenticator = Authenticator() + authenticator.get_access_token() + authenticator.get_api_key() + print(f"GitHub Copilot authentication stored in {token_directory}") + return + if not args.config: + parser.error("--config is required unless --authenticate is used") + model_alias, upstream_model = _load_model(args.config) + uvicorn.run( + create_app(model_alias=model_alias, upstream_model=upstream_model), + host=args.host, + port=args.port, + access_log=False, + ) + + +if __name__ == "__main__": + main()