defmodule ExAI.Graph.Validator do @moduledoc """ Graph definition validator. It validates graph wiring before execution and returns typed errors when invalid structure is detected. """ alias ExAI.Error alias ExAI.Graph.Definition alias ExAI.Graph.Definition.Node alias ExAI.Types @typedoc "A single validation issue." @type issue :: %{ required(:code) => atom(), required(:message) => String.t(), optional(:node_id) => Types.node_id(), optional(:edge) => map(), optional(:kind) => atom(), optional(:count) => non_neg_integer() } @spec validate(Definition.t()) :: :ok | {:error, Error.t()} def validate(%Definition{} = graph) do issues = [] |> append(duplicate_registration_issues(graph)) |> append(node_conflict_issues(graph)) |> append(edge_reference_issues(graph)) |> append(start_end_edge_issues(graph)) |> append(node_edge_presence_issues(graph)) |> append(unreachable_node_issues(graph)) |> append(join_reducer_wiring_issues(graph)) case issues do [] -> :ok _ -> {:error, Error.new(:invalid_graph, "graph validation failed", details: issues, context: %{issue_count: length(issues)} )} end end @spec validate!(Definition.t()) :: Definition.t() def validate!(%Definition{} = graph) do case validate(graph) do :ok -> graph {:error, %Error{} = error} -> raise error end end @spec duplicate_registration_issues(Definition.t()) :: [issue()] defp duplicate_registration_issues(%Definition{} = graph) do duplicate_issues = graph.registrations |> Enum.filter(fn {:step, _} -> true {:decision, _} -> true _ -> false end) |> Enum.group_by(fn {:step, id} -> {:step, id} {:decision, id} -> {:decision, id} end) |> Enum.flat_map(fn {{kind, id}, entries} -> if length(entries) > 1 do [ %{ code: :duplicate_node_registration, message: "#{kind} #{inspect(id)} registered multiple times", node_id: id, kind: kind, count: length(entries) } ] else [] end end) step_ids = MapSet.new(Map.keys(graph.steps)) decision_ids = MapSet.new(Map.keys(graph.decisions)) overlap_issues = step_ids |> MapSet.intersection(decision_ids) |> Enum.map(fn id -> %{ code: :duplicate_node_id, message: "node id #{inspect(id)} is registered as both step and decision", node_id: id } end) duplicate_issues ++ overlap_issues end @spec node_conflict_issues(Definition.t()) :: [issue()] defp node_conflict_issues(%Definition{} = graph) do Definition.node_ids(graph) |> Enum.flat_map(fn id -> case Definition.node(graph, id) do %Node{id: ^id} -> [] _ -> [%{code: :invalid_node_entry, message: "invalid node entry for #{inspect(id)}"}] end end) end @spec edge_reference_issues(Definition.t()) :: [issue()] defp edge_reference_issues(%Definition{} = graph) do known_nodes = MapSet.new(Definition.node_ids(graph)) graph.edges |> Enum.flat_map(fn edge -> [] |> maybe_add_invalid_direct_terminal_edge(edge) |> maybe_add_invalid_start_position(edge) |> maybe_add_invalid_end_position(edge) |> maybe_add_unknown_edge_endpoint(edge, :from, known_nodes) |> maybe_add_unknown_edge_endpoint(edge, :to, known_nodes) end) end @spec maybe_add_invalid_direct_terminal_edge([issue()], Definition.Edge.t()) :: [issue()] defp maybe_add_invalid_direct_terminal_edge(issues, %{from: :start, to: :end}) do issues ++ [ %{ code: :invalid_edge_endpoint, message: "graph cannot transition directly from :start to :end", edge: %{from: :start, to: :end} } ] end defp maybe_add_invalid_direct_terminal_edge(issues, _edge), do: issues @spec maybe_add_invalid_start_position([issue()], Definition.Edge.t()) :: [issue()] defp maybe_add_invalid_start_position(issues, %{to: :start, from: from}) do issues ++ [ %{ code: :invalid_edge_endpoint, message: ":start may only appear as an edge source", edge: %{from: from, to: :start} } ] end defp maybe_add_invalid_start_position(issues, _edge), do: issues @spec maybe_add_invalid_end_position([issue()], Definition.Edge.t()) :: [issue()] defp maybe_add_invalid_end_position(issues, %{from: :end, to: to}) do issues ++ [ %{ code: :invalid_edge_endpoint, message: ":end may only appear as an edge destination", edge: %{from: :end, to: to} } ] end defp maybe_add_invalid_end_position(issues, _edge), do: issues @spec maybe_add_unknown_edge_endpoint( [issue()], Definition.Edge.t(), :from | :to, MapSet.t(Types.node_id()) ) :: [issue()] defp maybe_add_unknown_edge_endpoint(issues, edge, endpoint, known_nodes) do value = Map.fetch!(edge, endpoint) cond do value in [:start, :end] -> issues MapSet.member?(known_nodes, value) -> issues true -> issues ++ [ %{ code: :unknown_edge_node, message: "edge references unknown #{endpoint} node #{inspect(value)}", node_id: value, edge: %{from: edge.from, to: edge.to} } ] end end @spec start_end_edge_issues(Definition.t()) :: [issue()] defp start_end_edge_issues(%Definition{} = graph) do [] |> maybe_add_missing_start_edge(graph.edges) |> maybe_add_missing_end_edge(graph.edges) end @spec maybe_add_missing_start_edge([issue()], [Definition.Edge.t()]) :: [issue()] defp maybe_add_missing_start_edge(issues, edges) do if Enum.any?(edges, &(&1.from == :start)) do issues else issues ++ [ %{ code: :missing_start_edge, message: "graph must include at least one edge from :start" } ] end end @spec maybe_add_missing_end_edge([issue()], [Definition.Edge.t()]) :: [issue()] defp maybe_add_missing_end_edge(issues, edges) do if Enum.any?(edges, &(&1.to == :end)) do issues else issues ++ [%{code: :missing_end_edge, message: "graph must include at least one edge to :end"}] end end @spec node_edge_presence_issues(Definition.t()) :: [issue()] defp node_edge_presence_issues(%Definition{} = graph) do node_ids = Definition.node_ids(graph) Enum.flat_map(node_ids, fn node_id -> incoming = Enum.count(graph.edges, &(&1.to == node_id)) outgoing = Enum.count(graph.edges, &(&1.from == node_id)) [] |> maybe_add_missing_incoming(node_id, incoming) |> maybe_add_missing_outgoing(node_id, outgoing) end) end @spec maybe_add_missing_incoming([issue()], Types.node_id(), non_neg_integer()) :: [issue()] defp maybe_add_missing_incoming(issues, _node_id, incoming) when incoming > 0, do: issues defp maybe_add_missing_incoming(issues, node_id, _incoming) do issues ++ [ %{ code: :missing_incoming_edge, message: "node #{inspect(node_id)} has no incoming edge", node_id: node_id } ] end @spec maybe_add_missing_outgoing([issue()], Types.node_id(), non_neg_integer()) :: [issue()] defp maybe_add_missing_outgoing(issues, _node_id, outgoing) when outgoing > 0, do: issues defp maybe_add_missing_outgoing(issues, node_id, _outgoing) do issues ++ [ %{ code: :missing_outgoing_edge, message: "node #{inspect(node_id)} has no outgoing edge", node_id: node_id } ] end @spec unreachable_node_issues(Definition.t()) :: [issue()] defp unreachable_node_issues(%Definition{} = graph) do adjacency = Enum.group_by(graph.edges, & &1.from, & &1.to) reachable_nodes = adjacency |> traverse_from_start() |> MapSet.delete(:start) |> MapSet.delete(:end) graph |> Definition.node_ids() |> Enum.reject(&MapSet.member?(reachable_nodes, &1)) |> Enum.map(fn node_id -> %{ code: :unreachable_node, message: "node #{inspect(node_id)} is unreachable from :start", node_id: node_id } end) end @spec traverse_from_start(%{ optional(Types.node_id() | :start | :end) => [Types.node_id() | :start | :end] }) :: MapSet.t(Types.node_id() | :start | :end) defp traverse_from_start(adjacency) do do_traverse(adjacency, [:start], MapSet.new([:start])) end @spec do_traverse(map(), [Types.node_id() | :start | :end], MapSet.t()) :: MapSet.t() defp do_traverse(_adjacency, [], visited), do: visited defp do_traverse(adjacency, [current | rest], visited) do neighbors = Map.get(adjacency, current, []) {queue_additions, next_visited} = Enum.reduce(neighbors, {[], visited}, fn node, {to_enqueue, acc_visited} -> if MapSet.member?(acc_visited, node) do {to_enqueue, acc_visited} else {[node | to_enqueue], MapSet.put(acc_visited, node)} end end) do_traverse(adjacency, rest ++ Enum.reverse(queue_additions), next_visited) end @spec join_reducer_wiring_issues(Definition.t()) :: [issue()] defp join_reducer_wiring_issues(%Definition{} = graph) do graph |> Definition.all_nodes() |> Map.values() |> Enum.flat_map(fn node -> join? = join_node?(node) reducer = reducer_ref(node) incoming = Enum.count(graph.edges, &(&1.to == node.id)) cond do join? and is_nil(reducer) -> [ %{ code: :join_missing_reducer, message: "join node #{inspect(node.id)} is missing reducer", node_id: node.id } ] not join? and not is_nil(reducer) -> [ %{ code: :reducer_on_non_join, message: "node #{inspect(node.id)} defines reducer but is not marked as join", node_id: node.id } ] join? and incoming < 2 -> [ %{ code: :join_requires_multiple_inputs, message: "join node #{inspect(node.id)} must have at least two incoming edges", node_id: node.id } ] true -> [] end end) end @spec join_node?(Node.t()) :: boolean() defp join_node?(%Node{} = node) do node.metadata[:join] == true or node.opts[:join] == true or node.metadata[:type] == :join or node.opts[:type] == :join end @spec reducer_ref(Node.t()) :: term() | nil defp reducer_ref(%Node{} = node) do node.opts[:reducer] || node.metadata[:reducer] end @spec append([issue()], [issue()]) :: [issue()] defp append(issues, new_issues), do: issues ++ new_issues end