# tests/test_briefs_db.py import pytest @pytest.fixture(autouse=True) def _patch_db(tmp_path, monkeypatch): import storage.briefs_db as m monkeypatch.setattr(m, "DB_PATH", tmp_path / "briefs.db") m.init_db() def test_save_and_get_round_trip(): from storage.briefs_db import save_brief, get_brief brief = {"ticker": "AAPL", "company_name": "Apple Inc.", "what_changed": ["Revenue up 5%"]} save_brief("AAPL", brief) result = get_brief("AAPL") # get_brief injects a default "language" key; strip it for the equality check result_without_lang = {k: v for k, v in result.items() if k != "language"} assert result_without_lang == brief def test_get_brief_returns_none_for_unknown_ticker(): from storage.briefs_db import get_brief assert get_brief("NVDA") is None def test_save_brief_overwrites_on_same_ticker(): from storage.briefs_db import save_brief, get_brief save_brief("AAPL", {"v": 1}) save_brief("AAPL", {"v": 2}) result = get_brief("AAPL") assert result["v"] == 2 def test_save_brief_normalizes_ticker_to_uppercase(): from storage.briefs_db import save_brief, get_brief save_brief("aapl", {"company_name": "Apple"}) assert get_brief("AAPL")["company_name"] == "Apple" assert get_brief("aapl")["company_name"] == "Apple" def test_list_tickers_returns_all_saved(): from storage.briefs_db import save_brief, list_tickers save_brief("AAPL", {}) save_brief("MSFT", {}) save_brief("NVDA", {}) assert list_tickers() == ["AAPL", "MSFT", "NVDA"] def test_list_tickers_empty_db(): from storage.briefs_db import list_tickers assert list_tickers() == [] def test_init_db_is_idempotent(): from storage.briefs_db import init_db, save_brief, get_brief init_db() init_db() save_brief("AAPL", {"x": 1}) assert get_brief("AAPL")["x"] == 1 def test_migration_adds_language_column(): """init_db() migrates an old briefs table that lacks the language column.""" import sqlite3 import storage.briefs_db as m # Create old-schema row manually (no language column) with sqlite3.connect(m.DB_PATH) as conn: conn.execute("DROP TABLE IF EXISTS briefs") conn.execute(""" CREATE TABLE briefs ( ticker TEXT PRIMARY KEY, filing_date TEXT, brief_json TEXT NOT NULL, saved_at TEXT NOT NULL ) """) conn.execute("INSERT INTO briefs VALUES ('AAPL', '2025-01-01', '{\"ticker\": \"AAPL\"}', '2025-01-01T00:00:00+00:00')") m.init_db() with sqlite3.connect(m.DB_PATH) as conn: cols = {r[1] for r in conn.execute("PRAGMA table_info(briefs)")} assert "language" in cols row = conn.execute("SELECT language FROM briefs WHERE ticker = 'AAPL'").fetchone() assert row[0] == "English" def test_save_brief_persists_language(): from storage.briefs_db import save_brief, get_brief save_brief("AAPL", {"ticker": "AAPL", "language": "French", "what_matters_most": "Bonjour"}) result = get_brief("AAPL") assert result["language"] == "French" def test_save_and_get_translation_round_trip(): from storage.briefs_db import save_brief, save_translation, get_translation save_brief("AAPL", {"ticker": "AAPL", "language": "English"}) translated = {"ticker": "AAPL", "language": "French", "what_matters_most": "Bonjour"} save_translation("AAPL", "French", translated) result = get_translation("AAPL", "French") assert result == translated assert get_translation("AAPL", "German") is None def test_save_brief_invalidates_translations(): from storage.briefs_db import save_brief, save_translation, get_translation save_brief("AAPL", {"ticker": "AAPL", "language": "English"}) save_translation("AAPL", "French", {"ticker": "AAPL", "language": "French"}) assert get_translation("AAPL", "French") is not None # Re-save brief → should delete cached translations save_brief("AAPL", {"ticker": "AAPL", "language": "English", "what_matters_most": "Updated"}) assert get_translation("AAPL", "French") is None def test_delete_brief_removes_legacy_snapshot_and_translations(): from storage.briefs_db import ( delete_brief, get_brief, get_translation, list_briefs, save_brief, save_translation, ) save_brief("NVDA", {"ticker": "NVDA", "generated_at": "2026-07-16T13:48:29Z"}) save_translation("NVDA", "French", {"ticker": "NVDA", "language": "French"}) assert delete_brief("nvda") >= 2 assert get_brief("NVDA") is None assert list_briefs() == [] assert get_translation("NVDA", "French") is None assert delete_brief("AAPL") == 0 def test_snapshots_preserve_history_and_support_as_of_lookup(): from storage.briefs_db import save_brief, get_brief, list_snapshots save_brief("AAPL", { "ticker": "AAPL", "filing_date": "2025-01-30", "generated_at": "2025-01-31T10:00:00+00:00", "what_matters_most": "First", }) save_brief("AAPL", { "ticker": "AAPL", "filing_date": "2025-04-30", "generated_at": "2025-05-01T10:00:00+00:00", "what_matters_most": "Second", }) assert get_brief("AAPL")["what_matters_most"] == "Second" historical = get_brief("AAPL", as_of="2025-02-01T00:00:00+00:00") assert historical["what_matters_most"] == "First" assert len(list_snapshots("AAPL")) == 2 def test_snapshot_identity_is_immutable_on_duplicate_save(): from storage.briefs_db import save_brief, get_brief, list_snapshots base = { "ticker": "AAPL", "filing_date": "2025-01-30", "generated_at": "2025-01-31T10:00:00+00:00", } save_brief("AAPL", {**base, "what_matters_most": "Original"}) save_brief("AAPL", {**base, "what_matters_most": "Mutated retry"}) assert len(list_snapshots("AAPL")) == 1 assert get_brief("AAPL")["what_matters_most"] == "Original" def test_legacy_row_is_fail_closed_when_no_snapshot_exists(): import json import sqlite3 import storage.briefs_db as m with sqlite3.connect(m.DB_PATH) as conn: conn.execute( "INSERT OR REPLACE INTO briefs (ticker, filing_date, brief_json, language, saved_at) " "VALUES (?, ?, ?, ?, ?)", ( "OLD", "2024-01-01", json.dumps({ "ticker": "OLD", "status": "COMPLETE", "display_policy": {"event_returns_aligned": True}, "market_expectations": { "d1_price_reaction_pct": 99.0, "event_aligned": True, "event_comparison_allowed": True, }, }), "English", "2024-01-02", ), ) brief = m.get_brief("OLD") assert brief["data_quality_status"] == "LEGACY_UNVERIFIED" assert brief["status"] == "PARTIAL" assert brief["display_policy"]["event_returns_aligned"] is False assert brief["market_expectations"]["d1_price_reaction_pct"] is None assert brief["market_expectations"]["event_aligned"] is False # ── get_previous_brief ──────────────────────────────────────────────────────── def test_get_previous_brief_none_when_only_one_snapshot(): from storage.briefs_db import save_brief, get_previous_brief save_brief("NVDA", {"ticker": "NVDA", "generated_at": "2026-01-01T00:00:00Z"}) assert get_previous_brief("NVDA") is None def test_get_previous_brief_returns_second_most_recent(): from storage.briefs_db import save_brief, get_previous_brief save_brief("NVDA", {"ticker": "NVDA", "generated_at": "2026-01-01T00:00:00Z", "v": "old"}) save_brief("NVDA", {"ticker": "NVDA", "generated_at": "2026-02-01T00:00:00Z", "v": "new"}) prev = get_previous_brief("NVDA") assert prev["v"] == "old" def test_get_previous_brief_none_for_unknown_ticker(): from storage.briefs_db import get_previous_brief assert get_previous_brief("UNKNOWN") is None def test_get_previous_brief_respects_explicit_current_generated_at(): from storage.briefs_db import save_brief, get_previous_brief save_brief("NVDA", {"ticker": "NVDA", "generated_at": "2026-01-01T00:00:00Z", "v": "oldest"}) save_brief("NVDA", {"ticker": "NVDA", "generated_at": "2026-02-01T00:00:00Z", "v": "middle"}) save_brief("NVDA", {"ticker": "NVDA", "generated_at": "2026-03-01T00:00:00Z", "v": "newest"}) # Asking "what came before the middle snapshot" should skip the newest one. prev = get_previous_brief("NVDA", current_generated_at="2026-02-01T00:00:00Z") assert prev["v"] == "oldest"