infer.py

"""
大模型推理轻量服务
独立插件式设计,不启动不占用资源,与交易核心隔离
预留 Ollama/FinGPT 接口
"""

import logging
from typing import Optional, Dict, Any, List
import importlib.util
import os


# 大模型推理模块日志
logger = logging.getLogger("AIQuant.llm")
logger.setLevel(logging.INFO)


class LLMInferenceService:
    """
    大模型推理服务
    独立插件式设计,不启动不占用资源,与交易核心隔离
    """
    def __init__(self, config: Dict[str, Any] = None):
        """
        初始化大模型推理服务

        Args:
            config: 推理服务配置
        """
        self.config = config or {}
        self.enabled = self.config.get("enabled", False)
        self.provider = self.config.get("provider", "ollama")
        self.model = self.config.get("model", "llama2")
        self.client = None

    def startup(self) -> bool:
        """
        启动大模型推理服务

        Returns:
            启动是否成功的布尔值
        """
        if not self.enabled:
            logger.warning("大模型推理服务未启用")
            return False

        logger.info(f"启动大模型推理服务,提供商: {self.provider}, 模型: {self.model}")

        try:
            if self.provider == "ollama":
                self.client = self._init_ollama_client()
            elif self.provider == "fingt":
                self.client = self._init_fingpt_client()
            else:
                logger.error(f"不支持的大模型提供商: {self.provider}")
                return False

            logger.info("大模型推理服务启动成功")
            return True
        except Exception as e:
            logger.error(f"大模型推理服务启动失败: {e}")
            return False

    def shutdown(self) -> bool:
        """
        关闭大模型推理服务

        Returns:
            关闭是否成功的布尔值
        """
        if not self.enabled:
            logger.warning("大模型推理服务未启用")
            return False

        logger.info("关闭大模型推理服务")
        try:
            self.client = None
            logger.info("大模型推理服务关闭成功")
            return True
        except Exception as e:
            logger.error(f"大模型推理服务关闭失败: {e}")
            return False

    def _init_ollama_client(self) -> Any:
        """
        初始化 Ollama 客户端

        Returns:
            Ollama 客户端对象
        """
        try:
            import ollama
            return ollama.Client()
        except ImportError:
            logger.error("Ollama 客户端库未安装,请执行: pip install ollama")
            raise

    def _init_fingpt_client(self) -> Any:
        """
        初始化 FinGPT 客户端

        Returns:
            FinGPT 客户端对象
        """
        try:
            # 这里需要根据 FinGPT 实际的 API 进行初始化
            logger.warning("FinGPT 客户端尚未完全实现,使用模拟对象")
            return None
        except ImportError as e:
            logger.error(f"FinGPT 客户端库未安装: {e}")
            raise

    def infer(self,
              prompt: str,
              temperature: float = 0.7,
              max_tokens: int = 512) -> Optional[str]:
        """
        执行大模型推理

        Args:
            prompt: 推理提示
            temperature: 推理温度
            max_tokens: 最大生成 tokens 数

        Returns:
            推理结果字符串(如果成功),否则 None
        """
        if not self.enabled or not self.client:
            logger.warning("大模型推理服务未启动")
            return None

        logger.debug(f"执行大模型推理,提示长度: {len(prompt)}")

        try:
            if self.provider == "ollama":
                response = self.client.generate(
                    model=self.model,
                    prompt=prompt,
                    options={
                        "temperature": temperature,
                        "num_predict": max_tokens
                    }
                )
                return response["response"]
            elif self.provider == "fingt":
                # 这里需要根据 FinGPT 实际的 API 进行推理
                logger.warning("FinGPT 推理尚未完全实现,返回模拟结果")
                return f"FinGPT 推理结果: {prompt}"
            else:
                logger.error(f"不支持的大模型提供商: {self.provider}")
                return None
        except Exception as e:
            logger.error(f"大模型推理失败: {e}")
            return None

    def analyze_stock_data(self,
                           symbol: str,
                           data: Dict[str, Any]) -> Optional[str]:
        """
        分析股票数据(专业方法)

        Args:
            symbol: 股票代码
            data: 股票数据

        Returns:
            分析结果字符串(如果成功),否则 None
        """
        prompt = f"""
        请分析股票 {symbol} 的以下数据:
        {data}

        请从以下几个方面进行分析:
        1. 价格趋势
        2. 成交量情况
        3. 技术指标
        4. 风险评估
        5. 投资建议

        要求:
        - 分析内容简洁明了
        - 使用专业术语
        - 突出关键要点
        - 给出具体的投资建议
        """

        return self.infer(prompt, temperature=0.4, max_tokens=1024)

    def generate_strategy(self, conditions: Dict[str, Any]) -> Optional[str]:
        """
        生成交易策略(专业方法)

        Args:
            conditions: 策略条件

        Returns:
            策略代码字符串(如果成功),否则 None
        """
        prompt = f"""
        请根据以下条件生成一个 A 股量化交易策略:
        {conditions}

        要求:
        - 策略代码使用 Python 语言
        - 遵循 VNpy 策略开发规范
        - 包含明确的买入和卖出信号
        - 包含风险控制机制
        - 代码结构清晰,注释详细

        返回:
        - 完整的策略代码
        - 策略说明文档
        """

        return self.infer(prompt, temperature=0.3, max_tokens=2048)


