This commit is contained in:
Mikan
2026-06-20 19:13:05 +03:00
parent 32575e217e
commit 8514c63ec6
193 changed files with 22105 additions and 11660 deletions

326
app/core/rag.py Normal file
View File

@@ -0,0 +1,326 @@
"""RAG — retrieval-augmented generation through Qdrant + PostgreSQL.
Two-stage retrieval:
1. Vector search in Qdrant (filtered by world_id).
2. Hydrate full entity/story-entry data from PostgreSQL by IDs.
"""
from __future__ import annotations
import time
import uuid
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import get_settings
from app.core.embeddings import (
HashEmbedder,
OpenAIEmbedder,
build_hash_embedder,
build_openai_embedder,
)
from app.core.logging import get_logger
from app.core.qdrant_client import get_qdrant_client
from app.models import Entity, StoryEntry
_logger = get_logger(__name__)
_embedder_cache: dict[str, Any] = {}
async def get_embedder():
"""Return the active Embedder based on settings.
Falls back to HashEmbedder if the OpenAI embedder cannot be built.
"""
from app.core.settings_service import get_all_settings
# We can't take a DB session here — use a module-level cache.
# On settings change the admin should hit "test embeddings" which clears the cache.
if "active" in _embedder_cache:
return _embedder_cache["active"]
cfg = get_settings()
provider = cfg.embeddings_provider
if provider == "offline_hash":
emb = build_hash_embedder(cfg.embeddings_dimension)
elif provider == "openai":
api_url = cfg.embeddings_api_url or cfg.llm_api_url
api_key = cfg.embeddings_api_key or cfg.llm_api_key
if not api_url:
_logger.warning("embeddings_openai_no_url_fallback_hash")
emb = build_hash_embedder(cfg.embeddings_dimension)
else:
emb = build_openai_embedder(
api_url=api_url,
api_key=api_key,
model=cfg.embeddings_model,
dimension=cfg.embeddings_dimension,
timeout=float(cfg.embeddings_timeout_seconds),
batch_size=cfg.embeddings_batch_size,
)
else:
_logger.warning("embeddings_unknown_provider_fallback_hash", provider=provider)
emb = build_hash_embedder(cfg.embeddings_dimension)
_embedder_cache["active"] = emb
return emb
def reset_embedder_cache() -> None:
"""Clear the cached embedder (used by admin test endpoints after settings change)."""
_embedder_cache.clear()
async def rag_query(
*,
db: AsyncSession,
world_id: uuid.UUID,
query: str,
limit: int = 5,
filter_type: str = "all",
min_score: float = 0.0,
) -> list[dict[str, Any]]:
"""Semantic search over entities + story_entries via Qdrant."""
cfg = get_settings()
prefix = cfg.qdrant_collection_prefix or ""
embedder = await get_embedder()
try:
vecs = await embedder.embed([query[: cfg.embeddings_max_text_chars]])
if not vecs:
return []
query_vec = vecs[0]
except Exception as e: # noqa: BLE001
_logger.warning("rag_query_embed_failed", error=str(e))
return []
client = get_qdrant_client()
from qdrant_client.http.models import (
FieldCondition,
Filter,
MatchValue,
)
world_filter = FieldCondition(
key="world_id", match=MatchValue(value=str(world_id))
)
raw_results: list[dict[str, Any]] = []
if filter_type in ("all", "entities"):
try:
ents = await client.search(
collection_name=f"{prefix}entities",
query_vector=query_vec,
query_filter=Filter(
must=[
world_filter,
FieldCondition(
key="deleted", match=MatchValue(value=False)
),
]
),
limit=limit,
score_threshold=min_score,
with_payload=True,
)
for p in ents:
raw_results.append({
"type": "entity",
"id": p.payload.get("entity_id"),
"score": float(p.score),
"name": p.payload.get("name"),
"entity_type": p.payload.get("entity_type"),
})
except Exception as e: # noqa: BLE001
_logger.warning("rag_query_entities_failed", error=str(e))
if filter_type in ("all", "story_entries"):
try:
sts = await client.search(
collection_name=f"{prefix}story_entries",
query_vector=query_vec,
query_filter=Filter(must=[world_filter]),
limit=limit,
score_threshold=min_score,
with_payload=True,
)
for p in sts:
raw_results.append({
"type": "story_entry",
"id": p.payload.get("entry_id"),
"score": float(p.score),
"entry_type": p.payload.get("entry_type"),
})
except Exception as e: # noqa: BLE001
_logger.warning("rag_query_stories_failed", error=str(e))
# Sort and truncate
raw_results.sort(key=lambda r: r["score"], reverse=True)
top = raw_results[:limit]
return await _hydrate(db, top, world_id)
async def _hydrate(
db: AsyncSession, items: list[dict[str, Any]], world_id: uuid.UUID
) -> list[dict[str, Any]]:
"""Stage 2: pull full records from PostgreSQL by IDs."""
entity_ids = [uuid.UUID(i["id"]) for i in items if i["type"] == "entity"]
story_ids = [uuid.UUID(i["id"]) for i in items if i["type"] == "story_entry"]
ents_map: dict[uuid.UUID, Entity] = {}
stories_map: dict[uuid.UUID, StoryEntry] = {}
if entity_ids:
rows = (
await db.execute(
select(Entity).where(
Entity.id.in_(entity_ids), Entity.world_id == world_id
)
)
).scalars().all()
ents_map = {r.id: r for r in rows}
if story_ids:
rows = (
await db.execute(
select(StoryEntry).where(
StoryEntry.id.in_(story_ids), StoryEntry.world_id == world_id
)
)
).scalars().all()
stories_map = {r.id: r for r in rows}
out: list[dict[str, Any]] = []
for i in items:
if i["type"] == "entity":
ent = ents_map.get(uuid.UUID(i["id"]))
if ent and ent.deleted_at is None:
out.append({
**i,
"content": {
"entity_type": ent.entity_type,
"name": ent.name,
"data": ent.data,
},
})
else:
se = stories_map.get(uuid.UUID(i["id"]))
if se:
out.append({
**i,
"content": {
"text": se.content,
"entry_type": se.entry_type,
"metadata": se.metadata_,
},
})
return out
async def rag_add(
*,
db: AsyncSession,
world_id: uuid.UUID,
content: str,
entry_type: str,
metadata: dict | None = None,
step_id: uuid.UUID | None = None,
) -> StoryEntry:
"""Add a story entry and index it in Qdrant (best-effort)."""
cfg = get_settings()
prefix = cfg.qdrant_collection_prefix or ""
entry = StoryEntry(
world_id=world_id,
content=content,
entry_type=entry_type,
metadata_=metadata or {},
embedding_status="pending",
)
db.add(entry)
await db.flush()
try:
embedder = await get_embedder()
vecs = await embedder.embed([content[: cfg.embeddings_max_text_chars]])
if vecs:
point_id = str(entry.id)
from qdrant_client.http.models import PointStruct
await get_qdrant_client().upsert(
collection_name=f"{prefix}story_entries",
points=[
PointStruct(
id=point_id,
vector=vecs[0],
payload={
"world_id": str(world_id),
"entry_id": point_id,
"entry_type": entry_type,
"step_id": str(step_id) if step_id else None,
"created_at": int(time.time()),
},
)
],
)
entry.qdrant_point_id = point_id
entry.embedding_status = "indexed"
except Exception as e: # noqa: BLE001
_logger.warning("rag_add_embed_failed", entry_id=str(entry.id), error=str(e))
entry.embedding_status = "failed"
await db.flush()
return entry
async def index_entity(
*,
db: AsyncSession,
entity: Entity,
) -> None:
"""Index (or re-index) an entity's vector in Qdrant."""
cfg = get_settings()
prefix = cfg.qdrant_collection_prefix or ""
text = entity.name + " " + _stringify(entity.data)
try:
embedder = await get_embedder()
vecs = await embedder.embed([text[: cfg.embeddings_max_text_chars]])
if not vecs:
return
point_id = str(entity.id)
from qdrant_client.http.models import PointStruct
await get_qdrant_client().upsert(
collection_name=f"{prefix}entities",
points=[
PointStruct(
id=point_id,
vector=vecs[0],
payload={
"world_id": str(entity.world_id),
"entity_id": point_id,
"entity_type": entity.entity_type,
"name": entity.name,
"deleted": entity.deleted_at is not None,
},
)
],
)
entity.qdrant_point_id = point_id
entity.embedding_status = "indexed"
except Exception as e: # noqa: BLE001
_logger.warning("entity_index_failed", entity_id=str(entity.id), error=str(e))
entity.embedding_status = "failed"
await db.flush()
def _stringify(obj: Any) -> str:
import json
try:
return json.dumps(obj, ensure_ascii=False, default=str)
except Exception: # noqa: BLE001
return str(obj)