232 lines
7.2 KiB
Python
232 lines
7.2 KiB
Python
from fastapi import FastAPI
|
|
from fastapi.middleware.cors import CORSMiddleware
|
|
from contextlib import asynccontextmanager
|
|
from sqlalchemy import text
|
|
import redis.asyncio as redis
|
|
import threading
|
|
import subprocess
|
|
import sys
|
|
import os
|
|
|
|
from app.config import settings
|
|
from app.database import engine
|
|
from app.routes import router
|
|
from app.logger import get_logger, LoggerManager
|
|
from app.celery_app import celery_app
|
|
|
|
|
|
# 初始化日志系统
|
|
LoggerManager.setup_logging()
|
|
logger = get_logger("main")
|
|
|
|
# Celery Worker 进程管理
|
|
celery_worker_process = None
|
|
|
|
|
|
def start_celery_worker():
|
|
"""启动 Celery Worker"""
|
|
global celery_worker_process
|
|
try:
|
|
logger.info("🌿 启动Celery Worker...")
|
|
|
|
# 启动Celery worker
|
|
celery_app.start([
|
|
'worker',
|
|
'--loglevel=info',
|
|
'--concurrency=4',
|
|
'--prefetch-multiplier=1',
|
|
'--max-tasks-per-child=1000',
|
|
'--time-limit=300', # 5分钟任务超时
|
|
'--soft-time-limit=240', # 4分钟软超时
|
|
])
|
|
except Exception as e:
|
|
logger.error(f"❌ Celery Worker 启动失败: {e}")
|
|
|
|
|
|
def start_celery_beat():
|
|
"""启动 Celery Beat 调度器"""
|
|
try:
|
|
logger.info("📅 启动Celery Beat调度器...")
|
|
|
|
# 启动Celery beat
|
|
celery_app.start([
|
|
'beat',
|
|
'--loglevel=info',
|
|
'--schedule=/tmp/celerybeat-schedule',
|
|
])
|
|
except Exception as e:
|
|
logger.error(f"❌ Celery Beat 启动失败: {e}")
|
|
|
|
|
|
@asynccontextmanager
|
|
async def lifespan(app: FastAPI):
|
|
# 启动时初始化
|
|
logger.info("🚀 应用启动中...")
|
|
|
|
# Redis 连接对象
|
|
redis_client = None
|
|
|
|
try:
|
|
# 验证数据库连接
|
|
logger.info("📊 验证数据库连接...")
|
|
try:
|
|
async with engine.begin() as conn:
|
|
# 执行简单的查询来验证连接
|
|
result = await conn.execute(text("SELECT 1"))
|
|
connection_test = result.scalar()
|
|
if connection_test == 1:
|
|
logger.info("✅ 数据库连接验证成功")
|
|
else:
|
|
raise Exception("数据库连接测试失败")
|
|
except Exception as db_error:
|
|
logger.error(f"❌ 数据库连接验证失败: {db_error}")
|
|
raise
|
|
|
|
# 验证 Redis 连接
|
|
logger.info("🔴 验证 Redis 连接...")
|
|
try:
|
|
redis_client = redis.from_url(settings.celery_broker_url)
|
|
# 执行 ping 命令验证连接
|
|
await redis_client.ping()
|
|
logger.info("✅ Redis 连接验证成功")
|
|
|
|
# 存储到应用状态中供其他组件使用
|
|
app.state.redis_client = redis_client
|
|
|
|
except Exception as redis_error:
|
|
logger.error(f"❌ Redis 连接验证失败: {redis_error}")
|
|
raise
|
|
|
|
# 检查启动模式
|
|
if len(sys.argv) > 1:
|
|
mode = sys.argv[1].replace("--mode=", "")
|
|
if mode == "all":
|
|
# 启动 Celery Worker (在后台线程中)
|
|
logger.info("🌿 启动Celery Worker...")
|
|
celery_worker_thread = threading.Thread(target=start_celery_worker, daemon=True)
|
|
celery_worker_thread.start()
|
|
|
|
# 启动 Celery Beat 调度器 (在后台线程中)
|
|
logger.info("📅 启动Celery Beat调度器...")
|
|
celery_beat_thread = threading.Thread(target=start_celery_beat, daemon=True)
|
|
celery_beat_thread.start()
|
|
|
|
logger.info(f"🎉 {settings.app_name} 启动完成!")
|
|
yield
|
|
|
|
except Exception as e:
|
|
logger.error(f"❌ 应用启动失败: {e}")
|
|
raise
|
|
|
|
finally:
|
|
# 关闭时清理
|
|
logger.info("🛑 应用关闭中...")
|
|
|
|
# 关闭 Redis 连接
|
|
if redis_client:
|
|
try:
|
|
await redis_client.close()
|
|
logger.info("🔴 Redis 连接已关闭")
|
|
except Exception as e:
|
|
logger.warning(f"⚠️ 关闭 Redis 连接时出现警告: {e}")
|
|
|
|
logger.info("👋 应用已关闭")
|
|
|
|
|
|
app = FastAPI(
|
|
title=settings.app_name,
|
|
version="1.0.0",
|
|
lifespan=lifespan,
|
|
docs_url=(
|
|
"/docs"
|
|
if not settings.disable_docs and settings.environment != "production"
|
|
else None
|
|
),
|
|
redoc_url=(
|
|
"/redoc"
|
|
if not settings.disable_docs and settings.environment != "production"
|
|
else None
|
|
),
|
|
openapi_url="/openapi.json" if not settings.disable_docs else None,
|
|
)
|
|
|
|
# 添加CORS中间件
|
|
app.add_middleware(
|
|
CORSMiddleware,
|
|
allow_origins=["*"],
|
|
allow_credentials=True,
|
|
allow_methods=["*"],
|
|
allow_headers=["*"],
|
|
)
|
|
|
|
# 注册路由
|
|
app.include_router(router)
|
|
|
|
|
|
@app.get("/")
|
|
async def root():
|
|
# 检查环境,生产环境下禁用根接口
|
|
if settings.environment == "production":
|
|
logger.warning("🚫 生产环境下禁止访问根接口")
|
|
from fastapi import HTTPException
|
|
|
|
raise HTTPException(status_code=404, detail="Not Found")
|
|
|
|
logger.info("📝 根接口被访问")
|
|
return {"message": f"Welcome to {settings.app_name}"}
|
|
|
|
|
|
@app.get("/health")
|
|
async def health_check():
|
|
# 检查环境,生产环境下禁用健康检查接口
|
|
if settings.environment == "production":
|
|
logger.warning("🚫 生产环境下禁止访问健康检查接口")
|
|
from fastapi import HTTPException
|
|
|
|
raise HTTPException(status_code=404, detail="Not Found")
|
|
|
|
logger.debug("💓 健康检查接口被访问")
|
|
|
|
health_status = {"status": "healthy"}
|
|
|
|
# 检查 Redis 连接状态
|
|
try:
|
|
redis_client = getattr(app.state, 'redis_client', None)
|
|
if redis_client:
|
|
await redis_client.ping()
|
|
health_status["redis"] = "connected"
|
|
else:
|
|
health_status["redis"] = "disconnected"
|
|
except Exception as e:
|
|
health_status["redis"] = f"error: {str(e)}"
|
|
health_status["status"] = "degraded"
|
|
|
|
return health_status
|
|
|
|
|
|
if __name__ == "__main__":
|
|
import uvicorn
|
|
import argparse
|
|
|
|
parser = argparse.ArgumentParser(description="AI Talk Callback API")
|
|
parser.add_argument("--mode", choices=["api", "worker", "beat", "all"],
|
|
default="api", help="启动模式: api(仅API), worker(仅Celery Worker), beat(仅Celery Beat), all(全部)")
|
|
args = parser.parse_args()
|
|
|
|
if args.mode == "api":
|
|
# 仅启动 FastAPI 应用
|
|
logger.info("🚀 启动FastAPI应用...")
|
|
uvicorn.run("main:app", host="0.0.0.0", port=8000, reload=settings.debug)
|
|
elif args.mode == "worker":
|
|
# 仅启动 Celery Worker
|
|
logger.info("🌿 启动Celery Worker...")
|
|
start_celery_worker()
|
|
elif args.mode == "beat":
|
|
# 仅启动 Celery Beat
|
|
logger.info("📅 启动Celery Beat调度器...")
|
|
start_celery_beat()
|
|
elif args.mode == "all":
|
|
# 启动 FastAPI + Celery Worker + Celery Beat
|
|
logger.info("🚀 启动完整服务栈(FastAPI + Celery Worker + Celery Beat)...")
|
|
uvicorn.run("main:app", host="0.0.0.0", port=8000, reload=settings.debug)
|