import pandas as pd
from typing import Literal

def calculate_atr(price_data: pd.DataFrame, period: int = 14, 
                 smoothing: Literal['SMA', 'EMA'] = 'SMA') -> pd.DataFrame:
    """
    计算ATR值
    :param price_data: DataFrame包含日期、最高价、最低价、收盘价
    :param period: ATR计算周期
    :param smoothing: 平滑方法 'SMA'或'EMA'
    :return: 包含TR和ATR的DataFrame
    """
    df = price_data.copy()
    
    # 计算TR
    df['prev_close'] = df['close'].shift(1)
    df['high_low'] = df['high'] - df['low']
    df['high_prev_close'] = (df['high'] - df['prev_close']).abs()
    df['low_prev_close'] = (df['low'] - df['prev_close']).abs()
    df['TR'] = df[['high_low', 'high_prev_close', 'low_prev_close']].max(axis=1)
    
    # 首日TR处理（确保有足够数据）
    if len(df) > 0:
        df.loc[0, 'TR'] = df.loc[0, 'high_low']
    else:
        raise ValueError("价格数据为空，无法计算ATR")
    
    # 计算ATR
    if smoothing == 'SMA':
        df['ATR'] = df['TR'].rolling(window=period).mean()
    elif smoothing == 'EMA':
        df['ATR'] = df['TR'].ewm(span=period, adjust=False).mean()
    
    # 保留所需列
    result = df[['date', 'high', 'low', 'close', 'TR', 'ATR']].copy()
    result = result.round(4)
    
    # 确保返回DataFrame类型
    return pd.DataFrame(result)