Skip to main content

af_agent/
testing.rs

1//! Test doubles for the Agent seams: a scripted [`ChatModel`] that replays a
2//! queue of responses and records every request it received.
3
4use std::collections::VecDeque;
5use std::sync::Mutex;
6
7use af_llm::{ChatMessage, CompletionRequest, CompletionResponse, LlmError};
8use async_trait::async_trait;
9use tokio::sync::mpsc::UnboundedSender;
10
11use crate::ChatModel;
12
13/// Replays queued responses in order and streams each response's text as one
14/// delta. An exhausted script fails the request with a decode-style
15/// [`LlmError`] so a test never hangs on a missing response.
16#[derive(Default)]
17pub struct ScriptedModel {
18    responses: Mutex<VecDeque<Result<CompletionResponse, LlmError>>>,
19    requests: Mutex<Vec<CompletionRequest>>,
20}
21
22impl ScriptedModel {
23    /// Model that replays `responses` in order.
24    pub fn new(responses: impl IntoIterator<Item = Result<CompletionResponse, LlmError>>) -> Self {
25        Self {
26            responses: Mutex::new(responses.into_iter().collect()),
27            requests: Mutex::new(Vec::new()),
28        }
29    }
30
31    /// Script one assistant reply per message, each finishing with `stop` (or
32    /// `tool_calls` when the message carries tool calls).
33    pub fn replies(messages: impl IntoIterator<Item = ChatMessage>) -> Self {
34        Self::new(
35            messages
36                .into_iter()
37                .map(|message| Ok(Self::response(message))),
38        )
39    }
40
41    /// Build a canonical single-choice response for `message`.
42    pub fn response(message: ChatMessage) -> CompletionResponse {
43        let finish_reason = if message
44            .tool_calls
45            .as_ref()
46            .is_some_and(|calls| !calls.is_empty())
47        {
48            af_llm::FinishReason::ToolCalls
49        } else {
50            af_llm::FinishReason::Stop
51        };
52        CompletionResponse {
53            id: "scripted".into(),
54            choices: vec![af_llm::Choice {
55                index: 0,
56                message,
57                finish_reason: Some(finish_reason),
58                output_blocks: Vec::new(),
59            }],
60            usage: Some(af_llm::Usage {
61                prompt_tokens: 1,
62                completion_tokens: 1,
63                total_tokens: 2,
64            }),
65        }
66    }
67
68    /// Responses still queued.
69    pub fn remaining(&self) -> usize {
70        self.responses
71            .lock()
72            .unwrap_or_else(std::sync::PoisonError::into_inner)
73            .len()
74    }
75
76    /// Every request the runtime sent, in order.
77    pub fn requests(&self) -> Vec<CompletionRequest> {
78        self.requests
79            .lock()
80            .unwrap_or_else(std::sync::PoisonError::into_inner)
81            .clone()
82    }
83}
84
85#[async_trait]
86impl ChatModel for ScriptedModel {
87    async fn complete_streaming(
88        &self,
89        request: &CompletionRequest,
90        delta_tx: UnboundedSender<(String, bool)>,
91    ) -> Result<CompletionResponse, LlmError> {
92        self.requests
93            .lock()
94            .unwrap_or_else(std::sync::PoisonError::into_inner)
95            .push(request.clone());
96        let next = self
97            .responses
98            .lock()
99            .unwrap_or_else(std::sync::PoisonError::into_inner)
100            .pop_front();
101        let response = match next {
102            Some(response) => response?,
103            None => return Err(LlmError::StreamProtocol("scripted model exhausted".into())),
104        };
105        if let Some(content) = response.first_content() {
106            let _ = delta_tx.send((content.into(), response.first_tool_calls().is_some()));
107        }
108        Ok(response)
109    }
110}
111
112#[cfg(test)]
113mod tests {
114    use super::*;
115
116    #[tokio::test]
117    async fn scripted_model_replays_records_and_exhausts_loudly() {
118        let model = ScriptedModel::replies([ChatMessage::assistant("hi")]);
119        assert_eq!(model.remaining(), 1);
120        let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel();
121        let request = CompletionRequest::new("m", vec![ChatMessage::user("hello")]);
122        let response = model.complete_streaming(&request, tx).await.unwrap();
123        assert_eq!(response.first_content(), Some("hi"));
124        assert_eq!(rx.recv().await.unwrap(), ("hi".to_string(), false));
125        assert_eq!(model.requests().len(), 1);
126        assert_eq!(model.remaining(), 0);
127        let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
128        assert!(matches!(
129            model.complete_streaming(&request, tx).await,
130            Err(LlmError::StreamProtocol(_))
131        ));
132        let failing = ScriptedModel::new([Err(LlmError::CircuitOpen)]);
133        let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
134        assert!(matches!(
135            failing.complete_streaming(&request, tx).await,
136            Err(LlmError::CircuitOpen)
137        ));
138    }
139}