Files
xiaobai-review/tests/test_market_insights.py

295 lines
15 KiB
Python

from __future__ import annotations
import tempfile
import unittest
from datetime import datetime, timedelta, timezone
from pathlib import Path
from database import ReviewDatabase
from backend.data.providers.tushare_client import TushareError
from backend.features.market.insights import MarketInsightsService
from backend.features.screener.engine import FACTOR_FIELDS, ScreenerEngine
class FakeMarketClient:
def resolve_trade_context(self, requested: str):
value = str(requested).replace("-", "")
return value, "20260723"
def query(self, api_name, params=None, fields=""):
params = params or {}
date = params.get("trade_date", "")
if api_name == "trade_cal":
return [
{"cal_date": f"202607{day:02d}", "is_open": 1}
for day in range(14, 24)
]
if api_name == "stock_basic":
return [
{"ts_code": "000001.SZ", "name": "平安银行", "industry": "银行", "market": "主板", "list_date": "19910403"},
{"ts_code": "000002.SZ", "name": "万科A", "industry": "房地产", "market": "主板", "list_date": "19910129"},
]
if api_name == "stk_auction":
if date == "20260724":
return []
return [
{"ts_code": "000001.SZ", "trade_date": date, "price": 10.5, "pre_close": 10, "vol": 20000, "amount": 5_000_000, "turnover_rate": 0.12, "volume_ratio": 1.8},
{"ts_code": "000002.SZ", "trade_date": date, "price": 9.8, "pre_close": 10, "vol": 10000, "amount": 2_000_000, "turnover_rate": 0.05, "volume_ratio": 0.8},
]
if api_name == "stk_limit":
return [
{"ts_code": "000001.SZ", "trade_date": date, "up_limit": 11, "down_limit": 9},
{"ts_code": "000002.SZ", "trade_date": date, "up_limit": 11, "down_limit": 9},
]
if api_name == "ths_index":
return [{"ts_code": "885001.TI", "name": "人工智能", "count": 2, "exchange": "A", "list_date": "20200101", "type": "N"}]
if api_name == "ths_daily":
if params.get("ts_code"):
return [
{"ts_code": "885001.TI", "trade_date": "20260722", "open": 99, "high": 102, "low": 98, "close": 101, "pct_change": 1, "vol": 100},
{"ts_code": "885001.TI", "trade_date": "20260723", "open": 101, "high": 104, "low": 100, "close": 103, "pct_change": 1.98, "vol": 120},
]
if date == "20260724":
return []
return [{"ts_code": "885001.TI", "trade_date": date, "close": 103, "pct_change": 1.98, "vol": 120, "turnover_rate": 2.3}]
if api_name == "ths_member":
return [
{"ts_code": "885001.TI", "con_code": "000001.SZ", "con_name": "平安银行"},
{"ts_code": "885001.TI", "con_code": "000002.SZ", "con_name": "万科A"},
]
if api_name == "daily":
return [
{"ts_code": "000001.SZ", "trade_date": date, "open": 10, "high": 11, "low": 9.8, "close": 10.5, "pct_chg": 5, "vol": 100, "amount": 200000},
{"ts_code": "000002.SZ", "trade_date": date, "open": 10, "high": 10, "low": 9.7, "close": 9.8, "pct_chg": -2, "vol": 100, "amount": 100000},
]
if api_name == "ths_hot":
if date == "20260724":
return []
return [
{"trade_date": date, "data_type": "热股", "ts_code": "000001.SZ", "ts_name": "平安银行", "rank": 1, "pct_change": 5, "current_price": 10.5, "hot": 1000, "concept": '["银行"]'},
{"trade_date": date, "data_type": "概念板块", "ts_code": "885001.TI", "ts_name": "人工智能", "rank": 1, "pct_change": 1.98, "hot": 800},
]
if api_name == "dc_hot":
if date == "20260724":
return []
return [{"trade_date": date, "data_type": "A股市场", "ts_code": "000001.SZ", "ts_name": "平安银行", "rank": 3, "pct_change": 5, "current_price": 10.5}]
return []
class ConfiguredIfind:
configured = True
class MarketInsightsTests(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.database = ReviewDatabase(Path(self.temp.name) / "review.db")
self.service = MarketInsightsService(
self.database,
FakeMarketClient(),
now_provider=lambda: datetime(
2026, 7, 24, 9, 20, tzinfo=timezone(timedelta(hours=8))
),
)
def tearDown(self):
self.temp.cleanup()
def test_auction_falls_back_and_normalizes_factors(self):
payload = self.service.auction_center("20260724")
self.assertEqual(payload["meta"]["trade_date"], "2026-07-23")
self.assertTrue(payload["meta"]["carried_forward"])
self.assertEqual(payload["meta"]["phase"], "observing")
self.assertEqual(payload["summary"]["stock_count"], 2)
self.assertEqual(payload["rows"][0]["amount_million"], 5)
self.assertEqual(payload["rows"][0]["change"], 5)
self.assertEqual(payload["rows"][0]["expectation"], "超预期")
self.assertEqual(set(payload["expectations"]), {"超预期", "符合预期", "低于预期"})
self.assertFalse(payload["news_feedback"]["available"])
self.assertEqual(len(payload["amount_history"]), 10)
self.assertEqual(payload["amount_history"][-1]["stock_count"], 2)
self.assertEqual(payload["focus_rows"][0]["code"], "000001")
def test_auction_amount_history_uses_the_same_a_share_universe_as_summary(self):
self.database.upsert_stock_master([
{"ts_code": "000001.SZ", "name": "平安银行", "industry": "银行", "market": "主板", "list_date": "19910403"},
{"ts_code": "688001.SH", "name": "首日上市", "industry": "半导体", "market": "科创板", "list_date": "20260723"},
])
self.database.upsert_auction_factors([
{"ts_code": "000001.SZ", "trade_date": "20260723", "price": 10.5, "pre_close": 10, "amount": 5_000_000, "vol": 20_000},
{"ts_code": "688001.SH", "trade_date": "20260723", "price": 50, "pre_close": 10, "amount": 150_000_000, "vol": 3_000_000},
{"ts_code": "159001.SZ", "trade_date": "20260723", "price": 1.1, "pre_close": 1, "amount": 90_000_000, "vol": 90_000_000},
])
history = self.service._auction_amount_history("20260723")
self.assertEqual(history[-1]["stock_count"], 1)
self.assertEqual(history[-1]["amount_billion"], 0.05)
def test_real_limit_price_is_isolated_from_scored_candidates(self):
class OnePriceClient(FakeMarketClient):
def query(self, api_name, params=None, fields=""):
if api_name == "stk_limit":
return [
{"ts_code": "000001.SZ", "trade_date": "20260723", "up_limit": 10.5, "down_limit": 9},
{"ts_code": "000002.SZ", "trade_date": "20260723", "up_limit": 11, "down_limit": 9},
]
return super().query(api_name, params, fields)
service = MarketInsightsService(self.database, OnePriceClient(), self.service._now_provider)
payload = service.auction_center("20260724", force=True)
self.assertEqual([row["code"] for row in payload["one_price_rows"]], ["000001"])
self.assertNotIn("000001", {row["code"] for row in payload["rows"]})
def test_core_broken_pool_and_watchlist_are_kept_separate(self):
self.database.save_snapshot(
"20260723",
"test",
{
"limits": [{"code": "000001", "name": "平安银行", "sector": "银行", "streak": 3, "amount_billion": 8}],
"broken": [{"code": "000002", "name": "万科A", "sector": "房地产", "streak": 1}],
"sectors": [{"name": "银行", "count": 1, "leader": "平安银行"}],
},
)
user = self.database.create_user("auction-user", "salt", "hash")
self.database.save_watchlist(user["id"], "000002", "万科A", "房地产", "red")
payload = self.service.auction_center("20260724", force=True, user_id=user["id"])
rows = {row["code"]: row for row in payload["rows"]}
self.assertIn("昨日炸板", rows["000002"]["candidate_sources"])
self.assertIn("三板以上", rows["000001"]["core_tags"])
self.assertIn("000001", {row["code"] for row in payload["focus_rows"]})
self.assertEqual([row["code"] for row in payload["watchlist_rows"]], ["000002"])
anonymous = self.service.auction_center("20260724", user_id=0)
self.assertEqual(anonymous["watchlist_rows"], [])
def test_selection_window_does_not_disguise_previous_day_as_current(self):
service = MarketInsightsService(
self.database,
FakeMarketClient(),
now_provider=lambda: datetime(
2026, 7, 24, 9, 26, tzinfo=timezone(timedelta(hours=8))
),
)
payload = service.auction_center("20260724", force=True)
self.assertEqual(payload["meta"]["phase"], "selection")
self.assertFalse(payload["meta"]["available"])
self.assertFalse(payload["meta"]["carried_forward"])
self.assertEqual(payload["rows"], [])
def test_finalized_window_uses_and_persists_ifind_closing_snapshot(self):
service = MarketInsightsService(
self.database,
FakeMarketClient(),
now_provider=lambda: datetime(
2026, 7, 24, 9, 31, tzinfo=timezone(timedelta(hours=8))
),
ifind=ConfiguredIfind(),
)
calls = []
service._dynamic_auction_rows = lambda trade_date, baseline_date, user_id: (
calls.append((trade_date, baseline_date, user_id))
or [
{
"ts_code": "000001.SZ",
"trade_date": trade_date,
"price": 10.5,
"pre_close": 10,
"vol": 20_000,
"amount": 5_000_000,
"turnover_rate": 0.12,
"volume_ratio": 1.8,
"dynamic": True,
}
]
)
payload = service.auction_center("20260724")
cached = service.auction_center("20260724")
self.assertEqual(payload["meta"]["phase"], "finalized")
self.assertTrue(payload["meta"]["available"])
self.assertEqual(payload["summary"]["stock_count"], 1)
self.assertEqual(payload["rows"][0]["code"], "000001")
self.assertEqual(len(calls), 1)
self.assertTrue(cached["meta"]["cached"])
self.assertEqual(cached["summary"]["stock_count"], 1)
def test_theme_library_detail_and_popularity(self):
library = self.service.theme_library("20260724")
self.assertEqual(library["meta"]["trade_date"], "2026-07-23")
self.assertEqual(library["items"][0]["hot_rank"], 1)
detail = self.service.theme_detail("885001.TI", "20260724")
self.assertEqual(detail["summary"]["member_count"], 2)
self.assertEqual(detail["members"][0]["code"], "000001")
hot = self.service.popularity("20260724")
self.assertEqual(hot["summary"]["dual_count"], 1)
self.assertEqual(hot["combined"][0]["name"], "平安银行")
def test_feature_pages_use_local_data_when_trade_context_is_offline(self):
class OfflineClient:
def resolve_trade_context(self, requested: str):
raise TushareError("offline")
def query(self, api_name, params=None, fields=""):
raise TushareError("offline")
self.database.save_snapshot(
"20260723", "test", {"meta": {"trade_date": "2026-07-23"}}
)
self.database.save_snapshot(
"20260724", "test", {"meta": {"trade_date": "2026-07-24"}}
)
self.database.upsert_stock_master([
{"ts_code": "000001.SZ", "name": "平安银行", "industry": "银行", "market": "主板", "list_date": "19910403"}
])
self.database.upsert_auction_factors([
{"ts_code": "000001.SZ", "trade_date": "20260724", "price": 10.5, "pre_close": 10, "amount": 5_000_000, "vol": 20_000, "turnover_rate": 0.12, "volume_ratio": 1.8}
])
self.database.save_data_snapshot(
"theme_library_v1", "20260724", "market",
{"meta": {"trade_date": "2026-07-24"}, "summary": {"theme_count": 1}, "items": [{"code": "885001.TI", "name": "人工智能", "member_count": 2}]},
)
self.database.save_data_snapshot(
"popularity_v1", "20260724", "market",
{"meta": {"trade_date": "2026-07-24"}, "summary": {"ths_count": 1, "dc_count": 0, "dual_count": 0}, "combined": [{"name": "平安银行"}], "ths": [], "dc": []},
)
service = MarketInsightsService(self.database, OfflineClient(), self.service._now_provider)
auction = service.auction_center("20260725")
self.assertEqual(auction["meta"]["trade_date"], "2026-07-24")
self.assertEqual(auction["summary"]["stock_count"], 1)
themes = service.theme_library("20260725")
self.assertTrue(themes["meta"]["cached"])
self.assertEqual(themes["items"][0]["name"], "人工智能")
popularity = service.popularity("20260725")
self.assertTrue(popularity["meta"]["cached"])
self.assertEqual(popularity["combined"][0]["name"], "平安银行")
class AuctionScreenerFactorTests(unittest.TestCase):
def test_auction_fields_are_available_to_formula_and_factor_rows(self):
with tempfile.TemporaryDirectory() as temporary:
database = ReviewDatabase(Path(temporary) / "review.db")
database.upsert_stock_master([
{"ts_code": "000001.SZ", "name": "平安银行", "industry": "银行", "market": "主板", "list_date": "19910403"}
])
dates = [f"202606{day:02d}" for day in range(1, 22)]
database.upsert_daily_bars([
{"ts_code": "000001.SZ", "trade_date": trade_date, "open": 10, "high": 11, "low": 9, "close": 10 + index * 0.1, "pct_chg": 1, "vol": 1000 + index, "amount": 200000}
for index, trade_date in enumerate(dates)
])
database.upsert_daily_indicators([
{"ts_code": "000001.SZ", "trade_date": dates[-1], "turnover_rate": 2, "volume_ratio": 1.2, "circ_mv": 100000, "total_mv": 120000}
])
database.upsert_auction_factors([
{"ts_code": "000001.SZ", "trade_date": dates[-1], "price": 12.6, "pre_close": 12, "amount": 8_000_000, "vol": 30000, "turnover_rate": 0.18, "volume_ratio": 2.1}
])
rows, actual_date = ScreenerEngine(database).build_factors(dates[-1])
self.assertEqual(actual_date, dates[-1])
self.assertEqual(rows[0]["auction_change"], 5)
self.assertEqual(rows[0]["auction_amount_million"], 8)
self.assertEqual(rows[0]["auction_volume_ratio"], 2.1)
self.assertIn("auction_change", FACTOR_FIELDS)
if __name__ == "__main__":
unittest.main()