Files
ai-talk-callback/main.py
2025-12-14 20:20:30 +08:00

321 lines
9.7 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from datetime import datetime
import os
import traceback
import signal
import tempfile
import time
import subprocess
import sys
from app.config import settings
from app.database import engine
from app.routes import router
from app.logger import LoggerManager, get_main_logger
from contextlib import asynccontextmanager
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
import redis.asyncio as redis
from sqlalchemy import text
# 初始化日志系统
LoggerManager.setup_logging()
logger = get_main_logger()
def global_exception_handler(exc_type, exc_value, exc_traceback):
if issubclass(exc_type, KeyboardInterrupt):
sys.__excepthook__(exc_type, exc_value, exc_traceback)
return
# 获取格式化的异常信息
error_msg = "捕获到未处理的异常:\n"
error_msg += f"异常类型: {exc_type.__name__}\n"
error_msg += f"异常信息: {exc_value}\n"
error_msg += "堆栈跟踪:\n"
error_msg += "".join(traceback.format_tb(exc_traceback))
error_msg += f"异常时间: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n"
error_msg += "-" * 50 + "\n"
# 输出到控制台
print(error_msg)
# 输出到文件
logger.error(error_msg)
# 注册全局异常处理器
sys.excepthook = global_exception_handler
@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()