commitcopilot/infer-001
0
1import asyncio2import html3import json4import os5import re6import time7import uuid8from contextlib import asynccontextmanager9from pathlib import Path10from typing import Annotated, Any11 12import httpx13from fastapi import Depends, FastAPI, HTTPException, Request, status14from fastapi.middleware.cors import CORSMiddleware15from fastapi.responses import JSONResponse16from pydantic import BaseModel, ConfigDict, Field17 18# ---------------------------------------------------------------------------19# Constants20# ---------------------------------------------------------------------------21 22CONFIG_PATH = Path(__file__).with_name("config.json")23 24BEARER_SCHEME = "bearer"25FINISH_REASON_STOP = "stop"26FINISH_REASON_TOOL_CALLS = "tool_calls"27TOOL_CALL_ID_PREFIX = "call_"28 29TOOL_CALL_BLOCK_RE = re.compile(r"<tool_call>\s*(.*?)\s*</tool_call>", re.DOTALL)30FUNCTION_BLOCK_RE = re.compile(r"<function=([^>\s]+)>\s*(.*?)\s*</function>", re.DOTALL)31PARAMETER_BLOCK_RE = re.compile(32 r"<parameter=([^>\s]+)>\s*(.*?)\s*</parameter>",33 re.DOTALL,34)35 36# ---------------------------------------------------------------------------37# Config38# ---------------------------------------------------------------------------39 40 41class LlamaConfig(BaseModel):42 server_bin: str = "/usr/local/bin/llama-server"43 n_ctx: int = 1638444 n_threads: int = 245 n_batch: int = 12846 n_parallel: int = 247 extra_args: list[str] = Field(default_factory=list)48 49 50class GenerationConfig(BaseModel):51 request_timeout_seconds: float = 300.052 default_max_tokens: int = 38453 default_temperature: float = 0.254 default_top_p: float = 0.9555 ready_wait_timeout_seconds: float = 180.056 ready_poll_interval_seconds: float = 1.057 58 59class AppConfig(BaseModel):60 model_id: str = "gemma-4-e2b"61 model_path: str = "/models/model.gguf"62 infer_api_key: str = ""63 llama: LlamaConfig = Field(default_factory=LlamaConfig)64 generation: GenerationConfig = Field(default_factory=GenerationConfig)65 66 67def load_config() -> AppConfig:68 if not CONFIG_PATH.exists():69 raise RuntimeError(f"Missing config file: {CONFIG_PATH}")70 with CONFIG_PATH.open("r", encoding="utf-8") as f:71 return AppConfig.model_validate(json.load(f))72 73 74CONFIG = load_config()75 76LLAMA_SERVER_HOST = os.getenv("LLAMA_SERVER_HOST", "127.0.0.1")77LLAMA_SERVER_PORT = int(os.getenv("LLAMA_SERVER_PORT", "8080"))78LLAMA_SERVER_BASE_URL = f"http://{LLAMA_SERVER_HOST}:{LLAMA_SERVER_PORT}"79 80# ---------------------------------------------------------------------------81# Process / client state82# ---------------------------------------------------------------------------83 84llama_process: asyncio.subprocess.Process | None = None85server_ready = False86client = httpx.AsyncClient(timeout=CONFIG.generation.request_timeout_seconds)87 88# ---------------------------------------------------------------------------89# Request / Response models90# ---------------------------------------------------------------------------91 92 93class ChatMessage(BaseModel):94 model_config = ConfigDict(extra="allow")95 96 role: str97 content: Any = None98 name: str | None = None99 tool_call_id: str | None = None100 tool_calls: list[dict[str, Any]] | None = None101 102 103class GenerateRequest(BaseModel):104 model_config = ConfigDict(extra="allow")105 106 request_id: str | None = None107 model: str = CONFIG.model_id108 messages: list[ChatMessage] = Field(default_factory=list)109 temperature: float | None = None110 top_p: float | None = None111 max_tokens: int | None = None112 stop: str | list[str] | None = None113 tools: list[dict[str, Any]] | None = None114 tool_choice: Any = None115 response_format: dict[str, Any] | None = None116 117 118class GenerateResponse(BaseModel):119 message: dict[str, Any]120 text: str | None = None121 model: str = CONFIG.model_id122 finish_reason: str = FINISH_REASON_STOP123 usage: dict[str, Any] | None = None124 request_id: str125 elapsed_seconds: float126 127 128# ---------------------------------------------------------------------------129# Auth130# ---------------------------------------------------------------------------131 132 133def require_infer_auth(request: Request) -> None:134 if not CONFIG.infer_api_key:135 return136 authorization = request.headers.get("authorization", "")137 scheme, _, token = authorization.partition(" ")138 if scheme.lower() != BEARER_SCHEME or token != CONFIG.infer_api_key:139 raise HTTPException(140 status_code=status.HTTP_401_UNAUTHORIZED,141 detail="Invalid or missing inference bearer token",142 )143 144 145InferAuth = Annotated[None, Depends(require_infer_auth)]146 147# ---------------------------------------------------------------------------148# llama-server lifecycle helpers149# ---------------------------------------------------------------------------150 151 152def build_llama_server_command() -> list[str]:153 path = Path(CONFIG.model_path)154 if not path.exists():155 raise RuntimeError(f"MODEL_PATH does not exist: {path}")156 157 llama = CONFIG.llama158 command = [159 llama.server_bin,160 "--model", str(path),161 "--host", LLAMA_SERVER_HOST,162 "--port", str(LLAMA_SERVER_PORT),163 "--ctx-size", str(llama.n_ctx),164 "--batch-size", str(llama.n_batch),165 "--parallel", str(llama.n_parallel),166 ]167 if llama.n_threads > 0:168 command.extend(["--threads", str(llama.n_threads)])169 command.extend(str(arg) for arg in llama.extra_args)170 return command171 172 173async def wait_for_server(timeout_seconds: float | None = None) -> bool:174 global server_ready175 176 if timeout_seconds is None:177 timeout_seconds = CONFIG.generation.ready_wait_timeout_seconds178 179 deadline = time.monotonic() + timeout_seconds180 while time.monotonic() < deadline:181 if llama_process and llama_process.returncode is not None:182 return False183 try:184 response = await client.get(f"{LLAMA_SERVER_BASE_URL}/health", timeout=5)185 if response.status_code < 500:186 server_ready = True187 return True188 except httpx.HTTPError:189 await asyncio.sleep(CONFIG.generation.ready_poll_interval_seconds)190 return False191 192 193# ---------------------------------------------------------------------------194# Tool-call parsing helpers195# ---------------------------------------------------------------------------196 197 198def normalize_tool_call(raw: dict[str, Any]) -> dict[str, Any] | None:199 function = raw.get("function")200 if isinstance(function, dict):201 name = function.get("name")202 arguments = function.get("arguments", "{}")203 else:204 name = raw.get("name")205 arguments = raw.get("arguments", "{}")206 207 if not isinstance(name, str) or not name:208 return None209 if not isinstance(arguments, str):210 arguments = json.dumps(arguments, ensure_ascii=False)211 212 return {213 "id": str(raw.get("id") or f"{TOOL_CALL_ID_PREFIX}{uuid.uuid4().hex}"),214 "type": "function",215 "function": {"name": name, "arguments": arguments},216 }217 218 219def strip_json_fence(text: str) -> str:220 stripped = text.strip()221 if stripped.startswith("```json") and stripped.endswith("```"):222 return stripped[len("```json"):-3].strip()223 if stripped.startswith("```") and stripped.endswith("```"):224 return stripped[3:-3].strip()225 return stripped226 227 228def coerce_json_tool_calls(parsed: Any) -> list[dict[str, Any]]:229 if isinstance(parsed, dict):230 if "name" in parsed or "function" in parsed:231 raw_calls: Any = [parsed]232 else:233 raw_calls = parsed.get("tool_calls") or parsed.get("tools") or parsed.get("calls")234 else:235 raw_calls = parsed236 237 if not isinstance(raw_calls, list):238 return []239 240 return [tc for raw in raw_calls if isinstance(raw, dict) for tc in [normalize_tool_call(raw)] if tc]241 242 243def parse_json_tool_calls(text: str) -> list[dict[str, Any]]:244 stripped = text.strip()245 if not stripped:246 return []247 try:248 parsed = json.loads(strip_json_fence(stripped))249 except json.JSONDecodeError:250 return []251 return coerce_json_tool_calls(parsed)252 253 254def parse_parameter_value(raw_value: str) -> Any:255 value = html.unescape(raw_value).strip()256 if not value:257 return ""258 try:259 return json.loads(value)260 except json.JSONDecodeError:261 return value262 263 264def parse_function_xml_tool_calls(text: str) -> list[dict[str, Any]]:265 tool_calls = []266 for function_match in FUNCTION_BLOCK_RE.finditer(text):267 name = html.unescape(function_match.group(1)).strip()268 body = function_match.group(2)269 if not name:270 continue271 arguments = {272 html.unescape(m.group(1)).strip(): parse_parameter_value(m.group(2))273 for m in PARAMETER_BLOCK_RE.finditer(body)274 if html.unescape(m.group(1)).strip()275 }276 tool_call = normalize_tool_call({"name": name, "arguments": arguments})277 if tool_call:278 tool_calls.append(tool_call)279 return tool_calls280 281 282def parse_tool_calls_from_text(text: str) -> list[dict[str, Any]]:283 tool_calls = parse_json_tool_calls(text)284 if tool_calls:285 return tool_calls286 287 xml_blocks = TOOL_CALL_BLOCK_RE.findall(text)288 if not xml_blocks:289 return []290 291 parsed: list[dict[str, Any]] = []292 for block in xml_blocks:293 parsed.extend(parse_json_tool_calls(block))294 parsed.extend(parse_function_xml_tool_calls(block))295 return parsed296 297 298# ---------------------------------------------------------------------------299# Response normalisation300# ---------------------------------------------------------------------------301 302 303def normalize_assistant_message(result: dict[str, Any]) -> tuple[dict[str, Any], str]:304 choices = result.get("choices") or []305 if not choices:306 return {"role": "assistant", "content": ""}, FINISH_REASON_STOP307 308 choice = choices[0]309 message = dict(choice.get("message") or {})310 if not message:311 message = {"role": "assistant", "content": choice.get("text") or ""}312 message["role"] = "assistant"313 message.setdefault("content", "")314 315 raw_tool_calls = message.get("tool_calls")316 if isinstance(raw_tool_calls, list):317 tool_calls = [318 tc319 for raw in raw_tool_calls320 if isinstance(raw, dict)321 for tc in [normalize_tool_call(raw)]322 if tc323 ]324 if tool_calls:325 message["tool_calls"] = tool_calls326 message["content"] = message.get("content") or None327 328 if not message.get("tool_calls") and isinstance(message.get("content"), str):329 tool_calls = parse_tool_calls_from_text(message["content"])330 if tool_calls:331 message["tool_calls"] = tool_calls332 message["content"] = None333 334 finish_reason = str(choice.get("finish_reason") or FINISH_REASON_STOP)335 if message.get("tool_calls") and finish_reason == FINISH_REASON_STOP:336 finish_reason = FINISH_REASON_TOOL_CALLS337 return message, finish_reason338 339 340# ---------------------------------------------------------------------------341# OpenAI payload builder342# ---------------------------------------------------------------------------343 344 345def build_openai_payload(payload: GenerateRequest) -> dict[str, Any]:346 gen = CONFIG.generation347 body: dict[str, Any] = {348 "model": payload.model or CONFIG.model_id,349 "messages": [m.model_dump(exclude_none=True) for m in payload.messages],350 "temperature": payload.temperature if payload.temperature is not None else gen.default_temperature,351 "top_p": payload.top_p if payload.top_p is not None else gen.default_top_p,352 "max_tokens": payload.max_tokens or gen.default_max_tokens,353 "stream": False,354 "cache_prompt": True,355 "reasoning_budget": 0,356 "chat_template_kwargs": {"enable_thinking": False},357 }358 for key, value in {359 "stop": payload.stop,360 "tools": payload.tools,361 "tool_choice": payload.tool_choice,362 "response_format": payload.response_format,363 }.items():364 if value is not None:365 body[key] = value366 return body367 368 369# ---------------------------------------------------------------------------370# Core generate function371# ---------------------------------------------------------------------------372 373 374async def generate(payload: GenerateRequest) -> GenerateResponse:375 if not payload.messages:376 raise HTTPException(377 status_code=status.HTTP_400_BAD_REQUEST,378 detail="messages must not be empty",379 )380 381 if not server_ready and not await wait_for_server(382 timeout_seconds=CONFIG.generation.ready_wait_timeout_seconds383 ):384 raise HTTPException(385 status_code=status.HTTP_503_SERVICE_UNAVAILABLE,386 detail="llama-server is not ready",387 )388 389 request_id = payload.request_id or uuid.uuid4().hex390 started = time.monotonic()391 try:392 response = await client.post(393 f"{LLAMA_SERVER_BASE_URL}/v1/chat/completions",394 json=build_openai_payload(payload),395 )396 response.raise_for_status()397 except httpx.HTTPStatusError as error:398 raise HTTPException(399 status_code=error.response.status_code,400 detail=error.response.text,401 ) from error402 except httpx.HTTPError as error:403 raise HTTPException(404 status_code=status.HTTP_503_SERVICE_UNAVAILABLE,405 detail=str(error),406 ) from error407 408 result = response.json()409 message, finish_reason = normalize_assistant_message(result)410 content = message.get("content")411 usage = result.get("usage") if isinstance(result.get("usage"), dict) else None412 return GenerateResponse(413 message=message,414 text=content if isinstance(content, str) else None,415 model=payload.model or CONFIG.model_id,416 finish_reason=finish_reason,417 usage=usage,418 request_id=request_id,419 elapsed_seconds=time.monotonic() - started,420 )421 422 423# ---------------------------------------------------------------------------424# App + lifespan425# ---------------------------------------------------------------------------426 427 428@asynccontextmanager429async def lifespan(_: FastAPI):430 global llama_process431 command = build_llama_server_command()432 llama_process = await asyncio.create_subprocess_exec(*command)433 asyncio.create_task(wait_for_server(timeout_seconds=600))434 yield435 if llama_process and llama_process.returncode is None:436 llama_process.terminate()437 try:438 await asyncio.wait_for(llama_process.wait(), timeout=10)439 except asyncio.TimeoutError:440 llama_process.kill()441 await llama_process.wait()442 await client.aclose()443 444 445app = FastAPI(446 title="Commit Copilot Cloud Inference Worker",447 version="3.0.0",448 lifespan=lifespan,449)450app.add_middleware(451 CORSMiddleware,452 allow_origins=["*"],453 allow_credentials=False,454 allow_methods=["*"],455 allow_headers=["*"],456)457 458# ---------------------------------------------------------------------------459# Routes460# ---------------------------------------------------------------------------461 462 463@app.get("/")464async def root() -> dict[str, Any]:465 return {"name": "Commit Copilot Cloud Inference Worker", "model": CONFIG.model_id}466 467 468@app.get("/health")469async def health() -> dict[str, Any]:470 ready = server_ready or await wait_for_server(timeout_seconds=1)471 process_running = llama_process is not None and llama_process.returncode is None472 return {473 "ok": process_running and ready,474 "model": CONFIG.model_id,475 "backend": "llama.cpp llama-server",476 "server_ready": ready,477 "server_running": process_running,478 }479 480 481@app.get("/ready")482async def ready() -> dict[str, Any]:483 if not server_ready and not await wait_for_server(timeout_seconds=1):484 raise HTTPException(485 status_code=status.HTTP_503_SERVICE_UNAVAILABLE,486 detail="llama-server is not ready",487 )488 return {"ready": True, "model": CONFIG.model_id}489 490 491@app.post("/generate")492async def generate_route(493 payload: GenerateRequest,494 _auth: InferAuth,495) -> GenerateResponse:496 return await generate(payload)497 498 499@app.post("/v1/chat/completions")500async def openai_compatible_chat(501 payload: GenerateRequest,502 _auth: InferAuth,503) -> JSONResponse:504 result = await generate(payload)505 finish_reason = result.finish_reason506 if result.message.get("tool_calls") and finish_reason == FINISH_REASON_STOP:507 finish_reason = FINISH_REASON_TOOL_CALLS508 return JSONResponse(509 {510 "id": f"chatcmpl-{uuid.uuid4().hex}",511 "object": "chat.completion",512 "created": int(time.time()),513 "model": result.model,514 "choices": [515 {516 "index": 0,517 "message": result.message,518 "finish_reason": finish_reason,519 }520 ],521 "usage": result.usage,522 }523 )524 