股票实体解析增强 (#2245)

Co-authored-by: zhulinsen <42829555+ZhuLinsen@users.noreply.github.com>
This commit is contained in:
Gach-Coder
2026-08-22 21:23:55 +08:00
committed by GitHub
parent cd1fd2229c
commit f6b719d1fe
4 changed files with 1218 additions and 39 deletions

View File

@@ -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:

View File

@@ -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 预算只作不可突破的外层 capresearch 路径不再传 `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()` 匹配美股 ticker1~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 PR2issue #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

View File

@@ -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 名掺入 PIDweb 与 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_STARTTTL 过期
的请求零等待;网络 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-revalidateTTL 的语义就是容忍陈旧,过期瞬间等待
# 毫无收益;返回旧值,下轮调用命中后台刷新出的新缓存。
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 全量列表已并入 stockDBextend_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.

View File

@@ -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):
# 兼容性契约:本地表唯一命中的名字直接返回本地代码,不做跨市场
# 合并判定(中国移动:本地仅港股 00941AkShare 有同名 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-revalidateTTL 过期必须立即返回旧值(零等待),刷新在后台
# 完成;网络故障的退避窗口内旧值继续服务(可用性),仅冷启动才退化。
# ---------------------------------------------------------------------------
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)
# 换成阻塞版拉取v2600000 改名为 新名称银行)
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"]) == ["浦发银行"]