Team Ai
Apppublic

commitcopilot/infer-004

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
app.py524 linesDownload Raw Back to root
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