File size: 8,829 Bytes
f97bc73
 
 
 
 
 
 
 
 
 
 
 
 
 
 
beec002
 
 
f97bc73
 
 
 
 
 
 
 
 
 
 
beec002
 
f97bc73
 
 
 
 
beec002
 
f97bc73
 
 
 
 
 
 
8b9291f
f97bc73
 
 
 
 
 
 
 
 
 
 
 
beec002
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7880373
 
b419b7b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7880373
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
# 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"