1use std::collections::HashMap;
2use std::fs::{File, OpenOptions};
3use std::io::Write;
4use std::net::{IpAddr, Ipv4Addr};
5use std::num::NonZeroUsize;
6use std::path::{Path, PathBuf};
7use std::sync::{Arc, OnceLock, Weak};
8use std::time::Duration;
9
10use tokio::io::{AsyncReadExt, AsyncWriteExt};
11use tokio::net::TcpListener;
12use tokio::task::JoinHandle;
13
14use ai_agents_core::{
15 ChatMessage, LLMChunk, LLMConfig, LLMError, LLMFeature, LLMProvider, LLMResponse,
16 LLMToolRequest, Tool, ToolChoice, ToolExecutionContext, ToolPolicyBindings, ToolResult,
17};
18use ai_agents_hitl::{ApprovalHandler, ApprovalRequest, ApprovalResult, ApprovalTrigger};
19use ai_agents_llm::providers::{ProviderType, UnifiedLLMProvider};
20use ai_agents_llm::{FinishReason, LLMRegistry};
21use ai_agents_runtime::spec::AgentSpec;
22use ai_agents_tools::{
23 CommandResponse, DiagnosticItem, DiagnosticsProvider, StaticCommandRunner,
24 StaticDiagnosticsProvider, StaticWebSearchProvider, ToolRegistry,
25 UnavailableDiagnosticsProvider, UnavailableWebSearchProvider, WebFetchResolver, WebFetchTool,
26 WebFetchTransport, WebFetchTransportRequest, WebFetchTransportResponse, WebSearchProvider,
27 WebSearchResponse, create_builtin_registry,
28};
29use async_trait::async_trait;
30use futures::Stream;
31use parking_lot::Mutex;
32use serde::{Deserialize, Serialize};
33use serde_json::{Value, json};
34use sha2::{Digest, Sha256};
35
36use crate::evidence::{ToolExecutionRecord, ToolExecutionSource};
37use crate::{EvalError, Result};
38
39#[derive(Debug, Clone, Serialize, Deserialize, Default)]
41#[serde(deny_unknown_fields)]
42pub struct FixturesConfig {
43 #[serde(default)]
45 pub context: Option<Value>,
46 #[serde(default)]
48 pub context_file: Option<PathBuf>,
49 #[serde(default)]
51 pub tools: HashMap<String, ToolMockConfig>,
52 #[serde(default)]
54 pub llm: LlmFixtureConfig,
55 #[serde(default)]
57 pub mock_server: Option<MockServerConfig>,
58 #[serde(default)]
60 pub diagnostics: Option<DiagnosticsFixtureConfig>,
61 #[serde(default)]
63 pub commands: Option<CommandsFixtureConfig>,
64 #[serde(default)]
66 pub web_search: Option<WebSearchFixtureConfig>,
67 #[serde(default)]
69 pub web_fetch_transport: Option<WebFetchTransportFixtureConfig>,
70 #[serde(default)]
72 pub approvals: Option<ApprovalFixtureConfig>,
73 #[serde(default)]
75 pub workspace_policy: Option<WorkspacePolicyFixtureConfig>,
76}
77
78#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)]
80#[serde(deny_unknown_fields)]
81pub struct WorkspacePolicyFixtureConfig {
82 #[serde(default)]
84 pub read_tools: Vec<String>,
85 #[serde(default)]
87 pub write_tools: Vec<String>,
88}
89
90#[derive(Debug, Clone, Serialize, Deserialize, Default)]
92#[serde(deny_unknown_fields)]
93pub struct ApprovalFixtureConfig {
94 #[serde(default)]
96 pub rules: Vec<ApprovalFixtureRule>,
97 #[serde(default)]
99 pub default: ApprovalFixtureOutcome,
100 #[serde(default)]
102 pub preferred_language: Option<String>,
103 #[serde(default)]
105 pub supported_languages: Option<Vec<String>>,
106}
107
108#[derive(Debug, Clone, Serialize, Deserialize)]
110pub struct ApprovalFixtureRule {
111 pub trigger: ApprovalFixtureTrigger,
113 #[serde(default)]
115 pub occurrence: Option<NonZeroUsize>,
116 #[serde(flatten)]
118 pub outcome: ApprovalFixtureOutcome,
119}
120
121#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
123#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
124pub enum ApprovalFixtureTrigger {
125 Tool {
126 name: String,
127 #[serde(default)]
128 args: Option<Value>,
129 },
130 Condition {
131 name: String,
132 #[serde(default)]
133 matched: Option<String>,
134 },
135 StateTransition {
136 #[serde(default)]
137 from: Option<String>,
138 to: String,
139 },
140 DisambiguationEscalation {
141 #[serde(default)]
142 reason: Option<String>,
143 },
144}
145
146#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq)]
148#[serde(tag = "outcome", rename_all = "snake_case", deny_unknown_fields)]
149pub enum ApprovalFixtureOutcome {
150 Approve,
151 Reject {
152 #[serde(default)]
153 reason: Option<String>,
154 },
155 Modify {
156 #[serde(default)]
157 changes: HashMap<String, Value>,
158 },
159 Timeout,
160 #[default]
161 Unavailable,
162}
163
164pub struct FixtureApprovalHandler {
166 config: ApprovalFixtureConfig,
167 occurrences: Mutex<Vec<usize>>,
168}
169
170impl FixtureApprovalHandler {
171 pub fn new(config: ApprovalFixtureConfig) -> Self {
172 let rule_count = config.rules.len();
173 Self {
174 config,
175 occurrences: Mutex::new(vec![0; rule_count]),
176 }
177 }
178
179 fn select_outcome(&self, request: &ApprovalRequest) -> ApprovalFixtureOutcome {
180 let mut occurrences = self.occurrences.lock();
181 for (index, rule) in self.config.rules.iter().enumerate() {
182 if !rule.trigger.matches(&request.trigger) {
183 continue;
184 }
185 occurrences[index] += 1;
186 if rule
187 .occurrence
188 .is_none_or(|occurrence| occurrence.get() == occurrences[index])
189 {
190 return rule.outcome.clone();
191 }
192 }
193 self.config.default.clone()
194 }
195}
196
197impl ApprovalFixtureTrigger {
198 fn matches(&self, trigger: &ApprovalTrigger) -> bool {
199 match (self, trigger) {
200 (
201 Self::Tool { name, args },
202 ApprovalTrigger::Tool {
203 name: actual_name,
204 args: actual_args,
205 },
206 ) => name == actual_name && args.as_ref().is_none_or(|args| args == actual_args),
207 (
208 Self::Condition { name, matched },
209 ApprovalTrigger::Condition {
210 name: actual_name,
211 matched: actual_matched,
212 },
213 ) => {
214 name != "disambiguation_escalation"
215 && name == actual_name
216 && matched
217 .as_ref()
218 .is_none_or(|matched| matched == actual_matched)
219 }
220 (
221 Self::StateTransition { from, to },
222 ApprovalTrigger::State {
223 from: actual_from,
224 to: actual_to,
225 },
226 ) => from == actual_from && to == actual_to,
227 (
228 Self::DisambiguationEscalation { reason },
229 ApprovalTrigger::Condition {
230 name,
231 matched: actual_reason,
232 },
233 ) => {
234 name == "disambiguation_escalation"
235 && reason.as_ref().is_none_or(|reason| reason == actual_reason)
236 }
237 _ => false,
238 }
239 }
240}
241
242impl From<ApprovalFixtureOutcome> for ApprovalResult {
243 fn from(outcome: ApprovalFixtureOutcome) -> Self {
244 match outcome {
245 ApprovalFixtureOutcome::Approve => Self::approved(),
246 ApprovalFixtureOutcome::Reject { reason } => Self::rejected(reason),
247 ApprovalFixtureOutcome::Modify { changes } => Self::modified(changes),
248 ApprovalFixtureOutcome::Timeout => Self::timeout(),
249 ApprovalFixtureOutcome::Unavailable => {
250 Self::rejected_with_reason("Approval fixture unavailable")
251 }
252 }
253 }
254}
255
256#[async_trait]
257impl ApprovalHandler for FixtureApprovalHandler {
258 async fn request_approval(&self, request: ApprovalRequest) -> ApprovalResult {
259 self.select_outcome(&request).into()
260 }
261
262 fn preferred_language(&self) -> Option<String> {
263 self.config.preferred_language.clone()
264 }
265
266 fn supported_languages(&self) -> Option<Vec<String>> {
267 self.config.supported_languages.clone()
268 }
269}
270
271pub fn build_approval_handler(config: &ApprovalFixtureConfig) -> Arc<dyn ApprovalHandler> {
273 Arc::new(FixtureApprovalHandler::new(config.clone()))
274}
275
276#[derive(Debug, Clone, Serialize, Deserialize)]
278#[serde(deny_unknown_fields)]
279pub struct DiagnosticsFixtureConfig {
280 #[serde(default = "default_true")]
282 pub available: bool,
283 #[serde(default)]
285 pub items: Vec<DiagnosticItem>,
286}
287
288#[derive(Debug, Clone, Serialize, Deserialize)]
290#[serde(deny_unknown_fields)]
291pub struct CommandsFixtureConfig {
292 #[serde(default = "default_true")]
294 pub available: bool,
295 #[serde(default)]
297 pub entries: Vec<CommandFixtureEntry>,
298}
299
300#[derive(Debug, Clone, Serialize, Deserialize)]
302#[serde(deny_unknown_fields)]
303pub struct CommandFixtureEntry {
304 #[serde(default)]
306 pub argv: Vec<String>,
307 #[serde(default)]
309 pub response: CommandResponse,
310}
311
312#[derive(Debug, Clone, Serialize, Deserialize)]
314#[serde(deny_unknown_fields)]
315pub struct WebSearchFixtureEntry {
316 pub query: String,
318 #[serde(default)]
320 pub response: WebSearchResponse,
321}
322
323#[derive(Debug, Clone, Serialize, Deserialize)]
325#[serde(deny_unknown_fields)]
326pub struct WebSearchFixtureConfig {
327 #[serde(default = "default_true")]
329 pub available: bool,
330 #[serde(default)]
332 pub entries: Vec<WebSearchFixtureEntry>,
333}
334
335#[derive(Debug, Clone, Serialize, Deserialize, Default)]
337#[serde(deny_unknown_fields)]
338pub struct WebFetchTransportFixtureConfig {
339 #[serde(default)]
341 pub routes: Vec<WebFetchTransportFixtureRoute>,
342}
343
344#[derive(Debug, Clone, Serialize, Deserialize)]
346#[serde(deny_unknown_fields)]
347pub struct WebFetchTransportFixtureRoute {
348 pub url: String,
350 #[serde(default = "default_status")]
352 pub status: u16,
353 #[serde(default)]
355 pub headers: HashMap<String, String>,
356 #[serde(default)]
358 pub body: Value,
359}
360
361#[derive(Debug, Clone, Serialize, Deserialize)]
363#[serde(deny_unknown_fields)]
364pub struct ToolMockConfig {
365 #[serde(default = "default_true")]
367 pub success: bool,
368 #[serde(default)]
370 pub output: Value,
371}
372
373impl Default for ToolMockConfig {
374 fn default() -> Self {
375 Self {
376 success: true,
377 output: Value::Null,
378 }
379 }
380}
381
382#[derive(Debug, Clone, Serialize, Deserialize, Default)]
384#[serde(deny_unknown_fields)]
385pub struct LlmFixtureConfig {
386 #[serde(default)]
388 pub mode: LlmFixtureMode,
389 #[serde(default)]
391 pub cassette: Option<PathBuf>,
392 #[serde(default)]
394 pub responses: Vec<String>,
395 #[serde(default)]
397 pub responses_by_alias: HashMap<String, Vec<String>>,
398 #[serde(default)]
400 pub errors_by_alias: HashMap<String, String>,
401 #[serde(default)]
403 pub outcomes_by_alias: HashMap<String, Vec<LlmFixtureOutcome>>,
404 #[serde(default)]
406 pub delays_by_alias: HashMap<String, u64>,
407}
408
409#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
411#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
412pub enum LlmFixtureOutcome {
413 Response {
414 content: String,
415 },
416 Error {
417 message: String,
418 #[serde(default)]
419 status: Option<u16>,
420 },
421}
422
423#[derive(Debug, Clone, Copy, Serialize, Deserialize, Default, PartialEq, Eq)]
425#[serde(rename_all = "snake_case")]
426pub enum LlmFixtureMode {
427 #[default]
428 Real,
429 Mock,
430 Replay,
431 Record,
432}
433
434#[derive(Debug, Clone, Serialize, Deserialize, Default)]
436#[serde(deny_unknown_fields)]
437pub struct MockServerConfig {
438 #[serde(default)]
440 pub enabled: bool,
441 #[serde(default)]
443 pub port: Option<u16>,
444 #[serde(default)]
446 pub routes: Vec<Value>,
447}
448
449#[derive(Debug, Clone, Deserialize)]
451#[serde(deny_unknown_fields)]
452struct MockRoute {
453 method: String,
455 path: String,
457 #[serde(default = "default_status")]
459 status: u16,
460 #[serde(default)]
462 headers: HashMap<String, String>,
463 #[serde(default)]
465 body: Value,
466}
467
468impl FixturesConfig {
469 pub(crate) fn validate(&self) -> Result<()> {
470 if let Some(mock_server) = &self.mock_server {
471 for (index, route) in mock_server.routes.iter().cloned().enumerate() {
472 serde_json::from_value::<MockRoute>(route).map_err(|error| {
473 EvalError::Config(format!(
474 "fixtures.mock_server.routes[{}] is invalid: {}",
475 index, error
476 ))
477 })?;
478 }
479 }
480 Ok(())
481 }
482}
483
484pub struct MockServerHandle {
486 base_url: String,
488 task: JoinHandle<()>,
490}
491
492impl MockServerHandle {
493 pub fn context(&self) -> HashMap<String, Value> {
494 let mut context = HashMap::new();
495 context.insert(
496 "mock_server".to_string(),
497 serde_json::json!({"base_url": self.base_url}),
498 );
499 context
500 }
501}
502
503impl Drop for MockServerHandle {
504 fn drop(&mut self) {
505 self.task.abort();
506 }
507}
508
509#[derive(Debug, PartialEq, Eq)]
511pub struct AttemptFixtureContext {
512 pub isolation_id: String,
514 pub workspace: PathBuf,
516 pub mock_server_base_url: Option<String>,
518}
519
520pub(crate) struct AttemptWorkspace {
522 context: AttemptFixtureContext,
523}
524
525impl AttemptWorkspace {
526 pub(crate) fn create(mock_server: Option<&MockServerHandle>) -> Result<Self> {
527 Ok(Self {
528 context: AttemptFixtureContext::create(mock_server)?,
529 })
530 }
531}
532
533impl std::ops::Deref for AttemptWorkspace {
534 type Target = AttemptFixtureContext;
535
536 fn deref(&self) -> &Self::Target {
537 &self.context
538 }
539}
540
541impl Drop for AttemptWorkspace {
542 fn drop(&mut self) {
543 if let Err(error) = std::fs::remove_dir_all(&self.context.workspace)
544 && error.kind() != std::io::ErrorKind::NotFound
545 {
546 tracing::warn!(
547 workspace = %self.context.workspace.display(),
548 error = %error,
549 "failed to clean eval attempt workspace"
550 );
551 }
552 }
553}
554
555impl AttemptFixtureContext {
556 pub fn create(mock_server: Option<&MockServerHandle>) -> Result<Self> {
557 let isolation_id = uuid::Uuid::new_v4().to_string();
558 let workspace = std::env::temp_dir().join(format!("ai_agents_eval_{}", isolation_id));
559 std::fs::create_dir_all(&workspace)?;
560 let workspace = workspace.canonicalize()?;
561 Ok(Self {
562 isolation_id,
563 workspace,
564 mock_server_base_url: mock_server.map(|server| server.base_url.clone()),
565 })
566 }
567
568 pub fn runtime_context(&self) -> HashMap<String, Value> {
569 let mut context = HashMap::from([(
570 "eval".to_string(),
571 json!({"workspace": self.workspace.display().to_string()}),
572 )]);
573 if let Some(base_url) = &self.mock_server_base_url {
574 context.insert("mock_server".to_string(), json!({"base_url": base_url}));
575 }
576 context
577 }
578
579 pub fn interpolate_llm_fixture(&self, config: &LlmFixtureConfig) -> Result<LlmFixtureConfig> {
580 let mut rewritten = config.clone();
581 for response in &mut rewritten.responses {
582 *response = self.interpolate_response(response)?;
583 }
584 for responses in rewritten.responses_by_alias.values_mut() {
585 for response in responses {
586 *response = self.interpolate_response(response)?;
587 }
588 }
589 for outcomes in rewritten.outcomes_by_alias.values_mut() {
590 for outcome in outcomes {
591 if let LlmFixtureOutcome::Response { content } = outcome {
592 *content = self.interpolate_response(content)?;
593 }
594 }
595 }
596 Ok(rewritten)
597 }
598
599 fn interpolate_response(&self, response: &str) -> Result<String> {
600 const WORKSPACE_TOKEN: &str = "{{ eval.workspace }}";
601 const MOCK_SERVER_TOKEN: &str = "{{ mock_server.base_url }}";
602 if !response.contains(WORKSPACE_TOKEN) && !response.contains(MOCK_SERVER_TOKEN) {
603 return Ok(response.to_string());
604 }
605 if response.contains(MOCK_SERVER_TOKEN) && self.mock_server_base_url.is_none() {
606 return Err(EvalError::Config(format!(
607 "mock LLM response uses {} without an enabled fixtures.mock_server",
608 MOCK_SERVER_TOKEN
609 )));
610 }
611 let workspace = self.workspace.display().to_string();
612 let mock_server = self.mock_server_base_url.as_deref();
613 if let Ok(mut value) = serde_json::from_str::<Value>(response) {
614 interpolate_json_strings(&mut value, &workspace, mock_server);
615 return serde_json::to_string(&value).map_err(EvalError::from);
616 }
617 Ok(interpolate_fixture_string(
618 response,
619 &workspace,
620 mock_server,
621 ))
622 }
623}
624
625fn interpolate_json_strings(value: &mut Value, workspace: &str, mock_server: Option<&str>) {
626 match value {
627 Value::String(text) => *text = interpolate_fixture_string(text, workspace, mock_server),
628 Value::Array(values) => {
629 for value in values {
630 interpolate_json_strings(value, workspace, mock_server);
631 }
632 }
633 Value::Object(values) => {
634 let previous = std::mem::take(values);
635 for (key, mut value) in previous {
636 interpolate_json_strings(&mut value, workspace, mock_server);
637 values.insert(
638 interpolate_fixture_string(&key, workspace, mock_server),
639 value,
640 );
641 }
642 }
643 _ => {}
644 }
645}
646
647fn interpolate_fixture_string(value: &str, workspace: &str, mock_server: Option<&str>) -> String {
648 let value = value.replace("{{ eval.workspace }}", workspace);
649 match mock_server {
650 Some(base_url) => value.replace("{{ mock_server.base_url }}", base_url),
651 None => value,
652 }
653}
654
655#[derive(Clone, Default)]
657pub struct RecordingToolLog {
658 inner: Arc<Mutex<Vec<ToolExecutionRecord>>>,
660}
661
662impl From<&ai_agents_core::ToolCallSource> for ToolExecutionSource {
663 fn from(source: &ai_agents_core::ToolCallSource) -> Self {
664 match source {
665 ai_agents_core::ToolCallSource::Model => Self::Llm,
666 ai_agents_core::ToolCallSource::Skill { .. } => Self::Skill,
667 ai_agents_core::ToolCallSource::Plan { .. } => Self::Plan,
668 ai_agents_core::ToolCallSource::StateAction { .. } => Self::StateAction,
669 ai_agents_core::ToolCallSource::Orchestration => Self::Orchestration,
670 ai_agents_core::ToolCallSource::Spawner => Self::Spawner,
671 ai_agents_core::ToolCallSource::EvalFixture => Self::Mock,
672 ai_agents_core::ToolCallSource::Fallback { .. }
673 | ai_agents_core::ToolCallSource::Task
674 | ai_agents_core::ToolCallSource::Manual => Self::Llm,
675 }
676 }
677}
678
679fn state_from_source(source: &ai_agents_core::ToolCallSource) -> Option<String> {
680 match source {
681 ai_agents_core::ToolCallSource::StateAction { state, .. } => state.clone(),
682 _ => None,
683 }
684}
685
686impl RecordingToolLog {
687 pub fn new() -> Self {
688 Self::default()
689 }
690
691 pub fn len(&self) -> usize {
692 self.inner.lock().len()
693 }
694
695 pub fn is_empty(&self) -> bool {
696 self.inner.lock().is_empty()
697 }
698
699 pub fn push(&self, record: ToolExecutionRecord) {
700 self.inner.lock().push(record);
701 }
702
703 pub fn push_executor_record(&self, record: &ai_agents_core::ToolExecutionRecord) {
705 let output = record.success.then(|| {
706 serde_json::from_str(&record.output).unwrap_or(Value::String(record.output.clone()))
707 });
708 let mut metadata = record.metadata.clone();
709 metadata.insert("executed".to_string(), Value::Bool(record.executed));
710 metadata.insert(
711 "policy".to_string(),
712 serde_json::to_value(&record.policy).unwrap_or(Value::Null),
713 );
714 metadata.insert(
715 "approval".to_string(),
716 serde_json::to_value(&record.approval).unwrap_or(Value::Null),
717 );
718 metadata.insert("timed_out".to_string(), Value::Bool(record.timed_out));
719 metadata.insert("cancelled".to_string(), Value::Bool(record.cancelled));
720 self.push(ToolExecutionRecord {
721 call_id: record.call_id.clone(),
722 tool_id: record.canonical_id.clone(),
723 requested_name: record.requested_name.clone(),
724 source: ToolExecutionSource::from(&record.source),
725 state: state_from_source(&record.source),
726 actor_id: None,
727 arguments_original: record.arguments.clone(),
728 arguments_executed: record.executed_arguments.clone(),
729 executed: record.executed,
730 success: record.success,
731 output,
732 error: (!record.success).then_some(record.output.clone()),
733 metadata: Some(serde_json::to_value(metadata).unwrap_or(Value::Null)),
734 started_at: record.started_at,
735 duration_ms: record.duration_ms,
736 observability_span_id: None,
737 });
738 }
739
740 pub fn records_since(&self, index: usize) -> Vec<ToolExecutionRecord> {
741 self.inner.lock().iter().skip(index).cloned().collect()
742 }
743}
744
745pub fn resolve_fixture_context(
746 config: &FixturesConfig,
747 base_dir: &Path,
748) -> Result<HashMap<String, Value>> {
749 let mut result = HashMap::new();
750 if let Some(path) = &config.context_file {
751 let resolved = resolve_path(base_dir, path);
752 let content = std::fs::read_to_string(&resolved).map_err(|error| {
753 EvalError::Config(format!(
754 "failed to read context_file '{}': {}",
755 resolved.display(),
756 error
757 ))
758 })?;
759 let value: Value = serde_json::from_str(&content).map_err(|error| {
760 EvalError::Config(format!(
761 "failed to parse context_file '{}': {}",
762 resolved.display(),
763 error
764 ))
765 })?;
766 merge_object_into_map(&mut result, value)?;
767 }
768 if let Some(value) = &config.context {
769 merge_object_into_map(&mut result, value.clone())?;
770 }
771 Ok(result)
772}
773
774fn merge_object_into_map(target: &mut HashMap<String, Value>, value: Value) -> Result<()> {
775 let Value::Object(map) = value else {
776 return Err(EvalError::Config(
777 "fixture context must be a JSON object".into(),
778 ));
779 };
780 for (key, value) in map {
781 target.insert(key, value);
782 }
783 Ok(())
784}
785
786fn resolve_path(base_dir: &Path, path: &Path) -> PathBuf {
787 if path.is_absolute() {
788 path.to_path_buf()
789 } else {
790 base_dir.join(path)
791 }
792}
793
794pub async fn start_mock_server(
795 config: Option<&MockServerConfig>,
796) -> Result<Option<MockServerHandle>> {
797 let Some(config) = config else {
798 return Ok(None);
799 };
800 if !config.enabled {
801 return Ok(None);
802 }
803 let routes = config
804 .routes
805 .iter()
806 .cloned()
807 .map(serde_json::from_value::<MockRoute>)
808 .collect::<std::result::Result<Vec<_>, _>>()?;
809 let port = config.port.unwrap_or(0);
810 let listener = TcpListener::bind(("127.0.0.1", port))
811 .await
812 .map_err(|error| EvalError::Runtime(format!("failed to start mock server: {}", error)))?;
813 let addr = listener.local_addr().map_err(|error| {
814 EvalError::Runtime(format!("failed to read mock server addr: {}", error))
815 })?;
816 let base_url = format!("http://{}", addr);
817 let task = tokio::spawn(async move {
818 loop {
819 let Ok((stream, _)) = listener.accept().await else {
820 break;
821 };
822 let routes = routes.clone();
823 tokio::spawn(async move {
824 let _ = handle_mock_connection(stream, routes).await;
825 });
826 }
827 });
828 Ok(Some(MockServerHandle { base_url, task }))
829}
830
831async fn handle_mock_connection(
832 mut stream: tokio::net::TcpStream,
833 routes: Vec<MockRoute>,
834) -> std::io::Result<()> {
835 let mut buffer = vec![0_u8; 8192];
836 let read = stream.read(&mut buffer).await?;
837 let request = String::from_utf8_lossy(&buffer[..read]);
838 let first_line = request.lines().next().unwrap_or_default();
839 let mut parts = first_line.split_whitespace();
840 let method = parts.next().unwrap_or_default();
841 let path = parts.next().unwrap_or_default();
842 let route = routes
843 .iter()
844 .find(|route| route.method.eq_ignore_ascii_case(method) && route.path == path);
845 let (status, headers, body) = if let Some(route) = route {
846 (
847 route.status,
848 route.headers.clone(),
849 mock_body_to_string(&route.body),
850 )
851 } else {
852 (
853 404,
854 HashMap::new(),
855 serde_json::json!({"error":"not found"}).to_string(),
856 )
857 };
858 let reason = match status {
859 200 => "OK",
860 201 => "Created",
861 204 => "No Content",
862 400 => "Bad Request",
863 401 => "Unauthorized",
864 403 => "Forbidden",
865 404 => "Not Found",
866 500 => "Internal Server Error",
867 _ => "OK",
868 };
869 let mut response = format!(
870 "HTTP/1.1 {} {}\r\nContent-Length: {}\r\nContent-Type: application/json\r\nConnection: close\r\n",
871 status,
872 reason,
873 body.len()
874 );
875 for (key, value) in headers {
876 response.push_str(&format!("{}: {}\r\n", key, value));
877 }
878 response.push_str("\r\n");
879 response.push_str(&body);
880 stream.write_all(response.as_bytes()).await?;
881 Ok(())
882}
883
884fn mock_body_to_string(body: &Value) -> String {
885 if let Some(text) = body.as_str() {
886 text.to_string()
887 } else {
888 serde_json::to_string(body).unwrap_or_else(|_| "null".to_string())
889 }
890}
891
892pub fn build_tool_registry(
893 fixtures: &FixturesConfig,
894 log: RecordingToolLog,
895) -> Result<ToolRegistry> {
896 let builtin = create_builtin_registry();
897 let mut registry = ToolRegistry::new();
898 let provider: Arc<dyn DiagnosticsProvider> = if let Some(diagnostics) = &fixtures.diagnostics {
899 Arc::new(StaticDiagnosticsProvider::with_availability(
900 diagnostics.items.clone(),
901 diagnostics.available,
902 ))
903 } else {
904 Arc::new(UnavailableDiagnosticsProvider)
905 };
906 builtin.set_diagnostics_provider(provider.clone());
907 registry.set_diagnostics_provider(provider);
908
909 if let Some(commands) = &fixtures.commands {
910 let responses = commands
911 .entries
912 .iter()
913 .map(|entry| (entry.argv.clone(), entry.response.clone()))
914 .collect();
915 let runner = Arc::new(StaticCommandRunner::with_availability(
916 responses,
917 commands.available,
918 ));
919 builtin.set_command_runner(runner.clone());
920 registry.set_command_runner(runner);
921 }
922
923 let search_provider: Arc<dyn WebSearchProvider> = if let Some(web_search) = &fixtures.web_search
924 {
925 let responses = web_search
926 .entries
927 .iter()
928 .map(|entry| (entry.query.clone(), entry.response.clone()))
929 .collect();
930 Arc::new(StaticWebSearchProvider::with_availability(
931 responses,
932 web_search.available,
933 ))
934 } else {
935 Arc::new(UnavailableWebSearchProvider)
936 };
937 builtin.set_web_search_provider(search_provider.clone());
938 registry.set_web_search_provider(search_provider);
939
940 for (id, mock) in &fixtures.tools {
941 let contract = builtin.get(id);
942 registry
943 .register(Arc::new(RecordingTool::new(
944 Arc::new(MockTool::new(id.clone(), mock.clone(), contract)),
945 log.clone(),
946 ToolExecutionSource::Mock,
947 )))
948 .map_err(|error| EvalError::Config(error.to_string()))?;
949 }
950
951 let web_fetch = fixtures
952 .web_fetch_transport
953 .as_ref()
954 .map(build_web_fetch_fixture_tool)
955 .transpose()?;
956 for id in builtin.list_ids() {
957 if fixtures.tools.contains_key(&id) {
958 continue;
959 }
960 let tool = if id == "web_fetch" {
961 web_fetch.clone().or_else(|| builtin.get(&id))
962 } else {
963 builtin.get(&id)
964 };
965 if let Some(tool) = tool {
966 registry
967 .register(Arc::new(RecordingTool::new(
968 tool,
969 log.clone(),
970 ToolExecutionSource::Llm,
971 )))
972 .map_err(|error| EvalError::Config(error.to_string()))?;
973 }
974 }
975
976 Ok(registry)
977}
978
979fn build_web_fetch_fixture_tool(config: &WebFetchTransportFixtureConfig) -> Result<Arc<dyn Tool>> {
980 let mut routes = HashMap::new();
981 for route in &config.routes {
982 let normalized_url = reqwest::Url::parse(&route.url)
983 .map_err(|error| {
984 EvalError::Config(format!(
985 "invalid fixtures.web_fetch_transport route URL '{}': {}",
986 route.url, error
987 ))
988 })?
989 .to_string();
990 let content_type = header_value(&route.headers, "content-type").map(str::to_string);
991 let location = header_value(&route.headers, "location").map(str::to_string);
992 let response = WebFetchTransportResponse {
993 status: route.status,
994 content_type,
995 location,
996 body: mock_body_to_string(&route.body).into_bytes(),
997 };
998 if routes.insert(normalized_url.clone(), response).is_some() {
999 return Err(EvalError::Config(format!(
1000 "duplicate fixtures.web_fetch_transport route URL '{}'",
1001 normalized_url
1002 )));
1003 }
1004 }
1005 Ok(Arc::new(WebFetchTool::with_transport_and_resolver(
1006 Arc::new(FixtureWebFetchTransport { routes }),
1007 Arc::new(FixtureWebFetchResolver),
1008 )))
1009}
1010
1011fn header_value<'a>(headers: &'a HashMap<String, String>, name: &str) -> Option<&'a str> {
1012 headers
1013 .iter()
1014 .find(|(key, _)| key.eq_ignore_ascii_case(name))
1015 .map(|(_, value)| value.as_str())
1016}
1017
1018struct FixtureWebFetchTransport {
1019 routes: HashMap<String, WebFetchTransportResponse>,
1020}
1021
1022#[async_trait]
1023impl WebFetchTransport for FixtureWebFetchTransport {
1024 async fn send(
1025 &self,
1026 request: WebFetchTransportRequest,
1027 ) -> std::result::Result<WebFetchTransportResponse, String> {
1028 self.routes
1029 .get(&request.url)
1030 .cloned()
1031 .ok_or_else(|| format!("No fixtures.web_fetch_transport route for {}", request.url))
1032 }
1033}
1034
1035struct FixtureWebFetchResolver;
1036
1037#[async_trait]
1038impl WebFetchResolver for FixtureWebFetchResolver {
1039 async fn resolve(&self, _host: &str, _port: u16) -> std::result::Result<Vec<IpAddr>, String> {
1040 Ok(vec![IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34))])
1041 }
1042}
1043
1044struct MockTool {
1046 id: String,
1048 config: ToolMockConfig,
1050 contract: Option<Arc<dyn Tool>>,
1052}
1053
1054impl MockTool {
1055 fn new(id: String, config: ToolMockConfig, contract: Option<Arc<dyn Tool>>) -> Self {
1056 Self {
1057 id,
1058 config,
1059 contract,
1060 }
1061 }
1062}
1063
1064#[async_trait]
1065impl Tool for MockTool {
1066 fn id(&self) -> &str {
1067 &self.id
1068 }
1069
1070 fn name(&self) -> &str {
1071 self.contract
1072 .as_ref()
1073 .map_or(self.id.as_str(), |tool| tool.name())
1074 }
1075
1076 fn description(&self) -> &str {
1077 self.contract
1078 .as_ref()
1079 .map_or("Evaluation mock tool", |tool| tool.description())
1080 }
1081
1082 fn input_schema(&self) -> Value {
1083 self.contract.as_ref().map_or_else(
1084 || serde_json::json!({"type": "object"}),
1085 |tool| tool.input_schema(),
1086 )
1087 }
1088
1089 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
1090 self.contract.as_ref().map_or_else(
1091 ai_agents_core::ToolSafetyMetadata::conservative_unknown,
1092 |tool| tool.safety_metadata(),
1093 )
1094 }
1095
1096 fn classify_call(&self, args: &Value) -> ai_agents_core::ToolCallClassification {
1097 self.contract.as_ref().map_or_else(
1098 || ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata()),
1099 |tool| tool.classify_call(args),
1100 )
1101 }
1102
1103 fn policy_bindings(&self) -> ToolPolicyBindings {
1104 self.contract
1105 .as_ref()
1106 .map_or_else(ToolPolicyBindings::default, |tool| tool.policy_bindings())
1107 }
1108
1109 async fn execute(
1110 &self,
1111 _args: Value,
1112 _ctx: ai_agents_core::ToolExecutionContext,
1113 ) -> ToolResult {
1114 let output = if self.config.output.is_string() {
1115 self.config.output.as_str().unwrap_or_default().to_string()
1116 } else {
1117 serde_json::to_string(&self.config.output).unwrap_or_else(|_| "null".to_string())
1118 };
1119 ToolResult {
1120 success: self.config.success,
1121 output,
1122 metadata: None,
1123 }
1124 }
1125}
1126
1127struct RecordingTool {
1129 inner: Arc<dyn Tool>,
1131}
1132
1133impl RecordingTool {
1134 fn new(inner: Arc<dyn Tool>, _log: RecordingToolLog, _source: ToolExecutionSource) -> Self {
1135 Self { inner }
1136 }
1137}
1138
1139#[async_trait]
1140impl Tool for RecordingTool {
1141 fn id(&self) -> &str {
1142 self.inner.id()
1143 }
1144
1145 fn name(&self) -> &str {
1146 self.inner.name()
1147 }
1148
1149 fn description(&self) -> &str {
1150 self.inner.description()
1151 }
1152
1153 fn input_schema(&self) -> Value {
1154 self.inner.input_schema()
1155 }
1156
1157 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
1158 self.inner.safety_metadata()
1159 }
1160
1161 fn classify_call(&self, args: &Value) -> ai_agents_core::ToolCallClassification {
1162 self.inner.classify_call(args)
1163 }
1164
1165 fn policy_bindings(&self) -> ToolPolicyBindings {
1166 self.inner.policy_bindings()
1167 }
1168
1169 async fn execute(&self, args: Value, ctx: ToolExecutionContext) -> ToolResult {
1170 self.inner.execute(args, ctx).await
1171 }
1172}
1173
1174pub fn build_llm_registry(
1175 spec: &AgentSpec,
1176 fixtures: &LlmFixtureConfig,
1177 base_dir: &Path,
1178) -> Result<(LLMRegistry, Option<Arc<dyn LLMProvider>>)> {
1179 let mut registry = LLMRegistry::new();
1180 let aliases = if spec.llms.is_empty() {
1181 vec![(
1182 "default".to_string(),
1183 spec.llm.as_config().cloned().unwrap_or_default(),
1184 )]
1185 } else {
1186 spec.llms
1187 .iter()
1188 .map(|(alias, config)| (alias.clone(), config.clone()))
1189 .collect()
1190 };
1191
1192 let cassette_records = if fixtures.mode == LlmFixtureMode::Replay {
1193 load_cassette_records(fixtures, base_dir)?
1194 } else {
1195 Vec::new()
1196 };
1197 let cassette_writer = if fixtures.mode == LlmFixtureMode::Record {
1198 let path = fixtures
1199 .cassette
1200 .as_ref()
1201 .map(|path| resolve_path(base_dir, path))
1202 .unwrap_or_else(|| base_dir.join("llm_cassette.jsonl"));
1203 Some(CassetteWriter::shared(path)?)
1204 } else {
1205 None
1206 };
1207 let mut judge_provider = None;
1208
1209 for (alias, config) in aliases {
1210 let fixture_delay_ms = fixtures.delays_by_alias.get(&alias).copied().unwrap_or(0);
1211 let provider = match fixtures.mode {
1212 LlmFixtureMode::Mock => {
1213 let fixture_responses =
1214 load_fixture_responses_for_alias(fixtures, base_dir, &alias)?;
1215 let fixture_outcomes =
1216 load_fixture_outcomes_for_alias(fixtures, &alias, &fixture_responses);
1217 Arc::new(
1218 SequenceLLMProvider::new(fixture_outcomes, fixture_delay_ms)
1219 .with_tool_choice(config.tool_choice.clone()),
1220 ) as Arc<dyn LLMProvider>
1221 }
1222 LlmFixtureMode::Replay => Arc::new(
1223 ReplayLLMProvider::new(
1224 alias.clone(),
1225 config.model.clone(),
1226 cassette_records.clone(),
1227 fixture_delay_ms,
1228 )
1229 .with_tool_choice(config.tool_choice.clone()),
1230 ) as Arc<dyn LLMProvider>,
1231 LlmFixtureMode::Real => build_real_provider(&config)?,
1232 LlmFixtureMode::Record => {
1233 let inner = build_real_provider(&config)?;
1234 Arc::new(RecordingLLMProvider::new(
1235 inner,
1236 alias.clone(),
1237 config.model.clone(),
1238 Arc::clone(
1239 cassette_writer
1240 .as_ref()
1241 .expect("record mode preflights a cassette writer"),
1242 ),
1243 )) as Arc<dyn LLMProvider>
1244 }
1245 };
1246 if judge_provider.is_none() {
1247 judge_provider = Some(provider.clone());
1248 }
1249 registry.register(alias, provider);
1250 }
1251
1252 let default_alias = spec.llm.get_default_alias();
1253 registry.set_default(default_alias);
1254 if let Some(router) = spec.llm.get_router_alias() {
1255 registry.set_router(router);
1256 }
1257
1258 Ok((registry, judge_provider))
1259}
1260
1261fn build_real_provider(
1262 config: &ai_agents_runtime::spec::LLMConfig,
1263) -> Result<Arc<dyn LLMProvider>> {
1264 use std::str::FromStr;
1265 let provider_type = ProviderType::from_str(&config.provider)
1266 .map_err(|error| EvalError::Config(error.to_string()))?;
1267 let core_config = ai_agents_core::LLMConfig {
1268 temperature: Some(config.temperature),
1269 max_tokens: Some(config.max_tokens),
1270 top_p: config.top_p,
1271 top_k: None,
1272 frequency_penalty: None,
1273 presence_penalty: None,
1274 stop_sequences: None,
1275 timeout_seconds: config.timeout_seconds,
1276 reasoning: config.reasoning,
1277 reasoning_effort: config.reasoning_effort.clone(),
1278 reasoning_budget_tokens: config.reasoning_budget_tokens,
1279 extra: config.extra.clone(),
1280 };
1281 let base_url = config.base_url.clone().or_else(|| {
1282 config
1283 .extra
1284 .get("base_url")
1285 .and_then(Value::as_str)
1286 .map(str::to_string)
1287 });
1288 let api_key = config
1289 .api_key_env
1290 .as_ref()
1291 .and_then(|env| std::env::var(env).ok());
1292 let mut provider = UnifiedLLMProvider::from_spec_config(
1293 provider_type,
1294 &config.model,
1295 api_key,
1296 base_url,
1297 core_config,
1298 )
1299 .map_err(|error| EvalError::Runtime(error.to_string()))?;
1300 if let Some(value) = config.function_calling {
1301 provider = provider.with_feature_override(LLMFeature::FunctionCalling, value);
1302 }
1303 if let Some(choice) = config.tool_choice.clone() {
1304 provider = provider.with_tool_choice(choice);
1305 }
1306 if let Some(value) = config.vision {
1307 provider = provider.with_feature_override(LLMFeature::Vision, value);
1308 }
1309 if let Some(value) = config.json_mode {
1310 provider = provider.with_feature_override(LLMFeature::JsonMode, value);
1311 }
1312 Ok(Arc::new(provider))
1313}
1314
1315fn load_fixture_responses_for_alias(
1316 config: &LlmFixtureConfig,
1317 base_dir: &Path,
1318 alias: &str,
1319) -> Result<Vec<LLMResponse>> {
1320 let configured = config
1321 .responses_by_alias
1322 .get(alias)
1323 .unwrap_or(&config.responses);
1324 let mut responses = Vec::new();
1325 for content in configured {
1326 responses.push(LLMResponse::new(content.clone(), FinishReason::Stop));
1327 }
1328 for record in load_cassette_records(config, base_dir)? {
1329 if record.alias == alias {
1330 responses.push(record.response);
1331 }
1332 }
1333 if responses.is_empty() {
1334 responses.push(LLMResponse::new("Mock response", FinishReason::Stop));
1335 }
1336 Ok(responses)
1337}
1338
1339fn load_fixture_outcomes_for_alias(
1340 config: &LlmFixtureConfig,
1341 alias: &str,
1342 responses: &[LLMResponse],
1343) -> Vec<SequenceOutcome> {
1344 if let Some(outcomes) = config.outcomes_by_alias.get(alias)
1345 && !outcomes.is_empty()
1346 {
1347 return outcomes
1348 .iter()
1349 .map(|outcome| match outcome {
1350 LlmFixtureOutcome::Response { content } => {
1351 SequenceOutcome::Response(LLMResponse::new(content.clone(), FinishReason::Stop))
1352 }
1353 LlmFixtureOutcome::Error { message, status } => SequenceOutcome::Error {
1354 message: message.clone(),
1355 status: *status,
1356 },
1357 })
1358 .collect();
1359 }
1360 if let Some(message) = config.errors_by_alias.get(alias) {
1361 return vec![SequenceOutcome::Error {
1362 message: message.clone(),
1363 status: None,
1364 }];
1365 }
1366 responses
1367 .iter()
1368 .cloned()
1369 .map(SequenceOutcome::Response)
1370 .collect()
1371}
1372
1373fn load_cassette_records(
1374 config: &LlmFixtureConfig,
1375 base_dir: &Path,
1376) -> Result<Vec<CassetteRecord>> {
1377 let mut records = Vec::new();
1378 if let Some(path) = &config.cassette {
1379 let resolved = resolve_path(base_dir, path);
1380 if resolved.exists() {
1381 let content = std::fs::read_to_string(&resolved)?;
1382 for line in content.lines().filter(|line| !line.trim().is_empty()) {
1383 records.push(serde_json::from_str(line)?);
1384 }
1385 }
1386 }
1387 Ok(records)
1388}
1389
1390#[derive(Debug, Clone, Serialize, Deserialize)]
1392struct CassetteRecord {
1393 alias: String,
1395 model: String,
1397 request_hash: String,
1399 #[serde(default)]
1401 request_hash_version: Option<String>,
1402 response: LLMResponse,
1404}
1405
1406#[derive(Clone)]
1407enum SequenceOutcome {
1408 Response(LLMResponse),
1409 Error {
1410 message: String,
1411 status: Option<u16>,
1412 },
1413}
1414
1415struct SequenceLLMProvider {
1417 outcomes: Arc<Vec<SequenceOutcome>>,
1419 index: Mutex<usize>,
1421 delay_ms: u64,
1423 tool_choice: Option<ToolChoice>,
1424}
1425
1426struct ReplayLLMProvider {
1428 alias: String,
1430 model: String,
1432 records: Arc<Vec<CassetteRecord>>,
1434 delay_ms: u64,
1436 tool_choice: Option<ToolChoice>,
1437}
1438
1439impl SequenceLLMProvider {
1440 fn new(outcomes: Vec<SequenceOutcome>, delay_ms: u64) -> Self {
1441 Self {
1442 outcomes: Arc::new(outcomes),
1443 index: Mutex::new(0),
1444 delay_ms,
1445 tool_choice: None,
1446 }
1447 }
1448
1449 fn with_tool_choice(mut self, tool_choice: Option<ToolChoice>) -> Self {
1450 self.tool_choice = tool_choice;
1451 self
1452 }
1453
1454 async fn wait_if_configured(&self) {
1455 if self.delay_ms > 0 {
1456 tokio::time::sleep(Duration::from_millis(self.delay_ms)).await;
1457 }
1458 }
1459
1460 fn next_outcome(&self) -> std::result::Result<LLMResponse, LLMError> {
1461 let mut index = self.index.lock();
1462 let outcome = self.outcomes.get(*index).or_else(|| self.outcomes.last());
1463 if *index + 1 < self.outcomes.len() {
1464 *index += 1;
1465 }
1466 match outcome {
1467 Some(SequenceOutcome::Response(response)) => Ok(response.clone()),
1468 Some(SequenceOutcome::Error { message, status }) => Err(LLMError::API {
1469 message: message.clone(),
1470 status: *status,
1471 }),
1472 None => Ok(LLMResponse::new("Mock response", FinishReason::Stop)),
1473 }
1474 }
1475}
1476
1477impl ReplayLLMProvider {
1478 fn new(alias: String, model: String, records: Vec<CassetteRecord>, delay_ms: u64) -> Self {
1479 Self {
1480 alias,
1481 model,
1482 records: Arc::new(records),
1483 delay_ms,
1484 tool_choice: None,
1485 }
1486 }
1487
1488 fn with_tool_choice(mut self, tool_choice: Option<ToolChoice>) -> Self {
1489 self.tool_choice = tool_choice;
1490 self
1491 }
1492
1493 async fn wait_if_configured(&self) {
1494 if self.delay_ms > 0 {
1495 tokio::time::sleep(Duration::from_millis(self.delay_ms)).await;
1496 }
1497 }
1498
1499 fn response_for(
1500 &self,
1501 messages: &[ChatMessage],
1502 config: Option<&LLMConfig>,
1503 ) -> std::result::Result<LLMResponse, LLMError> {
1504 let request_hash = hash_request(messages, config);
1505 self.response_for_hash(request_hash)
1506 }
1507
1508 fn response_for_tools(
1509 &self,
1510 messages: &[ChatMessage],
1511 config: Option<&LLMConfig>,
1512 request: &LLMToolRequest,
1513 ) -> std::result::Result<LLMResponse, LLMError> {
1514 let request_hash = hash_tool_request(messages, config, request);
1515 self.response_for_hash(request_hash)
1516 }
1517
1518 fn response_for_hash(
1519 &self,
1520 request_hash: String,
1521 ) -> std::result::Result<LLMResponse, LLMError> {
1522 self.records
1523 .iter()
1524 .find(|record| {
1525 record.alias == self.alias
1526 && record.request_hash == request_hash
1527 && (record.model == self.model || record.model.is_empty())
1528 })
1529 .map(|record| record.response.clone())
1530 .ok_or_else(|| LLMError::API {
1531 message: format!(
1532 "replay cassette miss for alias '{}' model '{}' request {}",
1533 self.alias, self.model, request_hash
1534 ),
1535 status: None,
1536 })
1537 }
1538}
1539
1540#[async_trait]
1541impl LLMProvider for ReplayLLMProvider {
1542 async fn complete(
1543 &self,
1544 messages: &[ChatMessage],
1545 config: Option<&LLMConfig>,
1546 ) -> std::result::Result<LLMResponse, LLMError> {
1547 self.wait_if_configured().await;
1548 self.response_for(messages, config)
1549 }
1550
1551 async fn complete_with_tools(
1552 &self,
1553 messages: &[ChatMessage],
1554 config: Option<&LLMConfig>,
1555 request: &LLMToolRequest,
1556 ) -> std::result::Result<LLMResponse, LLMError> {
1557 self.wait_if_configured().await;
1558 self.response_for_tools(messages, config, request)
1559 }
1560
1561 fn configured_tool_choice(&self) -> Option<ToolChoice> {
1562 self.tool_choice.clone()
1563 }
1564
1565 fn supports_tool_choice(&self, choice: &ToolChoice) -> bool {
1566 matches!(
1567 choice,
1568 ToolChoice::Auto | ToolChoice::Required | ToolChoice::Specific(_) | ToolChoice::None
1569 )
1570 }
1571
1572 async fn complete_stream(
1573 &self,
1574 messages: &[ChatMessage],
1575 config: Option<&LLMConfig>,
1576 ) -> std::result::Result<
1577 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
1578 LLMError,
1579 > {
1580 self.wait_if_configured().await;
1581 let response = self.response_for(messages, config)?;
1582 Ok(Box::new(futures::stream::iter(chunks_from_response(
1583 response,
1584 ))))
1585 }
1586
1587 fn provider_name(&self) -> &str {
1588 "eval-replay"
1589 }
1590
1591 fn supports(&self, feature: LLMFeature) -> bool {
1592 matches!(feature, LLMFeature::Streaming | LLMFeature::SystemMessages)
1593 }
1594}
1595
1596#[async_trait]
1597impl LLMProvider for SequenceLLMProvider {
1598 async fn complete(
1599 &self,
1600 _messages: &[ChatMessage],
1601 _config: Option<&LLMConfig>,
1602 ) -> std::result::Result<LLMResponse, LLMError> {
1603 self.wait_if_configured().await;
1604 self.next_outcome()
1605 }
1606
1607 fn configured_tool_choice(&self) -> Option<ToolChoice> {
1608 self.tool_choice.clone()
1609 }
1610
1611 async fn complete_stream(
1612 &self,
1613 _messages: &[ChatMessage],
1614 _config: Option<&LLMConfig>,
1615 ) -> std::result::Result<
1616 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
1617 LLMError,
1618 > {
1619 self.wait_if_configured().await;
1620 let response = self.next_outcome()?;
1621 Ok(Box::new(futures::stream::iter(chunks_from_response(
1622 response,
1623 ))))
1624 }
1625
1626 fn provider_name(&self) -> &str {
1627 "eval-sequence"
1628 }
1629
1630 fn supports(&self, feature: LLMFeature) -> bool {
1631 matches!(feature, LLMFeature::Streaming | LLMFeature::SystemMessages)
1632 }
1633}
1634
1635fn chunks_from_response(response: LLMResponse) -> Vec<std::result::Result<LLMChunk, LLMError>> {
1636 let deltas = split_stream_content(&response.content);
1637 if deltas.is_empty() {
1638 return vec![Ok(LLMChunk::final_chunk(
1639 "",
1640 response.finish_reason,
1641 response.usage,
1642 ))];
1643 }
1644 let last_index = deltas.len() - 1;
1645 deltas
1646 .into_iter()
1647 .enumerate()
1648 .map(|(index, delta)| {
1649 if index == last_index {
1650 Ok(LLMChunk::final_chunk(
1651 delta,
1652 response.finish_reason.clone(),
1653 response.usage,
1654 ))
1655 } else {
1656 Ok(LLMChunk::new(delta, false))
1657 }
1658 })
1659 .collect()
1660}
1661
1662fn split_stream_content(content: &str) -> Vec<String> {
1663 if content.is_empty() {
1664 return Vec::new();
1665 }
1666 let words: Vec<&str> = content.split_whitespace().collect();
1667 if words.len() > 1 {
1668 return words
1669 .into_iter()
1670 .enumerate()
1671 .map(|(index, word)| {
1672 if index == 0 {
1673 word.to_string()
1674 } else {
1675 format!(" {}", word)
1676 }
1677 })
1678 .collect();
1679 }
1680 let chars: Vec<char> = content.chars().collect();
1681 if chars.len() <= 1 {
1682 return vec![content.to_string()];
1683 }
1684 let chunk_size = 4.min(chars.len().div_ceil(2));
1685 chars
1686 .chunks(chunk_size)
1687 .map(|chunk| chunk.iter().collect::<String>())
1688 .collect()
1689}
1690
1691struct CassetteWriter {
1693 path: PathBuf,
1694 file: Mutex<File>,
1695}
1696
1697impl CassetteWriter {
1698 fn shared(path: PathBuf) -> Result<Arc<Self>> {
1699 static WRITERS: OnceLock<Mutex<HashMap<PathBuf, Weak<CassetteWriter>>>> = OnceLock::new();
1700 let writers = WRITERS.get_or_init(|| Mutex::new(HashMap::new()));
1701 let mut writers = writers.lock();
1702 if let Some(writer) = writers.get(&path).and_then(Weak::upgrade) {
1703 return Ok(writer);
1704 }
1705 let writer = Arc::new(Self::open(path.clone())?);
1706 writers.insert(path, Arc::downgrade(&writer));
1707 Ok(writer)
1708 }
1709
1710 fn open(path: PathBuf) -> Result<Self> {
1711 if let Some(parent) = path.parent()
1712 && !parent.as_os_str().is_empty()
1713 {
1714 std::fs::create_dir_all(parent).map_err(|error| {
1715 EvalError::Runtime(format!(
1716 "failed to prepare LLM cassette '{}': {error}",
1717 path.display()
1718 ))
1719 })?;
1720 }
1721 let file = OpenOptions::new()
1722 .create(true)
1723 .append(true)
1724 .open(&path)
1725 .map_err(|error| {
1726 EvalError::Runtime(format!(
1727 "failed to open LLM cassette '{}': {error}",
1728 path.display()
1729 ))
1730 })?;
1731 Ok(Self {
1732 path,
1733 file: Mutex::new(file),
1734 })
1735 }
1736
1737 fn append(&self, record: &CassetteRecord) -> std::result::Result<(), LLMError> {
1738 let encoded = serde_json::to_string(record).map_err(|error| self.write_error(error))?;
1739 let mut file = self.file.lock();
1740 writeln!(file, "{encoded}").map_err(|error| self.write_error(error))?;
1741 file.flush().map_err(|error| self.write_error(error))
1742 }
1743
1744 fn write_error(&self, error: impl std::fmt::Display) -> LLMError {
1745 LLMError::API {
1746 message: format!(
1747 "failed to write LLM cassette '{}': {error}",
1748 self.path.display()
1749 ),
1750 status: None,
1751 }
1752 }
1753}
1754
1755struct RecordingLLMProvider {
1757 inner: Arc<dyn LLMProvider>,
1759 alias: String,
1761 model: String,
1763 writer: Arc<CassetteWriter>,
1765}
1766
1767impl RecordingLLMProvider {
1768 fn new(
1769 inner: Arc<dyn LLMProvider>,
1770 alias: String,
1771 model: String,
1772 writer: Arc<CassetteWriter>,
1773 ) -> Self {
1774 Self {
1775 inner,
1776 alias,
1777 model,
1778 writer,
1779 }
1780 }
1781}
1782
1783#[async_trait]
1784impl LLMProvider for RecordingLLMProvider {
1785 async fn complete(
1786 &self,
1787 messages: &[ChatMessage],
1788 config: Option<&LLMConfig>,
1789 ) -> std::result::Result<LLMResponse, LLMError> {
1790 let response = self.inner.complete(messages, config).await?;
1791 let record = CassetteRecord {
1792 alias: self.alias.clone(),
1793 model: self.model.clone(),
1794 request_hash: hash_request(messages, config),
1795 request_hash_version: Some("sha256-v1".to_string()),
1796 response: response.clone(),
1797 };
1798 self.writer.append(&record)?;
1799 Ok(response)
1800 }
1801
1802 async fn complete_with_tools(
1803 &self,
1804 messages: &[ChatMessage],
1805 config: Option<&LLMConfig>,
1806 request: &LLMToolRequest,
1807 ) -> std::result::Result<LLMResponse, LLMError> {
1808 let response = self
1809 .inner
1810 .complete_with_tools(messages, config, request)
1811 .await?;
1812 let record = CassetteRecord {
1813 alias: self.alias.clone(),
1814 model: self.model.clone(),
1815 request_hash: hash_tool_request(messages, config, request),
1816 request_hash_version: Some("sha256-v2-tools".to_string()),
1817 response: response.clone(),
1818 };
1819 self.writer.append(&record)?;
1820 Ok(response)
1821 }
1822
1823 fn configured_tool_choice(&self) -> Option<ToolChoice> {
1824 self.inner.configured_tool_choice()
1825 }
1826
1827 fn supports_tool_choice(&self, choice: &ToolChoice) -> bool {
1828 self.inner.supports_tool_choice(choice)
1829 }
1830
1831 async fn complete_stream(
1832 &self,
1833 messages: &[ChatMessage],
1834 config: Option<&LLMConfig>,
1835 ) -> std::result::Result<
1836 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
1837 LLMError,
1838 > {
1839 let _ = (messages, config);
1840 Err(LLMError::API {
1841 message: "record mode does not support streaming because the stream cannot be recorded atomically"
1842 .to_string(),
1843 status: None,
1844 })
1845 }
1846
1847 fn provider_name(&self) -> &str {
1848 self.inner.provider_name()
1849 }
1850
1851 fn supports(&self, feature: LLMFeature) -> bool {
1852 feature != LLMFeature::Streaming && self.inner.supports(feature)
1853 }
1854}
1855
1856fn hash_request(messages: &[ChatMessage], config: Option<&LLMConfig>) -> String {
1857 let canonical_messages: Vec<Value> = messages
1858 .iter()
1859 .map(|message| {
1860 json!({
1861 "role": format!("{:?}", message.role),
1862 "content": message.content,
1863 "name": message.name,
1864 })
1865 })
1866 .collect();
1867 let canonical = json!({
1868 "version": "sha256-v1",
1869 "messages": canonical_messages,
1870 "config": config,
1871 });
1872 let encoded = serde_json::to_vec(&canonical).unwrap_or_default();
1873 let digest = Sha256::digest(encoded);
1874 format!("sha256-v1:{:x}", digest)
1875}
1876
1877fn hash_tool_request(
1878 messages: &[ChatMessage],
1879 config: Option<&LLMConfig>,
1880 request: &LLMToolRequest,
1881) -> String {
1882 let canonical = json!({
1883 "version": "sha256-v2-tools",
1884 "base_request": hash_request(messages, config),
1885 "tool_request": request,
1886 });
1887 let encoded = serde_json::to_vec(&canonical).unwrap_or_default();
1888 let digest = Sha256::digest(encoded);
1889 format!("sha256-v2-tools:{:x}", digest)
1890}
1891
1892fn default_status() -> u16 {
1893 200
1894}
1895
1896fn default_true() -> bool {
1897 true
1898}
1899
1900#[cfg(test)]
1901mod tests {
1902 use super::*;
1903 use std::sync::atomic::{AtomicUsize, Ordering};
1904
1905 struct CountingProvider {
1906 calls: Arc<AtomicUsize>,
1907 }
1908
1909 #[async_trait]
1910 impl LLMProvider for CountingProvider {
1911 async fn complete(
1912 &self,
1913 _messages: &[ChatMessage],
1914 _config: Option<&LLMConfig>,
1915 ) -> std::result::Result<LLMResponse, LLMError> {
1916 self.calls.fetch_add(1, Ordering::SeqCst);
1917 Ok(LLMResponse::new("counted response", FinishReason::Stop))
1918 }
1919
1920 async fn complete_stream(
1921 &self,
1922 _messages: &[ChatMessage],
1923 _config: Option<&LLMConfig>,
1924 ) -> std::result::Result<
1925 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
1926 LLMError,
1927 > {
1928 self.calls.fetch_add(1, Ordering::SeqCst);
1929 Err(LLMError::API {
1930 message: "unexpected stream call".to_string(),
1931 status: None,
1932 })
1933 }
1934
1935 fn provider_name(&self) -> &str {
1936 "counting"
1937 }
1938
1939 fn supports(&self, _feature: LLMFeature) -> bool {
1940 true
1941 }
1942 }
1943
1944 #[test]
1945 fn runtime_tool_sources_preserve_existing_labels_and_distinguish_plan() {
1946 use ai_agents_core::ToolCallSource;
1947
1948 let mappings = [
1949 (ToolCallSource::Model, ToolExecutionSource::Llm),
1950 (
1951 ToolCallSource::Skill {
1952 skill_id: "skill".to_string(),
1953 step_index: 0,
1954 },
1955 ToolExecutionSource::Skill,
1956 ),
1957 (
1958 ToolCallSource::Plan { step_index: 0 },
1959 ToolExecutionSource::Plan,
1960 ),
1961 (
1962 ToolCallSource::StateAction {
1963 state: Some("ready".to_string()),
1964 action_index: 0,
1965 },
1966 ToolExecutionSource::StateAction,
1967 ),
1968 (
1969 ToolCallSource::Orchestration,
1970 ToolExecutionSource::Orchestration,
1971 ),
1972 (ToolCallSource::Spawner, ToolExecutionSource::Spawner),
1973 (
1974 ToolCallSource::Fallback {
1975 original_tool: "missing".to_string(),
1976 },
1977 ToolExecutionSource::Llm,
1978 ),
1979 (ToolCallSource::Task, ToolExecutionSource::Llm),
1980 (ToolCallSource::Manual, ToolExecutionSource::Llm),
1981 (ToolCallSource::EvalFixture, ToolExecutionSource::Mock),
1982 ];
1983
1984 for (source, expected) in mappings {
1985 assert_eq!(ToolExecutionSource::from(&source), expected);
1986 }
1987 assert_eq!(
1988 serde_json::to_value(ToolExecutionSource::Plan).unwrap(),
1989 json!("plan")
1990 );
1991 }
1992
1993 #[test]
1994 fn llm_fixture_parses_tagged_outcomes_by_alias() {
1995 let config: FixturesConfig = serde_yaml::from_str(
1996 r#"
1997llm:
1998 mode: mock
1999 outcomes_by_alias:
2000 worker:
2001 - type: error
2002 message: transient overload
2003 - type: response
2004 content: recovered
2005"#,
2006 )
2007 .unwrap();
2008
2009 assert_eq!(
2010 config.llm.outcomes_by_alias["worker"],
2011 vec![
2012 LlmFixtureOutcome::Error {
2013 message: "transient overload".to_string(),
2014 status: None,
2015 },
2016 LlmFixtureOutcome::Response {
2017 content: "recovered".to_string(),
2018 },
2019 ]
2020 );
2021 }
2022
2023 #[tokio::test]
2024 async fn sequence_provider_mixes_outcomes_with_independent_alias_cursors() {
2025 let config: LlmFixtureConfig = serde_yaml::from_str(
2026 r#"
2027mode: mock
2028outcomes_by_alias:
2029 alpha:
2030 - type: error
2031 message: retry alpha
2032 - type: response
2033 content: alpha recovered
2034 beta:
2035 - type: response
2036 content: beta first
2037 - type: response
2038 content: beta second
2039"#,
2040 )
2041 .unwrap();
2042 let fallback = vec![LLMResponse::new("fallback", FinishReason::Stop)];
2043 let alpha = SequenceLLMProvider::new(
2044 load_fixture_outcomes_for_alias(&config, "alpha", &fallback),
2045 0,
2046 );
2047 let beta = SequenceLLMProvider::new(
2048 load_fixture_outcomes_for_alias(&config, "beta", &fallback),
2049 0,
2050 );
2051
2052 assert!(matches!(
2053 alpha.complete(&[], None).await,
2054 Err(LLMError::API { message, status: None }) if message == "retry alpha"
2055 ));
2056 assert_eq!(
2057 beta.complete(&[], None).await.unwrap().content,
2058 "beta first"
2059 );
2060 assert_eq!(
2061 alpha.complete(&[], None).await.unwrap().content,
2062 "alpha recovered"
2063 );
2064 assert_eq!(
2065 beta.complete(&[], None).await.unwrap().content,
2066 "beta second"
2067 );
2068 assert_eq!(
2069 alpha.complete(&[], None).await.unwrap().content,
2070 "alpha recovered"
2071 );
2072 assert_eq!(
2073 beta.complete(&[], None).await.unwrap().content,
2074 "beta second"
2075 );
2076 }
2077
2078 #[tokio::test]
2079 async fn sequence_provider_stream_errors_advance_to_recovery() {
2080 let provider = SequenceLLMProvider::new(
2081 vec![
2082 SequenceOutcome::Error {
2083 message: "stream retry".to_string(),
2084 status: None,
2085 },
2086 SequenceOutcome::Response(LLMResponse::new("stream recovered", FinishReason::Stop)),
2087 ],
2088 0,
2089 );
2090
2091 assert!(matches!(
2092 provider.complete_stream(&[], None).await,
2093 Err(LLMError::API { message, status: None }) if message == "stream retry"
2094 ));
2095 assert_eq!(
2096 provider.complete(&[], None).await.unwrap().content,
2097 "stream recovered"
2098 );
2099 }
2100
2101 #[tokio::test]
2102 async fn sequence_provider_preserves_legacy_responses_and_errors() {
2103 let config = LlmFixtureConfig {
2104 mode: LlmFixtureMode::Mock,
2105 responses: vec!["global".to_string()],
2106 responses_by_alias: HashMap::from([(
2107 "worker".to_string(),
2108 vec!["first".to_string(), "second".to_string()],
2109 )]),
2110 errors_by_alias: HashMap::from([(
2111 "failing".to_string(),
2112 "permanent failure".to_string(),
2113 )]),
2114 ..Default::default()
2115 };
2116 let worker_responses =
2117 load_fixture_responses_for_alias(&config, Path::new("."), "worker").unwrap();
2118 let failing_responses =
2119 load_fixture_responses_for_alias(&config, Path::new("."), "failing").unwrap();
2120 let worker = SequenceLLMProvider::new(
2121 load_fixture_outcomes_for_alias(&config, "worker", &worker_responses),
2122 0,
2123 );
2124 let failing = SequenceLLMProvider::new(
2125 load_fixture_outcomes_for_alias(&config, "failing", &failing_responses),
2126 0,
2127 );
2128
2129 assert_eq!(worker.complete(&[], None).await.unwrap().content, "first");
2130 assert_eq!(worker.complete(&[], None).await.unwrap().content, "second");
2131 assert_eq!(worker.complete(&[], None).await.unwrap().content, "second");
2132 for _ in 0..2 {
2133 assert!(matches!(
2134 failing.complete(&[], None).await,
2135 Err(LLMError::API { message, status: None }) if message == "permanent failure"
2136 ));
2137 }
2138 }
2139
2140 #[test]
2141 fn approval_fixture_parses_all_triggers_and_outcomes() {
2142 let config: FixturesConfig = serde_yaml::from_str(
2143 r#"
2144approvals:
2145 preferred_language: en
2146 supported_languages: [en, fr]
2147 rules:
2148 - trigger: { type: tool, name: transfer, args: { amount: 10 } }
2149 occurrence: 2
2150 outcome: approve
2151 - trigger: { type: condition, name: high_value, matched: "amount > 100" }
2152 outcome: reject
2153 reason: too expensive
2154 - trigger: { type: state_transition, from: review, to: complete }
2155 outcome: modify
2156 changes: { amount: 5 }
2157 - trigger: { type: disambiguation_escalation, reason: unclear intent }
2158 outcome: timeout
2159 default:
2160 outcome: unavailable
2161"#,
2162 )
2163 .unwrap();
2164 let approvals = config.approvals.unwrap();
2165
2166 assert_eq!(approvals.rules.len(), 4);
2167 assert_eq!(approvals.rules[0].occurrence.unwrap().get(), 2);
2168 assert_eq!(approvals.preferred_language.as_deref(), Some("en"));
2169 assert_eq!(
2170 approvals.supported_languages,
2171 Some(vec!["en".to_string(), "fr".to_string()])
2172 );
2173 assert_eq!(approvals.default, ApprovalFixtureOutcome::Unavailable);
2174 }
2175
2176 #[test]
2177 fn approval_fixture_rejects_zero_occurrence() {
2178 let error = serde_yaml::from_str::<ApprovalFixtureConfig>(
2179 r#"
2180rules:
2181 - trigger: { type: tool, name: transfer }
2182 occurrence: 0
2183 outcome: approve
2184"#,
2185 )
2186 .unwrap_err();
2187
2188 assert!(error.to_string().contains("nonzero"));
2189 }
2190
2191 #[tokio::test]
2192 async fn approval_handler_matches_order_occurrence_and_default() {
2193 let config = ApprovalFixtureConfig {
2194 rules: vec![
2195 ApprovalFixtureRule {
2196 trigger: ApprovalFixtureTrigger::Tool {
2197 name: "transfer".to_string(),
2198 args: Some(json!({"amount": 10})),
2199 },
2200 occurrence: NonZeroUsize::new(2),
2201 outcome: ApprovalFixtureOutcome::Modify {
2202 changes: HashMap::from([("amount".to_string(), json!(5))]),
2203 },
2204 },
2205 ApprovalFixtureRule {
2206 trigger: ApprovalFixtureTrigger::Tool {
2207 name: "transfer".to_string(),
2208 args: None,
2209 },
2210 occurrence: None,
2211 outcome: ApprovalFixtureOutcome::Reject {
2212 reason: Some("fallback rule".to_string()),
2213 },
2214 },
2215 ],
2216 default: ApprovalFixtureOutcome::Timeout,
2217 preferred_language: Some("ja".to_string()),
2218 supported_languages: Some(vec!["ja".to_string(), "en".to_string()]),
2219 };
2220 let handler = FixtureApprovalHandler::new(config);
2221 let request = || {
2222 ApprovalRequest::new(
2223 ApprovalTrigger::tool("transfer", json!({"amount": 10})),
2224 "Approve transfer?",
2225 )
2226 };
2227
2228 assert!(matches!(
2229 handler.request_approval(request()).await,
2230 ApprovalResult::Rejected { reason: Some(reason) } if reason == "fallback rule"
2231 ));
2232 assert!(matches!(
2233 handler.request_approval(request()).await,
2234 ApprovalResult::Modified { changes } if changes.get("amount") == Some(&json!(5))
2235 ));
2236 assert!(matches!(
2237 handler
2238 .request_approval(ApprovalRequest::new(
2239 ApprovalTrigger::tool("email", json!({})),
2240 "Approve email?",
2241 ))
2242 .await,
2243 ApprovalResult::Timeout
2244 ));
2245 assert_eq!(handler.preferred_language(), Some("ja".to_string()));
2246 assert_eq!(
2247 handler.supported_languages(),
2248 Some(vec!["ja".to_string(), "en".to_string()])
2249 );
2250 }
2251
2252 #[tokio::test]
2253 async fn approval_handler_matches_condition_state_and_escalation() {
2254 let config: ApprovalFixtureConfig = serde_yaml::from_str(
2255 r#"
2256rules:
2257 - trigger: { type: condition, name: high_value, matched: threshold }
2258 outcome: approve
2259 - trigger: { type: state_transition, from: review, to: complete }
2260 outcome: reject
2261 - trigger: { type: disambiguation_escalation, reason: unclear }
2262 outcome: timeout
2263default:
2264 outcome: unavailable
2265"#,
2266 )
2267 .unwrap();
2268 let handler = FixtureApprovalHandler::new(config);
2269
2270 let condition = handler
2271 .request_approval(ApprovalRequest::new(
2272 ApprovalTrigger::condition("high_value", "threshold"),
2273 "condition",
2274 ))
2275 .await;
2276 let state = handler
2277 .request_approval(ApprovalRequest::new(
2278 ApprovalTrigger::state(Some("review".to_string()), "complete"),
2279 "state",
2280 ))
2281 .await;
2282 let escalation = handler
2283 .request_approval(ApprovalRequest::new(
2284 ApprovalTrigger::condition("disambiguation_escalation", "unclear"),
2285 "escalation",
2286 ))
2287 .await;
2288 let unavailable = handler
2289 .request_approval(ApprovalRequest::new(
2290 ApprovalTrigger::condition("disambiguation_escalation", "other"),
2291 "other escalation",
2292 ))
2293 .await;
2294
2295 assert!(matches!(condition, ApprovalResult::Approved));
2296 assert!(matches!(state, ApprovalResult::Rejected { reason: None }));
2297 assert!(matches!(escalation, ApprovalResult::Timeout));
2298 assert!(matches!(
2299 unavailable,
2300 ApprovalResult::Rejected { reason: Some(reason) }
2301 if reason == "Approval fixture unavailable"
2302 ));
2303 }
2304
2305 #[tokio::test]
2306 async fn approval_handler_instances_have_fresh_concurrent_state() {
2307 let config: ApprovalFixtureConfig = serde_yaml::from_str(
2308 r#"
2309rules:
2310 - trigger: { type: tool, name: transfer }
2311 occurrence: 2
2312 outcome: approve
2313default:
2314 outcome: reject
2315"#,
2316 )
2317 .unwrap();
2318 let first = build_approval_handler(&config);
2319 let second = build_approval_handler(&config);
2320 let request =
2321 || ApprovalRequest::new(ApprovalTrigger::tool("transfer", json!({})), "transfer");
2322
2323 let (first_result, second_result) = tokio::join!(
2324 first.request_approval(request()),
2325 first.request_approval(request())
2326 );
2327 let approved =
2328 usize::from(first_result.is_approved()) + usize::from(second_result.is_approved());
2329 assert_eq!(approved, 1);
2330 assert!(second.request_approval(request()).await.is_rejected());
2331 assert!(second.request_approval(request()).await.is_approved());
2332 }
2333
2334 #[test]
2335 fn workspace_policy_is_optional_and_deserializes_narrow_tool_lists() {
2336 let absent: FixturesConfig = serde_yaml::from_str("llm: { mode: mock }").unwrap();
2337 assert_eq!(absent.workspace_policy, None);
2338
2339 let configured: FixturesConfig = serde_yaml::from_str(
2340 r#"
2341workspace_policy:
2342 read_tools: [file_read]
2343 write_tools: [file_write]
2344"#,
2345 )
2346 .unwrap();
2347 assert_eq!(
2348 configured.workspace_policy,
2349 Some(WorkspacePolicyFixtureConfig {
2350 read_tools: vec!["file_read".to_string()],
2351 write_tools: vec!["file_write".to_string()],
2352 })
2353 );
2354 }
2355
2356 #[test]
2357 fn attempt_context_uses_opaque_absolute_unique_workspaces() {
2358 let scenario_id = "customer-refund-scenario";
2359 let first = AttemptFixtureContext::create(None).unwrap();
2360 let second = AttemptFixtureContext::create(None).unwrap();
2361
2362 assert!(first.workspace.is_absolute());
2363 assert!(second.workspace.is_absolute());
2364 assert_ne!(first.workspace, second.workspace);
2365 assert!(!first.workspace.display().to_string().contains(scenario_id));
2366 let suffix = first
2367 .workspace
2368 .file_name()
2369 .unwrap()
2370 .to_string_lossy()
2371 .trim_start_matches("ai_agents_eval_")
2372 .to_string();
2373 assert_eq!(uuid::Uuid::parse_str(&suffix).unwrap().to_string(), suffix);
2374 assert_eq!(
2375 first.runtime_context()["eval"]["workspace"],
2376 json!(first.workspace.display().to_string())
2377 );
2378
2379 std::fs::remove_dir_all(&first.workspace).unwrap();
2380 std::fs::remove_dir_all(&second.workspace).unwrap();
2381 }
2382
2383 #[test]
2384 fn owned_attempt_workspace_cleans_up_after_drop_and_early_error() {
2385 let workspace = AttemptWorkspace::create(None).unwrap();
2386 let success_path = workspace.workspace.clone();
2387 std::fs::write(success_path.join("artifact.txt"), "temporary").unwrap();
2388 drop(workspace);
2389 assert!(!success_path.exists());
2390
2391 let captured = Arc::new(Mutex::new(None));
2392 let result: Result<()> = (|| {
2393 let workspace = AttemptWorkspace::create(None)?;
2394 *captured.lock() = Some(workspace.workspace.clone());
2395 Err(EvalError::Config("intentional early error".to_string()))
2396 })();
2397 assert!(result.is_err());
2398 assert!(!captured.lock().as_ref().unwrap().exists());
2399 }
2400
2401 #[tokio::test]
2402 async fn owned_attempt_workspace_cleans_up_when_future_is_cancelled() {
2403 let captured = Arc::new(Mutex::new(None));
2404 let future_captured = captured.clone();
2405 let pending = async move {
2406 let workspace = AttemptWorkspace::create(None).unwrap();
2407 *future_captured.lock() = Some(workspace.workspace.clone());
2408 std::future::pending::<()>().await;
2409 drop(workspace);
2410 };
2411
2412 assert!(
2413 tokio::time::timeout(Duration::from_millis(10), pending)
2414 .await
2415 .is_err()
2416 );
2417 assert!(!captured.lock().as_ref().unwrap().exists());
2418 }
2419
2420 #[test]
2421 fn attempt_context_interpolates_only_exact_tokens_json_safely() {
2422 let context = AttemptFixtureContext {
2423 isolation_id: "attempt-one".to_string(),
2424 workspace: PathBuf::from(r"C:\eval\attempt"),
2425 mock_server_base_url: Some("http://127.0.0.1:43123".to_string()),
2426 };
2427 let config = LlmFixtureConfig {
2428 mode: LlmFixtureMode::Mock,
2429 responses: vec![
2430 r#"{"{{ eval.workspace }}":{"workspace":"{{ eval.workspace }}"},"nested":{"workspace":"{{ eval.workspace }}"},"items":["{{ mock_server.base_url }}/ok",7]}"#.to_string(),
2431 "keep {{ unrelated.template }} and {{eval.workspace}}".to_string(),
2432 ],
2433 responses_by_alias: HashMap::from([(
2434 "worker".to_string(),
2435 vec!["write to {{ eval.workspace }}".to_string()],
2436 )]),
2437 outcomes_by_alias: HashMap::from([(
2438 "reviewer".to_string(),
2439 vec![
2440 LlmFixtureOutcome::Error {
2441 message: "keep {{ eval.workspace }} literal".to_string(),
2442 status: None,
2443 },
2444 LlmFixtureOutcome::Response {
2445 content: "review {{ eval.workspace }}".to_string(),
2446 },
2447 ],
2448 )]),
2449 ..Default::default()
2450 };
2451
2452 let rewritten = context.interpolate_llm_fixture(&config).unwrap();
2453 let json: Value = serde_json::from_str(&rewritten.responses[0]).unwrap();
2454 assert_eq!(json["nested"]["workspace"], json!(r"C:\eval\attempt"));
2455 assert!(json.get(r"C:\eval\attempt").is_some());
2456 assert_eq!(json["items"][0], json!("http://127.0.0.1:43123/ok"));
2457 assert_eq!(
2458 rewritten.responses[1],
2459 "keep {{ unrelated.template }} and {{eval.workspace}}"
2460 );
2461 assert_eq!(
2462 rewritten.responses_by_alias["worker"][0],
2463 r"write to C:\eval\attempt"
2464 );
2465 assert_eq!(
2466 rewritten.outcomes_by_alias["reviewer"][0],
2467 LlmFixtureOutcome::Error {
2468 message: "keep {{ eval.workspace }} literal".to_string(),
2469 status: None,
2470 }
2471 );
2472 assert_eq!(
2473 rewritten.outcomes_by_alias["reviewer"][1],
2474 LlmFixtureOutcome::Response {
2475 content: r"review C:\eval\attempt".to_string(),
2476 }
2477 );
2478 }
2479
2480 #[test]
2481 fn attempt_context_rejects_mock_server_token_without_server() {
2482 let context = AttemptFixtureContext {
2483 isolation_id: "attempt-one".to_string(),
2484 workspace: PathBuf::from("/tmp/attempt-one"),
2485 mock_server_base_url: None,
2486 };
2487 let config = LlmFixtureConfig {
2488 mode: LlmFixtureMode::Mock,
2489 responses: vec!["{{ mock_server.base_url }}/missing".to_string()],
2490 ..Default::default()
2491 };
2492
2493 let error = context.interpolate_llm_fixture(&config).unwrap_err();
2494 assert!(
2495 error
2496 .to_string()
2497 .contains("without an enabled fixtures.mock_server")
2498 );
2499 }
2500
2501 #[test]
2502 fn context_file_and_inline_context_merge() {
2503 let dir = std::env::temp_dir().join(format!(
2504 "ai_agents_eval_fixture_test_{}",
2505 uuid::Uuid::new_v4()
2506 ));
2507 std::fs::create_dir_all(&dir).unwrap();
2508 let context_path = dir.join("context.json");
2509 std::fs::write(
2510 &context_path,
2511 r#"{"user":{"tier":"basic"},"channel":"file"}"#,
2512 )
2513 .unwrap();
2514 let config = FixturesConfig {
2515 context: Some(serde_json::json!({"channel":"inline","feature":true})),
2516 context_file: Some(PathBuf::from("context.json")),
2517 ..Default::default()
2518 };
2519 let context = resolve_fixture_context(&config, &dir).unwrap();
2520 assert_eq!(context.get("channel"), Some(&serde_json::json!("inline")));
2521 assert_eq!(context.get("feature"), Some(&serde_json::json!(true)));
2522 assert!(context.contains_key("user"));
2523 let _ = std::fs::remove_dir_all(dir);
2524 }
2525
2526 #[tokio::test]
2527 async fn web_fetch_transport_uses_real_tool_without_sockets() {
2528 let fixtures: FixturesConfig = serde_yaml::from_str(
2529 r#"
2530web_fetch_transport:
2531 routes:
2532 - url: https://fixture.example/data
2533 status: 200
2534 headers:
2535 Content-Type: application/json
2536 body: { ok: true }
2537"#,
2538 )
2539 .unwrap();
2540 let registry = build_tool_registry(&fixtures, RecordingToolLog::new()).unwrap();
2541 let tool = registry.get("web_fetch").unwrap();
2542 let result = tool
2543 .execute(
2544 json!({"url":"https://fixture.example/data","cache_ttl_seconds":0}),
2545 ToolExecutionContext::test("web_fetch"),
2546 )
2547 .await;
2548
2549 assert!(result.success, "{}", result.output);
2550 let output: Value = serde_json::from_str(&result.output).unwrap();
2551 assert_eq!(output["status"], 200);
2552 assert_eq!(output["content"], "{\"ok\":true}");
2553
2554 let missing = tool
2555 .execute(
2556 json!({"url":"https://fixture.example/missing","cache_ttl_seconds":0}),
2557 ToolExecutionContext::test("web_fetch"),
2558 )
2559 .await;
2560 assert!(!missing.success);
2561 assert!(
2562 missing
2563 .output
2564 .contains("No fixtures.web_fetch_transport route")
2565 );
2566 }
2567
2568 #[tokio::test]
2569 async fn web_fetch_transport_preserves_redirect_policy_checks() {
2570 let fixtures: FixturesConfig = serde_yaml::from_str(
2571 r#"
2572web_fetch_transport:
2573 routes:
2574 - url: https://fixture.example/start
2575 status: 302
2576 headers:
2577 Location: https://blocked.example/secret
2578 - url: https://blocked.example/secret
2579 status: 200
2580 body: should-not-be-returned
2581"#,
2582 )
2583 .unwrap();
2584 let registry = build_tool_registry(&fixtures, RecordingToolLog::new()).unwrap();
2585 let tool = registry.get("web_fetch").unwrap();
2586 let mut execution_context = ToolExecutionContext::test("web_fetch");
2587 execution_context.policy_snapshot = json!({"blocked_domains":["blocked.example"]});
2588 let result = tool
2589 .execute(
2590 json!({"url":"https://fixture.example/start","cache_ttl_seconds":0}),
2591 execution_context,
2592 )
2593 .await;
2594
2595 assert!(!result.success);
2596 assert!(result.output.contains("blocked by policy"));
2597 }
2598
2599 #[test]
2600 fn mocked_builtin_preserves_tool_contract() {
2601 let fixtures = FixturesConfig {
2602 tools: HashMap::from([(
2603 "web_fetch".to_string(),
2604 ToolMockConfig {
2605 success: true,
2606 output: json!({"status": 200}),
2607 },
2608 )]),
2609 ..Default::default()
2610 };
2611 let registry = build_tool_registry(&fixtures, RecordingToolLog::new()).unwrap();
2612 let mocked = registry.get("web_fetch").unwrap();
2613 let builtin = ai_agents_tools::builtin::get_builtin_tool("web_fetch").unwrap();
2614
2615 assert_eq!(mocked.name(), builtin.name());
2616 assert_eq!(mocked.input_schema(), builtin.input_schema());
2617 assert_eq!(
2618 serde_json::to_value(mocked.safety_metadata()).unwrap(),
2619 serde_json::to_value(builtin.safety_metadata()).unwrap()
2620 );
2621 assert_eq!(mocked.policy_bindings(), builtin.policy_bindings());
2622 assert!(!mocked.policy_bindings().domain_fields.is_empty());
2623 }
2624
2625 #[test]
2626 fn mock_streaming_splits_response_into_multiple_chunks() {
2627 let response = LLMResponse::new(
2628 "Streaming hello from the mocked provider.".to_string(),
2629 FinishReason::Stop,
2630 );
2631 let chunks = chunks_from_response(response);
2632 assert!(chunks.len() > 1);
2633 let mut reconstructed = String::new();
2634 for (index, chunk) in chunks.into_iter().enumerate() {
2635 let chunk = chunk.unwrap();
2636 if index == 0 {
2637 assert!(!chunk.is_final);
2638 }
2639 reconstructed.push_str(&chunk.delta);
2640 if chunk.is_final {
2641 assert!(chunk.finish_reason.is_some());
2642 }
2643 }
2644 assert_eq!(reconstructed, "Streaming hello from the mocked provider.");
2645 }
2646
2647 #[test]
2648 fn mock_streaming_splits_single_word_response() {
2649 let chunks = split_stream_content("Hello");
2650 assert!(chunks.len() > 1);
2651 assert_eq!(chunks.join(""), "Hello");
2652 }
2653
2654 #[test]
2655 fn fixture_yaml_rejects_unknown_fields() {
2656 let error =
2657 serde_yaml::from_str::<FixturesConfig>("llm:\n mode: mock\n resposes: [fallback]\n")
2658 .unwrap_err()
2659 .to_string();
2660 assert!(error.contains("resposes"));
2661 }
2662
2663 #[tokio::test]
2664 async fn replay_miss_rejects_configured_fallback_responses() {
2665 let spec: AgentSpec = serde_yaml::from_str(
2666 r#"
2667name: ReplayAgent
2668system_prompt: test
2669llm:
2670 provider: openai
2671 model: replay-model
2672"#,
2673 )
2674 .unwrap();
2675 let fixtures = LlmFixtureConfig {
2676 mode: LlmFixtureMode::Replay,
2677 responses: vec!["must not be used".to_string()],
2678 ..Default::default()
2679 };
2680 let (registry, _) = build_llm_registry(&spec, &fixtures, Path::new(".")).unwrap();
2681 let provider = registry.default().unwrap();
2682
2683 let error = provider
2684 .complete(&[ChatMessage::user("uncassetted request")], None)
2685 .await
2686 .unwrap_err();
2687
2688 assert!(error.to_string().contains("replay cassette miss"));
2689 assert!(!error.to_string().contains("must not be used"));
2690 }
2691
2692 #[test]
2693 fn cassette_records_accept_unknown_persisted_fields() {
2694 let record: CassetteRecord = serde_json::from_value(json!({
2695 "alias": "default",
2696 "model": "model",
2697 "request_hash": "hash",
2698 "request_hash_version": "sha256-v1",
2699 "response": LLMResponse::new("response", FinishReason::Stop),
2700 "future_metadata": {"version": 2}
2701 }))
2702 .unwrap();
2703 assert_eq!(record.alias, "default");
2704 }
2705
2706 #[tokio::test]
2707 async fn recording_preflight_failure_skips_inner_provider() {
2708 let dir = std::env::temp_dir().join(format!(
2709 "ai_agents_eval_record_preflight_{}",
2710 uuid::Uuid::new_v4()
2711 ));
2712 std::fs::create_dir_all(&dir).unwrap();
2713 let blocked_parent = dir.join("not-a-directory");
2714 std::fs::write(&blocked_parent, "file").unwrap();
2715 let calls = Arc::new(AtomicUsize::new(0));
2716 let inner = Arc::new(CountingProvider {
2717 calls: Arc::clone(&calls),
2718 });
2719
2720 let provider =
2721 CassetteWriter::shared(blocked_parent.join("cassette.jsonl")).map(|writer| {
2722 RecordingLLMProvider::new(inner, "default".to_string(), "model".to_string(), writer)
2723 });
2724
2725 assert!(provider.is_err());
2726 assert_eq!(calls.load(Ordering::SeqCst), 0);
2727 let _ = std::fs::remove_dir_all(dir);
2728 }
2729
2730 #[tokio::test]
2731 async fn recording_write_errors_are_returned() {
2732 let dir = std::env::temp_dir().join(format!(
2733 "ai_agents_eval_record_error_{}",
2734 uuid::Uuid::new_v4()
2735 ));
2736 std::fs::create_dir_all(&dir).unwrap();
2737 let path = dir.join("read-only-cassette.jsonl");
2738 std::fs::write(&path, "").unwrap();
2739 let file = OpenOptions::new().read(true).open(&path).unwrap();
2740 let writer = Arc::new(CassetteWriter {
2741 path: path.clone(),
2742 file: Mutex::new(file),
2743 });
2744 let calls = Arc::new(AtomicUsize::new(0));
2745 let provider = RecordingLLMProvider::new(
2746 Arc::new(CountingProvider {
2747 calls: Arc::clone(&calls),
2748 }),
2749 "default".to_string(),
2750 "model".to_string(),
2751 writer,
2752 );
2753
2754 let error = provider
2755 .complete(&[ChatMessage::user("record me")], None)
2756 .await
2757 .unwrap_err();
2758
2759 assert!(error.to_string().contains("failed to write LLM cassette"));
2760 assert_eq!(calls.load(Ordering::SeqCst), 1);
2761 drop(provider);
2762 let _ = std::fs::remove_dir_all(dir);
2763 }
2764
2765 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2766 async fn aliases_share_synchronized_cassette_writes() {
2767 let dir = std::env::temp_dir().join(format!(
2768 "ai_agents_eval_shared_writer_{}",
2769 uuid::Uuid::new_v4()
2770 ));
2771 std::fs::create_dir_all(&dir).unwrap();
2772 let path = dir.join("cassette.jsonl");
2773 let first_writer = CassetteWriter::shared(path.clone()).unwrap();
2774 let second_writer = CassetteWriter::shared(path.clone()).unwrap();
2775 assert!(Arc::ptr_eq(&first_writer, &second_writer));
2776 let first = Arc::new(RecordingLLMProvider::new(
2777 Arc::new(SequenceLLMProvider::new(
2778 vec![SequenceOutcome::Response(LLMResponse::new(
2779 "first",
2780 FinishReason::Stop,
2781 ))],
2782 0,
2783 )),
2784 "first".to_string(),
2785 "model".to_string(),
2786 first_writer,
2787 ));
2788 let second = Arc::new(RecordingLLMProvider::new(
2789 Arc::new(SequenceLLMProvider::new(
2790 vec![SequenceOutcome::Response(LLMResponse::new(
2791 "second",
2792 FinishReason::Stop,
2793 ))],
2794 0,
2795 )),
2796 "second".to_string(),
2797 "model".to_string(),
2798 second_writer,
2799 ));
2800
2801 let first_calls = tokio::spawn(async move {
2802 for index in 0..50 {
2803 first
2804 .complete(&[ChatMessage::user(format!("first {index}"))], None)
2805 .await
2806 .unwrap();
2807 }
2808 });
2809 let second_calls = tokio::spawn(async move {
2810 for index in 0..50 {
2811 second
2812 .complete(&[ChatMessage::user(format!("second {index}"))], None)
2813 .await
2814 .unwrap();
2815 }
2816 });
2817 let (first_result, second_result) = tokio::join!(first_calls, second_calls);
2818 first_result.unwrap();
2819 second_result.unwrap();
2820
2821 let content = std::fs::read_to_string(&path).unwrap();
2822 let records: Vec<CassetteRecord> = content
2823 .lines()
2824 .map(|line| serde_json::from_str(line).unwrap())
2825 .collect();
2826 assert_eq!(records.len(), 100);
2827 assert_eq!(
2828 records
2829 .iter()
2830 .filter(|record| record.alias == "first")
2831 .count(),
2832 50
2833 );
2834 assert_eq!(
2835 records
2836 .iter()
2837 .filter(|record| record.alias == "second")
2838 .count(),
2839 50
2840 );
2841 let _ = std::fs::remove_dir_all(dir);
2842 }
2843
2844 #[tokio::test]
2845 async fn recording_streaming_fails_closed() {
2846 let dir = std::env::temp_dir().join(format!(
2847 "ai_agents_eval_record_stream_{}",
2848 uuid::Uuid::new_v4()
2849 ));
2850 std::fs::create_dir_all(&dir).unwrap();
2851 let writer = CassetteWriter::shared(dir.join("cassette.jsonl")).unwrap();
2852 let inner = Arc::new(SequenceLLMProvider::new(
2853 vec![SequenceOutcome::Response(LLMResponse::new(
2854 "streamed response",
2855 FinishReason::Stop,
2856 ))],
2857 0,
2858 ));
2859 let provider =
2860 RecordingLLMProvider::new(inner, "default".to_string(), "model".to_string(), writer);
2861
2862 let result = provider
2863 .complete_stream(&[ChatMessage::user("stream me")], None)
2864 .await;
2865
2866 let error = match result {
2867 Ok(_) => panic!("record streaming must fail closed"),
2868 Err(error) => error,
2869 };
2870 assert!(error.to_string().contains("does not support streaming"));
2871 assert!(!provider.supports(LLMFeature::Streaming));
2872 drop(provider);
2873 let _ = std::fs::remove_dir_all(dir);
2874 }
2875
2876 #[tokio::test]
2877 async fn mock_server_serves_configured_route() {
2878 let config = MockServerConfig {
2879 enabled: true,
2880 port: None,
2881 routes: vec![serde_json::json!({
2882 "method":"GET",
2883 "path":"/ok",
2884 "status":200,
2885 "body":{"ok":true}
2886 })],
2887 };
2888 let server = start_mock_server(Some(&config)).await.unwrap().unwrap();
2889 let context = server.context();
2890 let base_url = context
2891 .get("mock_server")
2892 .and_then(|value| value.get("base_url"))
2893 .and_then(Value::as_str)
2894 .unwrap()
2895 .trim_start_matches("http://")
2896 .to_string();
2897 let mut stream = tokio::net::TcpStream::connect(base_url).await.unwrap();
2898 stream
2899 .write_all(b"GET /ok HTTP/1.1\r\nHost: localhost\r\n\r\n")
2900 .await
2901 .unwrap();
2902 let mut response = String::new();
2903 stream.read_to_string(&mut response).await.unwrap();
2904 assert!(response.contains("200 OK"));
2905 assert!(response.contains("{\"ok\":true}"));
2906 }
2907}