""" 基于 smolagents 框架的 Hugging Face 数据集自然语言查询智能代理 功能: 1. 启动时自动获取数据集文件列表,作为系统提示词告知用户可用数据 2. 接收用户自然语言输入,从中提取数据集查询所需的参数信息 3. 使用提取的参数从 Hugging Face 数据集获取数据(支持 ZIP/CSV/Excel 等格式) 4. 返回结构化的查询结果 依赖安装: pip install 'smolagents[gradio]' requests pandas huggingface_hub openpyxl 使用前: 1. 将文件顶部 HF_TOKEN 常量的值替换为个人 Hugging Face Settings 下申请的 Read Token(用于数据集访问) 2. 将文件顶部 OPENAI_API_KEY 常量的值替换为你的 OpenAI API Key(用于 LLM 接入) 3. 如使用第三方兼容服务,修改 OPENAI_API_BASE 常量的地址 """ import os import json import logging import re from typing import Optional, List from huggingface_hub import HfApi from smolagents import Tool, ToolCallingAgent, OpenAIServerModel, GradioUI # 导入专用查询函数(从 query_tools 子模块导入) from query_tools.logbook_availability_query import query_logbook_availability from query_tools.query_gfw import query_gfw # --------------------------------------------------------------------------- # 日志配置(提前配置,以便在后续代码中使用) # --------------------------------------------------------------------------- logging.basicConfig( level=logging.INFO, format="%(asctime)s [%(levelname)s] %(name)s - %(message)s", ) logger = logging.getLogger("hf_data_agent") # =========================================================================== # Hugging Face 数据集访问 Token(从环境变量加载) # =========================================================================== HF_TOKEN = os.getenv('HF_TOKEN') # 从环境变量 HF_TOKEN 读取 # 启动时校验 HF_TOKEN 是否存在,并给出提示 if not HF_TOKEN: logger.warning( "未找到环境变量 HF_TOKEN,将无法访问 Hugging Face 数据集。" "请设置环境变量: export HF_TOKEN=hf_你的Token" ) # =========================================================================== # 数据集仓库 ID # =========================================================================== HF_DATASET_REPO = "squid-lab/squid_dataset" # Hugging Face 数据集仓库 ID # =========================================================================== # OpenAI 模型配置常量(从环境变量加载) # =========================================================================== OPENAI_API_KEY = os.getenv('OPENAI_API_KEY') # 从环境变量 OPENAI_API_KEY 读取 OPENAI_API_BASE = os.getenv('OPENAI_API_BASE', 'https://api.deepseek.com/v1') # 默认 DeepSeek API 地址 MODEL_ID = os.getenv('MODEL_ID', 'deepseek-v4-flash') # 默认使用 deepseek-v4-flash(推荐模型) # 启动时校验 OPENAI_API_KEY 是否存在,并给出提示 if not OPENAI_API_KEY: logger.warning( "未找到环境变量 OPENAI_API_KEY,Agent 将无法调用 LLM。" "请设置环境变量: export OPENAI_API_KEY=sk_你的Key" ) # 注意:deepseek-chat 和 deepseek-reasoner 将于 2026/07/24 弃用 # 推荐使用新模型:deepseek-v4-flash(快速)或 deepseek-v4-pro(专业) # =========================================================================== # DeepSeek 思考模式配置(已停用) # =========================================================================== # 注意:DeepSeek 思考模式已停用,原因如下: # 1. deepseek-chat 和 deepseek-reasoner 将于 2026/07/24 弃用 # 2. 新模型 deepseek-v4-flash/v4-pro 不需要特殊的思考模式适配 # 3. 思考模式不支持 tool_choice 参数,限制 Agent 工具调用能力 # 4. 使用标准 OpenAIServerModel 更稳定,兼容性更好 THINKING_MODE_ENABLED = False # 已停用思考模式 THINKING_EFFORT = "high" # 保留参数但不再使用 # =========================================================================== # 数据集文件列表获取函数 # =========================================================================== def get_dataset_file_list() -> List[str]: """ 从 Hugging Face 数据集仓库获取一级文件和文件夹名称列表。 使用 HfApi.dataset_info() 获取数据集元信息,提取 siblings 中的文件路径, 只保留一级目录/文件(不包含深层嵌套的文件)。 Returns: 一级文件和文件夹名称列表,如 ['中西太平洋WCPFC/', 'README.md', 'data.csv'] """ logger.info("正在获取数据集文件列表: %s", HF_DATASET_REPO) try: api = HfApi(token=HF_TOKEN) dataset_info = api.dataset_info(repo_id=HF_DATASET_REPO) # 提取所有文件路径 all_files = [sibling.rfilename for sibling in dataset_info.siblings] # 只保留一级目录/文件(路径中不包含 '/' 的项,或者一级文件夹) first_level_items = set() for file_path in all_files: parts = file_path.split("/") # 一级文件(无子目录) if len(parts) == 1: first_level_items.add(parts[0]) # 一级文件夹(取第一部分,添加 '/' 后缀标识为文件夹) else: first_level_items.add(parts[0] + "/") file_list = sorted(list(first_level_items)) logger.info("获取到 %d 个一级文件/文件夹: %s", len(file_list), file_list) return file_list except Exception as e: logger.error("获取数据集文件列表失败: %s", e, exc_info=True) return [] def build_system_prompt_with_file_list() -> str: """ 构建包含数据集文件列表的系统提示词。 在对话启动时告知用户当前可用的数据集文件,帮助用户了解可查询的内容。 Returns: 格式化的系统提示词字符串 """ file_list = get_dataset_file_list() if not file_list: return ( "你是一个 Hugging Face 数据集查询助手。" "用户可以询问数据集内容,你将帮助用户查询和获取数据。" "当前无法获取数据集文件列表,请检查 HF_TOKEN 配置是否正确。" ) # 格式化文件列表为提示词 files_str = "\n".join([f" - {item}" for item in file_list]) system_prompt = ( "你是一个渔业数据查询助手,拥有以下工具:\n\n" "1. **logbook_query** — 查询 logbook(捕捞日志)数据可用性,支持按海区、年份、物种、数据类型筛选\n" "2. **gfw_query** — 查询 GFW(Global Fishing Watch)渔船作业努力量,需要指定年份(2012-2024)、月份和空间范围\n" "3. **literature_cpue_query** — 查询 CPUE(单位捕捞努力量渔获量)相关文献数据,支持按海区、年份、物种、响应变量筛选\n" "4. **sprfmo_query** — 查询 SPRFMO(南太平洋区域渔业管理组织)数据,支持捕捞量或努力量查询,可按国家、年份、物种等筛选\n\n" "当前可查询的 Hugging Face 数据集仓库: " + HF_DATASET_REPO + "\n\n" "数据集一级文件和文件夹列表:\n" + files_str + "\n\n" "用户查询示例:\n" " - '2010-2020年东南太平洋有哪些鱿鱼捕捞日志数据?' → 使用 logbook_query\n" " - '查询2012年1月太平洋区域的渔船作业努力量' → 使用 gfw_query\n" " - '有哪些关于鱿鱼CPUE的研究文献?' → 使用 literature_cpue_query\n" " - '查询中国2015-2020年在南太平洋的渔获量' → 使用 sprfmo_query\n\n" "请根据用户需求选择合适的工具,不要混用。\n\n" "【重要】最终回答格式要求:\n" "在展示查询数据后,必须附加以下信息:\n" "---\n" "**查询执行详情:**\n" "- 使用工具:[本次调用的Tool名称]\n" "- 查询参数:\n" " - 参数1名称: 参数1值\n" " - 参数2名称: 参数2值\n" " - ...(列出所有实际传入的参数)\n" "---\n" "请严格遵守此格式,确保用户清楚了解每次查询的具体执行过程。" ) logger.info("系统提示词已构建,包含 %d 个文件/文件夹", len(file_list)) return system_prompt # =========================================================================== # DeepSeek 思考模式自定义模型类(已移除) # =========================================================================== # DeepSeekThinkingModel 类已被移除,原因如下: # 1. deepseek-chat 和 deepseek-reasoner 将于 2026/07/24 弃用 # 2. 新模型 deepseek-v4-flash/v4-pro 使用标准 OpenAI API,无需特殊适配 # 3. 思考模式不支持 tool_choice,限制了 Agent 的工具调用能力 # 4. 使用标准 OpenAIServerModel 更稳定,兼容性更好 # 现在使用标准 OpenAIServerModel,参见 create_hf_data_agent 函数 # =========================================================================== # Tool 1: 参数提取 —— 从自然语言中解析出数据集查询参数 # =========================================================================== class ParameterExtractionTool(Tool): """从用户的自然语言描述中提取 Hugging Face 数据集查询所需的结构化参数。""" name = "parameter_extractor" description = ( "从用户的自然语言查询请求中提取数据集查询参数,包括:" "region(区域,如中西太平洋)、time_range(时间范围,如1967-2024)、" "time_scale(时间尺度,如月尺度)、spatial_resolution(空间分辨率,如1x1)、" "file_name(文件名)。返回 JSON 格式的参数字典。" ) inputs = { "user_query": { "type": "string", "description": ( "用户的自然语言查询请求,例如'查询中西太平洋1967-2024年月尺度1x1分辨率的渔获数据'" ), } } output_type = "string" # --- 区域关键字映射表 --- REGION_KEYWORDS = { "中西太平洋": "中西太平洋WCPFC", "wcpfc": "中西太平洋WCPFC", "东太平洋": "东太平洋", "太平洋": "中西太平洋WCPFC", "印度洋": "印度洋", "大西洋": "大西洋", "南海": "南海", } # --- 时间尺度映射表 --- TIME_SCALE_KEYWORDS = { "月尺度": "月尺度", "月度": "月尺度", "月": "月尺度", "年尺度": "年尺度", "年度": "年尺度", "年": "年尺度", "日尺度": "日尺度", "日度": "日尺度", "日": "日尺度", } # --- 空间分辨率映射表 --- SPATIAL_RESOLUTION_KEYWORDS = { "1x1": "1x1", "1度": "1x1", "0.5x0.5": "0.5x0.5", "0.5度": "0.5x0.5", "0.1x0.1": "0.1x0.1", "0.1度": "0.1x0.1", } def forward(self, user_query: str) -> str: """ 解析自然语言,提取数据集查询参数。 参数提取流程: 1. 识别目标区域(优先匹配关键字映射表) 2. 识别时间范围(起始年份-结束年份) 3. 识别时间尺度(月尺度/年尺度/日尺度) 4. 识别空间分辨率(1x1/0.5x0.5等) 5. 识别文件名(用户指定的具体文件名) """ logger.info("开始提取参数,输入: %s", user_query) try: params = { "region": self._extract_region(user_query), "time_range": self._extract_time_range(user_query), "time_scale": self._extract_time_scale(user_query), "spatial_resolution": self._extract_spatial_resolution(user_query), "file_name": self._extract_file_name(user_query), } # 校验必要参数:区域不能为空 if not params["region"]: error_msg = ( "无法从输入中识别目标区域,请在查询中明确指定区域。" f"支持的区域关键字: {list(self.REGION_KEYWORDS.keys())}" ) logger.warning(error_msg) return json.dumps({"error": error_msg}, ensure_ascii=False) logger.info("参数提取成功: %s", params) return json.dumps(params, ensure_ascii=False) except Exception as e: logger.error("参数提取失败: %s", e, exc_info=True) return json.dumps( {"error": f"参数提取过程发生错误: {str(e)}"}, ensure_ascii=False, ) # ----- 以下为私有辅助方法,按提取维度拆分 ----- def _extract_region(self, query: str) -> Optional[str]: """ 从自然语言中识别目标区域。 优先级:关键字映射 → 直接匹配区域名称 """ # 优先通过关键字映射识别 for cn_keyword, region_name in self.REGION_KEYWORDS.items(): if cn_keyword.lower() in query.lower(): return region_name return None def _extract_time_range(self, query: str) -> Optional[str]: """ 从自然语言中识别时间范围。 匹配"XXXX-XXXX"格式或"XXXX年到XXXX年"格式。 """ # 匹配"1967-2024"格式 match = re.search(r"(\d{4})-(\d{4})", query) if match: return f"{match.group(1)}-{match.group(2)}" # 匹配"1967年到2024年"或"1967至2024"格式 match = re.search(r"(\d{4})\s*(?:年到|至|-)\s*(\d{4})", query) if match: return f"{match.group(1)}-{match.group(2)}" # 匹配单一年份"2024年" match = re.search(r"(\d{4})\s*年", query) if match: year = match.group(1) return f"{year}-{year}" return None def _extract_time_scale(self, query: str) -> Optional[str]: """ 从自然语言中识别时间尺度。 优先级:关键字映射 → 默认月尺度 """ for cn_keyword, scale_name in self.TIME_SCALE_KEYWORDS.items(): if cn_keyword in query: return scale_name return "月尺度" # 默认月尺度 def _extract_spatial_resolution(self, query: str) -> Optional[str]: """ 从自然语言中识别空间分辨率。 匹配"XxX"格式或"X度"格式。 """ # 优先通过关键字映射识别 for cn_keyword, resolution in self.SPATIAL_RESOLUTION_KEYWORDS.items(): if cn_keyword in query: return resolution # 匹配"1x1"格式 match = re.search(r"(\d+(?:\.\d+)?)x(\d+(?:\.\d+)?)", query) if match: return f"{match.group(1)}x{match.group(2)}" return "1x1" # 默认1x1分辨率 def _extract_file_name(self, query: str) -> Optional[str]: """ 从自然语言中识别具体文件名。 匹配"xxx.zip"或"xxx数据"模式。 """ # 匹配".zip"结尾的文件名 match = re.search(r"(\S+)\.zip", query) if match: return f"{match.group(1)}.zip" # 匹配"xxx文件"或"xxx数据"模式 match = re.search(r"(\w+)\s*(?:文件|数据)", query) if match: return f"{match.group(1)}.zip" return None # =========================================================================== # Tool 2: Hugging Face 数据集获取 —— 从云端获取并解析数据 # =========================================================================== # =========================================================================== # HFDataQueryTool 工具类(已移除) # =========================================================================== # HFDataQueryTool 已被移除,原因: # - 该工具的功能已整合到其他工具中,不再需要单独的数据获取工具 # - 移除了约 350 行代码,包括文件解析、ZIP 处理、多种格式支持等 # - Agent 现在主要使用 logbook_query 和 gfw_query 两个专用查询工具 # =========================================================================== # Tool 3: Logbook 数据可用性查询 # =========================================================================== class LogbookQueryTool(Tool): """查询 logbook 数据可用性,支持按海区、年份、物种、数据类型筛选。""" name = "logbook_query" description = ( "查询 logbook(捕捞日志)数据的可用性。" "支持按海区(region)、年份范围(year_start/year_end)、" "物种(species)、数据类型(data_type)筛选。" "返回统计摘要和预览数据。" ) inputs = { "region": { "type": "string", "description": "海区名称(模糊匹配),如 '东南太平洋'、'中西太平洋'。不指定则不筛选。", "nullable": True, }, "year_start": { "type": "integer", "description": "起始年份(包含),如 2010。不指定则不筛选。", "nullable": True, }, "year_end": { "type": "integer", "description": "结束年份(包含),如 2020。不指定则不筛选。", "nullable": True, }, "species": { "type": "string", "description": "物种名称(模糊匹配),如 '鱿鱼'。不指定则不筛选。", "nullable": True, }, "data_type": { "type": "string", "description": "数据类型(模糊匹配),如 '捕捞日志'。不指定则不筛选。", "nullable": True, }, } output_type = "string" def forward( self, region: Optional[str] = None, year_start: Optional[int] = None, year_end: Optional[int] = None, species: Optional[str] = None, data_type: Optional[str] = None, ) -> str: """ 调用 query_logbook_availability 查询 logbook 数据可用性。 处理流程: 1. 将参数传递给 query_logbook_availability 函数 2. 捕获异常并返回错误信息 3. 将结果转为 JSON 字符串返回 """ logger.info( "Logbook 查询: region=%s, year_start=%s, year_end=%s, species=%s, data_type=%s", region, year_start, year_end, species, data_type, ) try: result = query_logbook_availability( region=region if region else None, year_start=year_start, year_end=year_end, species=species if species else None, data_type=data_type if data_type else None, output_format="markdown", # 默认使用 markdown 输出,不写文件 ) # 格式化 summary 为可读字符串(兼容新旧版本) summary_data = result.get("summary", {}) if isinstance(summary_data, dict): # 新版本:summary 是字典,格式化为友好字符串 summary_str = ( f"找到 {summary_data.get('records_count', 0)} 条 logbook 数据记录\n" f"覆盖年份: {summary_data.get('year_range', ['未知', '未知'])[0]}-{summary_data.get('year_range', ['未知', '未知'])[1]}\n" f"海区: {', '.join(summary_data.get('regions', ['未知']))}\n" f"物种: {', '.join(summary_data.get('species', ['未知']))}" ) else: # 旧版本:summary 已经是字符串 summary_str = summary_data # 确保返回结构完整 complete_result = { "summary": summary_str, "records": result.get("records", []), "preview_markdown": result.get("preview_markdown", ""), "source_files": result.get("source_files", []), } return json.dumps(complete_result, ensure_ascii=False, default=str) except Exception as e: logger.error("Logbook 查询失败: %s", e, exc_info=True) return json.dumps( {"error": f"Logbook 查询失败: {str(e)}"}, ensure_ascii=False, ) # =========================================================================== # Tool 4: GFW 渔船作业努力量查询 # =========================================================================== class GfwQueryTool(Tool): """查询 GFW(Global Fishing Watch)渔船作业努力量数据。""" name = "gfw_query" description = ( "查询 GFW 渔船作业努力量数据。" "需要指定年份(仅支持 2012-2024)、月份和空间范围(经纬度)。" "可选按船旗国(flag)和渔具类型(geartype)筛选。" "注意:lon_min > lon_max 表示查询范围跨越 180° 经线。" ) inputs = { "year": { "type": "integer", "description": "查询年份(仅支持 2012-2024),如 2012。", }, "month": { "type": "integer", "description": "查询月份(1-12)。", }, "lat_min": { "type": "number", "description": "最小纬度,如 -10。", }, "lat_max": { "type": "number", "description": "最大纬度,如 10。", }, "lon_min": { "type": "number", "description": "最小经度,如 145。大于 lon_max 时表示跨越 180° 经线。", }, "lon_max": { "type": "number", "description": "最大经度,如 -175。", }, "flag": { "type": "string", "description": "船旗国代码(精确匹配,不区分大小写),如 'CN'。不指定则不筛选。", "nullable": True, }, "geartype": { "type": "string", "description": "渔具类型(精确匹配,不区分大小写),如 'drifting_longlines'。不指定则不筛选。", "nullable": True, }, "min_fishing_hours": { "type": "number", "description": "最小捕捞小时数阈值,如 10。不指定则不筛选。", "nullable": True, }, } output_type = "string" def forward( self, year: int, month: int, lat_min: float, lat_max: float, lon_min: float, lon_max: float, flag: Optional[str] = None, geartype: Optional[str] = None, min_fishing_hours: Optional[float] = None, ) -> str: """ 调用 query_gfw 查询 GFW 渔船作业努力量数据。 处理流程: 1. 将参数传递给 query_gfw 函数 2. 捕获异常并返回错误信息 3. 将 DataFrame 结果和 summary 转为 JSON 字符串返回 """ logger.info( "GFW 查询: year=%s, month=%s, lat=[%s,%s], lon=[%s,%s], flag=%s, geartype=%s, min_fishing_hours=%s", year, month, lat_min, lat_max, lon_min, lon_max, flag, geartype, min_fishing_hours, ) try: result_df, summary = query_gfw( year=year, month=month, lat_min=lat_min, lat_max=lat_max, lon_min=lon_min, lon_max=lon_max, flag=flag if flag else None, geartype=geartype if geartype else None, min_fishing_hours=min_fishing_hours, ) # 将 DataFrame 转为记录列表(限制预览条数) display_limit = 100 records = result_df.head(display_limit).to_dict(orient="records") output = { "summary": summary, "record_count": len(result_df), "display_records": min(len(result_df), display_limit), "data_preview": records, "columns": list(result_df.columns) if not result_df.empty else [], } return json.dumps(output, ensure_ascii=False, default=str) except FileNotFoundError as e: # GFW 数据文件未找到 logger.error("GFW 数据文件未找到: %s", e) return json.dumps( {"error": f"GFW 数据文件未找到: {str(e)}"}, ensure_ascii=False, ) except ValueError as e: # 参数校验失败(如月份越界、纬度范围错误) logger.error("GFW 查询参数错误: %s", e) return json.dumps( {"error": f"查询参数错误: {str(e)}"}, ensure_ascii=False, ) except Exception as e: logger.error("GFW 查询失败: %s", e, exc_info=True) return json.dumps( {"error": f"GFW 查询失败: {str(e)}"}, ensure_ascii=False, ) # =========================================================================== # Tool 5: 文献 CPUE 查询 # =========================================================================== class LiteratureCpueQueryTool(Tool): """查询 CPUE(单位捕捞努力量渔获量)相关文献数据。""" name = "literature_cpue_query" description = ( "查询 CPUE(单位捕捞努力量渔获量)相关文献数据。" "支持按海区、年份范围、物种、响应变量、论文类型筛选。" "返回文献统计摘要和预览数据。" ) inputs = { "region": { "type": "string", "description": "海区名称(模糊匹配),如 '东南太平洋'、'中西太平洋'。不指定则不筛选。", "nullable": True, }, "year_start": { "type": "integer", "description": "起始年份(包含),如 2010。不指定则不筛选。", "nullable": True, }, "year_end": { "type": "integer", "description": "结束年份(包含),如 2020。不指定则不筛选。", "nullable": True, }, "species": { "type": "string", "description": "物种名称(模糊匹配),如 '鱿鱼'。不指定则不筛选。", "nullable": True, }, "response_variable": { "type": "string", "description": "响应变量(模糊匹配),如 'CPUE'、'biomass'。不指定则不筛选。", "nullable": True, }, "paper_type": { "type": "string", "description": "论文类型(模糊匹配)。不指定则不筛选。", "nullable": True, }, } output_type = "string" def forward( self, region: Optional[str] = None, year_start: Optional[int] = None, year_end: Optional[int] = None, species: Optional[str] = None, response_variable: Optional[str] = None, paper_type: Optional[str] = None, ) -> str: """ 调用 query_literature_cpue 查询 CPUE 文献数据。 处理流程: 1. 将参数传递给 query_literature_cpue 函数 2. 捕获异常并返回错误信息 3. 将结果转为 JSON 字符串返回 """ logger.info( "CPUE 文献查询: region=%s, year_start=%s, year_end=%s, species=%s, response_variable=%s, paper_type=%s", region, year_start, year_end, species, response_variable, paper_type, ) try: # 动态导入以避免循环依赖 from query_tools.literature_cpue_query import query_literature_cpue result = query_literature_cpue( region=region if region else None, year_start=year_start, year_end=year_end, species=species if species else None, response_variable=response_variable if response_variable else None, paper_type=paper_type if paper_type else None, output_format="markdown", # 默认使用 markdown 输出,不写文件 ) # 格式化 summary 为可读字符串 summary_data = result.get("summary", {}) if isinstance(summary_data, dict): summary_str = ( f"找到 {summary_data.get('paper_count', 0)} 篇 CPUE 相关文献\n" f"覆盖年份: {summary_data.get('year_range', ['未知', '未知'])[0]}-{summary_data.get('year_range', ['未知', '未知'])[1]}\n" f"涉及海区: {', '.join(summary_data.get('regions', ['未知']))}\n" f"涉及物种: {', '.join(summary_data.get('species', ['未知']))}\n" f"使用模型: {', '.join(summary_data.get('models', ['未知']))}" ) else: summary_str = summary_data # 确保返回结构完整 complete_result = { "summary": summary_str, "records": result.get("records", []), "preview_markdown": result.get("preview_markdown", ""), "source_files": result.get("source_files", []), } return json.dumps(complete_result, ensure_ascii=False, default=str) except Exception as e: logger.error("CPUE 文献查询失败: %s", e, exc_info=True) return json.dumps( {"error": f"CPUE 文献查询失败: {str(e)}"}, ensure_ascii=False, ) # =========================================================================== # Tool 6: SPRFMO 南太平洋数据查询 # =========================================================================== class SprfmoQueryTool(Tool): """查询 SPRFMO(南太平洋区域渔业管理组织)数据。""" name = "sprfmo_query" description = ( "查询 SPRFMO(南太平洋区域渔业管理组织)数据。" "支持查询捕捞量(catch)或努力量(effort)数据。" "可按国家、年份范围、物种等筛选,支持按维度聚合统计。" "空间分辨率:5x5 度;时间分辨率:年度。" ) inputs = { "data_type": { "type": "string", "description": "数据类型,'catch'(捕捞量)或 'effort'(努力量)。必填。", }, "country": { "type": "string", "description": "国家代码或名称(如 'CHN'、'中国'、'JPN')。不指定则不筛选。", "nullable": True, }, "year_start": { "type": "integer", "description": "起始年份(包含)。不指定则不筛选。", "nullable": True, }, "year_end": { "type": "integer", "description": "结束年份(包含)。不指定则不筛选。", "nullable": True, }, "species": { "type": "string", "description": "物种名称(仅捕捞量数据)。不指定则不筛选。", "nullable": True, }, "gear_type": { "type": "string", "description": "渔具类型(仅努力量数据)。不指定则不筛选。", "nullable": True, }, "group_by": { "type": "string", "description": "聚合维度,如 'year'、'country'、'species'。多个维度用逗号分隔。不指定则不聚合。", "nullable": True, }, } output_type = "string" def forward( self, data_type: str, country: Optional[str] = None, year_start: Optional[int] = None, year_end: Optional[int] = None, species: Optional[str] = None, gear_type: Optional[str] = None, group_by: Optional[str] = None, ) -> str: """ 调用 query_sprfmo 查询 SPRFMO 数据。 处理流程: 1. 构建 filters 字典 2. 解析 group_by 参数 3. 调用 query_sprfmo 函数 4. 捕获异常并返回错误信息 5. 将结果转为 JSON 字符串返回 """ logger.info( "SPRFMO 查询: data_type=%s, country=%s, year=[%s,%s], species=%s, gear_type=%s, group_by=%s", data_type, country, year_start, year_end, species, gear_type, group_by, ) try: # 动态导入以避免循环依赖 from query_tools.query_sprfmo import query_sprfmo # 构建 filters 字典 filters = {} if country: filters["country"] = country if year_start is not None: filters["year_start"] = year_start if year_end is not None: filters["year_end"] = year_end if species and data_type == "catch": filters["species"] = species if gear_type and data_type == "effort": filters["gear_type"] = gear_type # 解析 group_by 参数 group_by_list = None if group_by: group_by_list = [g.strip() for g in group_by.split(",")] # 调用查询函数 result = query_sprfmo( filters=filters, group_by=group_by_list, metrics=data_type, # metrics 参数直接使用 data_type data_type=data_type, output_format="markdown", ) # 格式化 summary summary_data = result.get("summary", {}) summary_str = ( f"数据类型: {summary_data.get('data_type', '未知')}\n" f"总记录数: {summary_data.get('total_records', 0)}\n" f"数据来源: {summary_data.get('data_source', 'SPRFMO')}\n" ) if "year_range" in summary_data: summary_str += f"年份范围: {summary_data['year_range']}\n" if "total_catch" in summary_data: summary_str += f"总捕捞量: {summary_data['total_catch']:.2f} kg\n" if "total_effort" in summary_data: summary_str += f"总努力量: {summary_data['total_effort']:.2f} 天\n" # 确保返回结构完整 complete_result = { "summary": summary_str, "records": result.get("records", []), "preview_markdown": result.get("preview_markdown", ""), "source_files": result.get("source_files", []), "warnings": result.get("warnings", []), "metadata": result.get("metadata", {}), } return json.dumps(complete_result, ensure_ascii=False, default=str) except FileNotFoundError as e: # SPRFMO 数据文件未找到 logger.error("SPRFMO 数据文件未找到: %s", e) return json.dumps( {"error": f"SPRFMO 数据文件未找到: {str(e)}"}, ensure_ascii=False, ) except Exception as e: logger.error("SPRFMO 查询失败: %s", e, exc_info=True) return json.dumps( {"error": f"SPRFMO 查询失败: {str(e)}"}, ensure_ascii=False, ) # =========================================================================== # Agent 工厂函数:创建配置好的 ToolCallingAgent # =========================================================================== def create_hf_data_agent( model_id: str = MODEL_ID, api_key: Optional[str] = None, api_base: Optional[str] = None, max_steps: int = 5, ) -> ToolCallingAgent: """ 创建并返回一个配置好的 Hugging Face 数据集查询智能代理。 在 Agent 启动时会自动获取数据集文件列表并作为系统提示词告知用户。 注意:DeepSeek 思考模式已停用,现在使用标准 OpenAIServerModel。 原因:deepseek-chat/deepseek-reasoner 将于 2026/07/24 弃用, 新模型 deepseek-v4-flash/v4-pro 使用标准 OpenAI API。 Args: model_id: 模型 ID,默认使用文件顶部 MODEL_ID 常量 推荐: deepseek-v4-flash 或 deepseek-v4-pro api_key: API Key,默认使用文件顶部 OPENAI_API_KEY 常量 api_base: API 地址,默认使用文件顶部 OPENAI_API_BASE 常量 max_steps: Agent 最大推理步骤数 Returns: 配置好的 ToolCallingAgent 实例 """ key = api_key or OPENAI_API_KEY base = api_base or OPENAI_API_BASE # 初始化 LLM 模型(使用标准 OpenAIServerModel,已停用思考模式) model = OpenAIServerModel( model_id=model_id, api_key=key, api_base=base, ) # 实例化工具(包含新增的 CPUE 文献和 SPRFMO 查询工具) param_tool = ParameterExtractionTool() logbook_tool = LogbookQueryTool() gfw_tool = GfwQueryTool() literature_cpue_tool = LiteratureCpueQueryTool() sprfmo_tool = SprfmoQueryTool() # 构建包含数据集文件列表的指令(会被插入到系统提示词中) instructions = build_system_prompt_with_file_list() logger.info("Agent 自定义指令: %s", instructions[:200] + "...") # 只打印前200字符 # 创建 ToolCallingAgent,使用 instructions 参数设置自定义指令 # 工具列表包含5个查询工具(已移除 hf_data_tool) agent = ToolCallingAgent( tools=[param_tool, logbook_tool, gfw_tool, literature_cpue_tool, sprfmo_tool], model=model, max_steps=max_steps, instructions=instructions, ) logger.info("Hugging Face 数据集查询代理创建成功,模型: %s", model_id) return agent # =========================================================================== # 主入口:GradioUI 交互式 Web 界面 # =========================================================================== def main(): """ 使用 GradioUI 启动交互式 Web 界面。 用户在浏览器中输入自然语言查询,代理实时展示思考过程和数据获取结果。 界面基于 gr.ChatInterface,支持流式输出和步骤可视化。 """ agent = create_hf_data_agent() # GradioUI 将 agent 包装为 gr.ChatInterface Web 应用 # - share=True: 生成公网可访问的临时链接(72小时有效) # - reset_agent_memory=False: 保留对话上下文,允许多轮交互 demo = GradioUI( agent, reset_agent_memory=False, ) logger.info("正在启动 GradioUI Web 界面...") demo.launch(share=True) if __name__ == "__main__": main()