diff --git a/app/controller/add_websocket_voice_generate_route.py b/app/controller/add_websocket_voice_generate_route.py index 5c08e34..a35bee1 100644 --- a/app/controller/add_websocket_voice_generate_route.py +++ b/app/controller/add_websocket_voice_generate_route.py @@ -3,6 +3,7 @@ import threading from io import BufferedWriter from fastapi import FastAPI, Query +from pydantic import BaseModel, Field from starlette.websockets import WebSocket import dashscope @@ -14,14 +15,10 @@ dashscope.api_key = env.llm_key_bailian def add_websocket_voice_generate_route(app: FastAPI): - @app.post("/ws_voice_generate") - async def ws_voice_generate(): + @app.post("/voice_generate") + async def voice_generate(body: VoiceGenerateParamSchema): - # dashscope.api_key = "apiKey" - result = SpeechSynthesizer.call(model='sambert-zhichu-v1', - text='今天天气怎么样', - sample_rate=16000, - format='mp3') + result = SpeechSynthesizer.call(**body.model_dump()) if result.get_audio_data() is not None: with open('output.mp3', 'wb') as f: f.write(result.get_audio_data()) @@ -29,3 +26,12 @@ def add_websocket_voice_generate_route(app: FastAPI): (sys.getsizeof(result.get_audio_data()))) else: print('ERROR: response is %s' % (result.get_response())) + + +class VoiceGenerateParamSchema(BaseModel): + text: str = Field(..., description="要生成的文本") + model: str = Field(default="sambert-zhichu-v1", description="模型名称") + sample_rate: int = Field(default=16000, description="采样率") + format: str = Field(default="mp3", description="音频格式") + volume: int = Field(default=1, description="音量大小") + rate: int = Field(default=1, description="语速")