from __future__ import annotations import json from datetime import datetime from zoneinfo import ZoneInfo import httpx import pytest from backend.bootstrap.application import create_application from backend.bootstrap.settings import Settings from backend.data.contracts import ( DataSource, DataUsage, ObservationMetadata, ProviderResult, SnapshotState, ) from backend.data.gateway import DataGateway from backend.data.policy import DataPolicyError, DataSourcePolicy from backend.data.providers.ifind import IfindProvider from backend.data.providers.tushare import TushareProvider from backend.data.repository import MarketRepository from backend.database.connection import Database from backend.database.migrations import MIGRATIONS, MigrationRunner from backend.features.market.sentiment import calculate_sentiment from backend.features.market.snapshot import build_snapshot from backend.features.market.sync import MarketSnapshotService, SnapshotSyncError from tests.support import run_scenario SHANGHAI = ZoneInfo("Asia/Shanghai") class FakeProvider: source = DataSource.IFIND configured = True def calendar(self, start_date: str, end_date: str) -> ProviderResult: raise AssertionError("not used") def entities(self) -> ProviderResult: raise AssertionError("not used") def daily(self, entity_type: str, identifier: str, end_date: str) -> ProviderResult: return ProviderResult( ( { "time": "2026-07-28 15:00:00", "open": 10, "high": 11, "low": 9.8, "close": 10.8, "volume": 1000, "amount": 10800, }, { "time": "2026-07-29 15:00:00", "open": 11, "high": 11.2, "low": 10.7, "close": 11.1, "volume": 1200, "amount": 13320, }, { "time": "2026-07-30 09:15:00", "open": 0, "high": 0, "low": 0, "close": 11.1, "volume": 0, "amount": 0, }, ), metadata(DataSource.IFIND, SnapshotState.ARCHIVE), ) def minute(self, entity_type: str, identifier: str, trade_date: str) -> ProviderResult: rows = ( { "time": f"{trade_date} 09:30:00", "open": 11, "high": 11.1, "low": 10.9, "close": 11.05, "volume": 100, "amount": 1105, "avgPrice": 11.03, "preClose": 11, }, ) return ProviderResult(rows, metadata(DataSource.IFIND, SnapshotState.REALTIME)) def metadata(source: DataSource, state: SnapshotState) -> ObservationMetadata: return ObservationMetadata( source=source, observed_at=datetime(2026, 7, 30, 9, 15, tzinfo=SHANGHAI), unit="yuan/share", adjustment="forward", freshness_seconds=0, coverage=1, state=state, usage=DataUsage.DISPLAY, ) def gateway(tmp_path) -> DataGateway: database = Database(tmp_path / "market.db") MigrationRunner(database).upgrade(MIGRATIONS) repository = MarketRepository() with database.transaction() as connection: repository.replace_stocks( connection, ( { "ts_code": "000001.SZ", "symbol": "000001", "name": "平安银行", "industry": "银行", "list_status": "L", }, ), "tushare", "2026-07-29T15:00:00+08:00", ) connection.execute( """ INSERT INTO market_summaries (trade_date, observed_at, state, source, coverage, payload_json, created_at) VALUES (?, ?, 'final', 'tushare', 1, ?, ?) """, ( "2026-07-29", "2026-07-29T15:00:00+08:00", json.dumps({"limit_up": 46, "limit_down": 3}), "2026-07-29T15:05:00+08:00", ), ) return DataGateway(database, repository, (FakeProvider(),), DataSourcePolicy()) def test_public_sources_cannot_enter_calculations() -> None: policy = DataSourcePolicy() with pytest.raises(DataPolicyError): policy.assert_allowed(DataSource.EASTMONEY, DataUsage.CALCULATION) policy.assert_allowed(DataSource.EASTMONEY, DataUsage.DISPLAY) def test_ifind_top_level_tables_and_expired_access_token_are_handled() -> None: provider = IfindProvider("refresh-token", "expired-token") calls: list[tuple[str, str]] = [] def post(endpoint, body, access, refresh=""): calls.append((endpoint, access or refresh)) if endpoint == "get_access_token": return {"errorcode": 0, "data": {"access_token": "fresh-token"}} if access == "expired-token": return {"errorcode": -1302, "errmsg": "token expired"} return { "errorcode": 0, "tables": [ { "thscode": ["000001.SZ"], "time": ["2026-07-29 15:00:00"], "table": { "open": [10], "high": [11], "low": [9], "close": [10.5], "volume": [100], "amount": [1050], }, } ], } provider._post = post result = provider.daily("stock", "000001.SZ", "2026-07-29") assert result.rows[0]["time"] == "2026-07-29 15:00:00" assert result.rows[0]["thscode"] == "000001.SZ" assert calls == [ ("cmd_history_quotation", "expired-token"), ("get_access_token", "refresh-token"), ("cmd_history_quotation", "fresh-token"), ] def test_trade_context_keeps_real_snapshot_date(tmp_path) -> None: market = gateway(tmp_path) context = market.trade_context( "2026-07-30", datetime(2026, 7, 30, 9, 10, tzinfo=SHANGHAI) ) assert context.requested_date == "2026-07-30" assert context.actual_date == "2026-07-29" assert context.carried_forward is True assert context.observed_at.isoformat() == "2026-07-29T15:00:00+08:00" def test_latest_daily_chart_drops_empty_premarket_bar(tmp_path) -> None: series = gateway(tmp_path).chart( "stock", "000001.SZ", "day", datetime(2026, 7, 30, 9, 15, tzinfo=SHANGHAI) ) assert series.trade_date == "2026-07-29" assert [point.time for point in series.points] == ["2026-07-28", "2026-07-29"] assert series.points[-1].amount == 13320 def test_minute_chart_contract_has_real_session_bounds_and_hides_source(tmp_path) -> None: application = create_application(Settings.for_test(tmp_path)) async def scenario(client: httpx.AsyncClient) -> None: registered = await client.post( "/api/auth/register", json={"username": "market-user", "password": "Market-pass-123!"}, ) assert registered.status_code == 201 fake_market = type( "FakeMarketService", (), { "chart": lambda self, *_: { "entity_type": "stock", "identifier": "000001.SZ", "code": "000001", "name": "平安银行", "interval": "minute", "trade_date": "2026-07-29", "observed_at": "2026-07-29T15:00:00+08:00", "previous_close": 10.9, "range_start": "09:30", "range_end": "15:00", "points": [], } }, )() object.__setattr__(application.state.container, "market", fake_market) response = await client.get("/api/market/entities/stock/000001.SZ/charts/minute") assert response.status_code == 200 assert response.json()["range_start"] == "09:30" assert response.json()["range_end"] == "15:00" assert "source" not in response.text run_scenario(application, scenario) def test_search_groups_are_fixed_and_require_authentication(tmp_path) -> None: application = create_application(Settings.for_test(tmp_path)) async def scenario(client: httpx.AsyncClient) -> None: assert (await client.get("/api/market/search?q=上证")).status_code == 401 await client.post( "/api/auth/register", json={"username": "search-user", "password": "Search-pass-123!"}, ) response = await client.get("/api/market/search?q=上证") assert response.status_code == 200 groups = response.json()["groups"] assert [group["label"] for group in groups] == ["股票", "板块", "题材", "指数"] assert groups[-1]["items"][0]["name"] == "上证指数" run_scenario(application, scenario) def calculation_result(rows) -> ProviderResult: return ProviderResult( tuple(rows), ObservationMetadata( source=DataSource.TUSHARE, observed_at=datetime(2026, 7, 29, 15, tzinfo=SHANGHAI), unit="mixed", adjustment="not_applicable", freshness_seconds=0, coverage=1, state=SnapshotState.FINAL, usage=DataUsage.CALCULATION, ), ) def test_market_snapshot_units_and_yesterday_outcomes_are_deterministic() -> None: daily = [ { "ts_code": f"00000{index}.SZ", "close": 10 + index, "pct_chg": change, "amount": 100, } for index, change in enumerate((10, 4, -10, 2, -2), start=1) ] def event(index, streak=1): return { "ts_code": f"00000{index}.SZ", "name": f"样本{index}", "industry": "测试行业", "close": 10 + index, "pct_chg": daily[index - 1]["pct_chg"], "amount": 100, "limit_times": streak, } snapshot = build_snapshot( "2026-07-29", "2026-07-28", { "daily": calculation_result(daily), "limit_up": calculation_result([event(1, 2)]), "broken": calculation_result([event(2)]), "limit_down": calculation_result([event(3)]), "previous_limit_up": calculation_result([event(index) for index in range(1, 6)]), "price_limits": calculation_result( [{"ts_code": "000002.SZ", "up_limit": 15, "down_limit": 9}] ), }, ) assert snapshot["overview"]["amount"] == 500_000 assert snapshot["broken"][0]["distance_to_limit"] == 20 assert snapshot["ladders"][0]["level"] == 2 assert snapshot["ladders"][0]["stocks"][0]["code"] == "000001" assert snapshot["sector_rotation"][0]["name"] == "测试行业" assert snapshot["sector_rotation"][0]["delta"] == -4 assert [row["outcome"] for row in snapshot["yesterday_limits"]] == [ "晋级", "炸板", "跌停", "红盘", "断板", ] performance = snapshot["limit_performance"][0] assert performance == { "level": 1, "count": 5, "advanced": 1, "red": 1, "broken": 1, "opened": 1, "limit_down": 1, "advance_rate": 20.0, "positive_rate": 60.0, "average_change": 0.8, } def test_sector_members_are_normalized_and_persistently_cached(tmp_path) -> None: database = Database(tmp_path / "members.db") MigrationRunner(database).upgrade(MIGRATIONS) repository = MarketRepository() provider = TushareProvider("test-token") calls = 0 def members(*_): nonlocal calls calls += 1 return calculation_result( [ { "sector_code": "801082.SI", "sector_name": "元件", "ts_code": "002141.SZ", "name": "贤丰控股", "change": 3.2, "open": 4.1, "close": 4.3, "amount": 123_000_000, "quoted": True, } ] ) provider.sector_members = members market = DataGateway(database, repository, (provider,), DataSourcePolicy()) first = market.sector_members("2026-07-29", "元件", "002141.SZ") second = market.sector_members("2026-07-29", "元件", "002141.SZ") assert calls == 1 assert first == second assert first["items"][0]["code"] == "002141" assert first["quoted_count"] == 1 def test_sentiment_has_all_weighted_components_and_extreme_risk_cap() -> None: snapshot = { "overview": { "up_count": 10, "down_count": 90, "limit_up": 10, "limit_down": 100, "broken": 20, "seal_rate": 33.3, "amount": 100_000_000_000, }, "limits": [{"streak": 1, "amount": 100_000_000} for _ in range(10)], "yesterday_limits": [], } sentiment = calculate_sentiment(snapshot, []) assert sentiment["score"] <= 15 assert sentiment["phase"] == "冰点" assert {item["key"]: item["weight"] for item in sentiment["components"]} == { "breadth": 20, "limit_ecology": 25, "profit_effect": 30, "ladder_structure": 15, "liquidity": 10, } def test_incomplete_daily_snapshot_is_rejected_without_overwriting(tmp_path) -> None: database = Database(tmp_path / "sync.db") MigrationRunner(database).upgrade(MIGRATIONS) repository = MarketRepository() with database.transaction() as connection: repository.replace_calendar( connection, ( {"cal_date": "20260728", "is_open": 1, "pretrade_date": "20260727"}, {"cal_date": "20260729", "is_open": 1, "pretrade_date": "20260728"}, ), "tushare", "2026-07-29T15:00:00+08:00", ) repository.replace_stocks( connection, tuple( { "ts_code": f"{index:06d}.SZ", "symbol": f"{index:06d}", "name": f"样本{index}", "industry": "测试", "list_status": "L", } for index in range(1, 101) ), "tushare", "2026-07-29T15:00:00+08:00", ) provider = TushareProvider("test-token") provider.snapshot_inputs = lambda *_: { "daily": calculation_result([{"ts_code": "000001.SZ"}]), "limit_up": calculation_result([]), "limit_down": calculation_result([]), "broken": calculation_result([]), "previous_limit_up": calculation_result([]), "price_limits": calculation_result([]), } service = MarketSnapshotService( database, repository, DataGateway(database, repository, (provider,), DataSourcePolicy()), ) with pytest.raises(SnapshotSyncError, match="覆盖率"): service.sync("2026-07-29", datetime(2026, 7, 30, 16, tzinfo=SHANGHAI)) with database.read() as connection: assert repository.latest_summary(connection, "2026-07-29") is None