from __future__ import annotations import asyncio import json import socket import subprocess from collections.abc import AsyncIterator from pathlib import Path from typing import Any import httpx from ..config import ModelConfig, RuntimeConfig from ..types import BackendError, ChatChunk, ChatParams from .base import BaseBackend def _free_port() -> int: with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: sock.bind(("127.0.0.1", 0)) return int(sock.getsockname()[1]) class GGUFBackend(BaseBackend): """Run a GGUF model through an external CUDA-enabled llama-server.""" def __init__(self, model: ModelConfig, runtime: RuntimeConfig) -> None: super().__init__(model, runtime) self.port: int | None = None self.process: subprocess.Popen[str] | None = None self.log_path: Path | None = None self._log_handle: Any = None @property def endpoint(self) -> str: if self.port is None: raise BackendError("GGUF 后端尚未启动") return f"http://127.0.0.1:{self.port}" async def load(self) -> None: if self.loaded: return if not self.model.path: raise BackendError(f"GGUF 模型缺少 path: {self.model.id}") model_path = Path(self.model.path) if not model_path.exists(): raise BackendError(f"GGUF 文件不存在: {model_path}") server = Path(self.runtime.llama_server) if not server.exists(): raise BackendError( f"找不到 llama-server: {server}。请在 config.json 的 runtime.llama_server 中配置。" ) self.port = _free_port() runtime_dir = Path(self.runtime.runtime_dir) runtime_dir.mkdir(parents=True, exist_ok=True) self.log_path = runtime_dir / f"{self.model.id}.llama-server.log" self._log_handle = self.log_path.open("a", encoding="utf-8", errors="replace") options = self.model.options args = [ str(server), "--model", str(model_path), "--host", "127.0.0.1", "--port", str(self.port), "--alias", self.model.id, "--no-webui", "--jinja", "--cont-batching", "--cache-prompt", "--flash-attn", str(options.get("flash_attn", self.runtime.flash_attn)), "--gpu-layers", str(options.get("n_gpu_layers", self.runtime.n_gpu_layers)), "--ctx-size", str(options.get("ctx_size", self.runtime.ctx_size)), "--batch-size", str(options.get("batch_size", self.runtime.batch_size)), "--ubatch-size", str(options.get("ubatch_size", self.runtime.ubatch_size)), ] if self.runtime.fit_vram and options.get("fit_vram", True): args.extend(["--fit", "on", "--fit-ctx", str(options.get("fit_ctx", 4096))]) if options.get("mlock", False): args.append("--mlock") if options.get("no_mmap", False): args.append("--no-mmap") if options.get("mmproj"): args.extend(["--mmproj", str(options["mmproj"])]) if options.get("chat_template_file"): args.extend(["--chat-template-file", str(options["chat_template_file"])]) elif self.model.template: args.extend(["--chat-template", self.model.template]) creationflags = getattr(subprocess, "CREATE_NEW_PROCESS_GROUP", 0) self.process = subprocess.Popen( args, cwd=str(server.parent), stdout=self._log_handle, stderr=subprocess.STDOUT, text=True, creationflags=creationflags, ) try: await self._wait_ready() except Exception: await self.unload() raise self.loaded = True async def _wait_ready(self) -> None: deadline = asyncio.get_running_loop().time() + 90.0 last_error = "" async with httpx.AsyncClient(timeout=2.0) as client: while asyncio.get_running_loop().time() < deadline: if self.process and self.process.poll() is not None: tail = "" if self.log_path and self.log_path.exists(): tail = self.log_path.read_text(encoding="utf-8", errors="replace")[-2000:] if "check_tensor_dims" in tail or "wrong shape" in tail: tail = ( "检测到模型张量与当前 llama.cpp runner 不兼容。若 path 指向 Ollama 的 " "content-addressed blob,请改用 kind=ollama;否则换成与 runner 匹配的标准 GGUF 文件。\n" + tail ) raise BackendError(f"llama-server 启动失败。日志尾部:\n{tail}") try: response = await client.get(f"{self.endpoint}/health") if response.status_code == 200: return last_error = f"HTTP {response.status_code}" except httpx.HTTPError as exc: last_error = str(exc) await asyncio.sleep(0.25) raise BackendError(f"llama-server 在 90 秒内未就绪: {last_error}") async def unload(self) -> None: process, self.process = self.process, None if process is not None and process.poll() is None: process.terminate() try: await asyncio.to_thread(process.wait, 5) except subprocess.TimeoutExpired: process.kill() await asyncio.to_thread(process.wait) if self._log_handle is not None: self._log_handle.close() self._log_handle = None self.loaded = False self.port = None async def chat( self, messages: list[dict[str, Any]], params: ChatParams, stream: bool = False, ) -> AsyncIterator[ChatChunk]: if not self.loaded: await self.load() payload: dict[str, Any] = { "model": self.model.id, "messages": messages, "stream": stream, "max_tokens": params.max_tokens, "temperature": params.temperature, "top_p": params.top_p, "top_k": params.top_k, "min_p": params.min_p, "repeat_penalty": params.repeat_penalty, } if params.seed is not None: payload["seed"] = params.seed if params.stop: payload["stop"] = params.stop if params.enable_thinking is not None: payload["chat_template_kwargs"] = {"enable_thinking": params.enable_thinking} if params.response_format is not None: payload["response_format"] = ( {"type": params.response_format} if isinstance(params.response_format, str) else params.response_format ) if params.tools: payload["tools"] = params.tools if params.tool_choice is not None: payload["tool_choice"] = params.tool_choice payload.update(params.extra) timeout = httpx.Timeout(600.0, connect=10.0) async with httpx.AsyncClient(timeout=timeout) as client: if stream: async with client.stream("POST", f"{self.endpoint}/v1/chat/completions", json=payload) as response: if response.status_code >= 400: detail = (await response.aread()).decode("utf-8", errors="replace") raise BackendError(f"llama-server 请求失败 {response.status_code}: {detail}") async for line in response.aiter_lines(): if not line.startswith("data:"): continue raw = line[5:].strip() if raw == "[DONE]": yield ChatChunk(done=True, finish_reason="stop") return try: data = json.loads(raw) except json.JSONDecodeError: continue choice = (data.get("choices") or [{}])[0] delta = choice.get("delta") or {} text = delta.get("content") or "" finish = choice.get("finish_reason") if text or finish or len(delta) > 1: yield ChatChunk(text=text, finish_reason=finish, metadata=data, delta=delta) yield ChatChunk(done=True, finish_reason="stop") return response = await client.post(f"{self.endpoint}/v1/chat/completions", json=payload) if response.status_code >= 400: raise BackendError(f"llama-server 请求失败 {response.status_code}: {response.text}") data = response.json() choice = (data.get("choices") or [{}])[0] message = choice.get("message") or {} yield ChatChunk( text=message.get("content") or "", done=True, finish_reason=choice.get("finish_reason") or "stop", metadata=data, delta=message, ) def info(self) -> dict[str, Any]: return { "id": self.model.id, "kind": "gguf", "path": self.model.path, "loaded": self.loaded, "endpoint": self.endpoint if self.port else None, "pid": self.process.pid if self.process else None, "log": str(self.log_path) if self.log_path else None, }