First commit
This commit is contained in:
@@ -0,0 +1,432 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import sqlite3
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import aiohttp
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, File, Form, HTTPException, UploadFile
|
||||
from fastapi.responses import JSONResponse
|
||||
from llama_cpp import Llama
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# Configuration (env overrides)
|
||||
# ------------------------------------------------------------
|
||||
MODEL_PATH = os.getenv("MODEL_PATH", "models/gemma-3-1b-it-Q4_K_M.gguf")
|
||||
N_CTX = int(os.getenv("N_CTX", "2048"))
|
||||
N_THREADS = int(os.getenv("N_THREADS", "4"))
|
||||
ORANGE_PI_PORT = int(os.getenv("ORANGE_PI_PORT", "8080")) # default port
|
||||
UPLOAD_TIMEOUT = float(os.getenv("UPLOAD_TIMEOUT", "300.0")) # for uploads
|
||||
REQUEST_TIMEOUT = float(os.getenv("REQUEST_TIMEOUT", "30.0")) # for other requests
|
||||
DB_PATH = os.getenv("DB_PATH", "orangepi.db")
|
||||
|
||||
SYNTHESIS_TEMP = float(os.getenv("SYNTHESIS_TEMP", "0.5"))
|
||||
SYNTHESIS_MAX_TOKENS = int(os.getenv("SYNTHESIS_MAX_TOKENS", "512"))
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# Logging
|
||||
# ------------------------------------------------------------
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger("orchestrator")
|
||||
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# Database helpers
|
||||
# ------------------------------------------------------------
|
||||
def init_db():
|
||||
conn = sqlite3.connect(DB_PATH)
|
||||
c = conn.cursor()
|
||||
c.execute("""
|
||||
CREATE TABLE IF NOT EXISTS orangepi_devices (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
ip TEXT NOT NULL UNIQUE,
|
||||
port INTEGER DEFAULT 8080,
|
||||
enabled INTEGER DEFAULT 1,
|
||||
name TEXT
|
||||
)
|
||||
""")
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
|
||||
def get_all_enabled_ips() -> List[Dict[str, Any]]:
|
||||
conn = sqlite3.connect(DB_PATH)
|
||||
c = conn.cursor()
|
||||
c.execute("SELECT ip, port FROM orangepi_devices WHERE enabled = 1")
|
||||
rows = c.fetchall()
|
||||
conn.close()
|
||||
return [{"ip": row[0], "port": row[1]} for row in rows]
|
||||
|
||||
|
||||
def get_device(ip: str) -> Optional[Dict[str, Any]]:
|
||||
conn = sqlite3.connect(DB_PATH)
|
||||
c = conn.cursor()
|
||||
c.execute(
|
||||
"SELECT ip, port FROM orangepi_devices WHERE ip = ? AND enabled = 1", (ip,)
|
||||
)
|
||||
row = c.fetchone()
|
||||
conn.close()
|
||||
if row:
|
||||
return {"ip": row[0], "port": row[1]}
|
||||
return None
|
||||
|
||||
|
||||
def add_or_update_device(ip: str, port: int = None, name: str = None):
|
||||
if port is None:
|
||||
port = ORANGE_PI_PORT
|
||||
conn = sqlite3.connect(DB_PATH)
|
||||
c = conn.cursor()
|
||||
c.execute(
|
||||
"""
|
||||
INSERT INTO orangepi_devices (ip, port, name)
|
||||
VALUES (?, ?, ?)
|
||||
ON CONFLICT(ip) DO UPDATE SET
|
||||
port = excluded.port,
|
||||
name = excluded.name,
|
||||
enabled = 1
|
||||
""",
|
||||
(ip, port, name),
|
||||
)
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
|
||||
def remove_device(ip: str):
|
||||
conn = sqlite3.connect(DB_PATH)
|
||||
c = conn.cursor()
|
||||
c.execute("UPDATE orangepi_devices SET enabled = 0 WHERE ip = ?", (ip,))
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# Global model
|
||||
# ------------------------------------------------------------
|
||||
llm_model = None
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
global llm_model
|
||||
init_db()
|
||||
logger.info(f"Loading GGUF model from {MODEL_PATH} ...")
|
||||
try:
|
||||
llm_model = Llama(
|
||||
model_path=MODEL_PATH,
|
||||
n_ctx=N_CTX,
|
||||
n_threads=N_THREADS,
|
||||
verbose=False,
|
||||
)
|
||||
logger.info("Model loaded successfully.")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load model: {e}")
|
||||
raise RuntimeError("Model loading failed")
|
||||
yield
|
||||
del llm_model
|
||||
|
||||
|
||||
app = FastAPI(
|
||||
title="OrangePi Coordinator with File Upload & Document Listing",
|
||||
description="Orchestrate multiple Orange Pi RAG services: ask questions, upload files, list documents.",
|
||||
version="2.0",
|
||||
lifespan=lifespan,
|
||||
)
|
||||
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# Request/Response Schemas
|
||||
# ------------------------------------------------------------
|
||||
class AskRequest(BaseModel):
|
||||
question: str
|
||||
|
||||
|
||||
class DeviceRegister(BaseModel):
|
||||
ip: str
|
||||
port: int = Field(ORANGE_PI_PORT, ge=1, le=65535)
|
||||
name: str | None = None
|
||||
|
||||
|
||||
class DeviceRemove(BaseModel):
|
||||
ip: str
|
||||
|
||||
|
||||
class AskResponse(BaseModel):
|
||||
final_answer: str
|
||||
raw_responses: Dict[str, str]
|
||||
synthesis_used: bool = True
|
||||
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# Helper: forward requests to Orange Pi
|
||||
# ------------------------------------------------------------
|
||||
async def forward_to_orange_pi(
|
||||
ip: str,
|
||||
port: int,
|
||||
path: str,
|
||||
method: str = "GET",
|
||||
data: Any = None,
|
||||
files: Any = None,
|
||||
is_upload: bool = False,
|
||||
) -> tuple[dict, int]:
|
||||
url = f"http://{ip}:{port}{path}"
|
||||
timeout = UPLOAD_TIMEOUT if is_upload else REQUEST_TIMEOUT
|
||||
async with aiohttp.ClientSession() as session:
|
||||
try:
|
||||
if method.upper() == "GET":
|
||||
async with session.get(url, timeout=timeout) as resp:
|
||||
text = await resp.text()
|
||||
if resp.content_type == "application/json":
|
||||
return await resp.json(), resp.status
|
||||
else:
|
||||
return {
|
||||
"error": f"Non-JSON response: {text[:200]}"
|
||||
}, resp.status
|
||||
elif method.upper() == "POST":
|
||||
if files:
|
||||
form_data = aiohttp.FormData()
|
||||
for key, file_obj in files.items():
|
||||
# Stream the file without reading entirely into memory
|
||||
# file_obj is starlette.datastructures.UploadFile
|
||||
# We'll read in chunks, but aiohttp can handle streaming if we pass a file-like object
|
||||
# However, file_obj.file is a SpooledTemporaryFile; we can pass it directly
|
||||
form_data.add_field(
|
||||
key,
|
||||
file_obj.file,
|
||||
filename=file_obj.filename,
|
||||
content_type=file_obj.content_type
|
||||
or "application/octet-stream",
|
||||
)
|
||||
async with session.post(
|
||||
url, data=form_data, timeout=timeout
|
||||
) as resp:
|
||||
return await resp.json(), resp.status
|
||||
else:
|
||||
async with session.post(url, json=data, timeout=timeout) as resp:
|
||||
return await resp.json(), resp.status
|
||||
else:
|
||||
return {"error": "Unsupported method"}, 405
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"Timeout forwarding to {ip}:{port}{path}")
|
||||
return {"error": "Request timeout"}, 504
|
||||
except Exception as e:
|
||||
logger.error(f"Error forwarding to {ip}:{port}{path}: {e}", exc_info=True)
|
||||
return {"error": str(e)}, 500
|
||||
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# Existing: query all Orange Pis for an answer
|
||||
# ------------------------------------------------------------
|
||||
async def query_orange_pi_answer(
|
||||
ip: str, port: int, question: str, session: aiohttp.ClientSession
|
||||
) -> tuple[str, str | None, str | None]:
|
||||
url = f"http://{ip}:{port}/answer"
|
||||
payload = {"question": question}
|
||||
try:
|
||||
async with session.post(url, json=payload, timeout=REQUEST_TIMEOUT) as resp:
|
||||
if resp.status == 200:
|
||||
data = await resp.json()
|
||||
answer = data.get("answer", "")
|
||||
return (ip, answer, None)
|
||||
else:
|
||||
return (ip, None, f"HTTP {resp.status}")
|
||||
except asyncio.TimeoutError:
|
||||
return (ip, None, "Timeout")
|
||||
except Exception as e:
|
||||
return (ip, None, str(e))
|
||||
|
||||
|
||||
async def gather_responses(question: str) -> Dict[str, str]:
|
||||
devices = get_all_enabled_ips()
|
||||
if not devices:
|
||||
return {}
|
||||
async with aiohttp.ClientSession() as session:
|
||||
tasks = [
|
||||
query_orange_pi_answer(dev["ip"], dev["port"], question, session)
|
||||
for dev in devices
|
||||
]
|
||||
results = await asyncio.gather(*tasks)
|
||||
response_map = {}
|
||||
for ip, answer, error in results:
|
||||
if answer is not None:
|
||||
response_map[ip] = answer
|
||||
else:
|
||||
response_map[ip] = f"[ERROR: {error}]"
|
||||
return response_map
|
||||
|
||||
|
||||
def synthesize_answer(original_question: str, responses: Dict[str, str]) -> str:
|
||||
if not responses:
|
||||
return "No responses received from any Orange Pi device."
|
||||
responses_text = "\n".join(
|
||||
[f"- Device {ip}: {resp}" for ip, resp in responses.items()]
|
||||
)
|
||||
prompt = f"""You are an aggregator that synthesizes answers from multiple sources.
|
||||
|
||||
Original user question: question{original_question}
|
||||
|
||||
Responses from different Orange Pi devices:
|
||||
{responses_text}
|
||||
|
||||
Please provide a single, coherent, and concise final answer based on all the above responses.
|
||||
If there are conflicts, mention the different viewpoints. If some responses are errors, ignore them.
|
||||
Final answer:
|
||||
"""
|
||||
output = llm_model.create_completion(
|
||||
prompt=prompt,
|
||||
max_tokens=SYNTHESIS_MAX_TOKENS,
|
||||
temperature=SYNTHESIS_TEMP,
|
||||
stop=["\n\n", "User question:", "Original user question:"],
|
||||
echo=False,
|
||||
)
|
||||
return output["choices"][0]["text"].strip()
|
||||
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# NEW: Upload a PDF to a specific Orange Pi
|
||||
# ------------------------------------------------------------
|
||||
@app.post("/devices/{device_ip}/upload")
|
||||
async def upload_to_device(device_ip: str, file: UploadFile = File(...)):
|
||||
"""
|
||||
Upload a PDF file to a specific Orange Pi (by IP address).
|
||||
The Orange Pi must have the RAG service running with /upload endpoint.
|
||||
"""
|
||||
# Check if device is registered and enabled
|
||||
device = get_device(device_ip)
|
||||
if not device:
|
||||
raise HTTPException(
|
||||
status_code=404, detail=f"Device {device_ip} not registered or disabled"
|
||||
)
|
||||
|
||||
result, status = await forward_to_orange_pi(
|
||||
ip=device["ip"],
|
||||
port=device["port"],
|
||||
path="/upload",
|
||||
method="POST",
|
||||
files={"file": file},
|
||||
is_upload=True,
|
||||
)
|
||||
if status != 200:
|
||||
raise HTTPException(
|
||||
status_code=status,
|
||||
detail=result.get("detail", result.get("error", "Upload failed")),
|
||||
)
|
||||
return JSONResponse(content=result)
|
||||
|
||||
|
||||
@app.get("/devices/documents")
|
||||
async def list_all_documents(device_ip: Optional[str] = None):
|
||||
"""Retrieve the list of indexed documents from all registered Orange Pis."""
|
||||
if device_ip:
|
||||
devices = [get_device(device_ip)]
|
||||
if not devices[0]:
|
||||
raise HTTPException(status_code=404, detail=f"Device {device_ip} not found")
|
||||
else:
|
||||
devices = get_all_enabled_ips()
|
||||
if not devices:
|
||||
raise HTTPException(
|
||||
status_code=404, detail="No Orange Pi devices registered"
|
||||
)
|
||||
|
||||
results = {}
|
||||
async with aiohttp.ClientSession() as session:
|
||||
tasks = []
|
||||
for dev in devices:
|
||||
url = f"http://{dev['ip']}:{dev['port']}/documents"
|
||||
tasks.append(fetch_documents(session, dev["ip"], url))
|
||||
outcomes = await asyncio.gather(*tasks)
|
||||
for ip, docs, error in outcomes:
|
||||
if error:
|
||||
results[ip] = {"error": error}
|
||||
else:
|
||||
# Ensure docs is a dict before calling .get()
|
||||
if isinstance(docs, dict):
|
||||
results[ip] = docs.get("documents", [])
|
||||
else:
|
||||
results[ip] = []
|
||||
return {"documents_by_device": results}
|
||||
|
||||
|
||||
async def fetch_documents(session: aiohttp.ClientSession, ip: str, url: str):
|
||||
try:
|
||||
async with session.get(url, timeout=REQUEST_TIMEOUT) as resp:
|
||||
if resp.status == 200:
|
||||
data = await resp.json()
|
||||
return (ip, data, None)
|
||||
else:
|
||||
return (ip, None, f"HTTP {resp.status}")
|
||||
except Exception as e:
|
||||
return (ip, None, str(e))
|
||||
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# Existing endpoints (health, devices/register, devices/remove, devices/list, ask)
|
||||
# ------------------------------------------------------------
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "ok", "model_loaded": llm_model is not None}
|
||||
|
||||
|
||||
@app.post("/devices/register")
|
||||
async def register_device(device: DeviceRegister):
|
||||
add_or_update_device(device.ip, device.port, device.name)
|
||||
return {"message": f"Device {device.ip} registered/updated"}
|
||||
|
||||
|
||||
@app.post("/devices/remove")
|
||||
async def delete_device(device: DeviceRemove):
|
||||
remove_device(device.ip)
|
||||
return {"message": f"Device {device.ip} disabled"}
|
||||
|
||||
|
||||
@app.get("/devices/list")
|
||||
async def list_devices():
|
||||
devices = get_all_enabled_ips()
|
||||
return {"devices": devices}
|
||||
|
||||
|
||||
@app.post("/ask", response_model=AskResponse)
|
||||
async def ask_swarm(request: AskRequest):
|
||||
if llm_model is None:
|
||||
raise HTTPException(status_code=503, detail="Local GGUF model not loaded")
|
||||
responses = await gather_responses(request.question)
|
||||
if not responses:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="No Orange Pi devices available or all requests failed",
|
||||
)
|
||||
final_answer = synthesize_answer(request.question, responses)
|
||||
return AskResponse(
|
||||
final_answer=final_answer, raw_responses=responses, synthesis_used=True
|
||||
)
|
||||
|
||||
@app.delete("/devices/{device_ip}/documents/{doc_id}")
|
||||
async def delete_document_on_device(device_ip: str, doc_id: str):
|
||||
"""Delete a document from a specific Orange Pi."""
|
||||
device = get_device(device_ip)
|
||||
if not device:
|
||||
raise HTTPException(status_code=404, detail=f"Device {device_ip} not registered")
|
||||
|
||||
url = f"http://{device['ip']}:{device['port']}/documents/{doc_id}"
|
||||
async with aiohttp.ClientSession() as session:
|
||||
try:
|
||||
async with session.delete(url, timeout=REQUEST_TIMEOUT) as resp:
|
||||
if resp.status == 200:
|
||||
result = await resp.json()
|
||||
return result
|
||||
else:
|
||||
error_text = await resp.text()
|
||||
raise HTTPException(status_code=resp.status, detail=error_text)
|
||||
except Exception as e:
|
||||
logger.error(f"Error deleting document on {device_ip}: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# Run
|
||||
# ------------------------------------------------------------
|
||||
if __name__ == "__main__":
|
||||
uvicorn.run(
|
||||
"backend:app", host="0.0.0.0", port=8000, reload=False, log_level="info"
|
||||
)
|
||||
Reference in New Issue
Block a user