From 1964df4b42a3f47e8285a64fe79f77861fe2b94e Mon Sep 17 00:00:00 2001 From: Marvin Date: Thu, 24 Sep 2026 05:39:46 +0200 Subject: [PATCH] upgraded api server --- api/src/production/server.py | 457 ++++++++++++++++++++-------- api/src/production/start_server.bat | 2 +- api/src/production/start_server.sh | 2 +- 3 files changed, 325 insertions(+), 136 deletions(-) diff --git a/api/src/production/server.py b/api/src/production/server.py index 508c1f7..88b7023 100644 --- a/api/src/production/server.py +++ b/api/src/production/server.py @@ -1,32 +1,62 @@ +import asyncio +from collections import defaultdict, deque +from concurrent.futures import ThreadPoolExecutor +from contextlib import asynccontextmanager, closing +from datetime import datetime, timezone from enum import Enum +import io +import logging import os +from pathlib import Path +import secrets import sqlite3 import time -import cv2 from dotenv import load_dotenv -import numpy as np -from pydantic import BaseModel -import torch -from torchvision import transforms -from fastapi import Depends, FastAPI, File, HTTPException, Header, UploadFile, WebSocket, WebSocketDisconnect +from fastapi import Depends, FastAPI, File, Header, HTTPException, Query, Request, UploadFile, WebSocket, WebSocketDisconnect from fastapi.middleware.cors import CORSMiddleware from PIL import Image -import io +from pydantic import BaseModel, Field +import torch import torch.nn.functional as F - -import threading +from torchvision import transforms from yolo_model_production import convert_prediction, Yolo_model from bird_cnn_production import Bird_CNN -BUILD_PATH = "./build_models" +logger = logging.getLogger("server") + +BASE_DIR = Path(__file__).resolve().parent +BUILD_PATH = Path(os.getenv("BUILD_PATH", BASE_DIR / "build_models")) +DATABASE_DIR = Path(os.getenv("DATABASE_DIR", BASE_DIR / "database")) +DATABASE = DATABASE_DIR / "reviews.db" + +load_dotenv(DATABASE_DIR / ".env") +API_KEY = os.getenv("API_KEY") # if missing, protected endpoints reject every request + IMAGE_SIZE_CNN = 64 IMAGE_SIZE_YOLO = 64 -sem_ai = threading.Semaphore(1) +# ---- Resource limits (the server is weak, keep everything small and bounded) ---- +TORCH_THREADS = int(os.getenv("TORCH_THREADS", "1")) +MAX_PENDING_INFERENCES = int(os.getenv("MAX_PENDING_INFERENCES", "2")) # running + waiting +MAX_UPLOAD_BYTES = 5 * 1024 * 1024 +MAX_FRAME_BYTES = 1 * 1024 * 1024 +MAX_IMAGE_PIXELS = 20_000_000 +MAX_WEBSOCKETS = 3 +MAX_WEBSOCKETS_PER_IP = 1 +FACE_MIN_INTERVAL = 1 / 2.1 +DECODE_DRAFT_SIZE = (256, 256) # JPEG decodes at reduced scale, still larger than the 64px model input -class bird_species(Enum): +torch.set_num_threads(TORCH_THREADS) +Image.MAX_IMAGE_PIXELS = MAX_IMAGE_PIXELS + +origins = [ + "https://marvinkrausser.com", +] + + +class BirdSpecies(Enum): Common_Kingfisher = 0 Common_Myna = 1 House_Crow = 2 @@ -35,6 +65,7 @@ class bird_species(Enum): Ruddy_Shelduck = 5 Sarus_Crane = 6 + transform_bird = transforms.Compose([ transforms.Resize(IMAGE_SIZE_CNN), transforms.CenterCrop(IMAGE_SIZE_CNN), @@ -47,170 +78,328 @@ transform_face = transforms.Compose([ ]) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") +models = {} -model_bird = Bird_CNN(c_in=3, c_hidden=16, c_out=7) -full_path = os.path.join(BUILD_PATH, "bird_cnn") -model_bird.load_state_dict(torch.load(full_path, map_location=torch.device(device))) -model_bird.to(device) -model_bird.eval() -modeL_face = Yolo_model(c_in=3, boxes=1, grid=6, labels=1, c_hidden=16) -full_path = os.path.join(BUILD_PATH, "face_detection_yolo") -modeL_face.load_state_dict(torch.load(full_path, map_location=torch.device(device))) -modeL_face.to(device) -modeL_face.eval() +def load_models(): + model_bird = Bird_CNN(c_in=3, c_hidden=16, c_out=7) + model_bird.load_state_dict(torch.load(BUILD_PATH / "bird_cnn", map_location=device, weights_only=True)) + model_bird.to(device).eval() -app = FastAPI() + model_face = Yolo_model(c_in=3, boxes=1, grid=6, labels=1, c_hidden=16) + model_face.load_state_dict(torch.load(BUILD_PATH / "face_detection_yolo", map_location=device, weights_only=True)) + model_face.to(device).eval() -origins = [ - "https://marvinkrausser.com", - "https://api.marvinkrausser.com", -] + models["bird"] = model_bird + models["face"] = model_face + + +def init_db(): + DATABASE_DIR.mkdir(parents=True, exist_ok=True) + with closing(sqlite3.connect(DATABASE)) as conn: + conn.execute(""" + CREATE TABLE IF NOT EXISTS reviews ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + website TEXT NOT NULL, + rating INTEGER NOT NULL, + text TEXT NOT NULL, + created_at TEXT NOT NULL + ) + """) + conn.commit() + + +@asynccontextmanager +async def lifespan(app: FastAPI): + logging.basicConfig(level=logging.INFO) + if not API_KEY: + logger.warning("API_KEY is not set: GET/DELETE /review are disabled") + load_models() + init_db() + yield + + +app = FastAPI(lifespan=lifespan) app.add_middleware( CORSMiddleware, allow_origins=origins, - allow_credentials=True, - allow_methods=["*"], - allow_headers=["*"], + allow_methods=["GET", "POST", "DELETE"], + allow_headers=["Authorization", "Content-Type"], ) -@app.post("/predict") + +# -------------------------------------------------------------------------- +# Overload protection +# -------------------------------------------------------------------------- + +def client_ip(conn: Request | WebSocket) -> str: + # Caddy is the only proxy in front of this container and appends the real + # client address as the last X-Forwarded-For entry. + forwarded = conn.headers.get("x-forwarded-for") + if forwarded: + return forwarded.split(",")[-1].strip() + return conn.client.host if conn.client else "unknown" + + +class RateLimiter: + """Sliding-window limiter per key (in memory, single process).""" + + def __init__(self, limit: int, window: float): + self.limit = limit + self.window = window + self.hits = defaultdict(deque) + + def allow(self, key: str) -> bool: + now = time.monotonic() + if len(self.hits) > 10_000: + self.hits = defaultdict(deque, { + k: v for k, v in self.hits.items() if v and now - v[-1] < self.window + }) + q = self.hits[key] + while q and now - q[0] >= self.window: + q.popleft() + if len(q) >= self.limit: + return False + q.append(now) + return True + + +predict_limiter = RateLimiter(limit=10, window=60) +review_limiter = RateLimiter(limit=5, window=60) + + +def rate_limit(limiter: RateLimiter): + def dependency(request: Request): + if not limiter.allow(client_ip(request)): + raise HTTPException(status_code=429, detail="Too many requests", headers={"Retry-After": "60"}) + return dependency + + +class Busy(Exception): + pass + + +class InferenceGate: + """One worker thread runs all inference. At most `max_pending` jobs + (running + waiting) are admitted; everything else is rejected at once + instead of queueing up and eating memory.""" + + def __init__(self, max_pending: int): + self.max_pending = max_pending + self.pending = 0 + self.executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="inference") + + async def run(self, fn, *args): + if self.pending >= self.max_pending: + raise Busy + self.pending += 1 + try: + return await asyncio.get_running_loop().run_in_executor(self.executor, fn, *args) + finally: + self.pending -= 1 + + +gate = InferenceGate(MAX_PENDING_INFERENCES) + + +class InvalidImage(Exception): + pass + + +def open_image(data: bytes) -> tuple[Image.Image, tuple[int, int]]: + """Open lazily, reject oversized images before decoding, decode JPEGs at + reduced scale. Returns the image and its ORIGINAL (width, height).""" + try: + image = Image.open(io.BytesIO(data)) + original_size = image.size + if original_size[0] * original_size[1] > MAX_IMAGE_PIXELS: + raise InvalidImage("Image too large") + image.draft("RGB", DECODE_DRAFT_SIZE) + return image.convert("RGB"), original_size + except InvalidImage: + raise + except Exception as e: # PIL raises many different types for bad data + raise InvalidImage("Not a valid image") from e + + +# -------------------------------------------------------------------------- +# Bird classification +# -------------------------------------------------------------------------- + +def classify_bird(data: bytes) -> dict: + image, _ = open_image(data) + tensor = transform_bird(image).unsqueeze(0).to(device) + with torch.inference_mode(): + probs = F.softmax(models["bird"](tensor), dim=1) + confidence, cls = torch.max(probs, dim=1) + return { + "class": BirdSpecies(cls.item()).name, + "confidence": confidence.item() + } + + +async def read_limited(file: UploadFile, limit: int) -> bytes: + data = bytearray() + while chunk := await file.read(64 * 1024): + data += chunk + if len(data) > limit: + raise HTTPException(status_code=413, detail="File too large") + return bytes(data) + + +@app.post("/predict", dependencies=[Depends(rate_limit(predict_limiter))]) async def predict(file: UploadFile = File(...)): - with sem_ai: - image_bytes = await file.read() + data = await read_limited(file, MAX_UPLOAD_BYTES) + try: + return await gate.run(classify_bird, data) + except Busy: + raise HTTPException(status_code=503, detail="Server busy, try again shortly", headers={"Retry-After": "5"}) + except InvalidImage as e: + raise HTTPException(status_code=400, detail=str(e)) - image = Image.open(io.BytesIO(image_bytes)).convert("RGB") - image = transform_bird(image).unsqueeze(0).to(device) - with torch.no_grad(): - pred = model_bird(image) - probs = F.softmax(pred, dim=1) - confidence, cls = torch.max(probs, dim=1) +# -------------------------------------------------------------------------- +# Face detection (websocket) +# -------------------------------------------------------------------------- + +def detect_faces(data: bytes) -> list: + image, (W, H) = open_image(data) + scale_w = W / IMAGE_SIZE_YOLO + scale_h = H / IMAGE_SIZE_YOLO + + tensor = transform_face(image).unsqueeze(0).to(device) + with torch.inference_mode(): + pred = models["face"](tensor) + bboxes, _, _ = convert_prediction(pred.squeeze(0), tensor.squeeze(0), threshold=0.9) + + return [ + [int(b[0] * scale_w), int(b[1] * scale_h), int(b[2] * scale_w), int(b[3] * scale_h)] + for b in bboxes + ] + + +active_websockets = defaultdict(int) + - return { - "class": bird_species(cls.item()).name, - "confidence": confidence.item() - } - @app.websocket("/predict_face") async def predict_face(websocket: WebSocket): - await websocket.accept() + ip = client_ip(websocket) - last = time.time() + # CORS does not apply to websockets, so check the origin ourselves. + if websocket.headers.get("origin") not in origins: + await websocket.close(code=1008) + return + if (sum(active_websockets.values()) >= MAX_WEBSOCKETS + or active_websockets[ip] >= MAX_WEBSOCKETS_PER_IP): + await websocket.close(code=1013) # try again later + return + active_websockets[ip] += 1 try: + await websocket.accept() + last = 0.0 + while True: - jpg_bytes = await websocket.receive_bytes() + message = await websocket.receive() + if message["type"] == "websocket.disconnect": + break - if time.time() - last < 1 / 2.1: + data = message.get("bytes") + if data is None or len(data) > MAX_FRAME_BYTES: + await websocket.close(code=1009 if data else 1003) + break + + # Drop frames that arrive too fast or while the server is busy + # instead of queueing them. + now = time.monotonic() + if now - last < FACE_MIN_INTERVAL: continue + last = now - last = time.time() - with sem_ai: + try: + boxes = await gate.run(detect_faces, data) + except Busy: + continue + except InvalidImage: + await websocket.close(code=1003) + break - np_arr = np.frombuffer(jpg_bytes, np.uint8) - frame = cv2.imdecode(np_arr, cv2.IMREAD_COLOR) - - frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) - image = Image.fromarray(frame) - - H, W, _ = frame.shape - scale_w = W / IMAGE_SIZE_YOLO - scale_h = H / IMAGE_SIZE_YOLO - - image = transform_face(image).unsqueeze(0).to(device) - - with torch.no_grad(): - pred = modeL_face(image) - bboxes, _, _ = convert_prediction( - pred.squeeze(0), - image.squeeze(0), - threshold=0.9 - ) - - boxes_to_send = [] - for bbox in bboxes: - xmin = int(bbox[0] * scale_w) - ymin = int(bbox[1] * scale_h) - xmax = int(bbox[2] * scale_w) - ymax = int(bbox[3] * scale_h) - boxes_to_send.append([xmin, ymin, xmax, ymax]) - - await websocket.send_json({"bboxes": boxes_to_send}) + await websocket.send_json({"bboxes": boxes}) except WebSocketDisconnect: - print("Client disconnected") + pass + except Exception: + logger.exception("predict_face failed") + try: + await websocket.close(code=1011) + except Exception: + pass + finally: + active_websockets[ip] -= 1 + if active_websockets[ip] <= 0: + del active_websockets[ip] -load_dotenv("./database/.env") -API_KEY = os.getenv("API_KEY") -DATABASE = "./database/reviews.db" +# -------------------------------------------------------------------------- +# Reviews +# -------------------------------------------------------------------------- def get_api_key(authorization: str = Header(None)): - if authorization != f"Bearer {API_KEY}": + if not API_KEY or not authorization or not secrets.compare_digest( + authorization.encode(), f"Bearer {API_KEY}".encode() + ): raise HTTPException(status_code=401, detail="Unauthorized") return authorization -class Review(BaseModel): - website: str - rating: int - text: str - date: str -sem_db = threading.Semaphore(1) +class Review(BaseModel): + website: str = Field(min_length=1, max_length=100) + rating: int = Field(ge=1, le=5) + text: str = Field(max_length=1000) + # The client also sends a "date"; it is ignored, the server sets the time. + + +def connect_db() -> sqlite3.Connection: + conn = sqlite3.connect(DATABASE, timeout=5) + conn.row_factory = sqlite3.Row + return conn + @app.get("/review") -def get_reviews(auth=Depends(get_api_key)): - with sem_db: - conn = sqlite3.connect( - DATABASE, - check_same_thread=False - ) - cur = conn.cursor() +def get_reviews( + limit: int = Query(100, ge=1, le=500), + offset: int = Query(0, ge=0), + auth=Depends(get_api_key), +): + with closing(connect_db()) as conn: + rows = conn.execute( + "SELECT * FROM reviews ORDER BY id DESC LIMIT ? OFFSET ?", (limit, offset) + ).fetchall() + return {"data": [dict(row) for row in rows]} - cur.execute(""" - SELECT * FROM reviews; - """) - data = cur.fetchall() - - print(data) - - conn.close() - return {"data": data} - -@app.post("/review") +# Public on purpose: no credentials needed, only rate limited per IP. +@app.post("/review", dependencies=[Depends(rate_limit(review_limiter))]) def post_review(review: Review): - with sem_db: - conn = sqlite3.connect( - DATABASE, - check_same_thread=False + created_at = datetime.now(timezone.utc).isoformat() + with closing(connect_db()) as conn: + conn.execute( + "INSERT INTO reviews (website, rating, text, created_at) VALUES (?, ?, ?, ?)", + (review.website, review.rating, review.text.strip(), created_at), ) - cur = conn.cursor() - - cur.execute(""" - INSERT INTO reviews (website, rating, text, created_at) - VALUES (?, ?, ?, ?) - """, (review.website, review.rating, review.text, review.date)) - conn.commit() - conn.close() - return {"status": "ok"} + return {"status": "ok"} + @app.delete("/review") def delete_review(auth=Depends(get_api_key)): - with sem_db: - conn = sqlite3.connect( - DATABASE, - check_same_thread=False - ) - cur = conn.cursor() - - cur.execute(""" - DELETE FROM reviews; - """) - + with closing(connect_db()) as conn: + conn.execute("DELETE FROM reviews") conn.commit() - conn.close() - return {"status": "ok"} \ No newline at end of file + return {"status": "ok"} + + +@app.get("/health") +def health(): + return {"status": "ok"} diff --git a/api/src/production/start_server.bat b/api/src/production/start_server.bat index dc9811c..6f52117 100644 --- a/api/src/production/start_server.bat +++ b/api/src/production/start_server.bat @@ -1 +1 @@ -python -m uvicorn server:app --reload \ No newline at end of file +python -m uvicorn server:app --ws-max-size 1048576 --reload \ No newline at end of file diff --git a/api/src/production/start_server.sh b/api/src/production/start_server.sh index bdd00bb..484e4d1 100644 --- a/api/src/production/start_server.sh +++ b/api/src/production/start_server.sh @@ -1 +1 @@ -uvicorn server:app --host 0.0.0.0 --port 8000 --reload \ No newline at end of file +uvicorn server:app --host 0.0.0.0 --port 8000 --ws-max-size 1048576