xxx
This commit is contained in:
@@ -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}")
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user