use std::collections::HashMap;
use std::path::PathBuf;
use ai_agents_observability::ObservabilityConfig;
use ai_agents_observability::ObservabilityReport;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::assertion::{Assertion, AssertionResultDetail};
use crate::evidence::TurnEvidence;
use crate::fixtures::FixturesConfig;
use crate::redaction::RedactedString;
use crate::reset::ResetOptions;
use crate::{EvalError, Result};
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct EvalSuite {
pub name: String,
#[serde(default)]
pub agent: Option<PathBuf>,
#[serde(default)]
pub settings: EvalSettings,
#[serde(default)]
pub observability: Option<ObservabilityConfig>,
#[serde(default)]
pub fixtures: FixturesConfig,
#[serde(default)]
pub scenarios: Vec<Scenario>,
}
impl EvalSuite {
pub fn validate(&self, cli_agent: Option<&PathBuf>) -> Result<()> {
if self.name.trim().is_empty() {
return Err(EvalError::Config(
"eval suite name must not be empty".into(),
));
}
if cli_agent.is_none() && self.agent.is_none() {
return Err(EvalError::Config(
"agent path is required in suite or CLI".into(),
));
}
if self.scenarios.is_empty() {
return Err(EvalError::Config(
"eval suite must contain at least one scenario".into(),
));
}
self.fixtures.validate()?;
if self.settings.max_concurrent == 0 {
return Err(EvalError::Config(
"settings.max_concurrent must be greater than zero".into(),
));
}
if self.settings.timeout_per_turn_ms == 0 {
return Err(EvalError::Config(
"settings.timeout_per_turn_ms must be greater than zero".into(),
));
}
if matches!(
self.settings.isolation,
IsolationMode::Suite | IsolationMode::None
) {
return Err(EvalError::Config(
"settings.isolation currently supports scenario or turn".into(),
));
}
if self.settings.parallel && self.settings.isolation != IsolationMode::Scenario {
return Err(EvalError::Config(
"settings.parallel currently requires isolation: scenario".into(),
));
}
if self.settings.parallel
&& self
.scenarios
.iter()
.any(|scenario| !scenario.env.is_empty())
{
return Err(EvalError::Config(
"scenario.env cannot be used with parallel execution".into(),
));
}
if self
.scenarios
.iter()
.any(|scenario| scenario.budget.max_cost_usd.is_some())
{
let cost = self
.observability
.as_ref()
.map(|observability| &observability.cost)
.ok_or_else(|| {
EvalError::Config(
"budget.max_cost_usd requires suite observability.cost pricing".into(),
)
})?;
if !cost.enabled {
return Err(EvalError::Config(
"budget.max_cost_usd requires observability.cost.enabled: true".into(),
));
}
if cost.pricing.is_empty() && cost.pricing_file.is_none() {
return Err(EvalError::Config(
"budget.max_cost_usd requires observability.cost.pricing or pricing_file"
.into(),
));
}
}
let mut ids = std::collections::HashSet::new();
for scenario in &self.scenarios {
if scenario.id.trim().is_empty() {
return Err(EvalError::Config("scenario id must not be empty".into()));
}
if !ids.insert(scenario.id.clone()) {
return Err(EvalError::Config(format!(
"duplicate scenario id: {}",
scenario.id
)));
}
if !scenario.skip.is_skipped() && scenario.turns.is_empty() && scenario.steps.is_empty()
{
return Err(EvalError::Config(format!(
"scenario '{}' must define turns or steps",
scenario.id
)));
}
scenario.budget.validate(&scenario.id)?;
for (turn_index, turn) in scenario.turns.iter().enumerate() {
validate_turn_assertion(turn, &scenario.id, &format!("turns[{turn_index}]"))?;
}
for (step_index, step) in scenario.steps.iter().enumerate() {
if let ScenarioStep::Run(run) = step {
for (turn_index, turn) in run.turns.iter().enumerate() {
validate_turn_assertion(
turn,
&scenario.id,
&format!("steps[{step_index}].run.turns[{turn_index}]"),
)?;
}
}
}
}
Ok(())
}
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct EvalSettings {
#[serde(default)]
pub temperature: Option<f32>,
#[serde(default)]
pub seed: Option<u64>,
#[serde(default = "default_turn_timeout")]
pub timeout_per_turn_ms: u64,
#[serde(default)]
pub timeout_per_scenario_ms: Option<u64>,
#[serde(default)]
pub retries: u32,
#[serde(default = "default_retry_delay")]
pub retry_delay_ms: u64,
#[serde(default)]
pub isolation: IsolationMode,
#[serde(default)]
pub parallel: bool,
#[serde(default = "default_max_concurrent")]
pub max_concurrent: usize,
#[serde(default)]
pub fail_fast: bool,
#[serde(default = "default_true")]
pub redact_outputs: bool,
}
impl Default for EvalSettings {
fn default() -> Self {
Self {
temperature: None,
seed: None,
timeout_per_turn_ms: default_turn_timeout(),
timeout_per_scenario_ms: None,
retries: 0,
retry_delay_ms: default_retry_delay(),
isolation: IsolationMode::Scenario,
parallel: false,
max_concurrent: default_max_concurrent(),
fail_fast: false,
redact_outputs: true,
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize, Default, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum IsolationMode {
Turn,
#[default]
Scenario,
Suite,
None,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Scenario {
pub id: String,
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub tags: Vec<String>,
#[serde(default)]
pub language: Option<String>,
#[serde(default)]
pub actor: Option<String>,
#[serde(default)]
pub context: Value,
#[serde(default)]
pub env: HashMap<String, String>,
#[serde(default)]
pub skip: SkipConfig,
#[serde(default)]
pub budget: ScenarioBudget,
#[serde(default)]
pub turns: Vec<Turn>,
#[serde(default)]
pub steps: Vec<ScenarioStep>,
}
#[derive(Debug, Clone, Deserialize, Default)]
#[serde(deny_unknown_fields)]
pub struct ScenarioBudget {
#[serde(default)]
pub max_llm_calls: Option<u64>,
#[serde(default)]
pub max_total_tokens: Option<u64>,
#[serde(default)]
pub max_cost_usd: Option<f64>,
}
fn validate_turn_assertion(turn: &Turn, scenario_id: &str, location: &str) -> Result<()> {
if let Some(assertion) = &turn.assertions {
assertion.validate(&format!("scenario '{scenario_id}' {location}.assert"))?;
}
Ok(())
}
impl ScenarioBudget {
pub(crate) fn is_configured(&self) -> bool {
self.max_llm_calls.is_some()
|| self.max_total_tokens.is_some()
|| self.max_cost_usd.is_some()
}
fn validate(&self, scenario_id: &str) -> Result<()> {
if self.max_llm_calls == Some(0) {
return Err(EvalError::Config(format!(
"scenario '{scenario_id}' budget.max_llm_calls must be greater than zero"
)));
}
if self.max_total_tokens == Some(0) {
return Err(EvalError::Config(format!(
"scenario '{scenario_id}' budget.max_total_tokens must be greater than zero"
)));
}
if let Some(max_cost_usd) = self.max_cost_usd
&& (!max_cost_usd.is_finite() || max_cost_usd <= 0.0)
{
return Err(EvalError::Config(format!(
"scenario '{scenario_id}' budget.max_cost_usd must be finite and greater than zero"
)));
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct Turn {
pub input: String,
pub actor: Option<String>,
pub context: Value,
pub stream: Option<bool>,
pub timeout_ms: Option<u64>,
pub assertions: Option<Assertion>,
}
const EXPECT_ERROR_CONTEXT_KEY: &str = "__ai_agents_eval_expect_error";
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct TurnDefinition {
input: String,
#[serde(default)]
actor: Option<String>,
#[serde(default)]
context: Value,
#[serde(default)]
stream: Option<bool>,
#[serde(default)]
timeout_ms: Option<u64>,
#[serde(default)]
expect_error: Option<ExpectedError>,
#[serde(default, rename = "assert")]
assertions: Option<Assertion>,
}
impl<'de> Deserialize<'de> for Turn {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let definition = TurnDefinition::deserialize(deserializer)?;
let mut context = definition.context;
if let Some(expect_error) = definition.expect_error {
let object = match &mut context {
Value::Object(object) => object,
_ => {
context = Value::Object(serde_json::Map::new());
context
.as_object_mut()
.expect("context was replaced with an object")
}
};
object.insert(
EXPECT_ERROR_CONTEXT_KEY.to_string(),
serde_json::to_value(expect_error).map_err(serde::de::Error::custom)?,
);
}
Ok(Self {
input: definition.input,
actor: definition.actor,
context,
stream: definition.stream,
timeout_ms: definition.timeout_ms,
assertions: definition.assertions,
})
}
}
pub(crate) fn turn_expected_error(turn: &Turn) -> Option<ExpectedError> {
turn.context
.get(EXPECT_ERROR_CONTEXT_KEY)
.cloned()
.and_then(|value| serde_json::from_value(value).ok())
}
pub(crate) fn turn_runtime_context(turn: &Turn) -> Value {
let mut context = turn.context.clone();
if let Value::Object(object) = &mut context {
object.remove(EXPECT_ERROR_CONTEXT_KEY);
}
context
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(untagged)]
pub enum ExpectedError {
One(String),
Any(Vec<String>),
}
impl ExpectedError {
pub fn matches(&self, error: &str) -> bool {
self.items().iter().any(|expected| error.contains(expected))
}
pub fn items(&self) -> Vec<&str> {
match self {
Self::One(value) => vec![value.as_str()],
Self::Any(values) => values.iter().map(String::as_str).collect(),
}
}
}
#[derive(Debug, Clone, Deserialize)]
#[serde(untagged)]
pub enum SkipConfig {
Bool(bool),
Reason(String),
}
impl Default for SkipConfig {
fn default() -> Self {
Self::Bool(false)
}
}
impl SkipConfig {
pub fn is_skipped(&self) -> bool {
match self {
Self::Bool(value) => *value,
Self::Reason(_) => true,
}
}
pub fn reason(&self) -> Option<String> {
match self {
Self::Bool(_) => None,
Self::Reason(reason) => Some(reason.clone()),
}
}
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "snake_case", deny_unknown_fields)]
pub enum ScenarioStep {
Run(RunStep),
ResetAgent(ResetStepConfig),
SaveSession(String),
LoadSession(String),
SetContext { values: Value },
SetActor { actor: String },
CleanupExpired,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct RunStep {
#[serde(default)]
pub turns: Vec<Turn>,
#[serde(default)]
pub save_session: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(untagged)]
pub enum ResetStepConfig {
Bool(bool),
Options(ResetOptions),
}
#[derive(Debug, Clone, Serialize)]
pub struct EvalResult {
pub schema_version: u32,
pub suite: String,
pub agent: String,
pub total: usize,
pub passed: usize,
pub failed: usize,
pub skipped: usize,
pub duration_ms: u64,
pub scenarios: Vec<ScenarioResult>,
pub metrics: crate::metrics::EvalMetrics,
#[serde(skip_serializing_if = "Option::is_none")]
pub observability: Option<ObservabilityReport>,
}
#[derive(Debug, Clone, Serialize)]
pub struct ScenarioResult {
pub id: String,
pub name: Option<String>,
pub tags: Vec<String>,
pub language: Option<String>,
pub status: ScenarioStatus,
pub failure_category: Option<FailureCategory>,
pub flaky: bool,
pub attempts: Vec<AttemptResult>,
pub duration_ms: u64,
pub retries_used: u32,
}
#[derive(Debug, Clone, Serialize)]
pub struct AttemptResult {
pub attempt: u32,
pub turns: Vec<TurnResult>,
pub status: ScenarioStatus,
pub duration_ms: u64,
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ScenarioStatus {
Passed,
Failed { reason: String },
Skipped { reason: Option<String> },
Error { message: String },
}
impl ScenarioStatus {
pub fn is_passed(&self) -> bool {
matches!(self, Self::Passed)
}
pub fn is_failed(&self) -> bool {
matches!(self, Self::Failed { .. })
}
pub fn is_error(&self) -> bool {
matches!(self, Self::Error { .. })
}
}
#[derive(Debug, Clone, Serialize, PartialEq, Eq, Hash)]
#[serde(rename_all = "snake_case")]
pub enum FailureCategory {
ConfigError,
RuntimeError,
AssertionFailed,
JudgeError,
FlakyPass,
}
#[derive(Debug, Clone, Serialize)]
pub struct TurnResult {
pub index: usize,
pub input: RedactedString,
pub response: RedactedString,
pub response_present: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub runtime_error: Option<RedactedString>,
pub state: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<Value>,
#[serde(skip_serializing)]
pub evidence: TurnEvidence,
pub assertion_results: Vec<AssertionResultDetail>,
pub latency_ms: u64,
#[serde(skip_serializing_if = "Option::is_none")]
pub observability_span_id: Option<String>,
}
fn default_turn_timeout() -> u64 {
30_000
}
fn default_retry_delay() -> u64 {
1_000
}
fn default_max_concurrent() -> usize {
4
}
fn default_true() -> bool {
true
}
#[cfg(test)]
mod tests {
use super::*;
fn suite() -> EvalSuite {
EvalSuite {
name: "suite".to_string(),
agent: Some(PathBuf::from("agent.yaml")),
settings: EvalSettings::default(),
observability: None,
fixtures: FixturesConfig::default(),
scenarios: vec![Scenario {
id: "scenario-1".to_string(),
name: None,
tags: vec!["smoke".to_string()],
language: Some("en".to_string()),
actor: None,
context: Value::Null,
env: HashMap::new(),
skip: SkipConfig::default(),
budget: ScenarioBudget::default(),
turns: vec![Turn {
input: "hello".to_string(),
actor: None,
context: Value::Null,
stream: None,
timeout_ms: None,
assertions: None,
}],
steps: Vec::new(),
}],
}
}
#[test]
fn validation_accepts_minimal_suite() {
assert!(suite().validate(None).is_ok());
}
#[test]
fn validation_rejects_empty_suite() {
let mut suite = suite();
suite.scenarios.clear();
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("at least one scenario"));
}
#[test]
fn validation_rejects_invalid_deferred_mock_routes() {
let mut suite = suite();
suite.fixtures.mock_server = Some(crate::fixtures::MockServerConfig {
enabled: true,
port: None,
routes: vec![serde_json::json!({
"method": "GET",
"path": "/ok",
"statuz": 200
})],
});
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("fixtures.mock_server.routes[0]"));
assert!(error.contains("statuz"));
}
#[test]
fn validation_rejects_duplicate_ids() {
let mut suite = suite();
suite.scenarios.push(suite.scenarios[0].clone());
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("duplicate scenario id"));
}
#[test]
fn yaml_rejects_unknown_suite_and_assertion_fields() {
let suite_error = serde_yaml::from_str::<EvalSuite>(
"name: strict\nagent: agent.yaml\nunknown_setting: true\nscenarios: []\n",
)
.unwrap_err()
.to_string();
assert!(suite_error.contains("unknown field `unknown_setting`"));
let assertion_error = serde_yaml::from_str::<EvalSuite>(
r#"
name: strict
agent: agent.yaml
scenarios:
- id: strict
turns:
- input: hello
assert:
response_contians: hello
"#,
)
.unwrap_err()
.to_string();
assert!(assertion_error.contains("response_contians"));
}
#[test]
fn validation_rejects_empty_assertion_trees() {
let mut suite = suite();
suite.scenarios[0].turns[0].assertions = Some(Assertion::default());
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("assert must not be empty"));
suite.scenarios[0].turns[0].assertions = Some(Assertion {
all: Some(Vec::new()),
..Default::default()
});
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("assert.all must contain at least one assertion"));
suite.scenarios[0].turns[0].assertions = Some(Assertion {
any: Some(vec![Assertion {
all: Some(vec![Assertion::default()]),
..Default::default()
}]),
..Default::default()
});
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("assert.any[0].all[0] must not be empty"));
}
#[test]
fn validation_rejects_empty_nested_assertion_collections() {
use crate::assertion::{
ApprovalAssertion, ApprovalAssertionObject, LlmRequestAssertion,
OrchestrationAssertion, PathAssertion, StringList, ToolCalledAssertion,
ToolCalledObject,
};
use crate::judge::JudgeAssertion;
let cases = [
(
Assertion {
response_contains: Some(StringList::Many(Vec::new())),
..Default::default()
},
"response_contains",
),
(
Assertion {
state_in: Some(Vec::new()),
..Default::default()
},
"state_in",
),
(
Assertion {
llm_request: Some(LlmRequestAssertion {
system_contains: Some(StringList::Many(Vec::new())),
..Default::default()
}),
..Default::default()
},
"llm_request.system_contains",
),
(
Assertion {
approval_requested: Some(ApprovalAssertion::Object(ApprovalAssertionObject {
message_contains: Some(StringList::Many(Vec::new())),
..Default::default()
})),
..Default::default()
},
"approval_requested.message_contains",
),
(
Assertion {
tool_called: Some(ToolCalledAssertion::Object(ToolCalledObject {
source_in: Some(Vec::new()),
..Default::default()
})),
..Default::default()
},
"tool_called.source_in",
),
(
Assertion {
metadata_path: Some(PathAssertion {
path: "value".to_string(),
in_values: Some(Vec::new()),
..Default::default()
}),
..Default::default()
},
"metadata_path.in",
),
(
Assertion {
orchestration: Some(OrchestrationAssertion {
agents_include: Some(Vec::new()),
..Default::default()
}),
..Default::default()
},
"orchestration.agents_include",
),
(
Assertion {
metadata_contains: Some(HashMap::new()),
..Default::default()
},
"metadata_contains",
),
(
Assertion {
judge: Some(JudgeAssertion {
llm: None,
pass_threshold: 0.75,
criteria: Vec::new(),
}),
..Default::default()
},
"judge.criteria",
),
];
for (assertion, expected) in cases {
let mut suite = suite();
suite.scenarios[0].turns[0].assertions = Some(assertion);
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains(expected), "{error}");
}
}
#[test]
fn yaml_rejects_explicit_empty_observability_collections() {
let error = serde_yaml::from_str::<EvalSuite>(
r#"
name: strict
agent: agent.yaml
scenarios:
- id: strict
turns:
- input: hello
assert:
observability:
dimension_counts: []
"#,
)
.unwrap_err()
.to_string();
assert!(error.contains("assertion collection must contain at least one value"));
}
#[test]
fn expected_error_accepts_string_and_list() {
let one: ExpectedError = serde_yaml::from_str("timeout").unwrap();
let many: ExpectedError = serde_yaml::from_str("[timeout, unavailable]").unwrap();
assert!(one.matches("turn timeout"));
assert!(many.matches("service unavailable"));
assert!(!many.matches("permission denied"));
let turn: Turn =
serde_yaml::from_str("input: hello\nexpect_error: [timeout, unavailable]").unwrap();
assert!(turn_expected_error(&turn).unwrap().matches("turn timeout"));
assert_eq!(
turn_runtime_context(&turn),
Value::Object(serde_json::Map::new())
);
}
#[test]
fn validation_rejects_parallel_env() {
let mut suite = suite();
suite.settings.parallel = true;
suite.scenarios[0]
.env
.insert("TOKEN".to_string(), "secret".to_string());
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("scenario.env"));
}
#[test]
fn validation_rejects_invalid_scenario_budgets() {
let mut suite = suite();
suite.scenarios[0].budget.max_llm_calls = Some(0);
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("budget.max_llm_calls"));
suite.scenarios[0].budget.max_llm_calls = Some(1);
suite.scenarios[0].budget.max_total_tokens = Some(0);
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("budget.max_total_tokens"));
suite.scenarios[0].budget.max_total_tokens = Some(1);
suite.scenarios[0].budget.max_cost_usd = Some(f64::NAN);
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("budget.max_cost_usd"));
}
#[test]
fn validation_requires_pricing_for_cost_budgets() {
let mut suite = suite();
suite.scenarios[0].budget.max_cost_usd = Some(0.01);
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("observability.cost pricing"));
suite.observability = Some(ObservabilityConfig::default());
let error = suite.validate(None).unwrap_err().to_string();
assert!(error.contains("pricing or pricing_file"));
suite.observability.as_mut().unwrap().cost.pricing_file = Some("pricing.yaml".to_string());
assert!(suite.validate(None).is_ok());
}
}