"""Tests for database module.""" import json import os import tempfile import pytest from factsdb.database import DatabaseManager, get_db_manager from factsdb.config import DatabaseConfig @pytest.fixture def db_manager(tmp_path): """Create a DatabaseManager with a temp database.""" db_path = str(tmp_path / "test.db") config = DatabaseConfig(path=db_path) return DatabaseManager(config) class TestDatabaseManagerInit: def test_creates_tables(self, db_manager): tables = db_manager.get_table_names() assert tables == [] def test_creates_db_file(self, tmp_path): db_path = str(tmp_path / "init.db") config = DatabaseConfig(path=db_path) DatabaseManager(config) assert os.path.exists(db_path) def test_init_in_subdir(self, tmp_path): subdir = tmp_path / "sub" subdir.mkdir() db_path = str(subdir / "test.db") config = DatabaseConfig(path=db_path) mgr = DatabaseManager(config) assert mgr is not None class TestConnection: def test_get_connection(self, db_manager): with db_manager.get_connection() as conn: cur = conn.execute("SELECT 1") assert cur.fetchone()[0] == 1 class CreateTableTests: def test_create_table(self, db_manager): db_manager.create_table("test_table") names = db_manager.get_table_names() assert "test_table" in names def test_create_table_idempotent(self, db_manager): db_manager.create_table("t1") db_manager.create_table("t1") names = db_manager.get_table_names() assert names.count("t1") == 1 class TestGetTableInfo: def test_existing_table(self, db_manager): db_manager.create_table("info_test") info = db_manager.get_table_info("info_test") assert info is not None assert info["name"] == "info_test" assert info["record_count"] == 0 def test_nonexistent_table(self, db_manager): info = db_manager.get_table_info("no_such") assert info is None class TestUpdateTableCount: def test_update_count(self, db_manager): db_manager.create_table("cnt") db_manager.update_table_count("cnt", 42) info = db_manager.get_table_info("cnt") assert info["record_count"] == 42 class TestInsertFact: def test_insert_returns_id(self, db_manager): db_manager.create_table("ft") fid = db_manager.insert_fact("ft", {"fact": "F1", "key_entities": [], "key_dates": []}) assert fid > 0 def test_insert_stores_data(self, db_manager): db_manager.create_table("ft") db_manager.insert_fact("ft", {"fact": "F1", "key_entities": ["A"], "key_dates": ["2024-01-01"], "file_path": "/x"}) facts = db_manager.get_facts("ft") assert len(facts) == 1 assert facts[0]["fact"] == "F1" assert facts[0]["key_entities"] == ["A"] assert facts[0]["key_dates"] == ["2024-01-01"] def test_insert_updates_count(self, db_manager): db_manager.create_table("ft") db_manager.insert_fact("ft", {"fact": "F1"}) count = db_manager.get_table_count("ft") assert count == 1 class TestGetFacts: def test_empty(self, db_manager): db_manager.create_table("ft") assert db_manager.get_facts("ft") == [] def test_limit(self, db_manager): db_manager.create_table("ft") for i in range(5): db_manager.insert_fact("ft", {"fact": f"F{i}"}) facts = db_manager.get_facts("ft", limit=2) assert len(facts) == 2 def test_offset(self, db_manager): db_manager.create_table("ft") for i in range(5): db_manager.insert_fact("ft", {"fact": f"F{i}"}) facts = db_manager.get_facts("ft", limit=10, offset=3) assert len(facts) == 2 def test_json_decode_error(self, db_manager): db_manager.create_table("ft") db_manager.insert_fact("ft", {"fact": "F1", "key_entities": ["e"]}) with db_manager.get_connection() as conn: conn.execute("UPDATE facts SET key_entities = 'bad' WHERE id = 1") conn.commit() facts = db_manager.get_facts("ft") assert facts[0]["key_entities"] == [] class TestGetFactById: def test_existing(self, db_manager): db_manager.create_table("ft") fid = db_manager.insert_fact("ft", {"fact": "F1"}) fact = db_manager.get_fact_by_id(fid) assert fact is not None assert fact["fact"] == "F1" def test_nonexistent(self, db_manager): assert db_manager.get_fact_by_id(999) is None class TestGetAllTables: def test_multiple_tables(self, db_manager): db_manager.create_table("a") db_manager.create_table("b") tables = db_manager.get_all_tables() names = [t["name"] for t in tables] assert "a" in names assert "b" in names class TestFileTracking: def test_mark_processed(self, db_manager): db_manager.mark_file_processed("/path/file.txt", "ft", True) assert db_manager.is_file_processed("/path/file.txt") def test_mark_not_processed(self, db_manager): db_manager.mark_file_processed("/path/file.txt", "ft", False, "err") assert not db_manager.is_file_processed("/path/file.txt") def test_unprocessed_file(self, db_manager): assert db_manager.is_file_processed("/no/file") is False def test_get_processed_files(self, db_manager): db_manager.mark_file_processed("/a", "ft", True) db_manager.mark_file_processed("/b", "ft", False) processed = db_manager.get_processed_files("ft") assert "/a" in processed assert "/b" not in processed def test_get_unprocessed_files(self, db_manager): db_manager.mark_file_processed("/a", "ft", True) db_manager.mark_file_processed("/b", "ft", False) unprocessed = db_manager.get_unprocessed_files("ft") assert "/b" in unprocessed assert "/a" not in unprocessed class TestQueryAndAliases: def test_query_table(self, db_manager): db_manager.create_table("ft") db_manager.insert_fact("ft", {"fact": "F1"}) results = db_manager.query_table("ft") assert len(results) == 1 def test_get_table_record_count(self, db_manager): db_manager.create_table("ft") db_manager.insert_fact("ft", {"fact": "F1"}) assert db_manager.get_table_record_count("ft") == 1 def test_store_fact(self, db_manager): db_manager.create_table("ft") fid = db_manager.store_fact({"fact": "F1"}, "ft", "/x") assert fid > 0 class TestDatabaseStats: def test_stats(self, db_manager): db_manager.create_table("ft") db_manager.insert_fact("ft", {"fact": "F1"}) db_manager.mark_file_processed("/a", "ft", True) stats = db_manager.get_database_stats() assert stats["total_facts"] == 1 assert stats["processed_files"] == 1 assert "ft" in stats["table_counts"] class TestGetDbManager: def test_singleton(self, tmp_path, monkeypatch): import factsdb.database as db_module db_module._db_manager = None db_path = str(tmp_path / "sg.db") config = DatabaseConfig(path=db_path) m1 = get_db_manager(config) m2 = get_db_manager(config) assert m1 is m2 db_module._db_manager = None