370 lines
15 KiB
Python
370 lines
15 KiB
Python
from __future__ import annotations
|
||
|
||
import asyncio
|
||
import re
|
||
import time
|
||
from collections.abc import AsyncIterator
|
||
from dataclasses import asdict, dataclass
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
from .backends import GGUFBackend, ONNXBackend, OllamaBackend, RemoteBackend, TransformersBackend
|
||
from .config import AppConfig, ModelConfig, infer_kind, save_config
|
||
from .types import BackendError, ChatChunk, ChatParams
|
||
|
||
|
||
@dataclass
|
||
class ModelRuntimeState:
|
||
loaded_at: float = 0.0
|
||
last_used: float = 0.0
|
||
keep_alive_seconds: float = 300.0
|
||
last_load_duration_seconds: float = 0.0
|
||
active_requests: int = 0
|
||
total_requests: int = 0
|
||
last_duration_seconds: float = 0.0
|
||
idle_task: asyncio.Task[None] | None = None
|
||
|
||
|
||
def parse_keep_alive(value: str | int | float | None, default: float) -> float:
|
||
"""Parse Ollama-style keep_alive values such as 5m, 30s, 1h, 0, and -1."""
|
||
if value is None:
|
||
return default
|
||
if isinstance(value, (int, float)):
|
||
return float("inf") if value < 0 else max(0.0, float(value))
|
||
raw = str(value).strip().lower()
|
||
if raw in {"", "default"}:
|
||
return default
|
||
if raw in {"-1", "forever", "inf", "infinite"}:
|
||
return float("inf")
|
||
match = re.fullmatch(r"(-?\d+(?:\.\d+)?)\s*([smhd]?)", raw)
|
||
if not match:
|
||
return default
|
||
number = float(match.group(1))
|
||
if number < 0:
|
||
return float("inf")
|
||
multiplier = {"": 1.0, "s": 1.0, "m": 60.0, "h": 3600.0, "d": 86400.0}[match.group(2)]
|
||
return number * multiplier
|
||
|
||
|
||
class ModelManager:
|
||
def __init__(self, config: AppConfig, config_path: str | Path | None = None) -> None:
|
||
self.config = config
|
||
self.config_path = Path(config_path) if config_path else None
|
||
self.backends: dict[str, Any] = {}
|
||
self.states: dict[str, ModelRuntimeState] = {}
|
||
self.active_model_id: str | None = None
|
||
self._load_lock = asyncio.Lock()
|
||
|
||
def _state(self, model_id: str) -> ModelRuntimeState:
|
||
if model_id not in self.states:
|
||
self.states[model_id] = ModelRuntimeState(
|
||
keep_alive_seconds=self.config.runtime.keep_alive_seconds
|
||
)
|
||
return self.states[model_id]
|
||
|
||
@staticmethod
|
||
def _model_size(path: str | None) -> int:
|
||
if not path:
|
||
return 0
|
||
target = Path(path)
|
||
if target.is_file():
|
||
try:
|
||
return target.stat().st_size
|
||
except OSError:
|
||
return 0
|
||
if not target.is_dir():
|
||
return 0
|
||
total = 0
|
||
try:
|
||
for item in target.rglob("*"):
|
||
if item.is_file() and item.suffix.lower() in {".gguf", ".safetensors", ".bin", ".pt", ".onnx"}:
|
||
total += item.stat().st_size
|
||
except OSError:
|
||
return total
|
||
return total
|
||
|
||
def _model_record(self, spec: ModelConfig) -> dict[str, Any]:
|
||
backend = self.backends.get(spec.id)
|
||
state = self._state(spec.id)
|
||
size = self._model_size(spec.path)
|
||
modified = 0.0
|
||
if spec.path:
|
||
try:
|
||
modified = Path(spec.path).stat().st_mtime
|
||
except OSError:
|
||
pass
|
||
record: dict[str, Any] = {
|
||
"id": spec.id,
|
||
"name": spec.id,
|
||
"model": spec.id,
|
||
"object": "model",
|
||
"owned_by": "localpilot" if infer_kind(spec) not in {"cloud", "ollama"} else infer_kind(spec),
|
||
"kind": infer_kind(spec),
|
||
"path": spec.path,
|
||
"size": size,
|
||
"digest": f"local-{int(modified)}-{size}" if size else None,
|
||
"modified_at": modified,
|
||
"loaded": bool(backend and backend.loaded),
|
||
"active": spec.id == self.active_model_id,
|
||
"active_requests": state.active_requests,
|
||
"total_requests": state.total_requests,
|
||
"keep_alive_seconds": state.keep_alive_seconds,
|
||
"load_duration_seconds": state.last_load_duration_seconds,
|
||
"last_request_duration_seconds": state.last_duration_seconds,
|
||
"description": spec.description,
|
||
"capabilities": ["completion"],
|
||
}
|
||
if infer_kind(spec) in {"transformers", "ollama"}:
|
||
record["capabilities"].append("embedding")
|
||
if spec.options.get("mmproj"):
|
||
record["capabilities"].append("vision")
|
||
if backend:
|
||
record.update({key: value for key, value in backend.info().items() if key not in record})
|
||
return record
|
||
|
||
def list_models(self) -> list[dict[str, Any]]:
|
||
return [self._model_record(spec) for spec in self.config.models if spec.enabled]
|
||
|
||
def running_models(self) -> list[dict[str, Any]]:
|
||
return [
|
||
self._model_record(spec)
|
||
for spec in self.config.models
|
||
if (backend := self.backends.get(spec.id)) and backend.loaded
|
||
]
|
||
|
||
def show_model(self, model_id: str) -> dict[str, Any]:
|
||
spec = self.config.model_by_id(model_id)
|
||
return {
|
||
"model": model_id,
|
||
"details": {
|
||
"format": infer_kind(spec),
|
||
"family": spec.options.get("family", "unknown"),
|
||
"parameter_size": spec.options.get("parameter_size"),
|
||
"quantization_level": spec.options.get("quantization_level"),
|
||
},
|
||
"template": spec.template,
|
||
"system": spec.system,
|
||
"messages": spec.messages,
|
||
"parameters": spec.parameters,
|
||
"options": spec.options,
|
||
"runtime": self._model_record(spec),
|
||
}
|
||
|
||
def _make_backend(self, spec: ModelConfig) -> Any:
|
||
kind = infer_kind(spec)
|
||
if kind == "gguf":
|
||
return GGUFBackend(spec, self.config.runtime)
|
||
if kind == "transformers":
|
||
return TransformersBackend(spec, self.config.runtime)
|
||
if kind == "onnx":
|
||
return ONNXBackend(spec, self.config.runtime)
|
||
if kind == "ollama":
|
||
return OllamaBackend(spec, self.config.runtime)
|
||
if kind == "cloud":
|
||
if not spec.cloud_profile:
|
||
raise BackendError(f"云端模型 {spec.id} 缺少 cloud_profile")
|
||
return RemoteBackend(spec, self.config.cloud_by_id(spec.cloud_profile), self.config.runtime)
|
||
raise BackendError(f"无法识别模型格式: {spec.id} ({spec.path})")
|
||
|
||
async def _evict_if_needed(self, requested_id: str) -> None:
|
||
spec = self.config.model_by_id(requested_id)
|
||
if infer_kind(spec) in {"cloud", "ollama"}:
|
||
return
|
||
limit = max(1, self.config.runtime.max_loaded_models)
|
||
local_loaded = [
|
||
(model_id, self._state(model_id))
|
||
for model_id, backend in self.backends.items()
|
||
if backend.loaded and infer_kind(self.config.model_by_id(model_id)) not in {"cloud", "ollama"}
|
||
]
|
||
while len(local_loaded) >= limit:
|
||
candidates = [(mid, state) for mid, state in local_loaded if mid != requested_id and state.active_requests == 0]
|
||
if not candidates:
|
||
raise BackendError("模型加载槽位都在使用中,请等待当前请求结束或提高 runtime.max_loaded_models")
|
||
oldest_id = min(candidates, key=lambda item: item[1].last_used or item[1].loaded_at)[0]
|
||
await self.unload(oldest_id)
|
||
local_loaded = [item for item in local_loaded if item[0] != oldest_id]
|
||
|
||
async def load(self, model_id: str) -> Any:
|
||
async with self._load_lock:
|
||
if model_id in self.backends and self.backends[model_id].loaded:
|
||
self.active_model_id = model_id
|
||
self._state(model_id).last_used = time.time()
|
||
return self.backends[model_id]
|
||
spec = self.config.model_by_id(model_id)
|
||
await self._evict_if_needed(model_id)
|
||
backend = self.backends.get(model_id) or self._make_backend(spec)
|
||
started = time.perf_counter()
|
||
await backend.load()
|
||
state = self._state(model_id)
|
||
state.loaded_at = time.time()
|
||
state.last_used = state.loaded_at
|
||
state.last_load_duration_seconds = time.perf_counter() - started
|
||
self.backends[model_id] = backend
|
||
self.active_model_id = model_id
|
||
return backend
|
||
|
||
async def unload(self, model_id: str) -> None:
|
||
backend = self.backends.get(model_id)
|
||
state = self.states.get(model_id)
|
||
if state and state.idle_task and state.idle_task is not asyncio.current_task():
|
||
state.idle_task.cancel()
|
||
state.idle_task = None
|
||
if backend:
|
||
await backend.unload()
|
||
if self.active_model_id == model_id:
|
||
self.active_model_id = None
|
||
|
||
def _schedule_idle_unload(self, model_id: str, seconds: float) -> None:
|
||
state = self._state(model_id)
|
||
if state.idle_task:
|
||
state.idle_task.cancel()
|
||
state.keep_alive_seconds = seconds
|
||
state.last_used = time.time()
|
||
if seconds == float("inf"):
|
||
state.idle_task = None
|
||
return
|
||
|
||
async def expire() -> None:
|
||
try:
|
||
await asyncio.sleep(seconds)
|
||
current = self.states.get(model_id)
|
||
if not current or current.active_requests:
|
||
return
|
||
if time.time() - current.last_used >= seconds:
|
||
await self.unload(model_id)
|
||
except asyncio.CancelledError:
|
||
return
|
||
|
||
state.idle_task = asyncio.create_task(expire())
|
||
|
||
def _effective_params(self, spec: ModelConfig, params: ChatParams) -> ChatParams:
|
||
values = asdict(params)
|
||
defaults = ChatParams()
|
||
for key, value in spec.parameters.items():
|
||
if key in values and getattr(params, key) == getattr(defaults, key):
|
||
values[key] = value
|
||
return ChatParams(**values)
|
||
|
||
def _prepare_messages(
|
||
self,
|
||
spec: ModelConfig,
|
||
messages: list[dict[str, Any]],
|
||
params: ChatParams,
|
||
) -> list[dict[str, Any]]:
|
||
prepared = [dict(message) for message in spec.messages] + [dict(message) for message in messages]
|
||
if spec.system and not any(message.get("role") == "system" for message in prepared):
|
||
prepared.insert(0, {"role": "system", "content": spec.system})
|
||
if params.response_format and infer_kind(spec) in {"transformers", "onnx"}:
|
||
if isinstance(params.response_format, dict):
|
||
hint = "只输出符合以下 JSON Schema 的有效 JSON,不要输出 Markdown:" + str(params.response_format)
|
||
else:
|
||
hint = "只输出有效 JSON,不要输出 Markdown 或额外解释。"
|
||
prepared.insert(0, {"role": "system", "content": hint})
|
||
return prepared
|
||
|
||
async def chat(
|
||
self,
|
||
model_id: str | None,
|
||
messages: list[dict[str, Any]],
|
||
params: ChatParams,
|
||
stream: bool = False,
|
||
keep_alive: str | int | float | None = None,
|
||
) -> AsyncIterator[ChatChunk]:
|
||
chosen = model_id or self.active_model_id or self.config.runtime.default_model
|
||
if not chosen:
|
||
raise BackendError("没有可用模型,请在 config.json 设置 runtime.default_model 或请求中传 model")
|
||
spec = self.config.model_by_id(chosen)
|
||
effective_params = self._effective_params(spec, params)
|
||
effective_params.keep_alive = keep_alive
|
||
prepared_messages = self._prepare_messages(spec, messages, effective_params)
|
||
backend = await self.load(chosen)
|
||
state = self._state(chosen)
|
||
state.active_requests += 1
|
||
state.total_requests += 1
|
||
started = time.perf_counter()
|
||
try:
|
||
async for chunk in backend.chat(prepared_messages, effective_params, stream=stream):
|
||
yield chunk
|
||
finally:
|
||
state.active_requests = max(0, state.active_requests - 1)
|
||
state.last_duration_seconds = time.perf_counter() - started
|
||
self._schedule_idle_unload(chosen, parse_keep_alive(keep_alive, self.config.runtime.keep_alive_seconds))
|
||
|
||
async def embed(self, model_id: str | None, inputs: list[str], keep_alive: str | int | float | None = None) -> list[list[float]]:
|
||
chosen = model_id or self.active_model_id or self.config.runtime.default_model
|
||
if not chosen:
|
||
raise BackendError("没有可用 embedding 模型")
|
||
backend = await self.load(chosen)
|
||
state = self._state(chosen)
|
||
state.active_requests += 1
|
||
state.total_requests += 1
|
||
started = time.perf_counter()
|
||
try:
|
||
return await backend.embed(inputs)
|
||
finally:
|
||
state.active_requests = max(0, state.active_requests - 1)
|
||
state.last_duration_seconds = time.perf_counter() - started
|
||
self._schedule_idle_unload(chosen, parse_keep_alive(keep_alive, self.config.runtime.keep_alive_seconds))
|
||
|
||
def register_model(self, spec: ModelConfig) -> ModelConfig:
|
||
if any(item.id == spec.id for item in self.config.models):
|
||
raise BackendError(f"模型 ID 已存在: {spec.id}")
|
||
self.config.models.append(spec)
|
||
if self.config_path:
|
||
save_config(self.config_path, self.config)
|
||
return spec
|
||
|
||
def copy_model(self, source_id: str, destination_id: str) -> ModelConfig:
|
||
source = self.config.model_by_id(source_id)
|
||
copied = source.model_copy(update={"id": destination_id})
|
||
return self.register_model(copied)
|
||
|
||
async def pull_huggingface(
|
||
self,
|
||
model_id: str,
|
||
repo_id: str,
|
||
destination: str | Path | None = None,
|
||
revision: str | None = None,
|
||
allow_patterns: list[str] | None = None,
|
||
) -> ModelConfig:
|
||
try:
|
||
from huggingface_hub import snapshot_download
|
||
except ImportError as exc:
|
||
raise BackendError("当前 Conda LLM 环境缺少 huggingface_hub") from exc
|
||
if destination is None:
|
||
safe_repo = re.sub(r"[^A-Za-z0-9._-]+", "_", repo_id)
|
||
destination = Path(self.config.runtime.cache_dir) / "models" / safe_repo
|
||
target = Path(destination)
|
||
target.parent.mkdir(parents=True, exist_ok=True)
|
||
try:
|
||
downloaded = await asyncio.to_thread(
|
||
snapshot_download,
|
||
repo_id=repo_id,
|
||
revision=revision,
|
||
local_dir=str(target),
|
||
allow_patterns=allow_patterns,
|
||
)
|
||
except Exception as exc:
|
||
raise BackendError(f"Hugging Face 下载失败: {exc}") from exc
|
||
spec = ModelConfig(id=model_id, kind="auto", path=downloaded, description=f"Hugging Face: {repo_id}")
|
||
return self.register_model(spec)
|
||
|
||
async def remove_model(self, model_id: str) -> None:
|
||
await self.unload(model_id)
|
||
before = len(self.config.models)
|
||
self.config.models = [item for item in self.config.models if item.id != model_id]
|
||
if len(self.config.models) == before:
|
||
raise BackendError(f"模型不存在: {model_id}")
|
||
if self.config_path:
|
||
save_config(self.config_path, self.config)
|
||
|
||
async def shutdown(self) -> None:
|
||
for state in self.states.values():
|
||
if state.idle_task:
|
||
state.idle_task.cancel()
|
||
for backend in list(self.backends.values()):
|
||
if backend.loaded:
|
||
await backend.unload()
|
||
self.backends.clear()
|
||
self.active_model_id = None
|