feat: configure and select Ollama endpoints
This commit is contained in:
+163
-12
@@ -268,6 +268,104 @@ def _home() -> Path:
|
||||
return path
|
||||
|
||||
|
||||
CONNECTIONS_FILE = "connections.json"
|
||||
|
||||
|
||||
def _valid_ollama_url(value: str) -> str:
|
||||
value = str(value or "").strip().rstrip("/")
|
||||
parsed = urlparse(value)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.hostname or parsed.username or parsed.password:
|
||||
raise HTTPException(400, "Ollama URL must be an http(s) URL without credentials")
|
||||
if parsed.path not in {"", "/"} or parsed.query or parsed.fragment:
|
||||
raise HTTPException(400, "Ollama URL must be a base URL without a path, query, or fragment")
|
||||
try:
|
||||
if parsed.port is not None and not 1 <= parsed.port <= 65535:
|
||||
raise ValueError
|
||||
except ValueError as exc:
|
||||
raise HTTPException(400, "Ollama URL has an invalid port") from exc
|
||||
return value
|
||||
|
||||
|
||||
def _read_connections() -> dict[str, str]:
|
||||
try:
|
||||
value = json.loads((_home() / CONNECTIONS_FILE).read_text(encoding="utf-8"))
|
||||
except (OSError, ValueError):
|
||||
value = {}
|
||||
return {key: str(value.get(key) or "").strip().rstrip("/") for key in ("active_url", "local_url", "remote_url")}
|
||||
|
||||
|
||||
def _write_connections(value: dict[str, str]) -> None:
|
||||
path = _home() / CONNECTIONS_FILE
|
||||
path.write_text(json.dumps(value, indent=2) + "\n", encoding="utf-8")
|
||||
try:
|
||||
path.chmod(0o600)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _apply_saved_connection() -> None:
|
||||
global LOCAL_OLLAMA
|
||||
if os.environ.get("OLLAMA_HOST", "").strip():
|
||||
return
|
||||
saved = _read_connections().get("active_url")
|
||||
if saved:
|
||||
LOCAL_OLLAMA = saved
|
||||
|
||||
|
||||
def _target_endpoint(target: str = "local") -> str:
|
||||
target = str(target or "local").strip().lower()
|
||||
if target not in {"local", "remote"}:
|
||||
raise HTTPException(400, "Target must be local or remote")
|
||||
saved = _read_connections()
|
||||
if target == "remote":
|
||||
endpoint = saved.get("remote_url")
|
||||
if not endpoint:
|
||||
raise HTTPException(400, "No remote Ollama URL is configured")
|
||||
return _valid_ollama_url(endpoint)
|
||||
_apply_saved_connection()
|
||||
return _valid_ollama_url(saved.get("local_url") or LOCAL_OLLAMA)
|
||||
|
||||
|
||||
def _probe_endpoint(url: str, timeout: int = 5) -> dict[str, Any]:
|
||||
url = _valid_ollama_url(url)
|
||||
version_payload = _json_request(url + "/api/version", timeout=timeout)
|
||||
tags_payload = _json_request(url + "/api/tags", timeout=timeout)
|
||||
models = tags_payload.get("models", [])
|
||||
return {"available": True, "url": url, "version": str(version_payload.get("version") or "unknown"), "models": len(models) if isinstance(models, list) else 0}
|
||||
|
||||
|
||||
def _connection_snapshot() -> list[dict[str, Any]]:
|
||||
saved = _read_connections()
|
||||
configured_local = os.environ.get("OLLAMA_HOST", "").strip() or saved.get("local_url")
|
||||
candidates: list[tuple[str, str, str]] = []
|
||||
if configured_local:
|
||||
candidates.append(("local", "Configured local endpoint", configured_local))
|
||||
elif _running_in_container():
|
||||
candidates.extend((("local", "Docker Ollama service", "http://ollama:11434"), ("local", "Docker host Ollama", "http://host.docker.internal:11434")))
|
||||
else:
|
||||
candidates.append(("local", "Physical host Ollama", "http://localhost:11434"))
|
||||
if saved.get("remote_url"):
|
||||
candidates.append(("remote", "Configured remote endpoint", saved["remote_url"]))
|
||||
results: list[dict[str, Any]] = []
|
||||
seen: set[tuple[str, str]] = set()
|
||||
for kind, label, url in candidates:
|
||||
try:
|
||||
normalized = _valid_ollama_url(url)
|
||||
except HTTPException:
|
||||
continue
|
||||
key = (kind, normalized)
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
try:
|
||||
result = _probe_endpoint(normalized, timeout=3)
|
||||
result.update({"kind": kind, "label": label})
|
||||
except Exception as exc:
|
||||
result = {"available": False, "kind": kind, "label": label, "url": normalized, "error": str(exc)[:240]}
|
||||
results.append(result)
|
||||
return results
|
||||
|
||||
|
||||
def _json_request(url: str, method: str = "GET", payload: Any = None, timeout: int = 30) -> dict[str, Any]:
|
||||
data = None if payload is None else json.dumps(payload).encode("utf-8")
|
||||
headers = {"Accept": "application/json"}
|
||||
@@ -288,6 +386,7 @@ def _valid_name(name: str) -> str:
|
||||
|
||||
|
||||
def _local_tags() -> list[dict[str, Any]]:
|
||||
_apply_saved_connection()
|
||||
_discover_ollama_endpoint()
|
||||
try:
|
||||
payload = _json_request(LOCAL_OLLAMA + "/api/tags", timeout=15)
|
||||
@@ -902,10 +1001,11 @@ def _set_job(job_id: str, **values: Any) -> None:
|
||||
_jobs[job_id].update(values, updated_at=time.time())
|
||||
|
||||
|
||||
def _run_pull(job_id: str, name: str, action: str) -> None:
|
||||
def _run_pull(job_id: str, name: str, action: str, target: str) -> None:
|
||||
try:
|
||||
endpoint = _target_endpoint(target)
|
||||
payload = json.dumps({"name": name, "stream": True}).encode("utf-8")
|
||||
request = Request(LOCAL_OLLAMA + "/api/pull", data=payload, headers={"Content-Type": "application/json"}, method="POST")
|
||||
request = Request(endpoint + "/api/pull", data=payload, headers={"Content-Type": "application/json"}, method="POST")
|
||||
with urlopen(request, timeout=3600) as response:
|
||||
for raw_line in response:
|
||||
try:
|
||||
@@ -924,26 +1024,36 @@ def _run_pull(job_id: str, name: str, action: str) -> None:
|
||||
_set_job(job_id, state="failed", status="error", error=str(exc))
|
||||
|
||||
|
||||
def _run_delete(job_id: str, name: str) -> None:
|
||||
def _run_delete(job_id: str, name: str, target: str) -> None:
|
||||
try:
|
||||
_json_request(LOCAL_OLLAMA + "/api/delete", method="DELETE", payload={"name": name}, timeout=120)
|
||||
endpoint = _target_endpoint(target)
|
||||
_json_request(endpoint + "/api/delete", method="DELETE", payload={"name": name}, timeout=120)
|
||||
_set_job(job_id, state="completed", status="deleted", percent=100)
|
||||
except Exception as exc:
|
||||
_set_job(job_id, state="failed", status="error", error=str(exc))
|
||||
|
||||
|
||||
def _new_job(name: str, action: str) -> str:
|
||||
def _new_job(name: str, action: str, target: str = "local") -> str:
|
||||
target = str(target or "local").strip().lower()
|
||||
if target not in {"local", "remote"}:
|
||||
raise HTTPException(400, "Target must be local or remote")
|
||||
job_id = uuid.uuid4().hex
|
||||
with _jobs_lock:
|
||||
_jobs[job_id] = {"id": job_id, "name": name, "action": action, "state": "running", "status": "starting", "percent": 0, "created_at": time.time(), "updated_at": time.time()}
|
||||
target = _run_delete if action == "delete" else _run_pull
|
||||
args = (job_id, name) if action == "delete" else (job_id, name, action)
|
||||
threading.Thread(target=target, args=args, daemon=True, name=f"ollama-{action}-{job_id[:8]}").start()
|
||||
_jobs[job_id] = {"id": job_id, "name": name, "action": action, "target": target, "state": "running", "status": "starting", "percent": 0, "created_at": time.time(), "updated_at": time.time()}
|
||||
target_fn = _run_delete if action == "delete" else _run_pull
|
||||
args = (job_id, name, target) if action == "delete" else (job_id, name, action, target)
|
||||
threading.Thread(target=target_fn, args=args, daemon=True, name=f"ollama-{action}-{job_id[:8]}").start()
|
||||
return job_id
|
||||
|
||||
|
||||
class ModelRequest(BaseModel):
|
||||
name: str
|
||||
target: str = "local"
|
||||
|
||||
|
||||
class ConnectionRequest(BaseModel):
|
||||
url: str
|
||||
role: str = "local"
|
||||
|
||||
|
||||
class ModelsRequest(BaseModel):
|
||||
@@ -1362,6 +1472,39 @@ def chat(body: ChatRequest) -> dict[str, Any]:
|
||||
return {"ok": True, "request_id": request_id, "conversation_id": conversation_id, "model": selected[0], "models": selected, "message": {"role": "assistant", "content": combined}, "model_responses": {name: str((results.get(name, {}).get("message") or {}).get("content") or "") for name in selected if name in results}, "metrics": [results[name].get("metrics") for name in selected if name in results], "errors": errors, "done": True, "runtime": _runtime_snapshot()}
|
||||
|
||||
|
||||
@router.get("/connections")
|
||||
def connections() -> dict[str, Any]:
|
||||
saved = _read_connections()
|
||||
active = os.environ.get("OLLAMA_HOST", "").strip() or saved.get("active_url") or LOCAL_OLLAMA
|
||||
return {"active_url": active, "connections": _connection_snapshot(), "containerized": _running_in_container()}
|
||||
|
||||
|
||||
@router.post("/connections/test")
|
||||
def connections_test(body: ConnectionRequest) -> dict[str, Any]:
|
||||
result = _probe_endpoint(body.url, timeout=8)
|
||||
result["role"] = str(body.role or "local").strip().lower()
|
||||
return result
|
||||
|
||||
|
||||
@router.post("/connections/configure")
|
||||
def connections_configure(body: ConnectionRequest) -> dict[str, Any]:
|
||||
url = _valid_ollama_url(body.url)
|
||||
role = str(body.role or "local").strip().lower()
|
||||
if role not in {"local", "remote"}:
|
||||
raise HTTPException(400, "Connection role must be local or remote")
|
||||
saved = _read_connections()
|
||||
saved[f"{role}_url"] = url
|
||||
if role == "local" and not os.environ.get("OLLAMA_HOST", "").strip():
|
||||
saved["active_url"] = url
|
||||
_write_connections(saved)
|
||||
if role == "local" and not os.environ.get("OLLAMA_HOST", "").strip():
|
||||
global LOCAL_OLLAMA
|
||||
LOCAL_OLLAMA = url
|
||||
result = _probe_endpoint(url, timeout=8)
|
||||
result.update({"ok": True, "role": role, "message": f"Saved {role} Ollama endpoint"})
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/status")
|
||||
def status() -> dict[str, Any]:
|
||||
tags = _local_tags()
|
||||
@@ -1425,6 +1568,7 @@ def status() -> dict[str, Any]:
|
||||
row["variants"] = sorted(variants, key=lambda item: (item.get("size_bytes") or 0, item["name"]))
|
||||
|
||||
ollama_version = _ollama_version()
|
||||
connection_rows = _connection_snapshot()
|
||||
return {
|
||||
"ollama": {
|
||||
"available": bool(tags or ps_rows or ollama_version),
|
||||
@@ -1439,6 +1583,7 @@ def status() -> dict[str, Any]:
|
||||
else ""
|
||||
),
|
||||
},
|
||||
"connections": connection_rows,
|
||||
"models": local,
|
||||
"popular": popular,
|
||||
"popular_filter": {
|
||||
@@ -1484,19 +1629,25 @@ def catalog_refresh() -> dict[str, Any]:
|
||||
@router.post("/pull")
|
||||
def pull_model(body: ModelRequest) -> dict[str, Any]:
|
||||
name = _valid_name(body.name)
|
||||
return {"ok": True, "job_id": _new_job(name, "download"), "message": f"Downloading or updating {name}"}
|
||||
target = str(body.target or "local").strip().lower()
|
||||
_target_endpoint(target)
|
||||
return {"ok": True, "job_id": _new_job(name, "download", target), "message": f"Downloading or updating {name} on {target}"}
|
||||
|
||||
|
||||
@router.post("/redownload")
|
||||
def redownload_model(body: ModelRequest) -> dict[str, Any]:
|
||||
name = _valid_name(body.name)
|
||||
return {"ok": True, "job_id": _new_job(name, "redownload"), "message": f"Re-downloading or updating {name}"}
|
||||
target = str(body.target or "local").strip().lower()
|
||||
_target_endpoint(target)
|
||||
return {"ok": True, "job_id": _new_job(name, "redownload", target), "message": f"Re-downloading or updating {name} on {target}"}
|
||||
|
||||
|
||||
@router.delete("/model")
|
||||
def delete_model(body: ModelRequest) -> dict[str, Any]:
|
||||
name = _valid_name(body.name)
|
||||
return {"ok": True, "job_id": _new_job(name, "delete"), "message": f"Removing {name}"}
|
||||
target = str(body.target or "local").strip().lower()
|
||||
_target_endpoint(target)
|
||||
return {"ok": True, "job_id": _new_job(name, "delete", target), "message": f"Removing {name} from {target}"}
|
||||
|
||||
|
||||
def create_ollama_routes(app) -> None:
|
||||
|
||||
Reference in New Issue
Block a user