From afa6a735444f690eb86948d7b1ad09edb56ce852 Mon Sep 17 00:00:00 2001 From: wxg0103 <727495428@qq.com> Date: Tue, 15 Sep 2026 17:28:22 +0800 Subject: [PATCH] feat: update TTS API integration with new parameters and default values --- .../credential/tts.py | 84 ++++-- .../credential/ttv.py | 8 +- .../model/tti.py | 2 +- .../model/tts.py | 260 +++++++++--------- .../model/ttv.py | 60 ++-- .../volcanic_engine_model_provider.py | 22 +- 6 files changed, 249 insertions(+), 187 deletions(-) diff --git a/apps/models_provider/impl/volcanic_engine_model_provider/credential/tts.py b/apps/models_provider/impl/volcanic_engine_model_provider/credential/tts.py index fc12cbc2ab4..1f208e7e5c2 100644 --- a/apps/models_provider/impl/volcanic_engine_model_provider/credential/tts.py +++ b/apps/models_provider/impl/volcanic_engine_model_provider/credential/tts.py @@ -15,38 +15,74 @@ class VolcanicEngineTTSModelGeneralParams(BaseForm): TooltipLabel(_("timbre"), _("Chinese sounds can support mixed scenes of Chinese and English")), required=True, default_value="zh_female_cancan_mars_bigtts", - text_field="value", + text_field="label", value_field="value", option_list=[ - {"text": "灿灿/Shiny", "value": "zh_female_cancan_mars_bigtts"}, - {"text": "清新女声", "value": "zh_female_qingxinnvsheng_mars_bigtts"}, - {"text": "爽快思思/Skye", "value": "zh_female_shuangkuaisisi_moon_bigtts"}, - {"text": "湾区大叔", "value": "zh_female_wanqudashu_moon_bigtts"}, - {"text": "呆萌川妹", "value": "zh_female_daimengchuanmei_moon_bigtts"}, - {"text": "广州德哥", "value": "zh_male_guozhoudege_moon_bigtts"}, - {"text": "北京小爷", "value": "zh_male_beijingxiaoye_moon_bigtts"}, - {"text": "少年梓辛/Brayan", "value": "zh_male_shaonianzixin_moon_bigtts"}, - {"text": "魅力女友", "value": "zh_female_meilinvyou_moon_bigtts"}, + {"label": "灿灿/Shiny", "value": "zh_female_cancan_mars_bigtts"}, + {"label": "清新女声", "value": "zh_female_qingxinnvsheng_mars_bigtts"}, + {"label": "爽快思思/Skye", "value": "zh_female_shuangkuaisisi_moon_bigtts"}, + {"label": "湾区大叔", "value": "zh_female_wanqudashu_moon_bigtts"}, + {"label": "呆萌川妹", "value": "zh_female_daimengchuanmei_moon_bigtts"}, + {"label": "广州德哥", "value": "zh_male_guozhoudege_moon_bigtts"}, + {"label": "北京小爷", "value": "zh_male_beijingxiaoye_moon_bigtts"}, + {"label": "少年梓辛/Brayan", "value": "zh_male_shaonianzixin_moon_bigtts"}, + {"label": "魅力女友", "value": "zh_female_meilinvyou_moon_bigtts"}, ], ) - speed_ratio = forms.SliderField( - TooltipLabel(_("speaking speed"), _("[0.2,3], the default is 1, usually one decimal place is enough")), + format = forms.SingleSelect( + TooltipLabel(_("audio format"), _("The streaming scenario recommends pcm")), required=True, - default_value=1, - _min=0.2, - _max=3, - _step=0.1, - precision=1, + default_value="mp3", + text_field="label", + value_field="value", + option_list=[ + {"label": "mp3", "value": "mp3"}, + {"label": "pcm", "value": "pcm"}, + {"label": "ogg_opus", "value": "ogg_opus"}, + {"label": "wav", "value": "wav"}, + ], + ) + sample_rate = forms.SingleSelect( + TooltipLabel(_("sample rate"), _("ogg_opus only supports 48000")), + required=True, + default_value=24000, + text_field="label", + value_field="value", + option_list=[ + {"label": "8000", "value": 8000}, + {"label": "16000", "value": 16000}, + {"label": "22050", "value": 22050}, + {"label": "24000", "value": 24000}, + {"label": "32000", "value": 32000}, + {"label": "44100", "value": 44100}, + {"label": "48000", "value": 48000}, + ], + ) + speech_rate = forms.SliderField( + TooltipLabel(_("speaking speed"), _("[-50,100], 100 means 2x speed, -50 means 0.5x speed")), + required=True, + default_value=0, + _min=-50, + _max=100, + _step=1, + precision=0, + ) + loudness_rate = forms.SliderField( + TooltipLabel(_("volume"), _("[-50,100], 100 means 2x volume, -50 means 0.5x volume")), + required=True, + default_value=0, + _min=-50, + _max=100, + _step=1, + precision=0, ) class VolcanicEngineTTSModelCredential(BaseForm, BaseModelCredential): - volcanic_api_url = forms.TextInputField( - "API URL", required=True, default_value="wss://openspeech.bytedance.com/api/v1/tts/ws_binary" + api_url = forms.TextInputField( + "API URL", required=True, default_value="https://openspeech.bytedance.com/api/v3/tts/unidirectional" ) - volcanic_app_id = forms.TextInputField("App ID", required=True) - volcanic_token = forms.PasswordInputField("Access Token", required=True) - volcanic_cluster = forms.TextInputField("Cluster ID", required=True) + api_key = forms.PasswordInputField("API Key", required=True) def is_valid( self, @@ -64,7 +100,7 @@ def is_valid( gettext("{model_type} Model type is not supported").format(model_type=model_type), ) - for key in ["volcanic_api_url", "volcanic_app_id", "volcanic_token", "volcanic_cluster"]: + for key in ["api_url", "api_key"]: if key not in model_credential: if raise_exception: raise AppApiException(ValidCode.valid_error.value, gettext("{key} is required").format(key=key)) @@ -89,7 +125,7 @@ def is_valid( return True def encryption_dict(self, model: Dict[str, object]): - return {**model, "volcanic_token": super().encryption(model.get("volcanic_token", ""))} + return {**model, "api_key": super().encryption(model.get("api_key", ""))} def get_model_params_setting_form(self, model_name): return VolcanicEngineTTSModelGeneralParams() diff --git a/apps/models_provider/impl/volcanic_engine_model_provider/credential/ttv.py b/apps/models_provider/impl/volcanic_engine_model_provider/credential/ttv.py index 3b68e5614d1..ea6f5f7574b 100644 --- a/apps/models_provider/impl/volcanic_engine_model_provider/credential/ttv.py +++ b/apps/models_provider/impl/volcanic_engine_model_provider/credential/ttv.py @@ -15,11 +15,11 @@ class VolcanicEngineTTVModelGeneralParams(BaseForm): resolution = SingleSelect( TooltipLabel(_("Resolution"), _("Resolution")), required=True, - default_value="480P", + default_value="480p", option_list=[ - {"value": "480P", "label": "480P"}, - {"value": "720P", "label": "720P"}, - {"value": "1080P", "label": "1080P"}, + {"value": "480p", "label": "480p"}, + {"value": "720p", "label": "720p"}, + {"value": "1080p", "label": "1080p"}, ], text_field="label", value_field="value", diff --git a/apps/models_provider/impl/volcanic_engine_model_provider/model/tti.py b/apps/models_provider/impl/volcanic_engine_model_provider/model/tti.py index a7c4e3e9eb2..432d002da9e 100644 --- a/apps/models_provider/impl/volcanic_engine_model_provider/model/tti.py +++ b/apps/models_provider/impl/volcanic_engine_model_provider/model/tti.py @@ -41,7 +41,7 @@ def new_instance(model_type, model_name, model_credential: Dict[str, object], ** return VolcanicEngineTextToImage( model_version=model_name, api_key=model_credential.get("api_key"), - api_base=model_credential.get("volcanic_api_url") or "https://ark-api.volcengine.com", + api_base=model_credential.get("volcanic_api_url") or "https://ark.cn-beijing.volces.com/api/v3", **optional_params, ) diff --git a/apps/models_provider/impl/volcanic_engine_model_provider/model/tts.py b/apps/models_provider/impl/volcanic_engine_model_provider/model/tts.py index 1008256cd69..ba28788f218 100644 --- a/apps/models_provider/impl/volcanic_engine_model_provider/model/tts.py +++ b/apps/models_provider/impl/volcanic_engine_model_provider/model/tts.py @@ -1,178 +1,164 @@ # coding=utf-8 - """ -requires Python 3.6 or later - -pip install asyncio -pip install websockets +单向流式语音合成 HTTP 接口。 +接口文档: https://docs.volcengine.com/docs/6561/2528925 """ -import asyncio -import copy -import gzip +import base64 +import codecs import json -import re -import ssl - -import uuid_utils.compat as uuid from typing import Dict +from uuid import uuid4 -import websockets +import requests from django.utils.translation import gettext as _ from common.utils.common import _remove_empty_lines +from common.utils.logger import maxkb_logger from models_provider.base_model_provider import MaxKBBaseModel from models_provider.impl.base_tts import BaseTextToSpeech -MESSAGE_TYPES = {11: "audio-only server response", 12: "frontend server response", 15: "error message from server"} -MESSAGE_TYPE_SPECIFIC_FLAGS = { - 0: "no sequence number", - 1: "sequence number > 0", - 2: "last message from server (seq < 0)", - 3: "sequence number < 0", -} -MESSAGE_SERIALIZATION_METHODS = {0: "no serialization", 1: "JSON", 15: "custom type"} -MESSAGE_COMPRESSIONS = {0: "no compression", 1: "gzip", 15: "custom compression method"} - -# version: b0001 (4 bits) -# header size: b0001 (4 bits) -# message type: b0001 (Full client request) (4bits) -# message type specific flags: b0000 (none) (4bits) -# message serialization method: b0001 (JSON) (4 bits) -# message compression: b0001 (gzip) (4bits) -# reserved data: 0x00 (1 byte) -default_header = bytearray(b"\x11\x10\x11\x00") - -ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) -ssl_context.check_hostname = False -ssl_context.verify_mode = ssl.CERT_NONE +DEFAULT_API_URL = "https://openspeech.bytedance.com/api/v3/tts/unidirectional" +DEFAULT_VOICE_TYPE = "zh_female_cancan_mars_bigtts" +DEFAULT_FORMAT = "mp3" +DEFAULT_SAMPLE_RATE = 24000 + +# audio_params 中仅在用户显式设置时才透传的字段 +OPTIONAL_AUDIO_PARAM_KEYS = ("bit_rate",) + +# 音频发送完毕后服务端返回该结束码,属于正常结束 +STREAM_END_CODE = 20000000 + +REQUEST_TIMEOUT = (10, 600) class VolcanicEngineTextToSpeech(MaxKBBaseModel, BaseTextToSpeech): - volcanic_app_id: str - volcanic_cluster: str - volcanic_api_url: str - volcanic_token: str + api_url: str + api_key: str + model_name: str params: dict def __init__(self, **kwargs): + kwargs["api_url"] = kwargs.get("api_url") or DEFAULT_API_URL + kwargs["params"] = kwargs.get("params") or {} super().__init__(**kwargs) - self.volcanic_api_url = kwargs.get("volcanic_api_url") - self.volcanic_token = kwargs.get("volcanic_token") - self.volcanic_app_id = kwargs.get("volcanic_app_id") - self.volcanic_cluster = kwargs.get("volcanic_cluster") + self.api_url = kwargs.get("api_url") + self.api_key = kwargs.get("api_key") + self.model_name = kwargs.get("model_name") self.params = kwargs.get("params") + @staticmethod + def is_cache_model(): + return False + @staticmethod def new_instance(model_type, model_name, model_credential: Dict[str, object], **model_kwargs): - optional_params = {"params": {"voice_type": "zh_female_cancan_mars_bigtts", "speed_ratio": 1.0}} + optional_params = { + "params": { + "voice_type": DEFAULT_VOICE_TYPE, + "format": DEFAULT_FORMAT, + "sample_rate": DEFAULT_SAMPLE_RATE, + "speech_rate": 0, + "loudness_rate": 0, + } + } for key, value in model_kwargs.items(): if key not in ["model_id", "use_local", "streaming"]: optional_params["params"][key] = value return VolcanicEngineTextToSpeech( - volcanic_api_url=model_credential.get("volcanic_api_url"), - volcanic_token=model_credential.get("volcanic_token"), - volcanic_app_id=model_credential.get("volcanic_app_id"), - volcanic_cluster=model_credential.get("volcanic_cluster"), + api_url=model_credential.get("api_url"), + api_key=model_credential.get("api_key"), + model_name=model_name, **optional_params, ) def check_auth(self): self.text_to_speech(_("Hello")) - def text_to_speech(self, text): - request_json = { - "app": {"appid": self.volcanic_app_id, "token": "access_token", "cluster": self.volcanic_cluster}, - "user": {"uid": "uid"}, - "audio": { - "encoding": "mp3", - "volume_ratio": 1.0, - "pitch_ratio": 1.0, - } - | self.params, - "request": {"reqid": str(uuid.uuid7()), "text": "", "text_type": "plain", "operation": "xxx"}, + def _build_audio_params(self) -> dict: + params = self.params or {} + audio_params = { + "format": params.get("format") or DEFAULT_FORMAT, + "sample_rate": int(params.get("sample_rate") or DEFAULT_SAMPLE_RATE), + "speech_rate": int(params.get("speech_rate") or 0), + "loudness_rate": int(params.get("loudness_rate") or 0), } - text = _remove_empty_lines(text) - - return asyncio.run(self.submit(request_json, text)) - - def is_cache_model(self): - return False + for key in OPTIONAL_AUDIO_PARAM_KEYS: + value = params.get(key) + if value: + audio_params[key] = int(value) + return audio_params + + def _build_req_params(self, text: str) -> dict: + params = self.params or {} + req_params = { + "text": text, + "speaker": params.get("voice_type") or DEFAULT_VOICE_TYPE, + "audio_params": self._build_audio_params(), + } + # 仅当 speaker 为复刻音色时需要指定模型版本 + if params.get("model"): + req_params["model"] = params["model"] + return req_params - def token_auth(self): - return {"Authorization": "Bearer; {}".format(self.volcanic_token)} - - async def submit(self, request_json, text): - submit_request_json = copy.deepcopy(request_json) - submit_request_json["request"]["operation"] = "submit" - header = {"Authorization": f"Bearer; {self.volcanic_token}"} - result = b"" - async with websockets.connect( - self.volcanic_api_url, additional_headers=header, ping_interval=None, ssl=ssl_context - ) as ws: - lines = [text[i : i + 200] for i in range(0, len(text), 200)] - for line in lines: - if self.is_table_format_chars_only(line): + def text_to_speech(self, text): + headers = { + "X-Api-Key": self.api_key, + # 模型ID 即接口要求的 resource id(seed-tts-2.0 / seed-icl-2.0) + "X-Api-Resource-Id": self.model_name, + "X-Api-Request-Id": str(uuid4()), + "Content-Type": "application/json", + "Connection": "keep-alive", + } + payload = {"req_params": self._build_req_params(_remove_empty_lines(text))} + audio = bytearray() + buffer = "" + decoder = codecs.getincrementaldecoder("utf-8")() + with requests.post( + self.api_url, json=payload, headers=headers, stream=True, timeout=REQUEST_TIMEOUT + ) as response: + if response.status_code != 200: + raise Exception(f"语音合成请求失败: HTTP {response.status_code}, {response.text[:500]}") + for chunk in response.iter_content(chunk_size=None): + if not chunk: continue - submit_request_json["request"]["reqid"] = str(uuid.uuid7()) - submit_request_json["request"]["text"] = line - payload_bytes = str.encode(json.dumps(submit_request_json)) - payload_bytes = gzip.compress(payload_bytes) # if no compression, comment this line - full_client_request = bytearray(default_header) - full_client_request.extend((len(payload_bytes)).to_bytes(4, "big")) # payload size(4 bytes) - full_client_request.extend(payload_bytes) # payload - await ws.send(full_client_request) - result += await self.parse_response(ws) - return result - - @staticmethod - def is_table_format_chars_only(s): - # 检查是否仅包含 "|", "-", 和空格字符 - return bool(s) and re.fullmatch(r"[|\-\s]+", s) + buffer += decoder.decode(chunk) + buffer, finished = self._consume(buffer, audio) + if finished: + break + if buffer.strip(): + maxkb_logger.warning(f"语音合成响应存在未解析内容: {buffer[:200]}") + if not audio: + raise Exception("No audio data received") + return bytes(audio) @staticmethod - async def parse_response(ws): - result = b"" - while True: - res = await ws.recv() - protocol_version = res[0] >> 4 - header_size = res[0] & 0x0F - message_type = res[1] >> 4 - message_type_specific_flags = res[1] & 0x0F - serialization_method = res[2] >> 4 - message_compression = res[2] & 0x0F - reserved = res[3] - header_extensions = res[4 : header_size * 4] - payload = res[header_size * 4 :] - if header_size != 1: - # print(f" Header extensions: {header_extensions}") - pass - if message_type == 0xB: # audio-only server response - if message_type_specific_flags == 0: # no sequence number as ACK - continue - else: - sequence_number = int.from_bytes(payload[:4], "big", signed=True) - payload_size = int.from_bytes(payload[4:8], "big", signed=False) - payload = payload[8:] - result += payload - if sequence_number < 0: - break - else: - continue - elif message_type == 0xF: - code = int.from_bytes(payload[:4], "big", signed=False) - msg_size = int.from_bytes(payload[4:8], "big", signed=False) - error_msg = payload[8:] - if message_compression == 1: - error_msg = gzip.decompress(error_msg) - error_msg = str(error_msg, "utf-8") - raise Exception(f"Error code: {code}, message: {error_msg}") - elif message_type == 0xC: - msg_size = int.from_bytes(payload[:4], "big", signed=False) - payload = payload[4:] - if message_compression == 1: - payload = gzip.decompress(payload) - else: + def _consume(content: str, audio: bytearray) -> tuple: + """解析缓冲区中已完整的 JSON 分片,把 base64 音频累加到 audio。 + + 返回 (未解析完的剩余内容, 是否已收到结束标记) + """ + decoder = json.JSONDecoder() + index = 0 + length = len(content) + while index < length: + while index < length and content[index] in " \r\n\t": + index += 1 + if index >= length: + break + try: + chunk, end = decoder.raw_decode(content, index) + except ValueError: + # 分片不完整,等待后续内容 break - return result + code = chunk.get("code", 0) + if code == STREAM_END_CODE: + return "", True + if code > 0: + raise Exception(f"Error code: {code}, message: {chunk.get('message')}") + data = chunk.get("data") + if data: + audio.extend(base64.b64decode(data)) + index = end + return content[index:], False diff --git a/apps/models_provider/impl/volcanic_engine_model_provider/model/ttv.py b/apps/models_provider/impl/volcanic_engine_model_provider/model/ttv.py index ee42ceb113a..7693dbca331 100644 --- a/apps/models_provider/impl/volcanic_engine_model_provider/model/ttv.py +++ b/apps/models_provider/impl/volcanic_engine_model_provider/model/ttv.py @@ -5,6 +5,28 @@ from common.utils.logger import maxkb_logger from volcenginesdkarkruntime import Ark +# 视频生成接口支持的直接传参字段 +# 文档: https://www.volcengine.com/docs/82379/1520758 +VIDEO_PARAM_KEYS = ( + "resolution", + "ratio", + "duration", + "frames", + "watermark", + "camera_fixed", + "seed", + "generate_audio", + "draft", + "return_last_frame", + "service_tier", + "callback_url", + "execution_expires_after", + "priority", + "safety_identifier", +) + +INT_PARAM_KEYS = ("duration", "frames", "seed", "execution_expires_after", "priority") + class GenerationVideoModel(MaxKBBaseModel, BaseGenerationVideo): api_key: str @@ -42,20 +64,23 @@ def new_instance(model_type, model_name, model_credential: Dict[str, object], ** def check_auth(self): return True - def _build_prompt(self, prompt: str) -> str: - """拼接参数到 prompt 文本""" - param_map = { - "ratio": "rt", - "duration": "dur", - "framespersecond": "fps", - "resolution": "rs", - "watermark": "wm", - "camerafixed": "cf", - } - for key, value in self.params.items(): - if key in param_map: - prompt += f" --{param_map[key]} {value}" - return prompt + def _build_params(self) -> dict: + """把参数转换为视频生成接口的顶层字段,接口不支持的字段放入 extra_body""" + params = {} + extra_body = {} + for key, value in (self.params or {}).items(): + if value is None or value == "": + continue + name = str(key).replace(" ", "_").lower() + if name == "camerafixed": + name = "camera_fixed" + if name in VIDEO_PARAM_KEYS: + params[name] = int(value) if name in INT_PARAM_KEYS else value + else: + extra_body[name] = value + if extra_body: + params["extra_body"] = extra_body + return params def _poll_task(self, client: Ark, task_id: str, interval: int = 30): """轮询任务状态,直到完成""" @@ -72,9 +97,6 @@ def _poll_task(self, client: Ark, task_id: str, interval: int = 30): # --- 通用异步生成函数 --- def generate_video(self, prompt, negative_prompt=None, first_frame_url=None, last_frame_url=None, **kwargs): client = Ark(api_key=self.api_key, base_url=self.base_url) - # 根据params设置其他参数 豆包的参数和别的不一样 需要拼接在text里 - # --rt 16:9 --dur 5 --fps 24 --rs 720p --wm true --cf false - prompt = self._build_prompt(prompt) content = [{"type": "text", "text": prompt}] if first_frame_url: @@ -82,7 +104,7 @@ def generate_video(self, prompt, negative_prompt=None, first_frame_url=None, las if last_frame_url: content.append({"type": "image_url", "image_url": {"url": last_frame_url}, "role": "last_frame"}) - task = client.content_generation.tasks.create(model=self.model_name, content=content) + task = client.content_generation.tasks.create(model=self.model_name, content=content, **self._build_params()) task_id = task.id maxkb_logger.info(f"[ArkVideo] Created task {task_id}") @@ -98,5 +120,5 @@ def generate_video(self, prompt, negative_prompt=None, first_frame_url=None, las except Exception as e: maxkb_logger.error(f"[ArkVideo] Failed to delete task {task_id}: {e}") raise e - maxkb_logger.info("视频地址", result.content.video_url) + maxkb_logger.info(f"[ArkVideo] 视频地址 {result.content.video_url}") return result.content.video_url diff --git a/apps/models_provider/impl/volcanic_engine_model_provider/volcanic_engine_model_provider.py b/apps/models_provider/impl/volcanic_engine_model_provider/volcanic_engine_model_provider.py index ba27b890cce..f486d4e446e 100644 --- a/apps/models_provider/impl/volcanic_engine_model_provider/volcanic_engine_model_provider.py +++ b/apps/models_provider/impl/volcanic_engine_model_provider/volcanic_engine_model_provider.py @@ -70,7 +70,6 @@ ModelInfo( "bigmodel", "", ModelTypeConst.STT, volcanic_engine_big_stt_model_credential, VolcanicEngineBigModelSpeechToText ), - ModelInfo("tts", "", ModelTypeConst.TTS, volcanic_engine_tts_model_credential, VolcanicEngineTextToSpeech), ModelInfo( "doubao-seedream-3-0-t2i-250415", _(""), @@ -80,6 +79,24 @@ ), ] +# TTS 的模型ID 即接口的 resource id +model_info_tts_list = [ + ModelInfo( + "seed-tts-2.0", + _(""), + ModelTypeConst.TTS, + volcanic_engine_tts_model_credential, + VolcanicEngineTextToSpeech, + ), + ModelInfo( + "seed-icl-2.0", + _(""), + ModelTypeConst.TTS, + volcanic_engine_tts_model_credential, + VolcanicEngineTextToSpeech, + ), +] + open_ai_embedding_credential = VolcanicEmbeddingCredential() model_info_embedding_list = [ ModelInfo( @@ -113,7 +130,8 @@ .append_default_model_info(model_info_list[2]) .append_default_model_info(model_info_list[3]) .append_default_model_info(model_info_list[4]) - .append_default_model_info(model_info_list[5]) + .append_model_info_list(model_info_tts_list) + .append_default_model_info(model_info_tts_list[0]) .append_model_info_list(model_info_embedding_list) .append_default_model_info(model_info_embedding_list[0]) .append_model_info_list(model_info_ttv_list)