Files
xiaobaifupan/next/tests/test_market_data.py
T

605 lines
21 KiB
Python

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.data.sentiment import _phase, calculate_sentiment
from backend.database.connection import Database
from backend.database.migrations import MIGRATIONS, MigrationRunner
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_entity_detail_uses_selected_bar_previous_close_and_complete_stock_sections(
tmp_path,
) -> None:
market = gateway(tmp_path)
database = market._database
repository = MarketRepository()
with database.transaction() as connection:
connection.execute(
"UPDATE market_summaries SET payload_json = ? WHERE trade_date = ?",
(
json.dumps(
{
"overview": {},
"limits": [
{
"identifier": "000001.SZ", "code": "000001",
"reason": "银行板块走强", "streak": 1,
"first_time": "09:35", "last_time": "14:20", "open_times": 0,
}
],
}
),
"2026-07-29",
),
)
snapshot_id = connection.execute(
"""
INSERT INTO screener_factor_snapshots
(trade_date, version, observed_at, state, source_set_json,
coverage_json, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?)
"""
,
(
"2026-07-29",
"detail-v1",
"2026-07-29T15:05:00+08:00",
"final",
'["tushare"]',
"{}",
"2026-07-29T15:05:00+08:00",
),
).lastrowid
connection.execute(
"""
INSERT INTO screener_factor_values
(snapshot_id, identifier, code, name, sector, listed_days, is_st, payload_json)
VALUES (?, '000001.SZ', '000001', '平安银行', '银行', 1000, 0, ?)
""",
(
snapshot_id,
json.dumps(
{
"name": "平安银行", "sector": "银行", "turnover_rate": 2.5,
"return_5d": 3.2, "return_20d": 8.6, "total_mv_billion": 2100,
"circ_mv_billion": 1900, "net_flow_million": 12.5,
"large_flow_million": 8.1, "net_flow_5d_million": 35.2,
"flow_to_circ_mv_5d": 0.0185,
}
),
),
)
detail = MarketSnapshotService(database, repository, market).entity_detail(
"stock", "000001.SZ", "2026-07-29"
)
assert detail["trade_date"] == "2026-07-29"
assert detail["price"] == 11.1
assert detail["previous_close"] == 10.8
assert detail["change"] == pytest.approx(2.7778)
assert detail["entity"]["name"] == "平安银行"
assert detail["money_flow"]["available"] is True
assert detail["money_flow"]["net_million"] == 12.5
assert detail["event"] == {
"status": "涨停", "reason": "银行板块走强", "streak": 1,
"first_time": "09:35", "last_time": "14:20", "open_times": 0,
"seal_amount": None,
}
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 sentiment["stats"]["yesterday_count"] == 0
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_sentiment_phase_transitions_require_confirmed_recovery_and_fermentation() -> None:
phase, _reason = _phase(
{"phase": "冰点"}, 28, 3, 40, 45, 50, "修复", False
)
assert phase == "冰点"
phase, _reason = _phase(
{"phase": "修复", "fermentation_signal_count": 0},
55,
8,
60,
55,
60,
"发酵",
False,
)
assert phase == "修复"
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
def test_synced_snapshot_keeps_event_lists_separate_from_overview_counts(tmp_path) -> None:
database = Database(tmp_path / "sync-complete.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,
({
"ts_code": "000001.SZ", "symbol": "000001", "name": "平安银行",
"industry": "银行", "list_status": "L",
},),
"tushare",
"2026-07-29T15:00:00+08:00",
)
provider = TushareProvider("test-token")
stock = {
"ts_code": "000001.SZ", "name": "平安银行", "industry": "银行",
"close": 10, "pct_chg": 5, "amount": 100, "limit_times": 1,
}
provider.snapshot_inputs = lambda *_: {
"daily": calculation_result([stock]),
"limit_up": calculation_result([]),
"limit_down": calculation_result([]),
"broken": calculation_result([stock]),
"previous_limit_up": calculation_result([]),
"price_limits": calculation_result([{"ts_code": "000001.SZ", "up_limit": 11}]),
}
service = MarketSnapshotService(
database,
repository,
DataGateway(database, repository, (provider,), DataSourcePolicy()),
)
service.sync("2026-07-29", datetime(2026, 7, 30, 16, tzinfo=SHANGHAI))
broken = service.workspace("broken", "2026-07-29")
emotion = service.workspace("emotion", "2026-07-29")
assert len(broken["items"]) == 1
assert broken["items"][0]["identifier"] == "000001.SZ"
assert emotion["overview"]["broken"] == 1