memory-os/docker/worker/services/embedding.py

84 lines
2.9 KiB
Python

"""
Embedding client. Provider-agnostic — defaults to OpenRouter, configurable
via EMBEDDING_API_BASE and EMBEDDING_MODEL for local Ollama/vLLM/llama.cpp.
Mandatory dimension validation with in-memory LRU cache.
"""
import os
import logging
from collections import OrderedDict
import httpx
logger = logging.getLogger("cognitive-worker.embedding")
OPENROUTER_API_KEY = os.environ.get("OPENROUTER_API_KEY", "")
EMBEDDING_API_KEY = os.environ.get("EMBEDDING_API_KEY") # no fallback — explicit per-branch
EMBEDDING_DIMS = int(os.environ.get("EMBEDDING_DIMS", "4096"))
EMBEDDING_API_BASE = os.environ.get(
"EMBEDDING_API_BASE", "https://openrouter.ai/api/v1"
)
EMBEDDING_MODEL = os.environ.get(
"EMBEDDING_MODEL", "qwen/qwen3-embedding-8b"
)
# In-memory LRU cache — avoids recomputing embeddings for repeated text
_EMBEDDING_CACHE_MAXSIZE = 256
_embedding_cache: OrderedDict = OrderedDict()
async def get_embedding(text: str) -> list[float]:
"""
Generates embedding via the configured backend.
Validates that the returned dimensions match EMBEDDING_DIMS.
Results are cached in-memory (LRU, max 256 entries).
"""
# Check cache
if text in _embedding_cache:
_embedding_cache.move_to_end(text)
return _embedding_cache[text]
headers = {"Content-Type": "application/json"}
if "openrouter" in EMBEDDING_API_BASE.lower():
if not OPENROUTER_API_KEY:
raise RuntimeError("OPENROUTER_API_KEY is required for OpenRouter")
headers["Authorization"] = f"Bearer {OPENROUTER_API_KEY}"
headers["HTTP-Referer"] = "https://localhost"
headers["X-Title"] = "Cognitive-Agent-MaaS"
elif EMBEDDING_API_KEY:
headers["Authorization"] = f"Bearer {EMBEDDING_API_KEY}"
# else: no auth header — localhost or unauthenticated endpoint
payload = {
"model": EMBEDDING_MODEL,
"input": text,
"dimensions": EMBEDDING_DIMS, # OpenAI/OpenRouter-specific; ignored by Ollama/vLLM
}
async with httpx.AsyncClient(timeout=60) as client:
resp = await client.post(
f"{EMBEDDING_API_BASE}/embeddings",
headers=headers,
json=payload,
)
resp.raise_for_status()
data = resp.json()
vec = data["data"][0]["embedding"]
# ─── Critical dimension validation ──────────────────────────────────────
if len(vec) != EMBEDDING_DIMS:
raise ValueError(
f"Embedding dimension mismatch: "
f"expected {EMBEDDING_DIMS}, got {len(vec)}. "
f"Check EMBEDDING_DIMS in .env and the Qdrant collection."
)
logger.debug(f"Embedding generated: {len(vec)} dims")
# Store in LRU cache
_embedding_cache[text] = vec
if len(_embedding_cache) > _EMBEDDING_CACHE_MAXSIZE:
_embedding_cache.popitem(last=False)
return vec