simulate_test.py

#!/usr/bin/env python3
"""
模拟测试脚本
用于在没有账号的情况下进行模拟测试
"""

import sys
import os
import logging
import argparse
from datetime import datetime, timedelta

# 添加项目根目录到路径
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

# 配置日志
logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
    handlers=[
        logging.FileHandler("simulate_test.log"),
        logging.StreamHandler()
    ]
)

logger = logging.getLogger("AIQuant.simulate_test")


def main():
    """
    主入口函数
    """
    logger.info("AI Quant 模拟测试脚本启动")

    try:
        # 解析命令行参数
        parser = argparse.ArgumentParser(
            description="AI Quant 模拟测试脚本"
        )
        parser.add_argument(
            "--stock-count",
            type=int,
            default=10,
            help="测试的股票数量(默认:10)"
        )
        parser.add_argument(
            "--data-days",
            type=int,
            default=365,
            help="测试数据的天数(默认:365)"
        )
        parser.add_argument(
            "--frequency",
            choices=["d", "h", "m"],
            default="d",
            help="数据频率(d: 日线,h: 小时线,m: 分钟线,默认:d)"
        )
        parser.add_argument(
            "--save-cache",
            action="store_true",
            help="保存测试数据到本地缓存"
        )
        parser.add_argument(
            "--load-cache",
            action="store_true",
            help="从本地缓存加载测试数据"
        )
        parser.add_argument(
            "--clear-cache",
            action="store_true",
            help="测试完成后清理缓存"
        )

        args = parser.parse_args()

        logger.info(
            f"测试参数: 股票数量={args.stock_count}, "
            f"数据天数={args.data_days}, 频率={args.frequency}")

        # 导入所需模块
        from service.data.mock_data import (generate_random_kline_data,
                                           generate_stock_pool,
                                           generate_mock_account_data)
        from service.data.data_cache import (save_data, load_data,
                                           data_exists, clear_all_cache)
        from app.strategy.single_stock import SingleStockAnalyzer

        # 生成股票池
        logger.info("生成股票池...")
        stock_pool = generate_stock_pool(args.stock_count)
        logger.info(f"股票池: {stock_pool}")

        # 生成模拟账户数据
        logger.info("生成模拟账户数据...")
        account_data = generate_mock_account_data()
        logger.info(f"账户余额: {account_data['balance']}")

        # 初始化分析器
        logger.info("初始化单票量化分析器...")
        analyzer = SingleStockAnalyzer()

        # 测试结果
        test_results = []

        # 遍历股票池,进行分析
        for symbol in stock_pool:
            logger.info(f"\n分析股票: {symbol}")

            if args.load_cache and data_exists(symbol, args.frequency):
                # 从缓存加载数据
                logger.info("从缓存加载数据...")
                data = load_data(symbol, args.frequency)
            else:
                # 生成模拟数据
                logger.info("生成模拟数据...")
                data = generate_random_kline_data(
                    symbol,
                    start_date="2023-01-01",
                    end_date=(datetime(2023, 1, 1) +
                              timedelta(days=args.data_days - 1)
                              ).strftime("%Y-%m-%d"),
                    frequency=args.frequency
                )

                if args.save_cache:
                    logger.info("保存数据到缓存...")
                    save_data(symbol, data, args.frequency)

            # 分析股票
            result = analyzer.analyze(symbol, data)
            logger.info(f"分析结果: {result}")

            test_results.append(result)

        # 统计测试结果
        logger.info("\n测试结果统计:")
        logger.info(f"总股票数量: {len(test_results)}")

        # 按等级统计
        rating_counts = {
            "S": 0,
            "A": 0,
            "B": 0,
            "C": 0
        }

        for result in test_results:
            if result.rating in rating_counts:
                rating_counts[result.rating] += 1

        logger.info("标的等级分布:")
        for rating, count in rating_counts.items():
            logger.info(
                f"{rating}级: {count}只 ({count/len(test_results)*100:.1f}%)")

        # 按交易信号统计
        signal_counts = {
            "买入": 0,
            "持有": 0,
            "卖出": 0
        }

        for result in test_results:
            if result.signal in signal_counts:
                signal_counts[result.signal] += 1

        logger.info("交易信号分布:")
        for signal, count in signal_counts.items():
            logger.info(
                f"{signal}: {count}只 ({count/len(test_results)*100:.1f}%)")

        # 计算平均评分
        avg_score = sum(
            result.score for result in test_results) / len(test_results)
        logger.info(f"平均评分: {avg_score:.2f}")

        # 清理缓存
        if args.clear_cache:
            logger.info("清理数据缓存...")
            clear_all_cache()

        logger.info("模拟测试完成")
        return 0
    except Exception as e:
        logger.error(f"模拟测试过程中出错: {e}")
        logger.exception("详细错误信息:")
        return 1


if __name__ == "__main__":
    try:
        sys.exit(main())
    except KeyboardInterrupt:
        logger.info("程序被用户中断")
        sys.exit(0)