first commit
This commit is contained in:
@@ -0,0 +1,402 @@
|
||||
"""Async PostgreSQL database layer with pgvector support."""
|
||||
import json
|
||||
import uuid
|
||||
from typing import Any
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
import asyncpg
|
||||
import numpy as np
|
||||
|
||||
from app.config import get_settings
|
||||
|
||||
|
||||
class Database:
|
||||
"""Handles all PostgreSQL/pgvector operations."""
|
||||
|
||||
def __init__(self, pool: asyncpg.Pool):
|
||||
self.pool = pool
|
||||
|
||||
@classmethod
|
||||
async def create_pool(cls) -> asyncpg.Pool:
|
||||
settings = get_settings()
|
||||
pool = await asyncpg.create_pool(
|
||||
host=settings.db_host,
|
||||
port=settings.db_port,
|
||||
database=settings.db_name,
|
||||
user=settings.db_user,
|
||||
password=settings.db_password,
|
||||
min_size=2,
|
||||
max_size=10,
|
||||
)
|
||||
return pool
|
||||
|
||||
@classmethod
|
||||
@asynccontextmanager
|
||||
async def connection(cls):
|
||||
pool = await cls.create_pool()
|
||||
async with pool.acquire() as conn:
|
||||
yield conn
|
||||
await pool.close()
|
||||
|
||||
async def init_schema(self):
|
||||
"""Run initial SQL schema from migrations."""
|
||||
import os
|
||||
migrations_path = os.path.join(
|
||||
os.path.dirname(__file__), "..", "..", "migrations", "init.sql"
|
||||
)
|
||||
async with self.connection() as conn:
|
||||
with open(migrations_path) as f:
|
||||
await conn.execute(f.read())
|
||||
|
||||
async def upsert_document(
|
||||
self, filename: str, doc_id: str | None, mime_type: str,
|
||||
file_path: str, status: str, page_count: int,
|
||||
full_text: str, metadata: dict
|
||||
) -> str:
|
||||
async with self.connection() as conn:
|
||||
pk = doc_id or str(uuid.uuid4())
|
||||
await conn.execute(
|
||||
"""INSERT INTO documents (id, filename, doc_id, mime_type, file_path,
|
||||
status, page_count, full_text, metadata)
|
||||
VALUES ($1::uuid, $2, $1::text, $3, $4, $5, $6, $7, $8::jsonb)
|
||||
ON CONFLICT (filename) DO UPDATE SET
|
||||
status=EXCLUDED.status, page_count=EXCLUDED.page_count,
|
||||
full_text=EXCLUDED.full_text, metadata=EXCLUDED.metadata,
|
||||
updated_at=NOW()""",
|
||||
str(pk), filename, mime_type, file_path,
|
||||
status, page_count, full_text, json.dumps(metadata),
|
||||
)
|
||||
return pk
|
||||
|
||||
async def get_document(self, doc_id: str | uuid.UUID) -> dict | None:
|
||||
async with self.connection() as conn:
|
||||
row = await conn.fetchrow(
|
||||
"SELECT * FROM documents WHERE id=$1", doc_id
|
||||
)
|
||||
return dict(row) if row else None
|
||||
|
||||
async def list_documents(self, status: str | None = None) -> list[dict]:
|
||||
async with self.connection() as conn:
|
||||
if status:
|
||||
rows = await conn.fetch(
|
||||
"SELECT * FROM documents WHERE $1=ANY(string_to_array(status, ',')) ORDER BY created_at DESC",
|
||||
status,
|
||||
)
|
||||
else:
|
||||
rows = await conn.fetch("SELECT * FROM documents ORDER BY created_at DESC")
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
async def chunk_document(
|
||||
self, doc_id: str, content: str, polygon: dict | None,
|
||||
page_num: int, block_index: int, chunk_type: str = "text"
|
||||
) -> str:
|
||||
vec_str = "[" + ",".join(str(x) for x in self._extract_vector(content)) + "]"
|
||||
chunk_id = str(uuid.uuid4())
|
||||
async with self.connection() as conn:
|
||||
await conn.execute(
|
||||
"""INSERT INTO chunks (id, doc_id, content, vector, page_num,
|
||||
block_index, polygon, chunk_type)
|
||||
VALUES ($1, $2, $3, $4::vector, $5, $6, $7::jsonb, $8)""",
|
||||
chunk_id, str(doc_id), content, vec_str, page_num,
|
||||
block_index, json.dumps(polygon) if polygon else None, chunk_type,
|
||||
)
|
||||
return chunk_id
|
||||
|
||||
async def batch_chunk(self, doc_id: str, data: list[dict]):
|
||||
"""Insert multiple chunks at once."""
|
||||
vec_data = []
|
||||
for d in data:
|
||||
content = d.get("content", "")
|
||||
vec_str = "[" + ",".join(str(x) for x in self._extract_vector(content)) + "]"
|
||||
vec_data.append((
|
||||
str(doc_id), content, vec_str,
|
||||
d.get("page_num", 0), d.get("block_index", 0),
|
||||
json.dumps(d.get("polygon")) if d.get("polygon") else None,
|
||||
d.get("chunk_type", "text"),
|
||||
))
|
||||
async with self.connection() as conn:
|
||||
await conn.executemany(
|
||||
"""INSERT INTO chunks (doc_id, content, vector, page_num,
|
||||
block_index, polygon, chunk_type)
|
||||
VALUES ($1, $2, $3::vector, $4, $5, $6::jsonb, $7)""",
|
||||
vec_data
|
||||
)
|
||||
|
||||
async def vector_search(
|
||||
self, query_vector: list[float], doc_id: str | None = None,
|
||||
limit: int = 20, min_score: float = 0.0
|
||||
) -> list[dict]:
|
||||
"""Find similar chunks using pgvector cosine similarity."""
|
||||
query_vec = "[" + ",".join(str(x) for x in query_vector) + "]"
|
||||
async with self.connection() as conn:
|
||||
if doc_id:
|
||||
rows = await conn.fetch(
|
||||
"""SELECT id, content, doc_id, page_num, polygon,
|
||||
(1 - (vector <-> $4::vector) / 2) as similarity
|
||||
FROM chunks WHERE doc_id = $1
|
||||
AND (1 - (vector <-> $4::vector) / 2) >= $3
|
||||
ORDER BY vector <-> $4 LIMIT $2""",
|
||||
str(doc_id), limit, min_score, query_vec,
|
||||
)
|
||||
else:
|
||||
rows = await conn.fetch(
|
||||
"""SELECT id, content, doc_id, page_num, polygon,
|
||||
(1 - (vector <-> $3::vector) / 2) as similarity
|
||||
FROM chunks
|
||||
WHERE (1 - (vector <-> $3::vector) / 2) >= $2
|
||||
ORDER BY vector <-> $3 LIMIT $1""",
|
||||
limit, min_score, query_vec,
|
||||
)
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
async def semantic_search(
|
||||
self, query_text: str, limit: int = 20, min_score: float = 0.0
|
||||
) -> list[dict]:
|
||||
"""Semantic search by embedding the query text."""
|
||||
from app.core import get_embedding
|
||||
query_vec = get_embedding(query_text)
|
||||
return await self.vector_search(query_vec, limit=limit, min_score=min_score)
|
||||
|
||||
async def search_chunks_text(
|
||||
self, query: str, doc_id: str | None = None, limit: int = 20
|
||||
) -> list[dict]:
|
||||
"""Text-based search using trigram similarity."""
|
||||
async with self.connection() as conn:
|
||||
if doc_id:
|
||||
rows = await conn.fetch(
|
||||
"""SELECT id, content, doc_id, page_num,
|
||||
ts_rank(to_tsvector('simple', content),
|
||||
plainto_tsquery('simple', $4)) as rank
|
||||
FROM chunks WHERE doc_id = $1
|
||||
AND to_tsvector('simple', content) @@ plainto_tsquery('simple', $4)
|
||||
ORDER BY rank DESC LIMIT $2""",
|
||||
str(doc_id), limit, query,
|
||||
)
|
||||
else:
|
||||
rows = await conn.fetch(
|
||||
"""SELECT id, content, doc_id, page_num,
|
||||
ts_rank(to_tsvector('simple', content),
|
||||
plainto_tsquery('simple', $3)) as rank
|
||||
FROM chunks
|
||||
WHERE to_tsvector('simple', content) @@ plainto_tsquery('simple', $3)
|
||||
ORDER BY rank DESC LIMIT $1""",
|
||||
limit, query,
|
||||
)
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
async def store_memory(
|
||||
self, session_id: str, content: str, memory_type: str = "fact",
|
||||
importance: int = 3, source_doc_id: str | None = None
|
||||
) -> str:
|
||||
vec = self._extract_vector(content)
|
||||
mem_id = str(uuid.uuid4())
|
||||
async with self.connection() as conn:
|
||||
await conn.execute(
|
||||
"""INSERT INTO memories (id, session_id, content, vector,
|
||||
memory_type, importance, source_doc_id)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7)""",
|
||||
mem_id, str(session_id), content, vec, memory_type,
|
||||
importance, str(source_doc_id) if source_doc_id else None,
|
||||
)
|
||||
return mem_id
|
||||
|
||||
async def get_memories(
|
||||
self, session_id: str, limit: int = 50
|
||||
) -> list[dict]:
|
||||
async with self.connection() as conn:
|
||||
rows = await conn.fetch(
|
||||
"SELECT * FROM memories WHERE session_id=$1 ORDER BY importance DESC, created_at DESC LIMIT $2",
|
||||
str(session_id), limit,
|
||||
)
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
async def memory_similarity_search(
|
||||
self, query_text: str, limit: int = 10
|
||||
) -> list[dict]:
|
||||
from app.core import get_embedding
|
||||
query_vec = get_embedding(query_text)
|
||||
query_str = "[" + ",".join(str(x) for x in query_vec) + "]"
|
||||
async with self.connection() as conn:
|
||||
rows = await conn.fetch(
|
||||
"""SELECT id, content, memory_type, importance,
|
||||
(1 - (vector <-> $2::vector) / 2) as similarity
|
||||
FROM memories ORDER BY vector <-> $2 LIMIT $1""",
|
||||
limit, query_str,
|
||||
)
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
async def store_finding(
|
||||
self, session_id: str, question: str, answer: str,
|
||||
summary: str, agent_name: str, confidence: float,
|
||||
relevant_chunks: list | None = None
|
||||
) -> str:
|
||||
question_vec_str = "[" + ",".join(str(x) for x in self._extract_vector(answer)) + "]"
|
||||
finding_id = str(uuid.uuid4())
|
||||
async with self.connection() as conn:
|
||||
await conn.execute(
|
||||
"""INSERT INTO findings (id, session_id, question, answer,
|
||||
summary, vector, relevant_chunks, agent_name, confidence)
|
||||
VALUES ($1, $2, $3, $4, $5, $6::vector, $7::jsonb, $8, $9)""",
|
||||
finding_id, str(session_id), question, answer,
|
||||
summary, question_vec_str, json.dumps(relevant_chunks or []),
|
||||
agent_name, confidence,
|
||||
)
|
||||
return finding_id
|
||||
|
||||
async def get_findings(self, session_id: str) -> list[dict]:
|
||||
async with self.connection() as conn:
|
||||
rows = await conn.fetch(
|
||||
"SELECT * FROM findings WHERE session_id=$1 ORDER BY created_at DESC",
|
||||
str(session_id),
|
||||
)
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
async def create_session(self, query: str) -> str:
|
||||
session_id = str(uuid.uuid4())
|
||||
async with self.connection() as conn:
|
||||
await conn.execute(
|
||||
"INSERT INTO research_sessions (id, query, status, documents, findings) VALUES ($1, $2, $3, '[]', '[]')",
|
||||
session_id, query, "running",
|
||||
)
|
||||
return session_id
|
||||
|
||||
async def update_session(
|
||||
self, session_id: str, status: str | None = None,
|
||||
documents: list | dict | None = None, findings: list | dict | None = None
|
||||
):
|
||||
async with self.connection() as conn:
|
||||
if status:
|
||||
if documents or findings:
|
||||
await conn.execute(
|
||||
"UPDATE research_sessions SET status=$1, documents=$2::jsonb, findings=$3::jsonb, completed_at=NOW() WHERE id=$4",
|
||||
status, json.dumps(documents or []), json.dumps(findings or []),
|
||||
session_id,
|
||||
)
|
||||
else:
|
||||
await conn.execute(
|
||||
"UPDATE research_sessions SET status=$1 WHERE id=$2",
|
||||
status, session_id,
|
||||
)
|
||||
else:
|
||||
await conn.execute(
|
||||
"UPDATE research_sessions SET documents=$1::jsonb, findings=$2::jsonb WHERE id=$3",
|
||||
json.dumps(documents or []), json.dumps(findings or []),
|
||||
session_id,
|
||||
)
|
||||
|
||||
async def get_session(self, session_id: str) -> dict | None:
|
||||
async with self.connection() as conn:
|
||||
row = await conn.fetchrow(
|
||||
"SELECT * FROM research_sessions WHERE id=$1", session_id
|
||||
)
|
||||
return dict(row) if row else None
|
||||
|
||||
async def get_chunk(self, chunk_id: str) -> dict | None:
|
||||
async with self.connection() as conn:
|
||||
row = await conn.fetchrow(
|
||||
"SELECT * FROM chunks WHERE id=$1", chunk_id
|
||||
)
|
||||
return dict(row) if row else None
|
||||
|
||||
@staticmethod
|
||||
def _extract_vector(text: str) -> list[float]:
|
||||
"""Generate a lightweight embedding vector directly (no external call for speed)."""
|
||||
# Use a fast hash-based feature vector as fallback
|
||||
# In production this calls the server; locally we use a fast approximation
|
||||
import hashlib
|
||||
feature_dim = 4096
|
||||
vec = np.zeros(feature_dim, dtype=np.float32)
|
||||
# Create deterministic features from character trigrams
|
||||
trigrams = [text[i:i+3] for i in range(len(text)-2)]
|
||||
for i, tri in enumerate(trigrams):
|
||||
hash_val = hash(tri) & 0xFFFFFFFF
|
||||
# convert to signed 32-bit
|
||||
if hash_val >= 0x80000000:
|
||||
hash_val -= 0x100000000
|
||||
start_idx = (hash_val % feature_dim)
|
||||
end_idx = min(start_idx + 5, feature_dim)
|
||||
for j, byte in enumerate(hash_val.to_bytes(4, "big", signed=True)):
|
||||
idx = (start_idx + j) % feature_dim
|
||||
vec[idx] = (byte / 127.0) * np.sin(i * 0.1)
|
||||
norm = np.linalg.norm(vec)
|
||||
if norm > 0:
|
||||
vec = vec / norm
|
||||
return vec.tolist()
|
||||
|
||||
async def get_doc_chunks(self, doc_id: str) -> list[dict]:
|
||||
"""Get all chunks for a document."""
|
||||
async with self.connection() as conn:
|
||||
rows = await conn.fetch(
|
||||
"SELECT * FROM chunks WHERE doc_id=$1 ORDER BY page_num, block_index",
|
||||
str(doc_id),
|
||||
)
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
# ── Pipeline persistence ───────────────────────────────
|
||||
|
||||
async def save_pipeline_stage(
|
||||
self, session_id: str, stage: str, output: str, state: dict
|
||||
) -> str:
|
||||
"""Persist intermediate pipeline stage (triage/evidence/synthesis)."""
|
||||
stage_id = str(uuid.uuid4())
|
||||
async with self.connection() as conn:
|
||||
await conn.execute(
|
||||
"""INSERT INTO pipeline_stages (id, session_id, stage, output, state)
|
||||
VALUES ($1, $2, $3, $4, $5::jsonb)""",
|
||||
stage_id, str(session_id), stage, output, json.dumps(state),
|
||||
)
|
||||
return stage_id
|
||||
|
||||
async def get_pipeline_stages(self, session_id: str) -> list[dict]:
|
||||
"""Get all pipeline stages for a session in order."""
|
||||
async with self.connection() as conn:
|
||||
rows = await conn.fetch(
|
||||
"SELECT * FROM pipeline_stages WHERE session_id=$1 ORDER BY stage, created_at",
|
||||
str(session_id),
|
||||
)
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
async def save_structured_evidence(
|
||||
self, session_id: str,
|
||||
rows: list[dict]
|
||||
) -> list[str]:
|
||||
"""Save evidence-extraction rows with full traceability fields."""
|
||||
ids = []
|
||||
for row in rows:
|
||||
finding_id = str(uuid.uuid4())
|
||||
async with self.connection() as conn:
|
||||
await conn.execute(
|
||||
"""INSERT INTO findings (id, session_id, question, answer, summary,
|
||||
agent_name, confidence, relevant_chunks)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8)""",
|
||||
finding_id, str(session_id),
|
||||
row.get("topic", ""),
|
||||
row.get("evidence", ""),
|
||||
row.get("description", ""),
|
||||
"evidence_extraction",
|
||||
{"High": 0.9, "Medium": 0.6, "Low": 0.3}.get(row.get("confidence", "Medium"), 0.6),
|
||||
json.dumps({
|
||||
"evidence_type": row.get("evidence_type"),
|
||||
"trace_ref": row.get("trace_ref"),
|
||||
"review_needed": row.get("review_needed", False),
|
||||
"confidence": row.get("confidence"),
|
||||
"source_doc": row.get("source_doc"),
|
||||
}),
|
||||
)
|
||||
ids.append(finding_id)
|
||||
return ids
|
||||
|
||||
async def get_structured_findings(self, session_id: str) -> list[dict]:
|
||||
"""Get structured evidence findings for a session."""
|
||||
async with self.connection() as conn:
|
||||
rows = await conn.fetch(
|
||||
"""SELECT f.*, f.relevant_chunks::jsonb as meta
|
||||
FROM findings f
|
||||
WHERE f.session_id=$1 AND f.agent_name='evidence_extraction'
|
||||
ORDER BY f.created_at""",
|
||||
str(session_id),
|
||||
)
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
|
||||
db: Database | None = None
|
||||
Reference in New Issue
Block a user