Astraea: WA family-law multi-agent site (Flask + Ollama semantic RAG)
This commit is contained in:
185
rag.py
Normal file
185
rag.py
Normal file
@@ -0,0 +1,185 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Semantic RAG engine for Astraea — chunk the law 'books', embed, and retrieve.
|
||||
|
||||
stdlib-only (urllib + json + math), no numpy needed. Embeddings and generation
|
||||
run on the Ollama host (nightmare). The vector index is cached to disk and rebuilt
|
||||
incrementally only when a book changes.
|
||||
"""
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
import urllib.request
|
||||
|
||||
from agents import OLLAMA_URL, EMBED_MODEL, BOOKS_DIR, AGENTS
|
||||
|
||||
INDEX_PATH = "/opt/astraea/index.json"
|
||||
|
||||
|
||||
def _post(url, payload, timeout=120):
|
||||
"""POST JSON to Ollama with a no-proxy opener (safe on LAN)."""
|
||||
req = urllib.request.Request(
|
||||
url, data=json.dumps(payload).encode("utf-8"),
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
opener = urllib.request.build_opener(urllib.request.ProxyHandler({}))
|
||||
with opener.open(req, timeout=timeout) as r:
|
||||
return json.loads(r.read().decode("utf-8"))
|
||||
|
||||
|
||||
def embed(text):
|
||||
"""Return a float vector for `text` via the Ollama embeddings endpoint."""
|
||||
d = _post(f"{OLLAMA_URL}/api/embeddings", {"model": EMBED_MODEL, "prompt": text})
|
||||
v = d.get("embedding") or d.get("embeddings")
|
||||
if isinstance(v, list) and v and isinstance(v[0], list):
|
||||
v = v[0]
|
||||
if not v:
|
||||
raise RuntimeError("empty embedding returned")
|
||||
return v
|
||||
|
||||
|
||||
def _cosine(a, b):
|
||||
dot = sum(x * y for x, y in zip(a, b))
|
||||
na = math.sqrt(sum(x * x for x in a))
|
||||
nb = math.sqrt(sum(y * y for y in b))
|
||||
return dot / (na * nb) if na and nb else 0.0
|
||||
|
||||
|
||||
MAX_CHUNK = 1100 # chars — stay safely under the embedding model's context budget
|
||||
|
||||
|
||||
def _split_long(text, limit):
|
||||
"""Split text into pieces <= limit chars, preferring sentence boundaries."""
|
||||
if len(text) <= limit:
|
||||
return [text] if text.strip() else []
|
||||
parts = re.split(r"(?<=[.!?])\s+", text)
|
||||
out = []
|
||||
cur = ""
|
||||
for p in parts:
|
||||
if cur and len(cur) + len(p) + 1 > limit:
|
||||
out.append(cur)
|
||||
cur = p
|
||||
else:
|
||||
cur = (cur + " " + p).strip() if cur else p
|
||||
while len(cur) > limit:
|
||||
out.append(cur[:limit])
|
||||
cur = cur[limit:]
|
||||
if cur.strip():
|
||||
out.append(cur)
|
||||
return out
|
||||
|
||||
|
||||
def _chunk_markdown(text):
|
||||
"""Split markdown into self-contained (title, body) chunks, each <= MAX_CHUNK chars."""
|
||||
lines = text.splitlines()
|
||||
chunks = []
|
||||
cur_title = None
|
||||
cur_buf = []
|
||||
|
||||
def flush():
|
||||
nonlocal cur_title, cur_buf
|
||||
if not cur_buf:
|
||||
return
|
||||
body = "\n".join(cur_buf).strip()
|
||||
cur_buf = []
|
||||
if not body:
|
||||
return
|
||||
for piece in _split_long(body, MAX_CHUNK):
|
||||
chunks.append((cur_title or "Section", piece))
|
||||
|
||||
for line in lines:
|
||||
if re.match(r"^#{1,4}\s+", line):
|
||||
flush()
|
||||
cur_title = re.sub(r"^#{1,4}\s+", "", line).strip()
|
||||
else:
|
||||
cur_buf.append(line)
|
||||
flush()
|
||||
return chunks
|
||||
|
||||
|
||||
def build_index():
|
||||
"""Build (or load from cache) the vector index. Returns list of chunk dicts."""
|
||||
index = []
|
||||
book_files = []
|
||||
for a in AGENTS:
|
||||
for b in a["books"]:
|
||||
path = os.path.join(BOOKS_DIR, b)
|
||||
if os.path.exists(path):
|
||||
book_files.append((path, a["id"]))
|
||||
|
||||
# Decide whether to reuse cache: index.json exists, embed model matches, and no book changed.
|
||||
if os.path.exists(INDEX_PATH):
|
||||
try:
|
||||
with open(INDEX_PATH) as f:
|
||||
cached = json.load(f)
|
||||
if cached.get("embed_model") == EMBED_MODEL:
|
||||
newest_book = max(os.path.getmtime(p) for p, _ in book_files) if book_files else 0
|
||||
if os.path.getmtime(INDEX_PATH) >= newest_book:
|
||||
return cached["chunks"]
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for path, agent_id in book_files:
|
||||
try:
|
||||
with open(path, encoding="utf-8") as f:
|
||||
text = f.read()
|
||||
except Exception:
|
||||
continue
|
||||
base = os.path.basename(path)
|
||||
for title, body in _chunk_markdown(text):
|
||||
# Embed the title + body together for best semantic match.
|
||||
try:
|
||||
vec = embed(f"{title}\n{body}")
|
||||
except Exception:
|
||||
continue # skip any chunk that fails to embed
|
||||
index.append({
|
||||
"agent": agent_id,
|
||||
"source": base,
|
||||
"title": title,
|
||||
"text": body,
|
||||
"vector": vec,
|
||||
})
|
||||
|
||||
if index:
|
||||
try:
|
||||
with open(INDEX_PATH, "w") as f:
|
||||
json.dump({"embed_model": EMBED_MODEL, "chunks": index}, f)
|
||||
except Exception:
|
||||
pass
|
||||
return index
|
||||
|
||||
|
||||
def retrieve(query, agent_id, top_k=5):
|
||||
"""Return top-k relevant chunks for `agent_id` via cosine similarity."""
|
||||
qvec = embed(query)
|
||||
scored = []
|
||||
for c in _INDEX:
|
||||
if c["agent"] != agent_id:
|
||||
continue
|
||||
scored.append((_cosine(qvec, c["vector"]), c))
|
||||
scored.sort(key=lambda x: x[0], reverse=True)
|
||||
return [(s, c) for s, c in scored[:top_k]]
|
||||
|
||||
|
||||
# Global index, built lazily once in a background thread.
|
||||
_INDEX = []
|
||||
_index_lock = threading.Lock()
|
||||
_index_ready = False
|
||||
|
||||
|
||||
def start_background_index():
|
||||
def _run():
|
||||
global _index_ready
|
||||
try:
|
||||
idx = build_index()
|
||||
with _index_lock:
|
||||
_INDEX.clear()
|
||||
_INDEX.extend(idx)
|
||||
finally:
|
||||
_index_ready = True
|
||||
threading.Thread(target=_run, daemon=True).start()
|
||||
|
||||
|
||||
def index_status():
|
||||
return {"ready": _index_ready, "chunks": len(_INDEX)}
|
||||
Reference in New Issue
Block a user