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)