from typing import Literal
from atr_calculator import calculate_atr
from data_fetcher import get_historical_data
from utils import (
    validate_stock_code,
    is_trading_day,
    get_previous_trading_day,
    get_stock_name
)
import data_fetcher
import atr_calculator
import utils
import pandas as pd
import os
from datetime import datetime, timedelta

def main():
    # 获取用户输入
    stock_code = input("请输入股票代码(如600000.SH): ")
    base_date = input("请输入基准日期(YYYYMMDD格式，默认为今天): ") or datetime.now().strftime('%Y%m%d')
    atr_period = int(input("请输入ATR计算周期(1-300，默认14): ") or 14)
    smoothing_input = input("请选择平滑方法(SMA/EMA，默认SMA): ").upper() or "SMA"
    smoothing: Literal['SMA', 'EMA'] = 'SMA' if smoothing_input not in ('SMA', 'EMA') else smoothing_input

    # 校验输入
    if not utils.validate_stock_code(stock_code):
        print("无此股票数据，请检查输入")
        return

    if not utils.is_trading_day(base_date):
        base_date = utils.get_previous_trading_day(base_date)
        print(f"调整为前一交易日: {base_date}")

    # 获取数据
    try:
        price_data = data_fetcher.get_historical_data(stock_code, base_date, 60)
    except Exception as e:
        print(f"数据获取失败: {str(e)}")
        return

    # 计算ATR
    atr_data = atr_calculator.calculate_atr(price_data, atr_period, smoothing)
    
    # 导出CSV
    export_csv(stock_code, base_date, atr_data)

def export_csv(stock_code, base_date, atr_data):
    stock_name = utils.get_stock_name(stock_code)
    filename = f"{stock_code}_{stock_name}_ATR_{base_date}_60日.csv"
    atr_data.to_csv(filename, index=False, encoding='utf-8-sig')
    print(f"ATR数据已导出到: {os.path.abspath(filename)}")

if __name__ == "__main__":
    main()