File size: 1,483 Bytes
38830c1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

from langgraph.graph import END, START, StateGraph

from graph.nodes import WorkflowNodes
from graph.routes import (
    route_after_analysis,
    route_after_planner,
    route_after_reflection,
    route_after_sql,
    route_after_visualization,
)
from graph.state import BIState


class BIWorkflow:
    def __init__(self, nodes: WorkflowNodes) -> None:
        self.nodes = nodes
        self.graph = self._build_graph()

    def _build_graph(self):
        graph = StateGraph(BIState)
        graph.add_node("planner", self.nodes.run_planner)
        graph.add_node("sql", self.nodes.run_sql)
        graph.add_node("analysis", self.nodes.run_analysis)
        graph.add_node("reflection", self.nodes.run_reflection)
        graph.add_node("visualization", self.nodes.run_visualization)

        graph.add_edge(START, "planner")
        graph.add_conditional_edges("planner", route_after_planner, {"sql": "sql", "analysis": "analysis"})
        graph.add_conditional_edges("sql", route_after_sql, {"analysis": "analysis"})
        graph.add_conditional_edges("analysis", route_after_analysis, {"reflection": "reflection"})
        graph.add_conditional_edges("reflection", route_after_reflection, {"visualization": "visualization"})
        graph.add_conditional_edges("visualization", route_after_visualization, {END: END})
        return graph.compile()

    def run(self, state: BIState) -> BIState:
        return self.graph.invoke(state)