common.py
"""
QTrading 通用工具模块
"""
import datetime
import re
from typing import Any, List, Optional, Union
def validate_date(date_str: str) -> bool:
"""
验证日期格式(YYYYMMDD)
Args:
date_str: 日期字符串
Returns:
bool: 是否是有效日期格式
"""
if not isinstance(date_str, str) or len(date_str) != 8:
return False
try:
year = int(date_str[0:4])
month = int(date_str[4:6])
day = int(date_str[6:8])
datetime.date(year, month, day)
return True
except Exception:
return False
def validate_code(code: str) -> bool:
"""
验证股票代码格式
Args:
code: 股票代码
Returns:
bool: 是否是有效股票代码格式
"""
if not isinstance(code, str):
return False
# A股代码格式:6位数字,000/002/300/600/601/603开头
pattern = r'^(000|002|300|600|601|603)\d{3}$'
return bool(re.match(pattern, code))
def format_date(date: Union[str, datetime.date, datetime.datetime]) -> str:
"""
格式化日期为 YYYYMMDD 格式
Args:
date: 日期对象或字符串
Returns:
str: 格式化后的日期字符串
"""
if isinstance(date, datetime.datetime):
return date.strftime('%Y%m%d')
elif isinstance(date, datetime.date):
return date.strftime('%Y%m%d')
elif isinstance(date, str):
# 尝试解析常见日期格式
for fmt in ['%Y-%m-%d', '%Y/%m/%d', '%Y%m%d']:
try:
return datetime.datetime.strptime(date, fmt).strftime('%Y%m%d')
except Exception:
continue
raise ValueError(f"无效的日期格式: {date}")
else:
raise ValueError(f"无效的日期类型: {type(date)}")
def today_str() -> str:
"""
获取今天日期的字符串表示(YYYYMMDD)
Returns:
str: 今天日期字符串
"""
return datetime.datetime.now().strftime('%Y%m%d')
def get_date_range(start_date: str, end_date: str) -> List[str]:
"""
获取日期范围列表
Args:
start_date: 开始日期(YYYYMMDD)
end_date: 结束日期(YYYYMMDD)
Returns:
List[str]: 日期范围列表
"""
dates = []
start = datetime.datetime.strptime(start_date, '%Y%m%d')
end = datetime.datetime.strptime(end_date, '%Y%m%d')
current = start
while current <= end:
dates.append(current.strftime('%Y%m%d'))
current += datetime.timedelta(days=1)
return dates
def calculate_change_percent(current: float, previous: float) -> float:
"""
计算涨跌幅百分比
Args:
current: 当前值
previous: 前值
Returns:
float: 涨跌幅百分比(保留2位小数)
"""
if previous == 0:
return 0.0
return round(((current - previous) / previous) * 100, 2)
def calculate_volatility(data: List[float], period: int = 20) -> float:
"""
计算波动率(标准差)
Args:
data: 数据列表
period: 计算周期(默认20)
Returns:
float: 波动率
"""
if len(data) < period:
return 0.0
recent_data = data[-period:]
mean = sum(recent_data) / len(recent_data)
variance = sum((x - mean) ** 2 for x in recent_data) / len(recent_data)
return variance ** 0.5
def calculate_drawdown(prices: List[float]) -> float:
"""
计算最大回撤
Args:
prices: 价格序列
Returns:
float: 最大回撤(百分比)
"""
if not prices:
return 0.0
peak = prices[0]
max_drawdown = 0.0
for price in prices:
if price > peak:
peak = price
drawdown = (peak - price) / peak
if drawdown > max_drawdown:
max_drawdown = drawdown
return round(max_drawdown * 100, 2)
def normalize_score(score: float, min_score: float = 0, max_score: float = 100) -> float:
"""
归一化分数到 0-100 范围
Args:
score: 原始分数
min_score: 最小分数(默认0)
max_score: 最大分数(默认100)
Returns:
float: 归一化后的分数
"""
if max_score == min_score:
return 50.0
normalized = (score - min_score) / (max_score - min_score) * 100
return max(0, min(100, round(normalized, 2)))
def is_valid_number(value: Any) -> bool:
"""
检查值是否是有效的数字类型
Args:
value: 待检查的值
Returns:
bool: 是否是有效数字
"""
try:
float(value)
return True
except (TypeError, ValueError):
return False
def format_number(value: float, decimals: int = 2) -> str:
"""
格式化数字为字符串
Args:
value: 数字值
decimals: 小数位数(默认2)
Returns:
str: 格式化后的字符串
"""
return f"{value:.{decimals}f}"