Spaces:
Paused
Paused
Download gateway.py from Leon4gr45/llama: direct link, hf CLI and curl.
- Browser
- Download file 17.8 kB
-
https://huggingface.co/spaces/Leon4gr45/llama/resolve/main/gateway.py
- Command line
-
hf download hf://spaces/Leon4gr45/llama/gateway.py
-
curl -L -o gateway.py https://huggingface.co/spaces/Leon4gr45/llama/resolve/main/gateway.py
17.8 kB
| from __future__ import annotations | |
| import asyncio | |
| import json | |
| import os | |
| import sqlite3 | |
| import time | |
| import uuid | |
| from contextlib import asynccontextmanager | |
| from pathlib import Path | |
| from typing import Any | |
| import httpx | |
| from fastapi import Depends, FastAPI, Header, HTTPException, Request | |
| from fastapi.responses import JSONResponse, StreamingResponse | |
| LLAMA_URL = os.getenv("LLAMA_URL", "http://127.0.0.1:8080").rstrip("/") | |
| STATE_PATH = Path(os.getenv("STATE_PATH", "/data/agent-state.db")) | |
| MODEL_ALIAS = os.getenv("MODEL_ALIAS", "spark-x2.5-1.7b") | |
| DEFAULT_LANGUAGE = os.getenv("DEFAULT_LANGUAGE", "en") | |
| MAX_QUEUE_SIZE = int(os.getenv("MAX_QUEUE_SIZE", "32")) | |
| REQUEST_TIMEOUT = float(os.getenv("REQUEST_TIMEOUT_SECONDS", "900")) | |
| API_KEY = os.getenv("API_KEY", "").strip() | |
| LLAMA_API_KEY = os.getenv("LLAMA_API_KEY", "").strip() | |
| # Spark X2.5 is a reasoning model: with thinking on it routinely spends the | |
| # whole max_tokens budget reasoning and returns an empty answer, which breaks | |
| # agent callers. Default it off; a caller can still pass | |
| # chat_template_kwargs={"enable_thinking": true} explicitly. | |
| DEFAULT_ENABLE_THINKING = os.getenv("DEFAULT_ENABLE_THINKING", "false").lower() == "true" | |
| ENGLISH_GUARD = ( | |
| "Respond entirely in English unless the user explicitly requests another " | |
| "language. Preserve code, identifiers, and quoted source material exactly." | |
| ) | |
| inference_lock = asyncio.Lock() | |
| task_queue: asyncio.PriorityQueue[tuple[int, float, str]] = asyncio.PriorityQueue( | |
| maxsize=MAX_QUEUE_SIZE | |
| ) | |
| worker_task: asyncio.Task[None] | None = None | |
| def connect() -> sqlite3.Connection: | |
| STATE_PATH.parent.mkdir(parents=True, exist_ok=True) | |
| db = sqlite3.connect(STATE_PATH, timeout=30) | |
| db.row_factory = sqlite3.Row | |
| db.execute("PRAGMA journal_mode=WAL") | |
| db.execute("PRAGMA foreign_keys=ON") | |
| return db | |
| def initialise_database() -> None: | |
| with connect() as db: | |
| db.executescript( | |
| """ | |
| CREATE TABLE IF NOT EXISTS sessions ( | |
| id TEXT PRIMARY KEY, | |
| language TEXT NOT NULL, | |
| summary TEXT NOT NULL DEFAULT '', | |
| metadata TEXT NOT NULL DEFAULT '{}', | |
| created_at REAL NOT NULL, | |
| updated_at REAL NOT NULL | |
| ); | |
| CREATE TABLE IF NOT EXISTS messages ( | |
| id INTEGER PRIMARY KEY AUTOINCREMENT, | |
| session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE, | |
| role TEXT NOT NULL, | |
| content TEXT NOT NULL, | |
| created_at REAL NOT NULL | |
| ); | |
| CREATE INDEX IF NOT EXISTS messages_session_idx | |
| ON messages(session_id, id); | |
| CREATE TABLE IF NOT EXISTS jobs ( | |
| id TEXT PRIMARY KEY, | |
| session_id TEXT REFERENCES sessions(id) ON DELETE SET NULL, | |
| priority INTEGER NOT NULL, | |
| status TEXT NOT NULL, | |
| request_json TEXT NOT NULL, | |
| response_json TEXT, | |
| error TEXT, | |
| attempts INTEGER NOT NULL DEFAULT 0, | |
| created_at REAL NOT NULL, | |
| updated_at REAL NOT NULL | |
| ); | |
| CREATE INDEX IF NOT EXISTS jobs_status_idx | |
| ON jobs(status, priority, created_at); | |
| """ | |
| ) | |
| db.execute( | |
| "UPDATE jobs SET status='queued', updated_at=? WHERE status='running'", | |
| (time.time(),), | |
| ) | |
| def require_api_key(authorization: str | None = Header(default=None)) -> None: | |
| if not API_KEY: | |
| return | |
| if authorization != f"Bearer {API_KEY}": | |
| raise HTTPException(status_code=401, detail="Invalid API key") | |
| def apply_defaults(payload: dict[str, Any]) -> dict[str, Any]: | |
| kwargs = dict(payload.get("chat_template_kwargs") or {}) | |
| kwargs.setdefault("enable_thinking", DEFAULT_ENABLE_THINKING) | |
| payload = {**payload, "chat_template_kwargs": kwargs} | |
| return ensure_language_guard(payload) | |
| def ensure_language_guard(payload: dict[str, Any]) -> dict[str, Any]: | |
| if DEFAULT_LANGUAGE.lower() != "en": | |
| return payload | |
| messages = list(payload.get("messages") or []) | |
| guard = {"role": "system", "content": ENGLISH_GUARD} | |
| if messages and messages[0].get("role") == "system": | |
| messages[0] = { | |
| **messages[0], | |
| "content": f"{messages[0].get('content', '')}\n\n{ENGLISH_GUARD}", | |
| } | |
| else: | |
| messages.insert(0, guard) | |
| return {**payload, "model": payload.get("model") or MODEL_ALIAS, "messages": messages} | |
| async def wait_for_llama() -> None: | |
| deadline = time.monotonic() + REQUEST_TIMEOUT | |
| async with httpx.AsyncClient() as client: | |
| while time.monotonic() < deadline: | |
| try: | |
| response = await client.get(f"{LLAMA_URL}/health", timeout=5) | |
| if response.status_code == 200: | |
| return | |
| except httpx.HTTPError: | |
| pass | |
| await asyncio.sleep(2) | |
| raise RuntimeError("llama.cpp did not become healthy before the startup timeout") | |
| def llama_auth_headers() -> dict[str, str]: | |
| # llama-server reads its own API key from the LLAMA_API_KEY env var | |
| # (llama.cpp maps --api-key to that specific name, not the generic | |
| # LLAMA_ARG_* pattern), independently of this gateway's own API_KEY. | |
| # The gateway holds the same secret (it's set on the whole | |
| # container), so attach it on every request we make to llama-server | |
| # rather than relying on whatever the external caller supplied to us. | |
| return {"Authorization": f"Bearer {LLAMA_API_KEY}"} if LLAMA_API_KEY else {} | |
| async def execute_completion(payload: dict[str, Any]) -> dict[str, Any]: | |
| await wait_for_llama() | |
| async with inference_lock: | |
| async with httpx.AsyncClient(timeout=REQUEST_TIMEOUT) as client: | |
| response = await client.post( | |
| f"{LLAMA_URL}/v1/chat/completions", | |
| json=apply_defaults(payload), | |
| headers=llama_auth_headers(), | |
| ) | |
| response.raise_for_status() | |
| return response.json() | |
| async def job_worker() -> None: | |
| while True: | |
| _priority, _created_at, job_id = await task_queue.get() | |
| try: | |
| with connect() as db: | |
| row = db.execute("SELECT * FROM jobs WHERE id=?", (job_id,)).fetchone() | |
| if row is None or row["status"] != "queued": | |
| continue | |
| db.execute( | |
| "UPDATE jobs SET status='running', attempts=attempts+1, updated_at=? WHERE id=?", | |
| (time.time(), job_id), | |
| ) | |
| payload = json.loads(row["request_json"]) | |
| session_id = row["session_id"] | |
| if session_id: | |
| history = db.execute( | |
| "SELECT role, content FROM messages WHERE session_id=? ORDER BY id DESC LIMIT 20", | |
| (session_id,), | |
| ).fetchall() | |
| historical_messages = [dict(item) for item in reversed(history)] | |
| payload = { | |
| **payload, | |
| "messages": historical_messages + list(payload.get("messages") or []), | |
| } | |
| try: | |
| result = await execute_completion(payload) | |
| with connect() as db: | |
| db.execute( | |
| "UPDATE jobs SET status='completed', response_json=?, updated_at=? WHERE id=?", | |
| (json.dumps(result), time.time(), job_id), | |
| ) | |
| if session_id: | |
| for message in json.loads(row["request_json"]).get("messages", []): | |
| if message.get("role") in {"user", "assistant"} and isinstance( | |
| message.get("content"), str | |
| ): | |
| db.execute( | |
| "INSERT INTO messages(session_id, role, content, created_at) VALUES(?,?,?,?)", | |
| (session_id, message["role"], message["content"], time.time()), | |
| ) | |
| choices = result.get("choices") or [] | |
| assistant_content = ( | |
| choices[0].get("message", {}).get("content") if choices else None | |
| ) | |
| if isinstance(assistant_content, str): | |
| db.execute( | |
| "INSERT INTO messages(session_id, role, content, created_at) VALUES(?,?,?,?)", | |
| (session_id, "assistant", assistant_content, time.time()), | |
| ) | |
| db.execute( | |
| "UPDATE sessions SET updated_at=? WHERE id=?", | |
| (time.time(), session_id), | |
| ) | |
| except Exception as exc: | |
| with connect() as db: | |
| db.execute( | |
| "UPDATE jobs SET status='failed', error=?, updated_at=? WHERE id=?", | |
| (str(exc)[:2000], time.time(), job_id), | |
| ) | |
| finally: | |
| task_queue.task_done() | |
| async def lifespan(_app: FastAPI): | |
| global worker_task | |
| initialise_database() | |
| with connect() as db: | |
| queued = db.execute( | |
| "SELECT id, priority, created_at FROM jobs WHERE status='queued' ORDER BY priority, created_at" | |
| ).fetchall() | |
| for row in queued: | |
| await task_queue.put((row["priority"], row["created_at"], row["id"])) | |
| worker_task = asyncio.create_task(job_worker()) | |
| yield | |
| worker_task.cancel() | |
| try: | |
| await worker_task | |
| except asyncio.CancelledError: | |
| pass | |
| app = FastAPI(title="Spark X2.5 Agent Gateway", version="1.0.0", lifespan=lifespan) | |
| async def health() -> dict[str, Any]: | |
| try: | |
| async with httpx.AsyncClient() as client: | |
| response = await client.get(f"{LLAMA_URL}/health", timeout=5) | |
| model_ready = response.status_code == 200 | |
| except httpx.HTTPError: | |
| model_ready = False | |
| return { | |
| "status": "ok" if model_ready else "degraded", | |
| "model_ready": model_ready, | |
| "queue_depth": task_queue.qsize(), | |
| "inference_busy": inference_lock.locked(), | |
| } | |
| async def models() -> JSONResponse: | |
| async with httpx.AsyncClient() as client: | |
| response = await client.get( | |
| f"{LLAMA_URL}/v1/models", timeout=30, headers=llama_auth_headers() | |
| ) | |
| return JSONResponse(response.json(), status_code=response.status_code) | |
| async def chat_completions(request: Request): | |
| payload = await request.json() | |
| if not payload.get("stream", False): | |
| try: | |
| return await execute_completion(payload) | |
| except httpx.HTTPStatusError as exc: | |
| raise HTTPException(exc.response.status_code, exc.response.text) from exc | |
| payload = apply_defaults(payload) | |
| async def stream_response(): | |
| await wait_for_llama() | |
| await inference_lock.acquire() | |
| try: | |
| async with httpx.AsyncClient(timeout=REQUEST_TIMEOUT) as client: | |
| async with client.stream( | |
| "POST", | |
| f"{LLAMA_URL}/v1/chat/completions", | |
| json=payload, | |
| headers=llama_auth_headers(), | |
| ) as response: | |
| if response.status_code >= 400: | |
| body = await response.aread() | |
| yield body | |
| return | |
| async for chunk in response.aiter_bytes(): | |
| yield chunk | |
| finally: | |
| inference_lock.release() | |
| return StreamingResponse(stream_response(), media_type="text/event-stream") | |
| async def create_session(request: Request) -> dict[str, Any]: | |
| body = await request.json() | |
| now = time.time() | |
| session_id = f"ses_{uuid.uuid4().hex}" | |
| language = body.get("language", DEFAULT_LANGUAGE) | |
| metadata = body.get("metadata", {}) | |
| with connect() as db: | |
| db.execute( | |
| "INSERT INTO sessions(id, language, metadata, created_at, updated_at) VALUES(?,?,?,?,?)", | |
| (session_id, language, json.dumps(metadata), now, now), | |
| ) | |
| return {"session_id": session_id, "language": language, "created_at": now} | |
| async def get_session(session_id: str) -> dict[str, Any]: | |
| with connect() as db: | |
| row = db.execute("SELECT * FROM sessions WHERE id=?", (session_id,)).fetchone() | |
| if row is None: | |
| raise HTTPException(404, "Session not found") | |
| result = dict(row) | |
| result["metadata"] = json.loads(result["metadata"]) | |
| with connect() as db: | |
| messages = db.execute( | |
| "SELECT role, content, created_at FROM messages WHERE session_id=? ORDER BY id DESC LIMIT 20", | |
| (session_id,), | |
| ).fetchall() | |
| result["recent_messages"] = [dict(item) for item in reversed(messages)] | |
| return result | |
| async def create_task(request: Request) -> dict[str, Any]: | |
| if task_queue.full(): | |
| raise HTTPException(429, "Queue is full") | |
| body = await request.json() | |
| payload = body.get("request") | |
| if not isinstance(payload, dict) or not isinstance(payload.get("messages"), list): | |
| raise HTTPException(422, "request.messages must be provided") | |
| priority = max(0, min(int(body.get("priority", 1)), 3)) | |
| session_id = body.get("session_id") | |
| now = time.time() | |
| job_id = f"job_{uuid.uuid4().hex}" | |
| with connect() as db: | |
| if session_id and db.execute( | |
| "SELECT 1 FROM sessions WHERE id=?", (session_id,) | |
| ).fetchone() is None: | |
| raise HTTPException(404, "Session not found") | |
| db.execute( | |
| "INSERT INTO jobs(id, session_id, priority, status, request_json, created_at, updated_at) " | |
| "VALUES(?,?,?,?,?,?,?)", | |
| (job_id, session_id, priority, "queued", json.dumps(payload), now, now), | |
| ) | |
| await task_queue.put((priority, now, job_id)) | |
| return {"task_id": job_id, "status": "queued", "status_url": f"/api/tasks/{job_id}"} | |
| async def get_task(job_id: str) -> dict[str, Any]: | |
| with connect() as db: | |
| row = db.execute("SELECT * FROM jobs WHERE id=?", (job_id,)).fetchone() | |
| if row is None: | |
| raise HTTPException(404, "Task not found") | |
| result = dict(row) | |
| result["request"] = json.loads(result.pop("request_json")) | |
| response_json = result.pop("response_json") | |
| result["response"] = json.loads(response_json) if response_json else None | |
| return result | |
| async def cancel_task(job_id: str) -> dict[str, Any]: | |
| with connect() as db: | |
| row = db.execute("SELECT status FROM jobs WHERE id=?", (job_id,)).fetchone() | |
| if row is None: | |
| raise HTTPException(404, "Task not found") | |
| if row["status"] != "queued": | |
| raise HTTPException(409, "Only queued tasks can be cancelled") | |
| db.execute( | |
| "UPDATE jobs SET status='cancelled', updated_at=? WHERE id=?", | |
| (time.time(), job_id), | |
| ) | |
| return {"task_id": job_id, "status": "cancelled"} | |
| # Everything below is not part of this service's own API. It falls through to | |
| # llama.cpp's built-in web UI (served on LLAMA_URL, reachable only from | |
| # localhost inside the container), so that visiting the Space's public URL | |
| # renders the llama.cpp chat UI instead of a 404. | |
| _HOP_BY_HOP_HEADERS = { | |
| "connection", | |
| "keep-alive", | |
| "proxy-authenticate", | |
| "proxy-authorization", | |
| "te", | |
| "trailers", | |
| "transfer-encoding", | |
| "upgrade", | |
| "content-length", | |
| "host", | |
| } | |
| async def proxy_to_llama_webui(full_path: str, request: Request) -> Any: | |
| url = f"{LLAMA_URL}/{full_path}" | |
| headers = { | |
| key: value | |
| for key, value in request.headers.items() | |
| if key.lower() not in _HOP_BY_HOP_HEADERS | |
| } | |
| body = await request.body() | |
| client = httpx.AsyncClient(timeout=REQUEST_TIMEOUT) | |
| try: | |
| upstream_request = client.build_request( | |
| request.method, | |
| url, | |
| params=request.query_params, | |
| headers=headers, | |
| content=body, | |
| ) | |
| upstream_response = await client.send(upstream_request, stream=True) | |
| except httpx.HTTPError: | |
| await client.aclose() | |
| raise HTTPException(502, "llama.cpp server is unavailable") | |
| response_headers = { | |
| key: value | |
| for key, value in upstream_response.headers.items() | |
| if key.lower() not in _HOP_BY_HOP_HEADERS | |
| } | |
| async def stream_and_close(): | |
| try: | |
| async for chunk in upstream_response.aiter_raw(): | |
| yield chunk | |
| finally: | |
| await upstream_response.aclose() | |
| await client.aclose() | |
| return StreamingResponse( | |
| stream_and_close(), | |
| status_code=upstream_response.status_code, | |
| headers=response_headers, | |
| media_type=upstream_response.headers.get("content-type"), | |
| ) |