ForcePilot/server/main.py

128 lines
4.1 KiB
Python
Raw Normal View History

import asyncio
import time
from collections import defaultdict, deque
2024-10-02 20:11:28 +08:00
import uvicorn
from fastapi import FastAPI, Request, status
2024-10-02 20:11:28 +08:00
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
2025-05-02 23:56:59 +08:00
from starlette.middleware.base import BaseHTTPMiddleware
2025-04-04 00:16:18 +08:00
from server.routers import router
from server.utils.lifespan import lifespan
from server.utils.auth_middleware import is_public_path
from server.utils.common_utils import setup_logging
from server.utils.access_log_middleware import AccessLogMiddleware
2024-10-02 20:11:28 +08:00
# 设置日志配置
setup_logging()
2024-10-02 20:11:28 +08:00
RATE_LIMIT_MAX_ATTEMPTS = 10
RATE_LIMIT_WINDOW_SECONDS = 60
RATE_LIMIT_ENDPOINTS = {("/api/auth/token", "POST")}
# In-memory login attempt tracker to reduce brute-force exposure per worker
_login_attempts: defaultdict[str, deque[float]] = defaultdict(deque)
_attempt_lock = asyncio.Lock()
app = FastAPI(lifespan=lifespan)
app.include_router(router, prefix="/api")
2024-10-02 20:11:28 +08:00
# CORS 设置
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
def _extract_client_ip(request: Request) -> str:
forwarded_for = request.headers.get("x-forwarded-for")
if forwarded_for:
return forwarded_for.split(",")[0].strip()
if request.client:
return request.client.host
return "unknown"
class LoginRateLimitMiddleware(BaseHTTPMiddleware):
async def dispatch(self, request: Request, call_next):
normalized_path = request.url.path.rstrip("/") or "/"
request_signature = (normalized_path, request.method.upper())
if request_signature in RATE_LIMIT_ENDPOINTS:
client_ip = _extract_client_ip(request)
now = time.monotonic()
async with _attempt_lock:
attempt_history = _login_attempts[client_ip]
while attempt_history and now - attempt_history[0] > RATE_LIMIT_WINDOW_SECONDS:
attempt_history.popleft()
if len(attempt_history) >= RATE_LIMIT_MAX_ATTEMPTS:
retry_after = int(max(1, RATE_LIMIT_WINDOW_SECONDS - (now - attempt_history[0])))
return JSONResponse(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
content={"detail": "登录尝试过于频繁,请稍后再试"},
headers={"Retry-After": str(retry_after)},
)
attempt_history.append(now)
response = await call_next(request)
if response.status_code < 400:
async with _attempt_lock:
_login_attempts.pop(client_ip, None)
return response
return await call_next(request)
2025-05-02 23:56:59 +08:00
# 鉴权中间件
class AuthMiddleware(BaseHTTPMiddleware):
async def dispatch(self, request: Request, call_next):
# 获取请求路径
path = request.url.path
# 检查是否为公开路径,公开路径无需身份验证
if is_public_path(path):
return await call_next(request)
if not path.startswith("/api"):
2025-05-02 23:56:59 +08:00
# 非API路径可能是前端路由或静态资源
return await call_next(request)
# # 提取Authorization头
# auth_header = request.headers.get("Authorization")
# if not auth_header or not auth_header.startswith("Bearer "):
# return JSONResponse(
# status_code=status.HTTP_401_UNAUTHORIZED,
# content={"detail": f"请先登录。Path: {path}"},
# headers={"WWW-Authenticate": "Bearer"}
# )
2025-05-02 23:56:59 +08:00
# # 获取token
# token = auth_header.split("Bearer ")[1]
2025-05-02 23:56:59 +08:00
# # 添加token到请求状态后续路由可以直接使用
# request.state.token = token
2025-05-02 23:56:59 +08:00
# 继续处理请求
return await call_next(request)
# 添加访问日志中间件(记录请求处理时间)
app.add_middleware(AccessLogMiddleware)
2025-05-02 23:56:59 +08:00
# 添加鉴权中间件
app.add_middleware(LoginRateLimitMiddleware)
2025-05-02 23:56:59 +08:00
app.add_middleware(AuthMiddleware)
2024-10-02 20:11:28 +08:00
if __name__ == "__main__":
2025-04-28 22:53:13 +08:00
uvicorn.run(app, host="0.0.0.0", port=5050, threads=10, workers=10, reload=True)