diff --git a/backend/open_webui/routers/audio.py b/backend/open_webui/routers/audio.py index bba91c2d5a..51373caf3b 100644 --- a/backend/open_webui/routers/audio.py +++ b/backend/open_webui/routers/audio.py @@ -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,