mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-25 15:04:49 -06:00
feat: mistral tts
This commit is contained in:
@@ -3940,6 +3940,18 @@ AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT = PersistentConfig(
|
||||
os.getenv('AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT', 'audio-24khz-160kbitrate-mono-mp3'),
|
||||
)
|
||||
|
||||
AUDIO_TTS_MISTRAL_API_KEY = PersistentConfig(
|
||||
'AUDIO_TTS_MISTRAL_API_KEY',
|
||||
'audio.tts.mistral.api_key',
|
||||
os.getenv('AUDIO_TTS_MISTRAL_API_KEY', ''),
|
||||
)
|
||||
|
||||
AUDIO_TTS_MISTRAL_API_BASE_URL = PersistentConfig(
|
||||
'AUDIO_TTS_MISTRAL_API_BASE_URL',
|
||||
'audio.tts.mistral.api_base_url',
|
||||
os.getenv('AUDIO_TTS_MISTRAL_API_BASE_URL', 'https://api.mistral.ai/v1'),
|
||||
)
|
||||
|
||||
|
||||
####################################
|
||||
# LDAP
|
||||
|
||||
@@ -212,6 +212,8 @@ from open_webui.config import (
|
||||
AUDIO_TTS_AZURE_SPEECH_REGION,
|
||||
AUDIO_TTS_AZURE_SPEECH_BASE_URL,
|
||||
AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT,
|
||||
AUDIO_TTS_MISTRAL_API_KEY,
|
||||
AUDIO_TTS_MISTRAL_API_BASE_URL,
|
||||
PLAYWRIGHT_WS_URL,
|
||||
PLAYWRIGHT_TIMEOUT,
|
||||
FIRECRAWL_API_BASE_URL,
|
||||
@@ -1284,6 +1286,9 @@ app.state.config.TTS_AZURE_SPEECH_REGION = AUDIO_TTS_AZURE_SPEECH_REGION
|
||||
app.state.config.TTS_AZURE_SPEECH_BASE_URL = AUDIO_TTS_AZURE_SPEECH_BASE_URL
|
||||
app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT = AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT
|
||||
|
||||
app.state.config.TTS_MISTRAL_API_KEY = AUDIO_TTS_MISTRAL_API_KEY
|
||||
app.state.config.TTS_MISTRAL_API_BASE_URL = AUDIO_TTS_MISTRAL_API_BASE_URL
|
||||
|
||||
|
||||
app.state.faster_whisper_model = None
|
||||
app.state.speech_synthesiser = None
|
||||
|
||||
@@ -168,6 +168,8 @@ class TTSConfigForm(BaseModel):
|
||||
AZURE_SPEECH_REGION: str
|
||||
AZURE_SPEECH_BASE_URL: str
|
||||
AZURE_SPEECH_OUTPUT_FORMAT: str
|
||||
MISTRAL_API_KEY: str
|
||||
MISTRAL_API_BASE_URL: str
|
||||
|
||||
|
||||
class STTConfigForm(BaseModel):
|
||||
@@ -208,6 +210,8 @@ async def get_audio_config(request: Request, user=Depends(get_admin_user)):
|
||||
'AZURE_SPEECH_REGION': request.app.state.config.TTS_AZURE_SPEECH_REGION,
|
||||
'AZURE_SPEECH_BASE_URL': request.app.state.config.TTS_AZURE_SPEECH_BASE_URL,
|
||||
'AZURE_SPEECH_OUTPUT_FORMAT': request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT,
|
||||
'MISTRAL_API_KEY': request.app.state.config.TTS_MISTRAL_API_KEY,
|
||||
'MISTRAL_API_BASE_URL': request.app.state.config.TTS_MISTRAL_API_BASE_URL,
|
||||
},
|
||||
'stt': {
|
||||
'OPENAI_API_BASE_URL': request.app.state.config.STT_OPENAI_API_BASE_URL,
|
||||
@@ -242,6 +246,8 @@ async def update_audio_config(request: Request, form_data: AudioConfigUpdateForm
|
||||
request.app.state.config.TTS_AZURE_SPEECH_REGION = form_data.tts.AZURE_SPEECH_REGION
|
||||
request.app.state.config.TTS_AZURE_SPEECH_BASE_URL = form_data.tts.AZURE_SPEECH_BASE_URL
|
||||
request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT = form_data.tts.AZURE_SPEECH_OUTPUT_FORMAT
|
||||
request.app.state.config.TTS_MISTRAL_API_KEY = form_data.tts.MISTRAL_API_KEY
|
||||
request.app.state.config.TTS_MISTRAL_API_BASE_URL = form_data.tts.MISTRAL_API_BASE_URL
|
||||
|
||||
request.app.state.config.STT_OPENAI_API_BASE_URL = form_data.stt.OPENAI_API_BASE_URL
|
||||
request.app.state.config.STT_OPENAI_API_KEY = form_data.stt.OPENAI_API_KEY
|
||||
@@ -280,6 +286,8 @@ async def update_audio_config(request: Request, form_data: AudioConfigUpdateForm
|
||||
'AZURE_SPEECH_REGION': request.app.state.config.TTS_AZURE_SPEECH_REGION,
|
||||
'AZURE_SPEECH_BASE_URL': request.app.state.config.TTS_AZURE_SPEECH_BASE_URL,
|
||||
'AZURE_SPEECH_OUTPUT_FORMAT': request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT,
|
||||
'MISTRAL_API_KEY': request.app.state.config.TTS_MISTRAL_API_KEY,
|
||||
'MISTRAL_API_BASE_URL': request.app.state.config.TTS_MISTRAL_API_BASE_URL,
|
||||
},
|
||||
'stt': {
|
||||
'OPENAI_API_BASE_URL': request.app.state.config.STT_OPENAI_API_BASE_URL,
|
||||
@@ -551,6 +559,76 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
|
||||
return FileResponse(file_path)
|
||||
|
||||
elif request.app.state.config.TTS_ENGINE == 'mistral':
|
||||
api_key = request.app.state.config.TTS_MISTRAL_API_KEY
|
||||
api_base_url = request.app.state.config.TTS_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1'
|
||||
|
||||
if not api_key:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail='Mistral API key is required for Mistral TTS',
|
||||
)
|
||||
|
||||
try:
|
||||
timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT)
|
||||
async with aiohttp.ClientSession(timeout=timeout, trust_env=True) as session:
|
||||
mistral_payload = {
|
||||
'input': payload.get('input', ''),
|
||||
'model': request.app.state.config.TTS_MODEL or 'mistral-tts-latest',
|
||||
'voice_id': payload.get('voice', ''),
|
||||
'response_format': 'mp3',
|
||||
}
|
||||
|
||||
r = await session.post(
|
||||
url=f'{api_base_url}/audio/speech',
|
||||
json=mistral_payload,
|
||||
headers={
|
||||
'Content-Type': 'application/json',
|
||||
'Authorization': f'Bearer {api_key}',
|
||||
},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
)
|
||||
|
||||
r.raise_for_status()
|
||||
|
||||
res = await r.json()
|
||||
audio_data = res.get('audio_data', '')
|
||||
if not audio_data:
|
||||
raise ValueError('No audio_data in Mistral TTS response')
|
||||
|
||||
audio_bytes = base64.b64decode(audio_data)
|
||||
|
||||
async with aiofiles.open(file_path, 'wb') as f:
|
||||
await f.write(audio_bytes)
|
||||
|
||||
async with aiofiles.open(file_body_path, 'w') as f:
|
||||
await f.write(json.dumps(payload))
|
||||
|
||||
return FileResponse(file_path)
|
||||
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
detail = None
|
||||
|
||||
status_code = 500
|
||||
detail = 'Open WebUI: Server Connection Error'
|
||||
|
||||
if r is not None:
|
||||
status_code = r.status
|
||||
|
||||
try:
|
||||
res = await r.json()
|
||||
if 'error' in res:
|
||||
detail = f'External: {res["error"]}'
|
||||
elif 'message' in res:
|
||||
detail = f'External: {res["message"]}'
|
||||
except Exception:
|
||||
detail = f'External: {e}'
|
||||
|
||||
raise HTTPException(
|
||||
status_code=status_code,
|
||||
detail=detail,
|
||||
)
|
||||
|
||||
def transcription_handler(request, file_path, metadata, user=None):
|
||||
filename = os.path.basename(file_path)
|
||||
@@ -1238,6 +1316,8 @@ def get_available_models(request: Request) -> list[dict]:
|
||||
available_models = [{'name': model['name'], 'id': model['model_id']} for model in models]
|
||||
except requests.RequestException as e:
|
||||
log.error(f'Error fetching voices: {str(e)}')
|
||||
elif request.app.state.config.TTS_ENGINE == 'mistral':
|
||||
available_models = [{'id': 'mistral-tts-latest'}]
|
||||
return available_models
|
||||
|
||||
|
||||
@@ -1301,6 +1381,29 @@ def get_available_voices(request) -> dict:
|
||||
available_voices[voice['ShortName']] = f'{voice["DisplayName"]} ({voice["ShortName"]})'
|
||||
except requests.RequestException as e:
|
||||
log.error(f'Error fetching voices: {str(e)}')
|
||||
elif request.app.state.config.TTS_ENGINE == 'mistral':
|
||||
api_key = request.app.state.config.TTS_MISTRAL_API_KEY
|
||||
api_base_url = request.app.state.config.TTS_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1'
|
||||
|
||||
if api_key:
|
||||
try:
|
||||
response = requests.get(
|
||||
f'{api_base_url}/audio/voices',
|
||||
headers={
|
||||
'Authorization': f'Bearer {api_key}',
|
||||
},
|
||||
timeout=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST,
|
||||
)
|
||||
response.raise_for_status()
|
||||
voices_data = response.json()
|
||||
|
||||
for voice in voices_data:
|
||||
voice_id = voice.get('voice_id', voice.get('id', ''))
|
||||
voice_name = voice.get('name', voice_id)
|
||||
if voice_id:
|
||||
available_voices[voice_id] = voice_name
|
||||
except requests.RequestException as e:
|
||||
log.error(f'Error fetching Mistral voices: {str(e)}')
|
||||
|
||||
return available_voices
|
||||
|
||||
|
||||
@@ -37,6 +37,8 @@
|
||||
let TTS_AZURE_SPEECH_REGION = '';
|
||||
let TTS_AZURE_SPEECH_BASE_URL = '';
|
||||
let TTS_AZURE_SPEECH_OUTPUT_FORMAT = '';
|
||||
let TTS_MISTRAL_API_KEY = '';
|
||||
let TTS_MISTRAL_API_BASE_URL = '';
|
||||
|
||||
let STT_OPENAI_API_BASE_URL = '';
|
||||
let STT_OPENAI_API_KEY = '';
|
||||
@@ -124,6 +126,8 @@
|
||||
AZURE_SPEECH_REGION: TTS_AZURE_SPEECH_REGION,
|
||||
AZURE_SPEECH_BASE_URL: TTS_AZURE_SPEECH_BASE_URL,
|
||||
AZURE_SPEECH_OUTPUT_FORMAT: TTS_AZURE_SPEECH_OUTPUT_FORMAT,
|
||||
MISTRAL_API_KEY: TTS_MISTRAL_API_KEY,
|
||||
MISTRAL_API_BASE_URL: TTS_MISTRAL_API_BASE_URL,
|
||||
SPLIT_ON: TTS_SPLIT_ON
|
||||
},
|
||||
stt: {
|
||||
@@ -176,6 +180,8 @@
|
||||
TTS_AZURE_SPEECH_REGION = res.tts.AZURE_SPEECH_REGION;
|
||||
TTS_AZURE_SPEECH_BASE_URL = res.tts.AZURE_SPEECH_BASE_URL;
|
||||
TTS_AZURE_SPEECH_OUTPUT_FORMAT = res.tts.AZURE_SPEECH_OUTPUT_FORMAT;
|
||||
TTS_MISTRAL_API_KEY = res.tts.MISTRAL_API_KEY;
|
||||
TTS_MISTRAL_API_BASE_URL = res.tts.MISTRAL_API_BASE_URL;
|
||||
|
||||
STT_OPENAI_API_BASE_URL = res.stt.OPENAI_API_BASE_URL;
|
||||
STT_OPENAI_API_KEY = res.stt.OPENAI_API_KEY;
|
||||
@@ -517,6 +523,9 @@
|
||||
if (e.target?.value === 'openai') {
|
||||
TTS_VOICE = 'alloy';
|
||||
TTS_MODEL = 'tts-1';
|
||||
} else if (e.target?.value === 'mistral') {
|
||||
TTS_VOICE = '';
|
||||
TTS_MODEL = 'mistral-tts-latest';
|
||||
} else {
|
||||
TTS_VOICE = '';
|
||||
TTS_MODEL = '';
|
||||
@@ -528,6 +537,7 @@
|
||||
<option value="openai">{$i18n.t('OpenAI')}</option>
|
||||
<option value="elevenlabs">{$i18n.t('ElevenLabs')}</option>
|
||||
<option value="azure">{$i18n.t('Azure AI Speech')}</option>
|
||||
<option value="mistral">{$i18n.t('MistralAI')}</option>
|
||||
</select>
|
||||
</div>
|
||||
</div>
|
||||
@@ -585,6 +595,19 @@
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
{:else if TTS_ENGINE === 'mistral'}
|
||||
<div>
|
||||
<div class="mt-1 flex gap-2 mb-1">
|
||||
<input
|
||||
class="flex-1 w-full bg-transparent outline-hidden"
|
||||
placeholder={$i18n.t('API Base URL')}
|
||||
bind:value={TTS_MISTRAL_API_BASE_URL}
|
||||
required
|
||||
/>
|
||||
|
||||
<SensitiveInput placeholder={$i18n.t('API Key')} bind:value={TTS_MISTRAL_API_KEY} />
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
|
||||
<div class="mb-2">
|
||||
@@ -791,6 +814,47 @@
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
{:else if TTS_ENGINE === 'mistral'}
|
||||
<div class=" flex gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" mb-1.5 text-xs font-medium">{$i18n.t('TTS Voice')}</div>
|
||||
<div class="flex w-full">
|
||||
<div class="flex-1">
|
||||
<input
|
||||
list="voice-list"
|
||||
class="w-full rounded-lg py-2 px-4 text-sm bg-gray-50 dark:text-gray-300 dark:bg-gray-850 outline-hidden"
|
||||
bind:value={TTS_VOICE}
|
||||
placeholder={$i18n.t('Select a voice')}
|
||||
/>
|
||||
|
||||
<datalist id="voice-list">
|
||||
{#each voices as voice}
|
||||
<option value={voice.id}>{voice.name}</option>
|
||||
{/each}
|
||||
</datalist>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class=" mb-1.5 text-xs font-medium">{$i18n.t('TTS Model')}</div>
|
||||
<div class="flex w-full">
|
||||
<div class="flex-1">
|
||||
<input
|
||||
list="tts-model-list"
|
||||
class="w-full rounded-lg py-2 px-4 text-sm bg-gray-50 dark:text-gray-300 dark:bg-gray-850 outline-hidden"
|
||||
bind:value={TTS_MODEL}
|
||||
placeholder={$i18n.t('Select a model')}
|
||||
/>
|
||||
|
||||
<datalist id="tts-model-list">
|
||||
{#each models as model}
|
||||
<option value={model.id} class="bg-gray-50 dark:bg-gray-700" />
|
||||
{/each}
|
||||
</datalist>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
</div>
|
||||
|
||||
|
||||
Reference in New Issue
Block a user