FactsDB/factsdb/database.py

289 lines
11 KiB
Python

"""
Database module for FactsDB service
Handles SQLite database operations with proper locking
"""
import sqlite3
import threading
import os
from typing import List, Dict, Any, Optional
from contextlib import contextmanager
from .config import Config
class DatabaseManager:
"""Manages SQLite database operations with thread safety"""
def __init__(self, config):
self.config = config
self._lock = threading.Lock()
# Fix: Add error handling for database initialization
try:
self._init_database()
except Exception as e:
print(f"Warning: Database initialization failed: {e}")
print(f"Database path being used: {self.config.path}")
# Create the directory if it doesn't exist and the path has a directory component
import os
try:
# Check if path has a directory component
if os.path.dirname(self.config.path):
os.makedirs(os.path.dirname(self.config.path), exist_ok=True)
print(f"Created directory for database: {os.path.dirname(self.config.path)}")
else:
# If no directory, ensure current directory is writable
os.makedirs('.', exist_ok=True)
print("Verified current directory is writable")
# Try to test if we can create the file
test_path = os.path.abspath(self.config.path)
print(f"Testing database file access at: {test_path}")
# Try to create a simple test file to verify permissions
with open(test_path, 'w') as f:
f.write('')
os.remove(test_path)
print("Database file access test successful")
except Exception as dir_error:
print(f"Warning: Could not create database directory or test access: {dir_error}")
# Re-raise the original error
raise
def _init_database(self):
"""Initialize the database and create tables if they don't exist"""
with self._lock:
conn = sqlite3.connect(self.config.path, check_same_thread=False)
cursor = conn.cursor()
# Create facts table
cursor.execute('''
CREATE TABLE IF NOT EXISTS facts (
id INTEGER PRIMARY KEY AUTOINCREMENT,
table_name TEXT NOT NULL,
fact TEXT,
key_entities TEXT,
key_dates TEXT,
file_path TEXT,
extracted_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
processed BOOLEAN DEFAULT FALSE
)
''')
# Create tables metadata table
cursor.execute('''
CREATE TABLE IF NOT EXISTS tables (
name TEXT PRIMARY KEY,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
record_count INTEGER DEFAULT 0
)
''')
# Create files table to track processed files
cursor.execute('''
CREATE TABLE IF NOT EXISTS files (
id INTEGER PRIMARY KEY AUTOINCREMENT,
path TEXT UNIQUE NOT NULL,
table_name TEXT NOT NULL,
processed BOOLEAN DEFAULT FALSE,
processed_at TIMESTAMP,
error TEXT
)
''')
conn.commit()
conn.close()
@contextmanager
def get_connection(self):
"""Get a database connection with thread safety"""
with self._lock:
conn = sqlite3.connect(self.config.path, check_same_thread=False)
try:
yield conn
finally:
conn.close()
def create_table(self, table_name: str):
"""Create a new table entry"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute(
"INSERT OR REPLACE INTO tables (name, record_count) VALUES (?, 0)",
(table_name,)
)
conn.commit()
def get_table_names(self) -> List[str]:
"""Get all table names"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute("SELECT name FROM tables")
return [row[0] for row in cursor.fetchall()]
def get_table_info(self, table_name: str) -> Dict[str, Any]:
"""Get information about a specific table"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute("SELECT name, created_at, record_count FROM tables WHERE name = ?", (table_name,))
row = cursor.fetchone()
if row:
return {
'name': row[0],
'created_at': row[1],
'record_count': row[2]
}
return None
def update_table_count(self, table_name: str, count: int):
"""Update the record count for a table"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute(
"UPDATE tables SET record_count = ? WHERE name = ?",
(count, table_name)
)
conn.commit()
def insert_fact(self, table_name: str, fact_data: Dict[str, Any]) -> int:
"""Insert a fact into the database"""
with self.get_connection() as conn:
cursor = conn.cursor()
# Convert lists to JSON strings for storage
key_entities = str(fact_data.get('key_entities', []))
key_dates = str(fact_data.get('key_dates', []))
cursor.execute('''
INSERT INTO facts (
table_name, fact, key_entities, key_dates, file_path
) VALUES (?, ?, ?, ?, ?)
''', (
table_name,
fact_data.get('fact'),
key_entities,
key_dates,
fact_data.get('file_path')
))
fact_id = cursor.lastrowid
conn.commit()
# Update table record count
self._update_table_record_count(table_name)
return fact_id
def _update_table_record_count(self, table_name: str):
"""Update the record count for a table"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute("SELECT COUNT(*) FROM facts WHERE table_name = ?", (table_name,))
count = cursor.fetchone()[0]
cursor.execute("UPDATE tables SET record_count = ? WHERE name = ?", (count, table_name))
conn.commit()
def get_facts(self, table_name: str, limit: int = 100, offset: int = 0) -> List[Dict[str, Any]]:
"""Get facts from a specific table"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute('''
SELECT id, table_name, fact, key_entities, key_dates, file_path, extracted_at
FROM facts
WHERE table_name = ?
ORDER BY extracted_at DESC
LIMIT ? OFFSET ?
''', (table_name, limit, offset))
columns = [description[0] for description in cursor.description]
rows = cursor.fetchall()
facts = []
for row in rows:
fact = dict(zip(columns, row))
# Convert JSON strings back to lists
if fact['key_entities']:
fact['key_entities'] = eval(fact['key_entities'])
if fact['key_dates']:
fact['key_dates'] = eval(fact['key_dates'])
facts.append(fact)
return facts
def get_fact_by_id(self, fact_id: int) -> Optional[Dict[str, Any]]:
"""Get a specific fact by ID"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute('''
SELECT id, table_name, fact, key_entities, key_dates, file_path, extracted_at
FROM facts
WHERE id = ?
''', (fact_id,))
row = cursor.fetchone()
if row:
columns = [description[0] for description in cursor.description]
fact = dict(zip(columns, row))
# Convert JSON strings back to lists
if fact['key_entities']:
fact['key_entities'] = eval(fact['key_entities'])
if fact['key_dates']:
fact['key_dates'] = eval(fact['key_dates'])
return fact
return None
def get_all_tables(self) -> List[Dict[str, Any]]:
"""Get all tables with their information"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute("SELECT name, created_at, record_count FROM tables ORDER BY created_at DESC")
columns = [description[0] for description in cursor.description]
rows = cursor.fetchall()
return [dict(zip(columns, row)) for row in rows]
def get_table_count(self, table_name: str) -> int:
"""Get the count of facts in a specific table"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute("SELECT COUNT(*) FROM facts WHERE table_name = ?", (table_name,))
return cursor.fetchone()[0]
def mark_file_processed(self, file_path: str, table_name: str, success: bool = True, error: str = None):
"""Mark a file as processed"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute('''
INSERT OR REPLACE INTO files (path, table_name, processed, processed_at, error)
VALUES (?, ?, ?, datetime('now'), ?)
''', (file_path, table_name, success, error))
conn.commit()
def is_file_processed(self, file_path: str) -> bool:
"""Check if a file has been processed"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute("SELECT processed FROM files WHERE path = ?", (file_path,))
row = cursor.fetchone()
return row[0] if row else False
def get_processed_files(self, table_name: str) -> List[str]:
"""Get list of processed files for a table"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute("SELECT path FROM files WHERE table_name = ? AND processed = 1", (table_name,))
return [row[0] for row in cursor.fetchall()]
def get_unprocessed_files(self, table_name: str) -> List[str]:
"""Get list of unprocessed files for a table"""
with self.get_connection() as conn:
cursor = conn.cursor()
cursor.execute("SELECT path FROM files WHERE table_name = ? AND processed = 0", (table_name,))
return [row[0] for row in cursor.fetchall()]
# Global database manager instance
_db_manager = None
def get_db_manager(config) -> DatabaseManager:
"""Get the global database manager instance"""
global _db_manager
if _db_manager is None:
_db_manager = DatabaseManager(config)
return _db_manager