Initial commit: LocalPilot:本地模型运行时,复用 llama-server 并提供 Ollama 兼容 provider
This commit is contained in:
@@ -0,0 +1,127 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from ..config import ModelConfig, RuntimeConfig
|
||||
from ..types import BackendError, ChatChunk, ChatParams
|
||||
from .base import BaseBackend
|
||||
|
||||
|
||||
class OllamaBackend(BaseBackend):
|
||||
"""Compatibility provider for models already managed by a local Ollama daemon."""
|
||||
|
||||
def __init__(self, model: ModelConfig, runtime: RuntimeConfig) -> None:
|
||||
super().__init__(model, runtime)
|
||||
self.base_url = runtime.ollama_url.rstrip("/")
|
||||
self.remote_model = model.remote_model or model.id
|
||||
|
||||
async def load(self) -> None:
|
||||
if self.loaded:
|
||||
return
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=5.0) as client:
|
||||
response = await client.get(f"{self.base_url}/api/version")
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPError as exc:
|
||||
raise BackendError(f"无法连接本地 Ollama: {self.base_url}: {exc}") from exc
|
||||
self.loaded = True
|
||||
|
||||
async def unload(self) -> None:
|
||||
self.loaded = False
|
||||
|
||||
async def chat(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
params: ChatParams,
|
||||
stream: bool = False,
|
||||
) -> AsyncIterator[ChatChunk]:
|
||||
if not self.loaded:
|
||||
await self.load()
|
||||
options: dict[str, Any] = {
|
||||
"num_predict": 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:
|
||||
options["seed"] = params.seed
|
||||
if params.stop:
|
||||
options["stop"] = params.stop
|
||||
options.update(params.extra)
|
||||
payload: dict[str, Any] = {
|
||||
"model": self.remote_model,
|
||||
"messages": messages,
|
||||
"stream": stream,
|
||||
"options": options,
|
||||
"keep_alive": params.keep_alive if params.keep_alive is not None else "5m",
|
||||
}
|
||||
if params.response_format is not None:
|
||||
payload["format"] = params.response_format
|
||||
if params.enable_thinking is not None:
|
||||
payload["think"] = params.enable_thinking
|
||||
if params.tools:
|
||||
payload["tools"] = params.tools
|
||||
async with httpx.AsyncClient(timeout=httpx.Timeout(600.0, connect=10.0)) as client:
|
||||
if stream:
|
||||
async with client.stream("POST", f"{self.base_url}/api/chat", json=payload) as response:
|
||||
if response.status_code >= 400:
|
||||
detail = (await response.aread()).decode("utf-8", errors="replace")
|
||||
raise BackendError(f"Ollama 请求失败 {response.status_code}: {detail}")
|
||||
async for line in response.aiter_lines():
|
||||
if not line.strip():
|
||||
continue
|
||||
try:
|
||||
data = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
message = data.get("message") or {}
|
||||
text = message.get("content") or ""
|
||||
if data.get("done"):
|
||||
yield ChatChunk(done=True, finish_reason=data.get("done_reason") or "stop", metadata=data)
|
||||
return
|
||||
if text or len(message) > 1:
|
||||
yield ChatChunk(text=text, delta=message, metadata=data)
|
||||
yield ChatChunk(done=True, finish_reason="stop")
|
||||
return
|
||||
|
||||
payload["stream"] = False
|
||||
response = await client.post(f"{self.base_url}/api/chat", json=payload)
|
||||
if response.status_code >= 400:
|
||||
raise BackendError(f"Ollama 请求失败 {response.status_code}: {response.text}")
|
||||
data = response.json()
|
||||
message = data.get("message") or {}
|
||||
yield ChatChunk(
|
||||
text=message.get("content") or "",
|
||||
delta=message,
|
||||
done=True,
|
||||
finish_reason=data.get("done_reason") or "stop",
|
||||
metadata=data,
|
||||
)
|
||||
|
||||
async def embed(self, inputs: list[str]) -> list[list[float]]:
|
||||
if not self.loaded:
|
||||
await self.load()
|
||||
async with httpx.AsyncClient(timeout=120.0) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/embed",
|
||||
json={"model": self.remote_model, "input": inputs},
|
||||
)
|
||||
if response.status_code >= 400:
|
||||
raise BackendError(f"Ollama embeddings 失败 {response.status_code}: {response.text}")
|
||||
data = response.json()
|
||||
return data.get("embeddings", [])
|
||||
|
||||
def info(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.model.id,
|
||||
"kind": "ollama",
|
||||
"remote_model": self.remote_model,
|
||||
"loaded": self.loaded,
|
||||
"endpoint": self.base_url,
|
||||
}
|
||||
Reference in New Issue
Block a user