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() @asynccontextmanager 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) @app.get("/health") 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(), } @app.get("/v1/models", dependencies=[Depends(require_api_key)]) 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) @app.post("/v1/chat/completions", dependencies=[Depends(require_api_key)]) 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") @app.post("/api/sessions", dependencies=[Depends(require_api_key)]) 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} @app.get("/api/sessions/{session_id}", dependencies=[Depends(require_api_key)]) 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 @app.post("/api/tasks", status_code=202, dependencies=[Depends(require_api_key)]) 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}"} @app.get("/api/tasks/{job_id}", dependencies=[Depends(require_api_key)]) 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 @app.post("/api/tasks/{job_id}/cancel", dependencies=[Depends(require_api_key)]) 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", } @app.api_route( "/{full_path:path}", methods=["GET", "HEAD", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"], ) 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"), )