mirror of
https://github.com/ZhuLinsen/daily_stock_analysis.git
synced 2026-10-06 12:33:53 +08:00
fix:backtest stock identity and daily-window correctness refactor (#2073)
* fix: unify local daily window stock code resolution * fix: enforce authoritative daily window resolution * add test * fix: converge daily window resolution contract * fix: preserve daily stock identity compatibility * fix: rebuild legacy foreign market snapshots * fix(backtest): preserve legacy JP/KR bare-code compatibility * fix(backtest): disambiguate legacy offshore stock codes * fix(backtest): prevent cross-market alias collisions
This commit is contained in:
@@ -8,6 +8,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/).
|
||||
> For user-friendly release highlights, see the [GitHub Releases](https://github.com/ZhuLinsen/daily_stock_analysis/releases) page.
|
||||
|
||||
## [Unreleased]
|
||||
- [修复] 统一等价股票代码的本地日线候选与同源窗口解析;冲突沪深交易所代码不再降级匹配裸码,回测仅接受快照或交易日历确认的起点,并在同一起点中优先完整的单一代码窗口。
|
||||
<!-- 新条目格式:- [类型] 描述(类型取值:新功能/改进/修复/文档/测试/chore)-->
|
||||
<!-- 每条独立一行追加到本段末尾,无需分类标题,合并时冲突最小 -->
|
||||
|
||||
|
||||
@@ -248,6 +248,60 @@ def get_effective_trading_date(
|
||||
return fallback_date
|
||||
|
||||
|
||||
def resolve_historical_daily_bar_date(
|
||||
market: Optional[str],
|
||||
target_date: date,
|
||||
phase: Optional[str],
|
||||
) -> Optional[date]:
|
||||
"""Resolve the completed daily bar that a historical phase could consume.
|
||||
|
||||
A persisted ``effective_daily_bar_date`` remains the primary authority.
|
||||
This fallback is only for older snapshots that still contain a trustworthy
|
||||
phase. It fails closed for missing, unknown, or calendar-inconsistent phases.
|
||||
"""
|
||||
if market not in MARKET_EXCHANGE or not _XCALS_AVAILABLE:
|
||||
return None
|
||||
|
||||
normalized_phase = str(phase or "").strip().lower()
|
||||
if normalized_phase not in {
|
||||
"premarket",
|
||||
"intraday",
|
||||
"lunch_break",
|
||||
"closing_auction",
|
||||
"postmarket",
|
||||
"non_trading",
|
||||
}:
|
||||
return None
|
||||
|
||||
try:
|
||||
cal = xcals.get_calendar(MARKET_EXCHANGE[market])
|
||||
is_session = bool(cal.is_session(target_date))
|
||||
|
||||
if normalized_phase in {
|
||||
"premarket",
|
||||
"intraday",
|
||||
"lunch_break",
|
||||
"closing_auction",
|
||||
}:
|
||||
if not is_session:
|
||||
return None
|
||||
session = cal.date_to_session(target_date, direction="previous")
|
||||
return cal.previous_session(session).date()
|
||||
|
||||
if normalized_phase == "postmarket":
|
||||
return target_date if is_session else None
|
||||
|
||||
if is_session:
|
||||
return None
|
||||
return cal.date_to_session(target_date, direction="previous").date()
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"trading_calendar.resolve_historical_daily_bar_date fail-closed: %s",
|
||||
e,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _as_market_datetime(value: Any, tz_name: str) -> Optional[datetime]:
|
||||
"""
|
||||
Convert exchange-calendar timestamps into market-local datetimes.
|
||||
|
||||
@@ -20,6 +20,7 @@ logger = logging.getLogger(__name__)
|
||||
_STOCK_INDEX_FILENAME = "stocks.index.json"
|
||||
_STOCK_INDEX_CACHE: Dict[str, str] | None = None
|
||||
_STOCK_CODE_LOOKUP_CACHE: Dict[str, str] | None = None
|
||||
_STOCK_CODE_CANDIDATES_CACHE: Dict[str, tuple[str, ...]] | None = None
|
||||
_REMOTE_INDEX_VALIDITY_CACHE: tuple[Path, float, int, bool] | None = None
|
||||
_STOCK_INDEX_CACHE_LOCK = RLock()
|
||||
|
||||
@@ -118,40 +119,47 @@ def _is_jp_kr_index_code(code: str) -> bool:
|
||||
return get_suffix_market(code) in {"jp", "kr"}
|
||||
|
||||
|
||||
def _build_stock_code_lookup(raw_items: list) -> Dict[str, str]:
|
||||
exact_lookup: dict[str, set[str]] = {}
|
||||
suffix_base_lookup: dict[str, set[str]] = {}
|
||||
def _build_stock_code_candidates(raw_items: list) -> dict[str, set[str]]:
|
||||
candidates: dict[str, set[str]] = {}
|
||||
|
||||
for item in raw_items:
|
||||
if not isinstance(item, list) or len(item) < 2:
|
||||
continue
|
||||
|
||||
canonical_code = str(item[0] or "").strip()
|
||||
display_code = str(item[1] or "").strip()
|
||||
canonical_code = str(item[0] or "").strip().upper()
|
||||
display_code = str(item[1] or "").strip().upper()
|
||||
if not canonical_code:
|
||||
continue
|
||||
if not _is_jp_kr_index_code(canonical_code):
|
||||
continue
|
||||
if len(item) > 8 and item[8] is False:
|
||||
continue
|
||||
|
||||
_add_code_lookup(exact_lookup, canonical_code, canonical_code)
|
||||
_add_code_lookup(exact_lookup, display_code, canonical_code)
|
||||
indexed_market = (
|
||||
str(item[6] or "").strip().lower()
|
||||
if len(item) > 6
|
||||
else ""
|
||||
)
|
||||
if indexed_market not in {"cn", "hk", "jp", "kr"}:
|
||||
indexed_market = get_suffix_market(canonical_code) or ""
|
||||
if indexed_market not in {"cn", "hk", "jp", "kr"}:
|
||||
continue
|
||||
|
||||
canonical_upper = canonical_code.upper()
|
||||
if "." in canonical_upper and suffix_base_lookup_allowed(canonical_upper):
|
||||
base, _suffix = canonical_upper.rsplit(".", 1)
|
||||
if indexed_market in {"jp", "kr"}:
|
||||
_add_code_lookup(candidates, canonical_code, canonical_code)
|
||||
_add_code_lookup(candidates, display_code, canonical_code)
|
||||
if "." in canonical_code and suffix_base_lookup_allowed(canonical_code):
|
||||
base, _suffix = canonical_code.rsplit(".", 1)
|
||||
if base.isdigit():
|
||||
_add_code_lookup(candidates, base, canonical_code)
|
||||
elif indexed_market == "hk":
|
||||
base = canonical_code.removesuffix(".HK")
|
||||
if base.isdigit():
|
||||
_add_code_lookup(suffix_base_lookup, base, canonical_code)
|
||||
_add_code_lookup(candidates, base.lstrip("0") or "0", canonical_code)
|
||||
elif indexed_market == "cn":
|
||||
base = canonical_code.rsplit(".", 1)[0]
|
||||
if base.isdigit() and len(base) == 6:
|
||||
_add_code_lookup(candidates, base, canonical_code)
|
||||
|
||||
result: Dict[str, str] = {}
|
||||
for lookup in (exact_lookup, suffix_base_lookup):
|
||||
for key, codes in lookup.items():
|
||||
if key in result:
|
||||
continue
|
||||
if len(codes) == 1:
|
||||
result[key] = next(iter(codes))
|
||||
return result
|
||||
return candidates
|
||||
|
||||
|
||||
def _load_stock_index_file(index_path: Path) -> Dict[str, str]:
|
||||
@@ -292,7 +300,56 @@ def resolve_index_stock_code(query: str) -> str | None:
|
||||
if not code:
|
||||
return None
|
||||
|
||||
return get_stock_code_index_map().get(code)
|
||||
candidates = resolve_index_stock_code_candidates(code)
|
||||
if len(candidates) != 1:
|
||||
return None
|
||||
candidate = candidates[0]
|
||||
return candidate if _is_jp_kr_index_code(candidate) else None
|
||||
|
||||
|
||||
def resolve_index_stock_code_candidates(query: str) -> tuple[str, ...]:
|
||||
"""Return active indexed identities sharing one supported code alias."""
|
||||
code = str(query or "").strip().upper()
|
||||
if not code:
|
||||
return ()
|
||||
return get_stock_code_candidates_map().get(code, ())
|
||||
|
||||
|
||||
def get_stock_code_candidates_map() -> Dict[str, tuple[str, ...]]:
|
||||
"""Lazily load all indexed identities needed to detect bare-code ambiguity."""
|
||||
global _STOCK_CODE_CANDIDATES_CACHE
|
||||
|
||||
if _STOCK_CODE_CANDIDATES_CACHE is not None:
|
||||
return _STOCK_CODE_CANDIDATES_CACHE
|
||||
|
||||
with _STOCK_INDEX_CACHE_LOCK:
|
||||
if _STOCK_CODE_CANDIDATES_CACHE is not None:
|
||||
return _STOCK_CODE_CANDIDATES_CACHE
|
||||
|
||||
merged_candidates: dict[str, set[str]] = {}
|
||||
remote_path = get_remote_stock_index_cache_path()
|
||||
for index_path in _get_fresh_stock_index_candidates(
|
||||
get_stock_index_candidate_paths(),
|
||||
remote_path,
|
||||
):
|
||||
try:
|
||||
raw_items = _load_stock_index_payload(index_path)
|
||||
if _same_path(index_path, remote_path):
|
||||
validate_stock_index_payload(raw_items)
|
||||
for key, values in _build_stock_code_candidates(raw_items).items():
|
||||
merged_candidates.setdefault(key, set()).update(values)
|
||||
except (OSError, TypeError, ValueError) as exc:
|
||||
logger.debug(
|
||||
"[股票索引] 解析代码候选失败 %s: %s",
|
||||
index_path,
|
||||
exc,
|
||||
)
|
||||
|
||||
_STOCK_CODE_CANDIDATES_CACHE = {
|
||||
key: tuple(sorted(values))
|
||||
for key, values in merged_candidates.items()
|
||||
}
|
||||
return _STOCK_CODE_CANDIDATES_CACHE
|
||||
|
||||
|
||||
def get_stock_code_index_map() -> Dict[str, str]:
|
||||
@@ -306,48 +363,24 @@ def get_stock_code_index_map() -> Dict[str, str]:
|
||||
if _STOCK_CODE_LOOKUP_CACHE is not None:
|
||||
return _STOCK_CODE_LOOKUP_CACHE
|
||||
|
||||
merged_lookup: Dict[str, str] = {}
|
||||
remote_path = get_remote_stock_index_cache_path()
|
||||
for index_path in _get_fresh_stock_index_candidates(get_stock_index_candidate_paths(), remote_path):
|
||||
try:
|
||||
raw_items = _load_stock_index_payload(index_path)
|
||||
if _same_path(index_path, remote_path):
|
||||
validate_stock_index_payload(raw_items)
|
||||
for key, value in _build_stock_code_lookup(raw_items).items():
|
||||
merged_lookup.setdefault(key, value)
|
||||
except (OSError, TypeError, ValueError) as exc:
|
||||
logger.debug("[鑲$エ绱㈠紩] 瑙f瀽浠g爜绱㈠紩澶辫触 %s: %s", index_path, exc)
|
||||
merged_lookup = {
|
||||
key: values[0]
|
||||
for key, values in get_stock_code_candidates_map().items()
|
||||
if len(values) == 1 and _is_jp_kr_index_code(values[0])
|
||||
}
|
||||
|
||||
_STOCK_CODE_LOOKUP_CACHE = merged_lookup
|
||||
return _STOCK_CODE_LOOKUP_CACHE
|
||||
|
||||
|
||||
def _resolve_index_stock_code_uncached(query: str) -> str | None:
|
||||
code = str(query or "").strip().upper()
|
||||
if not code:
|
||||
return None
|
||||
|
||||
remote_path = get_remote_stock_index_cache_path()
|
||||
for index_path in _get_fresh_stock_index_candidates(get_stock_index_candidate_paths(), remote_path):
|
||||
try:
|
||||
raw_items = _load_stock_index_payload(index_path)
|
||||
if _same_path(index_path, remote_path):
|
||||
validate_stock_index_payload(raw_items)
|
||||
resolved = _build_stock_code_lookup(raw_items).get(code)
|
||||
if resolved:
|
||||
return resolved
|
||||
except (OSError, TypeError, ValueError) as exc:
|
||||
logger.debug("[股票索引] 解析代码索引失败 %s: %s", index_path, exc)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def clear_stock_index_cache() -> None:
|
||||
"""Clear the in-process stock index lookup cache."""
|
||||
global _REMOTE_INDEX_VALIDITY_CACHE, _STOCK_INDEX_CACHE, _STOCK_CODE_LOOKUP_CACHE
|
||||
global _REMOTE_INDEX_VALIDITY_CACHE
|
||||
global _STOCK_CODE_CANDIDATES_CACHE, _STOCK_CODE_LOOKUP_CACHE, _STOCK_INDEX_CACHE
|
||||
with _STOCK_INDEX_CACHE_LOCK:
|
||||
_STOCK_INDEX_CACHE = None
|
||||
_STOCK_CODE_LOOKUP_CACHE = None
|
||||
_STOCK_CODE_CANDIDATES_CACHE = None
|
||||
_REMOTE_INDEX_VALIDITY_CACHE = None
|
||||
|
||||
|
||||
|
||||
@@ -13,9 +13,8 @@ from typing import List, Optional, Tuple
|
||||
|
||||
from sqlalchemy import and_, delete, desc, func, or_, select
|
||||
|
||||
from data_provider.base import is_bse_code
|
||||
from src.core.backtest_engine import OVERALL_SENTINEL_CODE
|
||||
from src.services.stock_code_utils import normalize_code as normalize_backtest_code
|
||||
from src.services.stock_code_utils import resolve_daily_stock_identity
|
||||
|
||||
from src.storage import BacktestResult, BacktestSummary, DatabaseManager, AnalysisHistory
|
||||
|
||||
@@ -516,136 +515,16 @@ class BacktestRepository:
|
||||
else:
|
||||
raw_code = raw_code.upper()
|
||||
|
||||
normalized_code = normalize_backtest_code(raw_code)
|
||||
|
||||
candidates = [raw_code]
|
||||
if normalized_code and normalized_code != raw_code:
|
||||
candidates.append(normalized_code)
|
||||
candidates.extend(BacktestRepository._build_market_code_variants(raw_code, normalized_code))
|
||||
if raw_code == OVERALL_SENTINEL_CODE:
|
||||
candidates = [OVERALL_SENTINEL_CODE]
|
||||
else:
|
||||
identity = resolve_daily_stock_identity(raw_code)
|
||||
candidates = list(identity.code_candidates) if identity is not None else []
|
||||
if not candidates:
|
||||
return [column.in_([])]
|
||||
|
||||
if len(candidates) == 1:
|
||||
return [column == candidates[0]]
|
||||
|
||||
unique = list(dict.fromkeys(candidates))
|
||||
return [or_(*[column == candidate for candidate in unique])]
|
||||
|
||||
@staticmethod
|
||||
def _build_hk_market_variants(hk_digits: str) -> List[str]:
|
||||
"""Build normalized HK variants for padded/unpadded code shapes."""
|
||||
if not hk_digits.isdigit() or not hk_digits:
|
||||
return []
|
||||
|
||||
padded = hk_digits.zfill(5)
|
||||
unpadded = padded.lstrip("0") or "0"
|
||||
|
||||
variants: List[str] = [
|
||||
f"HK{padded}",
|
||||
f"{padded}.HK",
|
||||
padded,
|
||||
f"HK{unpadded}",
|
||||
f"{unpadded}.HK",
|
||||
f"HK.{padded}",
|
||||
]
|
||||
if unpadded == padded:
|
||||
variants.pop(3)
|
||||
variants.pop(3)
|
||||
|
||||
# Keep legacy no-leading-zero bare form for 1-3 digit inputs.
|
||||
if len(unpadded) <= 3 and unpadded != padded:
|
||||
variants.append(unpadded)
|
||||
variants.append(f"HK.{unpadded}")
|
||||
|
||||
return variants
|
||||
|
||||
@staticmethod
|
||||
def _build_market_code_variants(raw_code: str, normalized_code: str) -> List[str]:
|
||||
"""Return additional market-formatted variants for safe stock-code matching."""
|
||||
variants: List[str] = []
|
||||
if not raw_code:
|
||||
return variants
|
||||
|
||||
raw_code_upper = raw_code.upper()
|
||||
normalized_upper = normalized_code.upper() if normalized_code else ""
|
||||
|
||||
def _add_us_variants(code: str) -> None:
|
||||
if not code:
|
||||
return
|
||||
if code.endswith(".US"):
|
||||
bare = code[:-3]
|
||||
if bare.isalpha() and 1 <= len(bare) <= 5:
|
||||
variants.append(bare)
|
||||
return
|
||||
if "." not in code and code.isalpha() and 1 <= len(code) <= 5:
|
||||
variants.append(f"{code}.US")
|
||||
|
||||
_add_us_variants(raw_code_upper)
|
||||
if normalized_upper != raw_code_upper:
|
||||
_add_us_variants(normalized_upper)
|
||||
|
||||
def _explicit_exchange() -> Optional[str]:
|
||||
if raw_code_upper.startswith(("SH", "SS")) or raw_code_upper.endswith((".SH", ".SS")):
|
||||
return "SH"
|
||||
if raw_code_upper.startswith("SZ") or raw_code_upper.endswith(".SZ"):
|
||||
return "SZ"
|
||||
if raw_code_upper.startswith("BJ") or raw_code_upper.endswith(".BJ"):
|
||||
return "BJ"
|
||||
return None
|
||||
|
||||
def _exchange_by_code(base: str) -> str:
|
||||
if is_bse_code(base):
|
||||
return "BJ"
|
||||
if base.startswith(("5", "6")):
|
||||
return "SH"
|
||||
return "SZ"
|
||||
|
||||
if normalized_upper.isdigit() and len(normalized_upper) == 6:
|
||||
explicit_exchange = _explicit_exchange()
|
||||
if explicit_exchange is not None and explicit_exchange != _exchange_by_code(normalized_upper):
|
||||
return []
|
||||
|
||||
if raw_code_upper.startswith(("SH", "SS")) or raw_code_upper.endswith(".SH") or raw_code_upper.endswith(".SS"):
|
||||
exchange = "SH"
|
||||
elif raw_code_upper.startswith("SZ") or raw_code_upper.endswith(".SZ"):
|
||||
exchange = "SZ"
|
||||
elif raw_code_upper.startswith("BJ") or raw_code_upper.endswith(".BJ") or is_bse_code(normalized_upper):
|
||||
exchange = "BJ"
|
||||
elif normalized_upper.startswith(("5", "6", "9")):
|
||||
exchange = "SH"
|
||||
else:
|
||||
exchange = "SZ"
|
||||
|
||||
variants.append(f"{exchange}{normalized_upper}")
|
||||
variants.append(f"{normalized_upper}.{exchange}")
|
||||
variants.append(f"{exchange}.{normalized_upper}")
|
||||
if exchange == "SH":
|
||||
variants.append(f"SS{normalized_upper}")
|
||||
variants.append(f"{normalized_upper}.SS")
|
||||
variants.append(f"SS.{normalized_upper}")
|
||||
|
||||
if (
|
||||
normalized_upper.startswith("HK")
|
||||
and len(normalized_upper) > 2
|
||||
and normalized_upper[2:].isdigit()
|
||||
and len(normalized_upper[2:]) <= 5
|
||||
):
|
||||
variants.extend(BacktestRepository._build_hk_market_variants(normalized_upper[2:]))
|
||||
|
||||
if (
|
||||
raw_code_upper.startswith("HK.")
|
||||
and raw_code_upper[3:].isdigit()
|
||||
and len(raw_code_upper[3:]) <= 5
|
||||
):
|
||||
variants.extend(BacktestRepository._build_hk_market_variants(raw_code_upper[3:]))
|
||||
|
||||
if (
|
||||
raw_code_upper.endswith(".HK")
|
||||
and raw_code_upper[:-3].isdigit()
|
||||
and 1 <= len(raw_code_upper[:-3]) <= 5
|
||||
):
|
||||
hk_digits = raw_code_upper.rsplit(".", 1)[0]
|
||||
variants.extend(BacktestRepository._build_hk_market_variants(hk_digits))
|
||||
|
||||
if raw_code_upper.isdigit() and len(raw_code_upper) in (4, 5):
|
||||
variants.extend(BacktestRepository._build_hk_market_variants(raw_code_upper))
|
||||
|
||||
return variants
|
||||
|
||||
+131
-167
@@ -13,13 +13,22 @@ from sqlalchemy import and_, select
|
||||
from data_provider.base import canonical_stock_code, normalize_stock_code
|
||||
from src.config import get_config
|
||||
from src.core.backtest_engine import OVERALL_SENTINEL_CODE, BacktestEngine, EvaluationConfig
|
||||
from src.market_phase_summary import extract_market_phase_summary, normalize_analysis_phase_bucket
|
||||
from src.core.trading_calendar import resolve_historical_daily_bar_date
|
||||
from src.market_phase_summary import (
|
||||
extract_market_phase_summary,
|
||||
normalize_analysis_phase_bucket,
|
||||
rebuild_market_phase_summary_for_stock_code,
|
||||
)
|
||||
from src.repositories.backtest_repo import BacktestRepository
|
||||
from src.repositories.stock_repo import StockRepository
|
||||
from src.schemas.decision_action import build_action_fields
|
||||
from src.services.stock_code_utils import (
|
||||
normalize_code as normalize_backtest_code,
|
||||
resolve_daily_stock_identity,
|
||||
)
|
||||
from src.services.stock_daily_window_resolver import resolve_stock_daily_window
|
||||
from src.storage import BacktestResult, BacktestSummary, DatabaseManager
|
||||
from src.utils.data_processing import parse_json_field
|
||||
from src.services.stock_code_utils import normalize_code as normalize_backtest_code
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -55,6 +64,12 @@ class BacktestService:
|
||||
|
||||
if eval_window_days is None:
|
||||
eval_window_days = getattr(config, "backtest_eval_window_days", 10)
|
||||
if (
|
||||
isinstance(eval_window_days, bool)
|
||||
or not isinstance(eval_window_days, int)
|
||||
or eval_window_days <= 0
|
||||
):
|
||||
raise ValueError("eval_window_days must be a positive integer")
|
||||
if min_age_days is None:
|
||||
min_age_days = getattr(config, "backtest_min_age_days", 14)
|
||||
|
||||
@@ -109,25 +124,63 @@ class BacktestService:
|
||||
)
|
||||
)
|
||||
continue
|
||||
daily_code_candidates = self._build_daily_code_candidates(analysis.code)
|
||||
start_daily = self._get_start_daily_for_candidates(
|
||||
code_candidates=daily_code_candidates,
|
||||
analysis_date=analysis_date,
|
||||
phase_summary = extract_market_phase_summary(
|
||||
analysis.context_snapshot
|
||||
)
|
||||
persisted_market = (
|
||||
str(phase_summary.get("market") or "").strip().lower()
|
||||
if isinstance(phase_summary, dict)
|
||||
else None
|
||||
)
|
||||
daily_identity = resolve_daily_stock_identity(
|
||||
analysis.code,
|
||||
market_hint=persisted_market,
|
||||
)
|
||||
daily_code_candidates = (
|
||||
list(daily_identity.code_candidates)
|
||||
if daily_identity is not None
|
||||
else []
|
||||
)
|
||||
expected_start_date = self._resolve_expected_start_date(
|
||||
analysis=analysis,
|
||||
analysis_date=analysis_date,
|
||||
market=daily_identity.market if daily_identity is not None else None,
|
||||
stock_code=(
|
||||
daily_identity.normalized_code
|
||||
if daily_identity is not None
|
||||
else None
|
||||
),
|
||||
)
|
||||
daily_window = None
|
||||
if expected_start_date is not None:
|
||||
daily_window = resolve_stock_daily_window(
|
||||
stock_repo=self.stock_repo,
|
||||
code_candidates=daily_code_candidates,
|
||||
expected_start_date=expected_start_date,
|
||||
eval_window_days=int(eval_window_days),
|
||||
)
|
||||
|
||||
if start_daily is None or start_daily.close is None:
|
||||
refill_code = daily_code_candidates[0] if daily_code_candidates else analysis.code
|
||||
if (
|
||||
daily_identity is not None
|
||||
and expected_start_date is not None
|
||||
and (
|
||||
daily_window is None
|
||||
or len(daily_window.forward_bars) < int(eval_window_days)
|
||||
)
|
||||
):
|
||||
self._try_fill_daily_data(
|
||||
code=refill_code,
|
||||
analysis_date=analysis_date,
|
||||
code=daily_identity.refill_code,
|
||||
analysis_date=expected_start_date,
|
||||
eval_window_days=eval_window_days,
|
||||
)
|
||||
start_daily = self._get_start_daily_for_candidates(
|
||||
daily_window = resolve_stock_daily_window(
|
||||
stock_repo=self.stock_repo,
|
||||
code_candidates=daily_code_candidates,
|
||||
analysis_date=analysis_date,
|
||||
expected_start_date=expected_start_date,
|
||||
eval_window_days=int(eval_window_days),
|
||||
)
|
||||
|
||||
if start_daily is None or start_daily.close is None:
|
||||
if daily_window is None:
|
||||
insufficient += 1
|
||||
results_to_save.append(
|
||||
BacktestResult(
|
||||
@@ -143,40 +196,11 @@ class BacktestService:
|
||||
)
|
||||
continue
|
||||
|
||||
matched_daily_code = start_daily.code or (
|
||||
daily_code_candidates[0] if daily_code_candidates else analysis.code
|
||||
)
|
||||
forward_bars = self._get_forward_bars_by_candidates(
|
||||
code_candidates=daily_code_candidates,
|
||||
analysis_date=start_daily.date,
|
||||
eval_window_days=int(eval_window_days),
|
||||
preferred_code=matched_daily_code,
|
||||
)
|
||||
|
||||
if len(forward_bars) < int(eval_window_days):
|
||||
for fill_code in self._ordered_daily_refill_codes(
|
||||
code_candidates=daily_code_candidates,
|
||||
preferred_code=matched_daily_code,
|
||||
):
|
||||
self._try_fill_daily_data(
|
||||
code=fill_code,
|
||||
analysis_date=start_daily.date,
|
||||
eval_window_days=eval_window_days,
|
||||
)
|
||||
forward_bars = self._get_forward_bars_by_candidates(
|
||||
code_candidates=daily_code_candidates,
|
||||
analysis_date=start_daily.date,
|
||||
eval_window_days=int(eval_window_days),
|
||||
preferred_code=matched_daily_code,
|
||||
)
|
||||
if len(forward_bars) >= int(eval_window_days):
|
||||
break
|
||||
|
||||
evaluation = BacktestEngine.evaluate_single(
|
||||
operation_advice=analysis.operation_advice,
|
||||
analysis_date=start_daily.date,
|
||||
start_price=float(start_daily.close),
|
||||
forward_bars=forward_bars,
|
||||
analysis_date=daily_window.start_bar.date,
|
||||
start_price=float(daily_window.start_bar.close),
|
||||
forward_bars=daily_window.forward_bars,
|
||||
stop_loss=analysis.stop_loss,
|
||||
take_profit=analysis.take_profit,
|
||||
config=eval_config,
|
||||
@@ -412,49 +436,15 @@ class BacktestService:
|
||||
filtered.append(analysis)
|
||||
return filtered
|
||||
|
||||
def _get_start_daily_for_candidates(self, *, code_candidates: List[str], analysis_date: date):
|
||||
best_daily = None
|
||||
best_rank = len(code_candidates)
|
||||
for rank, candidate in enumerate(code_candidates):
|
||||
daily = self.stock_repo.get_start_daily(code=candidate, analysis_date=analysis_date)
|
||||
if daily is None:
|
||||
continue
|
||||
if best_daily is None or daily.date > best_daily.date or (
|
||||
daily.date == best_daily.date and rank < best_rank
|
||||
):
|
||||
best_daily = daily
|
||||
best_rank = rank
|
||||
return best_daily
|
||||
|
||||
@staticmethod
|
||||
def _build_daily_code_candidates(code: Optional[str]) -> List[str]:
|
||||
if not code:
|
||||
return []
|
||||
|
||||
raw_code = str(code).strip()
|
||||
if not raw_code:
|
||||
return []
|
||||
|
||||
raw_code = raw_code.upper()
|
||||
normalized_code = normalize_stock_code(raw_code)
|
||||
backtest_normalized_code = normalize_backtest_code(raw_code)
|
||||
candidates = [raw_code]
|
||||
for candidate in (normalized_code, backtest_normalized_code):
|
||||
if candidate and candidate != raw_code:
|
||||
candidates.append(candidate)
|
||||
for candidate in list(candidates):
|
||||
candidates.extend(BacktestRepository._build_market_code_variants(raw_code, candidate))
|
||||
return list(dict.fromkeys(candidate for candidate in candidates if candidate))
|
||||
|
||||
@staticmethod
|
||||
def _normalize_code(code: Optional[str]) -> Optional[str]:
|
||||
if not code:
|
||||
return None
|
||||
|
||||
normalized = normalize_backtest_code(str(code).strip())
|
||||
if normalized is None:
|
||||
identity = resolve_daily_stock_identity(str(code).strip())
|
||||
if identity is None:
|
||||
raise ValueError(f"非法股票代码格式: {code}")
|
||||
return normalized
|
||||
return identity.normalized_code
|
||||
|
||||
@staticmethod
|
||||
def _normalize_summary_code(code: Optional[str]) -> Optional[str]:
|
||||
@@ -469,90 +459,7 @@ class BacktestService:
|
||||
|
||||
@staticmethod
|
||||
def _normalize_code_for_display(code: Optional[str]) -> Optional[str]:
|
||||
if not code:
|
||||
return None
|
||||
|
||||
normalized = normalize_backtest_code(str(code).strip())
|
||||
if normalized is None:
|
||||
raise ValueError(f"非法股票代码格式: {code}")
|
||||
return normalized
|
||||
|
||||
@staticmethod
|
||||
def _ordered_candidate_codes(
|
||||
*,
|
||||
code_candidates: List[str],
|
||||
preferred_code: Optional[str] = None,
|
||||
) -> List[str]:
|
||||
ordered = list(dict.fromkeys(code_candidates))
|
||||
if not ordered:
|
||||
return []
|
||||
|
||||
if not preferred_code:
|
||||
return ordered
|
||||
|
||||
normalized_preferred = preferred_code.strip()
|
||||
if normalized_preferred and normalized_preferred in ordered:
|
||||
return [normalized_preferred] + [code for code in ordered if code != normalized_preferred]
|
||||
return ordered
|
||||
|
||||
@staticmethod
|
||||
def _normalize_daily_refill_code(code: Optional[str]) -> str:
|
||||
raw_code = str(code or "").strip()
|
||||
if not raw_code:
|
||||
return ""
|
||||
return canonical_stock_code(normalize_stock_code(raw_code))
|
||||
|
||||
@staticmethod
|
||||
def _ordered_daily_refill_codes(
|
||||
*,
|
||||
code_candidates: List[str],
|
||||
preferred_code: Optional[str] = None,
|
||||
) -> List[str]:
|
||||
ordered = BacktestService._ordered_candidate_codes(
|
||||
code_candidates=code_candidates,
|
||||
preferred_code=preferred_code,
|
||||
)
|
||||
refill_codes: List[str] = []
|
||||
seen: set[str] = set()
|
||||
for code in ordered:
|
||||
refill_code = BacktestService._normalize_daily_refill_code(code)
|
||||
if not refill_code or refill_code in seen:
|
||||
continue
|
||||
seen.add(refill_code)
|
||||
refill_codes.append(refill_code)
|
||||
return refill_codes
|
||||
|
||||
def _get_forward_bars_by_candidates(
|
||||
self,
|
||||
*,
|
||||
code_candidates: List[str],
|
||||
analysis_date: date,
|
||||
eval_window_days: int,
|
||||
preferred_code: Optional[str] = None,
|
||||
) -> List[Any]:
|
||||
ordered_codes = BacktestService._ordered_candidate_codes(
|
||||
code_candidates=code_candidates,
|
||||
preferred_code=preferred_code,
|
||||
)
|
||||
|
||||
if not ordered_codes:
|
||||
return []
|
||||
|
||||
best_bars: List[Any] = []
|
||||
for code in ordered_codes:
|
||||
if not code:
|
||||
continue
|
||||
|
||||
bars = self.stock_repo.get_forward_bars(
|
||||
code=code,
|
||||
analysis_date=analysis_date,
|
||||
eval_window_days=eval_window_days,
|
||||
)
|
||||
if len(bars) >= eval_window_days:
|
||||
return bars
|
||||
if len(bars) > len(best_bars):
|
||||
best_bars = bars
|
||||
return best_bars
|
||||
return BacktestService._normalize_code(code)
|
||||
|
||||
@staticmethod
|
||||
def _build_run_diagnostics(
|
||||
@@ -938,8 +845,65 @@ class BacktestService:
|
||||
logger.warning(f"无法确定分析日期,跳过记录: {analysis.code}#{getattr(analysis, 'id', '?')}")
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _resolve_expected_start_date(
|
||||
*,
|
||||
analysis,
|
||||
analysis_date: date,
|
||||
market: Optional[str],
|
||||
stock_code: Optional[str],
|
||||
) -> Optional[date]:
|
||||
phase_summary = extract_market_phase_summary(analysis.context_snapshot)
|
||||
snapshot_market = (
|
||||
str(phase_summary.get("market") or "").strip().lower()
|
||||
if isinstance(phase_summary, dict)
|
||||
else ""
|
||||
)
|
||||
if not market:
|
||||
return None
|
||||
if snapshot_market != market:
|
||||
phase_summary = rebuild_market_phase_summary_for_stock_code(
|
||||
stock_code,
|
||||
analysis.context_snapshot,
|
||||
)
|
||||
snapshot_market = (
|
||||
str(phase_summary.get("market") or "").strip().lower()
|
||||
if isinstance(phase_summary, dict)
|
||||
else ""
|
||||
)
|
||||
if snapshot_market != market:
|
||||
return None
|
||||
|
||||
effective_date_value = (
|
||||
phase_summary.get("effective_daily_bar_date")
|
||||
if isinstance(phase_summary, dict)
|
||||
else None
|
||||
)
|
||||
if effective_date_value:
|
||||
try:
|
||||
effective_date = datetime.strptime(
|
||||
str(effective_date_value),
|
||||
"%Y-%m-%d",
|
||||
).date()
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
if effective_date <= analysis_date:
|
||||
return effective_date
|
||||
return None
|
||||
|
||||
phase = (
|
||||
phase_summary.get("phase")
|
||||
if isinstance(phase_summary, dict)
|
||||
else None
|
||||
)
|
||||
return resolve_historical_daily_bar_date(
|
||||
market,
|
||||
analysis_date,
|
||||
phase,
|
||||
)
|
||||
|
||||
def _try_fill_daily_data(self, *, code: str, analysis_date: date, eval_window_days: int) -> None:
|
||||
refill_code = self._normalize_daily_refill_code(code)
|
||||
refill_code = str(code or "").strip()
|
||||
if not refill_code:
|
||||
return
|
||||
|
||||
|
||||
@@ -6,10 +6,16 @@ Shared stock code utilities.
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Optional
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
from data_provider.base import canonical_stock_code, is_bse_code
|
||||
from src.services.market_symbol_utils import normalize_suffix_market_symbol
|
||||
from data_provider.us_index_mapping import is_us_index_code
|
||||
from src.services.market_symbol_utils import (
|
||||
get_suffix_market,
|
||||
normalize_suffix_market_symbol,
|
||||
suffix_base_lookup_allowed,
|
||||
)
|
||||
|
||||
|
||||
# Known exchange prefixes (case-insensitive) and the digit lengths they accept.
|
||||
@@ -40,6 +46,43 @@ _SUFFIX_DIGIT_LENS: dict = {
|
||||
_PRESERVE_SUFFIXES = {".T", ".KS", ".KQ", ".TW", ".TWO"}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DailyStockIdentity:
|
||||
"""One parsed identity shared by daily-bar lookup, calendar, and refill."""
|
||||
|
||||
normalized_code: str
|
||||
market: str
|
||||
refill_code: str
|
||||
code_candidates: tuple[str, ...]
|
||||
|
||||
|
||||
def _filter_cross_market_numeric_aliases(
|
||||
*,
|
||||
raw_code: str,
|
||||
market: str,
|
||||
candidates: List[str],
|
||||
) -> tuple[str, ...]:
|
||||
"""Drop only derived numeric aliases known to collide across markets."""
|
||||
from src.core.trading_calendar import get_market_for_stock
|
||||
from src.data.stock_index_loader import resolve_index_stock_code_candidates
|
||||
|
||||
filtered: List[str] = []
|
||||
for candidate in dict.fromkeys(value for value in candidates if value):
|
||||
if candidate == raw_code or not candidate.isdigit():
|
||||
filtered.append(candidate)
|
||||
continue
|
||||
|
||||
indexed_markets = {
|
||||
indexed_market
|
||||
for indexed_code in resolve_index_stock_code_candidates(candidate)
|
||||
if (indexed_market := get_market_for_stock(indexed_code)) is not None
|
||||
}
|
||||
if indexed_markets and indexed_markets != {market}:
|
||||
continue
|
||||
filtered.append(candidate)
|
||||
return tuple(filtered)
|
||||
|
||||
|
||||
def _infer_cn_exchange(base: str) -> str:
|
||||
"""Infer CN exchange from a 6-digit A/B-share code."""
|
||||
if not (base.isdigit() and len(base) == 6):
|
||||
@@ -64,30 +107,38 @@ def _valid_exchange_code(exchange: str, base: str, digit_lens: tuple[int, ...])
|
||||
return True
|
||||
|
||||
|
||||
def _strip_exchange_prefix(text: str) -> Optional[str]:
|
||||
"""Strip leading exchange prefix (SH/SZ/HK etc.) and return the bare digits, or None."""
|
||||
def _split_explicit_exchange(
|
||||
text: str,
|
||||
) -> Optional[tuple[str, str, tuple[int, ...]]]:
|
||||
"""Return one recognized explicit exchange and its unvalidated base."""
|
||||
for suffix, digit_lens in _SUFFIX_DIGIT_LENS.items():
|
||||
if text.endswith(suffix):
|
||||
base = text[: -len(suffix)].strip()
|
||||
return suffix.lstrip("."), base, digit_lens
|
||||
|
||||
for prefix, digit_lens in _PREFIX_DIGIT_LENS.items():
|
||||
dotted_prefix = f"{prefix}."
|
||||
if text.startswith(dotted_prefix):
|
||||
base = text[len(dotted_prefix):]
|
||||
if _valid_exchange_code(prefix, base, digit_lens):
|
||||
return base.zfill(5) if prefix == "HK" else base
|
||||
return prefix, base, digit_lens
|
||||
if text.startswith(prefix):
|
||||
base = text[len(prefix):]
|
||||
if _valid_exchange_code(prefix, base, digit_lens):
|
||||
return base.zfill(5) if prefix == "HK" else base
|
||||
# Do not mistake US tickers such as SHOP/HKEX for exchange prefixes.
|
||||
if base.isdigit():
|
||||
return prefix, base, digit_lens
|
||||
return None
|
||||
|
||||
|
||||
def _strip_exchange_suffix(text: str) -> Optional[str]:
|
||||
"""Strip exchange suffix (.SH/.SZ/.SS/.HK) and return normalized bare digits, or None."""
|
||||
for suffix, digit_lens in _SUFFIX_DIGIT_LENS.items():
|
||||
if text.endswith(suffix):
|
||||
base = text[: -len(suffix)].strip()
|
||||
exchange = suffix.lstrip(".")
|
||||
if _valid_exchange_code(exchange, base, digit_lens):
|
||||
return base.zfill(5) if suffix == ".HK" else base
|
||||
return None
|
||||
def _normalize_explicit_exchange_parts(
|
||||
parts: Optional[tuple[str, str, tuple[int, ...]]],
|
||||
) -> Optional[str]:
|
||||
"""Return the normalized base from one previously parsed exchange."""
|
||||
if parts is None:
|
||||
return None
|
||||
exchange, base, digit_lens = parts
|
||||
if not _valid_exchange_code(exchange, base, digit_lens):
|
||||
return None
|
||||
return base.zfill(5) if exchange == "HK" else base
|
||||
|
||||
|
||||
def is_code_like(value: str) -> bool:
|
||||
@@ -97,13 +148,11 @@ def is_code_like(value: str) -> bool:
|
||||
return False
|
||||
if text.isdigit() and len(text) in (5, 6):
|
||||
return True
|
||||
if _strip_exchange_suffix(text) is not None:
|
||||
return True
|
||||
explicit_parts = _split_explicit_exchange(text)
|
||||
if explicit_parts is not None:
|
||||
return _normalize_explicit_exchange_parts(explicit_parts) is not None
|
||||
if re.match(r"^[A-Z]{1,5}(?:\.(?:US|[A-Z]))?$", text):
|
||||
return True
|
||||
# Support exchange-prefixed codes: SH600519, SZ000001, BJ920493, HK00700
|
||||
if _strip_exchange_prefix(text) is not None:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
@@ -116,26 +165,259 @@ def normalize_code(raw: str) -> Optional[str]:
|
||||
- Prefix format: SH600519, SH.600519, SZ000001, BJ920493, HK00700 (case-insensitive)
|
||||
- US ticker symbols: AAPL, TSLA
|
||||
"""
|
||||
normalized, _ = _normalize_code_and_exchange(raw)
|
||||
return normalized
|
||||
|
||||
|
||||
def _normalize_code_and_exchange(raw: str) -> tuple[Optional[str], str]:
|
||||
"""Normalize once and retain an explicit exchange for candidate expansion."""
|
||||
text = raw.strip().upper()
|
||||
if not text:
|
||||
return None
|
||||
return None, ""
|
||||
if text.isdigit() and len(text) in (5, 6):
|
||||
return text
|
||||
return text, ""
|
||||
explicit_parts = _split_explicit_exchange(text)
|
||||
explicit_exchange = explicit_parts[0] if explicit_parts is not None else ""
|
||||
explicit_code = _normalize_explicit_exchange_parts(explicit_parts)
|
||||
if explicit_parts is not None and explicit_code is None:
|
||||
return None, explicit_exchange
|
||||
suffix_symbol = normalize_suffix_market_symbol(text)
|
||||
if suffix_symbol is not None:
|
||||
return suffix_symbol
|
||||
return suffix_symbol, explicit_exchange
|
||||
if any(text.endswith(suffix) for suffix in _PRESERVE_SUFFIXES):
|
||||
return None
|
||||
return None, explicit_exchange
|
||||
if re.match(r"^[A-Z]{1,5}(?:\.(?:US|[A-Z]))?$", text):
|
||||
return text
|
||||
stripped_suffix = _strip_exchange_suffix(text)
|
||||
if stripped_suffix is not None:
|
||||
return stripped_suffix
|
||||
# Support exchange-prefixed codes: SH600519 -> 600519, BJ920493 -> 920493
|
||||
stripped = _strip_exchange_prefix(text)
|
||||
if stripped is not None:
|
||||
return stripped
|
||||
return None
|
||||
return text, explicit_exchange
|
||||
if explicit_code is not None:
|
||||
return explicit_code, explicit_exchange
|
||||
return None, explicit_exchange
|
||||
|
||||
|
||||
def _build_hk_market_variants(hk_digits: str) -> List[str]:
|
||||
"""Build normalized HK variants for padded and legacy code shapes."""
|
||||
if not hk_digits.isdigit() or not hk_digits:
|
||||
return []
|
||||
|
||||
padded = hk_digits.zfill(5)
|
||||
unpadded = padded.lstrip("0") or "0"
|
||||
variants = [
|
||||
f"HK{padded}",
|
||||
f"{padded}.HK",
|
||||
padded,
|
||||
f"HK{unpadded}",
|
||||
f"{unpadded}.HK",
|
||||
f"HK.{padded}",
|
||||
]
|
||||
if unpadded == padded:
|
||||
variants.pop(3)
|
||||
variants.pop(3)
|
||||
if len(unpadded) <= 4 and unpadded != padded:
|
||||
variants.extend([unpadded, f"HK.{unpadded}"])
|
||||
return variants
|
||||
|
||||
|
||||
def _build_market_code_variants(
|
||||
raw_code: str,
|
||||
normalized_code: str,
|
||||
explicit_exchange: str,
|
||||
) -> List[str]:
|
||||
"""Return additional market-formatted variants for stored-code matching."""
|
||||
variants: List[str] = []
|
||||
if not raw_code:
|
||||
return variants
|
||||
|
||||
raw_code_upper = raw_code.upper()
|
||||
normalized_upper = normalized_code.upper() if normalized_code else ""
|
||||
|
||||
def _add_us_variants(code: str) -> None:
|
||||
if not code:
|
||||
return
|
||||
if code.endswith(".US"):
|
||||
bare = code[:-3]
|
||||
if bare.isalpha() and 1 <= len(bare) <= 5:
|
||||
variants.append(bare)
|
||||
return
|
||||
if "." not in code and code.isalpha() and 1 <= len(code) <= 5:
|
||||
variants.append(f"{code}.US")
|
||||
|
||||
_add_us_variants(raw_code_upper)
|
||||
if normalized_upper != raw_code_upper:
|
||||
_add_us_variants(normalized_upper)
|
||||
|
||||
if normalized_upper.isdigit() and len(normalized_upper) == 6:
|
||||
if explicit_exchange in {"SH", "SS"}:
|
||||
exchange = "SH"
|
||||
elif explicit_exchange == "SZ":
|
||||
exchange = "SZ"
|
||||
elif explicit_exchange == "BJ" or is_bse_code(normalized_upper):
|
||||
exchange = "BJ"
|
||||
elif normalized_upper.startswith(("5", "6", "9")):
|
||||
exchange = "SH"
|
||||
else:
|
||||
exchange = "SZ"
|
||||
|
||||
variants.extend(
|
||||
[
|
||||
f"{exchange}{normalized_upper}",
|
||||
f"{normalized_upper}.{exchange}",
|
||||
f"{exchange}.{normalized_upper}",
|
||||
]
|
||||
)
|
||||
if exchange == "SH":
|
||||
variants.extend(
|
||||
[
|
||||
f"SS{normalized_upper}",
|
||||
f"{normalized_upper}.SS",
|
||||
f"SS.{normalized_upper}",
|
||||
]
|
||||
)
|
||||
|
||||
if explicit_exchange == "HK" and normalized_upper.isdigit():
|
||||
variants.extend(_build_hk_market_variants(normalized_upper))
|
||||
elif normalized_upper.startswith("HK") and normalized_upper[2:].isdigit() and len(normalized_upper[2:]) <= 5:
|
||||
variants.extend(_build_hk_market_variants(normalized_upper[2:]))
|
||||
if raw_code_upper.isdigit() and len(raw_code_upper) in (4, 5):
|
||||
variants.extend(_build_hk_market_variants(raw_code_upper))
|
||||
|
||||
return variants
|
||||
|
||||
|
||||
def resolve_daily_stock_identity(
|
||||
code: Optional[str],
|
||||
*,
|
||||
market_hint: Optional[str] = None,
|
||||
) -> Optional[DailyStockIdentity]:
|
||||
"""Parse one stock identity for every local daily-bar consumer.
|
||||
|
||||
Persisted market metadata and the stock index may disambiguate legacy bare
|
||||
JP/KR codes before numeric CN/HK defaults are applied.
|
||||
"""
|
||||
raw_code = str(code or "").strip().upper()
|
||||
if not raw_code:
|
||||
return None
|
||||
|
||||
identity_code = raw_code
|
||||
trusted_market = str(market_hint or "").strip().lower()
|
||||
if raw_code.isdigit() and len(raw_code) in {4, 5, 6}:
|
||||
from src.data.stock_index_loader import resolve_index_stock_code_candidates
|
||||
|
||||
indexed_candidates = resolve_index_stock_code_candidates(raw_code)
|
||||
indexed_identities = [
|
||||
(candidate, get_suffix_market(candidate))
|
||||
for candidate in indexed_candidates
|
||||
]
|
||||
indexed_offshore = [
|
||||
(candidate, market)
|
||||
for candidate, market in indexed_identities
|
||||
if market in {"jp", "kr"}
|
||||
]
|
||||
if trusted_market in {"jp", "kr"}:
|
||||
matching_candidates = [
|
||||
candidate
|
||||
for candidate, market in indexed_offshore
|
||||
if market == trusted_market
|
||||
]
|
||||
if len(matching_candidates) == 1:
|
||||
identity_code = matching_candidates[0]
|
||||
elif indexed_candidates:
|
||||
return None
|
||||
elif trusted_market == "jp" and len(raw_code) in {4, 5}:
|
||||
identity_code = f"{raw_code}.T"
|
||||
elif trusted_market == "kr" and len(raw_code) == 6:
|
||||
return DailyStockIdentity(
|
||||
normalized_code=raw_code,
|
||||
market="kr",
|
||||
refill_code="",
|
||||
code_candidates=(raw_code,),
|
||||
)
|
||||
else:
|
||||
return None
|
||||
elif trusted_market == "cn":
|
||||
if len(raw_code) == 6:
|
||||
pass
|
||||
elif len(indexed_candidates) == 1 and len(indexed_offshore) == 1:
|
||||
identity_code = indexed_offshore[0][0]
|
||||
else:
|
||||
return None
|
||||
elif trusted_market == "hk":
|
||||
if len(raw_code) not in {4, 5}:
|
||||
return None
|
||||
elif trusted_market:
|
||||
return None
|
||||
elif len(indexed_candidates) > 1:
|
||||
return None
|
||||
elif len(indexed_offshore) == 1:
|
||||
identity_code = indexed_offshore[0][0]
|
||||
|
||||
if is_us_index_code(identity_code):
|
||||
normalized_code, explicit_exchange = identity_code, ""
|
||||
elif identity_code.isdigit() and len(identity_code) == 4:
|
||||
normalized_code, explicit_exchange = identity_code.zfill(5), "HK"
|
||||
else:
|
||||
normalized_code, explicit_exchange = _normalize_code_and_exchange(identity_code)
|
||||
if normalized_code is None:
|
||||
return None
|
||||
|
||||
suffix_market = get_suffix_market(normalized_code)
|
||||
if explicit_exchange in {"SH", "SS", "SZ", "BJ"}:
|
||||
market = "cn"
|
||||
elif explicit_exchange == "HK":
|
||||
market = "hk"
|
||||
elif suffix_market:
|
||||
market = suffix_market
|
||||
elif is_us_index_code(normalized_code):
|
||||
market = "us"
|
||||
elif re.fullmatch(r"[A-Z]{1,5}(?:\.(?:US|[A-Z]))?", normalized_code):
|
||||
market = "us"
|
||||
elif normalized_code.isdigit() and len(normalized_code) == 6:
|
||||
market = "cn"
|
||||
elif normalized_code.isdigit() and len(normalized_code) == 5:
|
||||
market = "hk"
|
||||
else:
|
||||
return None
|
||||
|
||||
if market == "hk":
|
||||
normalized_code = normalized_code.zfill(5)
|
||||
refill_code = f"HK{normalized_code}"
|
||||
elif market == "us":
|
||||
normalized_code = normalized_code.removesuffix(".US")
|
||||
refill_code = normalized_code
|
||||
else:
|
||||
refill_code = normalized_code
|
||||
|
||||
if market == "hk":
|
||||
candidates = [raw_code]
|
||||
candidates.extend(_build_hk_market_variants(normalized_code))
|
||||
else:
|
||||
candidates = [raw_code, normalized_code, refill_code]
|
||||
if suffix_base_lookup_allowed(normalized_code):
|
||||
candidates.append(normalized_code.rsplit(".", 1)[0])
|
||||
if market not in {"jp", "kr", "tw"}:
|
||||
for candidate in list(candidates):
|
||||
candidates.extend(
|
||||
_build_market_code_variants(
|
||||
raw_code,
|
||||
candidate,
|
||||
explicit_exchange,
|
||||
)
|
||||
)
|
||||
unique_candidates = _filter_cross_market_numeric_aliases(
|
||||
raw_code=raw_code,
|
||||
market=market,
|
||||
candidates=candidates,
|
||||
)
|
||||
return DailyStockIdentity(
|
||||
normalized_code=normalized_code,
|
||||
market=market,
|
||||
refill_code=refill_code,
|
||||
code_candidates=unique_candidates,
|
||||
)
|
||||
|
||||
|
||||
def build_daily_code_candidates(code: Optional[str]) -> List[str]:
|
||||
"""Build ordered code variants used to locate locally stored daily bars."""
|
||||
identity = resolve_daily_stock_identity(code)
|
||||
return list(identity.code_candidates) if identity is not None else []
|
||||
|
||||
|
||||
def resolve_index_stock_code_for_analysis(raw: str) -> str:
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Resolve one coherent local daily-bar window across equivalent stock codes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import date
|
||||
from typing import List, Optional, Sequence, Tuple
|
||||
|
||||
from src.repositories.stock_repo import StockRepository
|
||||
from src.storage import StockDaily
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StockDailyWindow:
|
||||
"""A start bar and its forward bars from one stored stock-code shape."""
|
||||
|
||||
code: str
|
||||
start_bar: StockDaily
|
||||
forward_bars: List[StockDaily]
|
||||
|
||||
|
||||
def resolve_stock_daily_window(
|
||||
*,
|
||||
stock_repo: StockRepository,
|
||||
code_candidates: Sequence[str],
|
||||
expected_start_date: date,
|
||||
eval_window_days: int,
|
||||
) -> Optional[StockDailyWindow]:
|
||||
"""Choose one coherent window anchored to the expected trading session.
|
||||
|
||||
Only candidates with a bar on the authoritative expected date are eligible.
|
||||
Complete windows outrank partial ones; remaining ties prefer more forward
|
||||
bars and then candidate order. Start and forward bars are never combined
|
||||
across code shapes.
|
||||
"""
|
||||
best_window: Optional[StockDailyWindow] = None
|
||||
best_key: Optional[Tuple[bool, int, int]] = None
|
||||
if isinstance(eval_window_days, bool) or not isinstance(eval_window_days, int):
|
||||
raise ValueError("eval_window_days must be a positive integer")
|
||||
required_bars = eval_window_days
|
||||
if required_bars <= 0:
|
||||
raise ValueError("eval_window_days must be a positive integer")
|
||||
|
||||
for rank, code in enumerate(dict.fromkeys(code_candidates)):
|
||||
if not code:
|
||||
continue
|
||||
start_bar = stock_repo.get_daily_on_date(
|
||||
code=code,
|
||||
target_date=expected_start_date,
|
||||
)
|
||||
if start_bar is None or start_bar.close is None:
|
||||
continue
|
||||
|
||||
forward_bars = stock_repo.get_forward_bars(
|
||||
code=code,
|
||||
analysis_date=start_bar.date,
|
||||
eval_window_days=required_bars,
|
||||
)
|
||||
key = (
|
||||
len(forward_bars) >= required_bars,
|
||||
len(forward_bars),
|
||||
-rank,
|
||||
)
|
||||
if best_key is None or key > best_key:
|
||||
best_key = key
|
||||
best_window = StockDailyWindow(
|
||||
code=code,
|
||||
start_bar=start_bar,
|
||||
forward_bars=forward_bars,
|
||||
)
|
||||
|
||||
return best_window
|
||||
+804
-151
File diff suppressed because it is too large
Load Diff
@@ -9,12 +9,123 @@ import pytest
|
||||
from unittest.mock import patch
|
||||
|
||||
from src.services.stock_code_utils import (
|
||||
build_daily_code_candidates,
|
||||
is_code_like,
|
||||
normalize_code,
|
||||
resolve_daily_stock_identity,
|
||||
resolve_index_stock_code_for_analysis,
|
||||
)
|
||||
|
||||
|
||||
class TestBuildDailyCodeCandidates:
|
||||
@pytest.mark.parametrize(
|
||||
"code",
|
||||
["600519.SZ", "000001.SH", "920748.SH", "SH920748", "600519.HK"],
|
||||
)
|
||||
def test_rejects_conflicting_explicit_exchange_before_any_candidate(self, code):
|
||||
assert build_daily_code_candidates(code) == []
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("code", "required_candidates"),
|
||||
[
|
||||
(
|
||||
"600519.SH",
|
||||
{
|
||||
"600519.SH",
|
||||
"600519",
|
||||
"SH600519",
|
||||
"SH.600519",
|
||||
"SS600519",
|
||||
},
|
||||
),
|
||||
("600519", {"600519", "600519.SH"}),
|
||||
("000001.SZ", {"000001.SZ", "000001"}),
|
||||
("920748", {"920748", "BJ920748", "920748.BJ"}),
|
||||
("1810", {"1810", "01810", "HK01810", "01810.HK"}),
|
||||
("01810", {"1810", "01810", "HK01810", "01810.HK"}),
|
||||
("1810.HK", {"1810", "1810.HK", "01810", "HK01810", "01810.HK"}),
|
||||
("HK.01810", {"1810", "HK.01810", "01810", "HK01810", "01810.HK"}),
|
||||
("AAPL", {"AAPL", "AAPL.US"}),
|
||||
("AAPL.US", {"AAPL.US", "AAPL"}),
|
||||
("NASDAQ", {"NASDAQ"}),
|
||||
("^GSPC", {"^GSPC"}),
|
||||
],
|
||||
)
|
||||
def test_preserves_valid_explicit_and_legacy_bare_codes(
|
||||
self,
|
||||
code,
|
||||
required_candidates,
|
||||
):
|
||||
assert set(build_daily_code_candidates(code)) >= required_candidates
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("code", "normalized_code", "market", "refill_code"),
|
||||
[
|
||||
("600519.SH", "600519", "cn", "600519"),
|
||||
("1810", "01810", "hk", "HK01810"),
|
||||
("HK.01810", "01810", "hk", "HK01810"),
|
||||
("AAPL.US", "AAPL", "us", "AAPL"),
|
||||
("BRK.B", "BRK.B", "us", "BRK.B"),
|
||||
("NASDAQ", "NASDAQ", "us", "NASDAQ"),
|
||||
("^GSPC", "^GSPC", "us", "^GSPC"),
|
||||
("7203.T", "7203.T", "jp", "7203.T"),
|
||||
],
|
||||
)
|
||||
def test_one_identity_drives_candidates_market_and_refill(
|
||||
self,
|
||||
code,
|
||||
normalized_code,
|
||||
market,
|
||||
refill_code,
|
||||
):
|
||||
identity = resolve_daily_stock_identity(code)
|
||||
|
||||
assert identity is not None
|
||||
assert identity.normalized_code == normalized_code
|
||||
assert identity.market == market
|
||||
assert identity.refill_code == refill_code
|
||||
assert code in identity.code_candidates
|
||||
assert refill_code in identity.code_candidates
|
||||
|
||||
def test_kr_suffix_adds_only_its_legacy_bare_candidate(self):
|
||||
identity = resolve_daily_stock_identity("005930.KS")
|
||||
|
||||
assert identity is not None
|
||||
assert identity.market == "kr"
|
||||
assert identity.code_candidates == ("005930.KS", "005930")
|
||||
|
||||
def test_bare_code_with_unsupported_market_hint_fails_closed(self):
|
||||
assert resolve_daily_stock_identity("005930", market_hint="tw") is None
|
||||
|
||||
def test_cross_market_bare_code_without_hint_fails_closed(self):
|
||||
assert resolve_daily_stock_identity("8035") is None
|
||||
|
||||
def test_cross_market_suffix_code_does_not_add_ambiguous_bare_alias(self):
|
||||
identity = resolve_daily_stock_identity("8035.T")
|
||||
|
||||
assert identity is not None
|
||||
assert identity.market == "jp"
|
||||
assert identity.code_candidates == ("8035.T",)
|
||||
|
||||
def test_cross_market_hk_code_does_not_add_ambiguous_bare_alias(self):
|
||||
identity = resolve_daily_stock_identity("08035.HK")
|
||||
|
||||
assert identity is not None
|
||||
assert identity.market == "hk"
|
||||
assert "8035" not in identity.code_candidates
|
||||
assert "08035.HK" in identity.code_candidates
|
||||
assert "8035.HK" in identity.code_candidates
|
||||
|
||||
def test_trusted_legacy_jp_identity_keeps_its_raw_bare_code(self):
|
||||
identity = resolve_daily_stock_identity("8035", market_hint="jp")
|
||||
|
||||
assert identity is not None
|
||||
assert identity.market == "jp"
|
||||
assert identity.code_candidates[0] == "8035"
|
||||
assert "8035" in identity.code_candidates
|
||||
assert "8035.T" in identity.code_candidates
|
||||
|
||||
|
||||
class TestIsCodeLike:
|
||||
# --- Plain digit codes ---
|
||||
def test_plain_6_digit(self):
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Direct contract tests for coherent local daily-window resolution."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.stock_daily_window_resolver import resolve_stock_daily_window
|
||||
|
||||
|
||||
def _bar(day: date, close: float = 100.0):
|
||||
return SimpleNamespace(date=day, close=close)
|
||||
|
||||
|
||||
class _FakeStockRepository:
|
||||
def __init__(self, starts, forwards):
|
||||
self.starts = starts
|
||||
self.forwards = forwards
|
||||
self.selected_start_dates = {}
|
||||
|
||||
def get_daily_on_date(self, *, code, target_date):
|
||||
configured = self.starts.get(code)
|
||||
if configured is None:
|
||||
return None
|
||||
options = configured if isinstance(configured, list) else [configured]
|
||||
matching = [start for start in options if start.date == target_date]
|
||||
if not matching:
|
||||
return None
|
||||
start = matching[0]
|
||||
self.selected_start_dates[code] = start.date
|
||||
return start
|
||||
|
||||
def get_forward_bars(self, *, code, analysis_date, eval_window_days):
|
||||
assert self.selected_start_dates[code] == analysis_date
|
||||
return list(self.forwards.get(code, ()))[:eval_window_days]
|
||||
|
||||
|
||||
def _resolve(
|
||||
starts,
|
||||
forwards,
|
||||
candidates=("first", "second"),
|
||||
days=1,
|
||||
expected_start_date=date(2024, 1, 5),
|
||||
):
|
||||
return resolve_stock_daily_window(
|
||||
stock_repo=_FakeStockRepository(starts, forwards),
|
||||
code_candidates=candidates,
|
||||
expected_start_date=expected_start_date,
|
||||
eval_window_days=days,
|
||||
)
|
||||
|
||||
|
||||
def test_candidates_without_exact_start_return_none() -> None:
|
||||
window = _resolve(
|
||||
starts={
|
||||
"first": _bar(date(2020, 1, 2), 50.0),
|
||||
"second": _bar(date(2021, 1, 4), 60.0),
|
||||
},
|
||||
forwards={
|
||||
"first": [_bar(date(2024, 1, 8), 55.0)],
|
||||
"second": [_bar(date(2024, 1, 8), 65.0)],
|
||||
},
|
||||
)
|
||||
|
||||
assert window is None
|
||||
|
||||
|
||||
def test_same_date_complete_window_outranks_partial_window() -> None:
|
||||
window = _resolve(
|
||||
starts={
|
||||
"first": _bar(date(2024, 1, 5)),
|
||||
"second": _bar(date(2024, 1, 5)),
|
||||
},
|
||||
forwards={
|
||||
"first": [],
|
||||
"second": [_bar(date(2024, 1, 8))],
|
||||
},
|
||||
)
|
||||
|
||||
assert window.code == "second"
|
||||
|
||||
|
||||
def test_same_date_tie_preserves_candidate_order() -> None:
|
||||
window = _resolve(
|
||||
starts={
|
||||
"first": _bar(date(2024, 1, 5)),
|
||||
"second": _bar(date(2024, 1, 5)),
|
||||
},
|
||||
forwards={
|
||||
"first": [_bar(date(2024, 1, 8))],
|
||||
"second": [_bar(date(2024, 1, 8))],
|
||||
},
|
||||
)
|
||||
|
||||
assert window.code == "first"
|
||||
|
||||
|
||||
def test_partial_fallback_uses_more_bars_for_same_start_date() -> None:
|
||||
window = _resolve(
|
||||
starts={
|
||||
"first": _bar(date(2024, 1, 5)),
|
||||
"second": _bar(date(2024, 1, 5)),
|
||||
},
|
||||
forwards={
|
||||
"first": [_bar(date(2024, 1, 8))],
|
||||
"second": [
|
||||
_bar(date(2024, 1, 8)),
|
||||
_bar(date(2024, 1, 9)),
|
||||
],
|
||||
},
|
||||
days=3,
|
||||
)
|
||||
|
||||
assert window.code == "second"
|
||||
assert len(window.forward_bars) == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("days", [0, -1, 1.5, True, "1", "invalid"])
|
||||
def test_invalid_window_length_fails_closed(days) -> None:
|
||||
with pytest.raises(ValueError, match="positive integer"):
|
||||
_resolve(
|
||||
starts={"first": _bar(date(2024, 1, 5))},
|
||||
forwards={"first": []},
|
||||
candidates=("first",),
|
||||
days=days,
|
||||
)
|
||||
@@ -224,6 +224,31 @@ class TestStockIndexLoader(unittest.TestCase):
|
||||
self.assertEqual(stock_index_loader.resolve_index_stock_code("005930"), "005930.KS")
|
||||
self.assertEqual(stock_index_loader.resolve_index_stock_code("7203"), "7203.T")
|
||||
|
||||
def test_resolve_index_stock_code_rejects_cross_market_bare_alias(self):
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
bundled_path = Path(temp_dir) / "stocks.index.json"
|
||||
bundled_path.write_text(
|
||||
json.dumps(
|
||||
[
|
||||
["08035.HK", "08035", "HK 8035", "hk8035", "hk", [], "HK", "stock", True, 100],
|
||||
["8035.T", "8035.T", "JP 8035", "jp8035", "jp", [], "JP", "stock", True, 100],
|
||||
],
|
||||
ensure_ascii=False,
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
stock_index_loader,
|
||||
"get_remote_stock_index_cache_path",
|
||||
return_value=Path(temp_dir) / "missing.json",
|
||||
), patch.object(
|
||||
stock_index_loader,
|
||||
"get_stock_index_candidate_paths",
|
||||
return_value=(bundled_path,),
|
||||
):
|
||||
self.assertIsNone(stock_index_loader.resolve_index_stock_code("8035"))
|
||||
|
||||
def test_resolve_index_stock_code_reuses_cached_lookup(self):
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
bundled_path = Path(temp_dir) / "stocks.index.json"
|
||||
|
||||
@@ -145,6 +145,56 @@ class _CloseTimeCalendar(_FakeCalendar):
|
||||
return pd.Timestamp(local_close).tz_convert("UTC")
|
||||
|
||||
|
||||
class HistoricalDailyBarDateTestCase(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.calendar = _FakeCalendar(
|
||||
sessions=[date(2024, 1, 5), date(2024, 1, 8)],
|
||||
close_hour=15,
|
||||
tz_name="Asia/Shanghai",
|
||||
)
|
||||
|
||||
def _resolve(self, target_date: date, phase: Optional[str]) -> Optional[date]:
|
||||
with patch.object(trading_calendar, "_XCALS_AVAILABLE", True), patch.object(
|
||||
trading_calendar,
|
||||
"xcals",
|
||||
_calendar_namespace(self.calendar),
|
||||
create=True,
|
||||
):
|
||||
return trading_calendar.resolve_historical_daily_bar_date(
|
||||
"cn",
|
||||
target_date,
|
||||
phase,
|
||||
)
|
||||
|
||||
def test_open_session_phase_uses_previous_session(self):
|
||||
for phase in (
|
||||
"premarket",
|
||||
"intraday",
|
||||
"lunch_break",
|
||||
"closing_auction",
|
||||
):
|
||||
with self.subTest(phase=phase):
|
||||
self.assertEqual(
|
||||
self._resolve(date(2024, 1, 8), phase),
|
||||
date(2024, 1, 5),
|
||||
)
|
||||
|
||||
def test_postmarket_uses_current_session(self):
|
||||
self.assertEqual(
|
||||
self._resolve(date(2024, 1, 8), "postmarket"),
|
||||
date(2024, 1, 8),
|
||||
)
|
||||
|
||||
def test_non_session_and_unprovable_phase_fail_closed(self):
|
||||
self.assertEqual(
|
||||
self._resolve(date(2024, 1, 7), "non_trading"),
|
||||
date(2024, 1, 5),
|
||||
)
|
||||
for phase in (None, "unknown", "postmarket"):
|
||||
with self.subTest(phase=phase):
|
||||
self.assertIsNone(self._resolve(date(2024, 1, 7), phase))
|
||||
|
||||
|
||||
class EffectiveTradingDateTestCase(unittest.TestCase):
|
||||
def test_weekend_returns_previous_session(self):
|
||||
fake_calendar = _FakeCalendar(
|
||||
|
||||
Reference in New Issue
Block a user