File size: 18,815 Bytes
d82bbe4
 
 
 
 
 
 
8de7969
 
 
 
d82bbe4
8de7969
 
 
 
 
d82bbe4
8de7969
 
 
 
 
 
 
d82bbe4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
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<n<10K'
    dataset_num_limit: int = 5
    category: str = "PT"
    max_dataset_size: int = None  # 数据集大小限制(字节数),None表示不限制
    max_download_subtasks: Optional[int] = None  # 下载子任务执行数量上限,None 表示不限制
    rag_api_url: Optional[str] = None
    rag_api_key: Optional[str] = None
    rag_embed_model: Optional[str] = None
    tavily_api_key: Optional[str] = None


# ==================== 数据采集 State ====================
@dataclass
class DataCollectionState(MainState):
    """数据采集任务的State,继承自MainState"""
    # 重写request类型为DataCollectionRequest
    request: DataCollectionRequest = field(default_factory=DataCollectionRequest)
    
    # 数据采集特有的字段
    keywords: list[str] = field(default_factory=list)
    datasets: Dict[str, list] = field(default_factory=dict)
    downloads: Dict[str, list] = field(default_factory=dict)
    sources: Dict[str, Dict] = field(default_factory=dict)

# Iconagent相关 State 和 Request 定义
# ==================== Icon 生成 Request ====================
@dataclass
class IconGenRequest(MainRequest):      
    keywords: str = ""
    style: str = ""
    prev_image: str = ""
    edit_prompt: str = ""

# ==================== Icon 生成 State ======================
@dataclass
class IconGenState(MainState):
    request: IconGenRequest = field(default_factory=IconGenRequest)

    # 下面是 icongen 自己的产物 / 临时数据
    icon_prompt: str = ""                                 # 生成的图标提示词
    img_save_path: str = ""                              # 生成的图标保存路径


# ==================== Web 爬取/研究 Request ====================
@dataclass
class WebCrawlRequest(MainRequest):
    """Web 爬取任务的 Request,继承自 MainRequest"""
    # 初始需求与下载目录
    initial_request: str = ""
    download_dir: str = os.path.join(STATICS_DIR, "web_crawl")

    # 爬取/研究配置
    search_engine: str = "tavily"     # 'tavily' | 'duckduckgo' | 'jina'
    use_jina_reader: bool = False
    enable_rag: bool = True
    max_download_subtasks: Optional[int] = None


# ==================== Web 爬取/研究 State ====================
@dataclass
class WebCrawlState(MainState):
    """管理网络爬取与研究过程的状态"""
    # 重写 request 类型为 WebCrawlRequest
    request: WebCrawlRequest = field(default_factory=WebCrawlRequest)

    # 直通字段(为兼容调用方直接从 state 访问这些配置项)
    initial_request: str = ""
    download_dir: str = os.path.join(STATICS_DIR, "web_crawl")
    search_engine: str = "tavily"
    use_jina_reader: bool = False
    enable_rag: bool = True
    rag_manager: Any = None
    max_download_subtasks: Optional[int] = None

    # 研究/爬取过程中的临时与产出数据
    sub_tasks: list[Dict[str, Any]] = field(default_factory=list)
    completed_sub_tasks: list[Dict[str, Any]] = field(default_factory=list)
    research_summary: Dict[str, Any] = field(default_factory=dict)
    search_results_text: str = ""
    filtered_urls: list[str] = field(default_factory=list)
    crawled_data: list[Dict[str, Any]] = field(default_factory=list)
    visited_urls: set[str] = field(default_factory=set)
    url_queue: list[str] = field(default_factory=list)
    is_finished: bool = False
    supervisor_feedback: str = "Process has not started."
    # 控制参数
    max_crawl_cycles_per_task: int = 5
    max_crawl_cycles_for_research: int = 15
    max_dataset_size: Optional[int] = None
    current_cycle: int = 0
    download_successful_for_current_task: bool = False
    completed_download_tasks: int = 0

    def reset_for_new_task(self):
        self.search_results_text = ""
        self.filtered_urls = []
        self.visited_urls = set()
        self.url_queue = []
        self.current_cycle = 0
        self.download_successful_for_current_task = False
        

    
