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