284 lines
9.8 KiB
Python
284 lines
9.8 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
FastAPI应用创建和配置
|
||
"""
|
||
|
||
import warnings
|
||
import asyncio
|
||
import os
|
||
import logging
|
||
from contextlib import asynccontextmanager
|
||
from fastapi import FastAPI, HTTPException
|
||
from fastapi.exceptions import RequestValidationError
|
||
from fastapi_offline import FastAPIOffline
|
||
from fastapi.middleware.cors import CORSMiddleware
|
||
from fastapi.staticfiles import StaticFiles
|
||
|
||
from .core.config import settings
|
||
from .core.exceptions import (
|
||
APIException,
|
||
api_exception_handler,
|
||
general_exception_handler,
|
||
http_exception_handler,
|
||
validation_exception_handler,
|
||
)
|
||
from .core.logging import setup_logging, get_worker_id
|
||
from .core.executor import shutdown_executor
|
||
from .api.v1 import api_router
|
||
from .utils.boot_events import emit_boot_event
|
||
|
||
# 忽略 Pydantic V2 兼容性警告
|
||
warnings.filterwarnings("ignore", message="Valid config keys have changed in V2")
|
||
warnings.filterwarnings("ignore", message=".*has conflict with protected namespace.*")
|
||
warnings.filterwarnings("ignore", category=UserWarning, module="pydantic")
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
async def cleanup_task_state_loop():
|
||
"""定期清理过期离线任务状态文件。"""
|
||
while True:
|
||
await asyncio.sleep(600)
|
||
try:
|
||
from .core.task_store import cleanup_expired_tasks
|
||
|
||
cleanup_expired_tasks(force=True)
|
||
except asyncio.CancelledError:
|
||
raise
|
||
except Exception as exc:
|
||
logger.warning("定期清理离线任务状态失败: %s", exc)
|
||
|
||
|
||
def cleanup_temp_directory():
|
||
"""清理临时目录中的旧文件"""
|
||
import time
|
||
temp_dir = settings.TEMP_DIR
|
||
if not os.path.exists(temp_dir):
|
||
return
|
||
|
||
# 清理超过 1 小时的临时文件
|
||
max_age_seconds = 3600
|
||
current_time = time.time()
|
||
cleaned_count = 0
|
||
|
||
try:
|
||
for filename in os.listdir(temp_dir):
|
||
filepath = os.path.join(temp_dir, filename)
|
||
if os.path.isfile(filepath):
|
||
file_age = current_time - os.path.getmtime(filepath)
|
||
if file_age > max_age_seconds:
|
||
try:
|
||
os.remove(filepath)
|
||
cleaned_count += 1
|
||
except Exception:
|
||
pass
|
||
|
||
if cleaned_count > 0:
|
||
logger.info(f"已清理 {cleaned_count} 个过期临时文件")
|
||
except Exception as e:
|
||
logger.warning(f"清理临时目录时出错: {e}")
|
||
|
||
|
||
@asynccontextmanager
|
||
async def lifespan(app: FastAPI):
|
||
"""应用生命周期管理"""
|
||
workers = int(os.getenv("WORKERS", "1"))
|
||
worker_id = get_worker_id()
|
||
task_cleanup_task = None
|
||
|
||
# 启动时
|
||
logger.info(f"Worker [{worker_id}] 启动中...")
|
||
emit_boot_event("phase_start", phase="worker", total=1, message=f"Worker [{worker_id}] 启动中")
|
||
|
||
# 清理旧的临时文件(仅主 Worker 执行)
|
||
if worker_id == 0:
|
||
cleanup_temp_directory()
|
||
try:
|
||
from .core.task_store import recover_tasks_after_restart
|
||
|
||
recover_tasks_after_restart()
|
||
except Exception as exc:
|
||
logger.warning("恢复离线任务状态失败: %s", exc)
|
||
task_cleanup_task = asyncio.create_task(cleanup_task_state_loop())
|
||
|
||
if settings.SPEAKER_DB_ENABLED:
|
||
try:
|
||
from .core.database import pg_speaker_db
|
||
|
||
await pg_speaker_db.connect()
|
||
except Exception as exc:
|
||
logger.warning(
|
||
"声纹数据库连接失败,将保留原说话人编号并禁用声纹注册接口: %s",
|
||
exc,
|
||
)
|
||
|
||
from .utils.model_loader import (
|
||
preload_models,
|
||
verify_required_models_integrity,
|
||
)
|
||
|
||
integrity_result = verify_required_models_integrity()
|
||
if integrity_result["invalid_models"]:
|
||
emit_boot_event("error", phase="integrity", message="required model integrity check failed")
|
||
raise RuntimeError("required model integrity check failed")
|
||
|
||
logger.info(f"Worker [{worker_id}] 正在加载模型...")
|
||
preload_result = preload_models()
|
||
|
||
asr_results = preload_result.get("asr_models", {})
|
||
loaded_count = sum(1 for r in asr_results.values() if r.get("loaded"))
|
||
total_count = len(asr_results)
|
||
logger.info(f"Worker [{worker_id}] 模型加载完成: {loaded_count}/{total_count}")
|
||
failed_asr_models = {
|
||
model_id: status.get("error")
|
||
for model_id, status in asr_results.items()
|
||
if not status.get("loaded") and status.get("error")
|
||
}
|
||
try:
|
||
from .services.asr.model_plan import get_active_qwen_model, load_supported_model_ids
|
||
|
||
active_qwen_model = get_active_qwen_model(load_supported_model_ids())
|
||
except Exception as exc:
|
||
logger.error(f"Worker [{worker_id}] 无法解析当前应启用的 Qwen 模型: {exc}")
|
||
emit_boot_event("error", phase="preload", message=f"无法解析当前应启用的 Qwen 模型: {exc}")
|
||
raise
|
||
|
||
active_qwen_status = asr_results.get(active_qwen_model, {})
|
||
if not active_qwen_status.get("loaded"):
|
||
qwen_error = active_qwen_status.get("error") or "unknown error"
|
||
logger.error(
|
||
f"Worker [{worker_id}] Qwen 主模型预加载失败,拒绝启动: {active_qwen_model}, error={qwen_error}"
|
||
)
|
||
emit_boot_event(
|
||
"error",
|
||
phase="preload",
|
||
message=f"Qwen 主模型预加载失败,拒绝启动: {active_qwen_model}, error={qwen_error}",
|
||
)
|
||
raise RuntimeError(
|
||
f"required qwen model preload failed: {active_qwen_model}: {qwen_error}"
|
||
)
|
||
|
||
if failed_asr_models:
|
||
if loaded_count == 0:
|
||
logger.error(f"Worker [{worker_id}] ASR模型预加载失败详情: {failed_asr_models}")
|
||
emit_boot_event("error", phase="preload", message=f"ASR模型预加载失败详情: {failed_asr_models}")
|
||
raise RuntimeError(f"ASR model preload failed: {failed_asr_models}")
|
||
logger.warning(
|
||
f"Worker [{worker_id}] 部分ASR模型预加载失败,将以可用模型继续启动: {failed_asr_models}"
|
||
)
|
||
emit_boot_event(
|
||
"warning",
|
||
phase="preload",
|
||
message=f"部分ASR模型预加载失败,将以可用模型继续启动: {failed_asr_models}",
|
||
)
|
||
|
||
if settings.SPEAKER_DB_ENABLED:
|
||
try:
|
||
from .core.database import pg_speaker_db
|
||
from .services.speaker_registry import get_speaker_registry_service
|
||
|
||
if pg_speaker_db.is_connected:
|
||
logger.info(f"Worker [{worker_id}] 正在预加载声纹识别模型...")
|
||
get_speaker_registry_service().ensure_loaded()
|
||
except Exception as exc:
|
||
logger.warning("声纹识别模型预加载失败,将在首次请求时重试: %s", exc)
|
||
|
||
logger.info(f"Worker [{worker_id}] 已就绪")
|
||
emit_boot_event("ready", phase="worker", message=f"Worker [{worker_id}] 已就绪")
|
||
|
||
yield
|
||
|
||
# 关闭时
|
||
if task_cleanup_task is not None:
|
||
task_cleanup_task.cancel()
|
||
try:
|
||
await task_cleanup_task
|
||
except asyncio.CancelledError:
|
||
pass
|
||
|
||
if settings.SPEAKER_DB_ENABLED:
|
||
try:
|
||
from .core.database import pg_speaker_db
|
||
|
||
await pg_speaker_db.close()
|
||
except Exception as exc:
|
||
logger.warning("关闭声纹数据库连接时出错: %s", exc)
|
||
logger.info(f"Worker [{worker_id}] 正在关闭推理线程池...")
|
||
shutdown_executor()
|
||
logger.info(f"Worker [{worker_id}] 已关闭")
|
||
|
||
|
||
def create_app() -> FastAPI:
|
||
"""创建FastAPI应用"""
|
||
|
||
# 设置日志
|
||
setup_logging()
|
||
|
||
app = FastAPIOffline(
|
||
title=settings.APP_NAME,
|
||
description=settings.APP_DESCRIPTION,
|
||
version=settings.APP_VERSION,
|
||
docs_url=settings.docs_url,
|
||
redoc_url=settings.redoc_url,
|
||
lifespan=lifespan, # 添加生命周期管理
|
||
)
|
||
|
||
# 添加CORS中间件
|
||
app.add_middleware(
|
||
CORSMiddleware,
|
||
allow_origins=["*"],
|
||
allow_credentials=True,
|
||
allow_methods=["*"],
|
||
allow_headers=["*"],
|
||
)
|
||
|
||
# 注册异常处理器
|
||
app.add_exception_handler(APIException, api_exception_handler)
|
||
app.add_exception_handler(HTTPException, http_exception_handler)
|
||
app.add_exception_handler(RequestValidationError, validation_exception_handler)
|
||
app.add_exception_handler(Exception, general_exception_handler)
|
||
|
||
# 注册静态文件服务(用于临时文件)
|
||
app.mount("/tmp", StaticFiles(directory=settings.TEMP_DIR), name="temp_files")
|
||
app.mount(
|
||
"/test-web",
|
||
StaticFiles(directory=str(settings.BASE_DIR / "test_web"), html=True),
|
||
name="test_web",
|
||
)
|
||
|
||
# 注册API路由
|
||
app.include_router(api_router)
|
||
|
||
# 根路径
|
||
@app.get("/", summary="根路径", description="API服务根路径")
|
||
async def root():
|
||
return {
|
||
"message": settings.APP_NAME,
|
||
"version": settings.APP_VERSION,
|
||
"description": settings.APP_DESCRIPTION,
|
||
"endpoints": {
|
||
# 阿里云兼容 API
|
||
"asr": "/stream/v1/asr",
|
||
"asr_models": "/stream/v1/asr/models",
|
||
"asr_health": "/stream/v1/asr/health",
|
||
"ws_asr": "/ws/v1/asr/qwen",
|
||
"ws_qwen3_asr": "/ws/v1/asr/qwen",
|
||
# OpenAI 兼容 API
|
||
"openai_models": "/v1/models",
|
||
"openai_transcriptions": "/v1/audio/transcriptions",
|
||
# 独立会议离线 / 声纹管理 API
|
||
"meeting_files": "/api/v1/files",
|
||
"meeting_transcriptions": "/api/v1/asr/transcriptions",
|
||
"speakers": "/api/v1/speakers",
|
||
"test_web": "/test-web/",
|
||
# 文档
|
||
"docs": settings.docs_url or "禁用",
|
||
},
|
||
}
|
||
|
||
return app
|
||
|
||
|
||
# 创建全局应用实例
|
||
app = create_app()
|