1use std::{fmt, sync::Arc};
14
15use rig_core::{
16 completion::{
17 CompletionError, CompletionModel, CompletionRequest, CompletionResponse,
18 ProviderCapabilities,
19 },
20 streaming::StreamingCompletionResponse,
21 wasm_compat::{WasmBoxedFuture, WasmCompatSend, WasmCompatSync},
22};
23
24trait ErasedModel: WasmCompatSend + WasmCompatSync {
33 fn completion(
34 &self,
35 request: CompletionRequest,
36 ) -> WasmBoxedFuture<'_, Result<CompletionResponse, CompletionError>>;
37
38 fn stream(
39 &self,
40 request: CompletionRequest,
41 ) -> WasmBoxedFuture<'_, Result<StreamingCompletionResponse, CompletionError>>;
42}
43
44impl<M> ErasedModel for M
48where
49 M: CompletionModel + 'static,
50{
51 fn completion(
52 &self,
53 request: CompletionRequest,
54 ) -> WasmBoxedFuture<'_, Result<CompletionResponse, CompletionError>> {
55 Box::pin(CompletionModel::completion(self, request))
56 }
57
58 fn stream(
59 &self,
60 request: CompletionRequest,
61 ) -> WasmBoxedFuture<'_, Result<StreamingCompletionResponse, CompletionError>> {
62 Box::pin(CompletionModel::stream(self, request))
63 }
64}
65
66struct ModelDriver<M: ?Sized> {
70 capabilities: ProviderCapabilities,
72 label: Option<String>,
73 model: M,
74}
75
76#[derive(Clone)]
108pub struct ModelHandle {
109 inner: Arc<ModelDriver<dyn ErasedModel>>,
110}
111
112impl ModelHandle {
113 pub fn new<M>(model: M) -> Self
115 where
116 M: CompletionModel + 'static,
117 {
118 Self::from_parts(None, model)
119 }
120
121 pub fn named<M>(label: impl Into<String>, model: M) -> Self
126 where
127 M: CompletionModel + 'static,
128 {
129 Self::from_parts(Some(label.into()), model)
130 }
131
132 fn from_parts<M>(label: Option<String>, model: M) -> Self
133 where
134 M: CompletionModel + 'static,
135 {
136 let capabilities = model.capabilities();
140 Self {
141 inner: Arc::new(ModelDriver {
142 capabilities,
143 label,
144 model,
145 }),
146 }
147 }
148
149 pub fn label(&self) -> Option<&str> {
151 self.inner.label.as_deref()
152 }
153}
154
155impl CompletionModel for ModelHandle {
165 fn completion(
166 &self,
167 request: CompletionRequest,
168 ) -> impl Future<Output = Result<CompletionResponse, CompletionError>>
169 + rig_core::wasm_compat::WasmCompatSend {
170 self.inner.model.completion(request)
171 }
172
173 fn stream(
174 &self,
175 request: CompletionRequest,
176 ) -> impl Future<Output = Result<StreamingCompletionResponse, CompletionError>>
177 + rig_core::wasm_compat::WasmCompatSend {
178 self.inner.model.stream(request)
179 }
180
181 fn capabilities(&self) -> ProviderCapabilities {
182 self.inner.capabilities
183 }
184}
185
186impl fmt::Debug for ModelHandle {
187 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
188 formatter
189 .debug_struct("ModelHandle")
190 .field("label", &self.label())
191 .field("capabilities", &self.inner.capabilities)
192 .finish_non_exhaustive()
193 }
194}
195
196#[cfg(test)]
197mod tests {
198 use std::sync::atomic::{AtomicUsize, Ordering};
199
200 use super::*;
201 use crate::test_utils::{MockCompletionModel, MockTurn};
202
203 struct CloneCountingModel {
205 inner: MockCompletionModel,
206 clones: Arc<AtomicUsize>,
207 }
208
209 impl Clone for CloneCountingModel {
210 fn clone(&self) -> Self {
211 self.clones.fetch_add(1, Ordering::SeqCst);
212 Self {
213 inner: self.inner.clone(),
214 clones: Arc::clone(&self.clones),
215 }
216 }
217 }
218
219 impl CompletionModel for CloneCountingModel {
220 fn completion(
221 &self,
222 request: CompletionRequest,
223 ) -> impl Future<Output = Result<CompletionResponse, CompletionError>>
224 + rig_core::wasm_compat::WasmCompatSend {
225 CompletionModel::completion(&self.inner, request)
226 }
227
228 fn stream(
229 &self,
230 request: CompletionRequest,
231 ) -> impl Future<Output = Result<StreamingCompletionResponse, CompletionError>>
232 + rig_core::wasm_compat::WasmCompatSend {
233 CompletionModel::stream(&self.inner, request)
234 }
235 }
236
237 #[tokio::test]
241 async fn erasure_never_clones_the_model() {
242 let clones = Arc::new(AtomicUsize::new(0));
243 let model = CloneCountingModel {
244 inner: MockCompletionModel::from_turns([
245 MockTurn::text("one"),
246 MockTurn::text("two"),
247 MockTurn::text("three"),
248 ]),
249 clones: Arc::clone(&clones),
250 };
251
252 let handle = ModelHandle::new(model);
253 let request = handle.completion_request("go").build();
254 CompletionModel::completion(&handle, request.clone())
255 .await
256 .expect("first scripted turn");
257 CompletionModel::completion(&handle, request.clone())
258 .await
259 .expect("second scripted turn");
260 CompletionModel::completion(&handle, request)
261 .await
262 .expect("third scripted turn");
263
264 let stream_clones = Arc::new(AtomicUsize::new(0));
265 let stream_model = CloneCountingModel {
266 inner: MockCompletionModel::from_stream_turns([
267 vec![
268 crate::test_utils::MockStreamEvent::text("a"),
269 crate::test_utils::MockStreamEvent::final_response_with_default_usage(),
270 ],
271 vec![
272 crate::test_utils::MockStreamEvent::text("b"),
273 crate::test_utils::MockStreamEvent::final_response_with_default_usage(),
274 ],
275 ]),
276 clones: Arc::clone(&stream_clones),
277 };
278 let stream_handle = ModelHandle::new(stream_model);
279 let stream_request = stream_handle.completion_request("go").build();
280 CompletionModel::stream(&stream_handle, stream_request.clone())
281 .await
282 .expect("first scripted stream turn");
283 CompletionModel::stream(&stream_handle, stream_request)
284 .await
285 .expect("second scripted stream turn");
286
287 assert_eq!(
288 clones.load(Ordering::SeqCst),
289 0,
290 "erasure and unary attempts must never clone the model"
291 );
292 assert_eq!(
293 stream_clones.load(Ordering::SeqCst),
294 0,
295 "erasure and streaming attempts must never clone the model"
296 );
297 }
298
299 struct NonCloneModel;
303
304 impl CompletionModel for NonCloneModel {
305 fn completion(
306 &self,
307 _request: CompletionRequest,
308 ) -> impl Future<Output = Result<CompletionResponse, CompletionError>>
309 + rig_core::wasm_compat::WasmCompatSend {
310 std::future::ready(Err(CompletionError::ProviderError(
311 "compile-time probe".to_string(),
312 )))
313 }
314
315 fn stream(
316 &self,
317 _request: CompletionRequest,
318 ) -> impl Future<Output = Result<StreamingCompletionResponse, CompletionError>>
319 + rig_core::wasm_compat::WasmCompatSend {
320 std::future::ready(Err(CompletionError::ProviderError(
321 "compile-time probe".to_string(),
322 )))
323 }
324 }
325
326 #[test]
327 fn traits() {
328 fn assert_completion_model<M: CompletionModel>() {}
329
330 assert_completion_model::<NonCloneModel>();
331 assert_completion_model::<std::sync::Arc<NonCloneModel>>();
336
337 let _ = || {
340 let handle = ModelHandle::new(NonCloneModel);
341 let named = ModelHandle::named("probe", NonCloneModel);
342 let via_arc = std::sync::Arc::new(NonCloneModel).completion_request("go");
343 let builder = crate::AgentBuilder::new(NonCloneModel);
344 (handle, named, via_arc, builder)
345 };
346 }
347}