This commit is contained in:
Timothy Jaeryang Baek
2026-03-24 20:14:28 -05:00
parent 58dcc1b33f
commit 8b6fa1f4ab
+164 -185
View File
@@ -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'