Files
CNN_Website/api/src/server.py
T
2026-05-21 05:49:48 +02:00

151 lines
3.4 KiB
Python

from enum import Enum
import os
import sqlite3
from dotenv import load_dotenv
from pydantic import BaseModel
import torch
from torchvision import transforms
from fastapi import Depends, FastAPI, File, HTTPException, Header, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from PIL import Image
import io
import torch.nn.functional as F
from bird_cnn.bird_cnn import Bird_CNN
import threading
BUILD_PATH = "../build_models"
IMAGE_SIZE = 64
sem_ai = threading.Semaphore(1)
class bird_species(Enum):
Common_Kingfisher = 0
Common_Myna = 1
House_Crow = 2
Indian_Peacock = 3
Indian_Pitta = 4
Ruddy_Shelduck = 5
Sarus_Crane = 6
transform = transforms.Compose([
transforms.Resize(IMAGE_SIZE),
transforms.CenterCrop(IMAGE_SIZE),
transforms.ToTensor()
])
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = Bird_CNN(c_in=3, c_hidden=16, c_out=7)
full_path = os.path.join(BUILD_PATH, "bird_cnn")
model.load_state_dict(torch.load(full_path, map_location=torch.device(device)))
model.to(device)
model.eval()
app = FastAPI()
origins = [
"https://marvinkrausser.com",
"https://api.marvinkrausser.com",
]
app.add_middleware(
CORSMiddleware,
allow_origins=origins,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
@app.post("/predict")
async def predict(file: UploadFile = File(...)):
with sem_ai:
image_bytes = await file.read()
image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
image = transform(image).unsqueeze(0).to(device)
with torch.no_grad():
pred = model(image)
probs = F.softmax(pred, dim=1)
confidence, cls = torch.max(probs, dim=1)
return {
"class": bird_species(cls.item()).name,
"confidence": confidence.item()
}
load_dotenv("../database/.env")
API_KEY = os.getenv("API_KEY")
DATABASE = "../database/reviews.db"
def get_api_key(authorization: str = Header(None)):
if authorization != f"Bearer {API_KEY}":
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)
@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()
cur.execute("""
SELECT * FROM reviews;
""")
data = cur.fetchall()
print(data)
conn.close()
return {"data": data}
@app.post("/review")
def post_review(review: Review):
with sem_db:
conn = sqlite3.connect(
DATABASE,
check_same_thread=False
)
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"}
@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;
""")
conn.commit()
conn.close()
return {"status": "ok"}