This commit is contained in:
2026-07-30 10:30:50 +08:00
parent cf6d708566
commit 853c478657
85 changed files with 4711 additions and 1052 deletions
+212
View File
@@ -0,0 +1,212 @@
"""
shared/app_factory.py — FastAPI 应用工厂
将 moldinsight.py / inventory.py 两个入口的重复引导代码
(CORS、日志中间件、startup/shutdown、/health、SPA fallback)
收敛到一个工厂函数,消除漂移风险。
"""
import os
import time
from pathlib import Path
from typing import List, Optional, Callable, Awaitable
from fastapi import FastAPI, Request
from fastapi.staticfiles import StaticFiles
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse, FileResponse
from shared.config.settings import settings
from shared.utils.logger import setup_logging, get_logger, generate_request_id, set_request_id
setup_logging()
logger = get_logger(__name__)
def create_app(
*,
title: str,
service_name: str,
version: str = "4.0.0",
mount_html: bool = False,
startup_hooks: Optional[List[Callable[[], Awaitable[None]]]] = None,
register_routers: Optional[Callable[[FastAPI], None]] = None,
) -> FastAPI:
"""创建标准化的 FastAPI 应用实例。
Args:
title: 应用标题
service_name: 服务名(用于 /health 响应)
version: 版本号
mount_html: 是否挂载 /html 静态目录(moldinsight 需要)
startup_hooks: 额外的 startup 钩子列表(在数据库/RustFS/Redis 初始化后执行)
register_routers: 回调函数,用于注册业务路由
"""
app = FastAPI(title=title, version=version)
# ── CORS 白名单 ──────────────────────────────────────────────
cors_origins = settings.CORS_ORIGINS or ["*"]
if cors_origins == ["*"]:
logger.warning(
"CORS 使用通配符 ['*'],生产环境请设置 CORS_ORIGINS 环境变量"
)
app.add_middleware(
CORSMiddleware,
allow_origins=cors_origins,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# ── 请求日志中间件(结构化 + request_id 追踪)─────────────────
@app.middleware("http")
async def log_requests(request: Request, call_next):
# 生成/提取 request_id
rid = request.headers.get("X-Request-ID") or generate_request_id()
set_request_id(rid)
start_time = time.time()
response = await call_next(request)
duration_ms = round((time.time() - start_time) * 1000, 1)
# 跳过静态资源和健康检查的详细日志
path = request.url.path
is_static = path.startswith("/static") or path == "/health"
if not is_static:
log_level = "warning" if response.status_code >= 400 else "info"
extra = {
"method": request.method,
"path": path,
"status": response.status_code,
"duration_ms": duration_ms,
"client_ip": request.client.host if request.client else "-",
}
getattr(logger, log_level)(
f"{request.method} {path} -> {response.status_code} ({duration_ms}ms)",
extra=extra,
)
# 注入 X-Request-ID 响应头,方便前端/运维追踪
response.headers["X-Request-ID"] = rid
return response
# ── 目录准备 ─────────────────────────────────────────────────
Path("uploads").mkdir(exist_ok=True)
Path("static").mkdir(exist_ok=True)
if mount_html:
Path("html_output").mkdir(exist_ok=True)
# ── 静态文件挂载 ─────────────────────────────────────────────
app.mount(
"/static",
StaticFiles(directory=os.path.join(os.getcwd(), "static")),
name="static",
)
if mount_html:
app.mount(
"/html",
StaticFiles(directory=os.path.join(os.getcwd(), "html_output")),
name="html",
)
# ── Startup ──────────────────────────────────────────────────
@app.on_event("startup")
async def startup_event():
from shared.database.init_db import init_database
success = await init_database(keep_connected=True)
print(f"[{'OK' if success else 'FAIL'}] 数据库初始化")
# RustFS(仅 moldinsight 需要)
if mount_html:
try:
from moldinsight.storage.rustfs_storage import rustfs_manager
await rustfs_manager.connect(
endpoint=settings.RUSTFS_ENDPOINT,
access_key=settings.RUSTFS_ACCESS_KEY,
secret_key=settings.RUSTFS_SECRET_KEY,
timeout=settings.RUSTFS_TIMEOUT,
)
print("[OK] RustFS连接成功")
except Exception as e:
print(f"[WARN] RustFS连接失败: {e}")
# Redis
try:
from shared.services.redis_task_manager import redis_task_manager
await redis_task_manager.connect()
print(f"[{'OK' if redis_task_manager.is_connected else 'WARN'}] Redis")
except Exception as e:
print(f"[WARN] Redis异常: {e}")
# 额外钩子
for hook in (startup_hooks or []):
try:
await hook()
except Exception as e:
print(f"[WARN] startup hook 异常: {e}")
# ── Shutdown ─────────────────────────────────────────────────
@app.on_event("shutdown")
async def shutdown_event():
try:
from shared.services.redis_task_manager import redis_task_manager
await redis_task_manager.disconnect()
except Exception:
pass
# ── 认证路由 ─────────────────────────────────────────────────
from shared.services.auth_routes import router as auth_router
app.include_router(auth_router)
# ── 业务路由注册 ─────────────────────────────────────────────
if register_routers:
register_routers(app)
# ── /health 统一端点 ─────────────────────────────────────────
@app.get("/health")
@app.post("/health")
async def health():
from shared.database.database import db_manager
from sqlalchemy import text
db_ok = False
db_error = None
try:
if not db_manager.is_connected:
await db_manager.connect()
async with db_manager.engine.begin() as conn:
await conn.execute(text("SELECT 1"))
db_ok = True
except Exception as e:
db_error = str(e)
return {
"status": "healthy" if db_ok else "degraded",
"service": service_name,
"version": version,
"database_connected": db_ok,
"database_error": db_error,
}
# ── SPA fallback(排除 /api 前缀,避免吞掉 API 404)────────
@app.get("/{full_path:path}")
async def spa_fallback(full_path: str):
# API 路径不走 SPA fallback,让 FastAPI 正常返回 404 JSON
if full_path.startswith("api/") or full_path.startswith("api"):
raise _api_not_found(full_path)
# 健康检查 / 文档路径也排除
if full_path.startswith("docs") or full_path.startswith("openapi"):
raise _api_not_found(full_path)
static_index = os.path.join(os.getcwd(), "static", "index.html")
if os.path.exists(static_index):
return FileResponse(static_index)
return JSONResponse({"detail": "SPA index not found"}, status_code=404)
return app
def _api_not_found(path: str):
"""为 API 路径生成标准 404 异常"""
from fastapi import HTTPException
raise HTTPException(status_code=404, detail=f"Not Found: /{path}")
+15 -1
View File
@@ -1,6 +1,6 @@
import os
import urllib.parse
from typing import Dict, Any
from typing import Dict, Any, List
from dotenv import load_dotenv
load_dotenv()
@@ -75,6 +75,11 @@ class Settings:
self.REDIS_PASSWORD = os.getenv("REDIS_PASSWORD", "")
self.REDIS_DB = int(os.getenv("REDIS_DB", "0"))
# CORS 白名单(逗号分隔,默认允许本机开发地址)
self.CORS_ORIGINS = self._parse_cors_origins(
os.getenv("CORS_ORIGINS", "")
)
# LLM 增强分析配置(可选)
self.LLM_ENABLED = os.getenv("LLM_ENABLED", "false").lower() == "true"
self.LLM_API_URL = os.getenv("LLM_API_URL", "https://api.openai.com/v1")
@@ -95,5 +100,14 @@ class Settings:
def allowed_extensions_set(self) -> set:
return set(ext.strip() for ext in self.ALLOWED_EXTENSIONS.split(","))
@staticmethod
def _parse_cors_origins(raw: str) -> List[str]:
"""解析 CORS_ORIGINS 环境变量,逗号分隔。
为空时返回空列表(由 app_factory 决定是否降级为 ['*'])。
"""
if not raw or not raw.strip():
return []
return [o.strip().rstrip("/") for o in raw.split(",") if o.strip()]
settings = Settings()
+52 -10
View File
@@ -13,6 +13,24 @@ from shared.utils.logger import get_logger
logger = get_logger(__name__)
def _get_pool_config(role: str = "web") -> dict:
"""按角色返回连接池参数。
web 入口:适中并发;celery worker:少量长连接。
通过 DB_POOL_SIZE / DB_MAX_OVERFLOW 环境变量可覆盖默认值。
"""
defaults = {
"web": {"pool_size": 10, "max_overflow": 20},
"celery": {"pool_size": 5, "max_overflow": 10},
}
role_cfg = defaults.get(role, defaults["web"])
# 允许环境变量覆盖
pool_size = int(os.getenv("DB_POOL_SIZE", str(role_cfg["pool_size"])))
max_overflow = int(os.getenv("DB_MAX_OVERFLOW", str(role_cfg["max_overflow"])))
return {"pool_size": pool_size, "max_overflow": max_overflow}
class DatabaseManager:
"""数据库管理器"""
@@ -21,21 +39,27 @@ class DatabaseManager:
self.async_session = None
self.is_connected = False
async def connect(self):
"""连接数据库"""
async def connect(self, role: str = "web"):
"""连接数据库
Args:
role: 连接角色,"web" 或 "celery",决定连接池大小
"""
if not settings.DATABASE_URL:
logger.warning("未配置数据库连接,跳过数据库初始化")
self.is_connected = False
return
try:
pool_cfg = _get_pool_config(role)
# 创建异步引擎
self.engine = create_async_engine(
settings.DATABASE_URL,
echo=settings.DEBUG,
pool_size=20,
max_overflow=30,
pool_recycle=3600
pool_size=pool_cfg["pool_size"],
max_overflow=pool_cfg["max_overflow"],
pool_recycle=3600,
pool_pre_ping=True, # 自动检测失效连接,避免 PG 断连报错
)
# 创建异步会话工厂
@@ -50,7 +74,10 @@ class DatabaseManager:
await conn.execute(text("SELECT 1"))
self.is_connected = True
logger.info("数据库连接成功")
logger.info(
"数据库连接成功 (pool_size=%d, max_overflow=%d)",
pool_cfg["pool_size"], pool_cfg["max_overflow"],
)
except Exception as e:
logger.error(f"数据库连接失败: {e}")
@@ -66,18 +93,22 @@ class DatabaseManager:
@asynccontextmanager
async def session(self):
"""获取数据库会话的异步上下文管理器"""
"""获取数据库会话的异步上下文管理器(用于后台任务/Celery)"""
if not self.is_connected:
await self.connect()
session = self.async_session()
try:
yield session
await session.commit()
except Exception:
await session.rollback()
raise
finally:
await session.close()
async def get_session(self) -> AsyncSession:
"""获取数据库会话"""
"""获取数据库会话(非上下文管理器,配合 get_db_session 依赖使用)"""
if not self.is_connected:
await self.connect()
@@ -100,9 +131,20 @@ db_manager = DatabaseManager()
# 数据库依赖注入
async def get_db_session():
"""获取数据库会话的依赖函数"""
"""获取数据库会话的依赖函数
统一事务边界:
- 路由正常返回 → 自动 commit
- 路由抛出异常 → 自动 rollback
路由中应使用 flush() 代替 commit(),以便在提交前仍能 refresh()。
"""
session = await db_manager.get_session()
try:
yield session
await session.commit()
except Exception:
await session.rollback()
raise
finally:
await session.close()
await session.close()
+121 -13
View File
@@ -1,17 +1,125 @@
# utils/logger.py
"""
shared/utils/logger.py — 结构化日志 + 请求追踪
功能:
- JSON 结构化日志输出(生产友好,方便 ELK/Loki 采集)
- request_id 自动注入(通过 contextvars,跨 async 传播)
- 向后兼容:get_logger(name) / setup_logging() API 不变
- 支持 LOG_FORMAT 环境变量切换(json / text,默认 json)
"""
import json
import logging
import os
import sys
import uuid
from contextvars import ContextVar
from datetime import datetime, timezone
from typing import Optional
def setup_logging():
"""设置日志配置"""
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
handlers=[
logging.StreamHandler(sys.stdout)
]
)
# ── request_id 上下文变量(跨 async task 自动传播)────────────
request_id_var: ContextVar[Optional[str]] = ContextVar("request_id", default=None)
def get_logger(name: str):
"""获取日志器"""
return logging.getLogger(name)
def generate_request_id() -> str:
"""生成短 request_id(8 位 hex,便于日志阅读)"""
return uuid.uuid4().hex[:8]
def set_request_id(rid: Optional[str]) -> None:
"""设置当前请求的 request_id"""
request_id_var.set(rid)
def get_request_id() -> Optional[str]:
"""获取当前请求的 request_id"""
return request_id_var.get()
# ── JSON 结构化 Formatter ─────────────────────────────────────
class JSONFormatter(logging.Formatter):
"""将日志记录格式化为单行 JSON 字符串。
输出字段:
- timestamp: ISO-8601 UTC 时间戳
- level: 日志级别
- logger: logger 名称
- message: 日志消息
- request_id: 当前请求 ID(如果有)
- module/function/line: 代码位置
- exc_info: 异常信息(如果有)
"""
def format(self, record: logging.LogRecord) -> str:
log_entry = {
"timestamp": datetime.now(timezone.utc).isoformat(),
"level": record.levelname,
"logger": record.name,
"message": record.getMessage(),
"request_id": request_id_var.get(),
"module": record.module,
"func": record.funcName,
"line": record.lineno,
}
if record.exc_info and record.exc_info[0] is not None:
log_entry["exc_info"] = self.formatException(record.exc_info)
# 支持 extra 字段(通过 logger.info("msg", extra={"key": "val"}))
for key in ("method", "path", "status", "duration_ms", "client_ip",
"user_agent", "user_id"):
val = getattr(record, key, None)
if val is not None:
log_entry[key] = val
return json.dumps(log_entry, ensure_ascii=False)
# ── 文本 Formatter(开发环境友好)─────────────────────────────
class TextFormatter(logging.Formatter):
"""带 request_id 的文本格式,适合本地开发阅读。"""
def __init__(self):
super().__init__(
fmt="%(asctime)s [%(levelname)s] %(name)s [rid=%(request_id)s] %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
def format(self, record: logging.LogRecord) -> str:
if not hasattr(record, "request_id"):
record.request_id = request_id_var.get() or "-"
return super().format(record)
# ── 公共 API ──────────────────────────────────────────────────
def setup_logging(level: Optional[str] = None):
"""初始化日志系统。
Args:
level: 日志级别,默认从 LOG_LEVEL 环境变量读取(INFO)
环境变量:
LOG_FORMAT: json(默认)或 text
LOG_LEVEL: 日志级别(DEBUG/INFO/WARNING/ERROR)
"""
log_level = getattr(logging, (level or os.getenv("LOG_LEVEL", "INFO")).upper(), logging.INFO)
log_format = os.getenv("LOG_FORMAT", "json").lower()
formatter = JSONFormatter() if log_format == "json" else TextFormatter()
handler = logging.StreamHandler(sys.stdout)
handler.setFormatter(formatter)
root = logging.getLogger()
root.setLevel(log_level)
# 清除已有 handler 避免重复输出
root.handlers.clear()
root.addHandler(handler)
# 降低第三方库的日志级别
for noisy in ("uvicorn.access", "uvicorn.error", "httpx", "httpcore"):
logging.getLogger(noisy).setLevel(logging.WARNING)
def get_logger(name: str) -> logging.Logger:
"""获取带模块名的 logger(API 不变,向后兼容)。"""
return logging.getLogger(name)