1use std::collections::BTreeMap;
12
13use serde::{Deserialize, Serialize};
14use serde_json::Value;
15
16use gemini_adk_rs::flow::{Enforcement, FlowMonitor};
17use gemini_adk_rs::state::State;
18
19use super::SessionSpec;
20
21#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
23#[serde(rename_all = "snake_case")]
24pub enum SimEvent {
25 User(String),
28 Tool(String),
33 Set(BTreeMap<String, Value>),
36 Expect(TestExpectation),
38}
39
40#[derive(Debug, Clone, Default, Serialize, Deserialize, schemars::JsonSchema)]
43pub struct TestExpectation {
44 #[serde(default, skip_serializing_if = "Vec::is_empty")]
46 pub done: Vec<String>,
47 #[serde(default, skip_serializing_if = "Vec::is_empty")]
49 pub active: Vec<String>,
50 #[serde(default, skip_serializing_if = "Vec::is_empty")]
52 pub allowed: Vec<String>,
53 #[serde(default, skip_serializing_if = "Vec::is_empty")]
55 pub blocked: Vec<String>,
56 #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
58 pub state: BTreeMap<String, Value>,
59 #[serde(default, skip_serializing_if = "Option::is_none")]
61 pub complete: Option<bool>,
62}
63
64#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
66pub struct SpecTest {
67 pub name: String,
69 pub script: Vec<SimEvent>,
71}
72
73#[derive(Debug, Clone, Serialize)]
75pub struct TestStepResult {
76 pub index: usize,
78 pub event: String,
80 pub failures: Vec<String>,
83}
84
85#[derive(Debug, Clone, Serialize)]
87pub struct TestReport {
88 pub name: String,
90 pub passed: bool,
92 pub failures: Vec<TestStepResult>,
94 pub events: usize,
96}
97
98#[derive(Debug, Clone, Serialize)]
102pub struct SimSnapshot {
103 pub index: usize,
105 pub event: String,
107 #[serde(default, skip_serializing_if = "Vec::is_empty")]
109 pub failures: Vec<String>,
110 pub done: Vec<String>,
112 pub complete: bool,
114 #[serde(flatten)]
117 pub explanation: gemini_adk_rs::flow::FlowExplanation,
118}
119
120pub fn trace_test(spec: &SessionSpec, test_name: &str) -> Result<Vec<SimSnapshot>, Vec<String>> {
124 let flow = spec.effective_flow()?;
125 let test = spec
126 .tests
127 .iter()
128 .find(|t| t.name == test_name)
129 .ok_or_else(|| vec![format!("no test named '{test_name}' in the spec")])?;
130 Ok(replay(spec, flow, test))
131}
132
133fn replay(
136 spec: &SessionSpec,
137 flow: gemini_adk_rs::flow::Flow,
138 test: &SpecTest,
139) -> Vec<SimSnapshot> {
140 let state = State::new();
141 spec.seed_state_defaults(&state);
145 spec.recompute_computed(&state);
146 let mut monitor = FlowMonitor::new(flow, Enforcement::Enforce);
147 monitor.relatch(&state);
148
149 let snapshot = |index: usize,
150 event: String,
151 failures: Vec<String>,
152 monitor: &FlowMonitor,
153 state: &State| SimSnapshot {
154 index,
155 event,
156 failures,
157 done: monitor.marking().done.iter().cloned().collect(),
158 complete: monitor.is_complete(),
159 explanation: monitor.explain(state),
160 };
161
162 let mut snapshots = vec![snapshot(0, "start".into(), Vec::new(), &monitor, &state)];
163 for (index, event) in test.script.iter().enumerate() {
164 let mut failures = Vec::new();
165 let label = match event {
166 SimEvent::User(text) => {
167 spec.recompute_computed(&state);
168 monitor.on_turn(&state);
169 format!("user: {text}")
170 }
171 SimEvent::Tool(name) => {
172 match monitor.admits_tool(name, &state) {
173 Ok(()) => {
174 spec.apply_tool_state(name, &state);
175 spec.recompute_computed(&state);
176 monitor.on_tool_ok(name, &state);
177 }
178 Err(reason) => {
179 let anticipated = matches!(
180 test.script.get(index + 1),
181 Some(SimEvent::Expect(e)) if e.blocked.iter().any(|t| t == name)
182 );
183 if !anticipated {
184 failures.push(format!("tool '{name}' was blocked: {reason}"));
185 }
186 }
187 }
188 format!("tool: {name}")
189 }
190 SimEvent::Set(map) => {
191 for (key, value) in map {
192 let _ = state.set(key, value.clone());
193 }
194 spec.recompute_computed(&state);
195 monitor.relatch(&state);
196 format!(
197 "set: {}",
198 map.keys().cloned().collect::<Vec<_>>().join(", ")
199 )
200 }
201 SimEvent::Expect(expect) => {
202 check(expect, &monitor, &state, &mut failures);
203 "expect".to_string()
204 }
205 };
206 snapshots.push(snapshot(index + 1, label, failures, &monitor, &state));
207 }
208 snapshots
209}
210
211pub(crate) fn run_tests(spec: &SessionSpec) -> Vec<TestReport> {
213 let flow = match spec.effective_flow() {
214 Ok(flow) => flow,
215 Err(errors) => {
216 return spec
217 .tests
218 .iter()
219 .map(|t| TestReport {
220 name: t.name.clone(),
221 passed: false,
222 failures: vec![TestStepResult {
223 index: 0,
224 event: "setup".into(),
225 failures: errors.clone(),
226 }],
227 events: 0,
228 })
229 .collect();
230 }
231 };
232
233 spec.tests
234 .iter()
235 .map(|test| run_one(spec, flow.clone(), test))
236 .collect()
237}
238
239fn run_one(spec: &SessionSpec, flow: gemini_adk_rs::flow::Flow, test: &SpecTest) -> TestReport {
240 let failures: Vec<TestStepResult> = replay(spec, flow, test)
241 .into_iter()
242 .skip(1) .filter(|s| !s.failures.is_empty())
244 .map(|s| TestStepResult {
245 index: s.index - 1,
246 event: s.event,
247 failures: s.failures,
248 })
249 .collect();
250
251 TestReport {
252 name: test.name.clone(),
253 passed: failures.is_empty(),
254 failures,
255 events: test.script.len(),
256 }
257}
258
259fn check(
260 expect: &TestExpectation,
261 monitor: &FlowMonitor,
262 state: &State,
263 failures: &mut Vec<String>,
264) {
265 let explanation = monitor.explain(state);
266 for step in &expect.done {
267 if !monitor.marking().done.contains(step) {
268 failures.push(format!(
269 "expected step '{step}' done; done = [{}]",
270 join(&monitor.marking().done.iter().cloned().collect::<Vec<_>>())
271 ));
272 }
273 }
274 for step in &expect.active {
275 if !explanation.active.contains(step) {
276 failures.push(format!(
277 "expected step '{step}' active; active = [{}]",
278 join(&explanation.active)
279 ));
280 }
281 }
282 for tool in &expect.allowed {
283 if !explanation.allowed_tools.contains(tool) {
284 failures.push(format!(
285 "expected tool '{tool}' allowed; allowed = [{}]",
286 join(&explanation.allowed_tools)
287 ));
288 }
289 }
290 for tool in &expect.blocked {
291 if !explanation.blocked_tools.contains_key(tool) {
292 failures.push(format!(
293 "expected tool '{tool}' blocked; blocked = [{}]",
294 join(
295 &explanation
296 .blocked_tools
297 .keys()
298 .cloned()
299 .collect::<Vec<_>>()
300 )
301 ));
302 }
303 }
304 for (key, expected) in &expect.state {
305 let actual = state.get::<Value>(key);
306 if actual.as_ref() != Some(expected) {
307 failures.push(format!(
308 "expected state '{key}' = {expected}; got {}",
309 actual.map_or("<absent>".to_string(), |v| v.to_string())
310 ));
311 }
312 }
313 if let Some(complete) = expect.complete
314 && monitor.is_complete() != complete
315 {
316 failures.push(format!(
317 "expected complete = {complete}; got {}",
318 monitor.is_complete()
319 ));
320 }
321}
322
323fn join(items: &[String]) -> String {
324 items.join(", ")
325}
326
327#[cfg(test)]
328mod tests {
329 use super::super::SessionSpec;
330 use serde_json::json;
331
332 fn spec_with_tests() -> SessionSpec {
333 SessionSpec::from_value(json!({
334 "name": "collections",
335 "instruction": "Collect.",
336 "tools": [
337 {"name": "verify_identity", "set_state": {"identity_verified": true}},
338 {"name": "charge_card", "response": {"charged": true}}
339 ],
340 "flow": {
341 "steps": [
342 {"id": "verify", "posture": "Verify.", "allow": ["verify_identity"],
343 "done": {"is_true": "identity_verified"}},
344 {"id": "pay", "after": ["verify"], "posture": "Pay.",
345 "allow": ["charge_card"], "done": {"called_ok": "charge_card"}}
346 ],
347 "constraints": [
348 {"never_until": {"tool": "charge_card",
349 "until": {"is_true": "identity_verified"}}},
350 {"require": ["pay"]}
351 ]
352 },
353 "tests": [
354 {"name": "happy path", "script": [
355 {"expect": {"active": ["verify"], "blocked": ["charge_card"],
356 "complete": false}},
357 {"tool": "verify_identity"},
358 {"expect": {"done": ["verify"], "active": ["pay"],
359 "allowed": ["charge_card"],
360 "state": {"identity_verified": true}}},
361 {"tool": "charge_card"},
362 {"expect": {"done": ["pay"], "complete": true}}
363 ]},
364 {"name": "premature charge is blocked", "script": [
365 {"tool": "charge_card"},
366 {"expect": {"blocked": ["charge_card"], "complete": false}}
367 ]},
368 {"name": "deliberately wrong", "script": [
369 {"expect": {"done": ["pay"]}}
370 ]}
371 ]
372 }))
373 .expect("spec parses")
374 }
375
376 #[test]
377 fn scripted_tests_replay_through_the_real_monitor() {
378 let reports = spec_with_tests().run_tests();
379 assert_eq!(reports.len(), 3);
380 assert!(reports[0].passed, "happy path: {:?}", reports[0].failures);
381 assert!(
382 reports[1].passed,
383 "anticipated block passes: {:?}",
384 reports[1].failures
385 );
386 assert!(!reports[2].passed, "wrong expectation fails");
387 assert!(reports[2].failures[0].failures[0].contains("expected step 'pay' done"));
388 }
389
390 #[test]
391 fn trace_snapshots_every_event() {
392 let spec = spec_with_tests();
393 let snapshots = super::trace_test(&spec, "happy path").expect("traces");
394 assert_eq!(snapshots.len(), 6);
396 assert_eq!(snapshots[0].event, "start");
397 assert!(
398 snapshots[0]
399 .explanation
400 .active
401 .contains(&"verify".to_string())
402 );
403 assert!(snapshots[2].done.contains(&"verify".to_string()));
405 assert!(snapshots[2].explanation.active.contains(&"pay".to_string()));
406 assert!(snapshots[5].complete);
408 assert!(super::trace_test(&spec, "no such test").is_err());
409 }
410
411 #[test]
412 fn computed_variables_latch_guards_in_replay() {
413 let spec = SessionSpec::from_value(json!({
414 "instruction": "x",
415 "state": {"attempts": {"type": "number", "default": 0}},
416 "tools": [{"name": "record_score", "set_state": {"score": 0.9}}],
417 "computed": [{"key": "high_risk",
418 "from": {"gt": [{"key": "score"}, {"const": 0.5}]}}],
419 "flow": {"steps": [
420 {"id": "assess", "posture": "Assess.", "allow": ["record_score"],
421 "done": {"is_true": "high_risk"}},
422 {"id": "wrap", "after": ["assess"], "terminal": true}
423 ], "constraints": [{"require": ["wrap"]}]},
424 "tests": [{"name": "risk computes", "script": [
425 {"expect": {"active": ["assess"],
426 "state": {"attempts": 0}}},
427 {"tool": "record_score"},
428 {"expect": {"done": ["assess", "wrap"], "complete": true,
429 "state": {"derived:high_risk": true}}}
430 ]}]
431 }))
432 .expect("parses");
433 let reports = spec.run_tests();
434 assert!(
435 reports[0].passed,
436 "computed guard latches offline: {:?}",
437 reports[0].failures
438 );
439 }
440
441 #[test]
442 fn unanticipated_block_is_a_failure() {
443 let mut spec = spec_with_tests();
444 spec.tests = vec![super::SpecTest {
446 name: "unanticipated".into(),
447 script: vec![super::SimEvent::Tool("charge_card".into())],
448 }];
449 let reports = spec.run_tests();
450 assert!(!reports[0].passed);
451 assert!(reports[0].failures[0].failures[0].contains("was blocked"));
452 }
453}