starweaver_model/test/
scripted.rs1use 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#[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 #[must_use]
30 pub fn new() -> Self {
31 Self::with_responses(vec![ModelResponse::text("ok")])
32 }
33
34 #[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 #[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 #[must_use]
64 pub fn with_text(text: impl Into<String>) -> Self {
65 Self::with_responses(vec![ModelResponse::text(text)])
66 }
67
68 #[must_use]
70 pub fn with_json(value: &Value) -> Self {
71 Self::with_text(value.to_string())
72 }
73
74 #[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 #[must_use]
83 pub fn with_profile(mut self, profile: ModelProfile) -> Self {
84 self.profile = profile;
85 self
86 }
87
88 #[must_use]
90 pub fn with_default_settings(mut self, settings: ModelSettings) -> Self {
91 self.default_settings = Some(settings);
92 self
93 }
94
95 #[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 #[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}