Files
llm/scripts/rag_embedding_providers.py
T
2026-08-14 09:40:51 +03:00

132 lines
5.2 KiB
Python

from __future__ import annotations
import json
import math
import os
import urllib.request
from typing import Any
from common import hashing_embedding
LOCAL_HASHING_PROVIDER = "local-hashing"
OPENAI_COMPATIBLE_PROVIDER = "openai-compatible"
LOCAL_HASHING_MODEL = "local-hashing-v1"
def l2_normalize(vector: list[float]) -> list[float]:
norm = math.sqrt(sum(value * value for value in vector))
if norm <= 0:
return vector
return [value / norm for value in vector]
def openai_embeddings_url(base_url: str) -> str:
clean = str(base_url or "").strip().rstrip("/")
if not clean:
raise ValueError("embedding_base_url is required for openai-compatible embeddings")
if clean.endswith("/v1/embeddings"):
return clean
if clean.endswith("/embeddings"):
return clean
return f"{clean}/v1/embeddings"
def embed_texts_local_hashing(texts: list[str], *, dimensions: int) -> list[list[float]]:
return [hashing_embedding(text, dimensions=dimensions) for text in texts]
def embed_texts_openai_compatible(
texts: list[str],
*,
model: str,
base_url: str,
dimensions: int = 0,
api_key: str = "",
timeout_seconds: int = 120,
) -> list[list[float]]:
if not model:
raise ValueError("embedding_model is required")
request_payload: dict[str, Any] = {"model": model, "input": texts}
if int(dimensions or 0) > 0:
request_payload["dimensions"] = int(dimensions)
body = json.dumps(request_payload, ensure_ascii=False).encode("utf-8")
headers = {"Content-Type": "application/json"}
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
request = urllib.request.Request(openai_embeddings_url(base_url), data=body, headers=headers, method="POST")
with urllib.request.urlopen(request, timeout=timeout_seconds) as response:
payload = json.loads(response.read().decode("utf-8"))
rows = payload.get("data")
if not isinstance(rows, list):
raise ValueError("Embedding response must contain data[]")
by_index: dict[int, list[float]] = {}
for row in rows:
if not isinstance(row, dict):
continue
embedding = row.get("embedding")
index = row.get("index")
if isinstance(index, bool) or not isinstance(index, int):
index = len(by_index)
if not isinstance(embedding, list) or not embedding:
raise ValueError("Embedding response item has no embedding[]")
vector = [float(value) for value in embedding]
if int(dimensions or 0) > 0:
if len(vector) < int(dimensions):
raise ValueError(
f"Embedding response returned {len(vector)} dimensions, "
f"fewer than requested {int(dimensions)}"
)
vector = vector[: int(dimensions)]
by_index[int(index)] = l2_normalize(vector)
vectors = [by_index[index] for index in range(len(texts)) if index in by_index]
if len(vectors) != len(texts):
raise ValueError(f"Embedding response returned {len(vectors)} vector(s), expected {len(texts)}")
dimensions = len(vectors[0]) if vectors else 0
if any(len(vector) != dimensions for vector in vectors):
raise ValueError("Embedding response returned vectors with inconsistent dimensions")
return vectors
def embed_texts(
texts: list[str],
*,
provider: str,
model: str,
dimensions: int,
base_url: str = "",
api_key_env: str = "OPENAI_API_KEY",
timeout_seconds: int = 120,
) -> list[list[float]]:
normalized_provider = str(provider or LOCAL_HASHING_PROVIDER).strip().lower()
if normalized_provider in {LOCAL_HASHING_PROVIDER, "local_hashing"}:
return embed_texts_local_hashing(texts, dimensions=dimensions)
if normalized_provider in {OPENAI_COMPATIBLE_PROVIDER, "openai_compatible", "openai"}:
api_key = os.environ.get(api_key_env or "OPENAI_API_KEY", "")
return embed_texts_openai_compatible(
texts,
model=model,
base_url=base_url,
dimensions=dimensions,
api_key=api_key,
timeout_seconds=timeout_seconds,
)
raise ValueError(f"Unsupported embedding provider: {provider}")
def provider_metadata(*, provider: str, model: str, dimensions: int, base_url: str = "") -> dict[str, Any]:
normalized_provider = str(provider or LOCAL_HASHING_PROVIDER).strip().lower()
metadata = {
"embedding_provider": normalized_provider,
"embedding_model": model,
"embedding_dimensions": dimensions,
}
if normalized_provider in {OPENAI_COMPATIBLE_PROVIDER, "openai_compatible", "openai"}:
metadata["embedding_provider"] = OPENAI_COMPATIBLE_PROVIDER
metadata["embedding_base_url"] = base_url.rstrip("/") if base_url else ""
metadata["semantic_note"] = "openai-compatible embeddings are neural if the configured endpoint serves an embedding model; API keys are read from environment and not stored in the index."
else:
metadata["embedding_provider"] = "local_hashing"
metadata["semantic_note"] = "local-hashing-v1 is a deterministic lexical vector baseline, not a neural semantic embedding model."
return metadata