Compare commits
3
Commits
b2b43cb8cb
..
0.0.2
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b56b899595 | ||
|
|
716295b7fd | ||
|
|
9115c05573 |
@@ -3,8 +3,8 @@ import aiohttp
|
|||||||
from fastapi import FastAPI, HTTPException, Request, UploadFile
|
from fastapi import FastAPI, HTTPException, Request, UploadFile
|
||||||
from fastapi.responses import FileResponse
|
from fastapi.responses import FileResponse
|
||||||
from fastapi import Depends, FastAPI
|
from fastapi import Depends, FastAPI
|
||||||
from fastapi.security import OAuth2PasswordBearer
|
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||||
from pyrate_limiter import Duration, Limiter, Rate
|
from pyrate_limiter import AbstractBucket, BucketFactory, Duration, InMemoryBucket, Limiter, MonotonicClock, Rate, RateItem
|
||||||
from fastapi_limiter.depends import RateLimiter
|
from fastapi_limiter.depends import RateLimiter
|
||||||
from fastapi import Depends
|
from fastapi import Depends
|
||||||
import asyncmy
|
import asyncmy
|
||||||
@@ -33,7 +33,7 @@ logging.basicConfig(
|
|||||||
handlers=[file_handler, console_handler],
|
handlers=[file_handler, console_handler],
|
||||||
)
|
)
|
||||||
|
|
||||||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
|
security = HTTPBearer()
|
||||||
SECRET = os.getenv("SECRET")
|
SECRET = os.getenv("SECRET")
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
@@ -60,8 +60,27 @@ async def connect_db(app: FastAPI):
|
|||||||
app.state.pool.close()
|
app.state.pool.close()
|
||||||
await app.state.pool.wait_closed()
|
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)
|
app = FastAPI(lifespan=connect_db)
|
||||||
limiter = Limiter(Rate(50, Duration.MINUTE))
|
rates = [Rate(50, Duration.MINUTE)]
|
||||||
|
limiter = Limiter(MultiBucketFactory(rates,MonotonicClock()))
|
||||||
|
|
||||||
async def create_tables(pool):
|
async def create_tables(pool):
|
||||||
async with pool.acquire() as conn:
|
async with pool.acquire() as conn:
|
||||||
@@ -123,13 +142,13 @@ async def fetch_images():
|
|||||||
await asyncio.sleep(86400) # Sleep for 24 hours
|
await asyncio.sleep(86400) # Sleep for 24 hours
|
||||||
|
|
||||||
@app.post("/upload", dependencies=[Depends(RateLimiter(limiter=limiter))])
|
@app.post("/upload", dependencies=[Depends(RateLimiter(limiter=limiter))])
|
||||||
async def upload_image(file: UploadFile, request: Request, token: str = Depends(oauth2_scheme)):
|
async def upload_image(file: UploadFile, request: Request, credentials: HTTPAuthorizationCredentials = Depends(security)):
|
||||||
# if SECRET isn't set, return Unauthrorized
|
# if SECRET isn't set, return Unauthrorized
|
||||||
if token != SECRET or token is None:
|
if credentials.credentials != SECRET or credentials.credentials is None:
|
||||||
raise HTTPException(status_code=401, detail="Unauthorized")
|
raise HTTPException(status_code=401, detail="Unauthorized")
|
||||||
if not file:
|
if not file:
|
||||||
raise HTTPException(status_code=400, detail="No file uploaded")
|
raise HTTPException(status_code=400, detail="No file uploaded")
|
||||||
if not file.filename.lower().endswith(('.jpg', '.jpeg', '.png', '.gif')):
|
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")
|
raise HTTPException(status_code=400, detail="Unsupported file type")
|
||||||
try:
|
try:
|
||||||
content = await file.read()
|
content = await file.read()
|
||||||
|
|||||||
Reference in New Issue
Block a user