| import os |
| import sys |
| sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '../..'))) |
|
|
| from src.core.graph import medical_pipeline |
| from src.utils.logger import setup_logger |
|
|
| logger = setup_logger("Visualizer") |
|
|
| def generate_graphviz_dot(): |
| """Dynamically reads the medical_pipeline graph and returns a DOT string for Graphviz.""" |
| try: |
| |
| graph = medical_pipeline.get_graph() |
| |
| |
| dot = ["digraph G {"] |
| dot.append(' node [shape=box, style="filled, rounded", fontname="Arial", fillcolor="#f0f2f6", color="#b0b2b6"];') |
| dot.append(' edge [fontname="Arial", color="#505256"];') |
| dot.append(' rankdir=TB;') |
| |
| |
| for node_id, node in graph.nodes.items(): |
| label = node_id.replace("_", " ").title() |
| |
| color = "#e1f5fe" if "classifier" in node_id else "#f1f8e9" |
| if "validator" in node_id or "safety" in node_id: |
| color = "#fff3e0" |
| if "tools" in node_id or "search" in node_id: |
| color = "#f3e5f5" |
| dot.append(f' "{node_id}" [label="{label}", fillcolor="{color}"];') |
| |
| |
| for edge in graph.edges: |
| source = edge.source |
| target = edge.target |
| label = edge.data if edge.data else "" |
| |
| |
| if label: |
| dot.append(f' "{source}" -> "{target}" [label="{label}"];') |
| else: |
| dot.append(f' "{source}" -> "{target}";') |
| |
| dot.append("}") |
| return "\n".join(dot) |
| except Exception as e: |
| return f'digraph G {{ "Error" [label="Error generating graph: {str(e)}"]; }}' |
|
|
| def generate_pipeline_image(): |
| """Generates a PNG image of the LangGraph using its internal mermaid renderer.""" |
| try: |
| graph = medical_pipeline.get_graph() |
| |
| |
| node_colors = {} |
| for node_id in graph.nodes: |
| node_id_lower = str(node_id).lower() |
| if "classifier" in node_id_lower: |
| node_colors[node_id] = "#BBDEFB" |
| elif "validator" in node_id_lower or "safety" in node_id_lower: |
| node_colors[node_id] = "#FFE0B2" |
| elif "tools" in node_id_lower or "search" in node_id_lower: |
| node_colors[node_id] = "#E1BEE7" |
| elif "agent" in node_id_lower or "bot" in node_id_lower: |
| node_colors[node_id] = "#C8E6C9" |
| else: |
| node_colors[node_id] = "#F5F5F5" |
| |
| kwargs = {"node_colors": node_colors, "background_color": "#ffffff"} |
| |
| try: |
| from langchain_core.runnables.graph import CurveStyle |
| kwargs["curve_style"] = CurveStyle.BASIS |
| except ImportError: |
| pass |
| |
| try: |
| |
| return graph.draw_mermaid_png(**kwargs) |
| except TypeError: |
| try: |
| |
| kwargs.pop("background_color", None) |
| return graph.draw_mermaid_png(**kwargs) |
| except TypeError: |
| |
| return graph.draw_mermaid_png() |
| |
| except Exception as e: |
| logger.error(f"Error generating Mermaid image: {e}") |
| return None |
|
|
| def generate_graph_svg(): |
| """Alternative: Generates SVG bytes using Mermaid.""" |
| try: |
| |
| return medical_pipeline.get_graph().draw_mermaid_png() |
| except Exception as e: |
| return None |
|
|
| if __name__ == "__main__": |
| |
| image_data = generate_pipeline_image() |
| if image_data: |
| output_path = "medical_pipeline_graph.png" |
| with open(output_path, "wb") as f: |
| f.write(image_data) |
| |
| |
| assets_dir = os.path.join(os.path.dirname(__file__), '../../assets') |
| os.makedirs(assets_dir, exist_ok=True) |
| arch_path = os.path.join(assets_dir, 'assets/architecture.png') |
| with open(arch_path, "wb") as f: |
| f.write(image_data) |
| |
| logger.info(f"Success! Graph image saved to: {output_path} and {os.path.abspath(arch_path)}") |
| else: |
| |
| logger.warning("Mermaid failed. Printing DOT source instead.") |
| logger.info(generate_graphviz_dot()) |
|
|