refac: audio

This commit is contained in:
Timothy Jaeryang Baek
2026-05-19 21:24:58 +04:00
parent c306a7e16e
commit 2ca91ceeec
+325 -421
View File
@@ -8,13 +8,11 @@ import logging
import mimetypes
import os
import uuid
from concurrent.futures import ThreadPoolExecutor
from fnmatch import fnmatch
from typing import Optional
import aiofiles
import aiohttp
import requests
from fastapi import (
APIRouter,
Depends,
@@ -687,7 +685,230 @@ async def speech(request: Request, user=Depends(get_verified_user)):
)
def transcription_handler(request, file_path, metadata, user=None):
async def _transcribe_whisper(request, file_path, languages, file_dir, id):
if request.app.state.faster_whisper_model is None:
request.app.state.faster_whisper_model = set_faster_whisper_model(request.app.state.config.WHISPER_MODEL)
model = request.app.state.faster_whisper_model
def _run():
segments, info = model.transcribe(
file_path, beam_size=5, vad_filter=WHISPER_VAD_FILTER,
language=languages[0], multilingual=WHISPER_MULTILINGUAL,
)
log.info("Detected language '%s' with probability %f" % (info.language, info.language_probability))
return ''.join([segment.text for segment in list(segments)])
transcript = await asyncio.to_thread(_run)
data = {'text': transcript.strip()}
async with aiofiles.open(os.path.join(file_dir, f'{id}.json'), 'w') as f:
await f.write(json.dumps(data))
log.debug(data)
return data
async def _transcribe_openai(request, file_path, filename, languages, file_dir, id, user=None):
r = None
try:
timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT)
async with aiohttp.ClientSession(timeout=timeout, trust_env=True) as session:
for language in languages:
payload = {'model': request.app.state.config.STT_MODEL}
if language:
payload['language'] = language
headers = {'Authorization': f'Bearer {request.app.state.config.STT_OPENAI_API_KEY}'}
if user and ENABLE_FORWARD_USER_INFO_HEADERS:
headers = include_user_info_headers(headers, user)
form_data = aiohttp.FormData()
for key, value in payload.items():
form_data.add_field(key, str(value))
form_data.add_field('file', open(file_path, 'rb'), filename=filename)
r = await session.post(
url=f'{request.app.state.config.STT_OPENAI_API_BASE_URL}/audio/transcriptions',
headers=headers, data=form_data, ssl=AIOHTTP_CLIENT_SESSION_SSL,
)
if r.status == 200:
break
r.raise_for_status()
data = await r.json()
async with aiofiles.open(os.path.join(file_dir, f'{id}.json'), 'w') as f:
await f.write(json.dumps(data))
return data
except Exception as e:
log.exception(e)
detail = None
if r is not None:
try:
res = await r.json()
if 'error' in res:
detail = f'External: {res["error"].get("message", "")}'
except Exception:
detail = f'External: {e}'
raise Exception(detail if detail else 'Open WebUI: Server Connection Error')
async def _transcribe_deepgram(request, file_path, languages, file_dir, id):
r = None
try:
mime, _ = mimetypes.guess_type(file_path)
if not mime:
mime = 'audio/wav'
async with aiofiles.open(file_path, 'rb') as f:
file_data = await f.read()
headers = {
'Authorization': f'Token {request.app.state.config.DEEPGRAM_API_KEY}',
'Content-Type': mime,
}
timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT)
async with aiohttp.ClientSession(timeout=timeout, trust_env=True) as session:
for language in languages:
params = {}
if request.app.state.config.STT_MODEL:
params['model'] = request.app.state.config.STT_MODEL
if language:
params['language'] = language
r = await session.post(
'https://api.deepgram.com/v1/listen?smart_format=true',
headers=headers, params=params, data=file_data,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
)
if r.status == 200:
break
r.raise_for_status()
response_data = await r.json()
try:
transcript = response_data['results']['channels'][0]['alternatives'][0].get('transcript', '')
except (KeyError, IndexError) as e:
log.error(f'Malformed response from Deepgram: {str(e)}')
raise Exception('Failed to parse Deepgram response - unexpected response format')
data = {'text': transcript.strip()}
async with aiofiles.open(os.path.join(file_dir, f'{id}.json'), 'w') as f:
await f.write(json.dumps(data))
return data
except Exception as e:
log.exception(e)
detail = None
if r is not None:
try:
res = await r.json()
if 'error' in res:
detail = f'External: {res["error"].get("message", "")}'
except Exception:
detail = f'External: {e}'
raise Exception(detail if detail else 'Open WebUI: Server Connection Error')
async def _transcribe_azure(request, file_path, filename, file_dir, id):
if not os.path.exists(file_path):
raise HTTPException(status_code=400, detail='Audio file not found')
file_size = os.path.getsize(file_path)
if file_size > AZURE_MAX_FILE_SIZE:
raise HTTPException(
status_code=400,
detail=f"File size exceeds Azure's limit of {AZURE_MAX_FILE_SIZE_MB}MB",
)
api_key = request.app.state.config.AUDIO_STT_AZURE_API_KEY
region = request.app.state.config.AUDIO_STT_AZURE_REGION or 'eastus'
locales = request.app.state.config.AUDIO_STT_AZURE_LOCALES
base_url = request.app.state.config.AUDIO_STT_AZURE_BASE_URL
max_speakers = request.app.state.config.AUDIO_STT_AZURE_MAX_SPEAKERS or 3
if len(locales) < 2:
locales = ','.join([
'en-US', 'es-ES', 'es-MX', 'fr-FR', 'hi-IN', 'it-IT', 'de-DE',
'en-GB', 'en-IN', 'ja-JP', 'ko-KR', 'pt-BR', 'zh-CN',
])
if not api_key or not region:
raise HTTPException(status_code=400, detail='Azure API key is required for Azure STT')
r = None
try:
definition = json.dumps(
{'locales': locales.split(','), 'diarization': {'maxSpeakers': max_speakers, 'enabled': True}}
if locales else {}
)
url = (
base_url or f'https://{region}.api.cognitive.microsoft.com'
) + '/speechtotext/transcriptions:transcribe?api-version=2024-11-15'
form_data = aiohttp.FormData()
form_data.add_field('definition', definition)
form_data.add_field('audio', open(file_path, 'rb'), filename=filename)
timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT)
async with aiohttp.ClientSession(timeout=timeout, trust_env=True) as session:
r = await session.post(
url=url, data=form_data,
headers={'Ocp-Apim-Subscription-Key': api_key},
ssl=AIOHTTP_CLIENT_SESSION_SSL,
)
r.raise_for_status()
response = await r.json()
if not response.get('combinedPhrases'):
raise ValueError('No transcription found in response')
transcript = response['combinedPhrases'][0].get('text', '').strip()
if not transcript:
raise ValueError('Empty transcript in response')
data = {'text': transcript}
async with aiofiles.open(os.path.join(file_dir, f'{id}.json'), 'w') as f:
await f.write(json.dumps(data))
log.debug(data)
return data
except (KeyError, IndexError, ValueError) as e:
log.exception('Error parsing Azure response')
raise HTTPException(status_code=500, detail=f'Failed to parse Azure response: {str(e)}')
except aiohttp.ClientResponseError as e:
log.exception(e)
detail = None
try:
if r is not None and r.status != 200:
res = await r.json()
if 'code' in res and 'message' in res:
azure_code = res.get('innerError', {}).get('code', res['code'])
user_facing_codes = {
'EmptyAudioFile', 'AudioLengthLimitExceeded',
'NoLanguageIdentified', 'MultipleLanguagesIdentified',
}
if azure_code in user_facing_codes:
detail = res['message']
else:
log.error(f'Azure STT error [{azure_code}]: {res["message"]}')
detail = 'An error occurred during transcription.'
elif 'error' in res:
detail = f'External: {res["error"].get("message", "")}'
except Exception:
detail = f'External: {e}'
raise HTTPException(
status_code=e.status if e.status else 500,
detail=detail if detail else 'Open WebUI: Server Connection Error',
)
async def transcription_handler(request, file_path, metadata, user=None):
filename = os.path.basename(file_path)
file_dir = os.path.dirname(file_path)
id = filename.split('.')[0]
@@ -700,323 +921,48 @@ def transcription_handler(request, file_path, metadata, user=None):
]
if request.app.state.config.STT_ENGINE == '':
if request.app.state.faster_whisper_model is None:
request.app.state.faster_whisper_model = set_faster_whisper_model(request.app.state.config.WHISPER_MODEL)
model = request.app.state.faster_whisper_model
segments, info = model.transcribe(
file_path,
beam_size=5,
vad_filter=WHISPER_VAD_FILTER,
language=languages[0],
multilingual=WHISPER_MULTILINGUAL,
)
log.info("Detected language '%s' with probability %f" % (info.language, info.language_probability))
transcript = ''.join([segment.text for segment in list(segments)])
data = {'text': transcript.strip()}
# save the transcript to a json file
transcript_file = os.path.join(file_dir, f'{id}.json')
with open(transcript_file, 'w') as f:
json.dump(data, f)
log.debug(data)
return data
return await _transcribe_whisper(request, file_path, languages, file_dir, id)
elif request.app.state.config.STT_ENGINE == 'openai':
r = None
try:
for language in languages:
payload = {
'model': request.app.state.config.STT_MODEL,
}
if language:
payload['language'] = language
headers = {'Authorization': f'Bearer {request.app.state.config.STT_OPENAI_API_KEY}'}
if user and ENABLE_FORWARD_USER_INFO_HEADERS:
headers = include_user_info_headers(headers, user)
with open(file_path, 'rb') as audio_file:
r = requests.post(
url=f'{request.app.state.config.STT_OPENAI_API_BASE_URL}/audio/transcriptions',
headers=headers,
files={'file': (filename, audio_file)},
data=payload,
timeout=AIOHTTP_CLIENT_TIMEOUT,
)
if r.status_code == 200:
# Successful transcription
break
r.raise_for_status()
data = r.json()
# save the transcript to a json file
transcript_file = os.path.join(file_dir, f'{id}.json')
with open(transcript_file, 'w') as f:
json.dump(data, f)
return data
except Exception as e:
log.exception(e)
detail = None
if r is not None:
try:
res = r.json()
if 'error' in res:
detail = f'External: {res["error"].get("message", "")}'
except Exception:
detail = f'External: {e}'
raise Exception(detail if detail else 'Open WebUI: Server Connection Error')
return await _transcribe_openai(request, file_path, filename, languages, file_dir, id, user)
elif request.app.state.config.STT_ENGINE == 'deepgram':
try:
# Determine the MIME type of the file
mime, _ = mimetypes.guess_type(file_path)
if not mime:
mime = 'audio/wav' # fallback to wav if undetectable
# Read the audio file
with open(file_path, 'rb') as f:
file_data = f.read()
# Build headers and parameters
headers = {
'Authorization': f'Token {request.app.state.config.DEEPGRAM_API_KEY}',
'Content-Type': mime,
}
for language in languages:
params = {}
if request.app.state.config.STT_MODEL:
params['model'] = request.app.state.config.STT_MODEL
if language:
params['language'] = language
# Make request to Deepgram API
r = requests.post(
'https://api.deepgram.com/v1/listen?smart_format=true',
headers=headers,
params=params,
data=file_data,
timeout=AIOHTTP_CLIENT_TIMEOUT,
)
if r.status_code == 200:
# Successful transcription
break
r.raise_for_status()
response_data = r.json()
# Extract transcript from Deepgram response
try:
transcript = response_data['results']['channels'][0]['alternatives'][0].get('transcript', '')
except (KeyError, IndexError) as e:
log.error(f'Malformed response from Deepgram: {str(e)}')
raise Exception('Failed to parse Deepgram response - unexpected response format')
data = {'text': transcript.strip()}
# Save transcript
transcript_file = os.path.join(file_dir, f'{id}.json')
with open(transcript_file, 'w') as f:
json.dump(data, f)
return data
except Exception as e:
log.exception(e)
detail = None
if r is not None:
try:
res = r.json()
if 'error' in res:
detail = f'External: {res["error"].get("message", "")}'
except Exception:
detail = f'External: {e}'
raise Exception(detail if detail else 'Open WebUI: Server Connection Error')
return await _transcribe_deepgram(request, file_path, languages, file_dir, id)
elif request.app.state.config.STT_ENGINE == 'azure':
# Check file exists and size
if not os.path.exists(file_path):
raise HTTPException(status_code=400, detail='Audio file not found')
# Check file size (Azure has a larger limit of 200MB)
file_size = os.path.getsize(file_path)
if file_size > AZURE_MAX_FILE_SIZE:
raise HTTPException(
status_code=400,
detail=f"File size exceeds Azure's limit of {AZURE_MAX_FILE_SIZE_MB}MB",
)
api_key = request.app.state.config.AUDIO_STT_AZURE_API_KEY
region = request.app.state.config.AUDIO_STT_AZURE_REGION or 'eastus'
locales = request.app.state.config.AUDIO_STT_AZURE_LOCALES
base_url = request.app.state.config.AUDIO_STT_AZURE_BASE_URL
max_speakers = request.app.state.config.AUDIO_STT_AZURE_MAX_SPEAKERS or 3
# IF NO LOCALES, USE DEFAULTS
if len(locales) < 2:
locales = [
'en-US',
'es-ES',
'es-MX',
'fr-FR',
'hi-IN',
'it-IT',
'de-DE',
'en-GB',
'en-IN',
'ja-JP',
'ko-KR',
'pt-BR',
'zh-CN',
]
locales = ','.join(locales)
if not api_key or not region:
raise HTTPException(
status_code=400,
detail='Azure API key is required for Azure STT',
)
r = None
try:
# Prepare the request
data = {
'definition': json.dumps(
{
'locales': locales.split(','),
'diarization': {'maxSpeakers': max_speakers, 'enabled': True},
}
if locales
else {}
)
}
url = (
base_url or f'https://{region}.api.cognitive.microsoft.com'
) + '/speechtotext/transcriptions:transcribe?api-version=2024-11-15'
# Use context manager to ensure file is properly closed
with open(file_path, 'rb') as audio_file:
r = requests.post(
url=url,
files={'audio': audio_file},
data=data,
headers={
'Ocp-Apim-Subscription-Key': api_key,
},
timeout=AIOHTTP_CLIENT_TIMEOUT,
)
r.raise_for_status()
response = r.json()
# Extract transcript from response
if not response.get('combinedPhrases'):
raise ValueError('No transcription found in response')
# Get the full transcript from combinedPhrases
transcript = response['combinedPhrases'][0].get('text', '').strip()
if not transcript:
raise ValueError('Empty transcript in response')
data = {'text': transcript}
# Save transcript to json file (consistent with other providers)
transcript_file = os.path.join(file_dir, f'{id}.json')
with open(transcript_file, 'w') as f:
json.dump(data, f)
log.debug(data)
return data
except (KeyError, IndexError, ValueError) as e:
log.exception('Error parsing Azure response')
raise HTTPException(
status_code=500,
detail=f'Failed to parse Azure response: {str(e)}',
)
except requests.exceptions.RequestException as e:
log.exception(e)
detail = None
status_code = getattr(r, 'status_code', 500) if r else 500
try:
if r is not None and r.status_code != 200:
res = r.json()
# Azure returns {"code": "...", "message": "...", "innerError": {...}}
if 'code' in res and 'message' in res:
azure_code = res.get('innerError', {}).get('code', res['code'])
user_facing_codes = {
'EmptyAudioFile',
'AudioLengthLimitExceeded',
'NoLanguageIdentified',
'MultipleLanguagesIdentified',
}
if azure_code in user_facing_codes:
detail = res['message']
else:
log.error(f'Azure STT error [{azure_code}]: {res["message"]}')
detail = 'An error occurred during transcription.'
elif 'error' in res:
detail = f'External: {res["error"].get("message", "")}'
except Exception:
detail = f'External: {e}'
raise HTTPException(
status_code=status_code,
detail=detail if detail else 'Open WebUI: Server Connection Error',
)
return await _transcribe_azure(request, file_path, filename, file_dir, id)
elif request.app.state.config.STT_ENGINE == 'mistral':
# Check file exists
if not os.path.exists(file_path):
raise HTTPException(status_code=400, detail='Audio file not found')
return await _transcribe_mistral(request, file_path, filename, metadata, file_dir, id)
# Check file size
file_size = os.path.getsize(file_path)
if file_size > MAX_FILE_SIZE:
raise HTTPException(
status_code=400,
detail=f'File size exceeds limit of {MAX_FILE_SIZE_MB}MB',
)
api_key = request.app.state.config.AUDIO_STT_MISTRAL_API_KEY
api_base_url = request.app.state.config.AUDIO_STT_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1'
use_chat_completions = request.app.state.config.AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS
async def _transcribe_mistral(request, file_path, filename, metadata, file_dir, id):
if not os.path.exists(file_path):
raise HTTPException(status_code=400, detail='Audio file not found')
if not api_key:
raise HTTPException(
status_code=400,
detail='Mistral API key is required for Mistral STT',
)
file_size = os.path.getsize(file_path)
if file_size > MAX_FILE_SIZE:
raise HTTPException(status_code=400, detail=f'File size exceeds limit of {MAX_FILE_SIZE_MB}MB')
r = None
try:
# Use voxtral-mini-latest as the default model for transcription
model = request.app.state.config.STT_MODEL or 'voxtral-mini-latest'
api_key = request.app.state.config.AUDIO_STT_MISTRAL_API_KEY
api_base_url = request.app.state.config.AUDIO_STT_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1'
use_chat_completions = request.app.state.config.AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS
log.info(
f'Mistral STT - model: {model}, '
f'method: {"chat_completions" if use_chat_completions else "transcriptions"}'
)
if not api_key:
raise HTTPException(status_code=400, detail='Mistral API key is required for Mistral STT')
r = None
try:
model = request.app.state.config.STT_MODEL or 'voxtral-mini-latest'
log.info(
f'Mistral STT - model: {model}, '
f'method: {"chat_completions" if use_chat_completions else "transcriptions"}'
)
timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT)
async with aiohttp.ClientSession(timeout=timeout, trust_env=True) as session:
if use_chat_completions:
# Use chat completions API with audio input
# This method requires mp3 or wav format
audio_file_to_use = file_path
if is_audio_conversion_required(file_path):
log.debug('Converting audio to mp3 for chat completions API')
converted_path = convert_audio_to_mp3(file_path)
converted_path = await asyncio.to_thread(convert_audio_to_mp3, file_path)
if converted_path:
audio_file_to_use = converted_path
else:
@@ -1026,133 +972,99 @@ def transcription_handler(request, file_path, metadata, user=None):
detail='Audio conversion failed. Chat completions API requires mp3 or wav format.',
)
# Read and encode audio file as base64
with open(audio_file_to_use, 'rb') as audio_file:
async with aiofiles.open(audio_file_to_use, 'rb') as audio_file:
raw = await audio_file.read()
audio_base64 = {
'data': base64.b64encode(audio_file.read()).decode('utf-8'),
'data': base64.b64encode(raw).decode('utf-8'),
'format': mimetypes.guess_extension(mimetypes.guess_type(audio_file_to_use)[0]).lstrip('.'),
}
# Prepare chat completions request
url = f'{api_base_url}/chat/completions'
# Add language instruction if specified
language = metadata.get('language', None) if metadata else None
if language:
text_instruction = f'Transcribe this audio exactly as spoken in {language}. Do not translate it.'
else:
text_instruction = 'Transcribe this audio exactly as spoken in its original language. Do not translate it to another language.'
text_instruction = (
f'Transcribe this audio exactly as spoken in {language}. Do not translate it.'
if language
else 'Transcribe this audio exactly as spoken in its original language. Do not translate it to another language.'
)
payload = {
'model': model,
'messages': [
{
'role': 'user',
'content': [
{
'type': 'input_audio',
'input_audio': audio_base64,
},
{'type': 'text', 'text': text_instruction},
],
}
],
'messages': [{
'role': 'user',
'content': [
{'type': 'input_audio', 'input_audio': audio_base64},
{'type': 'text', 'text': text_instruction},
],
}],
}
r = requests.post(
url=url,
json=payload,
headers={
'Authorization': f'Bearer {api_key}',
'Content-Type': 'application/json',
},
timeout=AIOHTTP_CLIENT_TIMEOUT,
r = await session.post(
url=f'{api_base_url}/chat/completions', json=payload,
headers={'Authorization': f'Bearer {api_key}', 'Content-Type': 'application/json'},
ssl=AIOHTTP_CLIENT_SESSION_SSL,
)
r.raise_for_status()
response = r.json()
response = await r.json()
# Extract transcript from chat completion response
transcript = response.get('choices', [{}])[0].get('message', {}).get('content', '').strip()
if not transcript:
raise ValueError('Empty transcript in response')
data = {'text': transcript}
else:
# Use dedicated transcriptions API
url = f'{api_base_url}/audio/transcriptions'
# Determine the MIME type
mime_type, _ = mimetypes.guess_type(file_path)
if not mime_type:
mime_type = 'audio/webm'
# Use context manager to ensure file is properly closed
with open(file_path, 'rb') as audio_file:
files = {'file': (filename, audio_file, mime_type)}
data_form = {'model': model}
form_data = aiohttp.FormData()
form_data.add_field('model', model)
# Add language if specified in metadata
language = metadata.get('language', None) if metadata else None
if language:
data_form['language'] = language
language = metadata.get('language', None) if metadata else None
if language:
form_data.add_field('language', language)
r = requests.post(
url=url,
files=files,
data=data_form,
headers={
'Authorization': f'Bearer {api_key}',
},
timeout=AIOHTTP_CLIENT_TIMEOUT,
)
form_data.add_field('file', open(file_path, 'rb'), filename=filename, content_type=mime_type)
r = await session.post(
url=f'{api_base_url}/audio/transcriptions', data=form_data,
headers={'Authorization': f'Bearer {api_key}'},
ssl=AIOHTTP_CLIENT_SESSION_SSL,
)
r.raise_for_status()
response = r.json()
response = await r.json()
# Extract transcript from response
transcript = response.get('text', '').strip()
if not transcript:
raise ValueError('Empty transcript in response')
data = {'text': transcript}
# Save transcript to json file (consistent with other providers)
transcript_file = os.path.join(file_dir, f'{id}.json')
with open(transcript_file, 'w') as f:
json.dump(data, f)
async with aiofiles.open(os.path.join(file_dir, f'{id}.json'), 'w') as f:
await f.write(json.dumps(data))
log.debug(data)
return data
log.debug(data)
return data
except ValueError as e:
log.exception('Error parsing Mistral response')
raise HTTPException(
status_code=500,
detail=f'Failed to parse Mistral response: {str(e)}',
)
except requests.exceptions.RequestException as e:
log.exception(e)
detail = None
try:
if r is not None and r.status_code != 200:
res = r.json()
if 'error' in res:
detail = f'External: {res["error"].get("message", "")}'
else:
detail = f'External: {r.text}'
except Exception:
detail = f'External: {e}'
raise HTTPException(
status_code=getattr(r, 'status_code', 500) if r else 500,
detail=detail if detail else 'Open WebUI: Server Connection Error',
)
except ValueError as e:
log.exception('Error parsing Mistral response')
raise HTTPException(status_code=500, detail=f'Failed to parse Mistral response: {str(e)}')
except aiohttp.ClientResponseError as e:
log.exception(e)
detail = None
try:
if r is not None and r.status != 200:
res = await r.json()
if 'error' in res:
detail = f'External: {res["error"].get("message", "")}'
else:
detail = f'External: {await r.text()}'
except Exception:
detail = f'External: {e}'
raise HTTPException(
status_code=e.status if e.status else 500,
detail=detail if detail else 'Open WebUI: Server Connection Error',
)
def transcribe(request: Request, file_path: str, metadata: Optional[dict] = None, user=None):
async def transcribe(request: Request, file_path: str, metadata: Optional[dict] = None, user=None):
log.info(f'transcribe: {file_path} {metadata}')
if BYPASS_PYDUB_PREPROCESSING:
@@ -1160,7 +1072,7 @@ def transcribe(request: Request, file_path: str, metadata: Optional[dict] = None
chunk_paths = [file_path]
else:
if is_audio_conversion_required(file_path):
file_path = convert_audio_to_mp3(file_path)
file_path = await asyncio.to_thread(convert_audio_to_mp3, file_path)
if not file_path:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@@ -1168,13 +1080,13 @@ def transcribe(request: Request, file_path: str, metadata: Optional[dict] = None
)
try:
file_path = compress_audio(file_path)
file_path = await asyncio.to_thread(compress_audio, file_path)
except Exception as e:
log.exception(e)
# Always produce a list of chunk paths (could be one entry if small)
try:
chunk_paths = split_audio(file_path, MAX_FILE_SIZE)
chunk_paths = await asyncio.to_thread(split_audio, file_path, MAX_FILE_SIZE)
print(f'Chunk paths: {chunk_paths}')
except Exception as e:
log.exception(e)
@@ -1185,28 +1097,20 @@ def transcribe(request: Request, file_path: str, metadata: Optional[dict] = None
results = []
try:
if getattr(request.app.state.config, 'STT_ENGINE', '') == '':
max_workers = 1
else:
max_workers = None
with ThreadPoolExecutor(max_workers=max_workers) as executor:
# Submit tasks for each chunk_path
futures = [
executor.submit(transcription_handler, request, chunk_path, metadata, user)
for chunk_path in chunk_paths
]
# Gather results as they complete
for future in futures:
try:
results.append(future.result())
except HTTPException:
raise
except Exception as transcribe_exc:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f'Error transcribing chunk: {transcribe_exc}',
)
tasks = [
transcription_handler(request, chunk_path, metadata, user)
for chunk_path in chunk_paths
]
for coro in asyncio.as_completed(tasks):
try:
results.append(await coro)
except HTTPException:
raise
except Exception as transcribe_exc:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f'Error transcribing chunk: {transcribe_exc}',
)
finally:
# Clean up only the temporary chunks, never the original file
for chunk_path in chunk_paths:
@@ -1337,7 +1241,7 @@ async def transcription(
if language:
metadata = {'language': language}
result = await asyncio.to_thread(transcribe, request, file_path, metadata, user)
result = await transcribe(request, file_path, metadata, user)
return {
**result,