This commit is contained in:
Mikan
2026-06-19 11:28:04 +03:00
commit 53c89829a8
80 changed files with 12482 additions and 0 deletions

463
backend/app/core/rag.py Normal file
View 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}",
}