# -*- coding: utf-8 -*- """Per-user document index — ingest uploaded files into a vector array so the attorney can reference the user's actual documents alongside the case law. Reuses rag.py's embedder + chunker + cosine. stdlib + optional pymupdf (PDF). """ import json import os import re from rag import embed, _split_long, _cosine, MAX_CHUNK INDEX_DIR = "/opt/astraea/indexes" def _path(uid): return os.path.join(INDEX_DIR, f"user_{uid}.json") def _docx_text(path): import zipfile try: with zipfile.ZipFile(path) as z: xml = z.read("word/document.xml").decode("utf-8", "ignore") out = [] for p in re.split(r"", xml): ts = re.findall(r"]*>(.*?)", p) if ts: out.append("".join(ts)) return "\n".join(out) except Exception: return "" def _pdf_text(path): try: import fitz # pymupdf doc = fitz.open(path) parts = [] for page in doc: parts.append(page.get_text()) doc.close() return "\n".join(parts) except Exception: return "" def extract_text(orig_name, path): ext = os.path.splitext(orig_name)[1].lower() if ext == ".docx": return _docx_text(path) if ext == ".pdf": return _pdf_text(path) if ext in (".txt", ".md", ".csv", ".log", ".json", ".html", ".xml", ".rtf"): try: with open(path, "r", encoding="utf-8", errors="replace") as f: return f.read() except Exception: return "" return "" def build(uid, files, upload_dir): """Embed every file's text into a per-user vector index. Returns chunk count.""" chunks = [] for rec in files: path = os.path.join(upload_dir, rec["filename"]) if not os.path.exists(path): continue text = extract_text(rec.get("orig_name", rec["filename"]), path) if not text or not text.strip(): continue for piece in _split_long(text, MAX_CHUNK): chunks.append({"source": rec.get("orig_name", rec["filename"]), "text": piece, "vector": None}) for c in chunks: try: c["vector"] = embed(c["text"]) except Exception: c["vector"] = None chunks = [c for c in chunks if c["vector"]] os.makedirs(INDEX_DIR, exist_ok=True) try: with open(_path(uid), "w") as f: json.dump({"chunks": chunks}, f) except Exception: pass return len(chunks) def retrieve(uid, query, top_k=4): if not os.path.exists(_path(uid)): return [] try: with open(_path(uid)) as f: data = json.load(f) except Exception: return [] qvec = embed(query) scored = [] for c in data.get("chunks", []): v = c.get("vector") if not v: continue scored.append((_cosine(qvec, v), c)) scored.sort(key=lambda x: x[0], reverse=True) return [{"source": c["source"], "text": c["text"], "score": round(s, 3)} for s, c in scored[:top_k] if s > 0.15] def status(uid): if not os.path.exists(_path(uid)): return {"indexed": False, "chunks": 0} try: with open(_path(uid)) as f: data = json.load(f) return {"indexed": True, "chunks": len(data.get("chunks", []))} except Exception: return {"indexed": False, "chunks": 0}