use crate::on_error::ErrorPolicy;
use serde::de::Error as _;
use serde::Deserialize;
use super::types::{Expected, TestCase, TestInput};
const LEGACY_EXPECTED_KEYS: [&str; 4] = ["$ref", "must_contain", "sequence", "schema"];
#[derive(Default)]
enum RawExpected {
#[default]
Missing,
Present(serde_json::Value),
}
impl<'de> Deserialize<'de> for RawExpected {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
serde_json::Value::deserialize(deserializer).map(Self::Present)
}
}
fn value_kind(v: &serde_json::Value) -> &'static str {
match v {
serde_json::Value::Null => "null",
serde_json::Value::Bool(_) => "a boolean",
serde_json::Value::Number(_) => "a number",
serde_json::Value::String(_) => "a string",
serde_json::Value::Array(_) => "a list",
serde_json::Value::Object(_) => "a mapping",
}
}
pub(crate) fn parse_expected_entry(item: &serde_json::Value) -> Result<Expected, String> {
let strict_err = match serde_json::from_value::<Expected>(item.clone()) {
Ok(exp) => return reject_for_parse(exp),
Err(e) => e,
};
let Some(obj) = item.as_object() else {
return Err(format!(
"`expected:` must be a mapping (or a list of one mapping), found {}",
value_kind(item)
));
};
let matched_keys: Vec<&str> = LEGACY_EXPECTED_KEYS
.iter()
.copied()
.filter(|key| obj.contains_key(*key))
.collect();
if let Some(tag) = obj.get("type") {
let compatibility_key = match tag.as_str() {
Some("must_contain") => "must_contain",
Some("sequence") => "sequence",
_ => return Err(format!("invalid `expected:` block: {}", strict_err)),
};
if matched_keys.len() > 1 {
return Err(format!(
"ambiguous legacy `expected:` block contains multiple assertions {:?}; \
use one tagged assertion or move additional checks to `assertions:`",
matched_keys
));
}
let only_compatibility_key = matched_keys.as_slice() == [compatibility_key];
let has_unknown_key = obj
.keys()
.any(|key| key != "type" && key != compatibility_key);
if !only_compatibility_key || has_unknown_key {
return Err(format!("invalid `expected:` block: {}", strict_err));
}
} else {
let unknown_keys: Vec<&str> = obj
.keys()
.map(String::as_str)
.filter(|key| !LEGACY_EXPECTED_KEYS.contains(key))
.collect();
if !unknown_keys.is_empty() {
return Err(format!(
"unrecognized legacy `expected:` key(s) {:?}; supported keys are {:?}",
unknown_keys, LEGACY_EXPECTED_KEYS
));
}
}
let mut parsed = None;
if let Some(r) = obj.get("$ref") {
let path = r
.as_str()
.ok_or_else(|| format!("`$ref` must be a string, found {}", value_kind(r)))?;
parsed = Some(Expected::Reference {
path: path.to_string(),
});
}
if let Some(mc) = obj.get("must_contain") {
let val: Vec<String> = if let Some(s) = mc.as_str() {
vec![s.to_string()]
} else {
serde_json::from_value(mc.clone()).map_err(|e| {
format!(
"`must_contain` must be a string or a list of strings, found {}: {}",
value_kind(mc),
e
)
})?
};
if parsed.is_none() {
parsed = Some(Expected::MustContain { must_contain: val });
}
}
if let Some(seq) = obj.get("sequence") {
if parsed.is_none() {
let sequence: Vec<String> = serde_json::from_value(seq.clone()).map_err(|e| {
format!(
"`sequence` must be a list of strings, found {}: {}",
value_kind(seq),
e
)
})?;
parsed = Some(Expected::SequenceValid {
policy: None,
sequence: Some(sequence),
rules: None,
});
}
}
if obj.get("schema").is_some() && parsed.is_none() {
parsed = Some(Expected::ArgsValid {
policy: None,
schema: obj.get("schema").cloned(),
});
}
if matched_keys.len() > 1 {
return Err(format!(
"ambiguous legacy `expected:` block contains multiple assertions {:?}; \
use one tagged assertion or move additional checks to `assertions:`",
matched_keys
));
}
if let Some(p) = parsed {
return reject_for_parse(p);
}
if obj.contains_key("type") {
return Err(format!("invalid `expected:` block: {}", strict_err));
}
let found: Vec<&str> = obj.keys().map(String::as_str).collect();
Err(format!(
"unrecognized `expected:` block, found key(s) {:?}. Use the tagged form \
(e.g. `type: must_contain` with `must_contain: [...]`) or one of the legacy \
keys {:?}",
found, LEGACY_EXPECTED_KEYS
))
}
fn reject_vacuous(exp: Expected) -> Result<Expected, String> {
let Some(field) = super::validation::vacuous_expected_field(&exp) else {
return Ok(exp);
};
Err(format!(
"`{}` asserts nothing, so this test would pass for any response. \
Give it at least one entry, or remove the `expected:` block and put the \
test's checks in `assertions:`.",
field
))
}
fn reject_for_parse(exp: Expected) -> Result<Expected, String> {
let exp = reject_vacuous(exp)?;
if let Some(reason) = super::validation::non_executable_expected_reason(&exp) {
return Err(format!("expected block is not executable: {reason}"));
}
if let Some(reason) = super::validation::ineffective_expected_reason(&exp) {
return Err(reason.to_string());
}
Ok(exp)
}
fn parse_expected_value(test_id: &str, val: &serde_json::Value) -> Result<Expected, String> {
let Some(arr) = val.as_array() else {
return parse_expected_entry(val).map_err(|e| format!("test '{}': {}", test_id, e));
};
match arr.len() {
0 => Err(format!(
"test '{}': `expected:` is an empty list, which asserts nothing. \
Remove the key or give it an assertion.",
test_id
)),
1 => parse_expected_entry(&arr[0])
.map_err(|e| format!("test '{}': `expected:` entry 0 is invalid: {}", test_id, e)),
n => Err(format!(
"test '{}': `expected:` has {} entries but only one is supported \
(earlier versions silently dropped all but the first). \
Split them into separate tests, or move the extra checks to `assertions:`.",
test_id, n
)),
}
}
impl<'de> Deserialize<'de> for TestCase {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct RawTestCase {
id: String,
input: TestInput,
#[serde(default)]
expected: RawExpected,
assertions: Option<Vec<crate::agent_assertions::model::TraceAssertion>>,
#[serde(default)]
on_error: Option<ErrorPolicy>,
#[serde(default)]
tags: Vec<String>,
metadata: Option<serde_json::Value>,
}
let raw = RawTestCase::deserialize(deserializer)?;
let extra_assertions = raw.assertions.unwrap_or_default();
let expected_main = match &raw.expected {
RawExpected::Present(val) => {
parse_expected_value(&raw.id, val).map_err(D::Error::custom)?
}
RawExpected::Missing => Expected::default(),
};
Ok(TestCase {
id: raw.id,
input: raw.input,
expected: expected_main,
assertions: if extra_assertions.is_empty() {
None
} else {
Some(extra_assertions)
},
on_error: raw.on_error,
tags: raw.tags,
metadata: raw.metadata,
})
}
}
impl<'de> Deserialize<'de> for TestInput {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
struct TestInputVisitor;
impl<'de> serde::de::Visitor<'de> for TestInputVisitor {
type Value = TestInput;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("string or struct TestInput")
}
fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
Ok(TestInput {
prompt: value.to_owned(),
context: None,
})
}
fn visit_map<A>(self, map: A) -> Result<Self::Value, A::Error>
where
A: serde::de::MapAccess<'de>,
{
#[derive(Deserialize)]
struct Helper {
prompt: String,
#[serde(default)]
context: Option<Vec<String>>,
}
let helper =
Helper::deserialize(serde::de::value::MapAccessDeserializer::new(map))?;
Ok(TestInput {
prompt: helper.prompt,
context: helper.context,
})
}
}
deserializer.deserialize_any(TestInputVisitor)
}
}