mirror of
https://github.com/ZhuLinsen/daily_stock_analysis
synced 2026-09-20 10:53:33 +08:00
股票实体解析增强 (#2245)
Co-authored-by: zhulinsen <42829555+ZhuLinsen@users.noreply.github.com>
This commit is contained in:
@@ -289,6 +289,11 @@ async def app_lifespan(app: FastAPI):
|
||||
runtime_scheduler=app.state.runtime_scheduler_service,
|
||||
)
|
||||
_schedule_stock_index_background_refresh(app, "startup")
|
||||
# 名称解析器的 AkShare 缓存预热:命中磁盘缓存则零网络加载,否则发起
|
||||
# 后台单飞拉取。把冷启动等待从首个用户请求挪到进程启动窗口。
|
||||
from src.services.name_to_code_resolver import warmup_akshare_cache
|
||||
|
||||
warmup_akshare_cache()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
|
||||
@@ -24,6 +24,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/).
|
||||
- [修复] Agent 工具注册表(`src/agent/factory.get_tool_registry`)由模块级缓存改为按「类别超时映射的值」比对失效,规避 CPython 回收对象后地址复用(`id(config)` 相同)导致配置 reload 后的 `Config` 被误判为未变、沿用过期超时的真 bug;新增 `_coerce_config_timeout` 类型白名单,使调用方传入 `MagicMock` / 缺属性 stub / 脏字符串(如 `float(MagicMock())` 静默得到 1.0)时降级为「无类别限制」而非崩溃或强加 1 秒超时;`build_agent_executor(config)` / `build_agent_chat_executor(config)` 现已把调用方 `config` 透传给 `get_tool_registry(config)`(不再无参调用冻结首构 registry);`main._reload_runtime_config` 与 `SystemConfigService._reload_runtime_singletons`(及 `update()`→`reload_now` 路径)在配置热重载时调用 `reset_tool_registry()` 强制重建;回归测试补充「传入新 config 后 registry 重建」「reload 后新超时应生效」及「builder 透传 config」三类场景(#1890 的 review follow-up,闭环 OR-COM-dd1e8fa7 / OR-COM-bff42110)
|
||||
- [修复] Agent 工具超时 review 闭环(fixes #1890 的 4 个 blocker):超时解析由 min 契约改为 first-wins(显式 per-run `tool_call_timeout_seconds` > 单工具 `ToolDefinition.timeout_seconds` > 类别默认 > 无限制,剩余 wall-clock 预算只作不可突破的外层 cap;research 路径不再传 `tool_call_timeout_seconds` 以免覆盖类别限制);超时结果标记 `retriable: false` 并写入 `non_retriable_tool_results` 阻断 LLM 同调用重试重入,且超时触发时为仍在后台运行的 handler 武装协作取消信号(`is_tool_cancellation_requested()` 与既有 `check_tool_execution()` 检查点均响应,handler 从不轮询则行为不变),作为 review 要求的「handler 内协作取消」缓解,规避 Python 线程无法 force-stop 导致的重复执行与副作用;`_coerce_config_timeout` 对 `inf`/`nan`/负数降级为「无限制」,根绝 `future.result(timeout=inf)` 触发 `OverflowError`;`get_tool_registry` / `reset_tool_registry` 加 `threading.Lock` 双检锁,且重建后返回本次构建的局部 registry(而非全局缓存),消除并发重建竞态与跨调用超时串扰;`@tool` 装饰器将 `ToolPolicy.timeout_seconds` 折叠进 `ToolDefinition` 单一来源;统一单/并行工具超时包装(单一 executor + deadline 驱动的 wait loop,消除并行路径嵌套 executor 与线程翻倍,duration 精确到各工具自身超时值),并新增快慢工具混合并行回归;同步 `docs/full-guide_EN.md` 的超时环境变量文档;测试覆盖 first-wins、non-retriable、协作取消接线、finite 校验、缓存线程安全与快慢混合并行。
|
||||
- [修复] 按最新 review 复核收敛 3 处正确性问题(OR-COM-7f3d3f5b / 3d6b61f8 / a1e8b0c2):`BaseAgent._filtered_registry()` 携带源 registry 的类别超时映射(工具子集仍生效类别上限,不再绕过 #1890 类别超时);并行批次 >5 时排队调用的 per-tool 超时自 worker 实际开始起算(不再提交即烧预算导致对未启动调用的假超时);`get_tool_registry()` 缓存命中快路径在锁内读取一致对(消除与 `reset_tool_registry()` 竞态返回 `None` 或错配 registry)。新增对应回归测试。
|
||||
- [新功能] 股票名称解析引擎重构增强:新增 `resolver_name_to_code_list()` 公开 API,返回按市场排序(A 股→港股→美股)的 `Stock` 候选列表(最多 5 个),新增 `US_stock_code_match()` 匹配美股 ticker(1~5 位字母且仅限本地库已存在代码,避免 hello/open 等英文词误判为股票);AkShare 全量 A 股数据经幂等 `extend_AkShare()` 合并进全局 `stockDB`(30 分钟缓存 + 失败 5 分钟退避 + Future 单飞:TTL 过期 stale-while-revalidate 零等待、冷启动等待上界由拉取超时推导(拉取经子进程封顶 25s)、worker 先清账唤醒等待者再做日志/落盘(finally 兜底 BaseException)、成功拉取落盘 `data/cache` 跨重启复用,非中文输入跳过网络扩展),匹配策略升级为「精确→子串(≥2 汉字)→拼音子串(≥5 字母)→difflib 模糊(0.8,单字误写 0.7 兜底)」;`resolve_name_to_code()` 保持既有本地优先语义(本地精确命中零网络,调用方离线低延迟契约不变),跨市场候选能力由 `resolver_name_to_code_list()` 独立提供;解析全链路线程安全(`stockDB` 读写加锁、名称/拼音索引随库变更自动失效),新增 40 个单元测试覆盖精确/跨市场排序/子串/拼音/模糊/幂等扩展/失败退避/多候选场景。
|
||||
- [改进] `StockDaily` 表新增可空 `canonical_id` 列并支持双写(Expand-Contract PR2,issue #2207):自愈式迁移幂等加列 + 普通索引 `ix_stock_daily_canonical_id`,存量行与 `save_daily_data` 未显式传参时均经 index-aware 推导(裸指数码命中注册表时统一到指数 `canonical_id`,避免同一指数按输入形态分裂到不同桶——例如裸 `000300` 与显式 `sh000300` 现在都收敛到 `sh000300`,而非裸码被推导为 `sz000300`),推导失败写 NULL 降级;读路径仍用 `code` 列,`(code, date)` 唯一约束保留不变。显式登记契约漂移:PRD Glossary/FR-1/DD-3 与架构 AD-1/AD-7 中 canonical_id 的点分格式描述(`000016.SH`)已被 Phase 1 已合入代码的前缀格式(`sh000016`)取代,本变更遵循代码,PRD/架构文档的同步修正留待后续 PR 统一收敛。
|
||||
|
||||
## [3.30.0] - 2026-08-09
|
||||
|
||||
@@ -1,41 +1,95 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
===================================
|
||||
Name-to-Code Resolution Engine
|
||||
名称到代码解析引擎
|
||||
===================================
|
||||
|
||||
Resolve stock name to code: local mapping + pinyin + AkShare fallback + fuzzy matching.
|
||||
将股票名称解析为代码:全局 ``stockDB`` 上的精确匹配 + AkShare 扩充 +
|
||||
混合匹配(子串/拼音/模糊)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import concurrent.futures
|
||||
import difflib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
from typing import Dict, Optional, Set, Tuple
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Optional, Set, Tuple
|
||||
|
||||
from src.data.stock_mapping import STOCK_NAME_MAP
|
||||
from src.services.stock_code_utils import is_code_like, normalize_code
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# AkShare result cache: (timestamp, name_to_code_dict)
|
||||
# AkShare 结果缓存:(时间戳, name_to_code 字典)。TTL 过期后作为 stale
|
||||
# 数据继续服务(stale-while-revalidate),刷新由后台线程完成。
|
||||
_akshare_cache: Optional[tuple[float, Dict[str, str]]] = None
|
||||
_AKSHARE_CACHE_TTL = 1800 # 30 MIN
|
||||
|
||||
# 失败退避:最近一次拉取失败的时间戳。不写负缓存的话,网络异常期间每次
|
||||
# 解析都会重打网络请求,慢速超时直接放大为全量消息的用户延迟。
|
||||
# 有 stale 缓存时退避窗口内仍服务 stale(网络故障期间保持可用),
|
||||
# 仅冷启动(无任何缓存)才退化为本地库。
|
||||
_akshare_failure_cache: Optional[float] = None
|
||||
_AKSHARE_FAILURE_TTL = 300 # 5 MIN
|
||||
|
||||
# 拉取超时:akshare 对交易所的部分请求无超时(如上交所 GET、北交所
|
||||
# POST),挂起(非失败)时拉取线程永不返回——挂起时异常路径永不触发、
|
||||
# 失败退避永不武装。拉取经子进程包装封顶时长后,挂起子进程在超时处被
|
||||
# terminate/kill、以 TimeoutError 抛出,走常规失败路径写入失败退避。
|
||||
_AKSHARE_FETCH_TIMEOUT = 25.0
|
||||
|
||||
# 冷启动等待上限:拉取被子进程封顶在 _AKSHARE_FETCH_TIMEOUT 内,worker
|
||||
# 在其后有限的清账步内必然 resolve Future。等待上界从同一常数推导而非
|
||||
# 独立魔法数,结构上保证与在途拉取重叠的请求都能等到其结果——曾用独立
|
||||
# 常数 8s,实测拉取耗时 6.6~8.2s 与其同分布,冷启动请求确定性漏解。
|
||||
# 余量覆盖 DataFrame 解析与清账;落盘/日志在唤醒之后,不计入。TTL 过期
|
||||
# 走 stale-while-revalidate 零等待,此值仅影响冷启动场景。
|
||||
_AKSHARE_WAIT_COLD_START = _AKSHARE_FETCH_TIMEOUT + 5.0
|
||||
|
||||
# 状态锁:只保证缓存/退避/在途句柄读写的原子性。锁内唯一的 IO 例外是
|
||||
# _try_load_disk_cache_locked 的首次懒加载(每进程一次、本地小文件毫秒
|
||||
# 级);网络拉取与落盘都在锁外的后台刷新线程内。与 _db_lock 互不嵌套,
|
||||
# 不构成锁顺序环。
|
||||
_state_lock = threading.Lock()
|
||||
|
||||
# 在途拉取句柄:非 None 表示已有后台线程在刷新。并发等待者等 Future 的
|
||||
# 完成信号(成功/失败都会精确唤醒),而非排队抢锁再双重检查——等待与
|
||||
# 拉取耗时解耦,无 25<30 这类常数间隐式契约。
|
||||
_akshare_inflight: Optional[concurrent.futures.Future] = None
|
||||
|
||||
# 磁盘缓存:成功拉取后原子落盘,重启时懒加载为 stale 数据。股票名称表
|
||||
# 变化频率极低,跨重启复用把"冷启动等待"压缩到仅剩磁盘也无数据的场景。
|
||||
# 路径遵循 data/cache 约定(该目录已 gitignore)。
|
||||
_AKSHARE_DISK_CACHE_PATH = (
|
||||
Path(__file__).resolve().parents[2] / "data" / "cache" / "akshare_name_map.json"
|
||||
)
|
||||
_akshare_disk_checked = False
|
||||
|
||||
# stockDB 读写锁:意图解析在 asyncio.to_thread 工作线程中并发执行,
|
||||
# extend_AkShare 的原地合并与迭代 stockDB 的读路径不加锁会触发
|
||||
# "dictionary changed size during iteration"。RLock 允许嵌套获取。
|
||||
_db_lock = threading.RLock()
|
||||
|
||||
|
||||
def _contains_cjk(text: str) -> bool:
|
||||
"""Return True when text contains CJK characters."""
|
||||
"""当文本包含中日韩(CJK)字符时返回 True。"""
|
||||
return any("\u3400" <= ch <= "\u9fff" for ch in text)
|
||||
|
||||
|
||||
def _is_code_like(s: str) -> bool:
|
||||
"""Backward-compatible wrapper of shared code-like check."""
|
||||
"""共享的“形似代码”检查的向后兼容包装。"""
|
||||
return is_code_like(s)
|
||||
|
||||
|
||||
def _normalize_code(raw: str) -> Optional[str]:
|
||||
"""Backward-compatible wrapper of shared code normalization."""
|
||||
"""共享代码规范化的向后兼容包装。"""
|
||||
return normalize_code(raw)
|
||||
|
||||
|
||||
@@ -43,7 +97,7 @@ def _build_reverse_map_no_duplicates(
|
||||
code_to_name: Dict[str, str],
|
||||
) -> Dict[str, str]:
|
||||
"""
|
||||
Build name -> code map. If a name maps to multiple codes (ambiguous), exclude it.
|
||||
构建 name -> code 映射。若一个名称映射到多个代码(有歧义),则将其排除。
|
||||
"""
|
||||
name_to_codes: Dict[str, Set[str]] = {}
|
||||
for code, name in code_to_name.items():
|
||||
@@ -88,42 +142,348 @@ def _build_local_name_indexes(code_to_name: Dict[str, str]) -> Tuple[Dict[str, s
|
||||
_LOCAL_REVERSE_MAP, _LOCAL_AMBIGUOUS_NAMES = _build_local_name_indexes(STOCK_NAME_MAP)
|
||||
|
||||
|
||||
def _get_akshare_name_to_code() -> Optional[Dict[str, str]]:
|
||||
"""Fetch A-share name->code from AkShare, with cache."""
|
||||
global _akshare_cache
|
||||
now = time.time()
|
||||
if _akshare_cache is not None and (now - _akshare_cache[0]) < _AKSHARE_CACHE_TTL:
|
||||
return _akshare_cache[1]
|
||||
try:
|
||||
import akshare as ak
|
||||
def _akshare_stock_info_worker():
|
||||
"""子进程内执行:导入 akshare 并拉取全量 A 股代码表。
|
||||
|
||||
df = ak.stock_info_a_code_name()
|
||||
if df is None or df.empty:
|
||||
return None
|
||||
code_to_name = {}
|
||||
for _, row in df.iterrows():
|
||||
code = row.get("code")
|
||||
name = row.get("name")
|
||||
if code is None or name is None:
|
||||
continue
|
||||
code_str = str(code).strip()
|
||||
# Strip .SH/.SZ suffix
|
||||
if "." in code_str:
|
||||
base, suffix = code_str.rsplit(".", 1)
|
||||
if suffix.upper() in ("SH", "SZ", "SS") and base.isdigit():
|
||||
code_str = base
|
||||
code_to_name[code_str] = str(name).strip()
|
||||
result = _build_reverse_map_no_duplicates(code_to_name)
|
||||
_akshare_cache = (now, result)
|
||||
logger.info(f"[NameResolver] AkShare cache loaded: {len(result)} name->code mappings")
|
||||
return result
|
||||
必须保持模块级函数(spawn 经 pickle 按限定名传递);akshare 及其
|
||||
pandas 等重依赖只进子进程,父进程不加载。
|
||||
"""
|
||||
import akshare as ak
|
||||
|
||||
return ak.stock_info_a_code_name()
|
||||
|
||||
|
||||
def _fetch_akshare_df():
|
||||
"""AkShare 全量拉取的唯一网络出口(测试接缝)。
|
||||
|
||||
经 data_provider 的子进程超时包装(spawn)执行:挂起的网络请求在
|
||||
_AKSHARE_FETCH_TIMEOUT 内被杀死,父进程侧以 TimeoutError 抛出。
|
||||
akshare 内部的 @lru_cache 只存在于子进程,30 分钟 TTL 过期后的
|
||||
刷新在父进程侧始终拿到新数据。
|
||||
"""
|
||||
from data_provider.akshare_fetcher import _akshare_call_with_timeout
|
||||
|
||||
return _akshare_call_with_timeout(
|
||||
_akshare_stock_info_worker,
|
||||
timeout=_AKSHARE_FETCH_TIMEOUT,
|
||||
call_name="stock_info_a_code_name",
|
||||
)
|
||||
|
||||
|
||||
def _build_name_map_from_df(df) -> Dict[str, str]:
|
||||
"""把 AkShare 的 code/name DataFrame 转为去歧义的 name->code 映射。"""
|
||||
code_to_name = {}
|
||||
for _, row in df.iterrows():
|
||||
code = row.get("code")
|
||||
name = row.get("name")
|
||||
if code is None or name is None:
|
||||
continue
|
||||
code_str = str(code).strip()
|
||||
# Strip .SH/.SZ suffix
|
||||
if "." in code_str:
|
||||
base, suffix = code_str.rsplit(".", 1)
|
||||
if suffix.upper() in ("SH", "SZ", "SS") and base.isdigit():
|
||||
code_str = base
|
||||
code_to_name[code_str] = str(name).strip()
|
||||
return _build_reverse_map_no_duplicates(code_to_name)
|
||||
|
||||
|
||||
def _try_load_disk_cache_locked() -> None:
|
||||
"""懒加载磁盘缓存到内存(每进程至多尝试一次)。调用方必须持有 _state_lock。"""
|
||||
global _akshare_cache, _akshare_disk_checked
|
||||
if _akshare_disk_checked:
|
||||
return
|
||||
_akshare_disk_checked = True
|
||||
try:
|
||||
payload = json.loads(_AKSHARE_DISK_CACHE_PATH.read_text(encoding="utf-8"))
|
||||
ts = payload.get("ts")
|
||||
raw_map = payload.get("map")
|
||||
if isinstance(ts, (int, float)) and isinstance(raw_map, dict):
|
||||
name_map = {str(k): str(v) for k, v in raw_map.items() if k and v}
|
||||
if name_map:
|
||||
_akshare_cache = (float(ts), name_map)
|
||||
logger.info(
|
||||
"[NameResolver] AkShare 磁盘缓存已加载: "
|
||||
f"{len(name_map)} 条, age={time.time() - float(ts):.0f}s"
|
||||
)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.warning(f"[NameResolver] AkShare fallback failed: {e}")
|
||||
logger.debug(f"[NameResolver] AkShare 磁盘缓存不可用,忽略: {e}")
|
||||
|
||||
|
||||
def _persist_akshare_map(name_map: Dict[str, str]) -> None:
|
||||
"""原子落盘(每进程唯一 tmp 名 + os.replace)。失败非致命,仅记日志。
|
||||
|
||||
tmp 名掺入 PID:web 与 CLI 可能是两个进程,共用同一 tmp 名会在
|
||||
os.replace 前互相覆盖对方尚未写完的文件。
|
||||
"""
|
||||
try:
|
||||
path = _AKSHARE_DISK_CACHE_PATH
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
tmp = path.with_name(f"{path.name}.{os.getpid()}.tmp")
|
||||
tmp.write_text(
|
||||
json.dumps({"ts": time.time(), "map": name_map}, ensure_ascii=False),
|
||||
encoding="utf-8",
|
||||
)
|
||||
os.replace(tmp, path)
|
||||
except Exception as e:
|
||||
logger.debug(f"[NameResolver] AkShare 磁盘缓存写入失败(非致命): {e}")
|
||||
|
||||
|
||||
def _spawn_refresh_locked() -> concurrent.futures.Future:
|
||||
"""启动后台刷新线程并登记在途句柄。调用方必须持有 _state_lock。
|
||||
|
||||
网络拉取、DataFrame 解析、落盘全部发生在这个 daemon 线程内——请求
|
||||
路径只读取状态或等待 Future,永不执行网络 IO。
|
||||
"""
|
||||
global _akshare_inflight
|
||||
fut: concurrent.futures.Future = concurrent.futures.Future()
|
||||
# 先登记再启动:worker 理论上可能在 start() 返回前就完成,其清账
|
||||
# 分支 `is fut` 必须能看到已登记的句柄
|
||||
_akshare_inflight = fut
|
||||
|
||||
def _worker():
|
||||
global _akshare_cache, _akshare_failure_cache, _akshare_inflight
|
||||
result: Optional[Dict[str, str]] = None
|
||||
try:
|
||||
try:
|
||||
df = _fetch_akshare_df()
|
||||
if df is not None and not df.empty:
|
||||
# 空数据视同失败退避:短时间内重试只会得到同样的空结果
|
||||
result = _build_name_map_from_df(df)
|
||||
except Exception as e:
|
||||
logger.warning(f"[NameResolver] AkShare fallback failed: {e}")
|
||||
finally:
|
||||
# 先清账并唤醒全部等待者,再做日志/落盘等 best-effort 副作用:
|
||||
# 慢日志 handler 或慢磁盘不得拖住等待者拿结果(曾把落盘排在
|
||||
# set_result 之前,慢盘时缓存已就绪、等待者却超时漏解)。
|
||||
# finally 兜底:BaseException 逃逸时句柄/退避/Future 也必然
|
||||
# 清账,不留死句柄(否则之后每个冷启动请求都白等满超时)。
|
||||
with _state_lock:
|
||||
now = time.time()
|
||||
if result:
|
||||
_akshare_cache = (now, result)
|
||||
_akshare_failure_cache = None
|
||||
else:
|
||||
_akshare_failure_cache = now
|
||||
result = None
|
||||
if _akshare_inflight is fut:
|
||||
_akshare_inflight = None
|
||||
# 成功/失败都精确唤醒全部等待者;失败以 result=None 表达
|
||||
fut.set_result(result)
|
||||
if result:
|
||||
logger.info(
|
||||
f"[NameResolver] AkShare cache loaded: {len(result)} name->code mappings"
|
||||
)
|
||||
_persist_akshare_map(result)
|
||||
|
||||
try:
|
||||
threading.Thread(
|
||||
target=_worker, name="akshare-name-refresh", daemon=True
|
||||
).start()
|
||||
except RuntimeError as e:
|
||||
# 线程启动失败(如进程线程数耗尽):归还句柄并唤醒等待者
|
||||
_akshare_inflight = None
|
||||
fut.set_result(None)
|
||||
logger.warning(f"[NameResolver] AkShare 刷新线程启动失败: {e}")
|
||||
return fut
|
||||
|
||||
|
||||
def _get_akshare_name_to_code() -> Optional[Dict[str, str]]:
|
||||
"""获取 AkShare name->code:新鲜缓存直读;过期走 stale-while-revalidate
|
||||
(立即返回旧值 + 后台刷新);冷启动限时等待在途 Future;失败退避。
|
||||
|
||||
请求路径最坏只阻塞冷启动等待(_AKSHARE_WAIT_COLD_START);TTL 过期
|
||||
的请求零等待;网络 IO 全部在后台刷新线程。
|
||||
"""
|
||||
with _state_lock:
|
||||
if _akshare_cache is None:
|
||||
_try_load_disk_cache_locked()
|
||||
now = time.time()
|
||||
if _akshare_cache is not None and (now - _akshare_cache[0]) < _AKSHARE_CACHE_TTL:
|
||||
return _akshare_cache[1]
|
||||
stale_map = _akshare_cache[1] if _akshare_cache is not None else None
|
||||
in_backoff = (
|
||||
_akshare_failure_cache is not None
|
||||
and (now - _akshare_failure_cache) < _AKSHARE_FAILURE_TTL
|
||||
)
|
||||
fut = _akshare_inflight
|
||||
if fut is None and not in_backoff:
|
||||
fut = _spawn_refresh_locked()
|
||||
if stale_map is not None:
|
||||
# stale-while-revalidate:TTL 的语义就是容忍陈旧,过期瞬间等待
|
||||
# 毫无收益;返回旧值,下轮调用命中后台刷新出的新缓存。
|
||||
return stale_map
|
||||
if fut is None:
|
||||
# 冷启动且处于失败退避窗口:本地库照常服务
|
||||
return None
|
||||
try:
|
||||
result = fut.result(timeout=_AKSHARE_WAIT_COLD_START)
|
||||
except concurrent.futures.TimeoutError:
|
||||
# 正常情况下不可能到达:拉取被封顶、worker 在 finally 中必然
|
||||
# resolve。到达即 worker 违约(如子进程 kill 本身挂死)。
|
||||
logger.warning(
|
||||
"[NameResolver] 在途 AkShare 拉取超过死线未 resolve"
|
||||
f"({_AKSHARE_WAIT_COLD_START:g}s),本轮退化为本地库解析"
|
||||
)
|
||||
return None
|
||||
except Exception as e:
|
||||
# worker 以 result=None 表达失败,这里只可能是 Future 内部异常
|
||||
logger.warning(f"[NameResolver] 等待在途 AkShare 拉取异常: {e}")
|
||||
return None
|
||||
return result
|
||||
|
||||
|
||||
def warmup_akshare_cache() -> None:
|
||||
"""进程启动预热:后台线程触发一次解析链填充(幂等、非阻塞)。
|
||||
|
||||
命中磁盘缓存则零网络加载;否则发起后台单飞拉取。冷启动等待发生在
|
||||
本 daemon 线程内,不阻塞调用方。重复调用共享同一在途句柄,不会
|
||||
触发多次网络拉取。
|
||||
"""
|
||||
|
||||
def _warm():
|
||||
try:
|
||||
_get_akshare_name_to_code()
|
||||
except Exception as e: # noqa: BLE001 - 预热必须 best-effort
|
||||
logger.warning(f"[NameResolver] AkShare 预热失败: {e}")
|
||||
|
||||
threading.Thread(target=_warm, name="akshare-warmup", daemon=True).start()
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Stock — immutable code/name/market value object
|
||||
# =========================================================================
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Stock:
|
||||
"""股票条目:code / name / market("a" | "hk" | "us")。"""
|
||||
|
||||
code: str
|
||||
name: str = ""
|
||||
market: str = ""
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# stockDB — 全局 code->name 数据库,由 AkShare A 股数据扩充
|
||||
# =========================================================================
|
||||
|
||||
|
||||
stockDB: Dict[str, str] = dict(STOCK_NAME_MAP)
|
||||
|
||||
# 已存在代码被 AkShare 更新为当前官方名称时,旧名称作为别名保留,
|
||||
# 避免证券改名后旧名称彻底不可解析。
|
||||
stockAliases: Dict[str, Set[str]] = {}
|
||||
|
||||
_MARKET_ORDER: Dict[str, int] = {"a": 0, "hk": 1, "us": 2}
|
||||
|
||||
# 已合并的 AkShare 缓存对象(幂等短路:同一份缓存不重复合并)
|
||||
_akshare_merged: Optional[Dict[str, str]] = None
|
||||
|
||||
|
||||
def extend_AkShare() -> bool:
|
||||
"""用 AkShare 全量 A 股数据扩充 stockDB(幂等,30 分钟缓存)。
|
||||
|
||||
Returns:
|
||||
True — stockDB 实际并入了新条目或更新了已存在代码的名称;
|
||||
False — 无需扩展(同一份缓存已合并 / 拉取失败 / 返回空数据 /
|
||||
无任何变化)。
|
||||
"""
|
||||
global _akshare_merged
|
||||
akshare_map = _get_akshare_name_to_code()
|
||||
if not akshare_map:
|
||||
return False
|
||||
# 合并段整体持锁:与读路径的 stockDB 迭代串行,且"已合并对象"判定原子。
|
||||
# 网络抓取(上方)保持在锁外。
|
||||
with _db_lock:
|
||||
if _akshare_merged is akshare_map:
|
||||
return False
|
||||
changed = False
|
||||
# akshare_map 是 name->code 反向映射,stockDB 是 code->name
|
||||
for name, code in akshare_map.items():
|
||||
if code not in stockDB:
|
||||
stockDB[code] = name
|
||||
changed = True
|
||||
elif stockDB[code] != name:
|
||||
# 证券改名:保留旧名称作为别名,再更新为 AkShare 当前官方名称。
|
||||
stockAliases.setdefault(code, set()).add(stockDB[code])
|
||||
stockDB[code] = name
|
||||
changed = True
|
||||
_akshare_merged = akshare_map
|
||||
if changed:
|
||||
# stockDB 值更新或别名变化时 len(stockDB) 可能不变,
|
||||
# 必须显式失效名称/拼音缓存,否则改名后的新名称不可见。
|
||||
_names_cache[:] = [None, None, None]
|
||||
_pinyin_cache[:] = [None, None]
|
||||
return changed
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# 名称/拼音缓存 —— 以 (database 对象, len) 为键:对数据库的任何插入或
|
||||
# 删除都会自动使其失效,无需手动重置。
|
||||
# =========================================================================
|
||||
|
||||
|
||||
_names_cache: List = [None, None, None] # [database, len(database), names]
|
||||
_pinyin_cache: List = [None, None] # [names, pinyins]
|
||||
|
||||
|
||||
def _database_names(database: Dict[str, str]) -> List[str]:
|
||||
"""*database* 中去重、保持顺序的名称列表(带缓存)。
|
||||
|
||||
除了 canonical(规范)名称,还会包含 stockDB 的旧名称别名;这样旧名称也能参与
|
||||
精确、子串、拼音和模糊匹配。
|
||||
"""
|
||||
with _db_lock:
|
||||
if _names_cache[0] is database and _names_cache[1] == len(database):
|
||||
return _names_cache[2]
|
||||
names = list(dict.fromkeys(database.values()))
|
||||
if database is stockDB:
|
||||
for aliases in stockAliases.values():
|
||||
for alias in aliases:
|
||||
if alias not in names:
|
||||
names.append(alias)
|
||||
_names_cache[:] = [database, len(database), names]
|
||||
return names
|
||||
|
||||
|
||||
def _database_pinyins(
|
||||
database: Dict[str, str], names: Optional[List[str]] = None
|
||||
) -> List[str]:
|
||||
"""与 *names* 对齐的全拼小写拼音列表(带缓存)。
|
||||
|
||||
传入外层快照时按快照构建,保证与调用方本地 names 同源同序,
|
||||
避免并发合并期间两次取列表导致 zip 错位配对。
|
||||
"""
|
||||
if names is None:
|
||||
names = _database_names(database)
|
||||
with _db_lock:
|
||||
if _pinyin_cache[0] is not names:
|
||||
try:
|
||||
from pypinyin import lazy_pinyin
|
||||
|
||||
pinyins = ["".join(lazy_pinyin(name)).lower() for name in names]
|
||||
except Exception:
|
||||
pinyins = [""] * len(names)
|
||||
_pinyin_cache[:] = [names, pinyins]
|
||||
return _pinyin_cache[1]
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# 名称匹配辅助函数
|
||||
# =========================================================================
|
||||
|
||||
|
||||
def _exact_names(fragment: str, database: Dict[str, str]) -> List[str]:
|
||||
"""返回 database 名称中与 *fragment* 不区分大小写精确匹配的结果。"""
|
||||
fragment_lower = fragment.lower()
|
||||
return [name for name in _database_names(database) if name.lower() == fragment_lower]
|
||||
|
||||
|
||||
def _is_single_char_typo(input_name: str, candidate_name: str) -> bool:
|
||||
"""Return True when two names only differ by one character position."""
|
||||
"""当两个名称仅有一个字符位不同时返回 True。"""
|
||||
if not input_name or not candidate_name:
|
||||
return False
|
||||
if len(input_name) != len(candidate_name):
|
||||
@@ -135,6 +495,211 @@ def _is_single_char_typo(input_name: str, candidate_name: str) -> bool:
|
||||
return diff == 1
|
||||
|
||||
|
||||
# 拼音子串匹配的最短长度:拼音字母是逐音节的高频粒子,过短片段(o/k/ai/ma)
|
||||
# 会命中大量无关名称的拼音子串,把普通英文输入误解析成股票。
|
||||
_MIN_PINYIN_FRAGMENT_LEN = 5
|
||||
|
||||
|
||||
def _fix_name(fragment: str, database: Dict[str, str]) -> List[str]:
|
||||
"""通过混合匹配将名称片段解析为完整股票名称。
|
||||
|
||||
策略链(按顺序):
|
||||
1. 精确匹配(不区分大小写)
|
||||
2. 子串匹配(不区分大小写;片段含 >= 2 个 CJK 字符,如 茅台 -> 贵州茅台)
|
||||
3. 拼音子串匹配(如 "maotai" ⊂ "guizhoumaotai")
|
||||
4. difflib 模糊匹配(cutoff 0.8,另加 0.7 的保守单字符错别字回退,
|
||||
如 贵州茅苔 -> 贵州茅台)
|
||||
"""
|
||||
if not fragment:
|
||||
return []
|
||||
fragment_lower = fragment.lower()
|
||||
names = _database_names(database)
|
||||
|
||||
# 1. Exact match
|
||||
exact = [name for name in names if name.lower() == fragment_lower]
|
||||
if exact:
|
||||
return exact
|
||||
|
||||
# 2. Substring match
|
||||
if sum(1 for ch in fragment if "\u3400" <= ch <= "\u9fff") >= 2:
|
||||
matches = [name for name in names if fragment_lower in name.lower()]
|
||||
if matches:
|
||||
return matches
|
||||
|
||||
# 3. Pinyin substring match — 仅对纯 ASCII 片段生效
|
||||
# (设计场景是英文拼音输入:"maotai" ⊂ "guizhoumaotai")。CJK 片段
|
||||
# 转拼音后粒度失控:实义字 + 助词("阿里的" → "alide")恰好 5 字母
|
||||
# 过最短长度闸门,与不相干全名("alide" ⊂ "zhongdalide" 中大力德)
|
||||
# 发生子串碰撞;且 CJK 的合法场景已被其余三层覆盖(精确/子串/difflib)
|
||||
if _contains_cjk(fragment):
|
||||
fragment_pinyin = ""
|
||||
else:
|
||||
try:
|
||||
from pypinyin import lazy_pinyin
|
||||
|
||||
fragment_pinyin = "".join(lazy_pinyin(fragment)).lower()
|
||||
except Exception:
|
||||
fragment_pinyin = ""
|
||||
if len(fragment_pinyin) >= _MIN_PINYIN_FRAGMENT_LEN:
|
||||
matches = [
|
||||
name
|
||||
for name, pinyin in zip(names, _database_pinyins(database, names))
|
||||
if fragment_pinyin in pinyin
|
||||
]
|
||||
if matches:
|
||||
return matches
|
||||
|
||||
# 4. Fuzzy match
|
||||
if len(fragment) > 2:
|
||||
matches = difflib.get_close_matches(fragment, names, n=5, cutoff=0.8)
|
||||
if matches:
|
||||
return matches
|
||||
typo_matches = difflib.get_close_matches(fragment, names, n=5, cutoff=0.7)
|
||||
matches = [m for m in typo_matches if _is_single_char_typo(fragment, m)]
|
||||
if matches:
|
||||
return matches
|
||||
|
||||
return []
|
||||
|
||||
|
||||
def _infer_code_market(code: str) -> Optional[str]:
|
||||
"""根据代码格式推断市场:6 位数字→"a",5 位数字→"hk",字母→"us"。"""
|
||||
c = code.strip()
|
||||
if c.isdigit():
|
||||
return {5: "hk", 6: "a"}.get(len(c))
|
||||
return "us" if re.fullmatch(r"[A-Z]{1,5}(\.[A-Z])?", c) else None
|
||||
|
||||
|
||||
def _find_codes_for_name(name: str, database: Dict[str, str]) -> List[Tuple[str, str]]:
|
||||
"""查找 *database* 中映射到 *name* 的所有 (code, market) 组合。"""
|
||||
results: List[Tuple[str, str]] = []
|
||||
# 持锁迭代:与 extend_AkShare 的并发写串行
|
||||
with _db_lock:
|
||||
for code, mapped_name in database.items():
|
||||
if mapped_name == name or (
|
||||
database is stockDB and name in stockAliases.get(code, set())
|
||||
):
|
||||
market = _infer_code_market(code)
|
||||
if market:
|
||||
results.append((code, market))
|
||||
return results
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# 公开 API
|
||||
# =========================================================================
|
||||
|
||||
|
||||
def resolver_name_to_code_list(name: str) -> List[Stock]:
|
||||
"""将股票名称解析为匹配的 ``Stock`` 条目(最多 5 个)。暂不解析股票代码。
|
||||
|
||||
1. 对 CJK 输入,用 AkShare A 股数据扩充 ``stockDB``(幂等,带失败退避);
|
||||
非 CJK 输入完全跳过网络拉取。
|
||||
2. 在 ``stockDB`` 上做精确匹配。
|
||||
3. 在扩充后的数据库上做混合匹配(精确/子串/拼音/模糊)。
|
||||
|
||||
结果按 A 股 → 港股 → 美股 排序。
|
||||
|
||||
示例:
|
||||
"阿里巴巴" → [Stock("09988", "阿里巴巴", "hk"), Stock("BABA", "阿里巴巴", "us")]
|
||||
"茅台" → [Stock("600519", "贵州茅台", "a")]
|
||||
"你好世界" → []
|
||||
"""
|
||||
if not name or not isinstance(name, str):
|
||||
return []
|
||||
s = name.strip()
|
||||
# 单字符不可能是股票名(最短 2 字),短路避免全名表扫描
|
||||
if len(s) < 2:
|
||||
return []
|
||||
|
||||
# A 股名称全是 CJK —— 非 CJK 输入无法从网络扩充获益,但仍可参与本地
|
||||
# 混合(拼音)匹配。
|
||||
# 即使是本地精确命中也要先执行 AkShare 扩充:本地已有的同名港股/美股
|
||||
# 不能阻止 AkShare 中同名 A 股被并入,否则跨市场候选会不完整。
|
||||
if _contains_cjk(s):
|
||||
extend_AkShare()
|
||||
|
||||
full_names = _exact_names(s, stockDB)
|
||||
if not full_names:
|
||||
full_names = _fix_name(s, stockDB)
|
||||
if not full_names:
|
||||
return []
|
||||
|
||||
stocks: List[Stock] = []
|
||||
seen: Set[str] = set()
|
||||
for full_name in full_names:
|
||||
for code, market in _find_codes_for_name(full_name, stockDB):
|
||||
if code not in seen:
|
||||
seen.add(code)
|
||||
stocks.append(Stock(code=code, name=stockDB.get(code, full_name), market=market))
|
||||
|
||||
stocks.sort(key=lambda r: _MARKET_ORDER.get(r.market, 99))
|
||||
return stocks[:5]
|
||||
|
||||
|
||||
def US_stock_code_match(segment: str) -> List[Stock]:
|
||||
"""匹配美股代码:1~5 个英文字母的 ticker,仅本地库存在该代码时才返回。
|
||||
|
||||
避免把普通英文单词(hello/open/...)误判为股票代码。
|
||||
"""
|
||||
if not isinstance(segment, str) or not segment.isascii() or not segment.isalpha():
|
||||
return []
|
||||
if not 1 <= len(segment) <= 5:
|
||||
return []
|
||||
upper = segment.upper()
|
||||
name = stockDB.get(upper, "")
|
||||
return [Stock(code=upper, name=name, market="us")] if name else []
|
||||
|
||||
|
||||
def lookup_stock_by_code(code: str) -> Optional[Stock]:
|
||||
"""按规范化代码查库,返回完整 (code/name/market) 三元组;查不到返回 None。
|
||||
|
||||
格式合法不等于存在:stockDB 未命中即 None,绝不虚构(LLM 幻觉/过期
|
||||
代码不得直接注入 stocks)。本地库港股键为裸 5 位("00700"),HK 前缀
|
||||
形态("HK00700")自动补查去前缀裸键。
|
||||
"""
|
||||
c = (code or "").strip().upper()
|
||||
if not c:
|
||||
return None
|
||||
with _db_lock:
|
||||
name = stockDB.get(c)
|
||||
if not name and c.startswith("HK"):
|
||||
name = stockDB.get(c[2:])
|
||||
if not name:
|
||||
return None
|
||||
if c.startswith("HK"):
|
||||
market = "hk"
|
||||
else:
|
||||
market = _infer_code_market(c) or ""
|
||||
return Stock(code=c, name=name, market=market)
|
||||
|
||||
|
||||
def is_market_db_complete(market: str) -> bool:
|
||||
"""该市场名称库是否已扩展为全量(库未命中即可断定代码不存在)。
|
||||
|
||||
A 股:AkShare 全量列表已并入 stockDB(extend_AkShare 成功过,
|
||||
``_akshare_merged`` 非空)。港股/美股:本地精选库,永不视为全量,
|
||||
未命中只能存疑交下游判断。
|
||||
"""
|
||||
return market == "a" and _akshare_merged is not None
|
||||
|
||||
|
||||
def is_known_stock_name(name: str) -> bool:
|
||||
"""判断 *name* 是否为 stockDB 中已知的股票全名(含改名别名)。
|
||||
|
||||
只做本地名称表成员判定,绝不触发 AkShare 网络扩展。代码键
|
||||
(如 "600519")不算名称命中:分词管道(web_intent_tokenizer)用本
|
||||
函数保护"已识别名称 token",代码形文本必须继续进入代码提取步骤
|
||||
重新标注,而不是被当作名称跳过。
|
||||
"""
|
||||
if not isinstance(name, str):
|
||||
return False
|
||||
s = name.strip()
|
||||
if not s:
|
||||
return False
|
||||
return s in _database_names(stockDB)
|
||||
|
||||
|
||||
def resolve_name_to_code(name: str) -> Optional[str]:
|
||||
"""
|
||||
Resolve stock name to code.
|
||||
|
||||
@@ -8,19 +8,48 @@ Covers:
|
||||
- AkShare fallback (mocked)
|
||||
- Fuzzy match (difflib)
|
||||
- Ambiguous names return None
|
||||
- Stock dataclass / resolver_name_to_code_list / US_stock_code_match / extend_AkShare
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import threading
|
||||
import time
|
||||
from typing import Optional
|
||||
from unittest.mock import patch
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from src.data.stock_mapping import STOCK_NAME_MAP
|
||||
from src.services import name_to_code_resolver as ntc
|
||||
from src.services.name_to_code_resolver import (
|
||||
Stock,
|
||||
resolve_name_to_code,
|
||||
resolver_name_to_code_list,
|
||||
US_stock_code_match,
|
||||
_is_code_like,
|
||||
_normalize_code,
|
||||
_build_reverse_map_no_duplicates,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def clean_db(request):
|
||||
"""Isolate the global stockDB/caches; the AkShare fetch is mocked offline
|
||||
by default. Parametrize with ``indirect=True`` to inject a fake map."""
|
||||
fake_map = getattr(request, "param", None)
|
||||
with patch.object(ntc, "_get_akshare_name_to_code", return_value=fake_map):
|
||||
yield
|
||||
ntc.stockDB.clear()
|
||||
ntc.stockDB.update(STOCK_NAME_MAP)
|
||||
ntc._names_cache[:] = [None, None, None]
|
||||
ntc._pinyin_cache[:] = [None, None]
|
||||
ntc._akshare_merged = None
|
||||
ntc._akshare_cache = None
|
||||
ntc._akshare_failure_cache = None
|
||||
ntc._akshare_inflight = None
|
||||
ntc.stockAliases.clear()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _is_code_like
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -122,6 +151,23 @@ class TestResolveNameToCode:
|
||||
assert resolve_name_to_code("贵州茅台") == "600519"
|
||||
assert resolve_name_to_code("腾讯控股") == "00700"
|
||||
|
||||
@patch("src.services.name_to_code_resolver._get_akshare_name_to_code")
|
||||
def test_local_hit_does_not_trigger_akshare(self, mock_akshare):
|
||||
# 本地表精确命中必须零网络:既有调用方(API/Bot/导入)保持
|
||||
# 离线低延迟契约,不被 AkShare 冷启动等待拖住。
|
||||
assert resolve_name_to_code("贵州茅台") == "600519"
|
||||
assert resolve_name_to_code("腾讯控股") == "00700"
|
||||
mock_akshare.assert_not_called()
|
||||
|
||||
@patch("src.services.name_to_code_resolver._get_akshare_name_to_code")
|
||||
def test_local_hit_wins_over_akshare_same_name(self, mock_akshare):
|
||||
# 兼容性契约:本地表唯一命中的名字直接返回本地代码,不做跨市场
|
||||
# 合并判定(中国移动:本地仅港股 00941,AkShare 有同名 A 股 600941)。
|
||||
# 完整跨市场候选由 resolver_name_to_code_list 提供。
|
||||
mock_akshare.return_value = {"中国移动": "600941"}
|
||||
assert resolve_name_to_code("中国移动") == "00941"
|
||||
mock_akshare.assert_not_called()
|
||||
|
||||
def test_returns_none_for_empty_or_invalid_input(self):
|
||||
assert resolve_name_to_code("") is None
|
||||
assert resolve_name_to_code(" ") is None
|
||||
@@ -160,3 +206,565 @@ class TestResolveNameToCode:
|
||||
result = resolve_name_to_code("aaaaaaa")
|
||||
assert result is None
|
||||
mock_akshare.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize("clean_db", [{"三一重能": "688349"}], indirect=True)
|
||||
def test_akshare_exact_fallback_beats_fuzzy_for_non_local_name(self, clean_db):
|
||||
# 本地未收录的 CJK 名称:AkShare 精确命中(第 4 步)优先于模糊
|
||||
# 匹配(第 5 步),"三一重能" 不会被误配到相近的 三一重工。
|
||||
assert resolve_name_to_code("三一重能") == "688349"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stock dataclass
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestStock:
|
||||
def test_fields(self):
|
||||
s = Stock(code="600519", name="贵州茅台", market="a")
|
||||
assert (s.code, s.name, s.market) == ("600519", "贵州茅台", "a")
|
||||
|
||||
def test_value_equality(self):
|
||||
assert Stock("600519", "贵州茅台", "a") == Stock("600519", "贵州茅台", "a")
|
||||
assert Stock("600519", "贵州茅台", "a") != Stock("00700", "腾讯控股", "hk")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# resolver_name_to_code_list
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestResolverNameToCodeList:
|
||||
@pytest.mark.usefixtures("clean_db")
|
||||
def test_exact_match(self):
|
||||
assert resolver_name_to_code_list("贵州茅台") == [Stock("600519", "贵州茅台", "a")]
|
||||
|
||||
@pytest.mark.usefixtures("clean_db")
|
||||
def test_exact_match_cross_market_sorted(self):
|
||||
# 阿里巴巴 in STOCK_NAME_MAP: BABA (us) + 09988 (hk) → hk before us
|
||||
assert resolver_name_to_code_list("阿里巴巴") == [
|
||||
Stock("09988", "阿里巴巴", "hk"),
|
||||
Stock("BABA", "阿里巴巴", "us"),
|
||||
]
|
||||
|
||||
@pytest.mark.usefixtures("clean_db")
|
||||
def test_substring_match(self):
|
||||
assert resolver_name_to_code_list("茅台") == [Stock("600519", "贵州茅台", "a")]
|
||||
|
||||
@pytest.mark.usefixtures("clean_db")
|
||||
def test_pinyin_substring_match(self):
|
||||
# Non-CJK input: resolved locally via pinyin, without AkShare fetch
|
||||
assert resolver_name_to_code_list("maotai") == [Stock("600519", "贵州茅台", "a")]
|
||||
|
||||
@pytest.mark.usefixtures("clean_db")
|
||||
def test_cjk_fragment_skips_pinyin_layer(self):
|
||||
# 回归:CJK 片段不得走拼音子串层。"平果"拼音 pingguo 与苹果全名
|
||||
# 拼音完全相撞,预修复会经策略 3 误命中苹果;CJK 片段的拼音粒度
|
||||
# 失控(实义字+助词转拼音后可与不相干全名碰撞,如"阿里的"→
|
||||
# "alide" ⊂ "zhongdalide" 中大力德),仅 ASCII 拼音输入走该层。
|
||||
assert resolver_name_to_code_list("平果") == []
|
||||
|
||||
@pytest.mark.usefixtures("clean_db")
|
||||
def test_fuzzy_typo_match(self):
|
||||
assert resolver_name_to_code_list("贵州茅苔") == [Stock("600519", "贵州茅台", "a")]
|
||||
|
||||
@pytest.mark.usefixtures("clean_db")
|
||||
def test_no_match_returns_empty(self):
|
||||
assert resolver_name_to_code_list("你好世界") == []
|
||||
|
||||
@pytest.mark.usefixtures("clean_db")
|
||||
def test_invalid_input_returns_empty(self):
|
||||
assert resolver_name_to_code_list("") == []
|
||||
assert resolver_name_to_code_list(None) == [] # type: ignore
|
||||
assert resolver_name_to_code_list("茅") == [] # single char is never a name
|
||||
|
||||
@pytest.mark.parametrize("clean_db", [{"浦发银行": "600000"}], indirect=True)
|
||||
def test_akshare_extension_visible_after_retry(self, clean_db):
|
||||
# Exact match against the AkShare-extended database
|
||||
assert resolver_name_to_code_list("浦发银行") == [Stock("600000", "浦发银行", "a")]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("clean_db", [{"阿里巴巴": "600000"}], indirect=True)
|
||||
def test_local_exact_hit_still_merges_akshare_same_name_a_share(self, clean_db):
|
||||
# 本地已有同名港股/美股时,AkShare 中的同名 A 股也必须被并入。
|
||||
ntc.stockDB.clear()
|
||||
ntc.stockDB.update({"09988": "阿里巴巴", "BABA": "阿里巴巴"})
|
||||
ntc._names_cache[:] = [None, None, None]
|
||||
ntc._pinyin_cache[:] = [None, None]
|
||||
ntc._akshare_merged = None
|
||||
ntc.stockAliases.clear()
|
||||
assert resolver_name_to_code_list("阿里巴巴") == [
|
||||
Stock("600000", "阿里巴巴", "a"),
|
||||
Stock("09988", "阿里巴巴", "hk"),
|
||||
Stock("BABA", "阿里巴巴", "us"),
|
||||
]
|
||||
|
||||
@pytest.mark.parametrize("clean_db", [{"阿里巴巴": "600000"}], indirect=True)
|
||||
def test_local_single_candidate_gets_a_share_candidate_after_akshare_merge(self, clean_db):
|
||||
# 本地只有单一市场记录时,AkShare 补齐同名 A 股后候选变完整。
|
||||
ntc.stockDB.clear()
|
||||
ntc.stockDB.update({"09988": "阿里巴巴"})
|
||||
ntc._names_cache[:] = [None, None, None]
|
||||
ntc._pinyin_cache[:] = [None, None]
|
||||
ntc._akshare_merged = None
|
||||
ntc.stockAliases.clear()
|
||||
assert resolver_name_to_code_list("阿里巴巴") == [
|
||||
Stock("600000", "阿里巴巴", "a"),
|
||||
Stock("09988", "阿里巴巴", "hk"),
|
||||
]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# US_stock_code_match
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestUSStockCodeMatch:
|
||||
def test_known_ticker(self):
|
||||
assert US_stock_code_match("AAPL") == [Stock("AAPL", "苹果", "us")]
|
||||
assert US_stock_code_match("aapl") == [Stock("AAPL", "苹果", "us")]
|
||||
|
||||
def test_unknown_word_returns_empty(self):
|
||||
assert US_stock_code_match("HELLO") == [] # ordinary English word
|
||||
assert US_stock_code_match("TOOLONGTICKER") == []
|
||||
assert US_stock_code_match("贵州茅台") == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# extend_AkShare
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestExtendAkShare:
|
||||
@pytest.mark.parametrize("clean_db", [{"浦发银行": "600000"}], indirect=True)
|
||||
def test_merges_new_entries_idempotently(self, clean_db):
|
||||
assert ntc.extend_AkShare() is True
|
||||
assert ntc.stockDB["600000"] == "浦发银行"
|
||||
# Same cached map object is not merged twice
|
||||
assert ntc.extend_AkShare() is False
|
||||
|
||||
@pytest.mark.parametrize("clean_db", [{"贵州茅台": "600519"}], indirect=True)
|
||||
def test_no_new_entries_returns_false(self, clean_db):
|
||||
# All entries already in the local database
|
||||
assert ntc.extend_AkShare() is False
|
||||
|
||||
@pytest.mark.usefixtures("clean_db")
|
||||
def test_fetch_failure_returns_false(self):
|
||||
assert ntc.extend_AkShare() is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("clean_db", [{"新名称": "600000"}], indirect=True)
|
||||
def test_rename_existing_code_updates_canonical_name_and_keeps_alias(self, clean_db):
|
||||
ntc.stockDB.clear()
|
||||
ntc.stockDB.update({"600000": "旧名称"})
|
||||
ntc._names_cache[:] = [None, None, None]
|
||||
ntc._pinyin_cache[:] = [None, None]
|
||||
ntc._akshare_merged = None
|
||||
ntc.stockAliases.clear()
|
||||
assert ntc.extend_AkShare() is True
|
||||
assert ntc.stockDB["600000"] == "新名称"
|
||||
assert ntc.stockAliases["600000"] == {"旧名称"}
|
||||
# 新名称作为当前官方名称可解析
|
||||
assert resolver_name_to_code_list("新名称") == [Stock("600000", "新名称", "a")]
|
||||
# 旧名称作为别名仍然可解析,且展示当前官方名称
|
||||
assert resolver_name_to_code_list("旧名称") == [Stock("600000", "新名称", "a")]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AkShare 单飞并发:真实 _get_akshare_name_to_code + _fetch_akshare_df 假拉取
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeAkShareFetch:
|
||||
"""确定性的 _fetch_akshare_df 替身。
|
||||
|
||||
每次调用先置位 fetch_started(证明拉取已在后台线程开始),再阻塞在
|
||||
release_fetch 上——测试据此精确控制拉取窗口,不依赖 sleep 计时。
|
||||
"""
|
||||
|
||||
def __init__(self, fail: bool = False, rows: Optional[dict] = None):
|
||||
self.fetch_started = threading.Event()
|
||||
self.release_fetch = threading.Event()
|
||||
self.fail = fail
|
||||
self.rows = rows or {"code": ["600000"], "name": ["浦发银行"]}
|
||||
self.calls = 0
|
||||
|
||||
def fetch(self):
|
||||
self.calls += 1
|
||||
self.fetch_started.set()
|
||||
# 30s 兜底:用例逻辑正确时总会在收尾前 set,防自身缺陷挂死线程
|
||||
self.release_fetch.wait(timeout=30)
|
||||
if self.fail:
|
||||
raise RuntimeError("simulated akshare failure")
|
||||
return pd.DataFrame(self.rows)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def real_akshare_path(monkeypatch, tmp_path):
|
||||
"""复位 AkShare 拉取相关全局态,走真实 _get_akshare_name_to_code 路径。
|
||||
|
||||
磁盘缓存指向临时路径:既隔离仓库 data/cache 下的真实文件,也让
|
||||
落盘/懒加载用例可以安全读写。
|
||||
"""
|
||||
saved = (
|
||||
ntc._akshare_cache,
|
||||
ntc._akshare_failure_cache,
|
||||
ntc._akshare_merged,
|
||||
ntc._akshare_inflight,
|
||||
ntc._akshare_disk_checked,
|
||||
)
|
||||
ntc._akshare_cache = None
|
||||
ntc._akshare_failure_cache = None
|
||||
ntc._akshare_merged = None
|
||||
ntc._akshare_inflight = None
|
||||
ntc._akshare_disk_checked = False
|
||||
monkeypatch.setattr(ntc, "_AKSHARE_DISK_CACHE_PATH", tmp_path / "akshare_name_map.json")
|
||||
yield monkeypatch
|
||||
# 等待在途刷新收尾(正常用例在 finally 里已 release,这里毫秒级通过),
|
||||
# 避免迟到的 worker 把缓存写进下一个用例
|
||||
deadline = time.time() + 5
|
||||
while ntc._akshare_inflight is not None and time.time() < deadline:
|
||||
time.sleep(0.05)
|
||||
ntc._akshare_cache, ntc._akshare_failure_cache, ntc._akshare_merged = saved[:3]
|
||||
ntc._akshare_inflight, ntc._akshare_disk_checked = saved[3:]
|
||||
ntc.stockDB.clear()
|
||||
ntc.stockDB.update(STOCK_NAME_MAP)
|
||||
ntc._names_cache[:] = [None, None, None]
|
||||
ntc._pinyin_cache[:] = [None, None]
|
||||
ntc.stockAliases.clear()
|
||||
|
||||
|
||||
def _run_resolver_in_thread(query: str):
|
||||
"""在 daemon 线程里执行解析,捕获返回值/异常。"""
|
||||
outcome = {}
|
||||
|
||||
def _worker():
|
||||
try:
|
||||
outcome["value"] = ntc.resolver_name_to_code_list(query)
|
||||
except BaseException as exc: # noqa: BLE001 - 线程异常兜底
|
||||
outcome["error"] = exc
|
||||
|
||||
t = threading.Thread(target=_worker, daemon=True)
|
||||
t.start()
|
||||
return t, outcome
|
||||
|
||||
|
||||
def _wait_until(condition, timeout: float = 5.0) -> bool:
|
||||
"""有界轮询等待条件成立(后台刷新收尾是确定性事件,仅耗时不定)。"""
|
||||
deadline = time.time() + timeout
|
||||
while time.time() < deadline:
|
||||
if condition():
|
||||
return True
|
||||
time.sleep(0.02)
|
||||
return condition()
|
||||
|
||||
|
||||
class TestAkShareSingleFlightConcurrency:
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_waiter_resolves_after_inflight_fetch_completes(self, monkeypatch):
|
||||
fake = _FakeAkShareFetch()
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fake.fetch)
|
||||
|
||||
t1, r1 = _run_resolver_in_thread("浦发银行")
|
||||
t2 = None
|
||||
try:
|
||||
assert fake.fetch_started.wait(timeout=5)
|
||||
t2, r2 = _run_resolver_in_thread("浦发银行")
|
||||
# 旧实现(非阻塞放弃)在此窗口内立即返回空结果;新实现持续等待
|
||||
t2.join(timeout=2)
|
||||
# 释放前两个线程必须仍在等待(而非已带空结果返回)——漏解 bug 的核心
|
||||
assert t1.is_alive() and t2.is_alive()
|
||||
finally:
|
||||
fake.release_fetch.set()
|
||||
t1.join(timeout=10)
|
||||
if t2 is not None:
|
||||
t2.join(timeout=10)
|
||||
assert not t1.is_alive() and not t2.is_alive()
|
||||
# 等待者命中在途拉取的结果(本 bug 的核心断言),且全进程只拉取一次
|
||||
assert r1.get("value") == [Stock("600000", "浦发银行", "a")]
|
||||
assert r2.get("value") == [Stock("600000", "浦发银行", "a")]
|
||||
assert fake.calls == 1
|
||||
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_waiter_returns_empty_after_inflight_failure_and_no_refetch(self, monkeypatch):
|
||||
fake = _FakeAkShareFetch(fail=True)
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fake.fetch)
|
||||
|
||||
t1, r1 = _run_resolver_in_thread("浦发银行")
|
||||
t2 = None
|
||||
try:
|
||||
assert fake.fetch_started.wait(timeout=5)
|
||||
t2, r2 = _run_resolver_in_thread("浦发银行")
|
||||
t2.join(timeout=2)
|
||||
finally:
|
||||
fake.release_fetch.set()
|
||||
t1.join(timeout=10)
|
||||
if t2 is not None:
|
||||
t2.join(timeout=10)
|
||||
assert not t1.is_alive() and not t2.is_alive()
|
||||
# 在途拉取失败:等待者醒来命中失败退避(而非重试风暴),本地空库返回 []
|
||||
assert r1.get("value") == []
|
||||
assert r2.get("value") == []
|
||||
assert fake.calls == 1
|
||||
assert ntc._akshare_failure_cache is not None
|
||||
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_cold_waiters_degrade_after_timeout_but_fetch_lands(self, monkeypatch):
|
||||
# worker 违约(超清账余量仍未 resolve)时的兜底路径:等待者按死线
|
||||
# 自行降级,后台拉取不受影响仍单次完成落地
|
||||
monkeypatch.setattr(ntc, "_AKSHARE_WAIT_COLD_START", 0.2)
|
||||
fake = _FakeAkShareFetch()
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fake.fetch)
|
||||
|
||||
t1, r1 = _run_resolver_in_thread("浦发银行")
|
||||
t2 = None
|
||||
try:
|
||||
assert fake.fetch_started.wait(timeout=5)
|
||||
t2, r2 = _run_resolver_in_thread("浦发银行")
|
||||
# 拉取持续挂起时,所有冷启动等待者(含触发拉取的那一个)都必须
|
||||
# 在超时上界内自行返回,而非无限等待
|
||||
t2.join(timeout=10)
|
||||
finally:
|
||||
fake.release_fetch.set()
|
||||
t1.join(timeout=10)
|
||||
if t2 is not None:
|
||||
t2.join(timeout=10)
|
||||
assert not t1.is_alive() and not t2.is_alive()
|
||||
# 超时退化为本地库解析;后台拉取不受等待者超时影响,仍单次完成
|
||||
assert r1.get("value") == []
|
||||
assert r2.get("value") == []
|
||||
assert fake.calls == 1
|
||||
assert _wait_until(
|
||||
lambda: ntc._akshare_cache is not None
|
||||
and ntc._akshare_cache[1] == {"浦发银行": "600000"}
|
||||
)
|
||||
# 拉取落地后,后续请求立即命中新缓存
|
||||
assert ntc.resolver_name_to_code_list("浦发银行") == [Stock("600000", "浦发银行", "a")]
|
||||
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_slow_persist_does_not_delay_waiter_wakeup(self, monkeypatch):
|
||||
# 唤醒在落盘之前:慢磁盘不得拖住等待者拿结果(缓存已就绪、
|
||||
# 等待者却超时漏解曾是真实回归)
|
||||
fake = _FakeAkShareFetch()
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fake.fetch)
|
||||
release_persist = threading.Event()
|
||||
monkeypatch.setattr(
|
||||
ntc, "_persist_akshare_map", lambda map_: release_persist.wait(timeout=30)
|
||||
)
|
||||
|
||||
t1, r1 = _run_resolver_in_thread("浦发银行")
|
||||
try:
|
||||
assert fake.fetch_started.wait(timeout=5)
|
||||
fake.release_fetch.set()
|
||||
# 拉取完成即唤醒并返回,不等落盘
|
||||
t1.join(timeout=5)
|
||||
assert not t1.is_alive()
|
||||
assert r1.get("value") == [Stock("600000", "浦发银行", "a")]
|
||||
assert not release_persist.is_set() # 落盘仍挂着,等待者已返回
|
||||
finally:
|
||||
release_persist.set()
|
||||
t1.join(timeout=10)
|
||||
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_base_exception_still_clears_inflight_and_arms_backoff(self, monkeypatch):
|
||||
# finally 兜底:BaseException(如 KeyboardInterrupt/SystemExit)逃逸
|
||||
# 时也必须清在途句柄、武装退避并唤醒等待者,不留死句柄
|
||||
class _Bomb(BaseException):
|
||||
pass
|
||||
|
||||
def bomb():
|
||||
raise _Bomb
|
||||
|
||||
# 接管 excepthook:既静音 worker 死亡时的未处理异常输出,又给出
|
||||
# "worker 已带着 _Bomb 死亡"的确定性信号(不依赖全局线程状态,
|
||||
# 避免被其他用例泄漏的后台刷新线程干扰)
|
||||
hook_calls: list = []
|
||||
monkeypatch.setattr(threading, "excepthook", lambda args: hook_calls.append(args))
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", bomb)
|
||||
t1, r1 = _run_resolver_in_thread("浦发银行")
|
||||
t1.join(timeout=10)
|
||||
assert not t1.is_alive()
|
||||
assert _wait_until(lambda: hook_calls) # worker 已触发 excepthook
|
||||
assert r1.get("value") == [] # 等待者被 finally 唤醒后拿到 None
|
||||
assert ntc._akshare_inflight is None # 无死句柄
|
||||
assert ntc._akshare_failure_cache is not None # 退避已武装
|
||||
assert ntc._get_akshare_name_to_code() is None # 退避窗口内不再触网
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stale-while-revalidate:TTL 过期必须立即返回旧值(零等待),刷新在后台
|
||||
# 完成;网络故障的退避窗口内旧值继续服务(可用性),仅冷启动才退化。
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAkShareStaleWhileRevalidate:
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_stale_served_immediately_while_refresh_inflight(self, monkeypatch):
|
||||
# 先用立即返回的假拉取把缓存填充为 v1
|
||||
primed = {"浦发银行": "600000"}
|
||||
|
||||
def fetch_v1():
|
||||
return pd.DataFrame({"code": ["600000"], "name": ["浦发银行"]})
|
||||
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fetch_v1)
|
||||
assert ntc._get_akshare_name_to_code() == primed
|
||||
# 强制过期:时间戳拨回 TTL 之前
|
||||
ntc._akshare_cache = (time.time() - ntc._AKSHARE_CACHE_TTL - 10, primed)
|
||||
|
||||
# 换成阻塞版拉取(v2:600000 改名为 新名称银行)
|
||||
fake = _FakeAkShareFetch(rows={"code": ["600000"], "name": ["新名称银行"]})
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fake.fetch)
|
||||
|
||||
started = time.monotonic()
|
||||
result = ntc._get_akshare_name_to_code()
|
||||
elapsed = time.monotonic() - started
|
||||
try:
|
||||
# 核心断言:stale 值立即返回,没有等待在途拉取
|
||||
assert result == primed
|
||||
assert elapsed < 2
|
||||
# 后台刷新确实已发起且尚未完成(fake 仍被挂起)
|
||||
assert fake.fetch_started.wait(timeout=5)
|
||||
assert not fake.release_fetch.is_set()
|
||||
finally:
|
||||
fake.release_fetch.set()
|
||||
# 刷新收尾后新缓存可见,且全进程只拉取一次
|
||||
v2 = {"新名称银行": "600000"}
|
||||
assert _wait_until(lambda: ntc._akshare_cache is not None and ntc._akshare_cache[1] == v2)
|
||||
assert fake.calls == 1
|
||||
assert ntc._get_akshare_name_to_code() == v2
|
||||
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_backoff_window_serves_stale_during_outage(self, monkeypatch):
|
||||
# 预热缓存 v1 并强制过期
|
||||
primed = {"浦发银行": "600000"}
|
||||
|
||||
def fetch_ok():
|
||||
return pd.DataFrame({"code": ["600000"], "name": ["浦发银行"]})
|
||||
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fetch_ok)
|
||||
assert ntc._get_akshare_name_to_code() == primed
|
||||
ntc._akshare_cache = (time.time() - ntc._AKSHARE_CACHE_TTL - 10, primed)
|
||||
|
||||
# 换成立即失败的拉取:第一次 stale 调用会发起后台刷新并失败
|
||||
calls = {"n": 0}
|
||||
|
||||
def fetch_fail():
|
||||
calls["n"] += 1
|
||||
raise RuntimeError("simulated outage")
|
||||
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fetch_fail)
|
||||
assert ntc._get_akshare_name_to_code() == primed # 故障期间旧值仍可用
|
||||
assert _wait_until(lambda: ntc._akshare_failure_cache is not None)
|
||||
# 退避窗口内再次调用:继续服务 stale,且不再发起新的拉取
|
||||
assert ntc._get_akshare_name_to_code() == primed
|
||||
assert calls["n"] == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 磁盘缓存:成功拉取后原子落盘;重启(全局态复位)后懒加载,零网络。
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAkShareDiskCache:
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_persist_and_reload_without_network(self, monkeypatch):
|
||||
def fetch_ok():
|
||||
return pd.DataFrame({"code": ["600000"], "name": ["浦发银行"]})
|
||||
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fetch_ok)
|
||||
assert ntc._get_akshare_name_to_code() == {"浦发银行": "600000"}
|
||||
assert _wait_until(lambda: ntc._AKSHARE_DISK_CACHE_PATH.is_file())
|
||||
|
||||
# 模拟重启:内存态全部复位,磁盘保留
|
||||
ntc._akshare_cache = None
|
||||
ntc._akshare_failure_cache = None
|
||||
ntc._akshare_disk_checked = False
|
||||
|
||||
never = _FakeAkShareFetch()
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", never.fetch)
|
||||
# 懒加载直接命中(落盘时间戳新鲜),全程零网络
|
||||
assert ntc._get_akshare_name_to_code() == {"浦发银行": "600000"}
|
||||
assert never.calls == 0
|
||||
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_corrupt_disk_cache_ignored(self, monkeypatch):
|
||||
ntc._AKSHARE_DISK_CACHE_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
ntc._AKSHARE_DISK_CACHE_PATH.write_text("not-a-json{", encoding="utf-8")
|
||||
|
||||
calls = {"n": 0}
|
||||
|
||||
def fetch_ok():
|
||||
calls["n"] += 1
|
||||
return pd.DataFrame({"code": ["600000"], "name": ["浦发银行"]})
|
||||
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fetch_ok)
|
||||
# 损坏的磁盘缓存被忽略,正常走冷启动拉取
|
||||
assert ntc._get_akshare_name_to_code() == {"浦发银行": "600000"}
|
||||
assert calls["n"] == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 预热:幂等——并发/重复调用共享同一在途句柄,只触发一次网络拉取。
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWarmupIdempotent:
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_double_warmup_shares_single_fetch(self, monkeypatch):
|
||||
fake = _FakeAkShareFetch()
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", fake.fetch)
|
||||
|
||||
ntc.warmup_akshare_cache()
|
||||
assert fake.fetch_started.wait(timeout=5) # 在途句柄已登记
|
||||
ntc.warmup_akshare_cache() # 第二次:共享 Future,不再拉取
|
||||
fake.release_fetch.set()
|
||||
|
||||
assert _wait_until(
|
||||
lambda: ntc._akshare_cache is not None
|
||||
and ntc._akshare_cache[1] == {"浦发银行": "600000"}
|
||||
)
|
||||
assert fake.calls == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 挂起(非失败)场景:子进程超时包装以 TimeoutError 抛出后,必须走
|
||||
# 常规失败路径武装退避,而非无限持有单飞锁。
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAkShareHangBackoff:
|
||||
@pytest.mark.usefixtures("real_akshare_path")
|
||||
def test_timeout_error_arms_failure_backoff(self, monkeypatch):
|
||||
calls = []
|
||||
|
||||
def hang():
|
||||
calls.append(1)
|
||||
# 与 _akshare_call_with_timeout 超时抛出的异常同型
|
||||
raise TimeoutError("stock_info_a_code_name 调用超过 25s,已放弃等待")
|
||||
|
||||
monkeypatch.setattr(ntc, "_fetch_akshare_df", hang)
|
||||
assert ntc._get_akshare_name_to_code() is None
|
||||
assert ntc._akshare_failure_cache is not None # 退避已武装
|
||||
# 退避窗口内再次解析:命中失败缓存快路径,不再触网
|
||||
assert ntc._get_akshare_name_to_code() is None
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _fetch_akshare_df 接线:必须经子进程超时包装调用 worker(挂起封顶)。
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFetchAkshareDfWiring:
|
||||
def test_delegates_to_subprocess_timeout_wrapper(self, monkeypatch):
|
||||
import data_provider.akshare_fetcher as af
|
||||
|
||||
captured = {}
|
||||
|
||||
def fake_wrapper(func, *args, **kwargs):
|
||||
captured["func"] = func
|
||||
captured["timeout"] = kwargs.get("timeout")
|
||||
captured["call_name"] = kwargs.get("call_name")
|
||||
return pd.DataFrame({"code": ["600000"], "name": ["浦发银行"]})
|
||||
|
||||
monkeypatch.setattr(af, "_akshare_call_with_timeout", fake_wrapper)
|
||||
df = ntc._fetch_akshare_df()
|
||||
assert captured["func"] is ntc._akshare_stock_info_worker
|
||||
assert captured["timeout"] == ntc._AKSHARE_FETCH_TIMEOUT
|
||||
assert captured["call_name"] == "stock_info_a_code_name"
|
||||
assert list(df["name"]) == ["浦发银行"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user