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"
|