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}"