@dataclass
class PromptWritingState(MainState):
    """提示词生成任务的State,继承自MainState"""
    request: DFRequest = field(default_factory=DFRequest)
    
    # 提示词生成特有的字段
    prompt_op_name: str = ""
    prompt_args: Dict[str, Any] = field(default_factory=dict)
    prompt_output_format: Dict[str, Any] = field(default_factory=dict)
    delete_test_files: bool = True

# Paper2Video 相关 State 和 Request 定义
# ==================== Paper2Video 生成 Request ====================

@dataclass
class Paper2VideoRequest(MainRequest):
    paper_pdf_path: str = ""
    user_imgs_path: str = ""
    
    ref_audio_path: str = ""

# ==================== Paper2Video 生成 State ======================
@dataclass
class Paper2VideoState(MainState):
    # 重写 request
    request: Paper2VideoRequest = field(default_factory=Paper2VideoRequest)
    
    # paper2video 特有字段
    beamer_code_path: str = ""
    is_beamer_wrong: bool = False
    is_beamer_warning: bool = False
    code_debug_result: str = ""
    ppt_path: str = ""
    
    # 生成字幕 + cursor的位置信息
    slide_img_dir: str = ""
    subtitle_and_cursor: List[str] = field(default_factory=list)
    subtitle_and_cursor_path: str = ""
    
    # 生成的音频路径
    speech_save_dir: str = ""



# ==================== Planning Agent 相关 State ====================
@dataclass
class PlanningRequest(MainRequest):
    """Planning Agent 的 Request"""
    # 规划器配置
    planner_model: Optional[str] = None
    planner_temperature: float = 0.0
    
    # 执行器配置
    executor_model: Optional[str] = None
    executor_as_react: bool = True
    
    # 重规划器配置 (仅 Plan-and-Execute 模式)
    replanner_model: Optional[str] = None
    max_replanning_rounds: int = 3
    
    # Human-in-the-Loop 配置
    require_plan_approval: bool = True      # 是否需要计划审批
    interrupt_before_step: bool = True      # 每步执行前是否中断
    interrupt_after_step: bool = False      # 每步执行后是否中断
    
    # 执行配置
    max_plan_steps: int = 10
    planning_mode: str = "plan_solve"       # "plan_solve" | "plan_execute"


@dataclass
class PlanStep:
    """单个计划步骤"""
    index: int                              # 步骤索引
    description: str                        # 步骤描述
    status: str = "pending"                 # pending | running | completed | failed | skipped
    result: Optional[str] = None            # 执行结果
    error: Optional[str] = None             # 错误信息
    started_at: Optional[str] = None        # 开始时间
    completed_at: Optional[str] = None      # 完成时间


@dataclass
class PlanningState(MainState):
    """
    Planning Agent 的状态类
    
    支持两种模式:
    - Plan-and-Solve: 一次性生成计划,按顺序执行
    - Plan-and-Execute (Replanning): 动态调整计划
    """
    request: PlanningRequest = field(default_factory=PlanningRequest)
    
    # ===== 计划相关 =====
    plan: List[str] = field(default_factory=list)                    # 计划步骤列表 (简单字符串)
    plan_steps: List[Dict[str, Any]] = field(default_factory=list)   # 详细计划步骤
    current_step_index: int = 0                                       # 当前执行步骤索引
    past_steps: List[tuple] = field(default_factory=list)            # [(步骤描述, 执行结果), ...]
    
    # ===== 状态控制 =====
    plan_approved: bool = False                     # 计划是否已审批
    is_replanning_needed: bool = False              # 是否需要重新规划
    replanning_count: int = 0                       # 重规划次数
    final_response: str = ""                        # 最终响应
    is_finished: bool = False                       # 是否已完成
    
    # ===== Human-in-the-Loop 相关 =====
    awaiting_human_input: bool = False              # 是否等待人类输入
    human_feedback: Optional[str] = None            # 人类反馈
    interrupt_reason: Optional[str] = None          # 中断原因
    
    # ===== 执行上下文 =====
    original_task: str = ""                         # 原始任务描述
    executor_tools: List[str] = field(default_factory=list)  # 可用工具列表
    
    def get_current_step(self) -> 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 # 是否使用单页聚类器