amplegest / tests /test_briefs_db.py
Viney's picture
fix: keep fail-closed per fact but stop zero-verified synthesis from emptying the whole brief
b419b7b
Raw
History Blame Contribute Delete
8.83 kB
# 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"