Skip to main content

starweaver_model/test/
scripted.rs

1use std::sync::{Arc, Mutex};
2
3use async_trait::async_trait;
4use serde_json::Value;
5
6use crate::{
7    ModelAdapter, ModelError,
8    adapter::{ModelRequestContext, ModelRequestParameters, ModelResponseEventStream},
9    message::{ModelMessage, ModelResponse},
10    profile::{ModelProfile, ProtocolFamily},
11    settings::ModelSettings,
12    stream::ModelResponseStreamEvent,
13};
14
15/// Deterministic model that returns scripted responses.
16#[derive(Clone)]
17pub struct TestModel {
18    model_name: String,
19    profile: ModelProfile,
20    default_settings: Option<ModelSettings>,
21    responses: Arc<Mutex<Vec<ModelResponse>>>,
22    stream_events: Arc<Mutex<Vec<Vec<ModelResponseStreamEvent>>>>,
23    captured_messages: Arc<Mutex<Vec<Vec<ModelMessage>>>>,
24    captured_params: Arc<Mutex<Vec<ModelRequestParameters>>>,
25}
26
27impl TestModel {
28    /// Create a test model with a default text response.
29    #[must_use]
30    pub fn new() -> Self {
31        Self::with_responses(vec![ModelResponse::text("ok")])
32    }
33
34    /// Create a test model with scripted responses in call order.
35    #[must_use]
36    pub fn with_responses(responses: Vec<ModelResponse>) -> Self {
37        Self {
38            model_name: "test".to_string(),
39            profile: ModelProfile::for_protocol(ProtocolFamily::OpenAiChatCompletions),
40            default_settings: None,
41            responses: Arc::new(Mutex::new(responses.into_iter().rev().collect())),
42            stream_events: Arc::new(Mutex::new(Vec::new())),
43            captured_messages: Arc::new(Mutex::new(Vec::new())),
44            captured_params: Arc::new(Mutex::new(Vec::new())),
45        }
46    }
47
48    /// Create a test model with scripted stream event batches in call order.
49    #[must_use]
50    pub fn with_stream_events(events: Vec<Vec<ModelResponseStreamEvent>>) -> Self {
51        Self {
52            model_name: "test".to_string(),
53            profile: ModelProfile::for_protocol(ProtocolFamily::OpenAiChatCompletions),
54            default_settings: None,
55            responses: Arc::new(Mutex::new(Vec::new())),
56            stream_events: Arc::new(Mutex::new(events.into_iter().rev().collect())),
57            captured_messages: Arc::new(Mutex::new(Vec::new())),
58            captured_params: Arc::new(Mutex::new(Vec::new())),
59        }
60    }
61
62    /// Create a test model returning plain text.
63    #[must_use]
64    pub fn with_text(text: impl Into<String>) -> Self {
65        Self::with_responses(vec![ModelResponse::text(text)])
66    }
67
68    /// Create a test model returning JSON text.
69    #[must_use]
70    pub fn with_json(value: &Value) -> Self {
71        Self::with_text(value.to_string())
72    }
73
74    /// Set model name.
75    #[must_use]
76    pub fn with_model_name(mut self, model_name: impl Into<String>) -> Self {
77        self.model_name = model_name.into();
78        self
79    }
80
81    /// Set model profile.
82    #[must_use]
83    pub fn with_profile(mut self, profile: ModelProfile) -> Self {
84        self.profile = profile;
85        self
86    }
87
88    /// Set model default settings.
89    #[must_use]
90    pub fn with_default_settings(mut self, settings: ModelSettings) -> Self {
91        self.default_settings = Some(settings);
92        self
93    }
94
95    /// Return captured request histories.
96    #[must_use]
97    pub fn captured_messages(&self) -> Vec<Vec<ModelMessage>> {
98        self.captured_messages
99            .lock()
100            .map_or_else(|_| Vec::new(), |messages| messages.clone())
101    }
102
103    /// Return captured request parameters.
104    #[must_use]
105    pub fn captured_params(&self) -> Vec<ModelRequestParameters> {
106        self.captured_params
107            .lock()
108            .map_or_else(|_| Vec::new(), |params| params.clone())
109    }
110}
111
112impl Default for TestModel {
113    fn default() -> Self {
114        Self::new()
115    }
116}
117
118#[async_trait]
119impl ModelAdapter for TestModel {
120    fn model_name(&self) -> &str {
121        &self.model_name
122    }
123
124    fn provider_name(&self) -> Option<&str> {
125        Some("test")
126    }
127
128    fn profile(&self) -> &ModelProfile {
129        &self.profile
130    }
131
132    fn default_settings(&self) -> Option<&ModelSettings> {
133        self.default_settings.as_ref()
134    }
135
136    async fn request(
137        &self,
138        messages: Vec<ModelMessage>,
139        _settings: Option<ModelSettings>,
140        params: ModelRequestParameters,
141        _context: ModelRequestContext,
142    ) -> Result<ModelResponse, ModelError> {
143        if let Ok(mut captured) = self.captured_messages.lock() {
144            captured.push(messages);
145        }
146        if let Ok(mut captured) = self.captured_params.lock() {
147            captured.push(params);
148        }
149        self.responses
150            .lock()
151            .map_err(|err| ModelError::Transport(err.to_string()))?
152            .pop()
153            .ok_or_else(|| ModelError::Transport("test model script exhausted".to_string()))
154    }
155
156    async fn request_stream(
157        &self,
158        messages: Vec<ModelMessage>,
159        settings: Option<ModelSettings>,
160        params: ModelRequestParameters,
161        context: ModelRequestContext,
162    ) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
163        let mut stream = self
164            .request_stream_incremental(messages, settings, params, context)
165            .await?;
166        let mut events = Vec::new();
167        while let Some(event) = stream.recv().await {
168            events.push(event?);
169        }
170        Ok(events)
171    }
172
173    async fn request_stream_incremental(
174        &self,
175        messages: Vec<ModelMessage>,
176        _settings: Option<ModelSettings>,
177        params: ModelRequestParameters,
178        _context: ModelRequestContext,
179    ) -> Result<ModelResponseEventStream, ModelError> {
180        if let Ok(mut captured) = self.captured_messages.lock() {
181            captured.push(messages);
182        }
183        if let Ok(mut captured) = self.captured_params.lock() {
184            captured.push(params);
185        }
186        let stream_events = self
187            .stream_events
188            .lock()
189            .map_err(|err| ModelError::Transport(err.to_string()))?
190            .pop();
191        let events = if let Some(events) = stream_events {
192            events
193        } else {
194            let response = self
195                .responses
196                .lock()
197                .map_err(|err| ModelError::Transport(err.to_string()))?
198                .pop()
199                .ok_or_else(|| ModelError::Transport("test model script exhausted".to_string()))?;
200            vec![ModelResponseStreamEvent::FinalResult(Box::new(response))]
201        };
202        let (sender, receiver) = tokio::sync::mpsc::channel(events.len().max(1));
203        tokio::spawn(async move {
204            for event in events {
205                if sender.send(Ok(event)).await.is_err() {
206                    return;
207                }
208            }
209        });
210        Ok(ModelResponseEventStream::new(receiver))
211    }
212}