fix: 修复断点续传按自然日判断数据存在性的逻辑 (#880) (#900)

* fix(issue-880): [bug]-修复断点续传按自然日判断数据存在性的逻辑

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix: remove transient pr draft from pr-900

* fix(review-feedback-900): address latest review comments

* fix(review-feedback-900): address latest review comments

* fix(issue-880): [bug]-修复断点续传按自然日判断数据存在性的逻辑

* fix(review-feedback-900): address latest review comments

---------

Co-authored-by: AutoCode Bot <autocode@example.com>
This commit is contained in:
mumu
2026-04-01 11:22:17 +08:00
committed by GitHub
co-authored by AutoCode Bot
parent 4c3a48fcbf
commit 87d60e9daa
11 changed files with 504 additions and 27 deletions
+1
View File
@@ -241,6 +241,7 @@
>
> - **环境变量方式**:在 `.env` 或 GitHub Secrets 中设置,影响所有运行方式(定时触发、手动触发、本地运行)
> - **UI 勾选方式**:仅在 GitHub Actions 手动触发时可见,不影响定时任务,适合临时需求
> - **断点续传与 `--dry-run` 的数据存在性判断**:会按股票所属市场的本地时区和交易日历解析“最新可复用交易日”;周末/节假日复用最近交易日,交易日盘中复用上一已完成交易日,盘后若当日数据已落库则可直接跳过。详细规则见 [完整指南](docs/full-guide.md)。
### 方式二:本地运行 / Docker 部署
+1 -1
View File
@@ -19,9 +19,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/).
- [测试] 🧪 **补充前端变更验证命令** — 对应前端资源变更同步执行 `cd apps/dsa-web && npm ci && npm run lint && npm run build`,作为版本信息展示与 Docker 重建生效验证的最小验证闭环记录。
- [修复] 内置定时调度器现在会在运行中感知 WebUI 保存后的 `SCHEDULE_TIME` 变化,并在下一轮检查时重绑 daily job,避免 `python main.py --serve --schedule` 仍固定按启动时的 `18:00` 触发;`.env.example` 也同步删除了重复的定时任务配置示例。
- [修复] 🪟 **Windows Release 渠道编辑器保留 MiniMax 模型前缀** — 渠道模式下填写 `minimax/<模型名>` 时,后端归一化与 Web 设置页运行时模型列表都会保留该值原样,不再误改写成 `openai/minimax/<模型名>`,从而恢复 MiniMax 模型在 Win 客户端里的保存、选择与使用。
- [修复] 🗓️ **断点续传与 `--dry-run` 改按市场时区和交易日历判断可复用数据**(fixes #880)— 股票数据存在性检查不再直接使用服务器自然日,而是按 A 股 / 港股 / 美股各自市场时区解析“最新可复用交易日”;周末、节假日、跨时区以及盘中 / 盘后场景会统一复用最近已完成交易日,避免误判导致重复抓取或错误跳过。
- [修复] 🐳 **Docker WebUI 运行时优先复用预构建静态资源** — `prepare_webui_frontend_assets()` 现在会先检查镜像内已有的 `static/index.html` 是否可直接复用;当容器运行时不包含 `apps/dsa-web` 源码目录且未安装 `npm` 时,也不会误报“未找到前端项目,无法自动构建”,从而恢复 Docker 部署后的 WebUI 打开能力。
- [修复] 📨 **单股推送模式不再并发复用共享通知实例** — `StockAnalysisPipeline.run()` 现在会保留个股分析并发,但把 `SINGLE_STOCK_NOTIFY=true` 下的即时通知挪到结果收集侧串行发送;同时 `_send_single_stock_notification()` 为同一个 pipeline 实例补上实例级临界区,避免直接调用 `process_single_stock(..., single_stock_notify=True)` 时多个线程继续共享同一个 `NotificationService` 进入报告生成与发送链路,导致通知乱序、重复发送或状态污染。
- [修复] 📨 **单股推送模式不再并发复用共享通知实例** — `StockAnalysisPipeline.run()` 现在会保留个股分析并发,但把 `SINGLE_STOCK_NOTIFY=true` 下的即时通知挪到结果收集侧串行发送;同时 `_send_single_stock_notification()` 为同一个 pipeline 实例补上实例级临界区,避免直接调用 `process_single_stock(..., single_stock_notify=True)` 时多个线程继续共享同一个 `NotificationService` 进入报告生成与发送链路,导致通知乱序、重复发送或状态污染。补充了 `tests/test_pipeline_single_stock_notify.py` 与 `tests/test_pipeline_single_notify_thread_safety.py` 的回归场景以覆盖并发单股推送与直接单股入口的串行化行为。
- [修复] 🔇 **实时行情降级提示收口为单次告警** — 分析主流程获取股票名称时不再提前触发一次实时行情查询,避免每只股票重复命中 quote 链路;当某个前置实时数据源失败但后续 fallback 成功时,不再输出“实时行情获取失败”级别提示,只有在实时行情开关关闭或全部数据源都不可用时,才提示已降级为历史收盘价继续分析。
## [3.11.0] - 2026-03-27
+2
View File
@@ -154,6 +154,8 @@
默認每個工作日 **18:00(北京時間)** 自動執行
> 斷點續傳與 `--dry-run` 的資料存在性判斷,現在會按股票所屬市場的本地時區與交易日曆解析「最新可復用交易日」;週末 / 節假日會復用最近交易日,交易日盤中會復用上一個已完成交易日,盤後若當日資料已落庫則可直接跳過。詳細規則見 [完整配置指南](full-guide.md)。
### 方式二:本地運行 / Docker 部署
> 📖 本地運行、Docker 部署詳細步驟請參考 [完整配置指南](full-guide.md)
+2
View File
@@ -171,6 +171,8 @@ The system will:
- Send analysis reports to all configured channels
- Save reports locally
> Resume fetch and `--dry-run` data-existence checks now resolve the "latest reusable trading day" from each market's local timezone and trading calendar. Weekends and holidays reuse the most recent trading day, intraday runs reuse the last completed trading day, and after market close the run skips only if the current trading day's data is already stored. See [Full Guide](full-guide_EN.md) for the exact rules.
---
### Option 2: Local Deployment
+4
View File
@@ -522,6 +522,10 @@ docker run -e SCHEDULE_ENABLED=true -e SCHEDULE_RUN_IMMEDIATELY=false ...
- 使用 `exchange-calendars` 区分 A 股 / 港股 / 美股各自的交易日历(含节假日)
- 混合持仓时,每只股票只在其市场开市日分析,休市股票当日跳过
- 全部相关市场均为非交易日时,整体跳过执行(不启动 pipeline、不发推送)
- 断点续传和 `--dry-run` 的“数据已存在”判断共用同一套“最新可复用交易日”解析逻辑,不再直接使用服务器自然日
- `最新可复用交易日` 会按股票所属市场的本地时区解析:A 股使用 `Asia/Shanghai`,港股使用 `Asia/Hong_Kong`,美股使用 `America/New_York`
- 非交易日(周末 / 节假日)运行时,会回退到最近一个交易日检查本地数据;若该交易日数据已存在,则跳过重复抓取,否则继续补数
- 交易日盘中或收盘前运行时,会以上一个已完成交易日作为复用目标;交易日收盘后运行时,当日数据已存在则可直接跳过,不存在则继续抓取
- 覆盖方式:`TRADING_DAY_CHECK_ENABLED=false` 或 命令行 `--force-run`
#### 使用 Crontab
+63 -23
View File
@@ -17,7 +17,7 @@ import time
import uuid
from collections import defaultdict
from concurrent.futures import ThreadPoolExecutor, as_completed
from datetime import date, timedelta
from datetime import date, datetime, timedelta, timezone
from typing import List, Dict, Any, Optional, Tuple
import pandas as pd
@@ -39,7 +39,12 @@ from src.search_service import SearchService
from src.services.social_sentiment_service import SocialSentimentService
from src.enums import ReportType
from src.stock_analyzer import StockTrendAnalyzer, TrendAnalysisResult
from src.core.trading_calendar import get_market_for_stock, is_market_open
from src.core.trading_calendar import (
get_effective_trading_date,
get_market_for_stock,
get_market_now,
is_market_open,
)
from data_provider.us_index_mapping import is_us_stock_code
from bot.models import BotMessage
@@ -135,19 +140,21 @@ class StockAnalysisPipeline:
def fetch_and_save_stock_data(
self,
code: str,
force_refresh: bool = False
force_refresh: bool = False,
current_time: Optional[datetime] = None,
) -> Tuple[bool, Optional[str]]:
"""
获取并保存单只股票数据
断点续传逻辑:
1. 检查数据库是否已有今日数据
1. 检查数据库是否已有最新可复用交易日数据
2. 如果有且不强制刷新,则跳过网络请求
3. 否则从数据源获取并保存
Args:
code: 股票代码
force_refresh: 是否强制刷新(忽略本地缓存)
current_time: 本轮运行冻结的参考时间,用于统一断点续传目标交易日判断
Returns:
Tuple[是否成功, 错误信息]
@@ -157,16 +164,15 @@ class StockAnalysisPipeline:
# 首先获取股票名称
stock_name = self.fetcher_manager.get_stock_name(code, allow_realtime=False)
today = date.today()
# 注意:这里用自然日 date.today() 做“断点续传”判断。
# 若在周末/节假日/非交易日运行,或机器时区不在中国,可能出现:
# - 数据库已有最新交易日数据但仍会重复拉取(has_today_data 返回 False)
# - 或在跨日/时区偏移时误判“今日已有数据”
# 该行为目前保留(按需求不改逻辑),但如需更严谨可改为“最新交易日/数据源最新日期”判断。
# 断点续传检查:如果今日数据已存在,跳过
if not force_refresh and self.db.has_today_data(code, today):
logger.info(f"{stock_name}({code}) 今日数据已存在,跳过获取(断点续传)")
target_date = self._resolve_resume_target_date(
code, current_time=current_time
)
# 断点续传检查:如果最新可复用交易日的数据已存在,则跳过
if not force_refresh and self.db.has_today_data(code, target_date):
logger.info(
f"{stock_name}({code}) {target_date} 数据已存在,跳过获取(断点续传)"
)
return True, None
# 从数据源获取数据
@@ -295,7 +301,8 @@ class StockAnalysisPipeline:
# Step 3: 趋势分析(基于交易理念)— 在 Agent 分支之前执行,供两条路径共用
trend_result: Optional[TrendAnalysisResult] = None
try:
end_date = date.today()
_mkt = get_market_for_stock(normalize_stock_code(code))
end_date = get_market_now(_mkt).date()
start_date = end_date - timedelta(days=89) # ~60 trading days for MA60
historical_bars = self.db.get_data_range(code, start_date, end_date)
if historical_bars:
@@ -379,10 +386,13 @@ class StockAnalysisPipeline:
if context is None:
logger.warning(f"{stock_name}({code}) 无法获取历史行情数据,将仅基于新闻和实时行情分析")
_mkt_date = get_market_now(
get_market_for_stock(normalize_stock_code(code))
).date()
context = {
'code': code,
'stock_name': stock_name,
'date': date.today().isoformat(),
'date': _mkt_date.isoformat(),
'data_missing': True,
'today': {},
'yesterday': {}
@@ -566,7 +576,9 @@ class StockAnalysisPipeline:
enhanced['ma_status'] = self._compute_ma_status(
price, trend_result.ma5, trend_result.ma10, trend_result.ma20
)
enhanced['date'] = date.today().isoformat()
enhanced['date'] = get_market_now(
get_market_for_stock(normalize_stock_code(enhanced.get('code', '')))
).date().isoformat()
if yesterday_close is not None:
try:
yc = float(yesterday_close)
@@ -950,7 +962,8 @@ class StockAnalysisPipeline:
if not enable_realtime_tech:
return df
market = get_market_for_stock(code)
if market and not is_market_open(market, date.today()):
market_today = get_market_now(market).date()
if market and not is_market_open(market, market_today):
return df
last_val = df['date'].max()
@@ -968,7 +981,7 @@ class StockAnalysisPipeline:
amt = getattr(realtime_quote, 'amount', None)
pct = getattr(realtime_quote, 'change_pct', None)
if last_date >= date.today():
if last_date >= market_today:
# Update last row with realtime close (copy to avoid mutating caller's df)
df = df.copy()
idx = df.index[-1]
@@ -989,7 +1002,7 @@ class StockAnalysisPipeline:
# Append virtual today row
new_row = {
'code': code,
'date': date.today(),
'date': market_today,
'open': open_p,
'high': high_p,
'low': low_p,
@@ -1019,6 +1032,16 @@ class StockAnalysisPipeline:
"chip_distribution_raw": self._safe_to_dict(chip_data),
}
@staticmethod
def _resolve_resume_target_date(
code: str, current_time: Optional[datetime] = None
) -> date:
"""
Resolve the trading date used by checkpoint/resume checks.
"""
market = get_market_for_stock(normalize_stock_code(code))
return get_effective_trading_date(market, current_time=current_time)
@staticmethod
def _safe_to_dict(value: Any) -> Optional[Dict[str, Any]]:
"""
@@ -1092,6 +1115,7 @@ class StockAnalysisPipeline:
single_stock_notify: bool = False,
report_type: ReportType = ReportType.SIMPLE,
analysis_query_id: Optional[str] = None,
current_time: Optional[datetime] = None,
) -> Optional[AnalysisResult]:
"""
处理单只股票的完整流程
@@ -1110,6 +1134,7 @@ class StockAnalysisPipeline:
skip_analysis: 是否跳过 AI 分析
single_stock_notify: 是否启用单股推送模式(每分析完一只立即推送)
report_type: 报告类型枚举(从配置读取,Issue #119)
current_time: 本轮运行冻结的参考时间,用于统一断点续传目标交易日判断
Returns:
AnalysisResult 或 None
@@ -1118,7 +1143,9 @@ class StockAnalysisPipeline:
try:
# Step 1: 获取并保存数据
success, error = self.fetch_and_save_stock_data(code)
success, error = self.fetch_and_save_stock_data(
code, current_time=current_time
)
if not success:
logger.warning(f"[{code}] 数据获取失败: {error}")
@@ -1197,6 +1224,9 @@ class StockAnalysisPipeline:
logger.info(f"===== 开始分析 {len(stock_codes)} 只股票 =====")
logger.info(f"股票列表: {', '.join(stock_codes)}")
logger.info(f"并发数: {self.max_workers}, 模式: {'仅获取数据' if dry_run else '完整分析'}")
# 冻结本轮运行的统一参考时间,避免跨市场收盘边界时同批股票使用不同目标交易日。
resume_reference_time = datetime.now(timezone.utc)
# === 批量预取实时行情(优化:避免每只股票都触发全量拉取)===
# 只有股票数量 >= 5 时才进行预取,少量股票直接逐个查询更高效
@@ -1243,6 +1273,7 @@ class StockAnalysisPipeline:
single_stock_notify=False,
report_type=report_type, # Issue #119: 传递报告类型
analysis_query_id=uuid.uuid4().hex,
current_time=resume_reference_time,
): code
for code in stock_codes
}
@@ -1278,8 +1309,17 @@ class StockAnalysisPipeline:
# dry-run 模式下,数据获取成功即视为成功
if dry_run:
# 检查哪些股票的数据今天已存在
success_count = sum(1 for code in stock_codes if self.db.has_today_data(code))
# 检查哪些股票的最新可复用交易日数据已存在
success_count = sum(
1
for code in stock_codes
if self.db.has_today_data(
code,
self._resolve_resume_target_date(
code, current_time=resume_reference_time
),
)
)
fail_count = len(stock_codes) - success_count
else:
success_count = len(results)
+74 -1
View File
@@ -15,6 +15,7 @@
import logging
from datetime import date, datetime
from typing import Optional, Set
from zoneinfo import ZoneInfo
logger = logging.getLogger(__name__)
@@ -90,6 +91,79 @@ def is_market_open(market: str, check_date: date) -> bool:
return True
def get_market_now(
market: Optional[str], current_time: Optional[datetime] = None
) -> datetime:
"""
Return current time in the market's local timezone.
If current_time is naive, treat it as already expressed in the market timezone.
Unknown markets fall back to the given datetime (or local system time).
"""
tz_name = MARKET_TIMEZONE.get(market or "")
if current_time is None:
if tz_name:
return datetime.now(ZoneInfo(tz_name))
return datetime.now()
if not tz_name:
return current_time
tz = ZoneInfo(tz_name)
if current_time.tzinfo is None:
return current_time.replace(tzinfo=tz)
return current_time.astimezone(tz)
def get_effective_trading_date(
market: Optional[str], current_time: Optional[datetime] = None
) -> date:
"""
Resolve the latest reusable daily-bar date for checkpoint/resume logic.
Rules:
- Non-trading day / holiday: previous trading session
- Trading day before market close: previous completed trading session
- Trading day after market close: current trading session
- Calendar lookup failure: fail-open to market-local natural date
"""
market_now = get_market_now(market, current_time=current_time)
fallback_date = market_now.date()
if not _XCALS_AVAILABLE:
return fallback_date
ex = MARKET_EXCHANGE.get(market or "")
tz_name = MARKET_TIMEZONE.get(market or "")
if not ex or not tz_name:
return fallback_date
try:
cal = xcals.get_calendar(ex)
local_date = market_now.date()
if not cal.is_session(local_date):
return cal.date_to_session(local_date, direction="previous").date()
session = cal.date_to_session(local_date, direction="previous")
session_close = cal.session_close(session)
if hasattr(session_close, "tz_convert"):
close_local = session_close.tz_convert(tz_name).to_pydatetime()
elif session_close.tzinfo is not None:
close_local = session_close.astimezone(ZoneInfo(tz_name))
else:
close_local = session_close.replace(tzinfo=ZoneInfo(tz_name))
if market_now >= close_local:
return session.date()
return cal.previous_session(session).date()
except Exception as e:
logger.warning("trading_calendar.get_effective_trading_date fail-open: %s", e)
return fallback_date
def get_open_markets_today() -> Set[str]:
"""
Get markets that are open today (by each market's local timezone).
@@ -100,7 +174,6 @@ def get_open_markets_today() -> Set[str]:
if not _XCALS_AVAILABLE:
return {"cn", "hk", "us"}
result: Set[str] = set()
from zoneinfo import ZoneInfo
for mkt, tz_name in MARKET_TIMEZONE.items():
try:
tz = ZoneInfo(tz_name)
+92
View File
@@ -0,0 +1,92 @@
# -*- coding: utf-8 -*-
"""Tests that _augment_historical_with_realtime uses market-local date."""
import unittest
from datetime import date, datetime
from types import SimpleNamespace
from unittest.mock import patch
import pandas as pd
from src.core.pipeline import StockAnalysisPipeline
def _make_pipeline():
p = StockAnalysisPipeline.__new__(StockAnalysisPipeline)
p.config = SimpleNamespace(enable_realtime_technical_indicators=True)
return p
def _make_df(dates_and_closes):
rows = [
{"code": "AAPL", "date": d, "open": c, "high": c, "low": c, "close": c, "volume": 100, "amount": 0, "pct_chg": 0}
for d, c in dates_and_closes
]
return pd.DataFrame(rows)
class AugmentRealtimeMarketDateTestCase(unittest.TestCase):
"""Verify _augment_historical_with_realtime uses market-local date, not server date."""
@patch("src.core.pipeline.is_market_open", return_value=True)
@patch("src.core.pipeline.get_market_now")
@patch("src.core.pipeline.get_market_for_stock", return_value="us")
def test_appends_virtual_row_with_market_local_date(
self, _mock_market, mock_now, _mock_open
):
"""When server is UTC and US market date differs, virtual row uses market date."""
# Server UTC: 2026-03-28 01:00 => US ET: 2026-03-27 21:00
us_market_now = datetime(2026, 3, 27, 21, 0)
mock_now.return_value = us_market_now
df = _make_df([(date(2026, 3, 26), 150.0)])
quote = SimpleNamespace(price=155.0, open_price=151.0, high=156.0, low=149.0, volume=200, amount=None, change_pct=3.0, pre_close=None)
pipeline = _make_pipeline()
result = pipeline._augment_historical_with_realtime(df, quote, "AAPL")
self.assertEqual(len(result), 2)
appended_date = result.iloc[-1]["date"]
if hasattr(appended_date, "date"):
appended_date = appended_date.date()
self.assertEqual(appended_date, date(2026, 3, 27))
@patch("src.core.pipeline.is_market_open", return_value=True)
@patch("src.core.pipeline.get_market_now")
@patch("src.core.pipeline.get_market_for_stock", return_value="us")
def test_updates_existing_row_when_data_matches_market_date(
self, _mock_market, mock_now, _mock_open
):
"""When latest bar date >= market_today, update in place instead of appending."""
mock_now.return_value = datetime(2026, 3, 27, 17, 0)
df = _make_df([(date(2026, 3, 26), 150.0), (date(2026, 3, 27), 152.0)])
quote = SimpleNamespace(price=155.0, open_price=151.0, high=156.0, low=149.0, volume=200, amount=None, change_pct=3.0, pre_close=None)
pipeline = _make_pipeline()
result = pipeline._augment_historical_with_realtime(df, quote, "AAPL")
self.assertEqual(len(result), 2)
self.assertEqual(result.iloc[-1]["close"], 155.0)
@patch("src.core.pipeline.is_market_open", return_value=False)
@patch("src.core.pipeline.get_market_now")
@patch("src.core.pipeline.get_market_for_stock", return_value="cn")
def test_skips_augmentation_on_non_trading_day(
self, _mock_market, mock_now, _mock_open
):
"""Weekend/holiday: returns df unchanged."""
mock_now.return_value = datetime(2026, 3, 28, 10, 0)
df = _make_df([(date(2026, 3, 27), 30.0)])
quote = SimpleNamespace(price=31.0, open_price=30.5, high=31.5, low=29.5, volume=100, amount=None, change_pct=1.0, pre_close=None)
pipeline = _make_pipeline()
result = pipeline._augment_historical_with_realtime(df, quote, "600519")
self.assertEqual(len(result), 1)
self.assertEqual(result.iloc[0]["close"], 30.0)
if __name__ == "__main__":
unittest.main()
+42 -1
View File
@@ -1,8 +1,9 @@
# -*- coding: utf-8 -*-
"""Regression tests for pipeline data-fetch error handling."""
from datetime import date, datetime, timezone
import unittest
from unittest.mock import MagicMock
from unittest.mock import MagicMock, patch
from src.core.pipeline import StockAnalysisPipeline
@@ -21,6 +22,46 @@ class PipelineFetchErrorTestCase(unittest.TestCase):
self.assertFalse(success)
self.assertIn("name lookup failed", error or "")
@patch.object(
StockAnalysisPipeline,
"_resolve_resume_target_date",
return_value=date(2026, 3, 27),
)
def test_fetch_and_save_uses_effective_trading_date_for_resume_check(self, _mock_target):
pipeline = StockAnalysisPipeline.__new__(StockAnalysisPipeline)
pipeline.fetcher_manager = MagicMock()
pipeline.db = MagicMock()
pipeline.fetcher_manager.get_stock_name.return_value = "贵州茅台"
pipeline.db.has_today_data.return_value = True
current_time = datetime(2026, 3, 28, 1, 0, tzinfo=timezone.utc)
success, error = StockAnalysisPipeline.fetch_and_save_stock_data(
pipeline,
"600519",
current_time=current_time,
)
self.assertTrue(success)
self.assertIsNone(error)
_mock_target.assert_called_once_with("600519", current_time=current_time)
pipeline.db.has_today_data.assert_called_once_with("600519", date(2026, 3, 27))
pipeline.fetcher_manager.get_daily_data.assert_not_called()
def test_resolve_resume_target_date_normalizes_supported_a_share_formats(self):
with patch("src.core.pipeline.get_market_for_stock", return_value="cn") as mock_market, patch(
"src.core.pipeline.get_effective_trading_date",
return_value=date(2026, 3, 27),
) as mock_target:
for code in ("SH600519", "000001.SZ", "BJ920748"):
result = StockAnalysisPipeline._resolve_resume_target_date(code)
self.assertEqual(result, date(2026, 3, 27))
self.assertEqual(
[args.args[0] for args in mock_market.call_args_list],
["600519", "000001", "920748"],
)
self.assertEqual(mock_target.call_count, 3)
if __name__ == "__main__":
unittest.main()
+51 -1
View File
@@ -6,8 +6,9 @@ Regression tests for prefetch behavior in StockAnalysisPipeline.run().
import os
import sys
import unittest
from datetime import date
from types import SimpleNamespace
from unittest.mock import MagicMock
from unittest.mock import MagicMock, call
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
@@ -52,6 +53,55 @@ class TestPipelinePrefetchBehavior(unittest.TestCase):
["000001"], use_bulk=False
)
def test_run_dry_run_counts_existing_data_by_effective_trading_date(self):
pipeline = self._build_pipeline(process_result=None)
pipeline._resolve_resume_target_date = MagicMock(
side_effect=[date(2026, 3, 27), date(2026, 3, 26)]
)
pipeline.db.has_today_data.side_effect = [True, False]
pipeline.run(
stock_codes=["600519", "AAPL"],
dry_run=True,
send_notification=False,
)
self.assertEqual(
pipeline.db.has_today_data.call_args_list,
[
call("600519", date(2026, 3, 27)),
call("AAPL", date(2026, 3, 26)),
],
)
def test_run_uses_one_frozen_reference_time_for_tasks_and_dry_run_stats(self):
pipeline = self._build_pipeline(process_result=None)
pipeline._resolve_resume_target_date = MagicMock(
side_effect=[date(2026, 3, 27), date(2026, 3, 26)]
)
pipeline.db.has_today_data.side_effect = [True, False]
pipeline.run(
stock_codes=["600519", "AAPL"],
dry_run=True,
send_notification=False,
)
task_reference_times = [
call.kwargs["current_time"]
for call in pipeline.process_single_stock.call_args_list
]
stats_reference_times = [
call.kwargs["current_time"]
for call in pipeline._resolve_resume_target_date.call_args_list
]
self.assertEqual(len(task_reference_times), 2)
self.assertEqual(len(stats_reference_times), 2)
self.assertEqual(len({id(value) for value in task_reference_times}), 1)
self.assertEqual(len({id(value) for value in stats_reference_times}), 1)
self.assertIs(task_reference_times[0], stats_reference_times[0])
if __name__ == "__main__":
unittest.main()
+172
View File
@@ -0,0 +1,172 @@
# -*- coding: utf-8 -*-
"""Regression tests for effective trading date resolution."""
from datetime import date, datetime, time, timezone
from types import SimpleNamespace
import unittest
from unittest.mock import patch
from zoneinfo import ZoneInfo
import pandas as pd
from src.core import trading_calendar
class _FakeCalendar:
def __init__(self, sessions, close_hour: int, tz_name: str):
self._sessions = sorted(sessions)
self._close_hour = close_hour
self._tz_name = tz_name
def is_session(self, check_date: date) -> bool:
return check_date in self._sessions
def date_to_session(self, check_date: date, direction: str = "previous") -> pd.Timestamp:
if direction == "previous":
candidates = [d for d in self._sessions if d <= check_date]
elif direction == "next":
candidates = [d for d in self._sessions if d >= check_date]
else:
raise ValueError(f"unsupported direction: {direction}")
if not candidates:
raise ValueError(f"no session for {check_date} ({direction})")
return pd.Timestamp(candidates[-1] if direction == "previous" else candidates[0])
def previous_session(self, session: pd.Timestamp) -> pd.Timestamp:
session_date = session.date()
index = self._sessions.index(session_date)
if index == 0:
raise ValueError("no previous session")
return pd.Timestamp(self._sessions[index - 1])
def session_close(self, session: pd.Timestamp) -> pd.Timestamp:
local_close = datetime.combine(
session.date(),
time(self._close_hour, 0),
tzinfo=ZoneInfo(self._tz_name),
)
return pd.Timestamp(local_close).tz_convert("UTC")
class EffectiveTradingDateTestCase(unittest.TestCase):
def test_weekend_returns_previous_session(self):
fake_calendar = _FakeCalendar(
sessions=[date(2026, 3, 26), date(2026, 3, 27)],
close_hour=15,
tz_name="Asia/Shanghai",
)
current_time = datetime(2026, 3, 28, 10, 0, tzinfo=ZoneInfo("Asia/Shanghai"))
with patch.object(trading_calendar, "_XCALS_AVAILABLE", True), patch.object(
trading_calendar,
"xcals",
SimpleNamespace(get_calendar=lambda _ex: fake_calendar),
create=True,
):
result = trading_calendar.get_effective_trading_date("cn", current_time=current_time)
self.assertEqual(result, date(2026, 3, 27))
def test_holiday_returns_previous_session(self):
fake_calendar = _FakeCalendar(
sessions=[date(2025, 12, 31), date(2026, 1, 5)],
close_hour=15,
tz_name="Asia/Shanghai",
)
current_time = datetime(2026, 1, 1, 12, 0, tzinfo=ZoneInfo("Asia/Shanghai"))
with patch.object(trading_calendar, "_XCALS_AVAILABLE", True), patch.object(
trading_calendar,
"xcals",
SimpleNamespace(get_calendar=lambda _ex: fake_calendar),
create=True,
):
result = trading_calendar.get_effective_trading_date("cn", current_time=current_time)
self.assertEqual(result, date(2025, 12, 31))
def test_intraday_returns_previous_completed_session(self):
fake_calendar = _FakeCalendar(
sessions=[date(2026, 3, 26), date(2026, 3, 27)],
close_hour=16,
tz_name="America/New_York",
)
current_time = datetime(
2026,
3,
27,
15,
59,
tzinfo=ZoneInfo("America/New_York"),
)
with patch.object(trading_calendar, "_XCALS_AVAILABLE", True), patch.object(
trading_calendar,
"xcals",
SimpleNamespace(get_calendar=lambda _ex: fake_calendar),
create=True,
):
result = trading_calendar.get_effective_trading_date("us", current_time=current_time)
self.assertEqual(result, date(2026, 3, 26))
def test_after_close_returns_current_session(self):
fake_calendar = _FakeCalendar(
sessions=[date(2026, 3, 26), date(2026, 3, 27)],
close_hour=16,
tz_name="America/New_York",
)
current_time = datetime(
2026,
3,
27,
16,
1,
tzinfo=ZoneInfo("America/New_York"),
)
with patch.object(trading_calendar, "_XCALS_AVAILABLE", True), patch.object(
trading_calendar,
"xcals",
SimpleNamespace(get_calendar=lambda _ex: fake_calendar),
create=True,
):
result = trading_calendar.get_effective_trading_date("us", current_time=current_time)
self.assertEqual(result, date(2026, 3, 27))
def test_market_timezone_controls_cross_timezone_resolution(self):
fake_calendar = _FakeCalendar(
sessions=[date(2026, 3, 25), date(2026, 3, 26), date(2026, 3, 27)],
close_hour=16,
tz_name="America/New_York",
)
current_time = datetime(2026, 3, 27, 1, 0, tzinfo=timezone.utc)
with patch.object(trading_calendar, "_XCALS_AVAILABLE", True), patch.object(
trading_calendar,
"xcals",
SimpleNamespace(get_calendar=lambda _ex: fake_calendar),
create=True,
):
result = trading_calendar.get_effective_trading_date("us", current_time=current_time)
self.assertEqual(result, date(2026, 3, 26))
def test_calendar_error_falls_back_to_market_local_date(self):
current_time = datetime(2026, 3, 27, 18, 0, tzinfo=timezone.utc)
with patch.object(trading_calendar, "_XCALS_AVAILABLE", True), patch.object(
trading_calendar,
"xcals",
SimpleNamespace(get_calendar=lambda _ex: (_ for _ in ()).throw(RuntimeError("boom"))),
create=True,
):
result = trading_calendar.get_effective_trading_date("hk", current_time=current_time)
self.assertEqual(result, date(2026, 3, 28))
if __name__ == "__main__":
unittest.main()