initial
This commit is contained in:
463
backend/app/core/rag.py
Normal file
463
backend/app/core/rag.py
Normal file
@@ -0,0 +1,463 @@
|
||||
"""Qdrant RAG client: glossary / facts / history indexing and retrieval.
|
||||
|
||||
Embeddings are configurable via admin settings (see `embedding.*` keys):
|
||||
|
||||
* `embedding.provider = "hash"` — deterministic offline fallback (no semantic quality).
|
||||
* `embedding.provider = "openai"` — calls the OpenAI-compatible `/embeddings`
|
||||
endpoint of `embedding.base_url` (falls back to `llm.base_url` if empty).
|
||||
|
||||
Vector dimension (`embedding.dim`) is normally auto-probed from the endpoint on
|
||||
first use (set it to 0). When the configured dimension changes, the Qdrant
|
||||
collections are dropped and recreated — already-indexed points are lost, but
|
||||
they will be repopulated on the next RAG upsert from the engine.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import httpx
|
||||
from qdrant_client import AsyncQdrantClient
|
||||
from qdrant_client.http import models as qm
|
||||
|
||||
from app.config import settings
|
||||
from app.core.settings_service import cast_setting
|
||||
from app.logging_setup import get_logger
|
||||
|
||||
log = get_logger("rag")
|
||||
|
||||
|
||||
COLLECTION_GLOSSARY = "glossary"
|
||||
COLLECTION_HISTORY = "history"
|
||||
ALL_COLLECTIONS = (COLLECTION_GLOSSARY, COLLECTION_HISTORY)
|
||||
|
||||
# Fallback dimension for the hash embedder (kept stable across restarts).
|
||||
HASH_EMBED_DIM = 384
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Embedders
|
||||
# ---------------------------------------------------------------------------
|
||||
class _HashEmbedder:
|
||||
"""Deterministic lightweight embedder used as an offline fallback.
|
||||
|
||||
Not semantically rich, but provides stable vectors for retrieval by keyword
|
||||
overlap (bag-of-tokens hashed into a fixed-dim vector, L2-normalized).
|
||||
"""
|
||||
|
||||
def __init__(self, dim: int = HASH_EMBED_DIM):
|
||||
self.dim = dim
|
||||
|
||||
async def embed(self, text: str) -> List[float]:
|
||||
vec = [0.0] * self.dim
|
||||
tokens = [t for t in text.lower().split() if t]
|
||||
if not tokens:
|
||||
return vec
|
||||
for tok in tokens:
|
||||
h = abs(hash(tok)) % self.dim
|
||||
vec[h] += 1.0
|
||||
h2 = abs(hash(tok + "_b")) % self.dim
|
||||
vec[h2] += 0.5
|
||||
norm = sum(v * v for v in vec) ** 0.5
|
||||
if norm > 0:
|
||||
vec = [v / norm for v in vec]
|
||||
return vec
|
||||
|
||||
async def probe_dim(self) -> int:
|
||||
return self.dim
|
||||
|
||||
|
||||
class OpenAIEmbedder:
|
||||
"""Real embeddings via OpenAI-compatible `/embeddings` endpoint.
|
||||
|
||||
Falls back to `_HashEmbedder` per-call if the endpoint is unreachable or
|
||||
returns an error — so RAG keeps working even if the embeddings server is
|
||||
temporarily down.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
model: str,
|
||||
timeout: int = 60,
|
||||
fallback_dim: int = HASH_EMBED_DIM,
|
||||
):
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.model = model or "text-embedding-3-small"
|
||||
self.timeout = timeout
|
||||
self._fallback = _HashEmbedder(fallback_dim)
|
||||
|
||||
def _headers(self) -> Dict[str, str]:
|
||||
h = {"Content-Type": "application/json"}
|
||||
if self.api_key and self.api_key != "dummy":
|
||||
h["Authorization"] = f"Bearer {self.api_key}"
|
||||
return h
|
||||
|
||||
async def _raw_embed(self, text: str) -> Optional[List[float]]:
|
||||
url = f"{self.base_url}/embeddings"
|
||||
payload = {"model": self.model, "input": text}
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
resp = await client.post(url, json=payload, headers=self._headers())
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
arr = (data.get("data") or [{}])[0].get("embedding") or []
|
||||
if not arr:
|
||||
return None
|
||||
return [float(x) for x in arr]
|
||||
except Exception as e:
|
||||
log.warning("openai_embed_failed", model=self.model, error=f"{type(e).__name__}: {e}")
|
||||
return None
|
||||
|
||||
async def embed(self, text: str) -> List[float]:
|
||||
vec = await self._raw_embed(text)
|
||||
if vec:
|
||||
return vec
|
||||
# Network/endpoint failure — degrade gracefully to hash fallback
|
||||
return await self._fallback.embed(text)
|
||||
|
||||
async def probe_dim(self) -> int:
|
||||
"""Probe the endpoint with a short text and return the vector dimension.
|
||||
|
||||
Returns HASH_EMBED_DIM if the endpoint is unreachable so the system
|
||||
keeps working (with degraded retrieval quality).
|
||||
"""
|
||||
vec = await self._raw_embed("dimension probe")
|
||||
if vec:
|
||||
return len(vec)
|
||||
log.warning("embed_probe_failed_using_hash_dim", dim=HASH_EMBED_DIM)
|
||||
return HASH_EMBED_DIM
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# RAG client
|
||||
# ---------------------------------------------------------------------------
|
||||
class RagClient:
|
||||
"""Qdrant-backed RAG client with configurable embeddings."""
|
||||
|
||||
def __init__(self, url: str | None = None):
|
||||
url = url or settings.qdrant_url
|
||||
self.client = AsyncQdrantClient(url=url)
|
||||
# Cache of {collection_name: configured_dim}. Populated by ensure_collections.
|
||||
self._collection_dims: Dict[str, int] = {}
|
||||
# Lazily constructed embedder + its config signature (so we rebuild on settings change).
|
||||
self._embedder: Optional[Any] = None
|
||||
self._embedder_sig: Optional[str] = None
|
||||
self._configured_dim: Optional[int] = None # resolved dim (after probe)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_embedder_config(settings_map: Dict[str, Any]) -> Dict[str, Any]:
|
||||
provider = str(settings_map.get("embedding.provider", "hash")).lower().strip() or "hash"
|
||||
base_url = str(settings_map.get("embedding.base_url", "") or "").strip()
|
||||
if not base_url:
|
||||
base_url = str(settings_map.get("llm.base_url", "") or "").strip()
|
||||
api_key = str(settings_map.get("embedding.api_key", "") or "").strip()
|
||||
if not api_key:
|
||||
api_key = str(settings_map.get("llm.api_key", "") or "").strip()
|
||||
model = str(settings_map.get("embedding.model", "text-embedding-3-small") or "text-embedding-3-small")
|
||||
dim = int(cast_setting("embedding.dim", settings_map.get("embedding.dim", 0)) or 0)
|
||||
timeout = int(cast_setting("embedding.request_timeout", settings_map.get("embedding.request_timeout", 60)) or 60)
|
||||
return {
|
||||
"provider": provider,
|
||||
"base_url": base_url,
|
||||
"api_key": api_key,
|
||||
"model": model,
|
||||
"dim": dim,
|
||||
"timeout": timeout,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _build_embedder(cfg: Dict[str, Any]) -> Any:
|
||||
if cfg["provider"] == "openai" and cfg["base_url"]:
|
||||
return OpenAIEmbedder(
|
||||
base_url=cfg["base_url"],
|
||||
api_key=cfg["api_key"],
|
||||
model=cfg["model"],
|
||||
timeout=cfg["timeout"],
|
||||
fallback_dim=HASH_EMBED_DIM,
|
||||
)
|
||||
return _HashEmbedder(HASH_EMBED_DIM)
|
||||
|
||||
def _embedder_signature(self, cfg: Dict[str, Any]) -> str:
|
||||
# Only fields that affect the produced vector — `dim` is resolved via probe.
|
||||
return f"{cfg['provider']}|{cfg['base_url']}|{cfg['model']}"
|
||||
|
||||
async def get_embedder(self, settings_map: Optional[Dict[str, Any]] = None) -> Any:
|
||||
"""Return the current embedder, rebuilding it if settings changed.
|
||||
|
||||
If `settings_map` is provided and the provider/base_url/model changed,
|
||||
the embedder is rebuilt and Qdrant collections are reconfigured.
|
||||
"""
|
||||
if settings_map is None:
|
||||
# Caller has no DB context — return whatever is cached.
|
||||
if self._embedder is None:
|
||||
self._embedder = _HashEmbedder(HASH_EMBED_DIM)
|
||||
self._embedder_sig = "hash||"
|
||||
return self._embedder
|
||||
|
||||
cfg = RagClient._resolve_embedder_config(settings_map)
|
||||
sig = self._embedder_signature(cfg)
|
||||
if self._embedder is None or sig != self._embedder_sig:
|
||||
self._embedder = RagClient._build_embedder(cfg)
|
||||
self._embedder_sig = sig
|
||||
self._configured_dim = None # force re-probe on next ensure_collections
|
||||
await self.ensure_collections(settings_map)
|
||||
return self._embedder
|
||||
|
||||
async def _resolve_dim(self, embedder: Any, cfg: Dict[str, Any]) -> int:
|
||||
if cfg["dim"] and cfg["dim"] > 0:
|
||||
return cfg["dim"]
|
||||
if self._configured_dim is not None:
|
||||
return self._configured_dim
|
||||
# Auto-probe from the endpoint (or fallback to HASH_EMBED_DIM).
|
||||
dim = await embedder.probe_dim()
|
||||
self._configured_dim = dim
|
||||
log.info("rag_dim_probed", dim=dim, provider=cfg["provider"])
|
||||
return dim
|
||||
|
||||
async def ensure_collections(self, settings_map: Optional[Dict[str, Any]] = None) -> None:
|
||||
"""Create Qdrant collections if missing; recreate if dim changed.
|
||||
|
||||
Recreating drops all points — they will be repopulated by subsequent
|
||||
upserts from the engine (glossary tool, history indexing).
|
||||
"""
|
||||
cfg = RagClient._resolve_embedder_config(settings_map or {})
|
||||
embedder = await self.get_embedder(settings_map)
|
||||
desired_dim = await self._resolve_dim(embedder, cfg)
|
||||
|
||||
for name in ALL_COLLECTIONS:
|
||||
existing_dim = await self._get_collection_dim(name)
|
||||
if existing_dim is None:
|
||||
try:
|
||||
await self.client.create_collection(
|
||||
collection_name=name,
|
||||
vectors_config=qm.VectorParams(size=desired_dim, distance=qm.Distance.COSINE),
|
||||
)
|
||||
self._collection_dims[name] = desired_dim
|
||||
log.info("rag_collection_created", name=name, dim=desired_dim)
|
||||
except Exception as e:
|
||||
log.warning("rag_collection_create_failed", name=name, error=str(e))
|
||||
elif existing_dim != desired_dim:
|
||||
log.warning(
|
||||
"rag_collection_dim_mismatch_recreate",
|
||||
name=name,
|
||||
old=existing_dim,
|
||||
new=desired_dim,
|
||||
)
|
||||
try:
|
||||
await self.client.delete_collection(collection_name=name)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
await self.client.create_collection(
|
||||
collection_name=name,
|
||||
vectors_config=qm.VectorParams(size=desired_dim, distance=qm.Distance.COSINE),
|
||||
)
|
||||
self._collection_dims[name] = desired_dim
|
||||
except Exception as e:
|
||||
log.warning("rag_collection_recreate_failed", name=name, error=str(e))
|
||||
else:
|
||||
self._collection_dims[name] = existing_dim
|
||||
|
||||
async def _get_collection_dim(self, name: str) -> Optional[int]:
|
||||
try:
|
||||
info = await self.client.get_collection(collection_name=name)
|
||||
cfg = info.config.params.vectors
|
||||
# Qdrant returns either a single VectorParams or a NamedVectors dict
|
||||
if isinstance(cfg, qm.VectorParams):
|
||||
return cfg.size
|
||||
# NamedVectors: take first vector config
|
||||
if hasattr(cfg, "size") and isinstance(cfg.size, int):
|
||||
return cfg.size
|
||||
if isinstance(cfg, dict):
|
||||
for v in cfg.values():
|
||||
if hasattr(v, "size") and isinstance(v.size, int):
|
||||
return v.size
|
||||
except Exception:
|
||||
return None
|
||||
return None
|
||||
|
||||
async def embed(self, text: str, settings_map: Optional[Dict[str, Any]] = None) -> List[float]:
|
||||
embedder = await self.get_embedder(settings_map)
|
||||
return await embedder.embed(text)
|
||||
|
||||
async def upsert_glossary(
|
||||
self,
|
||||
world_id: uuid.UUID,
|
||||
entry_id: uuid.UUID,
|
||||
kind: str,
|
||||
name: str,
|
||||
description: str,
|
||||
payload: Dict[str, Any],
|
||||
settings_map: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
text = f"{kind}: {name}. {description}"
|
||||
vector = await self.embed(text, settings_map)
|
||||
await self.client.upsert(
|
||||
collection_name=COLLECTION_GLOSSARY,
|
||||
points=[
|
||||
qm.PointStruct(
|
||||
id=str(entry_id),
|
||||
vector=vector,
|
||||
payload={
|
||||
"world_id": str(world_id),
|
||||
"entry_id": str(entry_id),
|
||||
"kind": kind,
|
||||
"name": name,
|
||||
"description": description,
|
||||
"text": text,
|
||||
**payload,
|
||||
},
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
async def upsert_history(
|
||||
self,
|
||||
session_id: uuid.UUID,
|
||||
message_id: uuid.UUID,
|
||||
seq: int,
|
||||
text: str,
|
||||
kind: str,
|
||||
settings_map: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
vector = await self.embed(text, settings_map)
|
||||
await self.client.upsert(
|
||||
collection_name=COLLECTION_HISTORY,
|
||||
points=[
|
||||
qm.PointStruct(
|
||||
id=str(message_id),
|
||||
vector=vector,
|
||||
payload={
|
||||
"session_id": str(session_id),
|
||||
"message_id": str(message_id),
|
||||
"seq": seq,
|
||||
"kind": kind,
|
||||
"text": text,
|
||||
},
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
async def search_glossary(
|
||||
self,
|
||||
world_id: uuid.UUID,
|
||||
query: str,
|
||||
limit: int = 5,
|
||||
settings_map: Optional[Dict[str, Any]] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
vector = await self.embed(query, settings_map)
|
||||
res = await self.client.search(
|
||||
collection_name=COLLECTION_GLOSSARY,
|
||||
query_vector=vector,
|
||||
query_filter=qm.Filter(
|
||||
must=[qm.FieldCondition(key="world_id", match=qm.MatchValue(value=str(world_id)))]
|
||||
),
|
||||
limit=limit,
|
||||
with_payload=True,
|
||||
)
|
||||
return [r.payload for r in res]
|
||||
except Exception as e:
|
||||
log.warning("rag_search_glossary_failed", error=str(e))
|
||||
return []
|
||||
|
||||
async def search_history(
|
||||
self,
|
||||
session_id: uuid.UUID,
|
||||
query: str,
|
||||
limit: int = 5,
|
||||
settings_map: Optional[Dict[str, Any]] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
vector = await self.embed(query, settings_map)
|
||||
res = await self.client.search(
|
||||
collection_name=COLLECTION_HISTORY,
|
||||
query_vector=vector,
|
||||
query_filter=qm.Filter(
|
||||
must=[qm.FieldCondition(key="session_id", match=qm.MatchValue(value=str(session_id)))]
|
||||
),
|
||||
limit=limit,
|
||||
with_payload=True,
|
||||
)
|
||||
return [r.payload for r in res]
|
||||
except Exception as e:
|
||||
log.warning("rag_search_history_failed", error=str(e))
|
||||
return []
|
||||
|
||||
async def delete_history(self, session_id: uuid.UUID) -> None:
|
||||
try:
|
||||
await self.client.delete(
|
||||
collection_name=COLLECTION_HISTORY,
|
||||
points_selector=qm.FilterSelector(
|
||||
filter=qm.Filter(must=[qm.FieldCondition(key="session_id", match=qm.MatchValue(value=str(session_id)))])
|
||||
),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Singleton + cache invalidation
|
||||
# ---------------------------------------------------------------------------
|
||||
_rag: Optional[RagClient] = None
|
||||
|
||||
|
||||
async def get_rag(settings_map: Optional[Dict[str, Any]] = None) -> RagClient:
|
||||
"""Get the shared RagClient, ensuring collections are configured for the
|
||||
current embedding settings.
|
||||
|
||||
Pass `settings_map` from DB on the first call (or whenever settings may
|
||||
have changed) so the client can rebuild its embedder and reconfigure
|
||||
Qdrant collections if `embedding.provider` / `embedding.base_url` /
|
||||
`embedding.model` / `embedding.dim` changed.
|
||||
"""
|
||||
global _rag
|
||||
if _rag is None:
|
||||
_rag = RagClient()
|
||||
await _rag.ensure_collections(settings_map)
|
||||
elif settings_map is not None:
|
||||
# Re-check embedder signature; ensure_collections runs only if changed.
|
||||
await _rag.get_embedder(settings_map)
|
||||
return _rag
|
||||
|
||||
|
||||
async def reset_rag() -> None:
|
||||
"""Drop the cached RAG client so the next `get_rag()` rebuilds it from
|
||||
current settings. Call this after admin updates embedding.* settings.
|
||||
"""
|
||||
global _rag
|
||||
_rag = None
|
||||
|
||||
|
||||
async def probe_embeddings(settings_map: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Standalone probe used by the admin "Test embeddings" button.
|
||||
|
||||
Returns dict with: ok, provider, base_url, model, dim, sample_norm, error.
|
||||
Does not touch the shared singleton or Qdrant.
|
||||
"""
|
||||
cfg = RagClient._resolve_embedder_config(settings_map)
|
||||
embedder = RagClient._build_embedder(cfg)
|
||||
try:
|
||||
vec = await embedder.embed("RAG embedding probe: a brave adventurer enters a tavern.")
|
||||
if not vec:
|
||||
return {"ok": False, "provider": cfg["provider"], "error": "empty_vector"}
|
||||
norm = sum(v * v for v in vec) ** 0.5
|
||||
return {
|
||||
"ok": True,
|
||||
"provider": cfg["provider"],
|
||||
"base_url": cfg["base_url"],
|
||||
"model": cfg["model"],
|
||||
"dim": len(vec),
|
||||
"sample_norm": round(norm, 4),
|
||||
}
|
||||
except Exception as e:
|
||||
return {
|
||||
"ok": False,
|
||||
"provider": cfg["provider"],
|
||||
"base_url": cfg["base_url"],
|
||||
"model": cfg["model"],
|
||||
"error": f"{type(e).__name__}: {e}",
|
||||
}
|
||||
Reference in New Issue
Block a user