diff --git a/deepsearch/docs/feature/algorithm/report-generation/visualization-markdown.md b/deepsearch/docs/feature/algorithm/report-generation/visualization-markdown.md index dd8f3b44..151f229e 100644 --- a/deepsearch/docs/feature/algorithm/report-generation/visualization-markdown.md +++ b/deepsearch/docs/feature/algorithm/report-generation/visualization-markdown.md @@ -21,17 +21,21 @@ Markdown 可视化用于在报告正文中以 Mermaid 文本图表表达结构 - 同一章节可以插入多张 Mermaid 图表;多图来源于章节内多个高数据密度候选资料,而不是正文生成后的二次补图。 - 插入到报告正文的 Mermaid 图表会带有系统管理的居中图题,并在图题中保留对应 citation。 - 章节正文写作 Prompt 不允许模型直接输出 Mermaid 代码围栏、图表代码或手写图块;该约束不依赖 `visualization_enable`,因此 VLM 图表开启、Mermaid 可视化关闭时也不会允许正文草稿混入未受控 Mermaid。 +- VLM 图表开启时,如果正文草稿仍意外包含未受控 Mermaid 代码块,VLM 节点会在进入溯源前清理这些代码块;该兜底仅作用于 VLM 链路,不影响本文档描述的受控 Markdown Mermaid 可视化图表。 - 若某个候选资料抽取、归一化、合规校验或 Mermaid 生成失败,该候选会被跳过;系统不会使用本地正则从正文中硬抽图表数据。 +- 在进入多轮 LLM 生图链路前,系统会对候选资料做轻量预算控制:优先尝试数值结构清晰、来源可追溯、数据密度高且与其他候选不重复的资料;简短报告默认减少尝试数,用户明确要求图文并茂或多图时会适度放宽。 +- 当章节只有一张有效 Mermaid 图表时,插入位置规划使用本地规则完成,优先插到匹配 citation 的正文行后;只有多图场景才调用 LLM 规划多张图的相对位置。 ## 性能边界 Markdown 可视化会触发多轮 LLM 调用,因此当前实现只保留正文生成前的主链路: 1. 从章节的 `classified_content` 中选择数据密度较高的资料。 -2. 对每个候选资料执行图表数据抽取、校验、单位归一化和 Mermaid 生成。 -3. 子报告正文生成完成后,只执行插入位置规划和 Mermaid 片段渲染。 +2. 对候选资料做轻量排序和预算控制,排序依据是数据密度、原始资料中的数值数量、来源/citation 可追溯性,以及候选之间的重复度。 +3. 对预算内候选资料执行图表数据抽取、校验、单位归一化和 Mermaid 生成。 +4. 子报告正文生成完成后,只执行插入位置规划和 Mermaid 片段渲染;单图场景本地确定插入位置,多图场景再请求 LLM 规划。 -当前实现不在正文写完后再次扫描草稿正文、生成候选、重跑图表抽取或执行重复数据去重预算控制。 +当前实现不在正文写完后再次扫描草稿正文、生成候选或重跑图表抽取。预算控制发生在正文生成前的主可视化链路中,用于减少低价值候选进入昂贵 LLM 步骤。 ## 关键代码路径 @@ -56,13 +60,14 @@ Markdown 可视化会触发多轮 LLM 调用,因此当前实现只保留正文 ## 核心流程 1. 报告生成阶段根据 `classified_content` 的数据密度选择适合可视化的章节资料。 -2. 根据章节标题和章节大纲推断期望图型;该结果只作为软约束,不能覆盖真实数据形态。 -3. LLM 从候选原始资料中抽取图表标题、类型、records 和单位。 -4. 抽取结果通过 schema 校验;混合单位、空 records、字段缺失等结果会被拒绝。 -5. 对需要数值单位的图表执行单位归一化。 -6. 根据图表类型生成 Mermaid 片段。 -7. 合规校验确认 Mermaid 语法、图表类型、数据一致性、可读性和引用上下文满足要求。 -8. 子报告正文生成完成后,系统请求插入位置规划,将已生成的 Mermaid 片段插入正文,并在图题中保留 citation。 +2. 系统在本地对候选做去重、排序和预算控制,避免低价值或重复候选进入多轮 LLM 生图链路。 +3. 根据章节标题和章节大纲推断期望图型;该结果只作为软约束,不能覆盖真实数据形态。 +4. LLM 从预算内候选原始资料中抽取图表标题、类型、records 和单位。 +5. 抽取结果通过 schema 校验;混合单位、空 records、字段缺失等结果会被拒绝。 +6. 对需要数值单位的图表执行单位归一化。 +7. 根据图表类型生成 Mermaid 片段。 +8. 合规校验确认 Mermaid 语法、图表类型、数据一致性、可读性和引用上下文满足要求。 +9. 子报告正文生成完成后,系统将已生成的 Mermaid 片段插入正文,并在图题中保留 citation。单图章节优先使用本地 citation 锚点插入;多图章节请求 LLM 规划多张图的位置。 ## 数据契约与依赖 diff --git a/deepsearch/openjiuwen_deepsearch/algorithm/chart_generation/utils.py b/deepsearch/openjiuwen_deepsearch/algorithm/chart_generation/utils.py index bb1ec49a..564a9e43 100644 --- a/deepsearch/openjiuwen_deepsearch/algorithm/chart_generation/utils.py +++ b/deepsearch/openjiuwen_deepsearch/algorithm/chart_generation/utils.py @@ -2,6 +2,7 @@ # Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. import logging import json +import re from typing import List, Dict, NamedTuple, Optional import base64 @@ -17,6 +18,106 @@ logger = logging.getLogger(__name__) MAX_LLM_RETRY_TIMES = 3 +MERMAID_START_PATTERNS = ( + re.compile(r"^graph\s+(td|tb|bt|rl|lr)\b"), + re.compile(r"^flowchart\s+(td|tb|bt|rl|lr)\b"), + re.compile(r"^sequencediagram\b"), + re.compile(r"^classdiagram(?:-v2)?\b"), + re.compile(r"^statediagram(?:-v2)?\b"), + re.compile(r"^erdiagram\b"), + re.compile(r"^journey\b"), + re.compile(r"^gantt\b"), + re.compile(r"^pie\b"), + re.compile(r"^timeline\b"), + re.compile(r"^mindmap\b"), + re.compile(r"^quadrantchart\b"), + re.compile(r"^xychart-beta\b"), + re.compile(r"^sankey-beta\b"), +) + + +def _looks_like_mermaid_body(body_lines: List[str]) -> bool: + content_lines = [line.strip().lower() for line in body_lines if line.strip()] + if content_lines and content_lines[0] == "---": + try: + end_index = content_lines[1:].index("---") + 1 + content_lines = content_lines[end_index + 1:] + except ValueError: + return False + if not content_lines: + return False + first_line = content_lines[0] + return any(pattern.match(first_line) for pattern in MERMAID_START_PATTERNS) + + +def remove_mermaid_code_blocks(markdown_text: str) -> str: + """Remove fenced Mermaid source blocks that should not appear as report body text.""" + if not markdown_text: + return markdown_text + + lines = markdown_text.splitlines(keepends=True) + output_lines = [] + index = 0 + removed_block = False + + while index < len(lines): + stripped = lines[index].strip() + + if lines[index].startswith((" ", "\t")): + body_start = index + body_end = body_start + body_lines = [] + while body_end < len(lines): + current_line = lines[body_end] + if not current_line.strip(): + body_lines.append("") + body_end += 1 + elif current_line.startswith(" "): + body_lines.append(current_line[4:]) + body_end += 1 + elif current_line.startswith("\t"): + body_lines.append(current_line[1:]) + body_end += 1 + else: + break + + if _looks_like_mermaid_body(body_lines): + removed_block = True + index = body_end + continue + + output_lines.extend(lines[body_start:body_end]) + index = body_end + continue + + fence_match = re.match(r"^(```+|~~~+)\s*([^`]*)$", stripped) + if not fence_match: + output_lines.append(lines[index]) + index += 1 + continue + + fence = fence_match.group(1) + language = fence_match.group(2).strip().lower() + body_start = index + 1 + body_end = body_start + while body_end < len(lines) and not lines[body_end].strip().startswith(fence): + body_end += 1 + + body_lines = lines[body_start:body_end] + is_mermaid = "mermaid" in language or _looks_like_mermaid_body(body_lines) + if is_mermaid: + removed_block = True + index = body_end + 1 if body_end < len(lines) else body_end + continue + + output_lines.extend(lines[index:body_end + 1]) + index = body_end + 1 + + if not removed_block: + return markdown_text + + cleaned_text = "".join(output_lines) + return re.sub(r"\n{3,}", "\n\n", cleaned_text).strip() def type_check(result, expected_type): diff --git a/deepsearch/openjiuwen_deepsearch/algorithm/report/report.py b/deepsearch/openjiuwen_deepsearch/algorithm/report/report.py index 60e2ce8c..2777d5e6 100644 --- a/deepsearch/openjiuwen_deepsearch/algorithm/report/report.py +++ b/deepsearch/openjiuwen_deepsearch/algorithm/report/report.py @@ -10,6 +10,7 @@ import re import uuid from dataclasses import dataclass +from time import perf_counter from typing import Tuple, List, Dict from urllib.parse import urlparse @@ -47,6 +48,12 @@ validate_visualization_normalization_schema, ) from openjiuwen_deepsearch.algorithm.report.table_caption_utils import ensure_markdown_table_captions +from openjiuwen_deepsearch.algorithm.report.visualization_metrics import ( + VisualizationTaskMetrics, + build_visualization_generation_summary, + build_visualization_insert_summary, + elapsed_ms, +) from openjiuwen_deepsearch.common.exception import CustomValueException from openjiuwen_deepsearch.common.status_code import StatusCode, format_exception_info from openjiuwen_deepsearch.config.config import Config @@ -139,6 +146,16 @@ class VisualizationInsertRenderContext: language: str +@dataclass +class VisualizationMermaidContext: + visualization_content: dict + extracted_obj: dict + visualization_dict: dict + max_attempt_num: int + section_idx: int + metrics: VisualizationTaskMetrics | None = None + + @dataclass class DocSelectionContext: """Encapsulates doc-selection intermediate results for debug export.""" @@ -2707,10 +2724,13 @@ async def _validate_chart_compliance( section_idx: int, section_outline: str, max_attempt_num: int, + metrics: VisualizationTaskMetrics | None = None, ) -> dict: """Validate extracted chart data with compliance prompt.""" payload = (extracted_chart_json or "").strip() for attempt in range(max_attempt_num): + if attempt > 0 and metrics is not None: + metrics.record_retry("chart_compliance") try: llm_input = apply_system_prompt( "chart_compliance_validate", @@ -2781,11 +2801,14 @@ async def _validate_chart_traceability( origin_content: str, section_idx: int, max_attempt_num: int, + metrics: VisualizationTaskMetrics | None = None, ) -> dict: """Validate extracted chart data traceability with origin content.""" payload = (extracted_chart_json or "").strip() origin_text = (origin_content or "").strip() for attempt in range(max_attempt_num): + if attempt > 0 and metrics is not None: + metrics.record_retry("chart_traceability") try: llm_input = apply_system_prompt( "chart_data_traceability_check", @@ -2856,15 +2879,23 @@ async def _extract_visualization_data( visualization_content: dict, max_attempt_num: int, section_idx: int, + metrics: VisualizationTaskMetrics | None = None, ) -> tuple[bool, dict, dict | None]: extract_ok = False extracted_obj = None validation_error = "" previous_records: str | None = None for i in range(max_attempt_num): - visualization_content = await self._extract_data_from_text( - visualization_dict, validation_error, previous_records - ) + if i > 0 and metrics is not None: + metrics.record_retry("extract_data") + extract_started_at = perf_counter() + try: + visualization_content = await self._extract_data_from_text( + visualization_dict, validation_error, previous_records + ) + finally: + if metrics is not None: + metrics.record_stage("extract_data", extract_started_at) if not LogManager.is_sensitive(): logger.debug("%s [process_visualization_task] Extract data: %s.", EFFECT_SUB_REPORT_TAG, visualization_content) @@ -2914,12 +2945,16 @@ async def _extract_visualization_data( visualization_content[ "sub_section_visualization_content" ] = raw_payload + traceability_started_at = perf_counter() traceability = await self._validate_chart_traceability( raw_payload, visualization_dict.get("origin_content", ""), section_idx, max_attempt_num, + metrics, ) + if metrics is not None: + metrics.record_stage("chart_traceability", traceability_started_at) if not traceability.get("valid", False): traceability_error = ( traceability.get("error_msg", "") or "" @@ -2945,12 +2980,16 @@ async def _extract_visualization_data( previous_records = raw_payload or None extract_ok = False continue + compliance_started_at = perf_counter() compliance = await self._validate_chart_compliance( raw_payload, section_idx, visualization_dict.get("section_outline", ""), max_attempt_num, + metrics, ) + if metrics is not None: + metrics.record_stage("chart_compliance", compliance_started_at) if compliance.get("valid", False): validation_error = "" previous_records = None @@ -3002,24 +3041,101 @@ async def _extract_visualization_data( async def _build_visualization_mermaid( self, - visualization_content: dict, - extracted_obj: dict, - visualization_dict: dict, - max_attempt_num: int, - section_idx: int, + context: VisualizationMermaidContext, ) -> dict: - normalized = await self._normalize_visualization_content( - visualization_content, - extracted_obj, - visualization_dict, - max_attempt_num, - section_idx, - ) + visualization_content = context.visualization_content + section_idx = context.section_idx + metrics = context.metrics + normalize_started_at = perf_counter() + normalized = await self._normalize_visualization_content(context) + if metrics is not None: + metrics.record_stage("normalize_units", normalize_started_at) if not normalized: return visualization_content - if not self._precheck_value_variation(visualization_content, section_idx): + + precheck_started_at = perf_counter() + has_value_variation = self._precheck_value_variation( + visualization_content, section_idx + ) + if metrics is not None: + metrics.record_stage("value_variation_precheck", precheck_started_at) + if not has_value_variation: return visualization_content - return self._generate_mermaid_code(visualization_content, section_idx) + + render_started_at = perf_counter() + result = self._generate_mermaid_code(visualization_content, section_idx) + if metrics is not None: + metrics.record_stage("mermaid_render", render_started_at) + return result + + @staticmethod + def _parse_visualization_number(value: str) -> int | float | None: + normalized_value = value.strip().replace(",", "").replace(",", "") + try: + numeric_value = Decimal(normalized_value) + except (InvalidOperation, ValueError): + return None + if not numeric_value.is_finite(): + return None + if numeric_value == numeric_value.to_integral_value(): + return int(numeric_value) + return float(numeric_value) + + @staticmethod + def _scale_visualization_value(value: int | float, divisor: int) -> int | float: + scaled = Decimal(str(value)) / Decimal(divisor) + if scaled == scaled.to_integral_value(): + return int(scaled) + return float(scaled) + + @classmethod + def _normalize_same_unit_records_locally( + cls, + records: list, + image_type: str, + ) -> dict | None: + if image_type not in ("bar", "line", "pie"): + return None + + normalized_records = [] + normalized_unit = None + for row in records: + if not isinstance(row, list) or len(row) != 3: + return None + x_value, numeric_text, unit_text = row + if not ( + isinstance(x_value, str) + and isinstance(numeric_text, str) + and isinstance(unit_text, str) + ): + return None + x_value = x_value.strip() + unit_text = unit_text.strip() + if not x_value or not unit_text: + return None + if normalized_unit is None: + normalized_unit = unit_text + if unit_text != normalized_unit: + return None + + parsed_value = cls._parse_visualization_number(numeric_text) + if parsed_value is None: + return None + normalized_records.append([x_value, parsed_value]) + + if normalized_unit is None: + return None + + if normalized_unit.startswith("万"): + max_abs_value = max(abs(float(row[1])) for row in normalized_records) + if max_abs_value >= 10000: + normalized_unit = "亿" + normalized_unit[1:] + normalized_records = [ + [row[0], cls._scale_visualization_value(row[1], 10000)] + for row in normalized_records + ] + + return {"unit": normalized_unit, "records": normalized_records} @staticmethod def _parse_visualization_number(value: str) -> int | float | None: @@ -3092,12 +3208,14 @@ def _normalize_same_unit_records_locally( async def _normalize_visualization_content( self, - visualization_content: dict, - extracted_obj: dict, - visualization_dict: dict, - max_attempt_num: int, - section_idx: int, + context: VisualizationMermaidContext, ) -> bool: + visualization_content = context.visualization_content + extracted_obj = context.extracted_obj + visualization_dict = context.visualization_dict + max_attempt_num = context.max_attempt_num + section_idx = context.section_idx + metrics = context.metrics # Extracted schema is valid here. image_title = extracted_obj.get("image_title", "") image_type = extracted_obj.get("image_type", "") @@ -3157,6 +3275,8 @@ async def _normalize_visualization_content( "sub_section_visualization_normalize_units", normalize_context ) for j in range(max_attempt_num): + if j > 0 and metrics is not None: + metrics.record_retry("normalize_units") normalize_output = await ainvoke_llm_with_stats( llm=self._llm, messages=normalize_input, @@ -3212,12 +3332,16 @@ async def _process_visualization_task(self, visualization_dict: dict) -> dict: """Process one visualization task (LLM content + Mermaid generation)""" section_idx = visualization_dict.get("section_idx", 1) max_attempt_num = visualization_dict.get("max_attempt_num", 3) + metrics = VisualizationTaskMetrics(section_idx=section_idx) # Extract structured data visualization_content = dict(rs_success=True, visualization_content="") origin_content = (visualization_dict.get("origin_content") or "").strip() if not origin_content: visualization_content["rs_success"] = False visualization_content["error_msg"] = "origin_content_empty" + visualization_content["_visualization_metrics"] = metrics.finish( + False, "origin_content_empty" + ) return visualization_content extract_ok, visualization_content, extracted_obj = ( await self._extract_visualization_data( @@ -3225,22 +3349,34 @@ async def _process_visualization_task(self, visualization_dict: dict) -> dict: visualization_content, max_attempt_num, section_idx, + metrics, ) ) if not extract_ok: + visualization_content["_visualization_metrics"] = metrics.finish( + False, visualization_content.get("error_msg", "extract_data_failed") + ) return visualization_content - return await self._build_visualization_mermaid( - visualization_content, - extracted_obj, - visualization_dict, - max_attempt_num, - section_idx, + mermaid_context = VisualizationMermaidContext( + visualization_content=visualization_content, + extracted_obj=extracted_obj, + visualization_dict=visualization_dict, + max_attempt_num=max_attempt_num, + section_idx=section_idx, + metrics=metrics, + ) + visualization_content = await self._build_visualization_mermaid(mermaid_context) + visualization_content["_visualization_metrics"] = metrics.finish( + bool(visualization_content.get("rs_success")), + visualization_content.get("error_msg", ""), ) + return visualization_content async def _generate_content_for_visualization(self, current_inputs: dict) -> dict: """Generate content for visualization with concurrent LLM calls""" section_idx = current_inputs.get("section_idx", 1) + metrics_started_at = perf_counter() # Compliance validation depends on chapter outline; if outline is missing, skip visuals safely. section_outline = (current_inputs.get("sub_section_outline", "") or "").strip() if not section_outline: @@ -3275,9 +3411,31 @@ async def _generate_content_for_visualization(self, current_inputs: dict) -> dic visualization_content = self._select_visualization_from_classified_content( classified_content_for_visualization ) + source_candidate_count = len(classified_content_for_visualization) + pre_budget_candidate_count = len(visualization_content) + visualization_content = self._limit_visualization_candidates( + visualization_content, + current_inputs, + ) n = len(visualization_content) if n == 0: + summary = build_visualization_generation_summary( + section_idx=section_idx, + source_candidate_count=source_candidate_count, + selected_candidate_count=0, + generated_mermaid_count=0, + task_metrics=[], + exception_count=0, + wall_time_ms=elapsed_ms(metrics_started_at), + pre_budget_candidate_count=pre_budget_candidate_count, + candidate_budget=0, + ) + logger.info( + "%s [visualization_metrics] generation_summary: %s", + EFFECT_SUB_REPORT_TAG, + json.dumps(summary, ensure_ascii=False), + ) return dict(rs_success=True, visualization_content=visualization_content) # Build all async tasks tasks = [] @@ -3301,8 +3459,12 @@ async def _generate_content_for_visualization(self, current_inputs: dict) -> dic results = await asyncio.gather(*tasks, return_exceptions=True) # Aggregate results + task_metrics = [] + exception_count = 0 + generated_mermaid_count = 0 for i, res in enumerate(results): if isinstance(res, Exception): + exception_count += 1 logger.error( "%s [generate_sub_section_visualization_content] section_idx: [%s], " "error in task [%s]: %s", @@ -3314,11 +3476,16 @@ async def _generate_content_for_visualization(self, current_inputs: dict) -> dic visualization_content[i]["sub_section_visualization_content"] = "" visualization_content[i]["mermaid_content"] = "" else: + metric = res.get("_visualization_metrics") + if isinstance(metric, dict): + task_metrics.append(metric) if res.get("rs_success"): visualization_content[i]["sub_section_visualization_content"] = res[ "sub_section_visualization_content" ] visualization_content[i]["mermaid_content"] = res["mermaid_content"] + if res.get("mermaid_content"): + generated_mermaid_count += 1 else: visualization_content[i]["sub_section_visualization_content"] = "" visualization_content[i]["mermaid_content"] = "" @@ -3328,6 +3495,22 @@ async def _generate_content_for_visualization(self, current_inputs: dict) -> dic section_idx, res.get("error_msg", "Unknown"), ) + summary = build_visualization_generation_summary( + section_idx=section_idx, + source_candidate_count=source_candidate_count, + selected_candidate_count=n, + generated_mermaid_count=generated_mermaid_count, + task_metrics=task_metrics, + exception_count=exception_count, + wall_time_ms=elapsed_ms(metrics_started_at), + pre_budget_candidate_count=pre_budget_candidate_count, + candidate_budget=n, + ) + logger.info( + "%s [visualization_metrics] generation_summary: %s", + EFFECT_SUB_REPORT_TAG, + json.dumps(summary, ensure_ascii=False), + ) return dict(rs_success=True, visualization_content=visualization_content) @staticmethod @@ -3883,6 +4066,184 @@ def _select_visualization_from_classified_content( fallback_visualizations.append(item) return selected_visualizations or fallback_visualizations + @staticmethod + def _visualization_candidate_text(item: dict) -> str: + parts = [] + for key in ("title", "original_content", "chunk", "content"): + value = item.get(key) + if isinstance(value, str) and value.strip(): + parts.append(value.strip()) + + key_passages = item.get("key_passages") + if isinstance(key_passages, list): + for passage in key_passages: + passage_text = str(passage or "").strip() + if passage_text: + parts.append(passage_text) + + return "\n".join(parts) + + @staticmethod + def _count_visualization_numbers(item: dict) -> int: + text = Reporter._visualization_candidate_text(item) + if not text: + return 0 + return len( + re.findall( + r"(? bool: + citation_indices = Reporter._normalize_citation_indices( + item.get("citation_indices") + ) + if citation_indices: + return True + + citation_indices = Reporter._normalize_citation_indices( + [item.get("citation_index"), item.get("index")] + ) + if citation_indices: + return True + + return bool(str(item.get("url") or "").strip()) + + @staticmethod + def _visualization_candidate_signature(item: dict) -> str: + text = Reporter._visualization_candidate_text(item).lower() + text = re.sub(r"https?://\S+", " ", text) + text = re.sub(r"\[citation:\d+\]", " ", text) + text = re.sub(r"\s+", " ", text).strip() + return text[:240] + + @classmethod + def _deduplicate_visualization_candidates(cls, candidates: list[dict]) -> list[dict]: + deduplicated = [] + seen = set() + for item in candidates: + signature = cls._visualization_candidate_signature(item) + if signature and signature in seen: + continue + if signature: + seen.add(signature) + deduplicated.append(item) + return deduplicated + + @staticmethod + def _rich_visualization_requested(current_inputs: dict) -> bool: + text_parts = [] + for key in ( + "report_task", + "section_description", + "section_format_requirements", + "report_style", + "paragraph_style", + ): + text_parts.append(str(current_inputs.get(key, ""))) + text = " ".join(text_parts).lower() + rich_markers = ( + "图文并茂", + "多图", + "多张图", + "多张图表", + "图表丰富", + "可视化", + "visualization", + "visualisation", + "visualize", + "visualise", + "multiple charts", + "chart-rich", + "data charts", + ) + return any(marker in text for marker in rich_markers) + + @staticmethod + def _brief_visualization_context(current_inputs: dict) -> bool: + text_parts = [] + for key in ("report_type", "paragraph_style", "report_style", "report_task"): + text_parts.append(str(current_inputs.get(key, ""))) + text = " ".join(text_parts).lower() + brief_markers = ("brief", "short", "concise", "短篇", "简短") + return any(marker in text for marker in brief_markers) + + @classmethod + def _rank_visualization_candidates(cls, candidates: list[dict]) -> list[dict]: + ranked = [] + for order, item in enumerate(candidates): + data_density = get_numeric_score(item, "data_density") or 0.0 + number_count = cls._count_visualization_numbers(item) + trace_bonus = 1 if cls._has_visualization_trace(item) else 0 + score = data_density * 10 + min(number_count, 8) * 2 + trace_bonus * 12 + if number_count < 2: + score -= 12 + ranked.append((score, data_density, number_count, trace_bonus, -order, item)) + ranked.sort(reverse=True) + ordered = [] + for ranked_item in ranked: + ordered.append(ranked_item[-1]) + return ordered + + @classmethod + def _visualization_candidate_budget( + cls, + candidates: list[dict], + current_inputs: dict, + ) -> int: + candidate_count = len(candidates) + if candidate_count <= 2: + return candidate_count + + strong_candidate_count = 0 + for item in candidates: + data_density = get_numeric_score(item, "data_density") or 0.0 + number_count = cls._count_visualization_numbers(item) + if data_density >= 9.0 and number_count >= 3: + strong_candidate_count += 1 + + rich_visualization = cls._rich_visualization_requested(current_inputs) + brief_context = cls._brief_visualization_context(current_inputs) + + limit = 2 + if strong_candidate_count >= 3: + limit = 3 + if rich_visualization and strong_candidate_count >= 5: + limit = 4 + if brief_context and not rich_visualization: + limit = min(limit, 2) + + return max(1, min(candidate_count, limit)) + + @classmethod + def _limit_visualization_candidates( + cls, + candidates: list[dict], + current_inputs: dict, + ) -> list[dict]: + if len(candidates) <= 2: + return candidates + + deduplicated = cls._deduplicate_visualization_candidates(candidates) + ranked = cls._rank_visualization_candidates(deduplicated) + budget = cls._visualization_candidate_budget(ranked, current_inputs) + limited = ranked[:budget] + + if len(limited) < len(candidates): + logger.info( + "%s [visualization_budget] section_idx: [%s], candidates: %s, " + "deduplicated: %s, selected: %s.", + EFFECT_SUB_REPORT_TAG, + current_inputs.get("section_idx", 1), + len(candidates), + len(deduplicated), + len(limited), + ) + + return limited + async def _request_visualization_insert_plan( self, context: VisualizationInsertPlanContext ) -> dict: @@ -4084,10 +4445,77 @@ def _complete_visualization_insertions( ) return completed + @staticmethod + def _last_valid_visualization_anchor( + report_lines: list[str], + invalid_rows: set[int], + *, + allow_heading: bool = False, + ) -> int | None: + for row_idx in range(len(report_lines), 0, -1): + if row_idx in invalid_rows: + continue + stripped = report_lines[row_idx - 1].strip() + if not stripped: + continue + if not allow_heading and stripped.startswith("#"): + continue + return row_idx + return None + + @classmethod + def _single_visualization_insertion( + cls, + mermaid_map: dict[int, str], + title_meta_map: dict[int, dict], + report_lines: list[str], + invalid_rows: set[int], + ) -> list[dict]: + if len(mermaid_map) != 1: + return [] + + index = next(iter(mermaid_map)) + title_meta = title_meta_map.get(index, {}) + citation_indices = cls._normalize_citation_indices( + title_meta.get("citation_indices") + ) + if not citation_indices: + citation_indices = cls._normalize_citation_indices( + [title_meta.get("citation_index")] + ) + + if citation_indices: + citation_tokens = [f"[citation:{citation}]" for citation in citation_indices] + for row_idx, line in reversed(list(enumerate(report_lines, 1))): + if row_idx in invalid_rows: + continue + stripped = line.strip() + if not stripped or stripped.startswith("#"): + continue + if any(token in line for token in citation_tokens): + return [{"after_row": row_idx, "index": index}] + + fallback_row = cls._last_valid_visualization_anchor( + report_lines, + invalid_rows, + ) + if fallback_row is None: + fallback_row = cls._last_valid_visualization_anchor( + report_lines, + invalid_rows, + allow_heading=True, + ) + if fallback_row is None: + return [] + return [{"after_row": fallback_row, "index": index}] + async def _insert_visualization(self, current_inputs: Dict) -> dict: """ Insert placeholders for visualization content in the markdown report. """ + metrics_started_at = perf_counter() + section_idx = current_inputs.get("section_idx", 1) + original_report = "" try: report_markdown = current_inputs.get("sub_report_content", "") if not isinstance(report_markdown, str): @@ -4095,7 +4523,35 @@ async def _insert_visualization(self, current_inputs: Dict) -> dict: original_report = report_markdown visualization_list = current_inputs.get("visualization_result", []) + + def log_insert_metrics( + valid_mermaid_count: int, + planned_insertion_count: int, + inserted_mermaid_count: int, + planner_mode: str = "", + ) -> None: + visualization_count = ( + len(visualization_list) + if isinstance(visualization_list, list) + else 0 + ) + summary = build_visualization_insert_summary( + section_idx=section_idx, + visualization_count=visualization_count, + valid_mermaid_count=valid_mermaid_count, + planned_insertion_count=planned_insertion_count, + inserted_mermaid_count=inserted_mermaid_count, + wall_time_ms=elapsed_ms(metrics_started_at), + planner_mode=planner_mode, + ) + logger.info( + "%s [visualization_metrics] insert_summary: %s", + EFFECT_SUB_REPORT_TAG, + json.dumps(summary, ensure_ascii=False), + ) + if not isinstance(visualization_list, list) or not visualization_list: + log_insert_metrics(0, 0, 0) return dict(rs_success=False, result=original_report) report_lines = report_markdown.splitlines(keepends=True) @@ -4163,8 +4619,37 @@ async def _insert_visualization(self, current_inputs: Dict) -> dict: if not mermaid_map: # No valid visualization blocks, return original content. + log_insert_metrics(0, 0, 0) return dict(rs_success=False, result=original_report) + if len(mermaid_map) == 1: + insertions = self._single_visualization_insertion( + mermaid_map, + title_meta_map, + report_lines, + invalid_rows, + ) + if not insertions: + log_insert_metrics(len(mermaid_map), 0, 0, "local_single_chart") + return dict(rs_success=False, result=original_report) + rendered = self._apply_visualization_insertions( + VisualizationInsertRenderContext( + report_lines=report_lines, + insertions=insertions, + mermaid_map=mermaid_map, + title_meta_map=title_meta_map, + newline=newline, + language=current_inputs.get("language"), + ) + ) + log_insert_metrics( + len(mermaid_map), + len(insertions), + len(insertions), + "local_single_chart", + ) + return dict(rs_success=True, result=rendered) + llm_input_message = numbered_report.rstrip("\r\n") + "\n\n" llm_input_message += "=== VISUALIZATION DATA ===\n" for visualization_item in visualization_items: @@ -4185,6 +4670,7 @@ async def _insert_visualization(self, current_inputs: Dict) -> dict: ) ) if not plan_result.get("rs_success") or not plan_result.get("plan"): + log_insert_metrics(len(mermaid_map), 0, 0, "llm") return dict(rs_success=False, result=original_report) plan = plan_result["plan"] @@ -4207,11 +4693,12 @@ async def _insert_visualization(self, current_inputs: Dict) -> dict: language=current_inputs.get("language"), ) ) + log_insert_metrics(len(mermaid_map), len(insertions), len(insertions), "llm") return dict(rs_success=True, result=rendered) except Exception as e: logger.error( f"{EFFECT_SUB_REPORT_TAG} Unexpected error when inserting visualization for the section " - f"{current_inputs.get('section_idx', 1)}: {str(e)}", + f"{section_idx}: {str(e)}", exc_info=True, ) return dict(rs_success=False, result=original_report) diff --git a/deepsearch/openjiuwen_deepsearch/algorithm/report/visualization_metrics.py b/deepsearch/openjiuwen_deepsearch/algorithm/report/visualization_metrics.py new file mode 100644 index 00000000..f37df7bf --- /dev/null +++ b/deepsearch/openjiuwen_deepsearch/algorithm/report/visualization_metrics.py @@ -0,0 +1,128 @@ +# -*- coding: UTF-8 -*- +# Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved. +"""Runtime metrics helpers for report visualization generation.""" + +from collections import Counter +from dataclasses import dataclass, field +from time import perf_counter +from typing import Any + + +def elapsed_ms(started_at: float) -> int: + """Return non-negative elapsed milliseconds from a perf_counter timestamp.""" + return max(int((perf_counter() - started_at) * 1000), 0) + + +def _merge_counter(target: dict[str, int], source: dict[str, int]) -> None: + for key, value in source.items(): + if not isinstance(value, int): + continue + target[key] = target.get(key, 0) + value + + +@dataclass +class VisualizationTaskMetrics: + """Collect metrics for one visualization candidate.""" + + section_idx: int + stage_durations_ms: dict[str, int] = field(default_factory=dict) + retry_counts: dict[str, int] = field(default_factory=dict) + rs_success: bool = False + error_msg: str = "" + _started_at: float = field(default_factory=perf_counter, repr=False) + + def record_stage(self, stage: str, started_at: float) -> None: + self.stage_durations_ms[stage] = ( + self.stage_durations_ms.get(stage, 0) + elapsed_ms(started_at) + ) + + def record_retry(self, stage: str, count: int = 1) -> None: + self.retry_counts[stage] = self.retry_counts.get(stage, 0) + count + + def finish(self, rs_success: bool, error_msg: str = "") -> dict[str, Any]: + self.rs_success = bool(rs_success) + self.error_msg = str(error_msg or "") + self.record_stage("candidate_total", self._started_at) + return self.as_log_dict() + + def as_log_dict(self) -> dict[str, Any]: + return { + "section_idx": self.section_idx, + "rs_success": self.rs_success, + "error_msg": self.error_msg, + "stage_durations_ms": dict(self.stage_durations_ms), + "retry_counts": dict(self.retry_counts), + } + + +def build_visualization_generation_summary( + *, + section_idx: int, + source_candidate_count: int, + selected_candidate_count: int, + generated_mermaid_count: int, + task_metrics: list[dict[str, Any]], + exception_count: int, + wall_time_ms: int, + pre_budget_candidate_count: int | None = None, + candidate_budget: int | None = None, +) -> dict[str, Any]: + """Build a log-safe summary for one section's visualization generation.""" + failure_reasons: Counter[str] = Counter() + stage_durations_ms: dict[str, int] = {} + retry_counts: dict[str, int] = {} + successful_candidate_count = 0 + + for item in task_metrics: + if item.get("rs_success"): + successful_candidate_count += 1 + else: + reason = str(item.get("error_msg") or "unknown") + failure_reasons[reason] += 1 + _merge_counter(stage_durations_ms, item.get("stage_durations_ms", {})) + _merge_counter(retry_counts, item.get("retry_counts", {})) + + if exception_count: + failure_reasons["exception"] += exception_count + + summary = { + "section_idx": section_idx, + "source_candidate_count": source_candidate_count, + "selected_candidate_count": selected_candidate_count, + "attempted_candidate_count": len(task_metrics) + exception_count, + "successful_candidate_count": successful_candidate_count, + "generated_mermaid_count": generated_mermaid_count, + "failure_reasons": dict(failure_reasons), + "retry_counts": retry_counts, + "stage_durations_ms": stage_durations_ms, + "wall_time_ms": wall_time_ms, + } + if pre_budget_candidate_count is not None: + summary["pre_budget_candidate_count"] = pre_budget_candidate_count + if candidate_budget is not None: + summary["candidate_budget"] = candidate_budget + return summary + + +def build_visualization_insert_summary( + *, + section_idx: int, + visualization_count: int, + valid_mermaid_count: int, + planned_insertion_count: int, + inserted_mermaid_count: int, + wall_time_ms: int, + planner_mode: str | None = None, +) -> dict[str, Any]: + """Build a log-safe summary for one section's visualization insertion.""" + summary = { + "section_idx": section_idx, + "visualization_count": visualization_count, + "valid_mermaid_count": valid_mermaid_count, + "planned_insertion_count": planned_insertion_count, + "inserted_mermaid_count": inserted_mermaid_count, + "wall_time_ms": wall_time_ms, + } + if planner_mode: + summary["planner_mode"] = planner_mode + return summary diff --git a/deepsearch/openjiuwen_deepsearch/framework/openjiuwen/agent/main_graph_nodes.py b/deepsearch/openjiuwen_deepsearch/framework/openjiuwen/agent/main_graph_nodes.py index b89181a1..f2420169 100644 --- a/deepsearch/openjiuwen_deepsearch/framework/openjiuwen/agent/main_graph_nodes.py +++ b/deepsearch/openjiuwen_deepsearch/framework/openjiuwen/agent/main_graph_nodes.py @@ -17,6 +17,7 @@ from openjiuwen.core.workflow.components.flow.end_comp import End from openjiuwen.core.workflow.components.flow.start_comp import Start +from openjiuwen_deepsearch.algorithm.chart_generation.utils import remove_mermaid_code_blocks from openjiuwen_deepsearch.algorithm.chart_generation.vlm_chart_generator import VLMChartGenerator from openjiuwen_deepsearch.algorithm.search_nodes.utils import ( anonymize_config_for_logging, @@ -2128,10 +2129,12 @@ async def _do_invoke(self, inputs: Input, session: Session, context: ModelContex def _post_handle(self, inputs: Input, algorithm_output: dict, session: Session, context: ModelContext) -> dict: + current_inputs = algorithm_output.get("current_inputs", {}) vlm_chart_generator_output = algorithm_output.get("vlm_chart_generator_output", {}) chart_messages = [] - modified_report = algorithm_output.get("current_inputs", {}).get("report_content", "") - new_source_trace_datas = algorithm_output.get("current_inputs", {}).get("trace_source_datas", []) + modified_report = current_inputs.get("report_content", "") + new_source_trace_datas = current_inputs.get("trace_source_datas", []) + should_update_report = False if vlm_chart_generator_output.get("skip_node", False): logger.info("[VLMChartGeneratorNode] vlm_chart_generator_enable is False, skip VLMChartGeneratorNode.") @@ -2148,13 +2151,24 @@ def _post_handle(self, inputs: Input, algorithm_output: dict, session: Session, chart_messages = vlm_chart_generator_output.get("chart_messages", []) modified_report = vlm_chart_generator_output.get("modified_report", "") new_source_trace_datas = vlm_chart_generator_output.get("new_source_trace_datas", []) + should_update_report = True - current_report = session.get_global_state("search_context.current_report") - current_report.report_content = modified_report - current_report.merged_trace_source_datas = new_source_trace_datas - session.update_global_state({"search_context.current_report": current_report}) session.update_global_state({"search_context.final_result.chart_messages": chart_messages}) + if current_inputs.get("vlm_chart_generator_enable", False): + cleaned_report = remove_mermaid_code_blocks(modified_report) + if cleaned_report != modified_report: + logger.warning("[VLMChartGeneratorNode] Removed Mermaid code block(s) from report content.") + modified_report = cleaned_report + should_update_report = True + + if should_update_report: + current_report = session.get_global_state("search_context.current_report") + if current_report: + current_report.report_content = modified_report + current_report.merged_trace_source_datas = new_source_trace_datas + session.update_global_state({"search_context.current_report": current_report}) + add_debug_log_wrapper( session, NodeDebugData( diff --git a/deepsearch/tests/algorithm/chart_generation/test_utils.py b/deepsearch/tests/algorithm/chart_generation/test_utils.py new file mode 100644 index 00000000..227e9bfa --- /dev/null +++ b/deepsearch/tests/algorithm/chart_generation/test_utils.py @@ -0,0 +1,44 @@ +from openjiuwen_deepsearch.algorithm.chart_generation.utils import remove_mermaid_code_blocks + + +def test_remove_mermaid_code_blocks_keeps_non_mermaid_markdown_unchanged(): + markdown = "正文前\n\n\n```python\ngraph = {'A': 'B'}\n```\n\n正文后\n" + + assert remove_mermaid_code_blocks(markdown) == markdown + + +def test_remove_mermaid_code_blocks_removes_fenced_mermaid_and_preserves_text(): + markdown = ( + "正文前\n\n" + "```mermaid\n" + "flowchart TD\n" + " A --> B\n" + "```\n\n" + "正文后" + ) + + cleaned = remove_mermaid_code_blocks(markdown) + + assert "```mermaid" not in cleaned + assert "flowchart TD" not in cleaned + assert "正文前" in cleaned + assert "正文后" in cleaned + + +def test_remove_mermaid_code_blocks_removes_unlabeled_mermaid_block(): + markdown = "正文前\n\n```\nflowchart TD\n A --> B\n```\n\n正文后" + + cleaned = remove_mermaid_code_blocks(markdown) + + assert "```" not in cleaned + assert "flowchart TD" not in cleaned + assert cleaned == "正文前\n\n正文后" + + +def test_remove_mermaid_code_blocks_removes_indented_mermaid_block(): + markdown = "正文前\n\n flowchart TD\n A --> B\n\n正文后" + + cleaned = remove_mermaid_code_blocks(markdown) + + assert "flowchart TD" not in cleaned + assert cleaned == "正文前\n\n正文后" diff --git a/deepsearch/tests/node/test_agent_node.py b/deepsearch/tests/node/test_agent_node.py index c8c03431..96f9e175 100644 --- a/deepsearch/tests/node/test_agent_node.py +++ b/deepsearch/tests/node/test_agent_node.py @@ -23,6 +23,7 @@ OutlineNode, StartNode, UserFeedbackProcessorNode, + VLMChartGeneratorNode, ) from openjiuwen_deepsearch.algorithm.query_understanding.intent_recognition import IntentRecognitionResult from openjiuwen_deepsearch.config.config import OUTLINER_SECTION_NUM_MAX @@ -1134,6 +1135,93 @@ def _get_global_state(key): assert result_data["response_content"].endswith("This research report was generated by AI.") +def test_vlm_post_handle_removes_mermaid_blocks_from_successful_report(): + """VLM flow must not leave raw Mermaid source blocks in report body.""" + node = VLMChartGeneratorNode() + current_report = Mock(report_content="", merged_trace_source_datas=[]) + session = Mock(spec=Session) + session.get_global_state.side_effect = ( + lambda key: current_report if key == "search_context.current_report" else None + ) + session.update_global_state = Mock() + report_content = ( + "# RAG 系统\n\n" + "前文说明。\n\n" + "```mermaid\n" + "flowchart TD\n" + " A[用户查询] --> B[检索器]\n" + "```\n\n" + "```python\n" + "graph = {'A': 'B'}\n" + "```\n\n" + "后文说明。" + ) + algorithm_output = { + "current_inputs": { + "vlm_chart_generator_enable": True, + "report_content": report_content, + "trace_source_datas": [{"id": "old"}], + }, + "vlm_chart_generator_output": { + "chart_messages": [{"chart_id": "chart_1"}], + "modified_report": report_content, + "new_source_trace_datas": [{"id": "new"}], + }, + } + + with patch("openjiuwen_deepsearch.framework.openjiuwen.agent.main_graph_nodes.add_debug_log_wrapper"): + output = node._post_handle({}, algorithm_output, session, Context()) + + assert output == {"next_node": NodeId.SOURCE_TRACER.value} + assert "```mermaid" not in current_report.report_content + assert "flowchart TD" not in current_report.report_content + assert "graph = {'A': 'B'}" in current_report.report_content + assert "前文说明" in current_report.report_content + assert "后文说明" in current_report.report_content + assert current_report.merged_trace_source_datas == [{"id": "new"}] + session.update_global_state.assert_any_call( + {"search_context.final_result.chart_messages": [{"chart_id": "chart_1"}]} + ) + + +def test_vlm_post_handle_removes_existing_mermaid_blocks_when_generation_fails(): + """Even failed VLM generation sanitizes pre-existing body Mermaid blocks before source tracing.""" + node = VLMChartGeneratorNode() + current_report = Mock(report_content="", merged_trace_source_datas=[]) + session = Mock(spec=Session) + session.get_global_state.side_effect = ( + lambda key: current_report if key == "search_context.current_report" else None + ) + session.update_global_state = Mock() + report_content = ( + "节点解析如下。\n\n" + "```\n" + "flowchart TD\n" + " A --> B\n" + "```\n\n" + "继续分析各节点。" + ) + algorithm_output = { + "current_inputs": { + "vlm_chart_generator_enable": True, + "report_content": report_content, + "trace_source_datas": [{"id": "old"}], + }, + "vlm_chart_generator_output": {"error_msg": "vlm failed"}, + } + + with patch("openjiuwen_deepsearch.framework.openjiuwen.agent.main_graph_nodes.add_debug_log_wrapper"): + node._post_handle({}, algorithm_output, session, Context()) + + assert "```" not in current_report.report_content + assert "flowchart TD" not in current_report.report_content + assert "节点解析如下" in current_report.report_content + assert "继续分析各节点" in current_report.report_content + assert current_report.merged_trace_source_datas == [{"id": "old"}] + session.update_global_state.assert_any_call({"search_context.final_result.warning_info": "vlm failed"}) + session.update_global_state.assert_any_call({"search_context.current_report": current_report}) + + @pytest.mark.asyncio async def test_end_node_writes_workflow_llm_usage_when_stats_enabled(): """验证 EndNode 在开启统计时会写入 workflow 级 token 汇总。""" diff --git a/deepsearch/tests/report/test_sub_report.py b/deepsearch/tests/report/test_sub_report.py index 33fee44d..ef5d5cd6 100644 --- a/deepsearch/tests/report/test_sub_report.py +++ b/deepsearch/tests/report/test_sub_report.py @@ -14,6 +14,7 @@ from openjiuwen_deepsearch.algorithm.report.report import ( Reporter, VisualizationInsertPlanContext, + VisualizationMermaidContext, _get_classified_infos, ) from openjiuwen_deepsearch.algorithm.report.table_caption_utils import ensure_markdown_table_captions @@ -629,6 +630,94 @@ def test_select_visualization_uses_eight_point_fallback_when_no_high_density_doc assert [item["title"] for item in selected] == ["fallback density"] +def _visualization_candidate( + title: str, + density: float, + content: str, + *, + index: int | None = 1, + url: str = "https://example.com/source", +) -> dict: + item = { + "title": title, + "original_content": content, + "scores": {"data_density": density}, + } + if index is not None: + item["index"] = index + if url: + item["url"] = url + return item + + +def test_limit_visualization_candidates_keeps_small_candidate_sets(): + candidates = [ + _visualization_candidate("one", 9.5, "2020 10 2021 20"), + _visualization_candidate("two", 9.2, "A 30 B 40"), + ] + + selected = Reporter._limit_visualization_candidates(candidates, {}) + + assert selected == candidates + + +def test_limit_visualization_candidates_prefers_dense_numeric_sourced_items(): + candidates = [ + _visualization_candidate("weak numeric", 9.9, "only one value 2024"), + _visualization_candidate("trend", 9.5, "2019 425 2020 124 2021 213"), + _visualization_candidate("comparison", 9.4, "A 24% B 18% C 15%"), + _visualization_candidate("structure", 9.3, "2020 27% 2029 22%"), + _visualization_candidate("no source", 9.6, "Q1 10 Q2 20 Q3 30", index=None, url=""), + ] + + selected = Reporter._limit_visualization_candidates( + candidates, + {"section_idx": 2}, + ) + + assert [item["title"] for item in selected] == [ + "trend", + "structure", + "comparison", + ] + + +def test_limit_visualization_candidates_allows_extra_for_rich_visual_request(): + candidates = [ + _visualization_candidate( + f"candidate-{idx}", + 9.5, + f"metric-{idx} 2020 {idx} 2021 {idx + 1} 2022 {idx + 2}", + ) + for idx in range(6) + ] + + selected = Reporter._limit_visualization_candidates( + candidates, + {"report_task": "请生成图文并茂报告,尽量包含多张图表"}, + ) + + assert len(selected) == 4 + + +def test_limit_visualization_candidates_caps_brief_reports_without_rich_request(): + candidates = [ + _visualization_candidate( + f"candidate-{idx}", + 9.5, + f"metric-{idx} 2020 {idx} 2021 {idx + 1} 2022 {idx + 2}", + ) + for idx in range(5) + ] + + selected = Reporter._limit_visualization_candidates( + candidates, + {"report_type": "brief"}, + ) + + assert len(selected) == 2 + + def _visualization_reporter() -> Reporter: reporter = Reporter.__new__(Reporter) reporter._llm = object() @@ -776,11 +865,13 @@ async def test_visualization_normalization_uses_local_same_unit_fast_path(): new_callable=AsyncMock, ) as mocked_llm: normalized = await reporter._normalize_visualization_content( - visualization_content=visualization_content, - extracted_obj=extracted_obj, - visualization_dict={"language": "zh-CN"}, - max_attempt_num=3, - section_idx=1, + VisualizationMermaidContext( + visualization_content=visualization_content, + extracted_obj=extracted_obj, + visualization_dict={"language": "zh-CN"}, + max_attempt_num=3, + section_idx=1, + ) ) assert normalized is True @@ -981,21 +1072,75 @@ async def test_insert_visualization_renders_all_chart_citation_indices(): ], } + mock_ainvoke = AsyncMock( + return_value={"content": '{"insertions":[{"after_row":3,"index":1}]}'} + ) with patch( "openjiuwen_deepsearch.algorithm.report.report.ainvoke_llm_with_stats", - new=AsyncMock( - return_value={"content": '{"insertions":[{"after_row":3,"index":1}]}'} - ), + new=mock_ainvoke, ): result = await _visualization_reporter()._insert_visualization(current_inputs) assert result["rs_success"] is True + mock_ainvoke.assert_not_awaited() assert ( "**Vendor revenue comparison[citation:7][citation:8][citation:9]**" in result["result"] ) +@pytest.mark.asyncio +async def test_insert_visualization_single_chart_uses_local_citation_anchor(): + chart = { + "image_title": "Revenue trend", + "image_type": "line", + "unit": "million USD", + "records": [["2022", 10], ["2023", 20], ["2024", 30]], + } + current_inputs = { + "language": "en", + "section_idx": 2, + "max_generate_retry_num": 1, + "sub_report_content": ( + "# Section\n\n" + "Opening paragraph [citation:1].\n\n" + "Revenue rose across the period [citation:7].\n\n" + "Closing paragraph.\n" + ), + "visualization_result": [ + { + "url": "https://source.example/revenue", + "citation_indices": [7], + "sub_section_visualization_content": json.dumps(chart), + "mermaid_content": ( + 'xychart-beta\n x-axis ["2022", "2023", "2024"]\n' + " line [10, 20, 30]" + ), + } + ], + } + + mock_ainvoke = AsyncMock( + return_value={"content": '{"insertions":[{"after_row":7,"index":1}]}'} + ) + with patch( + "openjiuwen_deepsearch.algorithm.report.report.ainvoke_llm_with_stats", + new=mock_ainvoke, + ): + result = await _visualization_reporter()._insert_visualization(current_inputs) + + assert result["rs_success"] is True + mock_ainvoke.assert_not_awaited() + rendered = result["result"] + assert rendered.count("```mermaid") == 1 + assert ( + rendered.index("Revenue rose across the period") + < rendered.index("```mermaid") + < rendered.index("Closing paragraph.") + ) + assert "**Revenue trend[citation:7]**" in rendered + + @pytest.mark.asyncio async def test_insert_visualization_completes_missing_chart_indices_from_llm_plan(): chart_one = { diff --git a/deepsearch/tests/report/test_visualization_metrics.py b/deepsearch/tests/report/test_visualization_metrics.py new file mode 100644 index 00000000..65d98cee --- /dev/null +++ b/deepsearch/tests/report/test_visualization_metrics.py @@ -0,0 +1,91 @@ +"""Tests for Mermaid visualization runtime metrics.""" + +from openjiuwen_deepsearch.algorithm.report.visualization_metrics import ( + VisualizationTaskMetrics, + build_visualization_generation_summary, + build_visualization_insert_summary, +) + + +def test_visualization_task_metrics_records_stage_and_retry(): + metrics = VisualizationTaskMetrics(section_idx=2) + + metrics.stage_durations_ms["extract_data"] = 12 + metrics.record_retry("extract_data") + metrics.record_retry("extract_data") + result = metrics.finish(False, "extract_data_failed") + + assert result["section_idx"] == 2 + assert result["rs_success"] is False + assert result["error_msg"] == "extract_data_failed" + assert result["retry_counts"] == {"extract_data": 2} + assert result["stage_durations_ms"]["extract_data"] == 12 + assert result["stage_durations_ms"]["candidate_total"] >= 0 + + +def test_build_visualization_generation_summary_aggregates_candidates(): + task_metrics = [ + { + "rs_success": True, + "error_msg": "", + "stage_durations_ms": {"extract_data": 100, "candidate_total": 150}, + "retry_counts": {"extract_data": 1}, + }, + { + "rs_success": False, + "error_msg": "normalize_failed", + "stage_durations_ms": {"normalize_units": 40, "candidate_total": 70}, + "retry_counts": {"normalize_units": 2}, + }, + ] + + summary = build_visualization_generation_summary( + section_idx=3, + source_candidate_count=5, + pre_budget_candidate_count=4, + candidate_budget=2, + selected_candidate_count=2, + generated_mermaid_count=1, + task_metrics=task_metrics, + exception_count=1, + wall_time_ms=220, + ) + + assert summary["section_idx"] == 3 + assert summary["source_candidate_count"] == 5 + assert summary["pre_budget_candidate_count"] == 4 + assert summary["candidate_budget"] == 2 + assert summary["selected_candidate_count"] == 2 + assert summary["attempted_candidate_count"] == 3 + assert summary["successful_candidate_count"] == 1 + assert summary["generated_mermaid_count"] == 1 + assert summary["failure_reasons"] == {"normalize_failed": 1, "exception": 1} + assert summary["retry_counts"] == {"extract_data": 1, "normalize_units": 2} + assert summary["stage_durations_ms"] == { + "extract_data": 100, + "candidate_total": 220, + "normalize_units": 40, + } + assert summary["wall_time_ms"] == 220 + + +def test_build_visualization_insert_summary_keeps_log_safe_counts(): + summary = build_visualization_insert_summary( + section_idx=4, + visualization_count=3, + valid_mermaid_count=2, + planned_insertion_count=2, + inserted_mermaid_count=2, + wall_time_ms=80, + planner_mode="local_single_chart", + ) + + assert summary == { + "section_idx": 4, + "visualization_count": 3, + "valid_mermaid_count": 2, + "planned_insertion_count": 2, + "inserted_mermaid_count": 2, + "wall_time_ms": 80, + "planner_mode": "local_single_chart", + }