test/app/main.py

284 lines
9.8 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

# -*- 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()