refac
This commit is contained in:
@@ -10,12 +10,11 @@ from open_webui.env import GLOBAL_LOG_LEVEL, BYPASS_MODEL_ACCESS_CONTROL
|
||||
|
||||
from open_webui.routers.openai import embeddings as openai_embeddings
|
||||
from open_webui.routers.ollama import (
|
||||
embeddings as ollama_embeddings,
|
||||
GenerateEmbeddingsForm,
|
||||
embed as ollama_embed,
|
||||
GenerateEmbedForm,
|
||||
)
|
||||
|
||||
|
||||
from open_webui.utils.payload import convert_embedding_payload_openai_to_ollama
|
||||
from open_webui.utils.payload import convert_embed_payload_openai_to_ollama
|
||||
from open_webui.utils.response import convert_embedding_response_ollama_to_openai
|
||||
|
||||
logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
|
||||
@@ -71,12 +70,12 @@ async def generate_embeddings(
|
||||
if not bypass_filter and user.role == "user":
|
||||
check_model_access(user, model)
|
||||
|
||||
# Ollama backend
|
||||
# Ollama backend — use /api/embed which supports batch input natively
|
||||
if model.get("owned_by") == "ollama":
|
||||
ollama_payload = convert_embedding_payload_openai_to_ollama(form_data)
|
||||
response = await ollama_embeddings(
|
||||
ollama_payload = convert_embed_payload_openai_to_ollama(form_data)
|
||||
response = await ollama_embed(
|
||||
request=request,
|
||||
form_data=GenerateEmbeddingsForm(**ollama_payload),
|
||||
form_data=GenerateEmbedForm(**ollama_payload),
|
||||
user=user,
|
||||
)
|
||||
return convert_embedding_response_ollama_to_openai(response)
|
||||
@@ -87,3 +86,4 @@ async def generate_embeddings(
|
||||
form_data=form_data,
|
||||
user=user,
|
||||
)
|
||||
|
||||
|
||||
@@ -395,3 +395,30 @@ def convert_embedding_payload_openai_to_ollama(openai_payload: dict) -> dict:
|
||||
ollama_payload[optional_key] = openai_payload[optional_key]
|
||||
|
||||
return ollama_payload
|
||||
|
||||
|
||||
def convert_embed_payload_openai_to_ollama(openai_payload: dict) -> dict:
|
||||
"""
|
||||
Convert an embeddings request payload from OpenAI format to Ollama's
|
||||
/api/embed format, which supports batch input natively.
|
||||
|
||||
Args:
|
||||
openai_payload (dict): The original payload designed for OpenAI API usage.
|
||||
Expected keys: "model", "input" (str or list[str]).
|
||||
|
||||
Returns:
|
||||
dict: A payload compatible with the Ollama /api/embed endpoint.
|
||||
"""
|
||||
ollama_payload = {"model": openai_payload.get("model")}
|
||||
input_value = openai_payload.get("input")
|
||||
|
||||
# /api/embed accepts 'input' as a string or list of strings directly
|
||||
ollama_payload["input"] = input_value
|
||||
|
||||
# Optionally forward other fields if present
|
||||
for optional_key in ("truncate", "options", "keep_alive"):
|
||||
if optional_key in openai_payload:
|
||||
ollama_payload[optional_key] = openai_payload[optional_key]
|
||||
|
||||
return ollama_payload
|
||||
|
||||
|
||||
@@ -192,17 +192,29 @@ def convert_embedding_response_ollama_to_openai(response) -> dict:
|
||||
"model": "...",
|
||||
}
|
||||
"""
|
||||
# Ollama batch-style output
|
||||
# Ollama batch-style output from /api/embed
|
||||
# Response format: {"embeddings": [[0.1, 0.2, ...], [0.3, 0.4, ...]], "model": "..."}
|
||||
if isinstance(response, dict) and "embeddings" in response:
|
||||
openai_data = []
|
||||
for i, emb in enumerate(response["embeddings"]):
|
||||
openai_data.append(
|
||||
{
|
||||
"object": "embedding",
|
||||
"embedding": emb.get("embedding"),
|
||||
"index": emb.get("index", i),
|
||||
}
|
||||
)
|
||||
# /api/embed returns embeddings as plain float lists
|
||||
if isinstance(emb, list):
|
||||
openai_data.append(
|
||||
{
|
||||
"object": "embedding",
|
||||
"embedding": emb,
|
||||
"index": i,
|
||||
}
|
||||
)
|
||||
# Also handle dict format for robustness
|
||||
elif isinstance(emb, dict):
|
||||
openai_data.append(
|
||||
{
|
||||
"object": "embedding",
|
||||
"embedding": emb.get("embedding"),
|
||||
"index": emb.get("index", i),
|
||||
}
|
||||
)
|
||||
return {
|
||||
"object": "list",
|
||||
"data": openai_data,
|
||||
|
||||
Reference in New Issue
Block a user