rag-backend / retrieval /vectorstore.py
imtrt004
feat: line number and multi docs
391fc60
Raw
History Blame Contribute Delete
4.62 kB
from __future__ import annotations
from supabase import create_client, Client
from datetime import datetime
import os
def _client() -> Client:
return create_client(os.environ["SUPABASE_URL"], os.environ["SUPABASE_KEY"])
# ── ChunkResult: rich return type for similarity search ──────────────────────
class ChunkResult:
"""Holds chunk text, its page of origin, and source document."""
__slots__ = ("text", "page_number", "doc_id")
def __init__(self, text: str, page_number: int, doc_id: str):
self.text = text
self.page_number = page_number
self.doc_id = doc_id
# Behaves like a plain string so old code that does `"\n".join(chunks)` still works
def __str__(self) -> str: return self.text
def __repr__(self) -> str: return f"ChunkResult(doc={self.doc_id[:8]}, page={self.page_number})"
def store_chunks(
doc_id: str,
user_id: str,
chunks, # list[ChunkMeta] or list[str]
embeddings: list[list[float]],
expires_at: datetime,
) -> None:
client = _client()
rows = []
for i, (chunk, embedding) in enumerate(zip(chunks, embeddings)):
# Support both ChunkMeta objects (with .text/.page_number) and plain strings
text = chunk.text if hasattr(chunk, "text") else str(chunk)
page_number = chunk.page_number if hasattr(chunk, "page_number") else 1
rows.append({
"doc_id": doc_id,
"user_id": user_id,
"chunk_text": text,
"embedding": embedding,
"chunk_index": i,
"page_number": page_number,
"expires_at": expires_at.isoformat(),
})
# Insert in batches of 100 to avoid payload limits
for i in range(0, len(rows), 100):
client.table("chunks").insert(rows[i : i + 100]).execute()
def similarity_search(
doc_id: str,
query_embedding: list[float],
top_k: int = 5,
) -> list[ChunkResult]:
"""Search a single document and return rich ChunkResult objects."""
client = _client()
result = client.rpc(
"match_chunks",
{
"query_embedding": query_embedding,
"doc_id_filter": doc_id,
"match_count": top_k,
},
).execute()
return [
ChunkResult(
text = r["chunk_text"],
page_number = r.get("page_number", 1),
doc_id = str(r.get("doc_id", doc_id)),
)
for r in result.data
]
def similarity_search_multi(
doc_ids: list[str],
query_embedding: list[float],
top_k: int = 20,
) -> list[ChunkResult]:
"""Search across multiple documents and return rich ChunkResult objects."""
if not doc_ids:
return []
if len(doc_ids) == 1:
return similarity_search(doc_ids[0], query_embedding, top_k)
client = _client()
result = client.rpc(
"match_chunks_multi",
{
"query_embedding": query_embedding,
"doc_ids_filter": doc_ids,
"match_count": top_k,
},
).execute()
return [
ChunkResult(
text = r["chunk_text"],
page_number = r.get("page_number", 1),
doc_id = str(r["doc_id"]),
)
for r in result.data
]
def get_all_chunks(doc_id: str) -> list[ChunkResult]:
"""Return every chunk for a document in order, for full-context retrieval."""
client = _client()
result = (
client.table("chunks")
.select("chunk_text, page_number, doc_id")
.eq("doc_id", doc_id)
.order("chunk_index", desc=False)
.execute()
)
return [
ChunkResult(
text = r["chunk_text"],
page_number = r.get("page_number", 1),
doc_id = str(r.get("doc_id", doc_id)),
)
for r in result.data
]
def get_all_chunks_multi(doc_ids: list[str]) -> list[ChunkResult]:
"""Return all chunks for multiple documents in document+chunk order."""
if not doc_ids:
return []
client = _client()
result = (
client.table("chunks")
.select("chunk_text, page_number, doc_id, chunk_index")
.in_("doc_id", doc_ids)
.order("doc_id", desc=False)
.order("chunk_index", desc=False)
.execute()
)
return [
ChunkResult(
text = r["chunk_text"],
page_number = r.get("page_number", 1),
doc_id = str(r["doc_id"]),
)
for r in result.data
]