from dataclasses import dataclass, field import os from pathlib import Path from typing import Any, Dict, List, Optional current_file = Path(__file__).resolve() PROJDIR = current_file.parent.parent BASE_DIR = PROJDIR DATAFLOW_DIR = PROJDIR STATICS_DIR = PROJDIR / "static" # `open-dataflow` may pull in heavy deps (e.g. torch/MPI) via its CLI helpers. # In some environments, importing it can hard-crash the interpreter (SIGABRT), # which cannot be handled by try/except. Keep it opt-in. if os.getenv("DF_USE_OPEN_DATAFLOW_PATHS", "").strip().lower() in {"1", "true", "yes"}: # pragma: no cover try: from dataflow.cli_funcs.paths import DataFlowPath # type: ignore BASE_DIR = DataFlowPath.get_dataflow_dir() DATAFLOW_DIR = BASE_DIR.parent STATICS_DIR = DataFlowPath.get_dataflow_statics_dir() except Exception: BASE_DIR = PROJDIR DATAFLOW_DIR = PROJDIR STATICS_DIR = PROJDIR / "static" from typing_extensions import TypedDict, Annotated from langgraph.graph.message import add_messages from langchain_core.messages import BaseMessage # ==================== 最基础的 Request ==================== @dataclass class MainRequest: """所有Request的基类,只包含核心字段""" # ① 用户偏好的自然语言 language: str = "en" # "en" | "zh" | ... # ② LLM 接口 chat_api_url: str = os.getenv("DF_API_URL", "test") api_key: str = os.getenv("DF_API_KEY", "test") chat_api_key: str = os.getenv("DF_API_KEY", "test") #没区别,但是不想改之前代码了; # ③ 选用的 LLM 名称 model: str = "gpt-4o" # ④ 需求描述 target: str = "" def get(self, key, default=None): return getattr(self, key, default) def __setitem__(self, key, value): setattr(self, key, value) # ==================== 最基础的 State(所有State的祖先)==================== @dataclass class MainState: """所有State的基类,只包含核心字段""" request: MainRequest = field(default_factory=MainRequest) messages: Annotated[list[BaseMessage], add_messages] = field(default_factory=list) # 通用字段 agent_results: Dict[str, Any] = field(default_factory=dict) temp_data: Dict[str, Any] = field(default_factory=dict) def get(self, key, default=None): return getattr(self, key, default) def __setitem__(self, key, value): setattr(self, key, value) # ==================== 主流程 Request ==================== @dataclass class DFRequest(MainRequest): """主流程的Request,继承自MainRequest""" # ⑤ 测试样例文件(仅 CLI 批量跑用) json_file: str = "" # ⑥ Python 代码文件位置 python_file_path: str = "" # ⑦ Debug 相关 need_debug: bool = False max_debug_rounds: int = 3 # ⑧ 本地模型相关 use_local_model: bool = False local_model_path: str = "" # ⑨ 缓存和会话 cache_dir: str = f"{PROJDIR}/cache_dir" session_id: str = "default_session" # embeddings url chat_api_url_for_embeddings : str = "" embedding_model_name: str = "text-embedding-3-small" update_rag_content: bool = True # ==================== 主流程 State ==================== @dataclass class DFState(MainState): """主流程的State,继承自MainState""" # 重写request类型为DFRequest request: DFRequest = field(default_factory=DFRequest) # 主流程特有字段 category: Dict[str, Any] = field(default_factory=dict) recommendation: Dict[str, Any] = field(default_factory=dict) matched_ops: list[str] = field(default_factory=list) debug_mode: bool = False pipeline_structure_code: Dict[str, Any] = field(default_factory=dict) execution_result: Dict[str, Any] = field(default_factory=dict) code_debug_result: Dict[str, Any] = field(default_factory=dict) debug_history: Dict[Any, Dict[str, Any]] = field(default_factory=dict) opname_and_params: List[Dict[str, Dict[str, Any]]] = field(default_factory=list) # ==================== 数据采集 Request ==================== @dataclass class DataCollectionRequest(MainRequest): """数据采集任务的Request,继承自MainRequest""" # 重写language默认值 language: str = "English" # 数据采集特有的字段 download_dir: str = os.path.join(STATICS_DIR, "data_collection") dataset_size_category: str = '1K Optional[str]: """获取当前待执行的步骤""" if 0 <= self.current_step_index < len(self.plan): return self.plan[self.current_step_index] return None def get_remaining_steps(self) -> List[str]: """获取剩余未执行的步骤""" return self.plan[self.current_step_index:] def get_completed_steps(self) -> List[tuple]: """获取已完成的步骤及结果""" return self.past_steps def mark_step_complete(self, result: str): """标记当前步骤完成""" if self.current_step_index < len(self.plan): step = self.plan[self.current_step_index] self.past_steps.append((step, result)) self.current_step_index += 1 def reset_plan(self): """重置计划状态(用于重规划)""" self.plan = [] self.plan_steps = [] self.current_step_index = 0 self.is_replanning_needed = False # 保留 past_steps,因为重规划需要参考历史执行结果 def to_planning_context(self) -> Dict[str, Any]: """生成规划上下文(供 LLM 使用)""" return { "original_task": self.original_task or self.request.target, "past_steps": [ {"step": step, "result": result} for step, result in self.past_steps ], "remaining_steps": self.get_remaining_steps(), "replanning_count": self.replanning_count, "available_tools": self.executor_tools, } @dataclass class Paper2FigureRequest(MainRequest): gen_fig_model: str = "gemini-2.5-flash-image-preview" # gen_fig_model: str = "gemini-3-pro-image-preview" sam2_model: str = "models/facebook/sam2.1-hiera-tiny" bg_rm_model: str = "models/RMBG-2.0" input_type: str = "PDF" # 科研绘图复杂度 figure_complex: str = "hard" style: str = "kartoon" # PPT的页面数量 page_count: int = 10 # 是否编辑完毕,也就是是否需要重新生成完整的 PPT all_edited_down: bool = False # pdf2ppt是否使用AI编辑 use_ai_edit: bool = False @dataclass class Paper2FigureState(MainState): request: Paper2FigureRequest = field(default_factory=Paper2FigureRequest) fig_desc: str = '' aspect_ratio: str = '16:9' paper_file: str = '' # 原始带内容的图像路径 fig_draft_path: str = '' # MinerU 解析得到的内容元素(文本 / 图片 / 表格等) fig_mask: List[Dict[str, Any]] = field(default_factory=list) # 二次编辑后的空框模板图(仅外层矩形和箭头) fig_layout_path: str = '' # SAM + SVG + EMF 形成的布局元素(仅背景框架层) layout_items: List[Dict[str, Any]] = field(default_factory=list) result_path: str = '' ppt_path: str = '' mask_detail_level: int = 2 paper_idea: str = '' input_type: str = 'PDF' # 技术路线图使用属性 ============================== figure_tec_svg_content: str = "" svg_img_path: str = "" mineru_port: int = 8010 svg_file_path: str = "" # svg 带文字图的 地址 svg_bg_file_path: str = "" # 带文字版本的svg图片 svg_full_img_path: str = "" # 背景svg code: svg_bg_code : str = "" # 实验统计图使用属性 ============================== # ===== 输入 ===== pre_tool_results: Dict[str, Any] = field(default_factory=dict) # 前置工具结果注入 # ===== 中间结果 ===== paper_idea: str = '' # 论文核心思想 extracted_tables: List[Dict[str, Any]] = field(default_factory=list) # 从 MinerU 提取的表格列表 # 每个表格格式: {"table_id": str, "headers": List[str], "rows": List[List[str]], "caption": str} chart_configs: Dict[str, Dict[str, Any]] = field(default_factory=dict) # 图表配置字典 # 每个配置格式: {table_id: {"table_id": str, "chart_type": str, "x_column": str, "y_columns": List[str], ...}} generated_codes: Dict[str, Dict[str, Any]] = field(default_factory=dict) # 生成的代码字典 # 每个代码格式: {table_id: {"table_id": str, "code": str}} # ===== 输出 ===== generated_charts: Dict[str, str] = field(default_factory=dict) # 生成的图表路径字典 stylize_results: Dict[str, list] = field(default_factory=dict) # 风格化后的图表路径字典 svg_bg_code: str = "" # paper2ppt 专用 ============================== # 首次生成整套页面图是否已完成;False 走批量生成,True 走按页二次编辑 gen_down: bool = False # 0-based: 要二次编辑的页号 edit_page_num: int = -1 # 二次编辑提示词(用于 edit_page_num 对应页) edit_page_prompt: str = "" # 批量生成出来的页面图片路径(0-based 对齐 pagecontent) generated_pages: List[str] = field(default_factory=list) table_img_path: str = "" # pagecontent: 既可为结构化 slide 描述,也可为 [{"ppt_img_path": "..."}] 的图片列表 pagecontent: list[dict] = field(default_factory=list) minueru_output: str = "" mineru_root: str = "" text_content: str = "" # 生成的 PPT PDF 路径 ppt_pdf_path: str = "" ppt_pptx_path: str = "" # 长文PPT专用: long_text: str = "" target_pages: int = 60 pages_per_batch: int = 10 pages_to_generate: int = 12 max_rounds: int = 1 current_chunk: str = "" current_text: str = "" # pdf2ppt 专用 ============================== pdf_file: str = "" slide_images: List[str] = field(default_factory=list) ocr_pages: List[str] = field(default_factory=list) sam_pages: List[str] = field(default_factory=list) mineru_pages: List[Dict[str, Any]] = field(default_factory=list) # pdf2ppt是否使用AI编辑 use_ai_edit: bool = False use_global_font_clustering: bool = False # 是否使用单页聚类器