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#[derive(Clone, Debug)]
16pub struct FunctionModelInfo {
17 pub params: ModelRequestParameters,
19 pub context: ModelRequestContext,
21}
22
23pub type FunctionModelFn = dyn Send
25 + Sync
26 + Fn(
27 Vec<ModelMessage>,
28 Option<ModelSettings>,
29 FunctionModelInfo,
30 ) -> Result<ModelResponse, ModelError>;
31
32pub type FunctionModelStreamFn = dyn Send
34 + Sync
35 + Fn(
36 Vec<ModelMessage>,
37 Option<ModelSettings>,
38 FunctionModelInfo,
39 ) -> Result<Vec<ModelResponseStreamEvent>, ModelError>;
40
41#[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 #[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 #[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 #[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 #[must_use]
120 pub fn with_profile(mut self, profile: ModelProfile) -> Self {
121 self.profile = profile;
122 self
123 }
124
125 #[must_use]
127 pub fn with_default_settings(mut self, settings: ModelSettings) -> Self {
128 self.default_settings = Some(settings);
129 self
130 }
131
132 #[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 #[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}