test/app/api/v1/meeting.py

856 lines
33 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 -*-
"""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)