class LLMPluginManager:
    """
    大模型插件管理器
    用于动态加载和管理大模型插件
    """
    def __init__(self):
        """
        初始化插件管理器
        """
        self.plugins = {}
        self.plugin_dir = "plugins/llm"

    def load_plugins(self) -> List[str]:
        """
        加载所有大模型插件

        Returns:
            成功加载的插件名称列表
        """
        loaded_plugins = []

        # 检查插件目录是否存在
        if not os.path.exists(self.plugin_dir):
            logger.warning(f"插件目录 {self.plugin_dir} 不存在,创建目录")
            os.makedirs(self.plugin_dir, exist_ok=True)
            return loaded_plugins

        # 遍历插件目录
        for filename in os.listdir(self.plugin_dir):
            if filename.endswith(".py") and filename != "__init__.py":
                plugin_name = filename[:-3]
                try:
                    self.load_plugin(plugin_name)
                    loaded_plugins.append(plugin_name)
                except Exception as e:
                    logger.error(f"加载插件 {plugin_name} 失败: {e}")

        logger.info(f"成功加载 {len(loaded_plugins)} 个大模型插件")
        return loaded_plugins

    def load_plugin(self, plugin_name: str) -> bool:
        """
        加载单个大模型插件

        Args:
            plugin_name: 插件名称

        Returns:
            加载是否成功的布尔值
        """
        plugin_path = os.path.join(self.plugin_dir, f"{plugin_name}.py")

        if not os.path.exists(plugin_path):
            logger.error(f"插件文件不存在: {plugin_path}")
            return False

        try:
            # 导入插件模块
            spec = importlib.util.spec_from_file_location(
                plugin_name, plugin_path)
            module = importlib.util.module_from_spec(spec)
            spec.loader.exec_module(module)

            # 检查插件是否有正确的接口
            if hasattr(module, "LLMProvider"):
                self.plugins[plugin_name] = module.LLMProvider()
                logger.info(f"成功加载大模型插件: {plugin_name}")
                return True
            else:
                logger.error(f"插件 {plugin_name} 缺少 LLMProvider 类")
                return False
        except Exception as e:
            logger.error(f"加载插件 {plugin_name} 失败: {e}")
            return False

    def get_plugin(self, plugin_name: str) -> Optional[Any]:
        """
        获取插件实例

        Args:
            plugin_name: 插件名称

        Returns:
            插件实例(如果存在),否则 None
        """
        return self.plugins.get(plugin_name)


# 默认配置
LLM_DEFAULT_CONFIG = {
    "enabled": False,
    "provider": "ollama",
    "model": "llama2",
    "temperature": 0.7,
    "max_tokens": 512
}


# 工厂函数
def create_llm_service(config: Dict[str, Any] = None) -> LLMInferenceService:
    """
    创建大模型推理服务实例

    Args:
        config: 配置参数

    Returns:
        大模型推理服务实例
    """
    return LLMInferenceService(config)