Skip to main content

starweaver_model/test/
function_model.rs

1use std::sync::{Arc, Mutex};
2
3use async_trait::async_trait;
4
5use crate::{
6    ModelAdapter, ModelError,
7    adapter::{ModelRequestContext, ModelRequestParameters, ModelResponseEventStream},
8    message::{ModelMessage, ModelResponse},
9    profile::{ModelProfile, ProtocolFamily},
10    settings::ModelSettings,
11    stream::ModelResponseStreamEvent,
12};
13
14/// Information passed to a function-backed deterministic model.
15#[derive(Clone, Debug)]
16pub struct FunctionModelInfo {
17    /// Model request parameters visible to this model call.
18    pub params: ModelRequestParameters,
19    /// Runtime request context.
20    pub context: ModelRequestContext,
21}
22
23/// Function used by [`FunctionModel`] to produce responses.
24pub type FunctionModelFn = dyn Send
25    + Sync
26    + Fn(
27        Vec<ModelMessage>,
28        Option<ModelSettings>,
29        FunctionModelInfo,
30    ) -> Result<ModelResponse, ModelError>;
31
32/// Function used by [`FunctionModel`] to produce stream events.
33pub type FunctionModelStreamFn = dyn Send
34    + Sync
35    + Fn(
36        Vec<ModelMessage>,
37        Option<ModelSettings>,
38        FunctionModelInfo,
39    ) -> Result<Vec<ModelResponseStreamEvent>, ModelError>;
40
41/// Deterministic model backed by a caller-provided function.
42#[derive(Clone)]
43pub struct FunctionModel {
44    model_name: String,
45    profile: ModelProfile,
46    default_settings: Option<ModelSettings>,
47    function: Arc<FunctionModelFn>,
48    stream_function: Arc<FunctionModelStreamFn>,
49    captured_messages: Arc<Mutex<Vec<Vec<ModelMessage>>>>,
50    captured_params: Arc<Mutex<Vec<ModelRequestParameters>>>,
51}
52
53impl FunctionModel {
54    /// Create a function-backed model.
55    #[must_use]
56    pub fn new<F>(function: F) -> Self
57    where
58        F: Send
59            + Sync
60            + 'static
61            + Fn(
62                Vec<ModelMessage>,
63                Option<ModelSettings>,
64                FunctionModelInfo,
65            ) -> Result<ModelResponse, ModelError>,
66    {
67        let function: Arc<FunctionModelFn> = Arc::new(function);
68        let stream_function = function.clone();
69        Self {
70            model_name: "function".to_string(),
71            profile: ModelProfile::for_protocol(ProtocolFamily::OpenAiChatCompletions),
72            default_settings: None,
73            function,
74            stream_function: Arc::new(move |messages, settings, info| {
75                stream_function(messages, settings, info)
76                    .map(|response| vec![ModelResponseStreamEvent::FinalResult(Box::new(response))])
77            }),
78            captured_messages: Arc::new(Mutex::new(Vec::new())),
79            captured_params: Arc::new(Mutex::new(Vec::new())),
80        }
81    }
82
83    /// Create a function-backed streaming model.
84    #[must_use]
85    pub fn streaming<F>(function: F) -> Self
86    where
87        F: Send
88            + Sync
89            + 'static
90            + Fn(
91                Vec<ModelMessage>,
92                Option<ModelSettings>,
93                FunctionModelInfo,
94            ) -> Result<Vec<ModelResponseStreamEvent>, ModelError>,
95    {
96        Self {
97            model_name: "function".to_string(),
98            profile: ModelProfile::for_protocol(ProtocolFamily::OpenAiChatCompletions),
99            default_settings: None,
100            function: Arc::new(|_messages, _settings, _info| {
101                Err(ModelError::Transport(
102                    "function model response path is unavailable for streaming fixture".to_string(),
103                ))
104            }),
105            stream_function: Arc::new(function),
106            captured_messages: Arc::new(Mutex::new(Vec::new())),
107            captured_params: Arc::new(Mutex::new(Vec::new())),
108        }
109    }
110
111    /// Set model name.
112    #[must_use]
113    pub fn with_model_name(mut self, model_name: impl Into<String>) -> Self {
114        self.model_name = model_name.into();
115        self
116    }
117
118    /// Set model profile.
119    #[must_use]
120    pub fn with_profile(mut self, profile: ModelProfile) -> Self {
121        self.profile = profile;
122        self
123    }
124
125    /// Set model default settings.
126    #[must_use]
127    pub fn with_default_settings(mut self, settings: ModelSettings) -> Self {
128        self.default_settings = Some(settings);
129        self
130    }
131
132    /// Return captured request histories.
133    #[must_use]
134    pub fn captured_messages(&self) -> Vec<Vec<ModelMessage>> {
135        self.captured_messages
136            .lock()
137            .map_or_else(|_| Vec::new(), |messages| messages.clone())
138    }
139
140    /// Return captured request parameters.
141    #[must_use]
142    pub fn captured_params(&self) -> Vec<ModelRequestParameters> {
143        self.captured_params
144            .lock()
145            .map_or_else(|_| Vec::new(), |params| params.clone())
146    }
147}
148
149#[async_trait]
150impl ModelAdapter for FunctionModel {
151    fn model_name(&self) -> &str {
152        &self.model_name
153    }
154
155    fn provider_name(&self) -> Option<&str> {
156        Some("test")
157    }
158
159    fn profile(&self) -> &ModelProfile {
160        &self.profile
161    }
162
163    fn default_settings(&self) -> Option<&ModelSettings> {
164        self.default_settings.as_ref()
165    }
166
167    async fn request(
168        &self,
169        messages: Vec<ModelMessage>,
170        settings: Option<ModelSettings>,
171        params: ModelRequestParameters,
172        context: ModelRequestContext,
173    ) -> Result<ModelResponse, ModelError> {
174        if let Ok(mut captured) = self.captured_messages.lock() {
175            captured.push(messages.clone());
176        }
177        if let Ok(mut captured) = self.captured_params.lock() {
178            captured.push(params.clone());
179        }
180        (self.function)(messages, settings, FunctionModelInfo { params, context })
181    }
182
183    async fn request_stream(
184        &self,
185        messages: Vec<ModelMessage>,
186        settings: Option<ModelSettings>,
187        params: ModelRequestParameters,
188        context: ModelRequestContext,
189    ) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
190        let mut stream = self
191            .request_stream_incremental(messages, settings, params, context)
192            .await?;
193        let mut events = Vec::new();
194        while let Some(event) = stream.recv().await {
195            events.push(event?);
196        }
197        Ok(events)
198    }
199
200    async fn request_stream_incremental(
201        &self,
202        messages: Vec<ModelMessage>,
203        settings: Option<ModelSettings>,
204        params: ModelRequestParameters,
205        context: ModelRequestContext,
206    ) -> Result<ModelResponseEventStream, ModelError> {
207        if let Ok(mut captured) = self.captured_messages.lock() {
208            captured.push(messages.clone());
209        }
210        if let Ok(mut captured) = self.captured_params.lock() {
211            captured.push(params.clone());
212        }
213        let events =
214            (self.stream_function)(messages, settings, FunctionModelInfo { params, context })?;
215        let (sender, receiver) = tokio::sync::mpsc::channel(events.len().max(1));
216        tokio::spawn(async move {
217            for event in events {
218                if sender.send(Ok(event)).await.is_err() {
219                    return;
220                }
221            }
222        });
223        Ok(ModelResponseEventStream::new(receiver))
224    }
225}