132 lines
5.2 KiB
Python
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
|