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:
ObVious55
2026-07-27 21:40:38 +08:00
committed by GitHub
parent 905c339d80
commit f4d9956c52
12 changed files with 1790 additions and 536 deletions
+1
View File
@@ -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)-->
<!-- 每条独立一行追加到本段末尾,无需分类标题,合并时冲突最小 -->
+54
View File
@@ -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.
+87 -54
View File
@@ -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
+8 -129
View File
@@ -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
View File
@@ -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
+317 -35
View File
@@ -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
File diff suppressed because it is too large Load Diff
+111
View File
@@ -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):
+129
View File
@@ -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,
)
+25
View File
@@ -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"
+50
View File
@@ -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(