Files
bnuy-api/main.py
T
Michelle b56b899595
Build and publish bnuy api / build (push) Successful in 14m0s
Build and publish bnuy api / build (release) Successful in 13m46s
fix rate limit being global instead of per-IP
2026-05-16 22:37:00 +02:00

171 lines
6.6 KiB
Python

import uuid
import aiohttp
from fastapi import FastAPI, HTTPException, Request, UploadFile
from fastapi.responses import FileResponse
from fastapi import Depends, FastAPI
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from pyrate_limiter import AbstractBucket, BucketFactory, Duration, InMemoryBucket, Limiter, MonotonicClock, Rate, RateItem
from fastapi_limiter.depends import RateLimiter
from fastapi import Depends
import asyncmy
import asyncio
import os
from dotenv import load_dotenv
from contextlib import asynccontextmanager
import logging
load_dotenv()
log_level = os.getenv("LOG_LEVEL", "INFO").upper()
log_formatter = logging.Formatter(
fmt="%(asctime)s [%(levelname)s] %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
file_handler = logging.FileHandler("data/api.log", encoding="utf-8")
file_handler.setFormatter(log_formatter)
console_handler = logging.StreamHandler()
console_handler.setFormatter(log_formatter)
logging.basicConfig(
level=getattr(logging, log_level, logging.INFO),
handlers=[file_handler, console_handler],
)
security = HTTPBearer()
SECRET = os.getenv("SECRET")
@asynccontextmanager
async def connect_db(app: FastAPI):
app.state.pool = await asyncmy.create_pool(
host=os.getenv("DB_HOST", "localhost"),
port=int(os.getenv("DB_PORT", 3306)),
user=os.getenv("DB_USER"),
password=os.getenv("DB_PASSWORD"),
db=os.getenv("DB_NAME"),
minsize=5,
maxsize=20
)
await create_tables(app.state.pool)
task = asyncio.create_task(fetch_images())
try:
yield
finally:
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
app.state.pool.close()
await app.state.pool.wait_closed()
class MultiBucketFactory(BucketFactory):
def __init__(self, rates, clock):
self.clock = clock
self.rates = rates
self.buckets = {}
def wrap_item(self, name: str, weight: int = 1) -> RateItem:
"""Time-stamping item, return a RateItem"""
now = self.clock.now()
return RateItem(name, now, weight=weight)
def get(self, item: RateItem) -> AbstractBucket:
if item.name not in self.buckets:
new_bucket = self.create(InMemoryBucket, self.rates)
self.buckets.update({item.name: new_bucket})
return self.buckets[item.name]
app = FastAPI(lifespan=connect_db)
rates = [Rate(50, Duration.MINUTE)]
limiter = Limiter(MultiBucketFactory(rates,MonotonicClock()))
async def create_tables(pool):
async with pool.acquire() as conn:
async with conn.cursor() as cursor:
await cursor.execute("""
CREATE TABLE IF NOT EXISTS images (
id INT AUTO_INCREMENT PRIMARY KEY,
url VARCHAR(255) NOT NULL,
filename VARCHAR(255) NOT NULL,
source VARCHAR(255) NOT NULL,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)
""")
await conn.commit()
@app.get("/")
async def root():
return {"message": "yes the api works, maybe i will create a small landing page later here"}
@app.get("/random", dependencies=[Depends(RateLimiter(limiter=limiter))])
async def get_random_bnuy(request: Request):
async with app.state.pool.acquire() as conn:
async with conn.cursor() as cursor:
await cursor.execute("SELECT filename, source, url FROM images ORDER BY RAND() LIMIT 1;")
result = await cursor.fetchone()
if result:
filepath = os.path.join("data/images", result[0])
if os.path.exists(filepath):
return {"url": f"{request.base_url}images/{result[0]}", "source": result[1], "original_url": result[2]}
else:
raise HTTPException(status_code=404, detail="Image file not found")
else:
raise HTTPException(status_code=404, detail="No available images found")
@app.get("/images/{filename}", dependencies=[Depends(RateLimiter(limiter=limiter))])
async def get_image(filename: str):
async with app.state.pool.acquire() as conn:
async with conn.cursor() as cursor:
await cursor.execute("SELECT filename FROM images WHERE filename = %s", (filename,))
result = await cursor.fetchone()
if result:
filepath = os.path.join("data/images", result[0])
if os.path.exists(filepath):
return FileResponse(filepath)
else:
raise HTTPException(status_code=404, detail="Image file not found")
else:
raise HTTPException(status_code=404, detail="Image not found")
async def fetch_images():
from collector import save_picture
while True:
try:
logging.info("Starting image collection...")
await save_picture(app.state.pool)
logging.info("Image collection completed. Sleeping for 1 hour...")
except Exception as e:
logging.error(f"Error during image collection: {e}")
await asyncio.sleep(86400) # Sleep for 24 hours
@app.post("/upload", dependencies=[Depends(RateLimiter(limiter=limiter))])
async def upload_image(file: UploadFile, request: Request, credentials: HTTPAuthorizationCredentials = Depends(security)):
# if SECRET isn't set, return Unauthrorized
if credentials.credentials != SECRET or credentials.credentials is None:
raise HTTPException(status_code=401, detail="Unauthorized")
if not file:
raise HTTPException(status_code=400, detail="No file uploaded")
if not file.filename.lower().endswith(('.jpg', '.jpeg', '.png', '.gif')) or file.content_type not in ["image/jpeg", "image/png", "image/gif"]:
raise HTTPException(status_code=400, detail="Unsupported file type")
try:
content = await file.read()
generate_filename = str(uuid.uuid4()) + os.path.splitext(file.filename)[1]
filename = os.path.join("data/images", generate_filename)
with open(filename, "wb") as f:
f.write(content)
logging.info(f"Saved uploaded image to {filename}")
async with app.state.pool.acquire() as conn:
async with conn.cursor() as cursor:
await cursor.execute(
"INSERT INTO images (url, filename, source) VALUES (%s, %s, %s)",
(f"{request.base_url}images/{generate_filename}", generate_filename, "user_upload")
)
await conn.commit()
except Exception as e:
logging.error(f"Error saving uploaded image: {e}")
raise HTTPException(status_code=500, detail="Failed to save image")