Current section
Files
Jump to
Current section
Files
lib/ex_ai/graph/validator.ex
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