mirror of
https://github.com/ZhuLinsen/daily_stock_analysis.git
synced 2026-10-06 14:33:11 +08:00
* fix(issue-881): [bug]-修复并发执行时共享状态缺少统一加锁的问题
This commit is contained in:
+125
-60
@@ -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
|
||||
|
||||
@@ -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
@@ -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__":
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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)."""
|
||||
|
||||
Reference in New Issue
Block a user