fix: 修复并发执行时共享状态缺少统一加锁的问题 (#881) (#928)

* fix(issue-881): [bug]-修复并发执行时共享状态缺少统一加锁的问题
This commit is contained in:
mumu
2026-03-31 22:53:56 +08:00
committed by GitHub
parent 2134b47c3f
commit c89fd9cf20
8 changed files with 969 additions and 245 deletions
+125 -60
View File
@@ -490,6 +490,11 @@ class DataFetcherManager:
fetchers: 数据源列表(可选,默认按优先级自动创建)
"""
self._fetchers: List[BaseFetcher] = []
self._fetchers_lock = RLock()
self._fetcher_call_locks: Dict[int, RLock] = {}
self._fetcher_call_locks_lock = RLock()
self._stock_name_cache: Dict[str, str] = {}
self._stock_name_cache_lock = RLock()
if fetchers:
# 按优先级排序
@@ -506,6 +511,53 @@ class DataFetcherManager:
self._fundamental_timeout_worker_limit = 8
self._fundamental_timeout_slots = BoundedSemaphore(self._fundamental_timeout_worker_limit)
def _ensure_concurrency_guards(self) -> None:
"""Lazily initialize thread-safety primitives for test scaffolds using __new__."""
if not hasattr(self, "_fetchers_lock") or self._fetchers_lock is None:
self._fetchers_lock = RLock()
if not hasattr(self, "_fetcher_call_locks") or self._fetcher_call_locks is None:
self._fetcher_call_locks = {}
if not hasattr(self, "_fetcher_call_locks_lock") or self._fetcher_call_locks_lock is None:
self._fetcher_call_locks_lock = RLock()
if not hasattr(self, "_stock_name_cache") or self._stock_name_cache is None:
self._stock_name_cache = {}
if not hasattr(self, "_stock_name_cache_lock") or self._stock_name_cache_lock is None:
self._stock_name_cache_lock = RLock()
def _get_fetchers_snapshot(self) -> List[BaseFetcher]:
self._ensure_concurrency_guards()
with self._fetchers_lock:
return list(getattr(self, "_fetchers", []))
def _get_fetcher_call_lock(self, fetcher: BaseFetcher) -> RLock:
self._ensure_concurrency_guards()
fetcher_id = id(fetcher)
with self._fetcher_call_locks_lock:
lock = self._fetcher_call_locks.get(fetcher_id)
if lock is None:
lock = RLock()
self._fetcher_call_locks[fetcher_id] = lock
return lock
def _call_fetcher_method(self, fetcher: BaseFetcher, method_name: str, *args, **kwargs):
"""Serialize shared fetcher state access through manager-owned per-instance locks."""
method = getattr(fetcher, method_name)
with self._get_fetcher_call_lock(fetcher):
return method(*args, **kwargs)
def _get_cached_stock_name(self, stock_code: str) -> Optional[str]:
self._ensure_concurrency_guards()
with self._stock_name_cache_lock:
return self._stock_name_cache.get(stock_code)
def _cache_stock_name(self, stock_code: str, name: Optional[str]) -> Optional[str]:
if name is None:
return None
self._ensure_concurrency_guards()
with self._stock_name_cache_lock:
self._stock_name_cache[stock_code] = name
return name
def _get_tickflow_fetcher(self):
"""Lazily create a TickFlow fetcher for market-review-only calls."""
from src.config import get_config
@@ -767,26 +819,30 @@ class DataFetcherManager:
yfinance = YfinanceFetcher()
# 初始化数据源列表
self._fetchers = [
efinance,
akshare,
tushare,
pytdx,
baostock,
yfinance,
]
self._ensure_concurrency_guards()
with self._fetchers_lock:
self._fetchers = [
efinance,
akshare,
tushare,
pytdx,
baostock,
yfinance,
]
# 按优先级排序(Tushare 如果配置了 Token 且初始化成功,优先级为 0)
self._fetchers.sort(key=lambda f: f.priority)
# 按优先级排序(Tushare 如果配置了 Token 且初始化成功,优先级为 0)
self._fetchers.sort(key=lambda f: f.priority)
# 构建优先级说明
priority_info = ", ".join([f"{f.name}(P{f.priority})" for f in self._fetchers])
priority_info = ", ".join([f"{f.name}(P{f.priority})" for f in self._get_fetchers_snapshot()])
logger.info(f"已初始化 {len(self._fetchers)} 个数据源(按优先级): {priority_info}")
def add_fetcher(self, fetcher: BaseFetcher) -> None:
"""添加数据源并重新排序"""
self._fetchers.append(fetcher)
self._fetchers.sort(key=lambda f: f.priority)
self._ensure_concurrency_guards()
with self._fetchers_lock:
self._fetchers.append(fetcher)
self._fetchers.sort(key=lambda f: f.priority)
def get_daily_data(
self,
@@ -822,20 +878,23 @@ class DataFetcherManager:
# Normalize code (strip SH/SZ prefix etc.)
stock_code = normalize_stock_code(stock_code)
fetchers = self._get_fetchers_snapshot()
errors = []
total_fetchers = len(self._fetchers)
total_fetchers = len(fetchers)
request_start = time.time()
# 快速路径:美股指数与美股股票直接路由到 YfinanceFetcher
if is_us_index_code(stock_code) or is_us_stock_code(stock_code):
for attempt, fetcher in enumerate(self._fetchers, start=1):
for attempt, fetcher in enumerate(fetchers, start=1):
if fetcher.name == "YfinanceFetcher":
try:
logger.info(
f"[数据源尝试 {attempt}/{total_fetchers}] [{fetcher.name}] "
f"美股/美股指数 {stock_code} 直接路由..."
)
df = fetcher.get_daily_data(
df = self._call_fetcher_method(
fetcher,
"get_daily_data",
stock_code=stock_code,
start_date=start_date,
end_date=end_date,
@@ -863,10 +922,12 @@ class DataFetcherManager:
logger.error(f"[数据源终止] {stock_code} 获取失败: elapsed={elapsed:.2f}s\n{error_summary}")
raise DataFetchError(error_summary)
for attempt, fetcher in enumerate(self._fetchers, start=1):
for attempt, fetcher in enumerate(fetchers, start=1):
try:
logger.info(f"[数据源尝试 {attempt}/{total_fetchers}] [{fetcher.name}] 获取 {stock_code}...")
df = fetcher.get_daily_data(
df = self._call_fetcher_method(
fetcher,
"get_daily_data",
stock_code=stock_code,
start_date=start_date,
end_date=end_date,
@@ -890,7 +951,7 @@ class DataFetcherManager:
)
errors.append(error_msg)
if attempt < total_fetchers:
next_fetcher = self._fetchers[attempt]
next_fetcher = fetchers[attempt]
logger.info(f"[数据源切换] {stock_code}: [{fetcher.name}] -> [{next_fetcher.name}]")
# 继续尝试下一个数据源
continue
@@ -904,7 +965,7 @@ class DataFetcherManager:
@property
def available_fetchers(self) -> List[str]:
"""返回可用数据源名称列表"""
return [f.name for f in self._fetchers]
return [f.name for f in self._get_fetchers_snapshot()]
def prefetch_realtime_quotes(self, stock_codes: List[str]) -> int:
"""
@@ -1022,11 +1083,11 @@ class DataFetcherManager:
# 美股指数由 YfinanceFetcher 处理(在美股股票检查之前)
if is_us_index_code(stock_code):
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
if fetcher.name == "YfinanceFetcher":
if hasattr(fetcher, 'get_realtime_quote'):
try:
quote = fetcher.get_realtime_quote(stock_code)
quote = self._call_fetcher_method(fetcher, 'get_realtime_quote', stock_code)
if quote is not None:
logger.info(f"[实时行情] 美股指数 {stock_code} 成功获取 (来源: yfinance)")
return quote
@@ -1038,11 +1099,11 @@ class DataFetcherManager:
# 美股单独处理,使用 YfinanceFetcher
if _is_us_code(stock_code):
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
if fetcher.name == "YfinanceFetcher":
if hasattr(fetcher, 'get_realtime_quote'):
try:
quote = fetcher.get_realtime_quote(stock_code)
quote = self._call_fetcher_method(fetcher, 'get_realtime_quote', stock_code)
if quote is not None:
logger.info(f"[实时行情] 美股 {stock_code} 成功获取 (来源: yfinance)")
return quote
@@ -1055,13 +1116,13 @@ class DataFetcherManager:
# 港股实时行情只走港股专用入口,避免按 A 股 source_priority
# 反复触发同一个 ak.stock_hk_spot_em() 接口。
if _is_hk_market(stock_code):
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
if fetcher.name != "AkshareFetcher":
continue
if not hasattr(fetcher, 'get_realtime_quote'):
break
try:
quote = fetcher.get_realtime_quote(stock_code, source="hk")
quote = self._call_fetcher_method(fetcher, 'get_realtime_quote', stock_code, source="hk")
if quote is not None and quote.has_basic_data():
logger.info(f"[实时行情] 港股 {stock_code} 成功获取 (来源: akshare_hk)")
return quote
@@ -1088,42 +1149,42 @@ class DataFetcherManager:
if source == "efinance":
# 尝试 EfinanceFetcher
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
if fetcher.name == "EfinanceFetcher":
if hasattr(fetcher, 'get_realtime_quote'):
quote = fetcher.get_realtime_quote(stock_code)
quote = self._call_fetcher_method(fetcher, 'get_realtime_quote', stock_code)
break
elif source == "akshare_em":
# 尝试 AkshareFetcher 东财数据源
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
if fetcher.name == "AkshareFetcher":
if hasattr(fetcher, 'get_realtime_quote'):
quote = fetcher.get_realtime_quote(stock_code, source="em")
quote = self._call_fetcher_method(fetcher, 'get_realtime_quote', stock_code, source="em")
break
elif source == "akshare_sina":
# 尝试 AkshareFetcher 新浪数据源
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
if fetcher.name == "AkshareFetcher":
if hasattr(fetcher, 'get_realtime_quote'):
quote = fetcher.get_realtime_quote(stock_code, source="sina")
quote = self._call_fetcher_method(fetcher, 'get_realtime_quote', stock_code, source="sina")
break
elif source in ("tencent", "akshare_qq"):
# 尝试 AkshareFetcher 腾讯数据源
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
if fetcher.name == "AkshareFetcher":
if hasattr(fetcher, 'get_realtime_quote'):
quote = fetcher.get_realtime_quote(stock_code, source="tencent")
quote = self._call_fetcher_method(fetcher, 'get_realtime_quote', stock_code, source="tencent")
break
elif source == "tushare":
# 尝试 TushareFetcher(需要 Tushare Pro 积分)
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
if fetcher.name == "TushareFetcher":
if hasattr(fetcher, 'get_realtime_quote'):
quote = fetcher.get_realtime_quote(raw_stock_code or stock_code)
quote = self._call_fetcher_method(fetcher, 'get_realtime_quote', raw_stock_code or stock_code)
break
if quote is not None and quote.has_basic_data():
@@ -1231,7 +1292,7 @@ class DataFetcherManager:
circuit_breaker = get_chip_circuit_breaker()
# 直接遍历管理器已经按 priority 排好序的数据源列表
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
# 只处理实现了筹码分布逻辑的数据源
if not hasattr(fetcher, 'get_chip_distribution'):
continue
@@ -1246,11 +1307,14 @@ class DataFetcherManager:
continue
try:
chip = fetcher.get_chip_distribution(stock_code)
chip = self._call_fetcher_method(fetcher, 'get_chip_distribution', stock_code)
if chip is not None:
circuit_breaker.record_success(source_key)
logger.info(f"[筹码分布] {stock_code} 成功获取 (来源: {fetcher_name})")
return chip
else:
# 空结果:释放 HALF_OPEN 探测名额,避免卡死
circuit_breaker.record_inconclusive(source_key)
except Exception as e:
logger.warning(f"[筹码分布] {fetcher_name} 获取 {stock_code} 失败: {e}")
circuit_breaker.record_failure(source_key, str(e))
@@ -1283,33 +1347,29 @@ class DataFetcherManager:
static_name = STOCK_NAME_MAP.get(stock_code)
# 1. 先检查缓存
if hasattr(self, '_stock_name_cache') and stock_code in self._stock_name_cache:
return self._stock_name_cache[stock_code]
# 初始化缓存
if not hasattr(self, '_stock_name_cache'):
self._stock_name_cache = {}
cached_name = self._get_cached_stock_name(stock_code)
if cached_name is not None:
return cached_name
# 2. 尝试从实时行情中获取(最快,可按需禁用)
if allow_realtime:
quote = self.get_realtime_quote(raw_stock_code or stock_code)
if quote and hasattr(quote, 'name') and is_meaningful_stock_name(getattr(quote, 'name', ''), stock_code):
name = quote.name
self._stock_name_cache[stock_code] = name
self._cache_stock_name(stock_code, name)
logger.info(f"[股票名称] 从实时行情获取: {stock_code} -> {name}")
return name
if is_meaningful_stock_name(static_name, stock_code):
self._stock_name_cache[stock_code] = static_name
return static_name
return self._cache_stock_name(stock_code, static_name) or static_name
# 3. 依次尝试各个数据源
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
if hasattr(fetcher, 'get_stock_name'):
try:
name = fetcher.get_stock_name(stock_code)
name = self._call_fetcher_method(fetcher, 'get_stock_name', stock_code)
if is_meaningful_stock_name(name, stock_code):
self._stock_name_cache[stock_code] = name
self._cache_stock_name(stock_code, name)
logger.info(f"[股票名称] 从 {fetcher.name} 获取: {stock_code} -> {name}")
return name
except Exception as e:
@@ -1382,31 +1442,36 @@ class DataFetcherManager:
missing_codes = set(stock_codes)
# 1. 先检查缓存
if not hasattr(self, '_stock_name_cache'):
self._stock_name_cache = {}
for code in stock_codes:
if code in self._stock_name_cache:
result[code] = self._stock_name_cache[code]
missing_codes.discard(code)
self._ensure_concurrency_guards()
with self._stock_name_cache_lock:
for code in stock_codes:
cached_name = self._stock_name_cache.get(code)
if cached_name is not None:
result[code] = cached_name
missing_codes.discard(code)
if not missing_codes:
return result
# 2. 尝试批量获取股票列表
for fetcher in self._fetchers:
for fetcher in self._get_fetchers_snapshot():
if hasattr(fetcher, 'get_stock_list') and missing_codes:
try:
stock_list = fetcher.get_stock_list()
stock_list = self._call_fetcher_method(fetcher, 'get_stock_list')
if stock_list is not None and not stock_list.empty:
cache_updates: Dict[str, str] = {}
for _, row in stock_list.iterrows():
code = row.get('code')
name = row.get('name')
if code and name:
self._stock_name_cache[code] = name
cache_updates[code] = name
if code in missing_codes:
result[code] = name
missing_codes.discard(code)
if cache_updates:
with self._stock_name_cache_lock:
self._stock_name_cache.update(cache_updates)
if not missing_codes:
break
+92 -61
View File
@@ -16,6 +16,7 @@
import logging
import time
from threading import RLock
from dataclasses import dataclass, field
from typing import Optional, Dict, Any, Union
from enum import Enum
@@ -298,9 +299,10 @@ class CircuitBreaker:
# 各数据源状态 {source_name: {state, failures, last_failure_time, half_open_calls}}
self._states: Dict[str, Dict[str, Any]] = {}
self._lock = RLock()
def _get_state(self, source: str) -> Dict[str, Any]:
"""获取或初始化数据源状态"""
def _get_state_locked(self, source: str) -> Dict[str, Any]:
"""获取或初始化数据源状态(调用方需持有锁)。"""
if source not in self._states:
self._states[source] = {
'state': self.CLOSED,
@@ -317,79 +319,108 @@ class CircuitBreaker:
返回 True 表示可以尝试请求
返回 False 表示应跳过该数据源
"""
state = self._get_state(source)
current_time = time.time()
if state['state'] == self.CLOSED:
return True
if state['state'] == self.OPEN:
# 检查冷却时间
time_since_failure = current_time - state['last_failure_time']
if time_since_failure >= self.cooldown_seconds:
# 冷却完成,进入半开状态
state['state'] = self.HALF_OPEN
state['half_open_calls'] = 0
logger.info(f"[熔断器] {source} 冷却完成,进入半开状态")
with self._lock:
state = self._get_state_locked(source)
current_time = time.time()
if state['state'] == self.CLOSED:
return True
else:
remaining = self.cooldown_seconds - time_since_failure
logger.debug(f"[熔断器] {source} 处于熔断状态,剩余冷却时间: {remaining:.0f}s")
if state['state'] == self.OPEN:
# 检查冷却时间
time_since_failure = current_time - state['last_failure_time']
if time_since_failure >= self.cooldown_seconds:
# 冷却完成,进入半开状态(不预占名额,由 HALF_OPEN 分支统一管理)
state['state'] = self.HALF_OPEN
state['half_open_calls'] = 0
state['last_failure_time'] = current_time
logger.info(f"[熔断器] {source} 冷却完成,进入半开状态")
# Fall through to HALF_OPEN check below
else:
remaining = self.cooldown_seconds - time_since_failure
logger.debug(f"[熔断器] {source} 处于熔断状态,剩余冷却时间: {remaining:.0f}s")
return False
if state['state'] == self.HALF_OPEN:
if state['half_open_calls'] < self.half_open_max_calls:
state['half_open_calls'] += 1
return True
# 所有探测名额已用完;若冷却时间再次到期仍未收到
# record_success/record_failure 回调,重置名额允许重新探测,
# 避免永久卡在 HALF_OPEN。
time_since_failure = current_time - state['last_failure_time']
if time_since_failure >= self.cooldown_seconds:
state['half_open_calls'] = 1
state['last_failure_time'] = current_time
logger.info(f"[熔断器] {source} 半开状态探测超时,重新探测")
return True
return False
if state['state'] == self.HALF_OPEN:
# 半开状态下限制请求次数
if state['half_open_calls'] < self.half_open_max_calls:
return True
return False
return True
return True
def record_inconclusive(self, source: str) -> None:
"""记录不确定的探测结果(如返回 None)。
仅影响 HALF_OPEN 状态:将其转回 OPEN 以便冷却后重新探测。
CLOSED 状态下为空操作,不影响失败计数。
"""
with self._lock:
state = self._get_state_locked(source)
if state['state'] == self.HALF_OPEN:
state['state'] = self.OPEN
state['half_open_calls'] = 0
state['last_failure_time'] = time.time()
logger.info(f"[熔断器] {source} 半开探测结果不确定,重新进入冷却")
def record_success(self, source: str) -> None:
"""记录成功请求"""
state = self._get_state(source)
if state['state'] == self.HALF_OPEN:
# 半开状态下成功,完全恢复
logger.info(f"[熔断器] {source} 半开状态请求成功,恢复正常")
# 重置状态
state['state'] = self.CLOSED
state['failures'] = 0
state['half_open_calls'] = 0
with self._lock:
state = self._get_state_locked(source)
if state['state'] == self.HALF_OPEN:
# 半开状态下成功,完全恢复
logger.info(f"[熔断器] {source} 半开状态请求成功,恢复正常")
# 重置状态
state['state'] = self.CLOSED
state['failures'] = 0
state['half_open_calls'] = 0
def record_failure(self, source: str, error: Optional[str] = None) -> None:
"""记录失败请求"""
state = self._get_state(source)
current_time = time.time()
state['failures'] += 1
state['last_failure_time'] = current_time
if state['state'] == self.HALF_OPEN:
# 半开状态下失败,继续熔断
state['state'] = self.OPEN
state['half_open_calls'] = 0
logger.warning(f"[熔断器] {source} 半开状态请求失败,继续熔断 {self.cooldown_seconds}s")
elif state['failures'] >= self.failure_threshold:
# 达到阈值,进入熔断
state['state'] = self.OPEN
logger.warning(f"[熔断器] {source} 连续失败 {state['failures']} 次,进入熔断状态 "
f"(冷却 {self.cooldown_seconds}s)")
if error:
logger.warning(f"[熔断器] 最后错误: {error}")
with self._lock:
state = self._get_state_locked(source)
current_time = time.time()
state['failures'] += 1
state['last_failure_time'] = current_time
if state['state'] == self.HALF_OPEN:
# 半开状态下失败,继续熔断
state['state'] = self.OPEN
state['half_open_calls'] = 0
logger.warning(f"[熔断器] {source} 半开状态请求失败,继续熔断 {self.cooldown_seconds}s")
elif state['failures'] >= self.failure_threshold:
# 达到阈值,进入熔断
state['state'] = self.OPEN
logger.warning(f"[熔断器] {source} 连续失败 {state['failures']} 次,进入熔断状态 "
f"(冷却 {self.cooldown_seconds}s)")
if error:
logger.warning(f"[熔断器] 最后错误: {error}")
def get_status(self) -> Dict[str, str]:
"""获取所有数据源状态"""
return {source: info['state'] for source, info in self._states.items()}
with self._lock:
return {source: info['state'] for source, info in self._states.items()}
def reset(self, source: Optional[str] = None) -> None:
"""重置熔断器状态"""
if source:
if source in self._states:
del self._states[source]
else:
self._states.clear()
with self._lock:
if source:
if source in self._states:
del self._states[source]
else:
self._states.clear()
# 全局熔断器实例(实时行情专用)
+179 -117
View File
@@ -157,6 +157,7 @@ class BaseSearchProvider(ABC):
self._key_cycle = cycle(api_keys) if api_keys else None
self._key_usage: Dict[str, int] = {key: 0 for key in api_keys}
self._key_errors: Dict[str, int] = {key: 0 for key in api_keys}
self._state_lock = threading.RLock()
@property
def name(self) -> str:
@@ -173,32 +174,36 @@ class BaseSearchProvider(ABC):
策略:轮询 + 跳过错误过多的 key
"""
if not self._key_cycle:
return None
# 最多尝试所有 key
for _ in range(len(self._api_keys)):
key = next(self._key_cycle)
# 跳过错误次数过多的 key(超过 3 次)
if self._key_errors.get(key, 0) < 3:
return key
# 所有 key 都有问题,重置错误计数并返回第一个
logger.warning(f"[{self._name}] 所有 API Key 都有错误记录,重置错误计数")
self._key_errors = {key: 0 for key in self._api_keys}
return self._api_keys[0] if self._api_keys else None
with self._state_lock:
if not self._key_cycle:
return None
# 最多尝试所有 key
for _ in range(len(self._api_keys)):
key = next(self._key_cycle)
# 跳过错误次数过多的 key(超过 3 次)
if self._key_errors.get(key, 0) < 3:
return key
# 所有 key 都有问题,重置错误计数并返回第一个
logger.warning(f"[{self._name}] 所有 API Key 都有错误记录,重置错误计数")
self._key_errors = {key: 0 for key in self._api_keys}
return self._api_keys[0] if self._api_keys else None
def _record_success(self, key: str) -> None:
"""记录成功使用"""
self._key_usage[key] = self._key_usage.get(key, 0) + 1
# 成功后减少错误计数
if key in self._key_errors and self._key_errors[key] > 0:
self._key_errors[key] -= 1
with self._state_lock:
self._key_usage[key] = self._key_usage.get(key, 0) + 1
# 成功后减少错误计数
if key in self._key_errors and self._key_errors[key] > 0:
self._key_errors[key] -= 1
def _record_error(self, key: str) -> None:
"""记录错误"""
self._key_errors[key] = self._key_errors.get(key, 0) + 1
logger.warning(f"[{self._name}] API Key {key[:8]}... 错误计数: {self._key_errors[key]}")
with self._state_lock:
self._key_errors[key] = self._key_errors.get(key, 0) + 1
error_count = self._key_errors[key]
logger.warning(f"[{self._name}] API Key {key[:8]}... 错误计数: {error_count}")
@abstractmethod
def _do_search(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse:
@@ -833,30 +838,36 @@ class MiniMaxSearchProvider(BaseSearchProvider):
@property
def is_available(self) -> bool:
"""Check availability considering circuit breaker state."""
if not super().is_available:
return False
if self._consecutive_failures >= self._CB_FAILURE_THRESHOLD:
if time.time() < self._circuit_open_until:
with self._state_lock:
if not self._api_keys:
return False
# Cooldown expired -> half-open, allow one probe
return True
if self._consecutive_failures >= self._CB_FAILURE_THRESHOLD:
if time.time() < self._circuit_open_until:
return False
# Cooldown expired -> half-open, allow one probe
return True
def _record_success(self, key: str) -> None:
super()._record_success(key)
# Reset circuit breaker on success
self._consecutive_failures = 0
self._circuit_open_until = 0.0
with self._state_lock:
super()._record_success(key)
# Reset circuit breaker on success
self._consecutive_failures = 0
self._circuit_open_until = 0.0
def _record_error(self, key: str) -> None:
super()._record_error(key)
self._consecutive_failures += 1
if self._consecutive_failures >= self._CB_FAILURE_THRESHOLD:
self._circuit_open_until = time.time() + self._CB_COOLDOWN_SECONDS
logger.warning(
f"[MiniMax] Circuit breaker OPEN – "
f"{self._consecutive_failures} consecutive failures, "
f"cooldown {self._CB_COOLDOWN_SECONDS}s"
)
warning_message = None
with self._state_lock:
super()._record_error(key)
self._consecutive_failures += 1
if self._consecutive_failures >= self._CB_FAILURE_THRESHOLD:
self._circuit_open_until = time.time() + self._CB_COOLDOWN_SECONDS
warning_message = (
f"[MiniMax] Circuit breaker OPEN – "
f"{self._consecutive_failures} consecutive failures, "
f"cooldown {self._CB_COOLDOWN_SECONDS}s"
)
if warning_message:
logger.warning(warning_message)
# ------------------------------------------------------------------
# Time-range helpers
@@ -1724,6 +1735,8 @@ class SearchService:
# In-memory search result cache: {cache_key: (timestamp, SearchResponse)}
self._cache: Dict[str, Tuple[float, 'SearchResponse']] = {}
self._cache_lock = threading.RLock()
self._cache_inflight: Dict[str, threading.Event] = {}
# Default cache TTL in seconds (10 minutes)
self._cache_ttl: int = 600
logger.info(
@@ -1784,35 +1797,67 @@ class SearchService:
"""Build a cache key from query parameters."""
return f"{query}|{max_results}|{days}"
def _get_cached(self, key: str) -> Optional['SearchResponse']:
"""Return cached SearchResponse if still valid, else None."""
def _get_cached_locked(self, key: str) -> Optional['SearchResponse']:
entry = self._cache.get(key)
if entry is None:
return None
ts, response = entry
if time.time() - ts > self._cache_ttl:
del self._cache[key]
self._cache.pop(key, None)
return None
logger.debug(f"Search cache hit: {key[:60]}...")
return response
def _get_cached(self, key: str) -> Optional['SearchResponse']:
"""Return cached SearchResponse if still valid, else None."""
with self._cache_lock:
return self._get_cached_locked(key)
def _get_cached_or_reserve(
self,
key: str,
) -> Tuple[Optional['SearchResponse'], bool, Optional[threading.Event]]:
with self._cache_lock:
cached = self._get_cached_locked(key)
if cached is not None:
return cached, False, None
event = self._cache_inflight.get(key)
if event is None:
event = threading.Event()
self._cache_inflight[key] = event
return None, True, event
return None, False, event
def _release_cache_fill(self, key: str, event: threading.Event) -> None:
with self._cache_lock:
current = self._cache_inflight.get(key)
if current is event:
self._cache_inflight.pop(key, None)
event.set()
def _wait_for_cached(self, key: str, event: threading.Event) -> Optional['SearchResponse']:
event.wait(timeout=max(1.0, min(float(self._cache_ttl), 30.0)))
return self._get_cached(key)
def _put_cache(self, key: str, response: 'SearchResponse') -> None:
"""Store a successful SearchResponse in cache."""
# Hard cap: evict oldest entries when cache exceeds limit
_MAX_CACHE_SIZE = 500
if len(self._cache) >= _MAX_CACHE_SIZE:
now = time.time()
# First pass: remove expired entries
expired = [k for k, (ts, _) in self._cache.items() if now - ts > self._cache_ttl]
for k in expired:
del self._cache[k]
# Second pass: if still over limit, evict oldest entries (FIFO)
with self._cache_lock:
# Hard cap: evict oldest entries when cache exceeds limit
_MAX_CACHE_SIZE = 500
if len(self._cache) >= _MAX_CACHE_SIZE:
excess = len(self._cache) - _MAX_CACHE_SIZE + 1
oldest = sorted(self._cache.keys(), key=lambda k: self._cache[k][0])[:excess]
for k in oldest:
del self._cache[k]
self._cache[key] = (time.time(), response)
now = time.time()
# First pass: remove expired entries
expired = [k for k, (ts, _) in self._cache.items() if now - ts > self._cache_ttl]
for k in expired:
self._cache.pop(k, None)
# Second pass: if still over limit, evict oldest entries (FIFO)
if len(self._cache) >= _MAX_CACHE_SIZE:
excess = len(self._cache) - _MAX_CACHE_SIZE + 1
oldest = sorted(self._cache.keys(), key=lambda k: self._cache[k][0])[:excess]
for k in oldest:
self._cache.pop(k, None)
self._cache[key] = (time.time(), response)
def _effective_news_window_days(self) -> int:
"""Resolve effective news window from strategy profile and global max-age."""
@@ -2121,66 +2166,79 @@ class SearchService:
provider_max_results,
)
# Check cache first
cache_key = self._cache_key(query, max_results, search_days)
cached = self._get_cached(cache_key)
cached, cache_owner, cache_event = self._get_cached_or_reserve(cache_key)
if cached is not None:
logger.info(f"使用缓存搜索结果: {stock_name}({stock_code})")
return cached
# 依次尝试各个搜索引擎(若过滤后为空,继续尝试下一引擎)
had_provider_success = False
for provider in self._providers:
if not provider.is_available:
continue
if not cache_owner and cache_event is not None:
cached = self._wait_for_cached(cache_key, cache_event)
if cached is not None:
logger.info(f"使用并发填充后的缓存搜索结果: {stock_name}({stock_code})")
return cached
cached, cache_owner, cache_event = self._get_cached_or_reserve(cache_key)
if cached is not None:
logger.info(f"使用等待后命中的缓存搜索结果: {stock_name}({stock_code})")
return cached
search_kwargs: Dict[str, Any] = {}
if isinstance(provider, TavilySearchProvider):
search_kwargs["topic"] = "news"
try:
# 依次尝试各个搜索引擎(若过滤后为空,继续尝试下一引擎)
had_provider_success = False
for provider in self._providers:
if not provider.is_available:
continue
response = provider.search(query, provider_max_results, days=search_days, **search_kwargs)
filtered_response = self._filter_news_response(
response,
search_days=search_days,
max_results=max_results,
log_scope=f"{stock_code}:{provider.name}:stock_news",
)
had_provider_success = had_provider_success or bool(response.success)
search_kwargs: Dict[str, Any] = {}
if isinstance(provider, TavilySearchProvider):
search_kwargs["topic"] = "news"
if filtered_response.success and filtered_response.results:
logger.info(f"使用 {provider.name} 搜索成功")
self._put_cache(cache_key, filtered_response)
return filtered_response
else:
if response.success and not filtered_response.results:
logger.info(
"%s 搜索成功但过滤后无有效新闻,继续尝试下一引擎",
provider.name,
)
response = provider.search(query, provider_max_results, days=search_days, **search_kwargs)
filtered_response = self._filter_news_response(
response,
search_days=search_days,
max_results=max_results,
log_scope=f"{stock_code}:{provider.name}:stock_news",
)
had_provider_success = had_provider_success or bool(response.success)
if filtered_response.success and filtered_response.results:
logger.info(f"使用 {provider.name} 搜索成功")
self._put_cache(cache_key, filtered_response)
return filtered_response
else:
logger.warning(
"%s 搜索失败: %s,尝试下一个引擎",
provider.name,
response.error_message,
)
if response.success and not filtered_response.results:
logger.info(
"%s 搜索成功但过滤后无有效新闻,继续尝试下一引擎",
provider.name,
)
else:
logger.warning(
"%s 搜索失败: %s,尝试下一个引擎",
provider.name,
response.error_message,
)
if had_provider_success:
if had_provider_success:
return SearchResponse(
query=query,
results=[],
provider="Filtered",
success=True,
error_message=None,
)
# 所有引擎都失败
return SearchResponse(
query=query,
results=[],
provider="Filtered",
success=True,
error_message=None,
provider="None",
success=False,
error_message="所有搜索引擎都不可用或搜索失败"
)
# 所有引擎都失败
return SearchResponse(
query=query,
results=[],
provider="None",
success=False,
error_message="所有搜索引擎都不可用或搜索失败"
)
finally:
if cache_owner and cache_event is not None:
self._release_cache_fill(cache_key, cache_event)
def search_stock_events(
self,
@@ -2686,6 +2744,7 @@ class SearchService:
# === 便捷函数 ===
_search_service: Optional[SearchService] = None
_search_service_lock = threading.Lock()
def get_search_service() -> SearchService:
@@ -2693,20 +2752,22 @@ def get_search_service() -> SearchService:
global _search_service
if _search_service is None:
from src.config import get_config
config = get_config()
_search_service = SearchService(
bocha_keys=config.bocha_api_keys,
tavily_keys=config.tavily_api_keys,
brave_keys=config.brave_api_keys,
serpapi_keys=config.serpapi_keys,
minimax_keys=config.minimax_api_keys,
searxng_base_urls=config.searxng_base_urls,
searxng_public_instances_enabled=config.searxng_public_instances_enabled,
news_max_age_days=config.news_max_age_days,
news_strategy_profile=getattr(config, "news_strategy_profile", "short"),
)
with _search_service_lock:
if _search_service is None:
from src.config import get_config
config = get_config()
_search_service = SearchService(
bocha_keys=config.bocha_api_keys,
tavily_keys=config.tavily_api_keys,
brave_keys=config.brave_api_keys,
serpapi_keys=config.serpapi_keys,
minimax_keys=config.minimax_api_keys,
searxng_base_urls=config.searxng_base_urls,
searxng_public_instances_enabled=config.searxng_public_instances_enabled,
news_max_age_days=config.news_max_age_days,
news_strategy_profile=getattr(config, "news_strategy_profile", "short"),
)
return _search_service
@@ -2714,7 +2775,8 @@ def get_search_service() -> SearchService:
def reset_search_service() -> None:
"""重置搜索服务(用于测试)"""
global _search_service
_search_service = None
with _search_service_lock:
_search_service = None
if __name__ == "__main__":
+48 -7
View File
@@ -12,6 +12,7 @@ Only activates for US stock codes (AAPL, TSLA, etc.).
"""
import logging
import threading
import time
from typing import Any, Dict, List, Optional
@@ -34,6 +35,8 @@ _TRANSIENT_EXCEPTIONS = (
)
_REQUEST_TIMEOUT = 8 # seconds
_REQUEST_RETRY_ATTEMPTS = 2
_REQUEST_RETRY_WAIT_CAP = 5 # wait_exponential(..., max=5)
@retry(
@@ -71,6 +74,8 @@ class SocialSentimentService:
self._api_url = (api_url or "https://api.adanos.org").rstrip("/")
# Simple in-memory cache: {"key": (timestamp, data)}
self._cache: Dict[str, tuple] = {}
self._cache_lock = threading.RLock()
self._cache_inflight: Dict[str, threading.Event] = {}
@property
def is_available(self) -> bool:
@@ -97,16 +102,52 @@ class SocialSentimentService:
logger.warning("Social sentiment API %s unexpected error: %s", url, e)
return None
@classmethod
def _cache_wait_timeout_seconds(cls) -> float:
request_budget = (_REQUEST_TIMEOUT * _REQUEST_RETRY_ATTEMPTS) + _REQUEST_RETRY_WAIT_CAP
return max(1.0, min(float(cls._TRENDING_CACHE_TTL), float(request_budget), 30.0))
def _fetch_cached(self, cache_key: str, url: str, params: Optional[Dict[str, Any]] = None) -> Optional[Any]:
"""Fetch with simple TTL cache (for trending endpoints)."""
now = time.monotonic()
cached = self._cache.get(cache_key)
if cached and (now - cached[0]) < self._TRENDING_CACHE_TTL:
return cached[1]
data = self._fetch_json(url, params)
if data is not None:
self._cache[cache_key] = (now, data)
return data
with self._cache_lock:
cached = self._cache.get(cache_key)
if cached and (now - cached[0]) < self._TRENDING_CACHE_TTL:
return cached[1]
inflight = self._cache_inflight.get(cache_key)
if inflight is None:
inflight = threading.Event()
self._cache_inflight[cache_key] = inflight
owner = True
else:
owner = False
if not owner:
inflight.wait(timeout=self._cache_wait_timeout_seconds())
now = time.monotonic()
with self._cache_lock:
cached = self._cache.get(cache_key)
if cached and (now - cached[0]) < self._TRENDING_CACHE_TTL:
return cached[1]
data = self._fetch_json(url, params)
if data is not None:
with self._cache_lock:
self._cache[cache_key] = (time.monotonic(), data)
return data
try:
data = self._fetch_json(url, params)
if data is not None:
with self._cache_lock:
self._cache[cache_key] = (time.monotonic(), data)
return data
finally:
with self._cache_lock:
current = self._cache_inflight.get(cache_key)
if current is inflight:
self._cache_inflight.pop(cache_key, None)
inflight.set()
def fetch_reddit_report(self, ticker: str) -> Optional[Dict]:
"""Fetch detailed Reddit report for a single ticker."""
@@ -5,10 +5,14 @@ Regression tests for stock-name prefetch behavior.
import os
import sys
import threading
import time
import unittest
from unittest.mock import MagicMock, call, patch
from types import SimpleNamespace
import pandas as pd
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
from data_provider.base import DataFetcherManager
@@ -23,6 +27,30 @@ class _DummyFetcher:
return "测试股票"
class _ThreadUnsafeStockListFetcher:
name = "ThreadUnsafeStockListFetcher"
def __init__(self):
self._active = False
self.call_count = 0
def get_stock_list(self):
if self._active:
raise AssertionError("concurrent get_stock_list access")
self._active = True
self.call_count += 1
try:
time.sleep(0.05)
return pd.DataFrame(
[
{"code": "600519", "name": "贵州茅台"},
{"code": "000001", "name": "平安银行"},
]
)
finally:
self._active = False
class TestPrefetchStockNames(unittest.TestCase):
def test_prefetch_stock_names_calls_get_stock_name_without_realtime(self):
manager = DataFetcherManager.__new__(DataFetcherManager)
@@ -105,6 +133,37 @@ class TestPrefetchStockNames(unittest.TestCase):
self.assertEqual(fetcher._stock_list_cache["300750"], "宁德时代")
api.get_finance_info.assert_not_called()
def test_batch_get_stock_names_serializes_shared_fetcher_access(self):
manager = DataFetcherManager.__new__(DataFetcherManager)
manager._fetchers = [_ThreadUnsafeStockListFetcher()]
barrier = threading.Barrier(2)
errors = []
results = []
def worker():
try:
barrier.wait(timeout=1)
result = DataFetcherManager.batch_get_stock_names(manager, ["600519", "000001"])
results.append(result)
except Exception as exc: # pragma: no cover - thread collection
errors.append(exc)
threads = [threading.Thread(target=worker) for _ in range(2)]
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=2)
self.assertEqual(errors, [])
self.assertEqual(len(results), 2)
for result in results:
self.assertEqual(result["600519"], "贵州茅台")
self.assertEqual(result["000001"], "平安银行")
self.assertGreaterEqual(manager._fetchers[0].call_count, 1)
self.assertEqual(manager._stock_name_cache["600519"], "贵州茅台")
self.assertEqual(manager._stock_name_cache["000001"], "平安银行")
if __name__ == "__main__":
unittest.main()
+137
View File
@@ -0,0 +1,137 @@
# -*- coding: utf-8 -*-
"""Concurrency regression tests for realtime circuit-breaker state."""
import threading
import time
import unittest
from data_provider.realtime_types import CircuitBreaker
class CircuitBreakerConcurrencyTestCase(unittest.TestCase):
def test_half_open_allows_only_one_concurrent_probe(self):
breaker = CircuitBreaker(
failure_threshold=1,
cooldown_seconds=0.01,
half_open_max_calls=1,
)
breaker.record_failure("akshare_em", "boom")
time.sleep(0.02)
barrier = threading.Barrier(2)
allowed = []
errors = []
def worker():
try:
barrier.wait(timeout=1)
allowed.append(breaker.is_available("akshare_em"))
except Exception as exc: # pragma: no cover - thread collection
errors.append(exc)
threads = [threading.Thread(target=worker) for _ in range(2)]
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=2)
self.assertEqual(errors, [])
self.assertCountEqual(allowed, [True, False])
self.assertEqual(breaker.get_status()["akshare_em"], CircuitBreaker.HALF_OPEN)
def test_concurrent_record_updates_keep_state_consistent(self):
breaker = CircuitBreaker(
failure_threshold=3,
cooldown_seconds=60.0,
half_open_max_calls=1,
)
barrier = threading.Barrier(4)
errors = []
def record_success():
try:
barrier.wait(timeout=1)
for _ in range(100):
breaker.record_success("tushare")
except Exception as exc: # pragma: no cover - thread collection
errors.append(exc)
def record_failure():
try:
barrier.wait(timeout=1)
for _ in range(100):
breaker.record_failure("tushare", "network")
except Exception as exc: # pragma: no cover - thread collection
errors.append(exc)
threads = [
threading.Thread(target=record_success),
threading.Thread(target=record_success),
threading.Thread(target=record_failure),
threading.Thread(target=record_failure),
]
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=2)
self.assertEqual(errors, [])
state = breaker._states["tushare"]
self.assertIn(state["state"], {CircuitBreaker.CLOSED, CircuitBreaker.OPEN, CircuitBreaker.HALF_OPEN})
self.assertGreaterEqual(state["failures"], 0)
self.assertGreaterEqual(state["half_open_calls"], 0)
def test_half_open_not_stuck_after_inconclusive_probe(self):
"""Regression: a probe that returns None (no record_success/record_failure)
must not permanently block the source in HALF_OPEN."""
breaker = CircuitBreaker(
failure_threshold=1,
cooldown_seconds=0.01,
half_open_max_calls=1,
)
# Trip the breaker
breaker.record_failure("src", "boom")
time.sleep(0.02)
# Probe goes through (OPEN -> HALF_OPEN -> slot consumed)
self.assertTrue(breaker.is_available("src"))
self.assertEqual(breaker.get_status()["src"], CircuitBreaker.HALF_OPEN)
# Simulate ambiguous None result via record_inconclusive
breaker.record_inconclusive("src")
self.assertEqual(breaker.get_status()["src"], CircuitBreaker.OPEN)
# After cooldown, source becomes available again
time.sleep(0.02)
self.assertTrue(breaker.is_available("src"))
def test_half_open_self_heals_without_callback(self):
"""If neither record_success/record_failure/record_inconclusive is called,
the HALF_OPEN state self-heals after another cooldown period."""
breaker = CircuitBreaker(
failure_threshold=1,
cooldown_seconds=0.01,
half_open_max_calls=1,
)
breaker.record_failure("src", "boom")
time.sleep(0.02)
# First probe consumes the slot
self.assertTrue(breaker.is_available("src"))
# Slot exhausted, blocked
self.assertFalse(breaker.is_available("src"))
# After cooldown, self-healing allows another probe
time.sleep(0.02)
self.assertTrue(breaker.is_available("src"))
def test_record_inconclusive_noop_in_closed(self):
"""record_inconclusive must be a no-op when the breaker is CLOSED."""
breaker = CircuitBreaker(failure_threshold=3, cooldown_seconds=60.0)
breaker.record_success("src")
breaker.record_inconclusive("src")
self.assertEqual(breaker.get_status()["src"], CircuitBreaker.CLOSED)
if __name__ == "__main__":
unittest.main()
+276
View File
@@ -0,0 +1,276 @@
# -*- coding: utf-8 -*-
"""Concurrency regression tests for search service shared state."""
import sys
import threading
import time
import unittest
from datetime import datetime
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
# Mock newspaper before search_service import (optional dependency)
if "newspaper" not in sys.modules:
mock_np = MagicMock()
mock_np.Article = MagicMock()
mock_np.Config = MagicMock()
sys.modules["newspaper"] = mock_np
from src.search_service import (
BaseSearchProvider,
SearchResponse,
SearchResult,
SearchService,
get_search_service,
reset_search_service,
)
class _ThreadUnsafeCycle:
def __init__(self, values):
self._values = list(values)
self._index = 0
self._active = False
def __next__(self):
if self._active:
raise AssertionError("concurrent cycle access")
self._active = True
try:
time.sleep(0.05)
value = self._values[self._index % len(self._values)]
self._index += 1
return value
finally:
self._active = False
class _DummyProvider(BaseSearchProvider):
def __init__(self, api_keys):
super().__init__(api_keys, "DummyProvider")
def _do_search(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse:
return SearchResponse(
query=query,
results=[
SearchResult(
title=f"{api_key}:{query}",
snippet="snippet",
url=f"https://example.com/{api_key}",
source="example.com",
published_date=datetime.now().date().isoformat(),
)
],
provider=self.name,
success=True,
)
class SearchServiceConcurrencyTestCase(unittest.TestCase):
def tearDown(self) -> None:
reset_search_service()
def test_get_cached_or_reserve_prefers_cached_response(self):
service = SearchService(
searxng_public_instances_enabled=False,
news_max_age_days=3,
news_strategy_profile="short",
)
cache_key = "cached-query|3|3"
response = SearchResponse(
query="cached-query",
results=[
SearchResult(
title="cached-news",
snippet="snippet",
url="https://example.com/cached-news",
source="example.com",
published_date=datetime.now().date().isoformat(),
)
],
provider="Cache",
success=True,
)
service._put_cache(cache_key, response)
cached, owner, event = service._get_cached_or_reserve(cache_key)
self.assertIs(cached, response)
self.assertFalse(owner)
self.assertIsNone(event)
self.assertNotIn(cache_key, service._cache_inflight)
def test_provider_key_rotation_is_serialized(self):
provider = _DummyProvider(["key-1", "key-2"])
provider._key_cycle = _ThreadUnsafeCycle(["key-1", "key-2"])
barrier = threading.Barrier(2)
errors = []
def worker():
try:
barrier.wait(timeout=1)
provider.search("query", max_results=1)
except Exception as exc: # pragma: no cover - thread collection
errors.append(exc)
threads = [threading.Thread(target=worker) for _ in range(2)]
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=2)
self.assertEqual(errors, [])
self.assertEqual(sum(provider._key_usage.values()), 2)
def test_search_stock_news_coalesces_concurrent_cache_fill(self):
service = SearchService(
searxng_public_instances_enabled=False,
news_max_age_days=3,
news_strategy_profile="short",
)
call_count = 0
call_lock = threading.Lock()
def provider_search(query, max_results, days=7, **_kwargs):
nonlocal call_count
with call_lock:
call_count += 1
time.sleep(0.05)
return SearchResponse(
query=query,
results=[
SearchResult(
title="fresh-news",
snippet="snippet",
url="https://example.com/fresh-news",
source="example.com",
published_date=datetime.now().date().isoformat(),
)
],
provider="MockProvider",
success=True,
)
provider = SimpleNamespace(
is_available=True,
name="MockProvider",
search=MagicMock(side_effect=provider_search),
)
service._providers = [provider]
barrier = threading.Barrier(4)
errors = []
responses = []
def worker():
try:
barrier.wait(timeout=1)
responses.append(service.search_stock_news("600519", "贵州茅台", max_results=3))
except Exception as exc: # pragma: no cover - thread collection
errors.append(exc)
threads = [threading.Thread(target=worker) for _ in range(4)]
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=2)
self.assertEqual(errors, [])
self.assertEqual(call_count, 1)
self.assertEqual(len(responses), 4)
for response in responses:
self.assertTrue(response.success)
self.assertEqual([item.title for item in response.results], ["fresh-news"])
def test_search_stock_news_rechecks_cache_after_wait_before_provider_search(self):
service = SearchService(
searxng_public_instances_enabled=False,
news_max_age_days=3,
news_strategy_profile="short",
)
search_days = service._effective_news_window_days()
cache_key = service._cache_key("贵州茅台 600519 股票 最新消息", 3, search_days)
cached_response = SearchResponse(
query="贵州茅台 600519 股票 最新消息",
results=[
SearchResult(
title="cached-after-wait",
snippet="snippet",
url="https://example.com/cached-after-wait",
source="example.com",
published_date=datetime.now().date().isoformat(),
)
],
provider="Cache",
success=True,
)
service._cache_inflight[cache_key] = threading.Event()
provider = SimpleNamespace(
is_available=True,
name="MockProvider",
search=MagicMock(side_effect=AssertionError("provider search should not run after cache fills")),
)
service._providers = [provider]
def wait_for_cached(key, _event):
self.assertEqual(key, cache_key)
service._put_cache(cache_key, cached_response)
return None
with patch.object(service, "_wait_for_cached", side_effect=wait_for_cached):
response = service.search_stock_news("600519", "贵州茅台", max_results=3)
self.assertIs(response, cached_response)
provider.search.assert_not_called()
def test_get_search_service_initializes_singleton_once(self):
reset_search_service()
config = SimpleNamespace(
bocha_api_keys=[],
tavily_api_keys=[],
brave_api_keys=[],
serpapi_keys=[],
minimax_api_keys=[],
searxng_base_urls=[],
searxng_public_instances_enabled=False,
news_max_age_days=3,
news_strategy_profile="short",
)
created = []
def build_service(**kwargs):
time.sleep(0.05)
service = SimpleNamespace(kwargs=kwargs)
created.append(service)
return service
barrier = threading.Barrier(4)
errors = []
services = []
with patch("src.search_service.SearchService", side_effect=build_service) as mock_cls:
with patch("src.config.get_config", return_value=config):
def worker():
try:
barrier.wait(timeout=1)
services.append(get_search_service())
except Exception as exc: # pragma: no cover - thread collection
errors.append(exc)
threads = [threading.Thread(target=worker) for _ in range(4)]
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=2)
self.assertEqual(errors, [])
self.assertEqual(mock_cls.call_count, 1)
self.assertEqual(len(created), 1)
self.assertEqual(len({id(service) for service in services}), 1)
if __name__ == "__main__":
unittest.main()
+53
View File
@@ -1,6 +1,7 @@
# -*- coding: utf-8 -*-
"""Tests for SocialSentimentService."""
import threading
import time
import unittest
from unittest.mock import patch, MagicMock
@@ -106,6 +107,58 @@ class TestFetchTrending(unittest.TestCase):
self.assertEqual(mock_get.call_count, 1)
def test_trending_cache_is_thread_safe_for_same_key(self):
call_count = 0
lock = threading.Lock()
barrier = threading.Barrier(4)
errors = []
results = []
def fake_fetch_json(_url, _params=None):
nonlocal call_count
with lock:
call_count += 1
time.sleep(0.05)
return {"trending": [{"ticker": "AAPL"}]}
with patch.object(self.svc, "_fetch_json", side_effect=fake_fetch_json):
def worker():
try:
barrier.wait(timeout=1)
results.append(self.svc.fetch_x_trending())
except Exception as exc: # pragma: no cover - thread collection
errors.append(exc)
threads = [threading.Thread(target=worker) for _ in range(4)]
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=2)
self.assertEqual(errors, [])
self.assertEqual(call_count, 1)
self.assertEqual(len(results), 4)
for result in results:
self.assertEqual(result, [{"ticker": "AAPL"}])
def test_waiter_fallback_uses_capped_wait_and_populates_cache(self):
inflight = MagicMock()
inflight.wait.return_value = False
self.svc._cache_inflight["x_trending"] = inflight
payload = {"trending": [{"ticker": "AAPL"}]}
with patch.object(self.svc, "_fetch_json", return_value=payload) as mock_fetch:
first = self.svc.fetch_x_trending()
second = self.svc.fetch_x_trending()
self.assertEqual(first, [{"ticker": "AAPL"}])
self.assertEqual(second, [{"ticker": "AAPL"}])
inflight.wait.assert_called_once_with(timeout=self.svc._cache_wait_timeout_seconds())
self.assertLess(self.svc._cache_wait_timeout_seconds(), self.svc._TRENDING_CACHE_TTL)
self.assertLessEqual(self.svc._cache_wait_timeout_seconds(), 30.0)
self.assertEqual(mock_fetch.call_count, 1)
self.assertEqual(self.svc._cache["x_trending"][1], payload)
class TestGetSocialContext(unittest.TestCase):
"""Tests for get_social_context (main entry point)."""