Files
mosaic/storage/db.py
T

448 lines
17 KiB
Python

"""SQLite database interface for Mosaic with FTS5 and WAL mode."""
import sqlite3
import logging
import re
from pathlib import Path
from typing import Optional
import datetime
from .models import (
Document, ExtractedTool, TTP, Tradecraft, Entity,
Infrastructure, ExtractionStatus, TokenUsage
)
from .migrations import MigrationRunner
logger = logging.getLogger(__name__)
# FTS5 special characters that need escaping
FTS5_SPECIAL = re.compile(r'["\*\(\)\-\:]')
MAX_FTS_QUERY_LENGTH = 500
class MosaicDB:
"""SQLite database with FTS5 for Mosaic intelligence storage."""
def __init__(self, db_path: str = "mosaic.db"):
self.db_path = db_path
self.conn = self._connect()
self._setup()
def _connect(self) -> sqlite3.Connection:
conn = sqlite3.connect(self.db_path)
conn.row_factory = sqlite3.Row
conn.execute("PRAGMA journal_mode=WAL")
conn.execute("PRAGMA foreign_keys=ON")
conn.execute("PRAGMA busy_timeout=5000")
return conn
def _setup(self):
runner = MigrationRunner(self.conn)
runner.run()
# Crash recovery: reset any IN_PROGRESS states
self._reset_in_progress()
def _reset_in_progress(self):
"""Reset stale IN_PROGRESS/DOWNLOADING states on startup (crash recovery)."""
self.conn.execute(
"UPDATE documents SET status = 'PENDING' WHERE status = 'DOWNLOADING'"
)
self.conn.execute(
"UPDATE documents SET parse_status = 'PENDING' WHERE parse_status = 'IN_PROGRESS'"
)
self.conn.execute(
"UPDATE extraction_status SET status = 'PENDING' WHERE status = 'IN_PROGRESS'"
)
self.conn.commit()
logger.debug("Reset stale in-progress states (crash recovery)")
def close(self):
if self.conn:
self.conn.close()
# --- Document operations ---
def insert_document(self, doc: Document) -> bool:
"""Insert or update a document. Returns True if new, False if updated."""
existing = self.get_document(doc.id)
d = doc.to_dict()
if existing:
cols = ', '.join(f"{k} = ?" for k in d if k != 'id')
vals = [v for k, v in d.items() if k != 'id'] + [doc.id]
self.conn.execute(f"UPDATE documents SET {cols} WHERE id = ?", vals)
self.conn.commit()
return False
else:
cols = ', '.join(d.keys())
placeholders = ', '.join('?' * len(d))
self.conn.execute(
f"INSERT INTO documents ({cols}) VALUES ({placeholders})",
list(d.values())
)
self.conn.commit()
return True
def get_document(self, doc_id: str) -> Optional[Document]:
cursor = self.conn.execute(
"SELECT * FROM documents WHERE id = ?", (doc_id,)
)
row = cursor.fetchone()
return Document.from_row(row) if row else None
def get_document_by_url(self, url: str) -> Optional[Document]:
cursor = self.conn.execute(
"SELECT * FROM documents WHERE source_url = ?", (url,)
)
row = cursor.fetchone()
return Document.from_row(row) if row else None
def get_documents_by_source(self, source: str, status: Optional[str] = None,
parse_status: Optional[str] = None) -> list[Document]:
query = "SELECT * FROM documents WHERE source = ?"
params: list = [source]
if status:
query += " AND status = ?"
params.append(status)
if parse_status:
query += " AND parse_status = ?"
params.append(parse_status)
cursor = self.conn.execute(query, params)
return [Document.from_row(row) for row in cursor.fetchall()]
def get_documents_by_date_range(self, source: Optional[str] = None,
since: Optional[str] = None,
until: Optional[str] = None) -> list[Document]:
query = "SELECT * FROM documents WHERE 1=1"
params: list = []
if source:
query += " AND source = ?"
params.append(source)
if since:
query += " AND date >= ?"
params.append(since)
if until:
query += " AND date <= ?"
params.append(until)
cursor = self.conn.execute(query, params)
return [Document.from_row(row) for row in cursor.fetchall()]
def get_unextracted_documents(self, source: Optional[str] = None,
extractor: Optional[str] = None,
since_last_run: bool = False) -> list[Document]:
"""Get documents that need extraction."""
if extractor:
query = """
SELECT d.* FROM documents d
LEFT JOIN extraction_status es
ON d.id = es.doc_id AND es.extractor = ?
WHERE d.parse_status = 'PARSED'
AND (es.status IS NULL OR es.status = 'PENDING')
"""
params: list = [extractor]
else:
query = "SELECT * FROM documents WHERE parse_status = 'PARSED'"
params = []
if source:
query += " AND d.source = ?" if extractor else " AND source = ?"
params.append(source)
if since_last_run:
query += " AND d.last_extracted_at IS NULL" if extractor else " AND last_extracted_at IS NULL"
cursor = self.conn.execute(query, params)
return [Document.from_row(row) for row in cursor.fetchall()]
def update_document_status(self, doc_id: str, status: Optional[str] = None,
parse_status: Optional[str] = None):
updates = []
params: list = []
if status is not None:
updates.append("status = ?")
params.append(status)
if parse_status is not None:
updates.append("parse_status = ?")
params.append(parse_status)
if updates:
params.append(doc_id)
self.conn.execute(
f"UPDATE documents SET {', '.join(updates)} WHERE id = ?",
params
)
self.conn.commit()
def update_document_extraction_time(self, doc_id: str):
self.conn.execute(
"UPDATE documents SET last_extracted_at = ? WHERE id = ?",
(datetime.datetime.datetime.now(datetime.timezone.utc).isoformat(), doc_id)
)
self.conn.commit()
# --- Extracted data operations ---
def insert_extracted_tool(self, tool: ExtractedTool) -> bool:
"""Insert tool if not duplicate. Returns True if inserted."""
if not tool.dedup_key:
tool.compute_dedup_key()
d = tool.to_dict()
try:
cols = ', '.join(k for k in d if k != 'id')
placeholders = ', '.join('?' for k in d if k != 'id')
self.conn.execute(
f"INSERT INTO extracted_tools ({cols}) VALUES ({placeholders})",
[v for k, v in d.items() if k != 'id']
)
self.conn.commit()
return True
except sqlite3.IntegrityError:
# Duplicate dedup_key — merge source_doc_ids
cursor = self.conn.execute(
"SELECT source_doc_ids FROM extracted_tools WHERE dedup_key = ?",
(tool.dedup_key,)
)
row = cursor.fetchone()
if row:
import json
existing_ids = json.loads(row['source_doc_ids'])
merged = list(set(existing_ids + tool.source_doc_ids))
self.conn.execute(
"UPDATE extracted_tools SET source_doc_ids = ? WHERE dedup_key = ?",
(json.dumps(merged), tool.dedup_key)
)
self.conn.commit()
return False
def insert_ttp(self, ttp: TTP) -> bool:
if not ttp.dedup_key:
ttp.compute_dedup_key()
d = ttp.to_dict()
try:
cols = ', '.join(k for k in d if k != 'id')
placeholders = ', '.join('?' for k in d if k != 'id')
self.conn.execute(
f"INSERT INTO ttps ({cols}) VALUES ({placeholders})",
[v for k, v in d.items() if k != 'id']
)
self.conn.commit()
return True
except sqlite3.IntegrityError:
import json
cursor = self.conn.execute(
"SELECT source_doc_ids FROM ttps WHERE dedup_key = ?",
(ttp.dedup_key,)
)
row = cursor.fetchone()
if row:
existing_ids = json.loads(row['source_doc_ids'])
merged = list(set(existing_ids + ttp.source_doc_ids))
self.conn.execute(
"UPDATE ttps SET source_doc_ids = ? WHERE dedup_key = ?",
(json.dumps(merged), ttp.dedup_key)
)
self.conn.commit()
return False
def insert_tradecraft(self, tc: Tradecraft):
d = tc.to_dict()
cols = ', '.join(k for k in d if k != 'id')
placeholders = ', '.join('?' for k in d if k != 'id')
self.conn.execute(
f"INSERT INTO tradecraft ({cols}) VALUES ({placeholders})",
[v for k, v in d.items() if k != 'id']
)
self.conn.commit()
def insert_entity(self, entity: Entity):
d = entity.to_dict()
cols = ', '.join(k for k in d if k != 'id')
placeholders = ', '.join('?' for k in d if k != 'id')
self.conn.execute(
f"INSERT INTO entities ({cols}) VALUES ({placeholders})",
[v for k, v in d.items() if k != 'id']
)
self.conn.commit()
def insert_infrastructure(self, infra: Infrastructure):
d = infra.to_dict()
cols = ', '.join(k for k in d if k != 'id')
placeholders = ', '.join('?' for k in d if k != 'id')
self.conn.execute(
f"INSERT INTO infrastructure ({cols}) VALUES ({placeholders})",
[v for k, v in d.items() if k != 'id']
)
self.conn.commit()
# --- Extraction status ---
def get_extraction_status(self, doc_id: str, extractor: str) -> Optional[ExtractionStatus]:
cursor = self.conn.execute(
"SELECT * FROM extraction_status WHERE doc_id = ? AND extractor = ?",
(doc_id, extractor)
)
row = cursor.fetchone()
return ExtractionStatus.from_row(row) if row else None
def set_extraction_status(self, doc_id: str, extractor: str,
status: str, error: Optional[str] = None,
token_count: int = 0):
self.conn.execute("""
INSERT INTO extraction_status (doc_id, extractor, status, error, last_run, token_count)
VALUES (?, ?, ?, ?, ?, ?)
ON CONFLICT(doc_id, extractor)
DO UPDATE SET status = ?, error = ?, last_run = ?, token_count = ?
""", (
doc_id, extractor, status, error, datetime.datetime.now(datetime.timezone.utc).isoformat(), token_count,
status, error, datetime.datetime.now(datetime.timezone.utc).isoformat(), token_count
))
self.conn.commit()
# --- Token usage ---
def log_token_usage(self, usage: TokenUsage):
d = usage.to_dict()
cols = ', '.join(k for k in d if k != 'id')
placeholders = ', '.join('?' for k in d if k != 'id')
self.conn.execute(
f"INSERT INTO token_usage ({cols}) VALUES ({placeholders})",
[v for k, v in d.items() if k != 'id']
)
self.conn.commit()
def get_total_tokens(self, source: Optional[str] = None) -> dict:
if source:
cursor = self.conn.execute(
"SELECT SUM(input_tokens) as input, SUM(output_tokens) as output FROM token_usage WHERE source = ?",
(source,)
)
else:
cursor = self.conn.execute(
"SELECT SUM(input_tokens) as input, SUM(output_tokens) as output FROM token_usage"
)
row = cursor.fetchone()
return {
'input_tokens': row['input'] or 0,
'output_tokens': row['output'] or 0,
'total_tokens': (row['input'] or 0) + (row['output'] or 0)
}
# --- FTS5 search ---
def search_documents(self, query: str, source: Optional[str] = None,
limit: int = 50) -> list[Document]:
"""Full-text search with sanitized query."""
sanitized = self._sanitize_fts_query(query)
if not sanitized:
return []
if source:
sql = """
SELECT d.* FROM documents d
JOIN documents_fts fts ON d.rowid = fts.rowid
WHERE documents_fts MATCH ? AND d.source = ?
ORDER BY rank
LIMIT ?
"""
params = [sanitized, source, limit]
else:
sql = """
SELECT d.* FROM documents d
JOIN documents_fts fts ON d.rowid = fts.rowid
WHERE documents_fts MATCH ?
ORDER BY rank
LIMIT ?
"""
params = [sanitized, limit]
try:
cursor = self.conn.execute(sql, params)
return [Document.from_row(row) for row in cursor.fetchall()]
except sqlite3.OperationalError as e:
logger.warning("FTS query failed: %s (query: %s)", e, sanitized)
return []
def _sanitize_fts_query(self, query: str) -> str:
"""Sanitize FTS5 query to prevent injection and complexity issues."""
if len(query) > MAX_FTS_QUERY_LENGTH:
query = query[:MAX_FTS_QUERY_LENGTH]
# Escape special FTS5 characters
sanitized = FTS5_SPECIAL.sub(' ', query)
# Remove excess whitespace
sanitized = ' '.join(sanitized.split())
return sanitized.strip()
# --- Status / stats ---
def get_status_counts(self, source: Optional[str] = None) -> dict:
"""Get document counts by status."""
base = "SELECT status, parse_status, COUNT(*) as cnt FROM documents"
params: list = []
if source:
base += " WHERE source = ?"
params.append(source)
base += " GROUP BY status, parse_status"
cursor = self.conn.execute(base, params)
result = {
'collection': {},
'parsing': {},
'total': 0
}
for row in cursor.fetchall():
status = row['status']
parse_status = row['parse_status']
count = row['cnt']
result['collection'][status] = result['collection'].get(status, 0) + count
result['parsing'][parse_status] = result['parsing'].get(parse_status, 0) + count
result['total'] += count
return result
def get_extraction_counts(self, source: Optional[str] = None) -> dict:
"""Get extraction counts by extractor and status."""
if source:
cursor = self.conn.execute("""
SELECT es.extractor, es.status, COUNT(*) as cnt
FROM extraction_status es
JOIN documents d ON d.id = es.doc_id
WHERE d.source = ?
GROUP BY es.extractor, es.status
""", (source,))
else:
cursor = self.conn.execute("""
SELECT extractor, status, COUNT(*) as cnt
FROM extraction_status
GROUP BY extractor, status
""")
result = {}
for row in cursor.fetchall():
ext = row['extractor']
if ext not in result:
result[ext] = {}
result[ext][row['status']] = row['cnt']
return result
def get_all_tools(self) -> list[ExtractedTool]:
cursor = self.conn.execute("SELECT * FROM extracted_tools ORDER BY name")
return [ExtractedTool.from_row(row) for row in cursor.fetchall()]
def get_all_ttps(self) -> list[TTP]:
cursor = self.conn.execute("SELECT * FROM ttps ORDER BY category, technique")
return [TTP.from_row(row) for row in cursor.fetchall()]
def get_all_tradecraft(self) -> list[Tradecraft]:
cursor = self.conn.execute("SELECT * FROM tradecraft ORDER BY domain, method")
return [Tradecraft.from_row(row) for row in cursor.fetchall()]
def get_all_entities(self) -> list[Entity]:
cursor = self.conn.execute("SELECT * FROM entities ORDER BY type, name")
return [Entity.from_row(row) for row in cursor.fetchall()]
def get_all_infrastructure(self) -> list[Infrastructure]:
cursor = self.conn.execute("SELECT * FROM infrastructure ORDER BY indicator_type, indicator")
return [Infrastructure.from_row(row) for row in cursor.fetchall()]
# --- Verify integrity ---
def get_all_document_hashes(self) -> list[tuple[str, str, str]]:
"""Return (id, raw_path, content_hash) for all cached documents."""
cursor = self.conn.execute(
"SELECT id, raw_path, content_hash FROM documents WHERE status = 'CACHED' OR parse_status = 'PARSED'"
)
return [(row['id'], row['raw_path'], row['content_hash']) for row in cursor.fetchall()]