test_basic.py

#!/usr/bin/env python3
"""
AI Quant 项目基础测试脚本
用于验证项目核心功能是否正常运行
"""

import sys
import os
import logging
import unittest

# 添加项目根目录到路径
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.StreamHandler()
    ]
)

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


class TestBasicFunctionality(unittest.TestCase):
    """
    基础功能测试类
    """
    def test_imports(self):
        """
        测试核心模块导入
        """
        logger.info("测试核心模块导入...")

        # 测试配置模块
        from config import zs_sec_config
        self.assertIsNotNone(zs_sec_config.ZS_SEC_CONFIG)
        logger.info("配置模块导入成功")

        # 测试招商证券接口模块
        from core.api import zs_sec
        self.assertIsNotNone(zs_sec.ZSSecApi)
        self.assertIsNotNone(zs_sec.ZSSecConfig)
        logger.info("招商证券接口模块导入成功")

        # 测试单票量化策略模块
        from app.strategy import single_stock
        self.assertIsNotNone(single_stock.SingleStockAnalyzer)
        self.assertIsNotNone(single_stock.StockAnalysisResult)
        logger.info("单票量化策略模块导入成功")

        # 测试大模型推理模块
        from service.llm import infer
        self.assertIsNotNone(infer.LLMInferenceService)
        logger.info("大模型推理模块导入成功")

        logger.info("所有核心模块导入成功")

    def test_config_validation(self):
        """
        测试配置验证功能
        """
        logger.info("测试配置验证功能...")

        from core.api import zs_sec

        # 测试无效配置
        invalid_config = zs_sec.ZSSecConfig()
        self.assertFalse(invalid_config.validate())
        logger.info("无效配置验证成功")

        # 测试有效配置
        valid_config = zs_sec.ZSSecConfig(
            server_address="180.168.146.187",
            port=10100,
            account="123456",
            password="123456",
            broker_code="0000"
        )
        self.assertTrue(valid_config.validate())
        logger.info("有效配置验证成功")

    def test_analyzer_basic(self):
        """
        测试分析器基础功能
        """
        logger.info("测试分析器基础功能...")

        from app.strategy import single_stock
        import pandas as pd

        # 创建分析器
        analyzer = single_stock.SingleStockAnalyzer()
        self.assertIsNotNone(analyzer)
        logger.info("分析器创建成功")

        # 创建测试数据
        data = {
            "open": [10, 11, 12, 11, 13],
            "high": [12, 13, 14, 12, 15],
            "low": [9, 10, 11, 10, 12],
            "close": [11, 12, 13, 11, 14],
            "volume": [10000, 12000, 15000, 11000, 16000]
        }
        df = pd.DataFrame(data)

        # 分析股票
        result = analyzer.analyze("000001.SZ", df)
        self.assertIsNotNone(result)
        self.assertEqual(result.symbol, "000001.SZ")
        self.assertGreater(result.score, 0)
        logger.info(
            f"分析成功,评分: {result.score:.2f}, 等级: {result.rating}, "
            f"信号: {result.signal}")

    def test_llm_service_basic(self):
        """
        测试大模型推理服务基础功能
        """
        logger.info("测试大模型推理服务基础功能...")

        from service.llm import infer

        # 创建服务
        service = infer.LLMInferenceService()
        self.assertIsNotNone(service)
        logger.info("大模型推理服务创建成功")

        # 测试未启用状态
        response = service.infer("请分析股票 000001.SZ 的投资价值")
        self.assertIsNone(response)
        logger.info("未启用状态测试成功")

        # 测试启动和关闭(不依赖实际连接成功)
        service.config["enabled"] = True
        # 由于没有实际安装 ollama 或 fingt,我们不测试实际启动成功
        logger.info("大模型推理服务配置更新成功(实际启动需要安装相应的大模型库)")

        service.shutdown()
        logger.info("大模型推理服务关闭成功")


def main():
    """
    主入口函数
    """
    logger.info("AI Quant 项目基础测试开始")

    try:
        # 创建测试套件
        suite = unittest.TestLoader().loadTestsFromTestCase(
            TestBasicFunctionality)

        # 运行测试
        runner = unittest.TextTestRunner(verbosity=2)
        result = runner.run(suite)

        if result.wasSuccessful():
            logger.info("所有测试成功通过")
            return 0
        else:
            logger.error(
                f"测试失败: {len(result.failures)} 个失败, {len(result.errors)} 个错误")
            return 1
    except Exception as e:
        logger.error(f"测试过程中出错: {e}")
        logger.exception("详细错误信息:")
        return 2


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