275 lines
8.8 KiB
Python
275 lines
8.8 KiB
Python
import signal
|
||
import tempfile
|
||
import time
|
||
from fastapi import FastAPI
|
||
from fastapi.middleware.cors import CORSMiddleware
|
||
from contextlib import asynccontextmanager
|
||
from sqlalchemy import text
|
||
import redis.asyncio as redis
|
||
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 LoggerManager, get_main_logger
|
||
|
||
|
||
# 初始化日志系统
|
||
LoggerManager.setup_logging()
|
||
logger = get_main_logger()
|
||
|
||
|
||
@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
|
||
|
||
logger.info("✅ 服务初始化完成,FastAPI应用启动中...")
|
||
|
||
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("💓 健康检查接口被访问")
|
||
|
||
return {"status": "healthy"}
|
||
|
||
|
||
# 全局变量存储进程
|
||
processes = []
|
||
|
||
def signal_handler(signum, frame):
|
||
"""信号处理器,用于优雅关闭所有服务"""
|
||
print(f"\n🛑 接收到信号 {signum},正在关闭所有服务...")
|
||
|
||
# 逆序关闭进程(最后启动的最先关闭)
|
||
for i, process in enumerate(reversed(processes)):
|
||
if process and process.poll() is None: # 进程仍在运行
|
||
print(f"🔄 正在关闭进程 {len(processes) - i}...")
|
||
try:
|
||
process.terminate() # 发送 SIGTERM 信号
|
||
process.wait(timeout=10) # 等待最多10秒
|
||
print(f"✅ 进程已关闭")
|
||
except subprocess.TimeoutExpired:
|
||
print(f"⚠️ 进程未在10秒内响应,强制关闭...")
|
||
process.kill() # 强制杀死进程
|
||
except Exception as e:
|
||
print(f"❌ 关闭进程时出错: {e}")
|
||
|
||
print("👋 所有服务已关闭")
|
||
sys.exit(0)
|
||
|
||
def start_all_services():
|
||
"""启动所有服务"""
|
||
global processes
|
||
|
||
print("\n🚀 AI Talk Callback API 一键启动所有服务")
|
||
print("=" * 60)
|
||
|
||
# 注册信号处理器
|
||
signal.signal(signal.SIGINT, signal_handler) # Ctrl+C
|
||
signal.signal(signal.SIGTERM, signal_handler) # 终止信号
|
||
|
||
try:
|
||
# 1. 启动 Celery Worker
|
||
print("🌿 启动 Celery Worker...")
|
||
worker_process = subprocess.Popen([
|
||
sys.executable, "-m", "celery",
|
||
"-A", "app.celery_app",
|
||
"worker",
|
||
'--loglevel=info',
|
||
'--pool=solo',
|
||
'--concurrency=1',
|
||
'--time-limit=300', # 5分钟任务超时
|
||
'--soft-time-limit=240' # 4分钟软超时
|
||
])
|
||
processes.append(worker_process)
|
||
time.sleep(2) # 等待 Worker 启动
|
||
|
||
# 2. 启动 Celery Beat
|
||
print("\n📅 启动 Celery Beat...")
|
||
beat_process = subprocess.Popen([
|
||
sys.executable, "-m", "celery",
|
||
"-A", "app.celery_app",
|
||
"beat",
|
||
'--loglevel=info',
|
||
f'--schedule={os.path.join(tempfile.gettempdir(), "celerybeat-schedule")}'
|
||
])
|
||
processes.append(beat_process)
|
||
time.sleep(2) # 等待 Beat 启动
|
||
|
||
# 3. 启动 Flower 监控(如果启用)
|
||
if settings.flower_enabled:
|
||
print("\n📊 启动 Flower 监控服务...")
|
||
flower_cmd = [
|
||
sys.executable, "-m", "celery",
|
||
"-A", "app.celery_app",
|
||
f"--broker={settings.celery_broker_url}",
|
||
"flower",
|
||
f"--port={settings.flower_port}"
|
||
]
|
||
|
||
if settings.flower_basic_auth:
|
||
flower_cmd.append(f"--basic_auth={settings.flower_basic_auth}")
|
||
if settings.flower_url_prefix:
|
||
flower_cmd.append(f"--url_prefix={settings.flower_url_prefix}")
|
||
|
||
flower_process = subprocess.Popen(flower_cmd)
|
||
processes.append(flower_process)
|
||
time.sleep(2) # 等待 Flower 启动
|
||
|
||
# 4. 启动 FastAPI 应用
|
||
print("\n🚀 启动 FastAPI 应用...")
|
||
api_cmd = [
|
||
sys.executable, "-m", "uvicorn",
|
||
"main:app",
|
||
"--host", "0.0.0.0",
|
||
"--port", "8000"
|
||
]
|
||
|
||
# 添加调试模式(如果配置了)
|
||
if settings.debug:
|
||
api_cmd.append("--reload")
|
||
|
||
api_process = subprocess.Popen(api_cmd)
|
||
processes.append(api_process)
|
||
time.sleep(2) # 等待 API 启动
|
||
|
||
print("\n" + "=" * 60)
|
||
print("✅ 所有服务启动完成!")
|
||
print("\n🌐 服务地址:")
|
||
print(" 🚀 API 服务: http://localhost:8000")
|
||
if settings.environment != "production" and not settings.disable_docs:
|
||
print(" 📖 API 文档: http://localhost:8000/docs")
|
||
if settings.flower_enabled:
|
||
print(f" 📊 监控界面: {settings.flower_url}")
|
||
print("\n💡 使用 Ctrl+C 可以优雅关闭所有服务")
|
||
print("=" * 60)
|
||
|
||
# 等待所有进程
|
||
while True:
|
||
# 检查是否有进程异常退出
|
||
for i, process in enumerate(processes):
|
||
if process and process.poll() is not None:
|
||
print(f"❌ 进程 {i+1} 异常退出,退出码: {process.returncode}")
|
||
signal_handler(signal.SIGINT, None)
|
||
return
|
||
|
||
time.sleep(1) # 每秒检查一次
|
||
|
||
except KeyboardInterrupt:
|
||
signal_handler(signal.SIGINT, None)
|
||
except Exception as e:
|
||
print(f"❌ 启动服务时出错: {e}")
|
||
signal_handler(signal.SIGINT, None)
|
||
|
||
if __name__ == "__main__":
|
||
# 直接启动所有服务
|
||
start_all_services()
|