refac
This commit is contained in:
@@ -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'
|
||||
|
||||
Reference in New Issue
Block a user