Leon4gr45 commited on
Commit
e841cbd
·
verified ·
1 Parent(s): 88f4721

Upload 4 files

Browse files
Files changed (4) hide show
  1. Dockerfile +24 -3
  2. gateway.py +359 -0
  3. requirements.txt +3 -0
  4. start.sh +40 -0
Dockerfile CHANGED
@@ -1,7 +1,28 @@
1
  FROM ghcr.io/ggml-org/llama.cpp:server
2
 
3
- EXPOSE 7860
 
 
 
 
 
 
 
 
 
4
 
5
- ENTRYPOINT ["/app/llama-server"]
 
 
 
 
 
 
 
 
 
 
 
 
6
 
7
- CMD ["--host", "0.0.0.0", "--port", "7860", "--hf-repo", "XHToken/Spark-X2.5-1.7B-GGUF:Q8_0", "--alias", "spark-x2.5-1.7b", "--ctx-size", "32768", "--parallel", "1", "--threads", "2", "--threads-batch", "2", "--batch-size", "512", "--ubatch-size", "256", "--cache-type-k", "q8_0", "--cache-type-v", "q8_0", "--flash-attn", "auto", "--cont-batching", "--metrics"]
 
1
  FROM ghcr.io/ggml-org/llama.cpp:server
2
 
3
+ USER root
4
+
5
+ RUN apt-get update \
6
+ && apt-get install -y --no-install-recommends python3 python3-pip ca-certificates \
7
+ && rm -rf /var/lib/apt/lists/*
8
+
9
+ WORKDIR /srv
10
+
11
+ COPY requirements.txt /srv/requirements.txt
12
+ RUN python3 -m pip install --no-cache-dir --break-system-packages -r /srv/requirements.txt
13
 
14
+ COPY gateway.py /srv/gateway.py
15
+ COPY start.sh /srv/start.sh
16
+ RUN chmod +x /srv/start.sh \
17
+ && mkdir -p /data
18
+
19
+ ENV LLAMA_URL=http://127.0.0.1:8080
20
+ ENV STATE_PATH=/data/agent-state.db
21
+ ENV MODEL_ALIAS=spark-x2.5-1.7b
22
+ ENV DEFAULT_LANGUAGE=en
23
+ ENV MAX_QUEUE_SIZE=32
24
+ ENV REQUEST_TIMEOUT_SECONDS=900
25
+
26
+ EXPOSE 7860
27
 
28
+ ENTRYPOINT ["/srv/start.sh"]
gateway.py ADDED
@@ -0,0 +1,359 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import asyncio
4
+ import json
5
+ import os
6
+ import sqlite3
7
+ import time
8
+ import uuid
9
+ from contextlib import asynccontextmanager
10
+ from pathlib import Path
11
+ from typing import Any
12
+
13
+ import httpx
14
+ from fastapi import Depends, FastAPI, Header, HTTPException, Request
15
+ from fastapi.responses import JSONResponse, StreamingResponse
16
+
17
+
18
+ LLAMA_URL = os.getenv("LLAMA_URL", "http://127.0.0.1:8080").rstrip("/")
19
+ STATE_PATH = Path(os.getenv("STATE_PATH", "/data/agent-state.db"))
20
+ MODEL_ALIAS = os.getenv("MODEL_ALIAS", "spark-x2.5-1.7b")
21
+ DEFAULT_LANGUAGE = os.getenv("DEFAULT_LANGUAGE", "en")
22
+ MAX_QUEUE_SIZE = int(os.getenv("MAX_QUEUE_SIZE", "32"))
23
+ REQUEST_TIMEOUT = float(os.getenv("REQUEST_TIMEOUT_SECONDS", "900"))
24
+ API_KEY = os.getenv("API_KEY", "").strip()
25
+ ENGLISH_GUARD = (
26
+ "Respond entirely in English unless the user explicitly requests another "
27
+ "language. Preserve code, identifiers, and quoted source material exactly."
28
+ )
29
+
30
+ inference_lock = asyncio.Lock()
31
+ task_queue: asyncio.PriorityQueue[tuple[int, float, str]] = asyncio.PriorityQueue(
32
+ maxsize=MAX_QUEUE_SIZE
33
+ )
34
+ worker_task: asyncio.Task[None] | None = None
35
+
36
+
37
+ def connect() -> sqlite3.Connection:
38
+ STATE_PATH.parent.mkdir(parents=True, exist_ok=True)
39
+ db = sqlite3.connect(STATE_PATH, timeout=30)
40
+ db.row_factory = sqlite3.Row
41
+ db.execute("PRAGMA journal_mode=WAL")
42
+ db.execute("PRAGMA foreign_keys=ON")
43
+ return db
44
+
45
+
46
+ def initialise_database() -> None:
47
+ with connect() as db:
48
+ db.executescript(
49
+ """
50
+ CREATE TABLE IF NOT EXISTS sessions (
51
+ id TEXT PRIMARY KEY,
52
+ language TEXT NOT NULL,
53
+ summary TEXT NOT NULL DEFAULT '',
54
+ metadata TEXT NOT NULL DEFAULT '{}',
55
+ created_at REAL NOT NULL,
56
+ updated_at REAL NOT NULL
57
+ );
58
+ CREATE TABLE IF NOT EXISTS messages (
59
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
60
+ session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
61
+ role TEXT NOT NULL,
62
+ content TEXT NOT NULL,
63
+ created_at REAL NOT NULL
64
+ );
65
+ CREATE INDEX IF NOT EXISTS messages_session_idx
66
+ ON messages(session_id, id);
67
+ CREATE TABLE IF NOT EXISTS jobs (
68
+ id TEXT PRIMARY KEY,
69
+ session_id TEXT REFERENCES sessions(id) ON DELETE SET NULL,
70
+ priority INTEGER NOT NULL,
71
+ status TEXT NOT NULL,
72
+ request_json TEXT NOT NULL,
73
+ response_json TEXT,
74
+ error TEXT,
75
+ attempts INTEGER NOT NULL DEFAULT 0,
76
+ created_at REAL NOT NULL,
77
+ updated_at REAL NOT NULL
78
+ );
79
+ CREATE INDEX IF NOT EXISTS jobs_status_idx
80
+ ON jobs(status, priority, created_at);
81
+ """
82
+ )
83
+ db.execute(
84
+ "UPDATE jobs SET status='queued', updated_at=? WHERE status='running'",
85
+ (time.time(),),
86
+ )
87
+
88
+
89
+ def require_api_key(authorization: str | None = Header(default=None)) -> None:
90
+ if not API_KEY:
91
+ return
92
+ if authorization != f"Bearer {API_KEY}":
93
+ raise HTTPException(status_code=401, detail="Invalid API key")
94
+
95
+
96
+ def ensure_language_guard(payload: dict[str, Any]) -> dict[str, Any]:
97
+ if DEFAULT_LANGUAGE.lower() != "en":
98
+ return payload
99
+ messages = list(payload.get("messages") or [])
100
+ guard = {"role": "system", "content": ENGLISH_GUARD}
101
+ if messages and messages[0].get("role") == "system":
102
+ messages[0] = {
103
+ **messages[0],
104
+ "content": f"{messages[0].get('content', '')}\n\n{ENGLISH_GUARD}",
105
+ }
106
+ else:
107
+ messages.insert(0, guard)
108
+ return {**payload, "model": payload.get("model") or MODEL_ALIAS, "messages": messages}
109
+
110
+
111
+ async def wait_for_llama() -> None:
112
+ deadline = time.monotonic() + REQUEST_TIMEOUT
113
+ async with httpx.AsyncClient() as client:
114
+ while time.monotonic() < deadline:
115
+ try:
116
+ response = await client.get(f"{LLAMA_URL}/health", timeout=5)
117
+ if response.status_code == 200:
118
+ return
119
+ except httpx.HTTPError:
120
+ pass
121
+ await asyncio.sleep(2)
122
+ raise RuntimeError("llama.cpp did not become healthy before the startup timeout")
123
+
124
+
125
+ async def execute_completion(payload: dict[str, Any]) -> dict[str, Any]:
126
+ await wait_for_llama()
127
+ async with inference_lock:
128
+ async with httpx.AsyncClient(timeout=REQUEST_TIMEOUT) as client:
129
+ response = await client.post(
130
+ f"{LLAMA_URL}/v1/chat/completions",
131
+ json=ensure_language_guard(payload),
132
+ )
133
+ response.raise_for_status()
134
+ return response.json()
135
+
136
+
137
+ async def job_worker() -> None:
138
+ while True:
139
+ _priority, _created_at, job_id = await task_queue.get()
140
+ try:
141
+ with connect() as db:
142
+ row = db.execute("SELECT * FROM jobs WHERE id=?", (job_id,)).fetchone()
143
+ if row is None or row["status"] != "queued":
144
+ continue
145
+ db.execute(
146
+ "UPDATE jobs SET status='running', attempts=attempts+1, updated_at=? WHERE id=?",
147
+ (time.time(), job_id),
148
+ )
149
+ payload = json.loads(row["request_json"])
150
+ session_id = row["session_id"]
151
+ if session_id:
152
+ history = db.execute(
153
+ "SELECT role, content FROM messages WHERE session_id=? ORDER BY id DESC LIMIT 20",
154
+ (session_id,),
155
+ ).fetchall()
156
+ historical_messages = [dict(item) for item in reversed(history)]
157
+ payload = {
158
+ **payload,
159
+ "messages": historical_messages + list(payload.get("messages") or []),
160
+ }
161
+ try:
162
+ result = await execute_completion(payload)
163
+ with connect() as db:
164
+ db.execute(
165
+ "UPDATE jobs SET status='completed', response_json=?, updated_at=? WHERE id=?",
166
+ (json.dumps(result), time.time(), job_id),
167
+ )
168
+ if session_id:
169
+ for message in json.loads(row["request_json"]).get("messages", []):
170
+ if message.get("role") in {"user", "assistant"} and isinstance(
171
+ message.get("content"), str
172
+ ):
173
+ db.execute(
174
+ "INSERT INTO messages(session_id, role, content, created_at) VALUES(?,?,?,?)",
175
+ (session_id, message["role"], message["content"], time.time()),
176
+ )
177
+ choices = result.get("choices") or []
178
+ assistant_content = (
179
+ choices[0].get("message", {}).get("content") if choices else None
180
+ )
181
+ if isinstance(assistant_content, str):
182
+ db.execute(
183
+ "INSERT INTO messages(session_id, role, content, created_at) VALUES(?,?,?,?)",
184
+ (session_id, "assistant", assistant_content, time.time()),
185
+ )
186
+ db.execute(
187
+ "UPDATE sessions SET updated_at=? WHERE id=?",
188
+ (time.time(), session_id),
189
+ )
190
+ except Exception as exc:
191
+ with connect() as db:
192
+ db.execute(
193
+ "UPDATE jobs SET status='failed', error=?, updated_at=? WHERE id=?",
194
+ (str(exc)[:2000], time.time(), job_id),
195
+ )
196
+ finally:
197
+ task_queue.task_done()
198
+
199
+
200
+ @asynccontextmanager
201
+ async def lifespan(_app: FastAPI):
202
+ global worker_task
203
+ initialise_database()
204
+ with connect() as db:
205
+ queued = db.execute(
206
+ "SELECT id, priority, created_at FROM jobs WHERE status='queued' ORDER BY priority, created_at"
207
+ ).fetchall()
208
+ for row in queued:
209
+ await task_queue.put((row["priority"], row["created_at"], row["id"]))
210
+ worker_task = asyncio.create_task(job_worker())
211
+ yield
212
+ worker_task.cancel()
213
+ try:
214
+ await worker_task
215
+ except asyncio.CancelledError:
216
+ pass
217
+
218
+
219
+ app = FastAPI(title="Spark X2.5 Agent Gateway", version="1.0.0", lifespan=lifespan)
220
+
221
+
222
+ @app.get("/health")
223
+ async def health() -> dict[str, Any]:
224
+ try:
225
+ async with httpx.AsyncClient() as client:
226
+ response = await client.get(f"{LLAMA_URL}/health", timeout=5)
227
+ model_ready = response.status_code == 200
228
+ except httpx.HTTPError:
229
+ model_ready = False
230
+ return {
231
+ "status": "ok" if model_ready else "degraded",
232
+ "model_ready": model_ready,
233
+ "queue_depth": task_queue.qsize(),
234
+ "inference_busy": inference_lock.locked(),
235
+ }
236
+
237
+
238
+ @app.get("/v1/models", dependencies=[Depends(require_api_key)])
239
+ async def models() -> JSONResponse:
240
+ async with httpx.AsyncClient() as client:
241
+ response = await client.get(f"{LLAMA_URL}/v1/models", timeout=30)
242
+ return JSONResponse(response.json(), status_code=response.status_code)
243
+
244
+
245
+ @app.post("/v1/chat/completions", dependencies=[Depends(require_api_key)])
246
+ async def chat_completions(request: Request):
247
+ payload = await request.json()
248
+ if not payload.get("stream", False):
249
+ try:
250
+ return await execute_completion(payload)
251
+ except httpx.HTTPStatusError as exc:
252
+ raise HTTPException(exc.response.status_code, exc.response.text) from exc
253
+
254
+ payload = ensure_language_guard(payload)
255
+
256
+ async def stream_response():
257
+ await wait_for_llama()
258
+ await inference_lock.acquire()
259
+ try:
260
+ async with httpx.AsyncClient(timeout=REQUEST_TIMEOUT) as client:
261
+ async with client.stream(
262
+ "POST", f"{LLAMA_URL}/v1/chat/completions", json=payload
263
+ ) as response:
264
+ if response.status_code >= 400:
265
+ body = await response.aread()
266
+ yield body
267
+ return
268
+ async for chunk in response.aiter_bytes():
269
+ yield chunk
270
+ finally:
271
+ inference_lock.release()
272
+
273
+ return StreamingResponse(stream_response(), media_type="text/event-stream")
274
+
275
+
276
+ @app.post("/api/sessions", dependencies=[Depends(require_api_key)])
277
+ async def create_session(request: Request) -> dict[str, Any]:
278
+ body = await request.json()
279
+ now = time.time()
280
+ session_id = f"ses_{uuid.uuid4().hex}"
281
+ language = body.get("language", DEFAULT_LANGUAGE)
282
+ metadata = body.get("metadata", {})
283
+ with connect() as db:
284
+ db.execute(
285
+ "INSERT INTO sessions(id, language, metadata, created_at, updated_at) VALUES(?,?,?,?,?)",
286
+ (session_id, language, json.dumps(metadata), now, now),
287
+ )
288
+ return {"session_id": session_id, "language": language, "created_at": now}
289
+
290
+
291
+ @app.get("/api/sessions/{session_id}", dependencies=[Depends(require_api_key)])
292
+ async def get_session(session_id: str) -> dict[str, Any]:
293
+ with connect() as db:
294
+ row = db.execute("SELECT * FROM sessions WHERE id=?", (session_id,)).fetchone()
295
+ if row is None:
296
+ raise HTTPException(404, "Session not found")
297
+ result = dict(row)
298
+ result["metadata"] = json.loads(result["metadata"])
299
+ with connect() as db:
300
+ messages = db.execute(
301
+ "SELECT role, content, created_at FROM messages WHERE session_id=? ORDER BY id DESC LIMIT 20",
302
+ (session_id,),
303
+ ).fetchall()
304
+ result["recent_messages"] = [dict(item) for item in reversed(messages)]
305
+ return result
306
+
307
+
308
+ @app.post("/api/tasks", status_code=202, dependencies=[Depends(require_api_key)])
309
+ async def create_task(request: Request) -> dict[str, Any]:
310
+ if task_queue.full():
311
+ raise HTTPException(429, "Queue is full")
312
+ body = await request.json()
313
+ payload = body.get("request")
314
+ if not isinstance(payload, dict) or not isinstance(payload.get("messages"), list):
315
+ raise HTTPException(422, "request.messages must be provided")
316
+ priority = max(0, min(int(body.get("priority", 1)), 3))
317
+ session_id = body.get("session_id")
318
+ now = time.time()
319
+ job_id = f"job_{uuid.uuid4().hex}"
320
+ with connect() as db:
321
+ if session_id and db.execute(
322
+ "SELECT 1 FROM sessions WHERE id=?", (session_id,)
323
+ ).fetchone() is None:
324
+ raise HTTPException(404, "Session not found")
325
+ db.execute(
326
+ "INSERT INTO jobs(id, session_id, priority, status, request_json, created_at, updated_at) "
327
+ "VALUES(?,?,?,?,?,?,?)",
328
+ (job_id, session_id, priority, "queued", json.dumps(payload), now, now),
329
+ )
330
+ await task_queue.put((priority, now, job_id))
331
+ return {"task_id": job_id, "status": "queued", "status_url": f"/api/tasks/{job_id}"}
332
+
333
+
334
+ @app.get("/api/tasks/{job_id}", dependencies=[Depends(require_api_key)])
335
+ async def get_task(job_id: str) -> dict[str, Any]:
336
+ with connect() as db:
337
+ row = db.execute("SELECT * FROM jobs WHERE id=?", (job_id,)).fetchone()
338
+ if row is None:
339
+ raise HTTPException(404, "Task not found")
340
+ result = dict(row)
341
+ result["request"] = json.loads(result.pop("request_json"))
342
+ response_json = result.pop("response_json")
343
+ result["response"] = json.loads(response_json) if response_json else None
344
+ return result
345
+
346
+
347
+ @app.post("/api/tasks/{job_id}/cancel", dependencies=[Depends(require_api_key)])
348
+ async def cancel_task(job_id: str) -> dict[str, Any]:
349
+ with connect() as db:
350
+ row = db.execute("SELECT status FROM jobs WHERE id=?", (job_id,)).fetchone()
351
+ if row is None:
352
+ raise HTTPException(404, "Task not found")
353
+ if row["status"] != "queued":
354
+ raise HTTPException(409, "Only queued tasks can be cancelled")
355
+ db.execute(
356
+ "UPDATE jobs SET status='cancelled', updated_at=? WHERE id=?",
357
+ (time.time(), job_id),
358
+ )
359
+ return {"task_id": job_id, "status": "cancelled"}
requirements.txt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ fastapi==0.116.1
2
+ httpx==0.28.1
3
+ uvicorn[standard]==0.35.0
start.sh ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/bin/sh
2
+ set -eu
3
+
4
+ MODEL_REPO="${MODEL_REPO:-XHToken/Spark-X2.5-1.7B-GGUF:Q8_0}"
5
+ MODEL_ALIAS="${MODEL_ALIAS:-spark-x2.5-1.7b}"
6
+ CTX_SIZE="${CTX_SIZE:-32768}"
7
+ THREADS="${THREADS:-2}"
8
+
9
+ /app/llama-server \
10
+ --host 127.0.0.1 \
11
+ --port 8080 \
12
+ --hf-repo "$MODEL_REPO" \
13
+ --alias "$MODEL_ALIAS" \
14
+ --ctx-size "$CTX_SIZE" \
15
+ --parallel 1 \
16
+ --threads "$THREADS" \
17
+ --threads-batch "$THREADS" \
18
+ --batch-size 512 \
19
+ --ubatch-size 256 \
20
+ --cache-type-k q8_0 \
21
+ --cache-type-v q8_0 \
22
+ --flash-attn auto \
23
+ --cont-batching \
24
+ --metrics \
25
+ --no-webui &
26
+
27
+ LLAMA_PID=$!
28
+
29
+ cleanup() {
30
+ kill "$LLAMA_PID" 2>/dev/null || true
31
+ wait "$LLAMA_PID" 2>/dev/null || true
32
+ }
33
+ trap cleanup INT TERM EXIT
34
+
35
+ exec uvicorn gateway:app \
36
+ --app-dir /srv \
37
+ --host 0.0.0.0 \
38
+ --port 7860 \
39
+ --proxy-headers \
40
+ --forwarded-allow-ips='*'