856 lines
33 KiB
Python
856 lines
33 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""Independent meeting-style offline API compatible with Model-Test-New shapes."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import asyncio
|
||
import math
|
||
import time
|
||
import uuid
|
||
from pathlib import Path
|
||
from typing import Any, Optional
|
||
from urllib.parse import unquote, urlparse
|
||
|
||
import requests
|
||
from fastapi import APIRouter, BackgroundTasks, Body, File, HTTPException, UploadFile
|
||
from pydantic import BaseModel, ConfigDict, Field
|
||
|
||
from app.core.config import settings
|
||
from app.core.database import pg_speaker_db
|
||
from app.core.exceptions import InvalidParameterException
|
||
from app.core.task_store import create_task_record, get_task_record, update_task_record
|
||
from app.models.common import SampleRate
|
||
from app.services.asr.model_selection import validate_offline_model_id
|
||
from app.services.asr.offline_transcription_service import (
|
||
OfflineTranscriptionOptions,
|
||
PreparedAudio,
|
||
get_offline_transcription_service,
|
||
)
|
||
from app.services.speaker_registry import get_speaker_registry_service
|
||
|
||
router = APIRouter(prefix=settings.API_PREFIX)
|
||
|
||
_UPLOAD_DIR = Path(settings.TEMP_DIR) / "api_uploads"
|
||
_UPLOAD_DIR.mkdir(parents=True, exist_ok=True)
|
||
_uploaded_files: dict[str, dict[str, str]] = {}
|
||
|
||
|
||
def _recognition_request_config_example() -> dict[str, Any]:
|
||
example: dict[str, Any] = {
|
||
"enable_speaker": True,
|
||
"match_speaker_registry": True,
|
||
"enable_text_cleanup": True,
|
||
"speaker_threshold": 0.6,
|
||
"hotwords": [
|
||
{"hotword": "阿里巴巴", "weight": 2.0},
|
||
{"hotword": "腾讯", "weight": 1.5},
|
||
],
|
||
}
|
||
if settings.ASR_ENABLE_WORD_TIMESTAMPS:
|
||
example["word_timestamps"] = False
|
||
return example
|
||
|
||
|
||
def _meeting_request_examples() -> dict[str, Any]:
|
||
return {
|
||
"url_with_hotwords": {
|
||
"summary": "URL + 结构化热词",
|
||
"description": "推荐的生产调用格式。",
|
||
"value": {
|
||
"model": "qwen3-asr-0.6b",
|
||
"audio_address": "https://example.com/media/meeting.mp4",
|
||
"config": _recognition_request_config_example(),
|
||
},
|
||
},
|
||
"local_path": {
|
||
"summary": "服务端本地路径",
|
||
"value": {
|
||
"audio_address": "/data/audio/meeting.m4a",
|
||
"config": {
|
||
"enable_speaker": True,
|
||
"match_speaker_registry": False,
|
||
"enable_text_cleanup": True,
|
||
},
|
||
},
|
||
},
|
||
}
|
||
|
||
|
||
def _recognition_request_config_schema_extra(schema: dict[str, Any], _model: Any) -> None:
|
||
schema["example"] = _recognition_request_config_example()
|
||
if settings.ASR_ENABLE_WORD_TIMESTAMPS:
|
||
return
|
||
properties = schema.get("properties")
|
||
if isinstance(properties, dict):
|
||
properties.pop("word_timestamps", None)
|
||
required = schema.get("required")
|
||
if isinstance(required, list):
|
||
schema["required"] = [item for item in required if item != "word_timestamps"]
|
||
|
||
|
||
class HotwordItem(BaseModel):
|
||
model_config = ConfigDict(json_schema_extra={"example": {"hotword": "汇智", "weight": 2.0}})
|
||
|
||
hotword: str = Field(
|
||
...,
|
||
title="热词内容",
|
||
description="需要优先匹配或纠正的业务词、专有名词、人名、产品名等。Qwen3-ASR 无原生热词能力,服务会在识别后做规则化热词纠正。",
|
||
examples=["通义千问"],
|
||
)
|
||
weight: float = Field(
|
||
1.0,
|
||
title="热词权重",
|
||
description="热词权重,兼容老项目 Model-Test-New 的传参格式;当前规则纠正主要使用热词文本,权重保留用于接口兼容和后续扩展。",
|
||
examples=[2.0],
|
||
)
|
||
|
||
|
||
class RecognitionRequestConfig(BaseModel):
|
||
model_config = ConfigDict(
|
||
json_schema_extra=_recognition_request_config_schema_extra
|
||
)
|
||
enable_speaker: bool = Field(
|
||
True,
|
||
title="是否启用说话人分离",
|
||
description="是否执行说话人分离。开启后返回的分段会带 speaker 字段;关闭后不做说话人分离和声纹库匹配。",
|
||
examples=[True],
|
||
)
|
||
match_speaker_registry: bool = Field(
|
||
True,
|
||
title="是否匹配已注册声纹库",
|
||
description="是否在说话人分离后匹配已注册声纹库。只有 enable_speaker=true 时生效。",
|
||
examples=[True],
|
||
)
|
||
enable_text_cleanup: bool = Field(
|
||
True,
|
||
title="是否启用文字去重",
|
||
description="是否启用识别结果去重/边界重叠裁剪,适合长音频分段识别后的重复文本清理。",
|
||
examples=[True],
|
||
)
|
||
speaker_threshold: Optional[float] = Field(
|
||
None,
|
||
title="声纹匹配阈值",
|
||
description="本次声纹库匹配阈值,范围通常为 0~1;不传则使用服务默认配置。",
|
||
examples=[0.6],
|
||
)
|
||
word_timestamps: bool = Field(
|
||
False,
|
||
title="是否返回字词级时间戳",
|
||
description="是否返回字词级时间戳。长会议场景建议默认关闭以降低耗时和结果体积。",
|
||
examples=[False],
|
||
)
|
||
hotwords: str | list[HotwordItem] = Field(
|
||
default_factory=list,
|
||
title="热词",
|
||
description=(
|
||
"热词参数。推荐传老项目兼容的结构化数组:"
|
||
"[{\"hotword\":\"通义千问\",\"weight\":2.0}];也兼容字符串:\"通义千问 2.0 ModelScope 1.5\"。"
|
||
"Qwen3-ASR 本身没有原生热词参数,服务会在识别后使用热词规则做文本纠正。"
|
||
),
|
||
examples=[[{"hotword": "通义千问", "weight": 2.0}, {"hotword": "ModelScope", "weight": 1.5}]],
|
||
)
|
||
sample_rate: int = Field(
|
||
16000,
|
||
title="采样率",
|
||
description="音频处理采样率,默认 16000。",
|
||
examples=[16000],
|
||
)
|
||
|
||
|
||
class TranscriptionCreateRequest(BaseModel):
|
||
model_config = ConfigDict(
|
||
json_schema_extra={
|
||
"example": {
|
||
"model": "qwen3-asr-0.6b",
|
||
"audio_address": "https://example.com/meeting.m4a",
|
||
"config": {
|
||
"enable_speaker": True,
|
||
"match_speaker_registry": True,
|
||
"enable_text_cleanup": True,
|
||
"speaker_threshold": 0.6,
|
||
"hotwords": [
|
||
{"hotword": "汇智", "weight": 2.0},
|
||
{"hotword": "通义千问", "weight": 2.5},
|
||
],
|
||
},
|
||
}
|
||
}
|
||
)
|
||
model: Optional[str] = Field(
|
||
None,
|
||
title="离线 ASR 模型 ID",
|
||
description="可选。不传则使用服务当前默认模型;可传当前可用的本地模型 ID,如 qwen3-asr-0.6b 或 qwen3-asr-1.7b。",
|
||
examples=["qwen3-asr-0.6b"],
|
||
max_length=128,
|
||
)
|
||
audio_address: str = Field(
|
||
...,
|
||
title="音视频地址",
|
||
description="必填。音频/视频文件地址,支持 HTTP/HTTPS URL、file:// 地址或服务端本地音视频路径。该接口生产调用不使用 file_id/file_url。",
|
||
examples=["https://example.com/meeting.m4a"],
|
||
)
|
||
config: RecognitionRequestConfig = Field(
|
||
default_factory=RecognitionRequestConfig,
|
||
title="识别配置",
|
||
description=(
|
||
"离线会议识别配置,包括说话人分离、声纹库匹配、文字去重和热词规则。"
|
||
if not settings.ASR_ENABLE_WORD_TIMESTAMPS
|
||
else "离线会议识别配置,包括说话人分离、声纹库匹配、文字去重、字词时间戳和热词规则。"
|
||
)
|
||
)
|
||
|
||
|
||
class SpeakerRegisterRequest(BaseModel):
|
||
model_config = ConfigDict(
|
||
json_schema_extra={"example": {"name": "张三", "user_id": "u001", "audio_address": "/data/speakers/zhangsan.wav"}}
|
||
)
|
||
name: str
|
||
audio_address: Optional[str] = None
|
||
file_url: Optional[str] = None
|
||
file_id: Optional[str] = None
|
||
user_id: Optional[str] = None
|
||
|
||
|
||
class SpeakerIdentifyRequest(BaseModel):
|
||
model_config = ConfigDict(json_schema_extra={"example": {"audio_address": "/data/audio/sample.wav", "threshold": 0.6}})
|
||
audio_address: Optional[str] = None
|
||
file_url: Optional[str] = None
|
||
file_id: Optional[str] = None
|
||
threshold: Optional[float] = None
|
||
|
||
|
||
def _get_meeting_transcription_description() -> str:
|
||
return """创建离线会议转写任务。该接口采用 **JSON 请求体**,适合业务系统直接提交会议音频/视频地址并异步轮询结果。
|
||
|
||
**音频输入方式:**
|
||
1. **HTTP/HTTPS URL**:`audio_address="https://example.com/meeting.mp4"`
|
||
2. **file:// 地址**:`audio_address="file:///data/audio/meeting.m4a"`
|
||
3. **服务端本地路径**:`audio_address="/data/audio/meeting.m4a"`
|
||
|
||
> 生产调用只使用 `audio_address`。`file_id` / `file_url` 不属于该接口参数,上传接口仅用于测试页面调试。
|
||
|
||
**请求参数:**
|
||
| 字段 | 类型 | 必填 | 默认值 | 说明 |
|
||
|------|------|------|--------|------|
|
||
| `model` | string/null | 否 | 服务默认模型 | 离线 ASR 模型 ID;不传则使用 `QWEN3_ASR_MODEL` 或自动选择结果 |
|
||
| `audio_address` | string | 是 | - | 音频/视频文件地址,支持 HTTP/HTTPS、`file://`、服务端本地路径 |
|
||
| `config` | object | 否 | 默认配置 | 识别配置对象 |
|
||
| `config.enable_speaker` | boolean | 否 | `true` | 是否启用说话人分离 |
|
||
| `config.match_speaker_registry` | boolean | 否 | `true` | 是否匹配已注册声纹库,仅在 `enable_speaker=true` 时生效 |
|
||
| `config.enable_text_cleanup` | boolean | 否 | `true` | 是否启用文字去重、边界重叠裁剪 |
|
||
| `config.speaker_threshold` | number/null | 否 | 服务默认值 | 本次声纹库匹配阈值,通常为 `0~1` |
|
||
""" + (
|
||
"| `config.word_timestamps` | boolean | 否 | `false` | 是否返回字词级时间戳;开启后会调用 Qwen3 ForcedAligner |\n"
|
||
if settings.ASR_ENABLE_WORD_TIMESTAMPS
|
||
else ""
|
||
) + """
|
||
| `config.hotwords` | array/string | 否 | `[]` | 热词。推荐老项目兼容数组格式:`[{\"hotword\":\"通义千问\",\"weight\":2.0}]` |
|
||
| `config.sample_rate` | integer | 否 | `16000` | 音频处理采样率 |
|
||
|
||
**热词说明:**
|
||
- Qwen3-ASR 本身没有原生热词参数。
|
||
- 本服务兼容 Model-Test-New 的热词格式,识别完成后通过规则做热词纠正。
|
||
- 推荐格式:`[{\"hotword\":\"通义千问\",\"weight\":2.0},{\"hotword\":\"ModelScope\",\"weight\":1.5}]`
|
||
- 兼容字符串格式:`\"通义千问 2.0 ModelScope 1.5\"`
|
||
|
||
**创建任务返回:**
|
||
| 字段 | 类型 | 说明 |
|
||
|------|------|------|
|
||
| `code` | integer | 业务状态码,`0` 表示成功 |
|
||
| `message` | string | 业务消息 |
|
||
| `data.task_id` | string | 任务 ID,用于查询识别进度与结果 |
|
||
| `data.status` | string | 初始状态,通常为 `queued` |
|
||
| `data.stage` | string | 当前阶段 |
|
||
| `data.percentage` | number | 当前进度百分比 |
|
||
|
||
**查询任务返回:**
|
||
调用 `GET /api/v1/asr/transcriptions/{task_id}` 查询任务。完成后 `data.result` 包含完整结果。
|
||
|
||
| 字段 | 类型 | 说明 |
|
||
|------|------|------|
|
||
| `data.status` | string | `queued` / `processing` / `completed` / `failed` |
|
||
| `data.result.text` | string | 完整转写文本 |
|
||
| `data.result.segments` | array | 分段结果 |
|
||
| `data.result.segments[].start` | number | 分段开始时间,秒 |
|
||
| `data.result.segments[].end` | number | 分段结束时间,秒 |
|
||
| `data.result.segments[].text` | string | 分段文本 |
|
||
| `data.result.segments[].speaker_id` | string | 说话人 ID,启用说话人分离时返回 |
|
||
| `data.result.segments[].speaker_name` | string | 匹配到声纹库时返回姓名,否则通常等于 speaker_id |
|
||
| `data.result.segments[].user_id` | string | 匹配到声纹库用户时返回 |
|
||
| `data.result.segments[].word_tokens` | array | 字词级时间戳,仅 `word_timestamps=true` 且模型支持时返回 |
|
||
| `data.result.audio_duration_seconds` | number | 音频总时长,秒 |
|
||
| `data.result.speaker_count` | integer | 说话人数量 |
|
||
| `data.result.speakers` | array | 说话人列表 |
|
||
"""
|
||
|
||
|
||
def _ok(data: Any) -> dict[str, Any]:
|
||
return {"code": 0, "message": "ok", "data": data or {}}
|
||
|
||
|
||
def _public_task_payload(task_id: str, task: dict[str, Any]) -> dict[str, Any]:
|
||
payload: dict[str, Any] = {"task_id": task_id, "status": task["status"]}
|
||
for key in (
|
||
"stage",
|
||
"message",
|
||
"percentage",
|
||
"detail",
|
||
"created_at",
|
||
"updated_at",
|
||
):
|
||
if task.get(key) is not None:
|
||
payload[key] = task[key]
|
||
if (
|
||
task["status"] == "processing"
|
||
and task.get("created_at") is not None
|
||
and task.get("percentage") not in (None, 0)
|
||
):
|
||
elapsed_seconds = max(1.0, time.time() - float(task["created_at"]))
|
||
percentage = max(1.0, float(task.get("percentage") or 0))
|
||
payload["eta_seconds"] = max(
|
||
0,
|
||
int(math.ceil(elapsed_seconds * (100.0 - percentage) / percentage)),
|
||
)
|
||
elif task["status"] == "completed":
|
||
payload["eta_seconds"] = 0
|
||
if task.get("result") is not None:
|
||
payload["result"] = task["result"]
|
||
if task.get("error"):
|
||
payload["error"] = task["error"]
|
||
return payload
|
||
|
||
|
||
def _require_db() -> None:
|
||
if not pg_speaker_db.is_connected:
|
||
raise HTTPException(
|
||
status_code=503,
|
||
detail={"code": 1001, "message": "speaker database is not connected"},
|
||
)
|
||
|
||
|
||
def _normalize_hotwords_for_asr(hotwords: str | list[HotwordItem] | None) -> str:
|
||
if isinstance(hotwords, str):
|
||
return hotwords.strip()
|
||
parts: list[str] = []
|
||
for item in hotwords or []:
|
||
text = item.hotword.strip()
|
||
if not text:
|
||
continue
|
||
parts.extend([text, f"{float(item.weight):g}"])
|
||
return " ".join(parts)
|
||
|
||
|
||
async def _download_url(url: str) -> bytes:
|
||
loop = asyncio.get_running_loop()
|
||
|
||
def _download() -> bytes:
|
||
response = requests.get(url, timeout=60)
|
||
response.raise_for_status()
|
||
return response.content
|
||
|
||
return await loop.run_in_executor(None, _download)
|
||
|
||
|
||
def _is_http_url(value: str) -> bool:
|
||
return value.lower().startswith(("http://", "https://"))
|
||
|
||
|
||
def _source_filename(source: str) -> str:
|
||
parsed = urlparse(source)
|
||
path = unquote(parsed.path or source)
|
||
return Path(path).name or "audio"
|
||
|
||
|
||
async def _read_local_file(path_value: str) -> bytes:
|
||
parsed = urlparse(path_value)
|
||
local_path = Path(unquote(parsed.path)) if parsed.scheme == "file" else Path(path_value).expanduser()
|
||
if not local_path.is_file():
|
||
raise HTTPException(status_code=404, detail={"code": 1001, "message": "local media file not found"})
|
||
loop = asyncio.get_running_loop()
|
||
return await loop.run_in_executor(None, local_path.read_bytes)
|
||
|
||
|
||
async def _resolve_source_bytes(audio_address: Optional[str], file_id: Optional[str]) -> tuple[bytes, Optional[str]]:
|
||
if audio_address:
|
||
source = audio_address.strip()
|
||
if _is_http_url(source):
|
||
return await _download_url(source), _source_filename(source)
|
||
return await _read_local_file(source), _source_filename(source)
|
||
if file_id:
|
||
item = _uploaded_files.get(file_id)
|
||
if not item:
|
||
matched_files = sorted(_UPLOAD_DIR.glob(f"{file_id}.*"))
|
||
if not matched_files:
|
||
raise HTTPException(status_code=404, detail={"code": 1001, "message": "file_id not found"})
|
||
matched_path = matched_files[0]
|
||
item = {"path": str(matched_path), "filename": matched_path.name}
|
||
_uploaded_files[file_id] = item
|
||
path = item["path"]
|
||
with open(path, "rb") as file_obj:
|
||
return file_obj.read(), item.get("filename")
|
||
raise HTTPException(status_code=400, detail={"code": 1001, "message": "audio_address or file_id is required"})
|
||
|
||
|
||
def _cleanup_uploaded_file(file_id: Optional[str]) -> None:
|
||
if not file_id:
|
||
return
|
||
item = _uploaded_files.pop(file_id, None)
|
||
candidate_paths: list[Path] = []
|
||
if item and item.get("path"):
|
||
candidate_paths.append(Path(item["path"]))
|
||
candidate_paths.extend(sorted(_UPLOAD_DIR.glob(f"{file_id}.*")))
|
||
for path in candidate_paths:
|
||
try:
|
||
path.unlink(missing_ok=True)
|
||
except Exception:
|
||
pass
|
||
|
||
|
||
def _format_segments(asr_result: Any) -> list[dict[str, Any]]:
|
||
segments: list[dict[str, Any]] = []
|
||
for segment in asr_result.segments:
|
||
item: dict[str, Any] = {
|
||
"start": round(float(segment.start_time), 2),
|
||
"end": round(float(segment.end_time), 2),
|
||
"duration": round(float(segment.end_time - segment.start_time), 2),
|
||
"text": segment.text,
|
||
}
|
||
if segment.speaker_id:
|
||
item["speaker_id"] = segment.speaker_id
|
||
if segment.speaker_name:
|
||
item["speaker_name"] = segment.speaker_name
|
||
elif segment.speaker_id:
|
||
item["speaker_name"] = segment.speaker_id
|
||
if segment.user_id:
|
||
item["user_id"] = segment.user_id
|
||
if segment.word_tokens:
|
||
item["word_tokens"] = [
|
||
{
|
||
"text": token.text,
|
||
"start": round(float(token.start_time), 3),
|
||
"end": round(float(token.end_time), 3),
|
||
}
|
||
for token in segment.word_tokens
|
||
]
|
||
segments.append(item)
|
||
return segments
|
||
|
||
|
||
def _build_transcription_result(asr_result: Any) -> dict[str, Any]:
|
||
segments = _format_segments(asr_result)
|
||
speakers: list[str] = []
|
||
for segment in segments:
|
||
name = str(segment.get("speaker_name") or segment.get("speaker_id") or "").strip()
|
||
if name and name not in speakers:
|
||
speakers.append(name)
|
||
return {
|
||
"text": asr_result.text,
|
||
"segments": segments,
|
||
"audio_duration_seconds": round(float(asr_result.duration), 2),
|
||
"speaker_count": len(speakers),
|
||
"speakers": speakers,
|
||
}
|
||
|
||
|
||
def _update_task_progress(
|
||
task_id: str,
|
||
*,
|
||
status: str,
|
||
stage: str,
|
||
message: str,
|
||
percentage: int,
|
||
detail: Optional[dict[str, Any]] = None,
|
||
) -> None:
|
||
current_task = get_task_record(task_id) or {}
|
||
current_status = str(current_task.get("status") or "")
|
||
current_percentage = int(current_task.get("percentage") or 0)
|
||
incoming_percentage = max(0, min(100, int(percentage)))
|
||
if (
|
||
current_status in {"queued", "processing"}
|
||
and status in {"queued", "processing"}
|
||
and incoming_percentage < current_percentage
|
||
):
|
||
return
|
||
payload: dict[str, Any] = {
|
||
"status": status,
|
||
"stage": stage,
|
||
"message": message,
|
||
"percentage": incoming_percentage,
|
||
}
|
||
if detail is not None:
|
||
payload["detail"] = detail
|
||
update_task_record(task_id, payload)
|
||
|
||
|
||
async def _run_transcription_task(
|
||
*,
|
||
task_id: str,
|
||
model: Optional[str],
|
||
audio_address: Optional[str],
|
||
config: RecognitionRequestConfig,
|
||
) -> None:
|
||
service = get_offline_transcription_service()
|
||
prepared_audio: Optional[PreparedAudio] = None
|
||
_update_task_progress(
|
||
task_id,
|
||
status="processing",
|
||
stage="preparing",
|
||
message="任务已开始,正在准备音频文件。",
|
||
percentage=5,
|
||
detail={"segment_total": 0, "segment_completed": 0},
|
||
)
|
||
try:
|
||
audio_data, filename = await _resolve_source_bytes(audio_address, None)
|
||
_update_task_progress(
|
||
task_id,
|
||
status="processing",
|
||
stage="normalizing",
|
||
message="音频来源已读取,正在转换为识别格式。",
|
||
percentage=8,
|
||
detail={"segment_total": 0, "segment_completed": 0},
|
||
)
|
||
prepared_audio = await service.prepare_upload(
|
||
audio_data=audio_data,
|
||
filename=filename,
|
||
task_id=task_id,
|
||
sample_rate=config.sample_rate or int(SampleRate.RATE_16000),
|
||
)
|
||
_update_task_progress(
|
||
task_id,
|
||
status="processing",
|
||
stage="transcribing",
|
||
message="音频已准备完成,正在执行离线识别。",
|
||
percentage=12,
|
||
detail={
|
||
"segment_total": 0,
|
||
"segment_completed": 0,
|
||
"audio_duration_seconds": round(float(prepared_audio.duration), 2),
|
||
},
|
||
)
|
||
|
||
def progress_callback(
|
||
stage: str,
|
||
message: str,
|
||
percentage: int,
|
||
detail: Optional[dict[str, Any]] = None,
|
||
) -> None:
|
||
_update_task_progress(
|
||
task_id,
|
||
status="processing",
|
||
stage=stage,
|
||
message=message,
|
||
percentage=percentage,
|
||
detail=detail,
|
||
)
|
||
|
||
asr_result = await service.transcribe(
|
||
prepared_audio,
|
||
OfflineTranscriptionOptions(
|
||
model_id=model,
|
||
sample_rate=config.sample_rate or int(SampleRate.RATE_16000),
|
||
hotwords=_normalize_hotwords_for_asr(config.hotwords),
|
||
enable_speaker_diarization=config.enable_speaker,
|
||
enable_speaker_identification=(
|
||
config.enable_speaker and config.match_speaker_registry
|
||
),
|
||
enable_text_cleanup=config.enable_text_cleanup,
|
||
word_timestamps=(
|
||
settings.ASR_ENABLE_WORD_TIMESTAMPS and config.word_timestamps
|
||
),
|
||
task_id=task_id,
|
||
progress_callback=progress_callback,
|
||
),
|
||
)
|
||
if (
|
||
config.enable_speaker
|
||
and config.match_speaker_registry
|
||
and config.speaker_threshold is not None
|
||
):
|
||
_update_task_progress(
|
||
task_id,
|
||
status="processing",
|
||
stage="speaker_matching",
|
||
message="正在使用本次阈值匹配已注册声纹库。",
|
||
percentage=95,
|
||
detail={
|
||
"segment_total": len(asr_result.segments),
|
||
"segment_completed": len(asr_result.segments),
|
||
"audio_duration_seconds": round(float(asr_result.duration), 2),
|
||
},
|
||
)
|
||
await get_speaker_registry_service().apply_registered_speakers(
|
||
asr_result,
|
||
threshold=config.speaker_threshold,
|
||
)
|
||
result = _build_transcription_result(asr_result)
|
||
_update_task_progress(
|
||
task_id,
|
||
status="completed",
|
||
stage="completed",
|
||
message="离线识别已经完成。",
|
||
percentage=100,
|
||
detail={
|
||
"segment_total": len(result.get("segments") or []),
|
||
"segment_completed": len(result.get("segments") or []),
|
||
"audio_duration_seconds": result.get("audio_duration_seconds"),
|
||
"speaker_count": result.get("speaker_count"),
|
||
"speakers": result.get("speakers"),
|
||
},
|
||
)
|
||
update_task_record(task_id, {"result": result})
|
||
except Exception as exc:
|
||
_update_task_progress(
|
||
task_id,
|
||
status="failed",
|
||
stage="failed",
|
||
message="离线识别执行失败。",
|
||
percentage=100,
|
||
)
|
||
update_task_record(task_id, {"error": str(exc)})
|
||
finally:
|
||
service.cleanup(prepared_audio)
|
||
|
||
|
||
@router.post("/files", tags=["会议离线接口"], summary="上传音频文件")
|
||
async def upload_file(file: UploadFile = File(...)) -> dict[str, Any]:
|
||
file_id = uuid.uuid4().hex
|
||
suffix = Path(file.filename or "").suffix or ".audio"
|
||
path = _UPLOAD_DIR / f"{file_id}{suffix}"
|
||
content = await file.read()
|
||
path.write_bytes(content)
|
||
_uploaded_files[file_id] = {"path": str(path), "filename": file.filename or path.name}
|
||
return _ok({"file_id": file_id, "file_path": str(path)})
|
||
|
||
|
||
@router.post(
|
||
"/asr/transcriptions",
|
||
tags=["会议离线接口"],
|
||
summary="创建离线会议识别任务",
|
||
description=_get_meeting_transcription_description(),
|
||
response_description="任务创建成功,返回 task_id;后续调用 GET /api/v1/asr/transcriptions/{task_id} 查询进度和结果。",
|
||
responses={
|
||
200: {
|
||
"description": "任务创建成功",
|
||
"content": {
|
||
"application/json": {
|
||
"example": {
|
||
"code": 0,
|
||
"message": "ok",
|
||
"data": {
|
||
"task_id": "task_6fd90daffd2f446f944b64fdf6c17572",
|
||
"status": "queued",
|
||
"stage": "queued",
|
||
"message": "任务已创建,等待进入处理流程。",
|
||
"percentage": 0,
|
||
},
|
||
}
|
||
}
|
||
},
|
||
},
|
||
400: {
|
||
"description": "请求参数错误",
|
||
"content": {
|
||
"application/json": {
|
||
"example": {
|
||
"detail": {
|
||
"code": 1001,
|
||
"message": "audio_address is required",
|
||
}
|
||
}
|
||
}
|
||
},
|
||
},
|
||
},
|
||
)
|
||
async def create_transcription_task(
|
||
background_tasks: BackgroundTasks,
|
||
item: TranscriptionCreateRequest = Body(
|
||
...,
|
||
description="离线会议识别任务 JSON 请求体。必须提供 audio_address,可在 config 中配置说话人分离、声纹库匹配、文字去重、热词等。",
|
||
examples=_meeting_request_examples(),
|
||
),
|
||
) -> dict[str, Any]:
|
||
if not item.audio_address:
|
||
raise HTTPException(status_code=400, detail={"code": 1001, "message": "audio_address is required"})
|
||
try:
|
||
model_id = validate_offline_model_id(item.model)
|
||
except InvalidParameterException as exc:
|
||
raise HTTPException(status_code=400, detail={"code": 1001, "message": exc.message}) from exc
|
||
task_id = f"task_{uuid.uuid4().hex}"
|
||
create_task_record(
|
||
task_id,
|
||
{
|
||
"status": "queued",
|
||
"stage": "queued",
|
||
"message": "任务已创建,等待进入处理流程。",
|
||
"percentage": 0,
|
||
"result": None,
|
||
"detail": {"segment_total": 0, "segment_completed": 0},
|
||
},
|
||
)
|
||
background_tasks.add_task(
|
||
_run_transcription_task,
|
||
task_id=task_id,
|
||
model=model_id,
|
||
audio_address=item.audio_address,
|
||
config=item.config,
|
||
)
|
||
return _ok(
|
||
{
|
||
"task_id": task_id,
|
||
"status": "queued",
|
||
"stage": "queued",
|
||
"message": "任务已创建,等待进入处理流程。",
|
||
"percentage": 0,
|
||
}
|
||
)
|
||
|
||
|
||
@router.get(
|
||
"/asr/transcriptions/{task_id}",
|
||
tags=["会议离线接口"],
|
||
summary="查询离线会议识别任务",
|
||
description=(
|
||
"根据创建任务时返回的 `task_id` 查询离线会议识别进度和结果。"
|
||
"`status=completed` 时,`data.result` 会包含完整文本、分段、说话人和可选字词级时间戳。"
|
||
),
|
||
responses={
|
||
200: {
|
||
"description": "查询成功",
|
||
"content": {
|
||
"application/json": {
|
||
"examples": {
|
||
"processing": {
|
||
"summary": "处理中",
|
||
"value": {
|
||
"code": 0,
|
||
"message": "ok",
|
||
"data": {
|
||
"task_id": "task_6fd90daffd2f446f944b64fdf6c17572",
|
||
"status": "processing",
|
||
"stage": "asr",
|
||
"message": "正在识别分段 80/127",
|
||
"percentage": 62,
|
||
"detail": {"segment_total": 127, "segment_completed": 80},
|
||
"eta_seconds": 38,
|
||
},
|
||
},
|
||
},
|
||
"completed": {
|
||
"summary": "已完成",
|
||
"value": {
|
||
"code": 0,
|
||
"message": "ok",
|
||
"data": {
|
||
"task_id": "task_6fd90daffd2f446f944b64fdf6c17572",
|
||
"status": "completed",
|
||
"stage": "completed",
|
||
"message": "识别完成。",
|
||
"percentage": 100,
|
||
"eta_seconds": 0,
|
||
"result": {
|
||
"text": "大家好,今天我们讨论通义千问和 ModelScope 的接入方案。",
|
||
"segments": [
|
||
{
|
||
"start": 0.0,
|
||
"end": 4.28,
|
||
"duration": 4.28,
|
||
"text": "大家好,今天我们讨论通义千问和 ModelScope 的接入方案。",
|
||
"speaker_id": "speaker_1",
|
||
"speaker_name": "张三",
|
||
"user_id": "u001",
|
||
"word_tokens": [
|
||
{"text": "大家好", "start": 0.12, "end": 0.68}
|
||
],
|
||
}
|
||
],
|
||
"audio_duration_seconds": 4669.05,
|
||
"speaker_count": 1,
|
||
"speakers": ["张三"],
|
||
},
|
||
},
|
||
},
|
||
},
|
||
}
|
||
}
|
||
},
|
||
},
|
||
404: {
|
||
"description": "任务不存在",
|
||
"content": {
|
||
"application/json": {
|
||
"example": {
|
||
"detail": {
|
||
"code": 1004,
|
||
"message": "task not found",
|
||
}
|
||
}
|
||
}
|
||
},
|
||
},
|
||
},
|
||
)
|
||
async def get_transcription_task(task_id: str) -> dict[str, Any]:
|
||
task = get_task_record(task_id, include_result=True)
|
||
if not task:
|
||
raise HTTPException(status_code=404, detail={"code": 1001, "message": "Task not found"})
|
||
return _ok(_public_task_payload(task_id, task))
|
||
|
||
|
||
@router.post("/speakers", tags=["声纹管理"], summary="注册声纹")
|
||
async def register_speaker(item: SpeakerRegisterRequest) -> dict[str, Any]:
|
||
_require_db()
|
||
temp_path: Optional[str] = None
|
||
try:
|
||
audio_data, filename = await _resolve_source_bytes(item.audio_address or item.file_url, item.file_id)
|
||
suffix = Path(filename or "").suffix or ".wav"
|
||
temp_path = await get_speaker_registry_service().save_upload_to_temp(audio_data, suffix=suffix)
|
||
result = await get_speaker_registry_service().register_file(
|
||
name=item.name,
|
||
file_path=temp_path,
|
||
user_id=item.user_id,
|
||
)
|
||
return _ok(result)
|
||
finally:
|
||
get_speaker_registry_service().cleanup_file(temp_path)
|
||
_cleanup_uploaded_file(item.file_id)
|
||
|
||
|
||
@router.get("/speakers", tags=["声纹管理"], summary="获取声纹人员列表")
|
||
async def list_speakers() -> dict[str, Any]:
|
||
_require_db()
|
||
return _ok(
|
||
{
|
||
"items": await get_speaker_registry_service().list_speakers(),
|
||
"speaker_model": "CampPlus",
|
||
}
|
||
)
|
||
|
||
|
||
@router.delete("/speakers/{speaker_id}", tags=["声纹管理"], summary="删除声纹人员")
|
||
async def delete_speaker(speaker_id: str) -> dict[str, Any]:
|
||
_require_db()
|
||
deleted = await get_speaker_registry_service().delete_speaker(speaker_id)
|
||
if not deleted:
|
||
raise HTTPException(status_code=404, detail={"code": 1001, "message": "Speaker not found"})
|
||
return _ok({"id": speaker_id})
|
||
|
||
|
||
@router.post("/speakers/identify", tags=["声纹管理"], summary="识别音频中的说话人")
|
||
async def identify_speaker(item: SpeakerIdentifyRequest) -> dict[str, Any]:
|
||
_require_db()
|
||
temp_path: Optional[str] = None
|
||
try:
|
||
audio_data, filename = await _resolve_source_bytes(item.audio_address or item.file_url, item.file_id)
|
||
suffix = Path(filename or "").suffix or ".wav"
|
||
temp_path = await get_speaker_registry_service().save_upload_to_temp(audio_data, suffix=suffix)
|
||
return _ok(
|
||
await get_speaker_registry_service().identify_file(
|
||
file_path=temp_path,
|
||
threshold=item.threshold,
|
||
)
|
||
)
|
||
finally:
|
||
get_speaker_registry_service().cleanup_file(temp_path)
|
||
_cleanup_uploaded_file(item.file_id)
|