Current section

Files

Jump to
langchain lib trajectory assertions.ex
Raw

lib/trajectory/assertions.ex

defmodule LangChain.Trajectory.Assertions do
@moduledoc """
ExUnit assertion helpers for trajectory comparison.
## Usage
use LangChain.Trajectory.Assertions
test "agent calls the right tools" do
trajectory = Trajectory.from_chain(chain)
assert_trajectory trajectory, [
%{name: "search", arguments: %{"query" => "weather"}},
%{name: "get_forecast", arguments: nil}
]
end
test "agent does not call dangerous tool" do
trajectory = Trajectory.from_chain(chain)
refute_trajectory trajectory, [
%{name: "delete_all", arguments: nil}
], mode: :superset
end
"""
@doc """
Assert that a trajectory matches the expected tool call sequence.
Accepts the same options as `LangChain.Trajectory.matches?/3`:
* `:mode` — `:strict` (default), `:unordered`, `:superset`
* `:args` — `:exact` (default), `:subset`
On failure, raises `ExUnit.AssertionError` with a diff showing expected vs
actual tool calls.
"""
defmacro assert_trajectory(actual, expected, opts \\ []) do
quote do
actual_val = unquote(actual)
expected_val = unquote(expected)
opts_val = unquote(opts)
unless LangChain.Trajectory.matches?(actual_val, expected_val, opts_val) do
actual_calls = LangChain.Trajectory.Assertions.extract_tool_calls(actual_val)
expected_calls = LangChain.Trajectory.Assertions.extract_tool_calls(expected_val)
raise ExUnit.AssertionError,
left: actual_calls,
right: expected_calls,
message:
"Trajectory mismatch (mode: #{Keyword.get(opts_val, :mode, :strict)}, args: #{Keyword.get(opts_val, :args, :exact)})"
end
end
end
@doc """
Assert that a trajectory does NOT match the expected tool call sequence.
Useful for verifying that specific tools were not called or that a
particular call pattern did not occur.
Accepts the same options as `assert_trajectory/3`.
On failure, raises `ExUnit.AssertionError` indicating an unexpected match.
"""
defmacro refute_trajectory(actual, expected, opts \\ []) do
quote do
actual_val = unquote(actual)
expected_val = unquote(expected)
opts_val = unquote(opts)
if LangChain.Trajectory.matches?(actual_val, expected_val, opts_val) do
matched_calls = LangChain.Trajectory.Assertions.extract_tool_calls(expected_val)
actual_calls = LangChain.Trajectory.Assertions.extract_tool_calls(actual_val)
raise ExUnit.AssertionError,
left: actual_calls,
right: matched_calls,
message:
"Unexpected trajectory match (mode: #{Keyword.get(opts_val, :mode, :strict)}, args: #{Keyword.get(opts_val, :args, :exact)})"
end
end
end
@doc """
Assert that tool `name_a` was called before tool `name_b`.
Wraps `LangChain.Trajectory.called_before?/4` and accepts the same options:
* `:require_both` — when `true`, a missing tool raises rather than failing
the ordering check (see `LangChain.Trajectory.called_before?/4`)
## Examples
assert_called_before trajectory, "search", "answer"
assert_called_before trajectory, "search", "answer", require_both: true
"""
defmacro assert_called_before(actual, name_a, name_b, opts \\ []) do
quote do
actual_val = unquote(actual)
name_a_val = unquote(name_a)
name_b_val = unquote(name_b)
opts_val = unquote(opts)
unless LangChain.Trajectory.called_before?(actual_val, name_a_val, name_b_val, opts_val) do
actual_calls = LangChain.Trajectory.Assertions.extract_tool_calls(actual_val)
raise ExUnit.AssertionError,
left: actual_calls,
right: [name_a_val, name_b_val],
message: "Expected #{inspect(name_a_val)} to be called before #{inspect(name_b_val)}"
end
end
end
@doc """
Assert that tool `name_a` was NOT called before tool `name_b`.
Accepts the same options as `assert_called_before/4`. Note that with the
default options a missing tool makes this pass vacuously; pass
`require_both: true` to surface a missing tool as an error instead.
## Examples
refute_called_before trajectory, "write_file", "read_file"
"""
defmacro refute_called_before(actual, name_a, name_b, opts \\ []) do
quote do
actual_val = unquote(actual)
name_a_val = unquote(name_a)
name_b_val = unquote(name_b)
opts_val = unquote(opts)
if LangChain.Trajectory.called_before?(actual_val, name_a_val, name_b_val, opts_val) do
actual_calls = LangChain.Trajectory.Assertions.extract_tool_calls(actual_val)
raise ExUnit.AssertionError,
left: actual_calls,
right: [name_a_val, name_b_val],
message:
"Did not expect #{inspect(name_a_val)} to be called before #{inspect(name_b_val)}"
end
end
end
@doc false
def extract_tool_calls(%LangChain.Chains.LLMChain{} = chain) do
chain |> LangChain.Trajectory.from_chain() |> extract_tool_calls()
end
def extract_tool_calls(%LangChain.Trajectory{tool_calls: calls}), do: calls
def extract_tool_calls(calls) when is_list(calls), do: calls
defmacro __using__(_opts) do
quote do
import LangChain.Trajectory.Assertions
end
end
end