diff --git a/backend/features/market/charts.py b/backend/features/market/charts.py index 97cd8b7..8b40c94 100644 --- a/backend/features/market/charts.py +++ b/backend/features/market/charts.py @@ -12,6 +12,7 @@ from datetime import datetime, time as dt_time, timedelta from threading import Lock from typing import Any, ClassVar +from backend.bootstrap.config import tushare_code as _stock_market_code from backend.data.providers.ifind_client import IfindError, IfindHttpClient @@ -475,16 +476,6 @@ def _ifind_point(row: dict[str, Any]) -> dict[str, Any] | None: } -def _stock_market_code(code: str) -> str: - if code.startswith(("4", "8", "9")): - suffix = "BJ" - elif code.startswith("6"): - suffix = "SH" - else: - suffix = "SZ" - return f"{code}.{suffix}" - - def _number(value: Any) -> float: try: return float(value or 0) diff --git a/tests/test_market_symbol_normalization.py b/tests/test_market_symbol_normalization.py new file mode 100644 index 0000000..9ae7cfa --- /dev/null +++ b/tests/test_market_symbol_normalization.py @@ -0,0 +1,26 @@ +from __future__ import annotations + +import unittest + +from backend.bootstrap.config import tushare_code +from backend.features.market import charts + + +class MarketSymbolNormalizationTests(unittest.TestCase): + def test_chart_and_market_services_share_one_suffix_converter(self) -> None: + self.assertIs(charts._stock_market_code, tushare_code) + + def test_existing_exchange_mapping_is_preserved(self) -> None: + cases = { + "000001": "000001.SZ", + "600000": "600000.SH", + "430047": "430047.BJ", + "830799": "830799.BJ", + } + for code, expected in cases.items(): + with self.subTest(code=code): + self.assertEqual(charts._stock_market_code(code), expected) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_preservation_slice_market.py b/tests/test_preservation_slice_market.py index cd101bd..cc370a6 100644 --- a/tests/test_preservation_slice_market.py +++ b/tests/test_preservation_slice_market.py @@ -9,6 +9,7 @@ import chart_data_provider import ifind_client import realtime_aggregator import tushare_client +from backend.bootstrap import config as bootstrap_config from backend.data import realtime from backend.data.providers import ifind_client as canonical_ifind from backend.data.providers import tushare_client as canonical_tushare @@ -101,6 +102,18 @@ def top_level_definitions(path: Path) -> dict[str, str]: } +def function_contract(path: Path, name: str) -> tuple[str, str]: + tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + function = next( + node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == name + ) + body = ast.Module(body=function.body, type_ignores=[]) + return ( + ast.dump(function.args, include_attributes=False), + ast.dump(body, include_attributes=False), + ) + + class MarketSliceSourceEquivalenceTests(unittest.TestCase): def test_market_service_methods_are_exact_original_ast(self) -> None: original = class_methods(ORIGINAL_ROOT / "server.py", "DashboardService") @@ -145,10 +158,21 @@ class MarketSliceSourceEquivalenceTests(unittest.TestCase): top_level_definitions(ORIGINAL_ROOT / "tushare_client.py"), top_level_definitions(APP_ROOT / "backend/data/providers/tushare_client.py"), ) + original_charts = top_level_definitions(ORIGINAL_ROOT / "chart_data_provider.py") + original_charts.pop("_stock_market_code") self.assertEqual( - top_level_definitions(ORIGINAL_ROOT / "chart_data_provider.py"), + original_charts, top_level_definitions(APP_ROOT / "backend/features/market/charts.py"), ) + self.assertEqual( + function_contract( + ORIGINAL_ROOT / "chart_data_provider.py", "_stock_market_code" + ), + function_contract( + APP_ROOT / "backend/bootstrap/config.py", "tushare_code" + ), + ) + self.assertIs(charts._stock_market_code, bootstrap_config.tushare_code) def test_relocated_frontend_preserves_original_market_runtime_and_styles(self) -> None: assert_frontend_runtime_matches_audited_baseline(self)