main
py 197 lines 10.6 KB
Raw
1 from datetime import date, datetime, timedelta, timezone
2 from decimal import Decimal
3 import sqlite3
4 from uuid import UUID
5
6 import pytest
7 from pydantic import ValidationError
8
9 from app.models import DailyMarketBar, MarketPriceObservation
10 from app.persistence import SqliteResearchPersistence, DisabledResearchPersistence, persistence_from_settings
11 from app.postgres_persistence import PostgresResearchPersistence, _PostgresConnectionAdapter
12 from app.repository import ResearchRepository
13 from app.settings import Settings
14
15
16 KEY = UUID(int=1)
17 DAY = date(2026, 9, 1)
18 NOW = datetime(2026, 9, 2, tzinfo=timezone.utc)
19
20
21 def bar(**updates):
22 return DailyMarketBar.model_validate(dict(global_instrument_id=KEY, trading_date=DAY,
23 open="100.123456789012", high="105.123456789012", low="99.123456789012", close="103.123456789012",
24 previous_close=None, volume=None, turnover=None, currency="INR", provider="PUBLIC_SOURCE",
25 provider_symbol="ABC", source_mode="REAL", source_url="https://example.test/bars", retrieved_at=NOW) | updates)
26
27
28 def test_fresh_sqlite_bootstrap_schema_indexes_and_repeat_startup(tmp_path):
29 path = tmp_path / "fresh.sqlite"
30 store = persistence_from_settings(Settings(research_persistence_enabled=True, research_database_backend="sqlite", research_database_name=str(path)))
31 columns = store._connection.execute("PRAGMA table_info(global_daily_market_bars)").fetchall()
32 assert {r['name'] for r in columns} == {"global_instrument_id", "trading_date", "open_price", "high_price", "low_price",
33 "close_price", "previous_close", "volume", "turnover", "currency", "provider", "provider_symbol", "source_mode", "source_url", "retrieved_at"}
34 assert [r['name'] for r in sorted(columns, key=lambda r:r['pk']) if r['pk']] == ["global_instrument_id", "trading_date", "provider"]
35 indexes = store._connection.execute("PRAGMA index_list(global_daily_market_bars)").fetchall()
36 assert any(r['unique'] for r in indexes)
37 assert any(r['name'] == 'idx_daily_market_bars_date_instrument' for r in indexes)
38 store.upsert_daily_market_bar(bar())
39 store.migrate()
40 assert SqliteResearchPersistence(path).load_daily_market_bars({KEY}) == [bar()]
41
42
43 def test_existing_sqlite_database_receives_table_without_manual_ddl(tmp_path, monkeypatch):
44 import app.persistence as persistence
45 original = persistence._sqlite_schema
46 schema = original()
47 # Bootstrap an earlier SQLite schema using its normal mechanism, then reopen
48 # with the current bootstrap function. Existing close-only rows survive.
49 previous = schema[schema.index(' CREATE TABLE IF NOT EXISTS research_acquisition_observations'):]
50 path = tmp_path / 'upgrade.sqlite'
51 monkeypatch.setattr(persistence, '_sqlite_schema', lambda: previous)
52 old = SqliteResearchPersistence(path)
53 price = MarketPriceObservation(instrument_id=KEY, observed_at=NOW, retrieved_at=NOW, price=Decimal(100),
54 currency='INR', provider='EXISTING', source_url='https://example.test')
55 old.upsert_market_price_observation(price)
56 old._connection.close()
57 monkeypatch.setattr(persistence, '_sqlite_schema', original)
58 current = SqliteResearchPersistence(path)
59 assert current.load_daily_market_bars({KEY}) == []
60 assert current.load_market_price_observations({KEY}) == [price]
61
62
63 def test_insert_round_trip_preserves_decimals_date_utc_and_nulls():
64 store = SqliteResearchPersistence()
65 original = bar(retrieved_at=NOW.astimezone(timezone(timedelta(hours=5, minutes=30))))
66 store.upsert_daily_market_bar(original)
67 loaded = store.load_daily_market_bars({KEY})[0]
68 assert loaded == original
69 assert loaded.open == Decimal('100.123456789012')
70 assert type(loaded.trading_date) is date and loaded.retrieved_at.utcoffset() == timedelta(0)
71 assert loaded.volume is loaded.turnover is loaded.previous_close is None
72 assert loaded.model_dump(by_alias=True)['globalInstrumentId'] == KEY
73 assert store.load_market_price_observations({KEY}) == []
74
75
76 def test_batch_idempotence_correction_and_provider_symbol_not_identity():
77 store = SqliteResearchPersistence()
78 first, second = bar(), bar(provider='SECOND_SOURCE')
79 assert store.upsert_daily_market_bars([first, second]) == 2
80 assert store.upsert_daily_market_bars([first, second]) == 2
81 corrected = bar(open=101, high=111, low=98, close=108, previous_close=102, volume=0,
82 turnover=Decimal('1234567890123.123456789012'), provider_symbol='RENAMED', source_url='https://example.test/corrected',
83 source_mode='DEMO', retrieved_at=NOW+timedelta(days=1))
84 store.upsert_daily_market_bar(corrected)
85 assert store.load_daily_market_bars({KEY}) == [corrected, second]
86 assert store.load_daily_market_bars({KEY})[0].volume == 0
87 # A partial correction clears missing fields rather than preserving stale values.
88 store.upsert_daily_market_bar(bar(provider_symbol=None))
89 loaded = store.load_daily_market_bars({KEY})[0]
90 assert loaded.volume is loaded.turnover is loaded.previous_close is loaded.provider_symbol is None
91
92
93 def test_large_batches_filter_order_and_no_n_plus_one():
94 store = SqliteResearchPersistence()
95 rows = [bar(global_instrument_id=UUID(int=i), provider=provider, trading_date=DAY+timedelta(days=day))
96 for i in range(1, 504) for provider in ('B', 'A') for day in (1, 0)]
97 queries = []
98 store._connection.set_trace_callback(queries.append)
99 store.upsert_daily_market_bars(list(reversed(rows)))
100 assert sum(q.startswith('INSERT INTO global_daily_market_bars') for q in queries) == (len(rows)+49)//50
101 queries.clear()
102 ids = {UUID(int=i) for i in range(1, 502)}
103 loaded = store.load_daily_market_bars(ids, start_date=DAY, end_date=DAY, provider='B')
104 assert [b.global_instrument_id.int for b in loaded] == list(range(1, 502))
105 assert all(b.trading_date == DAY and b.provider == 'B' for b in loaded)
106 assert len(queries) == 2 and all('global_instrument_id IN (' in q for q in queries)
107 all_rows = store.load_daily_market_bars({KEY, UUID(int=2)})
108 assert [(r.global_instrument_id.int,r.trading_date,r.provider) for r in all_rows] == sorted((r.global_instrument_id.int,r.trading_date,r.provider) for r in all_rows)
109
110
111 def test_empty_ids_and_empty_batch_issue_no_queries():
112 store = SqliteResearchPersistence()
113 store.upsert_daily_market_bar(bar())
114 queries = []
115 store._connection.set_trace_callback(queries.append)
116 assert store.load_daily_market_bars(set()) == []
117 assert store.upsert_daily_market_bars([]) == 0
118 assert queries == []
119 with pytest.raises(ValueError): store.load_daily_market_bars(None)
120 with pytest.raises(ValueError): store.load_daily_market_bars({KEY}, start_date=DAY+timedelta(days=1), end_date=DAY)
121 with pytest.raises(ValueError): store.load_daily_market_bars({KEY}, provider=' ')
122 with pytest.raises(ValueError): store.load_daily_market_bars({KEY}, start_date=NOW)
123
124
125 @pytest.mark.parametrize('updates', [
126 {'volume': -1}, {'volume': 1.5}, {'volume': True}, {'volume': 9223372036854775808},
127 {'high': 90, 'low': 100}, {'open': 0}, {'close': -1}, {'previous_close': 0}, {'turnover': -1},
128 {'currency': ' '}, {'provider': ' '}, {'trading_date': NOW}, {'trading_date': '2026-09-01T00:00:00Z'},
129 {'retrieved_at': NOW.replace(tzinfo=None)}, {'close': 'NaN'}, {'turnover': 'Infinity'},
130 {'close': '1.1234567890123'}, {'source_mode': 'UNKNOWN'}, {'global_instrument_id': None},
131 ])
132 def test_invalid_domain_evidence_rejected(updates):
133 with pytest.raises(ValidationError): bar(**updates)
134
135
136 def test_partial_rows_zero_turnover_and_bigint_volume_are_valid():
137 store = SqliteResearchPersistence()
138 value = bar(open=None, high=None, low=None, close=None, turnover=0, volume=9223372036854775807)
139 store.upsert_daily_market_bar(value)
140 assert store.load_daily_market_bars({KEY}) == [value]
141
142
143 @pytest.mark.parametrize('column,value', [('volume', -1), ('high_price', '1'), ('close_price', '0'),
144 ('currency',' '), ('provider',' '), ('turnover','-1')])
145 def test_database_constraints_reject_invalid_updates(column, value):
146 store = SqliteResearchPersistence()
147 store.upsert_daily_market_bar(bar())
148 with pytest.raises(sqlite3.IntegrityError):
149 store._connection.execute(f'UPDATE global_daily_market_bars SET {column}=?', (value,))
150 store._connection.rollback()
151 assert store.load_daily_market_bars({KEY}) == [bar()]
152
153
154 def test_batch_validates_all_inputs_before_writing_and_last_correction_wins():
155 store = SqliteResearchPersistence()
156 invalid = bar().model_copy(update={'volume': -1})
157 with pytest.raises(ValidationError): store.upsert_daily_market_bars([bar(), invalid])
158 assert store.load_daily_market_bars({KEY}) == []
159 corrected = bar(close=104)
160 assert store.upsert_daily_market_bars([bar(), corrected]) == 1
161 assert store.load_daily_market_bars({KEY}) == [corrected]
162
163
164 def test_postgres_uses_same_parameterized_batch_boundary():
165 queries = []
166 class Connection:
167 def execute(self, sql, params=None): queries.append((sql, params)); return self
168 def fetchall(self): return []
169 def commit(self): pass
170 def rollback(self): pass
171 store = PostgresResearchPersistence.__new__(PostgresResearchPersistence)
172 store._connection = _PostgresConnectionAdapter(Connection())
173 store.upsert_daily_market_bars([bar(global_instrument_id=UUID(int=i)) for i in range(1, 102)])
174 assert [len(params) for _,params in queries] == [750, 750, 15]
175 assert all('%s' in sql and '?' not in sql for sql,_ in queries)
176 queries.clear()
177 assert store.load_daily_market_bars({KEY}, start_date=DAY, end_date=DAY, provider="a' OR 1=1 --") == []
178 assert "a' OR 1=1 --" not in queries[0][0]
179 assert queries[0][1] == [str(KEY), DAY.isoformat(), DAY.isoformat(), "a' OR 1=1 --"]
180
181
182 @pytest.mark.asyncio
183 async def test_repository_boundary_and_disabled_mode_are_provider_free(monkeypatch):
184 import socket
185 monkeypatch.setattr(socket, 'create_connection', lambda *a,**k: pytest.fail('provider/network invoked'))
186 store = SqliteResearchPersistence()
187 repo = ResearchRepository.__new__(ResearchRepository)
188 repo._persistence = store
189 await repo.upsert_daily_market_bar_async(bar())
190 assert await repo.upsert_daily_market_bars_async([bar(provider='SECOND')]) == 1
191 assert await repo.daily_market_bars_for_instruments({KEY,UUID(int=2)}, start_date=DAY, end_date=DAY, provider='SECOND') == {
192 KEY: [bar(provider='SECOND')], UUID(int=2): []}
193 assert repo.daily_market_bars_for({KEY}) == {KEY: store.load_daily_market_bars({KEY})}
194 disabled = DisabledResearchPersistence()
195 assert disabled.load_daily_market_bars({KEY}) == []
196 assert disabled.upsert_daily_market_bars([bar()]) == 0
197 assert disabled.upsert_daily_market_bar(bar()) is None