diff --git a/backend_api_python/app/config/api_keys.py b/backend_api_python/app/config/api_keys.py index 165d234..f4eebfa 100644 --- a/backend_api_python/app/config/api_keys.py +++ b/backend_api_python/app/config/api_keys.py @@ -68,6 +68,42 @@ class MetaAPIKeys(type): from app.utils.config_loader import load_addon_config val = load_addon_config().get('grok', {}).get('api_key') return val if val else '' + + @property + def TAVILY_API_KEYS(cls): + """Tavily Search API keys (comma-separated for rotation)""" + env_val = os.getenv('TAVILY_API_KEYS', '').strip() + if env_val: + return [k.strip() for k in env_val.split(',') if k.strip()] + from app.utils.config_loader import load_addon_config + val = load_addon_config().get('tavily', {}).get('api_keys', '') + if val: + return [k.strip() for k in val.split(',') if k.strip()] + return [] + + @property + def BOCHA_API_KEYS(cls): + """Bocha Search API keys (comma-separated for rotation)""" + env_val = os.getenv('BOCHA_API_KEYS', '').strip() + if env_val: + return [k.strip() for k in env_val.split(',') if k.strip()] + from app.utils.config_loader import load_addon_config + val = load_addon_config().get('bocha', {}).get('api_keys', '') + if val: + return [k.strip() for k in val.split(',') if k.strip()] + return [] + + @property + def SERPAPI_KEYS(cls): + """SerpAPI keys (comma-separated for rotation)""" + env_val = os.getenv('SERPAPI_KEYS', '').strip() + if env_val: + return [k.strip() for k in env_val.split(',') if k.strip()] + from app.utils.config_loader import load_addon_config + val = load_addon_config().get('serpapi', {}).get('api_keys', '') + if val: + return [k.strip() for k in val.split(',') if k.strip()] + return [] class APIKeys(metaclass=MetaAPIKeys): diff --git a/backend_api_python/app/data_sources/__init__.py b/backend_api_python/app/data_sources/__init__.py index 22ac2a6..9ef889b 100644 --- a/backend_api_python/app/data_sources/__init__.py +++ b/backend_api_python/app/data_sources/__init__.py @@ -1,8 +1,61 @@ """ 数据源模块 支持多种市场的K线数据获取 + +改进版本(参考 daily_stock_analysis 项目): +- 熔断器保护 (circuit_breaker) +- 数据缓存 (cache_manager) +- 防封禁策略 (rate_limiter) +- 多数据源自动切换 (data_manager) """ from app.data_sources.factory import DataSourceFactory +from app.data_sources.circuit_breaker import ( + CircuitBreaker, + get_ashare_circuit_breaker, + get_realtime_circuit_breaker +) +from app.data_sources.cache_manager import ( + DataCache, + get_realtime_cache, + get_kline_cache, + get_stock_info_cache +) +from app.data_sources.rate_limiter import ( + RateLimiter, + get_eastmoney_limiter, + get_tencent_limiter, + get_akshare_limiter, + get_random_user_agent, + random_sleep, + retry_with_backoff +) +from app.data_sources.data_manager import ( + AShareDataManager, + get_ashare_data_manager +) -__all__ = ['DataSourceFactory'] +__all__ = [ + # 工厂 + 'DataSourceFactory', + # 熔断器 + 'CircuitBreaker', + 'get_ashare_circuit_breaker', + 'get_realtime_circuit_breaker', + # 缓存 + 'DataCache', + 'get_realtime_cache', + 'get_kline_cache', + 'get_stock_info_cache', + # 限流器 + 'RateLimiter', + 'get_eastmoney_limiter', + 'get_tencent_limiter', + 'get_akshare_limiter', + 'get_random_user_agent', + 'random_sleep', + 'retry_with_backoff', + # 数据管理器 + 'AShareDataManager', + 'get_ashare_data_manager', +] diff --git a/backend_api_python/app/data_sources/cache_manager.py b/backend_api_python/app/data_sources/cache_manager.py new file mode 100644 index 0000000..fdb8878 --- /dev/null +++ b/backend_api_python/app/data_sources/cache_manager.py @@ -0,0 +1,232 @@ +# -*- coding: utf-8 -*- +""" +=================================== +数据缓存管理模块 +=================================== + +参考 daily_stock_analysis 项目实现 +用于缓存实时行情和K线数据,减少重复请求 + +特性: +1. TTL (Time To Live) 过期机制 +2. LRU (Least Recently Used) 淘汰策略 +3. 按数据类型分区管理 +""" + +import time +import logging +from typing import Dict, Any, Optional, List +from collections import OrderedDict +from dataclasses import dataclass, field +from datetime import datetime +import threading + +logger = logging.getLogger(__name__) + + +@dataclass +class CacheEntry: + """缓存条目""" + data: Any + timestamp: float + ttl: float + hit_count: int = 0 + + def is_expired(self) -> bool: + """检查是否过期""" + return time.time() - self.timestamp > self.ttl + + def age(self) -> float: + """返回缓存年龄(秒)""" + return time.time() - self.timestamp + + +class DataCache: + """ + 数据缓存管理器 + + 特性: + - TTL 过期机制 + - 最大容量限制 + - LRU 淘汰策略 + - 线程安全 + """ + + def __init__( + self, + name: str = "default", + default_ttl: float = 600.0, # 默认10分钟 + max_size: int = 1000 # 最大缓存条目数 + ): + self.name = name + self.default_ttl = default_ttl + self.max_size = max_size + self._cache: OrderedDict[str, CacheEntry] = OrderedDict() + self._lock = threading.RLock() + + # 统计信息 + self._hits = 0 + self._misses = 0 + + def get(self, key: str) -> Optional[Any]: + """ + 获取缓存数据 + + Returns: + 缓存的数据,不存在或过期返回 None + """ + with self._lock: + if key not in self._cache: + self._misses += 1 + return None + + entry = self._cache[key] + + # 检查是否过期 + if entry.is_expired(): + del self._cache[key] + self._misses += 1 + logger.debug(f"[缓存] {self.name}:{key} 已过期,删除") + return None + + # 更新访问顺序(LRU) + self._cache.move_to_end(key) + entry.hit_count += 1 + self._hits += 1 + + logger.debug(f"[缓存命中] {self.name}:{key} (年龄: {entry.age():.0f}s/{entry.ttl:.0f}s)") + return entry.data + + def set( + self, + key: str, + data: Any, + ttl: Optional[float] = None + ) -> None: + """ + 设置缓存数据 + + Args: + key: 缓存键 + data: 缓存数据 + ttl: 过期时间(秒),None 使用默认值 + """ + with self._lock: + # 检查容量,执行 LRU 淘汰 + while len(self._cache) >= self.max_size: + oldest_key, _ = self._cache.popitem(last=False) + logger.debug(f"[缓存] {self.name} 容量已满,淘汰: {oldest_key}") + + actual_ttl = ttl if ttl is not None else self.default_ttl + self._cache[key] = CacheEntry( + data=data, + timestamp=time.time(), + ttl=actual_ttl + ) + + logger.debug(f"[缓存更新] {self.name}:{key} TTL={actual_ttl}s") + + def delete(self, key: str) -> bool: + """删除缓存条目""" + with self._lock: + if key in self._cache: + del self._cache[key] + logger.debug(f"[缓存] {self.name}:{key} 已删除") + return True + return False + + def clear(self) -> int: + """清空缓存""" + with self._lock: + count = len(self._cache) + self._cache.clear() + logger.info(f"[缓存] {self.name} 已清空 {count} 条记录") + return count + + def cleanup_expired(self) -> int: + """清理过期条目""" + with self._lock: + expired_keys = [ + key for key, entry in self._cache.items() + if entry.is_expired() + ] + for key in expired_keys: + del self._cache[key] + + if expired_keys: + logger.debug(f"[缓存] {self.name} 清理 {len(expired_keys)} 条过期记录") + return len(expired_keys) + + def stats(self) -> Dict[str, Any]: + """获取缓存统计信息""" + with self._lock: + total_requests = self._hits + self._misses + hit_rate = self._hits / total_requests if total_requests > 0 else 0 + + return { + 'name': self.name, + 'size': len(self._cache), + 'max_size': self.max_size, + 'hits': self._hits, + 'misses': self._misses, + 'hit_rate': f"{hit_rate:.1%}", + 'default_ttl': self.default_ttl + } + + +# ============================================ +# 全局缓存实例 +# ============================================ + +# A股实时行情缓存(20分钟TTL,全市场数据量大) +_ashare_realtime_cache = DataCache( + name="ashare_realtime", + default_ttl=1200.0, # 20分钟 + max_size=6000 # 约5000+股票 +) + +# K线数据缓存(5分钟TTL,按需缓存) +_kline_cache = DataCache( + name="kline", + default_ttl=300.0, # 5分钟 + max_size=500 # 最多500个交易对 +) + +# 股票基本信息缓存(1天TTL) +_stock_info_cache = DataCache( + name="stock_info", + default_ttl=86400.0, # 24小时 + max_size=6000 +) + + +def get_realtime_cache() -> DataCache: + """获取实时行情缓存""" + return _ashare_realtime_cache + + +def get_kline_cache() -> DataCache: + """获取K线数据缓存""" + return _kline_cache + + +def get_stock_info_cache() -> DataCache: + """获取股票信息缓存""" + return _stock_info_cache + + +def generate_kline_cache_key( + symbol: str, + timeframe: str, + limit: int, + before_time: Optional[int] = None +) -> str: + """ + 生成K线缓存键 + + 格式: symbol:timeframe:limit[:before_time] + """ + key = f"{symbol}:{timeframe}:{limit}" + if before_time: + key += f":{before_time}" + return key diff --git a/backend_api_python/app/data_sources/circuit_breaker.py b/backend_api_python/app/data_sources/circuit_breaker.py new file mode 100644 index 0000000..b0e3e1c --- /dev/null +++ b/backend_api_python/app/data_sources/circuit_breaker.py @@ -0,0 +1,186 @@ +# -*- coding: utf-8 -*- +""" +=================================== +熔断器模块 (Circuit Breaker) +=================================== + +参考 daily_stock_analysis 项目实现 +用于管理数据源的熔断/冷却状态,避免连续失败时反复请求 + +状态机: +CLOSED(正常) --失败N次--> OPEN(熔断)--冷却时间到--> HALF_OPEN(半开) +HALF_OPEN --成功--> CLOSED +HALF_OPEN --失败--> OPEN +""" + +import time +import logging +from typing import Dict, Any, Optional +from enum import Enum + +logger = logging.getLogger(__name__) + + +class CircuitState(Enum): + """熔断器状态""" + CLOSED = "closed" # 正常状态 + OPEN = "open" # 熔断状态(不可用) + HALF_OPEN = "half_open" # 半开状态(试探性请求) + + +class CircuitBreaker: + """ + 熔断器 - 管理数据源的熔断/冷却状态 + + 策略: + - 连续失败 N 次后进入熔断状态 + - 熔断期间跳过该数据源 + - 冷却时间后自动恢复半开状态 + - 半开状态下单次成功则完全恢复,失败则继续熔断 + """ + + def __init__( + self, + failure_threshold: int = 3, # 连续失败次数阈值 + cooldown_seconds: float = 300.0, # 冷却时间(秒),默认5分钟 + half_open_max_calls: int = 1 # 半开状态最大尝试次数 + ): + self.failure_threshold = failure_threshold + self.cooldown_seconds = cooldown_seconds + self.half_open_max_calls = half_open_max_calls + + # 各数据源状态 {source_name: {state, failures, last_failure_time, half_open_calls}} + self._states: Dict[str, Dict[str, Any]] = {} + + def _get_state(self, source: str) -> Dict[str, Any]: + """获取或初始化数据源状态""" + if source not in self._states: + self._states[source] = { + 'state': CircuitState.CLOSED, + 'failures': 0, + 'last_failure_time': 0.0, + 'half_open_calls': 0, + 'last_error': None + } + return self._states[source] + + def is_available(self, source: str) -> bool: + """ + 检查数据源是否可用 + + 返回 True 表示可以尝试请求 + 返回 False 表示应跳过该数据源 + """ + state = self._get_state(source) + current_time = time.time() + + if state['state'] == CircuitState.CLOSED: + return True + + if state['state'] == CircuitState.OPEN: + # 检查冷却时间 + time_since_failure = current_time - state['last_failure_time'] + if time_since_failure >= self.cooldown_seconds: + # 冷却完成,进入半开状态 + state['state'] = CircuitState.HALF_OPEN + state['half_open_calls'] = 0 + logger.info(f"[熔断器] {source} 冷却完成,进入半开状态") + return True + else: + remaining = self.cooldown_seconds - time_since_failure + logger.debug(f"[熔断器] {source} 处于熔断状态,剩余冷却时间: {remaining:.0f}s") + return False + + if state['state'] == CircuitState.HALF_OPEN: + # 半开状态下限制请求次数 + if state['half_open_calls'] < self.half_open_max_calls: + return True + return False + + return True + + def record_success(self, source: str) -> None: + """记录成功请求""" + state = self._get_state(source) + + if state['state'] == CircuitState.HALF_OPEN: + # 半开状态下成功,完全恢复 + logger.info(f"[熔断器] {source} 半开状态请求成功,恢复正常") + + # 重置状态 + state['state'] = CircuitState.CLOSED + state['failures'] = 0 + state['half_open_calls'] = 0 + state['last_error'] = None + + 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 + state['last_error'] = error + + if state['state'] == CircuitState.HALF_OPEN: + # 半开状态下失败,继续熔断 + state['state'] = CircuitState.OPEN + state['half_open_calls'] = 0 + logger.warning(f"[熔断器] {source} 半开状态请求失败,继续熔断 {self.cooldown_seconds}s") + elif state['failures'] >= self.failure_threshold: + # 达到阈值,进入熔断 + state['state'] = CircuitState.OPEN + logger.warning(f"[熔断器] {source} 连续失败 {state['failures']} 次,进入熔断状态 " + f"(冷却 {self.cooldown_seconds}s)") + if error: + logger.warning(f"[熔断器] 最后错误: {error}") + + def get_status(self) -> Dict[str, Dict[str, Any]]: + """获取所有数据源状态""" + return { + source: { + 'state': info['state'].value, + 'failures': info['failures'], + 'last_error': info['last_error'] + } + 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] + logger.info(f"[熔断器] 已重置 {source} 的熔断状态") + else: + self._states.clear() + logger.info("[熔断器] 已重置所有数据源的熔断状态") + + +# ============================================ +# 全局熔断器实例 +# ============================================ + +# A股数据源熔断器(标准策略) +_ashare_circuit_breaker = CircuitBreaker( + failure_threshold=3, # 连续失败3次熔断 + cooldown_seconds=300.0, # 冷却5分钟 + half_open_max_calls=1 +) + +# 实时行情熔断器(更严格的策略) +_realtime_circuit_breaker = CircuitBreaker( + failure_threshold=2, # 连续失败2次熔断 + cooldown_seconds=180.0, # 冷却3分钟 + half_open_max_calls=1 +) + + +def get_ashare_circuit_breaker() -> CircuitBreaker: + """获取A股数据源熔断器""" + return _ashare_circuit_breaker + + +def get_realtime_circuit_breaker() -> CircuitBreaker: + """获取实时行情熔断器""" + return _realtime_circuit_breaker diff --git a/backend_api_python/app/data_sources/cn_stock.py b/backend_api_python/app/data_sources/cn_stock.py index 6e25362..8f3ae01 100644 --- a/backend_api_python/app/data_sources/cn_stock.py +++ b/backend_api_python/app/data_sources/cn_stock.py @@ -1,11 +1,18 @@ """ CN/HK stock data source. Supports A-Share and H-Share with multiple public sources. -Priority (AShare): Eastmoney (intraday/daily) > yfinance (daily) > akshare (daily, optional). -Priority (HShare): Tencent (intraday) > Eastmoney/Tencent (daily) > yfinance (daily) > akshare (daily, optional). + +改进版本(参考 daily_stock_analysis 项目): +- 多数据源自动切换(按优先级) +- 熔断器保护 +- 数据缓存 +- 防封禁策略(随机休眠+UA轮换) + +Priority (AShare): Eastmoney > Tencent > Sina > Akshare > yfinance +Priority (HShare): Tencent > Eastmoney > yfinance > akshare """ import json -from typing import Dict, List, Any, Optional +from typing import Dict, List, Any, Optional, Tuple from datetime import datetime, timedelta import requests @@ -13,6 +20,9 @@ import yfinance as yf from app.data_sources.base import BaseDataSource from app.data_sources.us_stock import USStockDataSource +from app.data_sources.data_manager import get_ashare_data_manager, AShareDataManager +from app.data_sources.circuit_breaker import get_ashare_circuit_breaker, get_realtime_circuit_breaker +from app.data_sources.rate_limiter import get_request_headers, get_tencent_limiter, get_eastmoney_limiter from app.utils.logger import get_logger from app.utils.http import get_retry_session @@ -166,7 +176,14 @@ class TencentDataMixin: class AShareDataSource(BaseDataSource, TencentDataMixin): - """A-Share data source.""" + """ + A-Share data source. + + 改进版本:使用 AShareDataManager 实现多数据源自动切换 + - 熔断器保护 + - 数据缓存 + - 防封禁策略 + """ name = "AShare" @@ -190,6 +207,11 @@ class AShareDataSource(BaseDataSource, TencentDataMixin): def __init__(self): self.us_stock_source = USStockDataSource() + # 使用新的数据管理器 + self._data_manager = get_ashare_data_manager() + # 熔断器和限流器 + self._circuit_breaker = get_ashare_circuit_breaker() + self._em_limiter = get_eastmoney_limiter() def get_kline( self, @@ -198,11 +220,28 @@ class AShareDataSource(BaseDataSource, TencentDataMixin): limit: int, before_time: Optional[int] = None ) -> List[Dict[str, Any]]: - """Fetch A-Share Kline data.""" - klines = [] + """ + Fetch A-Share Kline data. - # Prefer Eastmoney (supports most intraday timeframes) - klines = self._fetch_eastmoney_ashare(symbol, timeframe, limit) + 改进版本:使用数据管理器自动切换数据源 + """ + # 使用新的数据管理器获取数据(自动切换数据源) + klines, source = self._data_manager.get_kline( + symbol=symbol, + timeframe=timeframe, + limit=limit, + before_time=before_time + ) + + if klines: + self.log_result(symbol, klines, timeframe) + return klines + + # 如果数据管理器失败,使用传统方式作为最后备选 + logger.warning(f"[AShare] 数据管理器获取 {symbol} 失败,尝试传统方式") + + # 传统方式:直接调用东方财富 + klines = self._fetch_eastmoney_ashare_legacy(symbol, timeframe, limit) if klines: klines = self.filter_and_limit(klines, limit, before_time) self.log_result(symbol, klines, timeframe) @@ -212,20 +251,24 @@ class AShareDataSource(BaseDataSource, TencentDataMixin): if timeframe in ('1D', '1W'): yahoo_symbol = self._to_yahoo_symbol(symbol) if yahoo_symbol: - # logger.info(f"尝试使用 yfinance 获取A股: {yahoo_symbol}") klines = self.us_stock_source.get_kline(yahoo_symbol, timeframe, limit, before_time) if klines: - # logger.info(f"yfinance 成功获取 {len(klines)} 条A股数据") return klines - # Fallback: akshare (daily/weekly) - if HAS_AKSHARE and timeframe in self.AKSHARE_PERIOD_MAP: - klines = self._fetch_akshare(symbol, timeframe, limit, before_time) - if klines: - return klines - logger.warning(f"AShare {symbol} data fetch failed") - return klines + return [] + + def _fetch_eastmoney_ashare_legacy( + self, + symbol: str, + timeframe: str, + limit: int + ) -> List[Dict[str, Any]]: + """ + 传统方式获取东方财富数据(兜底用) + 不使用新的熔断器和限流器,保持原有逻辑 + """ + return self._fetch_eastmoney_ashare(symbol, timeframe, limit) def _to_tencent_symbol(self, symbol: str) -> Optional[str]: """转换为腾讯财经格式""" @@ -391,7 +434,10 @@ class AShareDataSource(BaseDataSource, TencentDataMixin): """ 获取A股实时报价 - 使用东方财富实时行情API获取实时报价 + 改进版本:使用数据管理器自动切换数据源 + - 熔断器保护 + - 数据缓存(60秒TTL) + - 多数据源自动切换 Returns: dict: { @@ -406,6 +452,21 @@ class AShareDataSource(BaseDataSource, TencentDataMixin): """ symbol = (symbol or '').strip() + # 使用数据管理器获取实时报价(自动切换数据源) + quote, source = self._data_manager.get_realtime_quote(symbol) + if quote and quote.get('last', 0) > 0: + return quote + + # 如果数据管理器失败,使用传统方式作为兜底 + logger.debug(f"[AShare] 数据管理器获取 {symbol} 实时报价失败,尝试传统方式") + return self._get_ticker_legacy(symbol) + + def _get_ticker_legacy(self, symbol: str) -> Dict[str, Any]: + """ + 传统方式获取实时报价(兜底用) + + 保持原有逻辑,不使用熔断器和限流器 + """ # 优先使用东方财富实时行情 API try: # 判断市场 @@ -423,8 +484,6 @@ class AShareDataSource(BaseDataSource, TencentDataMixin): params = { 'secid': secid, 'fields': 'f43,f44,f45,f46,f47,f48,f57,f58,f60,f169,f170', - # f43=最新价, f44=最高价, f45=最低价, f46=开盘价 - # f60=昨收价, f169=涨跌额, f170=涨跌幅 } session = get_retry_session() @@ -434,9 +493,8 @@ class AShareDataSource(BaseDataSource, TencentDataMixin): if data and data.get('data'): d = data['data'] last_price = d.get('f43', 0) - # 东方财富返回的价格是整数(分),需要除以100 if last_price and last_price > 0: - divisor = 100 if last_price > 1000 else 1 # 价格超过10元时用分表示 + divisor = 100 if last_price > 1000 else 1 return { 'last': last_price / divisor, 'high': d.get('f44', 0) / divisor, @@ -444,7 +502,7 @@ class AShareDataSource(BaseDataSource, TencentDataMixin): 'open': d.get('f46', 0) / divisor, 'previousClose': d.get('f60', 0) / divisor, 'change': d.get('f169', 0) / divisor, - 'changePercent': d.get('f170', 0) / 100 # 涨跌幅是整数(%*100) + 'changePercent': d.get('f170', 0) / 100 } except Exception as e: logger.debug(f"Eastmoney ticker failed for {symbol}: {e}") @@ -473,28 +531,6 @@ class AShareDataSource(BaseDataSource, TencentDataMixin): except Exception as e: logger.debug(f"Tencent ticker failed for {symbol}: {e}") - # 第三备选: Akshare - try: - import akshare as ak - # 使用 akshare 获取实时行情 - df = ak.stock_zh_a_spot_em() - if df is not None and not df.empty: - # 在数据中查找对应股票 - row = df[df['代码'] == symbol] - if not row.empty: - row = row.iloc[0] - return { - 'last': float(row.get('最新价', 0) or 0), - 'change': float(row.get('涨跌额', 0) or 0), - 'changePercent': float(row.get('涨跌幅', 0) or 0), - 'high': float(row.get('最高', 0) or 0), - 'low': float(row.get('最低', 0) or 0), - 'open': float(row.get('今开', 0) or 0), - 'previousClose': float(row.get('昨收', 0) or 0) - } - except Exception as e: - logger.debug(f"Akshare ticker failed for {symbol}: {e}") - return {'last': 0, 'symbol': symbol} diff --git a/backend_api_python/app/data_sources/data_manager.py b/backend_api_python/app/data_sources/data_manager.py new file mode 100644 index 0000000..20a2b5e --- /dev/null +++ b/backend_api_python/app/data_sources/data_manager.py @@ -0,0 +1,629 @@ +# -*- coding: utf-8 -*- +""" +=================================== +A股数据源管理器 (Data Manager) +=================================== + +参考 daily_stock_analysis 项目的 DataFetcherManager 实现 +统一管理多个A股数据源,实现自动故障切换 + +数据源优先级: +1. 东方财富 (Eastmoney) - 数据最全,首选 +2. 腾讯财经 (Tencent) - 稳定可靠 +3. 新浪财经 (Sina) - 轻量级 +4. Akshare - 功能丰富,但容易被封 +5. yfinance - 国际数据源,兜底 +""" + +import logging +from typing import Dict, List, Any, Optional, Tuple +from datetime import datetime +import requests + +from app.data_sources.circuit_breaker import ( + CircuitBreaker, + get_ashare_circuit_breaker, + get_realtime_circuit_breaker +) +from app.data_sources.cache_manager import ( + DataCache, + get_realtime_cache, + get_kline_cache, + generate_kline_cache_key +) +from app.data_sources.rate_limiter import ( + RateLimiter, + get_eastmoney_limiter, + get_tencent_limiter, + get_akshare_limiter, + get_request_headers, + retry_with_backoff +) +from app.utils.logger import get_logger +from app.utils.http import get_retry_session + +logger = get_logger(__name__) + + +# ============================================ +# 数据源常量 +# ============================================ + +class DataSource: + """数据源标识""" + EASTMONEY = "eastmoney" + TENCENT = "tencent" + SINA = "sina" + AKSHARE = "akshare" + YFINANCE = "yfinance" + + +# 数据源优先级(数字越小优先级越高) +DATA_SOURCE_PRIORITY = { + DataSource.EASTMONEY: 0, + DataSource.TENCENT: 1, + DataSource.SINA: 2, + DataSource.AKSHARE: 3, + DataSource.YFINANCE: 4, +} + + +# ============================================ +# A股数据管理器 +# ============================================ + +class AShareDataManager: + """ + A股数据源管理器 + + 功能: + 1. 多数据源自动切换(按优先级) + 2. 熔断器保护 + 3. 数据缓存 + 4. 防封禁策略 + """ + + # 东方财富 K 线周期映射 + EM_PERIOD_MAP = { + '1m': '1', + '5m': '5', + '15m': '15', + '30m': '30', + '1H': '60', + '4H': '240', + '1D': '101', + '1W': '102', + } + + def __init__(self): + # 熔断器 + self._circuit_breaker = get_ashare_circuit_breaker() + self._realtime_cb = get_realtime_circuit_breaker() + + # 缓存 + self._realtime_cache = get_realtime_cache() + self._kline_cache = get_kline_cache() + + # 限流器 + self._em_limiter = get_eastmoney_limiter() + self._tencent_limiter = get_tencent_limiter() + self._akshare_limiter = get_akshare_limiter() + + # Akshare 可用性检查 + self._has_akshare = self._check_akshare() + + def _check_akshare(self) -> bool: + """检查 akshare 是否可用""" + try: + import akshare + return True + except ImportError: + logger.debug("akshare 未安装,相关功能已禁用") + return False + + def get_kline( + self, + symbol: str, + timeframe: str, + limit: int, + before_time: Optional[int] = None, + use_cache: bool = True + ) -> Tuple[List[Dict[str, Any]], str]: + """ + 获取K线数据(自动切换数据源) + + Args: + symbol: 股票代码 + timeframe: 时间周期 + limit: 数据条数 + before_time: 获取此时间之前的数据 + use_cache: 是否使用缓存 + + Returns: + (K线数据列表, 成功的数据源名称) + """ + # 检查缓存 + if use_cache: + cache_key = generate_kline_cache_key(symbol, timeframe, limit, before_time) + cached = self._kline_cache.get(cache_key) + if cached: + logger.debug(f"[缓存命中] K线数据 {symbol}:{timeframe}") + return cached, "cache" + + errors = [] + + # 按优先级尝试各个数据源 + sources = [ + (DataSource.EASTMONEY, self._fetch_eastmoney_kline), + (DataSource.TENCENT, self._fetch_tencent_kline), + (DataSource.AKSHARE, self._fetch_akshare_kline), + (DataSource.YFINANCE, self._fetch_yfinance_kline), + ] + + for source_name, fetch_func in sources: + # 检查熔断器 + if not self._circuit_breaker.is_available(source_name): + logger.debug(f"[熔断] {source_name} 处于熔断状态,跳过") + continue + + # 跳过不可用的 akshare + if source_name == DataSource.AKSHARE and not self._has_akshare: + continue + + try: + logger.debug(f"[数据源] 尝试 {source_name} 获取 {symbol}") + klines = fetch_func(symbol, timeframe, limit, before_time) + + if klines: + self._circuit_breaker.record_success(source_name) + + # 更新缓存 + if use_cache: + cache_key = generate_kline_cache_key(symbol, timeframe, limit, before_time) + self._kline_cache.set(cache_key, klines) + + logger.info(f"[数据源] {source_name} 成功获取 {symbol} {len(klines)} 条数据") + return klines, source_name + + except Exception as e: + error_msg = f"{source_name}: {str(e)}" + errors.append(error_msg) + self._circuit_breaker.record_failure(source_name, str(e)) + logger.warning(f"[数据源] {error_msg}") + + # 所有数据源都失败 + logger.error(f"[数据源] 所有数据源获取 {symbol} 失败: {errors}") + return [], "" + + def _fetch_eastmoney_kline( + self, + symbol: str, + timeframe: str, + limit: int, + before_time: Optional[int] = None + ) -> List[Dict[str, Any]]: + """使用东方财富获取K线数据""" + klines = [] + + period = self.EM_PERIOD_MAP.get(timeframe) + if not period: + raise ValueError(f"Eastmoney 不支持时间周期: {timeframe}") + + # 限流 + self._em_limiter.wait() + + # 确定市场代码 + if symbol.startswith('6'): + secid = f"1.{symbol}" # 上海 + else: + secid = f"0.{symbol}" # 深圳 + + url = "https://push2his.eastmoney.com/api/qt/stock/kline/get" + params = { + 'secid': secid, + 'fields1': 'f1,f2,f3,f4,f5,f6', + 'fields2': 'f51,f52,f53,f54,f55,f56,f57', + 'klt': period, + 'fqt': '1', # 前复权 + 'end': '20500101', + 'lmt': limit, + } + + headers = get_request_headers(referer='https://quote.eastmoney.com/') + + session = get_retry_session() + response = session.get(url, params=params, headers=headers, timeout=15) + + if response.status_code != 200: + raise ConnectionError(f"HTTP {response.status_code}") + + data = response.json() + + if data.get('data') and data['data'].get('klines'): + for line in data['data']['klines']: + try: + parts = line.split(',') + if len(parts) >= 6: + time_str = parts[0] + if ' ' in time_str: + dt = datetime.strptime(time_str, '%Y-%m-%d %H:%M') + else: + dt = datetime.strptime(time_str, '%Y-%m-%d') + + klines.append({ + 'time': int(dt.timestamp()), + 'open': round(float(parts[1]), 4), + 'high': round(float(parts[3]), 4), + 'low': round(float(parts[4]), 4), + 'close': round(float(parts[2]), 4), + 'volume': round(float(parts[5]), 2) + }) + except (ValueError, IndexError): + continue + + # 过滤和排序 + klines.sort(key=lambda x: x['time']) + if before_time: + klines = [k for k in klines if k['time'] < before_time] + if len(klines) > limit: + klines = klines[-limit:] + + return klines + + def _fetch_tencent_kline( + self, + symbol: str, + timeframe: str, + limit: int, + before_time: Optional[int] = None + ) -> List[Dict[str, Any]]: + """使用腾讯财经获取K线数据""" + # 腾讯周期映射 + period_map = { + '1m': 1, '5m': 5, '15m': 15, '30m': 30, + '1H': 60, '1D': 'day', '1W': 'week' + } + + period = period_map.get(timeframe) + if period is None: + raise ValueError(f"腾讯财经不支持时间周期: {timeframe}") + + # 转换代码格式 + if symbol.startswith('6'): + tencent_symbol = f"sh{symbol}" + elif symbol.startswith('0') or symbol.startswith('3'): + tencent_symbol = f"sz{symbol}" + else: + tencent_symbol = f"bj{symbol}" + + # 限流 + self._tencent_limiter.wait() + + # 构建URL + if isinstance(period, int): + url = f"http://ifzq.gtimg.cn/appstock/app/kline/mkline?param={tencent_symbol},m{period},,{limit}" + else: + url = f"http://web.ifzq.gtimg.cn/appstock/app/fqkline/get?param={tencent_symbol},{period},,,{limit},qfq" + + response = requests.get(url, timeout=10) + + if response.status_code != 200: + raise ConnectionError(f"HTTP {response.status_code}") + + data = response.json() + klines = [] + + if data.get('code') == 0 and 'data' in data: + stock_data = data['data'].get(tencent_symbol) + if stock_data: + if isinstance(period, int): + candles = stock_data.get(f'm{period}', []) + else: + candles = stock_data.get('qfqday', stock_data.get('day', [])) + + for candle in candles: + if len(candle) >= 5: + try: + time_str = str(candle[0]) + if len(time_str) == 12: + dt = datetime.strptime(time_str, '%Y%m%d%H%M') + elif len(time_str) == 10: + dt = datetime.strptime(time_str, '%Y-%m-%d') + else: + continue + + klines.append({ + 'time': int(dt.timestamp()), + 'open': round(float(candle[1]), 4), + 'high': round(float(candle[3]), 4), + 'low': round(float(candle[4]), 4), + 'close': round(float(candle[2]), 4), + 'volume': round(float(candle[5]), 2) if len(candle) > 5 else 0 + }) + except (ValueError, IndexError): + continue + + # 过滤和排序 + klines.sort(key=lambda x: x['time']) + if before_time: + klines = [k for k in klines if k['time'] < before_time] + if len(klines) > limit: + klines = klines[-limit:] + + return klines + + def _fetch_akshare_kline( + self, + symbol: str, + timeframe: str, + limit: int, + before_time: Optional[int] = None + ) -> List[Dict[str, Any]]: + """使用 Akshare 获取K线数据""" + if not self._has_akshare: + raise RuntimeError("akshare 未安装") + + import akshare as ak + from datetime import timedelta + + # Akshare 只支持日线/周线 + period_map = {'1D': 'daily', '1W': 'weekly'} + period = period_map.get(timeframe) + if not period: + raise ValueError(f"Akshare 不支持时间周期: {timeframe}") + + # 限流 + self._akshare_limiter.wait() + + # 计算日期范围 + if before_time: + end_date = datetime.fromtimestamp(before_time).strftime('%Y%m%d') + else: + end_date = datetime.now().strftime('%Y%m%d') + + days = limit * 2 if timeframe == '1D' else limit * 10 + start_date = (datetime.now() - timedelta(days=days)).strftime('%Y%m%d') + + df = ak.stock_zh_a_hist( + symbol=symbol, + period=period, + start_date=start_date, + end_date=end_date, + adjust="qfq" + ) + + klines = [] + if df is not None and not df.empty: + df = df.tail(limit) + for _, row in df.iterrows(): + ts = int(datetime.strptime(str(row['日期']), '%Y-%m-%d').timestamp()) + klines.append({ + 'time': ts, + 'open': round(float(row['开盘']), 4), + 'high': round(float(row['最高']), 4), + 'low': round(float(row['最低']), 4), + 'close': round(float(row['收盘']), 4), + 'volume': round(float(row['成交量']), 2) + }) + + return klines + + def _fetch_yfinance_kline( + self, + symbol: str, + timeframe: str, + limit: int, + before_time: Optional[int] = None + ) -> List[Dict[str, Any]]: + """使用 yfinance 获取K线数据(兜底)""" + import yfinance as yf + + # 转换为 Yahoo 格式 + if symbol.startswith('6'): + yahoo_symbol = f"{symbol}.SS" + elif symbol.startswith('0') or symbol.startswith('3'): + yahoo_symbol = f"{symbol}.SZ" + else: + yahoo_symbol = f"{symbol}.SS" + + # 周期映射 + period_map = { + '1D': ('1d', f'{limit}d'), + '1W': ('1wk', f'{limit * 7}d'), + '1H': ('1h', f'{limit}d'), + } + + interval, period = period_map.get(timeframe, ('1d', f'{limit}d')) + + ticker = yf.Ticker(yahoo_symbol) + df = ticker.history(period=period, interval=interval) + + klines = [] + if df is not None and not df.empty: + df = df.tail(limit) + for idx, row in df.iterrows(): + ts = int(idx.timestamp()) + klines.append({ + 'time': ts, + 'open': round(float(row['Open']), 4), + 'high': round(float(row['High']), 4), + 'low': round(float(row['Low']), 4), + 'close': round(float(row['Close']), 4), + 'volume': round(float(row['Volume']), 2) + }) + + return klines + + def get_realtime_quote( + self, + symbol: str, + use_cache: bool = True + ) -> Tuple[Dict[str, Any], str]: + """ + 获取实时报价(自动切换数据源) + + Args: + symbol: 股票代码 + use_cache: 是否使用缓存 + + Returns: + (报价数据字典, 成功的数据源名称) + """ + # 检查缓存 + if use_cache: + cached = self._realtime_cache.get(f"quote:{symbol}") + if cached: + return cached, "cache" + + errors = [] + + # 按优先级尝试各个数据源 + sources = [ + (DataSource.EASTMONEY, self._fetch_eastmoney_quote), + (DataSource.TENCENT, self._fetch_tencent_quote), + (DataSource.AKSHARE, self._fetch_akshare_quote), + ] + + for source_name, fetch_func in sources: + if not self._realtime_cb.is_available(source_name): + continue + + if source_name == DataSource.AKSHARE and not self._has_akshare: + continue + + try: + quote = fetch_func(symbol) + if quote and quote.get('last', 0) > 0: + self._realtime_cb.record_success(source_name) + + # 更新缓存 + if use_cache: + self._realtime_cache.set(f"quote:{symbol}", quote, ttl=60.0) + + return quote, source_name + + except Exception as e: + errors.append(f"{source_name}: {str(e)}") + self._realtime_cb.record_failure(source_name, str(e)) + + logger.warning(f"[实时报价] 所有数据源获取 {symbol} 失败") + return {'last': 0, 'symbol': symbol}, "" + + def _fetch_eastmoney_quote(self, symbol: str) -> Dict[str, Any]: + """使用东方财富获取实时报价""" + if symbol.startswith('6'): + secid = f"1.{symbol}" + else: + secid = f"0.{symbol}" + + url = "https://push2.eastmoney.com/api/qt/stock/get" + params = { + 'secid': secid, + 'fields': 'f43,f44,f45,f46,f47,f48,f57,f58,f60,f169,f170', + } + + self._em_limiter.wait() + + session = get_retry_session() + response = session.get(url, params=params, headers=get_request_headers(), timeout=10) + + if response.status_code == 200: + data = response.json() + if data and data.get('data'): + d = data['data'] + last_price = d.get('f43', 0) + if last_price and last_price > 0: + divisor = 100 if last_price > 1000 else 1 + return { + 'last': last_price / divisor, + 'high': d.get('f44', 0) / divisor, + 'low': d.get('f45', 0) / divisor, + 'open': d.get('f46', 0) / divisor, + 'previousClose': d.get('f60', 0) / divisor, + 'change': d.get('f169', 0) / divisor, + 'changePercent': d.get('f170', 0) / 100 + } + + return {} + + def _fetch_tencent_quote(self, symbol: str) -> Dict[str, Any]: + """使用腾讯财经获取实时报价""" + if symbol.startswith('6'): + tencent_symbol = f"sh{symbol}" + else: + tencent_symbol = f"sz{symbol}" + + url = f"http://qt.gtimg.cn/q={tencent_symbol}" + + self._tencent_limiter.wait() + + response = requests.get(url, timeout=10) + content = response.content.decode('gbk', errors='ignore') + + if '="' in content: + data_str = content.split('="')[1].strip('";\n') + if data_str: + parts = data_str.split('~') + if len(parts) > 32: + return { + 'last': float(parts[3]) if parts[3] else 0, + 'change': float(parts[31]) if parts[31] else 0, + 'changePercent': float(parts[32]) if parts[32] else 0, + 'high': float(parts[33]) if len(parts) > 33 and parts[33] else 0, + 'low': float(parts[34]) if len(parts) > 34 and parts[34] else 0, + 'open': float(parts[5]) if len(parts) > 5 and parts[5] else 0, + 'previousClose': float(parts[4]) if parts[4] else 0 + } + + return {} + + def _fetch_akshare_quote(self, symbol: str) -> Dict[str, Any]: + """使用 Akshare 获取实时报价""" + if not self._has_akshare: + return {} + + import akshare as ak + + self._akshare_limiter.wait() + + df = ak.stock_zh_a_spot_em() + if df is not None and not df.empty: + row = df[df['代码'] == symbol] + if not row.empty: + row = row.iloc[0] + return { + 'last': float(row.get('最新价', 0) or 0), + 'change': float(row.get('涨跌额', 0) or 0), + 'changePercent': float(row.get('涨跌幅', 0) or 0), + 'high': float(row.get('最高', 0) or 0), + 'low': float(row.get('最低', 0) or 0), + 'open': float(row.get('今开', 0) or 0), + 'previousClose': float(row.get('昨收', 0) or 0) + } + + return {} + + def get_status(self) -> Dict[str, Any]: + """获取数据管理器状态""" + return { + 'circuit_breaker': self._circuit_breaker.get_status(), + 'realtime_circuit_breaker': self._realtime_cb.get_status(), + 'realtime_cache_stats': self._realtime_cache.stats(), + 'kline_cache_stats': self._kline_cache.stats(), + 'has_akshare': self._has_akshare, + } + + +# ============================================ +# 全局实例 +# ============================================ + +_ashare_data_manager: Optional[AShareDataManager] = None + + +def get_ashare_data_manager() -> AShareDataManager: + """获取A股数据管理器单例""" + global _ashare_data_manager + if _ashare_data_manager is None: + _ashare_data_manager = AShareDataManager() + return _ashare_data_manager diff --git a/backend_api_python/app/data_sources/rate_limiter.py b/backend_api_python/app/data_sources/rate_limiter.py new file mode 100644 index 0000000..9987eb5 --- /dev/null +++ b/backend_api_python/app/data_sources/rate_limiter.py @@ -0,0 +1,272 @@ +# -*- coding: utf-8 -*- +""" +=================================== +防封禁工具模块 (Rate Limiter) +=================================== + +参考 daily_stock_analysis 项目实现 +提供反爬虫策略: +1. 随机休眠(Jitter) +2. 随机 User-Agent 轮换 +3. 指数退避重试 +4. 请求频率限制 +""" + +import time +import random +import logging +from typing import Optional, Callable, Any, Type, Tuple +from functools import wraps + +logger = logging.getLogger(__name__) + + +# ============================================ +# User-Agent 池 +# ============================================ + +USER_AGENTS = [ + # Chrome Windows + 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36', + 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/119.0.0.0 Safari/537.36', + 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/121.0.0.0 Safari/537.36', + # Chrome Mac + 'Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36', + 'Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/119.0.0.0 Safari/537.36', + # Firefox + 'Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:121.0) Gecko/20100101 Firefox/121.0', + 'Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:120.0) Gecko/20100101 Firefox/120.0', + 'Mozilla/5.0 (Macintosh; Intel Mac OS X 10.15; rv:121.0) Gecko/20100101 Firefox/121.0', + # Safari + 'Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/17.2 Safari/605.1.15', + 'Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/17.1 Safari/605.1.15', + # Edge + 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36 Edg/120.0.0.0', + # Linux Chrome + 'Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36', +] + + +def get_random_user_agent() -> str: + """获取随机 User-Agent""" + return random.choice(USER_AGENTS) + + +def get_request_headers(referer: Optional[str] = None) -> dict: + """ + 获取带有随机 User-Agent 的请求头 + + Args: + referer: 可选的 Referer 头 + + Returns: + 请求头字典 + """ + headers = { + 'User-Agent': get_random_user_agent(), + 'Accept': 'application/json, text/plain, */*', + 'Accept-Language': 'zh-CN,zh;q=0.9,en;q=0.8', + 'Accept-Encoding': 'gzip, deflate', + 'Connection': 'keep-alive', + } + + if referer: + headers['Referer'] = referer + + return headers + + +# ============================================ +# 随机休眠 +# ============================================ + +def random_sleep( + min_seconds: float = 1.0, + max_seconds: float = 3.0, + log: bool = False +) -> None: + """ + 随机休眠(Jitter) + + 防封禁策略:模拟人类行为的随机延迟 + 在请求之间加入不规则的等待时间 + + Args: + min_seconds: 最小休眠时间(秒) + max_seconds: 最大休眠时间(秒) + log: 是否记录日志 + """ + sleep_time = random.uniform(min_seconds, max_seconds) + if log: + logger.debug(f"随机休眠 {sleep_time:.2f} 秒...") + time.sleep(sleep_time) + + +# ============================================ +# 请求频率限制器 +# ============================================ + +class RateLimiter: + """ + 请求频率限制器 + + 确保请求之间有最小间隔时间 + """ + + def __init__( + self, + min_interval: float = 1.0, + jitter_min: float = 0.5, + jitter_max: float = 1.5 + ): + """ + 初始化频率限制器 + + Args: + min_interval: 最小请求间隔(秒) + jitter_min: 随机抖动最小值(秒) + jitter_max: 随机抖动最大值(秒) + """ + self.min_interval = min_interval + self.jitter_min = jitter_min + self.jitter_max = jitter_max + self._last_request_time: Optional[float] = None + + def wait(self) -> float: + """ + 等待直到可以发起下一次请求 + + Returns: + 实际等待的时间(秒) + """ + wait_time = 0.0 + + if self._last_request_time is not None: + elapsed = time.time() - self._last_request_time + if elapsed < self.min_interval: + # 补充休眠到最小间隔 + wait_time = self.min_interval - elapsed + time.sleep(wait_time) + + # 添加随机抖动 + jitter = random.uniform(self.jitter_min, self.jitter_max) + time.sleep(jitter) + wait_time += jitter + + # 记录本次请求时间 + self._last_request_time = time.time() + + return wait_time + + def reset(self) -> None: + """重置限制器""" + self._last_request_time = None + + +# ============================================ +# 指数退避重试装饰器 +# ============================================ + +def retry_with_backoff( + max_attempts: int = 3, + base_delay: float = 2.0, + max_delay: float = 30.0, + exponential_base: float = 2.0, + exceptions: Tuple[Type[Exception], ...] = (Exception,), + on_retry: Optional[Callable[[int, Exception], None]] = None +): + """ + 指数退避重试装饰器 + + Args: + max_attempts: 最大重试次数 + base_delay: 基础延迟时间(秒) + max_delay: 最大延迟时间(秒) + exponential_base: 指数基数 + exceptions: 需要重试的异常类型 + on_retry: 重试时的回调函数 + + 使用示例: + @retry_with_backoff(max_attempts=3, exceptions=(ConnectionError, TimeoutError)) + def fetch_data(): + ... + """ + def decorator(func: Callable) -> Callable: + @wraps(func) + def wrapper(*args, **kwargs) -> Any: + last_exception = None + + for attempt in range(1, max_attempts + 1): + try: + return func(*args, **kwargs) + except exceptions as e: + last_exception = e + + if attempt == max_attempts: + logger.error(f"[重试] {func.__name__} 已达最大重试次数 ({max_attempts}),放弃") + raise + + # 计算退避延迟: base_delay * (exponential_base ^ (attempt - 1)) + delay = min( + base_delay * (exponential_base ** (attempt - 1)), + max_delay + ) + # 添加随机抖动 (±20%) + delay *= random.uniform(0.8, 1.2) + + logger.warning( + f"[重试] {func.__name__} 第 {attempt}/{max_attempts} 次失败: {e}, " + f"等待 {delay:.1f}s 后重试..." + ) + + if on_retry: + on_retry(attempt, e) + + time.sleep(delay) + + # 不应该到达这里 + raise last_exception + + return wrapper + return decorator + + +# ============================================ +# 全局限流器实例 +# ============================================ + +# 东方财富接口限流器(较严格) +_eastmoney_limiter = RateLimiter( + min_interval=2.0, + jitter_min=1.0, + jitter_max=3.0 +) + +# 腾讯财经接口限流器(较宽松) +_tencent_limiter = RateLimiter( + min_interval=1.0, + jitter_min=0.5, + jitter_max=1.5 +) + +# Akshare 接口限流器 +_akshare_limiter = RateLimiter( + min_interval=2.0, + jitter_min=1.5, + jitter_max=3.5 +) + + +def get_eastmoney_limiter() -> RateLimiter: + """获取东方财富限流器""" + return _eastmoney_limiter + + +def get_tencent_limiter() -> RateLimiter: + """获取腾讯财经限流器""" + return _tencent_limiter + + +def get_akshare_limiter() -> RateLimiter: + """获取 Akshare 限流器""" + return _akshare_limiter diff --git a/backend_api_python/app/routes/settings.py b/backend_api_python/app/routes/settings.py index cf90538..4234bcc 100644 --- a/backend_api_python/app/routes/settings.py +++ b/backend_api_python/app/routes/settings.py @@ -658,9 +658,9 @@ CONFIG_SCHEMA = { 'key': 'SEARCH_PROVIDER', 'label': 'Search Provider', 'type': 'select', - 'options': ['google', 'bing', 'none'], - 'default': 'google', - 'description': 'Web search provider for AI research features' + 'options': ['bocha', 'tavily', 'google', 'bing', 'none'], + 'default': 'bocha', + 'description': 'Web search provider for AI research features. Bocha recommended for A-share news' }, { 'key': 'SEARCH_MAX_RESULTS', @@ -669,6 +669,36 @@ CONFIG_SCHEMA = { 'default': '10', 'description': 'Maximum search results to return' }, + # Tavily Search API + { + 'key': 'TAVILY_API_KEYS', + 'label': 'Tavily API Keys', + 'type': 'password', + 'required': False, + 'link': 'https://tavily.com/', + 'link_text': 'settings.link.getApiKey', + 'description': 'Tavily Search API keys, comma-separated for rotation. Free 1000 requests/month' + }, + # Bocha Search API + { + 'key': 'BOCHA_API_KEYS', + 'label': 'Bocha API Keys', + 'type': 'password', + 'required': False, + 'link': 'https://bochaai.com/', + 'link_text': 'settings.link.getApiKey', + 'description': 'Bocha Search API keys, comma-separated for rotation. Best for A-share news' + }, + # SerpAPI + { + 'key': 'SERPAPI_KEYS', + 'label': 'SerpAPI Keys', + 'type': 'password', + 'required': False, + 'link': 'https://serpapi.com/', + 'link_text': 'settings.link.getApiKey', + 'description': 'SerpAPI keys for Google/Bing search, comma-separated for rotation' + }, { 'key': 'SEARCH_GOOGLE_API_KEY', 'label': 'Google API Key', diff --git a/backend_api_python/app/services/market_data_collector.py b/backend_api_python/app/services/market_data_collector.py index 7ed2005..962ec77 100644 --- a/backend_api_python/app/services/market_data_collector.py +++ b/backend_api_python/app/services/market_data_collector.py @@ -890,15 +890,16 @@ class MarketDataCollector: """ 获取新闻和情绪数据 - 策略: - 1. 使用结构化API (Finnhub) - 无需深度阅读 - 2. 只获取标题和摘要 - 不读取全文 - 3. 多来源聚合 - Finnhub + 市场特定来源 + 策略(按优先级): + 1. 结构化API (Finnhub) - 美股首选 + 2. akshare 多源 - A股首选(东方财富/新浪/同花顺/雪球) + 3. 搜索引擎 (Bocha/Tavily) - 补充搜索 + 4. 情绪分析 - Finnhub 社交媒体情绪 """ news_list = [] sentiment = {} - # 1) Finnhub 新闻 (最可靠) + # === 1) Finnhub 新闻 (美股首选) === if self._finnhub_client: try: end_date = datetime.now().strftime('%Y-%m-%d') @@ -908,29 +909,32 @@ class MarketDataCollector: if market == 'USStock': raw_news = self._finnhub_client.company_news(symbol, _from=start_date, to=end_date) - else: - # 通用新闻 + elif market == 'Crypto': + # 加密货币通用新闻 + raw_news = self._finnhub_client.general_news('crypto', min_id=0) + elif market not in ('AShare', 'HShare'): + # 其他市场通用新闻 raw_news = self._finnhub_client.general_news('general', min_id=0) if raw_news: - for item in raw_news[:10]: # 最多10条 + for item in raw_news[:10]: if not item.get('headline'): continue news_list.append({ "datetime": datetime.fromtimestamp(item.get('datetime', 0)).strftime('%Y-%m-%d %H:%M'), "headline": item.get('headline', ''), - "summary": item.get('summary', '')[:300] if item.get('summary') else '', # 截断摘要 + "summary": item.get('summary', '')[:300] if item.get('summary') else '', "source": item.get('source', 'Finnhub'), "url": item.get('url', ''), - "sentiment": item.get('sentiment', 'neutral'), # Finnhub有时提供情绪 + "sentiment": item.get('sentiment', 'neutral'), }) + logger.info(f"Finnhub 新闻获取成功: {len(news_list)} 条") except Exception as e: logger.debug(f"Finnhub news fetch failed: {e}") - # 2) Finnhub 情绪分数 (如果可用) + # === 2) Finnhub 情绪分数 (美股社交媒体情绪) === if self._finnhub_client and market == 'USStock': try: - # Finnhub 提供社交媒体情绪 social = self._finnhub_client.stock_social_sentiment(symbol) if social: sentiment['reddit'] = social.get('reddit', {}) @@ -938,31 +942,183 @@ class MarketDataCollector: except Exception as e: logger.debug(f"Finnhub sentiment fetch failed: {e}") - # 3) A股特定新闻 (akshare) + # === 3) A股多源新闻 (akshare) === if market == 'AShare' and self._ak: - try: - # 个股新闻 - df = self._ak.stock_news_em(symbol=symbol) - if df is not None and not df.empty: - for _, row in df.head(10).iterrows(): - news_list.append({ - "datetime": str(row.get('发布时间', ''))[:16], - "headline": row.get('新闻标题', ''), - "summary": row.get('新闻内容', '')[:200] if row.get('新闻内容') else '', - "source": row.get('文章来源', 'eastmoney'), - "url": row.get('新闻链接', ''), - "sentiment": 'neutral', - }) - except Exception as e: - logger.debug(f"akshare news fetch failed: {e}") + ashare_news = self._get_ashare_news_multi_source(symbol) + news_list.extend(ashare_news) + + # === 4) 港股新闻 (akshare) === + if market == 'HShare' and self._ak: + hshare_news = self._get_hshare_news(symbol) + news_list.extend(hshare_news) + + # === 5) 搜索引擎补充 (如果新闻太少) === + if len(news_list) < 5: + search_news = self._get_news_from_search(market, symbol, company_name) + news_list.extend(search_news) + + # 去重(按标题) + seen_titles = set() + unique_news = [] + for item in news_list: + title = item.get('headline', '') + if title and title not in seen_titles: + seen_titles.add(title) + unique_news.append(item) # 按时间排序 - news_list.sort(key=lambda x: x.get('datetime', ''), reverse=True) + unique_news.sort(key=lambda x: x.get('datetime', ''), reverse=True) return { - "news": news_list[:15], # 最多15条 + "news": unique_news[:15], # 最多15条 "sentiment": sentiment, } + + def _get_ashare_news_multi_source(self, symbol: str) -> List[Dict[str, Any]]: + """ + A股多源新闻获取 + + 来源(按优先级): + 1. 东方财富个股新闻 (stock_news_em) + 2. 新浪财经滚动新闻 (stock_news_sina) + 3. 同花顺个股新闻 (stock_news_ths) + 4. 雪球热帖 (stock_xuqiu) + """ + news_list = [] + + # 1) 东方财富个股新闻 + try: + df = self._ak.stock_news_em(symbol=symbol) + if df is not None and not df.empty: + for _, row in df.head(8).iterrows(): + news_list.append({ + "datetime": str(row.get('发布时间', ''))[:16], + "headline": row.get('新闻标题', ''), + "summary": row.get('新闻内容', '')[:200] if row.get('新闻内容') else '', + "source": "东方财富", + "url": row.get('新闻链接', ''), + "sentiment": 'neutral', + }) + logger.debug(f"东方财富新闻: {len(df)} 条") + except Exception as e: + logger.debug(f"东方财富新闻获取失败: {e}") + + # 2) 新浪财经个股新闻 + try: + # 注意:akshare 的新浪新闻接口可能需要股票名称而非代码 + df = self._ak.stock_news_sina(symbol=symbol) + if df is not None and not df.empty: + for _, row in df.head(5).iterrows(): + title = row.get('title', '') or row.get('新闻标题', '') + if title and title not in [n.get('headline') for n in news_list]: + news_list.append({ + "datetime": str(row.get('time', row.get('发布时间', '')))[:16], + "headline": title, + "summary": row.get('content', row.get('新闻内容', ''))[:200] if row.get('content') or row.get('新闻内容') else '', + "source": "新浪财经", + "url": row.get('url', row.get('新闻链接', '')), + "sentiment": 'neutral', + }) + logger.debug(f"新浪财经新闻: {len(df)} 条") + except Exception as e: + logger.debug(f"新浪财经新闻获取失败: {e}") + + # 3) 同花顺个股新闻 + try: + df = self._ak.stock_news_ths(symbol=symbol) + if df is not None and not df.empty: + for _, row in df.head(5).iterrows(): + title = row.get('标题', '') or row.get('title', '') + if title and title not in [n.get('headline') for n in news_list]: + news_list.append({ + "datetime": str(row.get('发布时间', row.get('time', '')))[:16], + "headline": title, + "summary": row.get('内容', row.get('content', ''))[:200] if row.get('内容') or row.get('content') else '', + "source": "同花顺", + "url": row.get('链接', row.get('url', '')), + "sentiment": 'neutral', + }) + logger.debug(f"同花顺新闻: {len(df)} 条") + except Exception as e: + logger.debug(f"同花顺新闻获取失败: {e}") + + # 4) 雪球热帖(社区讨论) + try: + df = self._ak.stock_xuqiu(symbol=symbol) + if df is not None and not df.empty: + for _, row in df.head(3).iterrows(): + title = row.get('标题', '') or row.get('title', '') + if title and title not in [n.get('headline') for n in news_list]: + news_list.append({ + "datetime": str(row.get('发布时间', row.get('time', '')))[:16], + "headline": title, + "summary": row.get('内容摘要', row.get('content', ''))[:200] if row.get('内容摘要') or row.get('content') else '', + "source": "雪球", + "url": row.get('链接', row.get('url', '')), + "sentiment": 'neutral', + }) + logger.debug(f"雪球热帖: {len(df)} 条") + except Exception as e: + logger.debug(f"雪球热帖获取失败: {e}") + + return news_list + + def _get_hshare_news(self, symbol: str) -> List[Dict[str, Any]]: + """港股新闻获取""" + news_list = [] + + try: + # 港股新闻 (如果 akshare 支持) + df = self._ak.stock_hk_spot_em() + # 港股一般没有专门的新闻接口,可以通过搜索补充 + except Exception as e: + logger.debug(f"港股新闻获取失败: {e}") + + return news_list + + def _get_news_from_search( + self, market: str, symbol: str, company_name: str = None + ) -> List[Dict[str, Any]]: + """ + 从搜索引擎获取新闻 + + 使用增强的搜索服务 (Bocha/Tavily/SerpAPI) + """ + news_list = [] + + try: + from app.services.search import get_search_service + search_service = get_search_service() + + if not search_service.is_available: + return news_list + + # 构建搜索名称 + search_name = company_name or symbol + + # 搜索股票新闻 + response = search_service.search_stock_news( + stock_code=symbol, + stock_name=search_name, + market=market, + max_results=5 + ) + + if response.success and response.results: + for result in response.results: + news_list.append({ + "datetime": result.published_date or datetime.now().strftime('%Y-%m-%d'), + "headline": result.title, + "summary": result.snippet[:200] if result.snippet else '', + "source": f"搜索:{result.source}", + "url": result.url, + "sentiment": result.sentiment, + }) + logger.info(f"搜索引擎新闻补充: {len(news_list)} 条 (来源: {response.provider})") + except Exception as e: + logger.debug(f"搜索引擎新闻获取失败: {e}") + + return news_list # 全局实例 diff --git a/backend_api_python/app/services/search.py b/backend_api_python/app/services/search.py index f471246..ceb03b0 100644 --- a/backend_api_python/app/services/search.py +++ b/backend_api_python/app/services/search.py @@ -1,12 +1,28 @@ """ -Search service. -Integrates Google Custom Search (CSE), Bing Search API, and DuckDuckGo (free fallback). -Configuration is provided via environment variables (see env.example) through config_loader. +Search service v2.0 - 增强版搜索服务 +整合多个搜索引擎,支持 API Key 轮换和故障转移 + +支持的搜索引擎(按优先级): +1. Bocha (博查) - 国内搜索优化,A股新闻推荐 +2. Tavily - 专为AI设计,免费1000次/月 +3. SerpAPI - Google/Bing 结果抓取 +4. Google CSE - 自定义搜索引擎 +5. Bing Search API +6. DuckDuckGo - 免费兜底 + +参考:daily_stock_analysis-main/src/search_service.py """ import requests import json import time +import re +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from datetime import datetime from typing import List, Dict, Any, Optional +from itertools import cycle +from urllib.parse import urlparse + from app.utils.logger import get_logger from app.utils.config_loader import load_addon_config @@ -17,103 +33,615 @@ _google_quota_exhausted = False _google_quota_reset_time = 0 -class SearchService: - """Search service with automatic fallback.""" +@dataclass +class SearchResult: + """搜索结果数据类""" + title: str + snippet: str # 摘要 + url: str + source: str # 来源网站 + published_date: Optional[str] = None + sentiment: str = 'neutral' # 情绪标签 - def __init__(self): - self._config = {} - self._load_config() + def to_text(self) -> str: + """转换为文本格式""" + date_str = f" ({self.published_date})" if self.published_date else "" + return f"【{self.source}】{self.title}{date_str}\n{self.snippet}" + + def to_dict(self) -> Dict[str, Any]: + """转换为字典""" + return { + 'title': self.title, + 'link': self.url, + 'snippet': self.snippet, + 'source': self.source, + 'published': self.published_date or '', + 'sentiment': self.sentiment, + } - def _load_config(self): - """Load config (re-read env-config on each call for local hot-reload).""" - config = load_addon_config() - self._config = config.get('search', {}) - self.provider = self._config.get('provider', 'google') - self.max_results = int(self._config.get('max_results', 10)) - def search(self, query: str, num_results: int = None, date_restrict: str = None) -> List[Dict[str, Any]]: +@dataclass +class SearchResponse: + """搜索响应""" + query: str + results: List[SearchResult] + provider: str # 使用的搜索引擎 + success: bool = True + error_message: Optional[str] = None + search_time: float = 0.0 # 搜索耗时(秒) + + def to_context(self, max_results: int = 5) -> str: + """将搜索结果转换为可用于 AI 分析的上下文""" + if not self.success or not self.results: + return f"搜索 '{self.query}' 未找到相关结果。" + + lines = [f"【{self.query} 搜索结果】(来源:{self.provider})"] + for i, result in enumerate(self.results[:max_results], 1): + lines.append(f"\n{i}. {result.to_text()}") + + return "\n".join(lines) + + def to_list(self) -> List[Dict[str, Any]]: + """转换为列表格式(兼容旧接口)""" + return [r.to_dict() for r in self.results] + + +class BaseSearchProvider(ABC): + """搜索引擎基类""" + + def __init__(self, api_keys: List[str], name: str): """ - Execute a web search with automatic fallback. + 初始化搜索引擎 Args: - query: Search query - num_results: Override default max results - date_restrict: Time restriction like 'd7' (past 7 days), Google only + api_keys: API Key 列表(支持多个 key 负载均衡) + name: 搜索引擎名称 + """ + self._api_keys = api_keys + self._name = name + 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} + + @property + def name(self) -> str: + return self._name + + @property + def is_available(self) -> bool: + """检查是否有可用的 API Key""" + return bool(self._api_keys) + + def _get_next_key(self) -> Optional[str]: + """ + 获取下一个可用的 API Key(负载均衡) + + 策略:轮询 + 跳过错误过多的 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 + + 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 + + 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]}") + + @abstractmethod + def _do_search(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse: + """执行搜索(子类实现)""" + pass + + def search(self, query: str, max_results: int = 5, days: int = 7) -> SearchResponse: + """ + 执行搜索 + + Args: + query: 搜索关键词 + max_results: 最大返回结果数 + days: 搜索最近几天的时间范围(默认7天) Returns: - List of search results + SearchResponse 对象 """ - global _google_quota_exhausted, _google_quota_reset_time + api_key = self._get_next_key() + if not api_key: + return SearchResponse( + query=query, + results=[], + provider=self._name, + success=False, + error_message=f"{self._name} 未配置 API Key" + ) - # 重新加载配置以支持热更新 - self._load_config() - - limit = num_results if num_results else self.max_results - - # Check if Google quota has reset (after midnight UTC typically) - if _google_quota_exhausted and time.time() > _google_quota_reset_time: - _google_quota_exhausted = False - logger.info("Google API quota reset, re-enabling Google search") - - results = [] - - if self.provider == 'bing': - results = self._search_bing(query, limit) - elif self.provider == 'duckduckgo': - results = self._search_duckduckgo(query, limit) - else: - # Google with fallback - if not _google_quota_exhausted: - results = self._search_google(query, limit, date_restrict) + start_time = time.time() + try: + response = self._do_search(query, api_key, max_results, days=days) + response.search_time = time.time() - start_time - # If Google failed or returned empty, try fallbacks - if not results: - logger.info("Google search failed or empty, trying fallback search engines...") - # Try Bing first if configured - results = self._search_bing(query, limit) - - # If Bing also failed, try DuckDuckGo (free, no API key needed) - if not results: - results = self._search_duckduckgo(query, limit) - - return results + if response.success: + self._record_success(api_key) + logger.info(f"[{self._name}] 搜索 '{query}' 成功,返回 {len(response.results)} 条结果,耗时 {response.search_time:.2f}s") + else: + self._record_error(api_key) + + return response + + except Exception as e: + self._record_error(api_key) + elapsed = time.time() - start_time + logger.error(f"[{self._name}] 搜索 '{query}' 失败: {e}") + return SearchResponse( + query=query, + results=[], + provider=self._name, + success=False, + error_message=str(e), + search_time=elapsed + ) + + @staticmethod + def _extract_domain(url: str) -> str: + """从 URL 提取域名作为来源""" + try: + parsed = urlparse(url) + domain = parsed.netloc.replace('www.', '') + return domain or '未知来源' + except: + return '未知来源' - def _search_google(self, query: str, num_results: int, date_restrict: str = None) -> List[Dict[str, Any]]: - """Google Custom Search (CSE).""" - global _google_quota_exhausted, _google_quota_reset_time - - api_key = self._config.get('google', {}).get('api_key') - cx = self._config.get('google', {}).get('cx') - - if not api_key or not cx: - logger.warning("Google Search is not configured (missing api_key or cx).") - return [] - - url = "https://www.googleapis.com/customsearch/v1" - params = { - 'key': api_key, - 'cx': cx, - 'q': query, - 'num': min(num_results, 10), # Google API 限制每次最多10条 - 'gl': 'cn' if any(c in query for c in ['A股', '利好', '利空', '财报']) else None # 针对中文内容优化地区 - } - - # 添加时间限制参数 - if date_restrict: - params['dateRestrict'] = date_restrict + +class TavilySearchProvider(BaseSearchProvider): + """ + Tavily 搜索引擎 + + 特点: + - 专为 AI/LLM 优化的搜索 API + - 免费版每月 1000 次请求 + - 返回结构化的搜索结果 + + 文档:https://docs.tavily.com/ + """ + + def __init__(self, api_keys: List[str]): + super().__init__(api_keys, "Tavily") + + def _do_search(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse: + """执行 Tavily 搜索""" + try: + from tavily import TavilyClient + except ImportError: + # 如果未安装 tavily-python,使用 REST API + return self._do_search_rest(query, api_key, max_results, days) try: + client = TavilyClient(api_key=api_key) + + # 执行搜索 + response = client.search( + query=query, + search_depth="advanced", + max_results=max_results, + include_answer=False, + include_raw_content=False, + days=days, + ) + + # 解析结果 + results = [] + for item in response.get('results', []): + results.append(SearchResult( + title=item.get('title', ''), + snippet=item.get('content', '')[:500], + url=item.get('url', ''), + source=self._extract_domain(item.get('url', '')), + published_date=item.get('published_date'), + )) + + return SearchResponse( + query=query, + results=results, + provider=self.name, + success=True, + ) + + except Exception as e: + error_msg = str(e) + if 'rate limit' in error_msg.lower() or 'quota' in error_msg.lower(): + error_msg = f"API 配额已用尽: {error_msg}" + + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=error_msg + ) + + def _do_search_rest(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse: + """使用 REST API 执行 Tavily 搜索(备选方案)""" + try: + url = "https://api.tavily.com/search" + headers = { + 'Content-Type': 'application/json', + } + payload = { + 'api_key': api_key, + 'query': query, + 'search_depth': 'advanced', + 'max_results': max_results, + 'include_answer': False, + 'include_raw_content': False, + } + + response = requests.post(url, headers=headers, json=payload, timeout=15) + + if response.status_code != 200: + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=f"HTTP {response.status_code}: {response.text}" + ) + + data = response.json() + results = [] + for item in data.get('results', []): + results.append(SearchResult( + title=item.get('title', ''), + snippet=item.get('content', '')[:500], + url=item.get('url', ''), + source=self._extract_domain(item.get('url', '')), + published_date=item.get('published_date'), + )) + + return SearchResponse( + query=query, + results=results, + provider=self.name, + success=True, + ) + + except Exception as e: + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=str(e) + ) + + +class BochaSearchProvider(BaseSearchProvider): + """ + 博查搜索引擎 + + 特点: + - 专为AI优化的中文搜索API + - 结果准确、摘要完整 + - 支持时间范围过滤和AI摘要 + + 文档:https://bocha-ai.feishu.cn/wiki/RXEOw02rFiwzGSkd9mUcqoeAnNK + """ + + def __init__(self, api_keys: List[str]): + super().__init__(api_keys, "Bocha") + + def _do_search(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse: + """执行博查搜索""" + try: + url = "https://api.bochaai.com/v1/web-search" + + headers = { + 'Authorization': f'Bearer {api_key}', + 'Content-Type': 'application/json' + } + + # 确定时间范围 + freshness = "oneWeek" + if days <= 1: + freshness = "oneDay" + elif days <= 7: + freshness = "oneWeek" + elif days <= 30: + freshness = "oneMonth" + else: + freshness = "oneYear" + + payload = { + "query": query, + "freshness": freshness, + "summary": True, + "count": min(max_results, 50) + } + + response = requests.post(url, headers=headers, json=payload, timeout=15) + + if response.status_code != 200: + error_message = response.text + try: + if response.headers.get('content-type', '').startswith('application/json'): + error_data = response.json() + error_message = error_data.get('message', response.text) + except: + pass + + if response.status_code == 403: + error_msg = f"余额不足: {error_message}" + elif response.status_code == 401: + error_msg = f"API KEY无效: {error_message}" + elif response.status_code == 429: + error_msg = f"请求频率达到限制: {error_message}" + else: + error_msg = f"HTTP {response.status_code}: {error_message}" + + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=error_msg + ) + + data = response.json() + + if data.get('code') != 200: + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=data.get('msg') or f"API返回错误码: {data.get('code')}" + ) + + results = [] + web_pages = data.get('data', {}).get('webPages', {}) + value_list = web_pages.get('value', []) + + for item in value_list[:max_results]: + snippet = item.get('summary') or item.get('snippet', '') + if snippet: + snippet = snippet[:500] + + results.append(SearchResult( + title=item.get('name', ''), + snippet=snippet, + url=item.get('url', ''), + source=item.get('siteName') or self._extract_domain(item.get('url', '')), + published_date=item.get('datePublished'), + )) + + return SearchResponse( + query=query, + results=results, + provider=self.name, + success=True, + ) + + except requests.exceptions.Timeout: + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message="请求超时" + ) + except Exception as e: + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=str(e) + ) + + +class SerpAPISearchProvider(BaseSearchProvider): + """ + SerpAPI 搜索引擎 + + 特点: + - 支持 Google、Bing、百度等多种搜索引擎 + - 免费版每月 100 次请求 + + 文档:https://serpapi.com/ + """ + + def __init__(self, api_keys: List[str]): + super().__init__(api_keys, "SerpAPI") + + def _do_search(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse: + """执行 SerpAPI 搜索""" + try: + from serpapi import GoogleSearch + except ImportError: + return self._do_search_rest(query, api_key, max_results, days) + + try: + tbs = "qdr:w" + if days <= 1: + tbs = "qdr:d" + elif days <= 7: + tbs = "qdr:w" + elif days <= 30: + tbs = "qdr:m" + else: + tbs = "qdr:y" + + params = { + "engine": "google", + "q": query, + "api_key": api_key, + "google_domain": "google.com.hk", + "hl": "zh-cn", + "gl": "cn", + "tbs": tbs, + "num": max_results + } + + search = GoogleSearch(params) + response = search.get_dict() + + results = [] + organic_results = response.get('organic_results', []) + + for item in organic_results[:max_results]: + results.append(SearchResult( + title=item.get('title', ''), + snippet=item.get('snippet', '')[:500], + url=item.get('link', ''), + source=item.get('source', self._extract_domain(item.get('link', ''))), + published_date=item.get('date'), + )) + + return SearchResponse( + query=query, + results=results, + provider=self.name, + success=True, + ) + + except Exception as e: + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=str(e) + ) + + def _do_search_rest(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse: + """使用 REST API 执行 SerpAPI 搜索""" + try: + tbs = "qdr:w" + if days <= 1: + tbs = "qdr:d" + elif days <= 7: + tbs = "qdr:w" + elif days <= 30: + tbs = "qdr:m" + + url = "https://serpapi.com/search" + params = { + "engine": "google", + "q": query, + "api_key": api_key, + "hl": "zh-cn", + "gl": "cn", + "tbs": tbs, + "num": max_results + } + + response = requests.get(url, params=params, timeout=15) + + if response.status_code != 200: + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=f"HTTP {response.status_code}" + ) + + data = response.json() + results = [] + + for item in data.get('organic_results', [])[:max_results]: + results.append(SearchResult( + title=item.get('title', ''), + snippet=item.get('snippet', '')[:500], + url=item.get('link', ''), + source=self._extract_domain(item.get('link', '')), + published_date=item.get('date'), + )) + + return SearchResponse( + query=query, + results=results, + provider=self.name, + success=True, + ) + + except Exception as e: + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=str(e) + ) + + +class GoogleSearchProvider(BaseSearchProvider): + """Google Custom Search (CSE) 搜索引擎""" + + def __init__(self, api_key: str, cx: str): + super().__init__([api_key] if api_key else [], "Google") + self._cx = cx + + def _do_search(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse: + """执行 Google 搜索""" + global _google_quota_exhausted, _google_quota_reset_time + + if not self._cx: + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message="Google Search 未配置 CX" + ) + + try: + url = "https://www.googleapis.com/customsearch/v1" + params = { + 'key': api_key, + 'cx': self._cx, + 'q': query, + 'num': min(max_results, 10), + } + + # 添加时间限制 + if days <= 1: + params['dateRestrict'] = 'd1' + elif days <= 7: + params['dateRestrict'] = 'w1' + elif days <= 30: + params['dateRestrict'] = 'm1' + response = requests.get(url, params=params, timeout=10) - # Check for quota exceeded (429) if response.status_code == 429: - logger.warning("Google Search API quota exceeded (429). Switching to fallback search engines.") _google_quota_exhausted = True - # Set reset time to next day midnight UTC import datetime tomorrow = datetime.datetime.utcnow().replace(hour=0, minute=0, second=0, microsecond=0) + datetime.timedelta(days=1) _google_quota_reset_time = tomorrow.timestamp() - return [] + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message="Google API 配额已用尽" + ) response.raise_for_status() data = response.json() @@ -121,53 +649,49 @@ class SearchService: results = [] if 'items' in data: for item in data['items']: - logger.debug(f"Search Item: {item.get('title')} - {item.get('link')}") - results.append({ - 'title': item.get('title'), - 'link': item.get('link'), - 'snippet': item.get('snippet'), - 'source': 'Google', - 'published': item.get('pagemap', {}).get('metatags', [{}])[0].get('article:published_time', '') - }) - else: - logger.warning(f"Google Search returned no 'items'. Full response: {json.dumps(data, ensure_ascii=False)}") - - return results + results.append(SearchResult( + title=item.get('title', ''), + snippet=item.get('snippet', ''), + url=item.get('link', ''), + source='Google', + published_date=item.get('pagemap', {}).get('metatags', [{}])[0].get('article:published_time', ''), + )) + + return SearchResponse( + query=query, + results=results, + provider=self.name, + success=True, + ) - except requests.exceptions.HTTPError as e: - if hasattr(e, 'response') and e.response is not None and e.response.status_code == 429: - logger.warning("Google Search API quota exceeded. Switching to fallback.") - _google_quota_exhausted = True - import datetime - tomorrow = datetime.datetime.utcnow().replace(hour=0, minute=0, second=0, microsecond=0) + datetime.timedelta(days=1) - _google_quota_reset_time = tomorrow.timestamp() - else: - logger.error(f"Google search failed: {e}") - if hasattr(e, 'response') and e.response is not None: - logger.error(f"Response: {e.response.text}") - return [] except Exception as e: - logger.error(f"Google search failed: {e}") - return [] + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=str(e) + ) - def _search_bing(self, query: str, num_results: int) -> List[Dict[str, Any]]: - """Bing search.""" - api_key = self._config.get('bing', {}).get('api_key') - - if not api_key: - logger.warning("Bing Search is not configured (missing api_key).") - return [] - - url = "https://api.bing.microsoft.com/v7.0/search" - headers = {"Ocp-Apim-Subscription-Key": api_key} - params = { - "q": query, - "count": num_results, - "textDecorations": True, - "textFormat": "HTML" - } - + +class BingSearchProvider(BaseSearchProvider): + """Bing Search API 搜索引擎""" + + def __init__(self, api_key: str): + super().__init__([api_key] if api_key else [], "Bing") + + def _do_search(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse: + """执行 Bing 搜索""" try: + url = "https://api.bing.microsoft.com/v7.0/search" + headers = {"Ocp-Apim-Subscription-Key": api_key} + params = { + "q": query, + "count": max_results, + "textDecorations": True, + "textFormat": "HTML" + } + response = requests.get(url, headers=headers, params=params, timeout=10) response.raise_for_status() data = response.json() @@ -175,26 +699,41 @@ class SearchService: results = [] if 'webPages' in data and 'value' in data['webPages']: for item in data['webPages']['value']: - results.append({ - 'title': item.get('name'), - 'link': item.get('url'), - 'snippet': item.get('snippet'), - 'source': 'Bing', - 'published': item.get('datePublished', '') - }) - return results + results.append(SearchResult( + title=item.get('name', ''), + snippet=item.get('snippet', ''), + url=item.get('url', ''), + source='Bing', + published_date=item.get('datePublished', ''), + )) + + return SearchResponse( + query=query, + results=results, + provider=self.name, + success=True, + ) except Exception as e: - logger.error(f"Bing search failed: {e}") - return [] + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=str(e) + ) - def _search_duckduckgo(self, query: str, num_results: int) -> List[Dict[str, Any]]: - """ - DuckDuckGo search (free, no API key required). - Uses the DuckDuckGo HTML search endpoint. - """ + +class DuckDuckGoSearchProvider(BaseSearchProvider): + """DuckDuckGo 搜索引擎(免费,无需 API Key)""" + + def __init__(self): + super().__init__(['free'], "DuckDuckGo") + + def _do_search(self, query: str, api_key: str, max_results: int, days: int = 7) -> SearchResponse: + """执行 DuckDuckGo 搜索""" try: - # Use DuckDuckGo Instant Answer API + # 使用 DuckDuckGo Instant Answer API url = "https://api.duckduckgo.com/" params = { 'q': query, @@ -209,65 +748,64 @@ class SearchService: results = [] - # Get results from RelatedTopics + # 获取 RelatedTopics related_topics = data.get('RelatedTopics', []) - for topic in related_topics[:num_results]: + for topic in related_topics[:max_results]: if isinstance(topic, dict): if 'FirstURL' in topic: - results.append({ - 'title': topic.get('Text', '')[:100], - 'link': topic.get('FirstURL', ''), - 'snippet': topic.get('Text', ''), - 'source': 'DuckDuckGo', - 'published': '' - }) - # Handle nested topics + results.append(SearchResult( + title=topic.get('Text', '')[:100], + snippet=topic.get('Text', ''), + url=topic.get('FirstURL', ''), + source='DuckDuckGo', + )) elif 'Topics' in topic: for sub_topic in topic['Topics']: - if len(results) >= num_results: + if len(results) >= max_results: break if 'FirstURL' in sub_topic: - results.append({ - 'title': sub_topic.get('Text', '')[:100], - 'link': sub_topic.get('FirstURL', ''), - 'snippet': sub_topic.get('Text', ''), - 'source': 'DuckDuckGo', - 'published': '' - }) + results.append(SearchResult( + title=sub_topic.get('Text', '')[:100], + snippet=sub_topic.get('Text', ''), + url=sub_topic.get('FirstURL', ''), + source='DuckDuckGo', + )) - # Also check AbstractURL and AbstractText - if data.get('AbstractURL') and len(results) < num_results: - results.insert(0, { - 'title': data.get('Heading', query), - 'link': data.get('AbstractURL', ''), - 'snippet': data.get('AbstractText', ''), - 'source': 'DuckDuckGo', - 'published': '' - }) + # 检查 AbstractURL + if data.get('AbstractURL') and len(results) < max_results: + results.insert(0, SearchResult( + title=data.get('Heading', query), + snippet=data.get('AbstractText', ''), + url=data.get('AbstractURL', ''), + source='DuckDuckGo', + )) - # If no results from Instant Answer, try HTML scraping as fallback + # 如果没有结果,尝试 HTML 版本 if not results: - results = self._search_duckduckgo_html(query, num_results) + results = self._search_html(query, max_results) - if results: - logger.info(f"DuckDuckGo search returned {len(results)} results") - - return results[:num_results] + return SearchResponse( + query=query, + results=results[:max_results], + provider=self.name, + success=len(results) > 0, + ) except Exception as e: - logger.error(f"DuckDuckGo search failed: {e}") - # Try HTML fallback - return self._search_duckduckgo_html(query, num_results) - - def _search_duckduckgo_html(self, query: str, num_results: int) -> List[Dict[str, Any]]: - """ - DuckDuckGo HTML search fallback. - Scrapes the lite HTML version for better results. - """ + return SearchResponse( + query=query, + results=[], + provider=self.name, + success=False, + error_message=str(e) + ) + + def _search_html(self, query: str, max_results: int) -> List[SearchResult]: + """DuckDuckGo HTML 搜索备选""" try: url = "https://lite.duckduckgo.com/lite/" headers = { - 'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36' + 'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36' } data = {'q': query} @@ -275,33 +813,234 @@ class SearchService: response.raise_for_status() results = [] - - # Simple HTML parsing without BeautifulSoup html = response.text - # Find all result links (they have class="result-link") - import re - - # Pattern to find result entries link_pattern = r']*class="result-link"[^>]*href="([^"]*)"[^>]*>([^<]*)' snippet_pattern = r'