From c764d43a5419ad283bfa91e4468b53f5dcf286ca Mon Sep 17 00:00:00 2001 From: wxg0103 <727495428@qq.com> Date: Tue, 25 Aug 2026 16:53:49 +0800 Subject: [PATCH] feat: add Tencent WAND ASR support with model credential and parameter handling --- .../credential/tokenhub_stt.py | 83 +++++++++++++++++++ .../impl/tencent_model_provider/model/stt.py | 73 +++++++++++++++- .../tencent_model_provider.py | 21 ++++- 3 files changed, 175 insertions(+), 2 deletions(-) create mode 100644 apps/models_provider/impl/tencent_model_provider/credential/tokenhub_stt.py diff --git a/apps/models_provider/impl/tencent_model_provider/credential/tokenhub_stt.py b/apps/models_provider/impl/tencent_model_provider/credential/tokenhub_stt.py new file mode 100644 index 00000000000..c4051f9ee2d --- /dev/null +++ b/apps/models_provider/impl/tencent_model_provider/credential/tokenhub_stt.py @@ -0,0 +1,83 @@ +# coding=utf-8 +""" +@project: MaxKB +@desc: Tencent Tokenhub ASR sync_transcribe credential (model: wand-asr-v1 / hy-asr-3.0-preview) +""" + +from django.utils.translation import gettext_lazy as _, gettext + +from common import forms +from common.exception.app_exception import AppApiException +from common.forms import BaseForm, TooltipLabel +from common.utils.logger import maxkb_logger +from models_provider.base_model_provider import BaseModelCredential, ValidCode + + +class TencentTokenhubSTTModelParams(BaseForm): + source = forms.SingleSelect( + label=TooltipLabel(_("Recognition language"), _("Recognition language: zh / en, auto detected when omitted")), + text_field="value", + value_field="value", + option_list=[ + {"value": "", "label": _("Auto detect")}, + {"value": "zh", "label": _("Chinese")}, + {"value": "en", "label": _("English")}, + ], + required=False, + default_value="", + ) + voice_encode_format = forms.SingleSelect( + label=TooltipLabel(_("Audio encoding"), _("pcm / wav / ogg / mp3, auto detected when omitted")), + text_field="value", + value_field="value", + option_list=[ + {"value": "", "label": _("Auto")}, + {"value": "pcm", "label": "pcm"}, + {"value": "wav", "label": "wav"}, + {"value": "ogg", "label": "ogg"}, + {"value": "mp3", "label": "mp3"}, + ], + required=False, + default_value="", + ) + + +class TencentTokenhubSTTModelCredential(BaseForm, BaseModelCredential): + def is_valid(self, model_type, model_name, model_credential, model_params, provider, raise_exception=False): + model_type_list = provider.get_model_type_list() + if not any(list(filter(lambda mt: mt.get("value") == model_type, model_type_list))): + raise AppApiException( + ValidCode.valid_error.value, + gettext("{model_type} Model type is not supported").format(model_type=model_type), + ) + if "api_key" not in model_credential: + if raise_exception: + raise AppApiException(ValidCode.valid_error.value, gettext("{key} is required").format(key="api_key")) + return False + try: + model = provider.get_model(model_type, model_name, model_credential, **model_params) + model.check_auth() + except Exception as e: + maxkb_logger.error(f"Exception: {e}", exc_info=True) + if raise_exception: + raise AppApiException( + ValidCode.valid_error.value, + gettext("Verification failed, please check whether the parameters are correct: {error}").format( + error=str(e) + ), + ) + return False + return True + + def encryption_dict(self, model): + return {**model, "api_key": super().encryption(model.get("api_key", ""))} + + base_url = forms.TextInputField( + label=TooltipLabel(_("API URL"), _("Tokenhub sync_transcribe endpoint")), + required=False, + default_value="https://tokenhub.tencentmaas.com/v1/wand/asrproxy/sync_transcribe", + ) + api_key = forms.PasswordInputField(_("API Key"), required=True) + + def get_model_params_setting_form(self, model_name): + return TencentTokenhubSTTModelParams() diff --git a/apps/models_provider/impl/tencent_model_provider/model/stt.py b/apps/models_provider/impl/tencent_model_provider/model/stt.py index f8157735c71..f09ced2c7b0 100644 --- a/apps/models_provider/impl/tencent_model_provider/model/stt.py +++ b/apps/models_provider/impl/tencent_model_provider/model/stt.py @@ -2,7 +2,9 @@ import json import os import traceback -from typing import Dict + +import requests +from typing import Dict, Optional from tencentcloud.asr.v20190614 import asr_client, models from tencentcloud.common import credential @@ -82,3 +84,72 @@ def speech_to_text(self, audio_file): except TencentCloudSDKException as err: maxkb_logger.error(f":Error: {str(err)}: {traceback.format_exc()}") raise err + + +DEFAULT_WAND_BASE_URL = 'https://tokenhub.tencentmaas.com/v1/wand/asrproxy/sync_transcribe' + + +class TencentWandSpeechToText(MaxKBBaseModel, BaseSpeechToText): + api_key: str + model: str + params: dict + base_url: Optional[str] = DEFAULT_WAND_BASE_URL + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.api_key = kwargs.get('api_key') + self.model = kwargs.get('model') + self.params = kwargs.get('params') or {} + self.base_url = kwargs.get('base_url') or DEFAULT_WAND_BASE_URL + + @staticmethod + def is_cache_model(): + return False + + @staticmethod + def new_instance(model_type, model_name, model_credential: Dict[str, object], **model_kwargs): + instance_kwargs = { + "api_key": model_credential.get('api_key'), + "model": model_name, + "params": model_kwargs, + **model_kwargs, + } + base_url = model_credential.get('base_url') + if base_url: + instance_kwargs["base_url"] = base_url + return TencentWandSpeechToText(**instance_kwargs) + + def check_auth(self): + cwd = os.path.dirname(os.path.abspath(__file__)) + with open(f'{cwd}/iat_mp3_16k.mp3', 'rb') as f: + self.speech_to_text(f) + + def speech_to_text(self, audio_file): + try: + payload = {"model": self.model} + # 仅使用上传音频文件的 base64 data,不提供 input_url 兜底 + audio_data = audio_file.read() + payload["data"] = base64.b64encode(audio_data).decode('utf-8') + for key in ("source", "voice_encode_format"): + if self.params.get(key): + payload[key] = self.params[key] + + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + } + response = requests.post(self.base_url, headers=headers, json=payload, timeout=300) + response.raise_for_status() + result = response.json() + if result.get("status") != "completed": + maxkb_logger.error(f"WAND ASR task not completed: {result}") + raise Exception(f"WAND ASR task not completed: {result}") + output = result.get("output") or {} + text = output.get("text") + if not text: + sentences = output.get("sentences") or [] + text = " ".join([s.get('text', '') for s in sentences if s.get('text')]) + return text + except Exception as e: + maxkb_logger.error(f"WAND ASR Error: {str(e)}: {traceback.format_exc()}") + raise e diff --git a/apps/models_provider/impl/tencent_model_provider/tencent_model_provider.py b/apps/models_provider/impl/tencent_model_provider/tencent_model_provider.py index 1b4877f9242..7541b60bc01 100644 --- a/apps/models_provider/impl/tencent_model_provider/tencent_model_provider.py +++ b/apps/models_provider/impl/tencent_model_provider/tencent_model_provider.py @@ -10,11 +10,12 @@ from models_provider.impl.tencent_model_provider.credential.image import TencentVisionModelCredential from models_provider.impl.tencent_model_provider.credential.llm import TencentLLMModelCredential from models_provider.impl.tencent_model_provider.credential.stt import TencentSTTModelCredential +from models_provider.impl.tencent_model_provider.credential.tokenhub_stt import TencentTokenhubSTTModelCredential from models_provider.impl.tencent_model_provider.credential.tti import TencentTTIModelCredential from models_provider.impl.tencent_model_provider.model.embedding import TencentEmbeddingModel from models_provider.impl.tencent_model_provider.model.image import TencentVision from models_provider.impl.tencent_model_provider.model.llm import TencentModel -from models_provider.impl.tencent_model_provider.model.stt import TencentSpeechToText +from models_provider.impl.tencent_model_provider.model.stt import TencentSpeechToText, TencentWandSpeechToText from models_provider.impl.tencent_model_provider.model.tti import TencentTextToImageModel from maxkb.conf import PROJECT_DIR from django.utils.translation import gettext as _ @@ -78,6 +79,18 @@ def _initialize_model_info(): ModelTypeConst.STT, TencentSTTModelCredential, TencentSpeechToText), + _create_model_info( + 'wand-asr-v1', + _('Tencent WAND ASR online recognition, supports long audio/video. Outputs text with sentence-level timestamps.'), + ModelTypeConst.STT, + TencentTokenhubSTTModelCredential, + TencentWandSpeechToText), + _create_model_info( + 'hy-asr-3.0-preview', + _('Tencent Hunyuan ASR 3.0 (preview) online recognition. Accepts the same request/response format as wand-asr-v1.'), + ModelTypeConst.STT, + TencentTokenhubSTTModelCredential, + TencentWandSpeechToText), ] tencent_embedding_model_info = _create_model_info( @@ -125,6 +138,12 @@ def __init__(self): def get_model_info_manage(self): return self._model_info_manage + def get_model(self, model_type, model_name, model_credential, **model_kwargs): + # STT 模型:模型名不以 asr- 开头的一律走 Tencent Tokenhub WAND 识别 + if model_type == ModelTypeConst.STT.name and not model_name.startswith('asr-'): + return TencentWandSpeechToText.new_instance(model_type, model_name, model_credential, **model_kwargs) + return super().get_model(model_type, model_name, model_credential, **model_kwargs) + def get_model_provide_info(self): icon_path = _get_tencent_icon_path() icon_data = get_file_content(icon_path)