From 8b6fa1f4ab6099a305de08706621075c205f65c4 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Tue, 24 Mar 2026 20:14:28 -0500 Subject: [PATCH] refac --- backend/open_webui/retrieval/utils.py | 349 ++++++++++++-------------- 1 file changed, 164 insertions(+), 185 deletions(-) diff --git a/backend/open_webui/retrieval/utils.py b/backend/open_webui/retrieval/utils.py index c80e52209..1c974fc40 100644 --- a/backend/open_webui/retrieval/utils.py +++ b/backend/open_webui/retrieval/utils.py @@ -550,34 +550,30 @@ def generate_openai_batch_embeddings( key: str = '', prefix: str = None, user: UserModel = None, -) -> Optional[list[list[float]]]: - try: - log.debug(f'generate_openai_batch_embeddings:model {model} batch size: {len(texts)}') - json_data = {'input': texts, 'model': model} - if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): - json_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix +) -> list[list[float]]: + log.debug(f'generate_openai_batch_embeddings:model {model} batch size: {len(texts)}') + json_data = {'input': texts, 'model': model} + if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): + json_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix - headers = { - 'Content-Type': 'application/json', - 'Authorization': f'Bearer {key}', - } - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) + headers = { + 'Content-Type': 'application/json', + 'Authorization': f'Bearer {key}', + } + if ENABLE_FORWARD_USER_INFO_HEADERS and user: + headers = include_user_info_headers(headers, user) - r = requests.post( - f'{url}/embeddings', - headers=headers, - json=json_data, - ) - r.raise_for_status() - data = r.json() - if 'data' in data: - return [elem['embedding'] for elem in data['data']] - else: - raise ValueError("Unexpected OpenAI embeddings response: missing 'data' key") - except Exception as e: - log.exception(f'Error generating openai batch embeddings: {e}') - return None + r = requests.post( + f'{url}/embeddings', + headers=headers, + json=json_data, + ) + r.raise_for_status() + data = r.json() + if 'data' in data: + return [elem['embedding'] for elem in data['data']] + else: + raise ValueError("Unexpected OpenAI embeddings response: missing 'data' key") async def agenerate_openai_batch_embeddings( @@ -587,38 +583,34 @@ async def agenerate_openai_batch_embeddings( key: str = '', prefix: str = None, user: UserModel = None, -) -> Optional[list[list[float]]]: - try: - log.debug(f'agenerate_openai_batch_embeddings:model {model} batch size: {len(texts)}') - form_data = {'input': texts, 'model': model} - if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): - form_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix +) -> list[list[float]]: + log.debug(f'agenerate_openai_batch_embeddings:model {model} batch size: {len(texts)}') + form_data = {'input': texts, 'model': model} + if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): + form_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix - headers = { - 'Content-Type': 'application/json', - 'Authorization': f'Bearer {key}', - } - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) + headers = { + 'Content-Type': 'application/json', + 'Authorization': f'Bearer {key}', + } + if ENABLE_FORWARD_USER_INFO_HEADERS and user: + headers = include_user_info_headers(headers, user) - async with aiohttp.ClientSession( - trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) - ) as session: - async with session.post( - f'{url}/embeddings', - headers=headers, - json=form_data, - ssl=AIOHTTP_CLIENT_SESSION_SSL, - ) as r: - r.raise_for_status() - data = await r.json() - if 'data' in data: - return [item['embedding'] for item in data['data']] - else: - raise Exception('Something went wrong :/') - except Exception as e: - log.exception(f'Error generating openai batch embeddings: {e}') - return None + async with aiohttp.ClientSession( + trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) + ) as session: + async with session.post( + f'{url}/embeddings', + headers=headers, + json=form_data, + ssl=AIOHTTP_CLIENT_SESSION_SSL, + ) as r: + r.raise_for_status() + data = await r.json() + if 'data' in data: + return [item['embedding'] for item in data['data']] + else: + raise ValueError("Unexpected OpenAI embeddings response: missing 'data' key") def generate_azure_openai_batch_embeddings( @@ -629,42 +621,38 @@ def generate_azure_openai_batch_embeddings( version: str = '', prefix: str = None, user: UserModel = None, -) -> Optional[list[list[float]]]: - try: - log.debug(f'generate_azure_openai_batch_embeddings:deployment {model} batch size: {len(texts)}') - json_data = {'input': texts} - if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): - json_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix +) -> list[list[float]]: + log.debug(f'generate_azure_openai_batch_embeddings:deployment {model} batch size: {len(texts)}') + json_data = {'input': texts} + if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): + json_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix - url = f'{url}/openai/deployments/{model}/embeddings?api-version={version}' + url = f'{url}/openai/deployments/{model}/embeddings?api-version={version}' - for _ in range(5): - headers = { - 'Content-Type': 'application/json', - 'api-key': key, - } - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) + for _ in range(5): + headers = { + 'Content-Type': 'application/json', + 'api-key': key, + } + if ENABLE_FORWARD_USER_INFO_HEADERS and user: + headers = include_user_info_headers(headers, user) - r = requests.post( - url, - headers=headers, - json=json_data, - ) - if r.status_code == 429: - retry = float(r.headers.get('Retry-After', '1')) - time.sleep(retry) - continue - r.raise_for_status() - data = r.json() - if 'data' in data: - return [elem['embedding'] for elem in data['data']] - else: - raise Exception('Something went wrong :/') - return None - except Exception as e: - log.exception(f'Error generating azure openai batch embeddings: {e}') - return None + r = requests.post( + url, + headers=headers, + json=json_data, + ) + if r.status_code == 429: + retry = float(r.headers.get('Retry-After', '1')) + time.sleep(retry) + continue + r.raise_for_status() + data = r.json() + if 'data' in data: + return [elem['embedding'] for elem in data['data']] + else: + raise ValueError("Unexpected Azure OpenAI embeddings response: missing 'data' key") + raise Exception('Azure OpenAI embedding request failed: max retries (429) exceeded') async def agenerate_azure_openai_batch_embeddings( @@ -675,40 +663,36 @@ async def agenerate_azure_openai_batch_embeddings( version: str = '', prefix: str = None, user: UserModel = None, -) -> Optional[list[list[float]]]: - try: - log.debug(f'agenerate_azure_openai_batch_embeddings:deployment {model} batch size: {len(texts)}') - form_data = {'input': texts} - if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): - form_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix +) -> list[list[float]]: + log.debug(f'agenerate_azure_openai_batch_embeddings:deployment {model} batch size: {len(texts)}') + form_data = {'input': texts} + if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): + form_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix - full_url = f'{url}/openai/deployments/{model}/embeddings?api-version={version}' + full_url = f'{url}/openai/deployments/{model}/embeddings?api-version={version}' - headers = { - 'Content-Type': 'application/json', - 'api-key': key, - } - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) + headers = { + 'Content-Type': 'application/json', + 'api-key': key, + } + if ENABLE_FORWARD_USER_INFO_HEADERS and user: + headers = include_user_info_headers(headers, user) - async with aiohttp.ClientSession( - trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) - ) as session: - async with session.post( - full_url, - headers=headers, - json=form_data, - ssl=AIOHTTP_CLIENT_SESSION_SSL, - ) as r: - r.raise_for_status() - data = await r.json() - if 'data' in data: - return [item['embedding'] for item in data['data']] - else: - raise Exception('Something went wrong :/') - except Exception as e: - log.exception(f'Error generating azure openai batch embeddings: {e}') - return None + async with aiohttp.ClientSession( + trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) + ) as session: + async with session.post( + full_url, + headers=headers, + json=form_data, + ssl=AIOHTTP_CLIENT_SESSION_SSL, + ) as r: + r.raise_for_status() + data = await r.json() + if 'data' in data: + return [item['embedding'] for item in data['data']] + else: + raise ValueError("Unexpected Azure OpenAI embeddings response: missing 'data' key") def generate_ollama_batch_embeddings( @@ -718,37 +702,33 @@ def generate_ollama_batch_embeddings( key: str = '', prefix: str = None, user: UserModel = None, -) -> Optional[list[list[float]]]: - try: - log.debug(f'generate_ollama_batch_embeddings:model {model} batch size: {len(texts)}') - json_data = {'input': texts, 'model': model, 'truncate': True} - if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): - json_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix +) -> list[list[float]]: + log.debug(f'generate_ollama_batch_embeddings:model {model} batch size: {len(texts)}') + json_data = {'input': texts, 'model': model, 'truncate': True} + if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): + json_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix - headers = { - 'Content-Type': 'application/json', - 'Authorization': f'Bearer {key}', - } - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) + headers = { + 'Content-Type': 'application/json', + 'Authorization': f'Bearer {key}', + } + if ENABLE_FORWARD_USER_INFO_HEADERS and user: + headers = include_user_info_headers(headers, user) - r = requests.post( - f'{url}/api/embed', - headers=headers, - json=json_data, - ) - if r.status_code != 200: - error_detail = r.json().get('error', r.text) - raise Exception(f'Ollama embed error ({r.status_code}): {error_detail}') - data = r.json() + r = requests.post( + f'{url}/api/embed', + headers=headers, + json=json_data, + ) + if r.status_code != 200: + error_detail = r.json().get('error', r.text) + raise Exception(f'Ollama embed error ({r.status_code}): {error_detail}') + data = r.json() - if 'embeddings' in data: - return data['embeddings'] - else: - raise ValueError("Unexpected Ollama embeddings response: missing 'embeddings' key") - except Exception as e: - log.exception(f'Error generating ollama batch embeddings: {e}') - return None + if 'embeddings' in data: + return data['embeddings'] + else: + raise ValueError("Unexpected Ollama embeddings response: missing 'embeddings' key") async def agenerate_ollama_batch_embeddings( @@ -758,41 +738,37 @@ async def agenerate_ollama_batch_embeddings( key: str = '', prefix: str = None, user: UserModel = None, -) -> Optional[list[list[float]]]: - try: - log.debug(f'agenerate_ollama_batch_embeddings:model {model} batch size: {len(texts)}') - form_data = {'input': texts, 'model': model, 'truncate': True} - if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): - form_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix +) -> list[list[float]]: + log.debug(f'agenerate_ollama_batch_embeddings:model {model} batch size: {len(texts)}') + form_data = {'input': texts, 'model': model, 'truncate': True} + if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): + form_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix - headers = { - 'Content-Type': 'application/json', - 'Authorization': f'Bearer {key}', - } - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) + headers = { + 'Content-Type': 'application/json', + 'Authorization': f'Bearer {key}', + } + if ENABLE_FORWARD_USER_INFO_HEADERS and user: + headers = include_user_info_headers(headers, user) - async with aiohttp.ClientSession( - trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) - ) as session: - async with session.post( - f'{url}/api/embed', - headers=headers, - json=form_data, - ssl=AIOHTTP_CLIENT_SESSION_SSL, - ) as r: - if r.status != 200: - error_data = await r.json() - error_detail = error_data.get('error', str(error_data)) - raise Exception(f'Ollama embed error ({r.status}): {error_detail}') - data = await r.json() - if 'embeddings' in data: - return data['embeddings'] - else: - raise Exception('Something went wrong :/') - except Exception as e: - log.exception(f'Error generating ollama batch embeddings: {e}') - return None + async with aiohttp.ClientSession( + trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) + ) as session: + async with session.post( + f'{url}/api/embed', + headers=headers, + json=form_data, + ssl=AIOHTTP_CLIENT_SESSION_SSL, + ) as r: + if r.status != 200: + error_data = await r.json() + error_detail = error_data.get('error', str(error_data)) + raise Exception(f'Ollama embed error ({r.status}): {error_detail}') + data = await r.json() + if 'embeddings' in data: + return data['embeddings'] + else: + raise ValueError("Unexpected Ollama embeddings response: missing 'embeddings' key") def get_embedding_function( @@ -860,11 +836,14 @@ def get_embedding_function( for batch in batches: batch_results.append(await embedding_function(batch, prefix=prefix, user=user)) - # Flatten results + # Flatten results — raise if any batch failed embeddings = [] - for batch_embeddings in batch_results: - if isinstance(batch_embeddings, list): - embeddings.extend(batch_embeddings) + for i, batch_embeddings in enumerate(batch_results): + if batch_embeddings is None: + raise Exception( + f'Embedding generation failed for batch {i + 1}/{len(batches)}' + ) + embeddings.extend(batch_embeddings) log.debug( f'generate_multiple_async: Generated {len(embeddings)} embeddings from {len(batches)} parallel batches'