use std::{fmt, sync::Arc};
use rig_core::{
completion::{
CompletionError, CompletionModel, CompletionRequest, CompletionResponse,
ProviderCapabilities,
},
streaming::StreamingCompletionResponse,
wasm_compat::{WasmBoxedFuture, WasmCompatSend, WasmCompatSync},
};
trait ErasedModel: WasmCompatSend + WasmCompatSync {
fn completion(
&self,
request: CompletionRequest,
) -> WasmBoxedFuture<'_, Result<CompletionResponse, CompletionError>>;
fn stream(
&self,
request: CompletionRequest,
) -> WasmBoxedFuture<'_, Result<StreamingCompletionResponse, CompletionError>>;
}
impl<M> ErasedModel for M
where
M: CompletionModel + 'static,
{
fn completion(
&self,
request: CompletionRequest,
) -> WasmBoxedFuture<'_, Result<CompletionResponse, CompletionError>> {
Box::pin(CompletionModel::completion(self, request))
}
fn stream(
&self,
request: CompletionRequest,
) -> WasmBoxedFuture<'_, Result<StreamingCompletionResponse, CompletionError>> {
Box::pin(CompletionModel::stream(self, request))
}
}
struct ModelDriver<M: ?Sized> {
capabilities: ProviderCapabilities,
label: Option<String>,
model: M,
}
#[derive(Clone)]
pub struct ModelHandle {
inner: Arc<ModelDriver<dyn ErasedModel>>,
}
impl ModelHandle {
pub fn new<M>(model: M) -> Self
where
M: CompletionModel + 'static,
{
Self::from_parts(None, model)
}
pub fn named<M>(label: impl Into<String>, model: M) -> Self
where
M: CompletionModel + 'static,
{
Self::from_parts(Some(label.into()), model)
}
fn from_parts<M>(label: Option<String>, model: M) -> Self
where
M: CompletionModel + 'static,
{
let capabilities = model.capabilities();
Self {
inner: Arc::new(ModelDriver {
capabilities,
label,
model,
}),
}
}
pub fn label(&self) -> Option<&str> {
self.inner.label.as_deref()
}
}
impl CompletionModel for ModelHandle {
fn completion(
&self,
request: CompletionRequest,
) -> impl Future<Output = Result<CompletionResponse, CompletionError>>
+ rig_core::wasm_compat::WasmCompatSend {
self.inner.model.completion(request)
}
fn stream(
&self,
request: CompletionRequest,
) -> impl Future<Output = Result<StreamingCompletionResponse, CompletionError>>
+ rig_core::wasm_compat::WasmCompatSend {
self.inner.model.stream(request)
}
fn capabilities(&self) -> ProviderCapabilities {
self.inner.capabilities
}
}
impl fmt::Debug for ModelHandle {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ModelHandle")
.field("label", &self.label())
.field("capabilities", &self.inner.capabilities)
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use super::*;
use crate::test_utils::{MockCompletionModel, MockTurn};
struct CloneCountingModel {
inner: MockCompletionModel,
clones: Arc<AtomicUsize>,
}
impl Clone for CloneCountingModel {
fn clone(&self) -> Self {
self.clones.fetch_add(1, Ordering::SeqCst);
Self {
inner: self.inner.clone(),
clones: Arc::clone(&self.clones),
}
}
}
impl CompletionModel for CloneCountingModel {
fn completion(
&self,
request: CompletionRequest,
) -> impl Future<Output = Result<CompletionResponse, CompletionError>>
+ rig_core::wasm_compat::WasmCompatSend {
CompletionModel::completion(&self.inner, request)
}
fn stream(
&self,
request: CompletionRequest,
) -> impl Future<Output = Result<StreamingCompletionResponse, CompletionError>>
+ rig_core::wasm_compat::WasmCompatSend {
CompletionModel::stream(&self.inner, request)
}
}
#[tokio::test]
async fn erasure_never_clones_the_model() {
let clones = Arc::new(AtomicUsize::new(0));
let model = CloneCountingModel {
inner: MockCompletionModel::from_turns([
MockTurn::text("one"),
MockTurn::text("two"),
MockTurn::text("three"),
]),
clones: Arc::clone(&clones),
};
let handle = ModelHandle::new(model);
let request = handle.completion_request("go").build();
CompletionModel::completion(&handle, request.clone())
.await
.expect("first scripted turn");
CompletionModel::completion(&handle, request.clone())
.await
.expect("second scripted turn");
CompletionModel::completion(&handle, request)
.await
.expect("third scripted turn");
let stream_clones = Arc::new(AtomicUsize::new(0));
let stream_model = CloneCountingModel {
inner: MockCompletionModel::from_stream_turns([
vec![
crate::test_utils::MockStreamEvent::text("a"),
crate::test_utils::MockStreamEvent::final_response_with_default_usage(),
],
vec![
crate::test_utils::MockStreamEvent::text("b"),
crate::test_utils::MockStreamEvent::final_response_with_default_usage(),
],
]),
clones: Arc::clone(&stream_clones),
};
let stream_handle = ModelHandle::new(stream_model);
let stream_request = stream_handle.completion_request("go").build();
CompletionModel::stream(&stream_handle, stream_request.clone())
.await
.expect("first scripted stream turn");
CompletionModel::stream(&stream_handle, stream_request)
.await
.expect("second scripted stream turn");
assert_eq!(
clones.load(Ordering::SeqCst),
0,
"erasure and unary attempts must never clone the model"
);
assert_eq!(
stream_clones.load(Ordering::SeqCst),
0,
"erasure and streaming attempts must never clone the model"
);
}
struct NonCloneModel;
impl CompletionModel for NonCloneModel {
fn completion(
&self,
_request: CompletionRequest,
) -> impl Future<Output = Result<CompletionResponse, CompletionError>>
+ rig_core::wasm_compat::WasmCompatSend {
std::future::ready(Err(CompletionError::ProviderError(
"compile-time probe".to_string(),
)))
}
fn stream(
&self,
_request: CompletionRequest,
) -> impl Future<Output = Result<StreamingCompletionResponse, CompletionError>>
+ rig_core::wasm_compat::WasmCompatSend {
std::future::ready(Err(CompletionError::ProviderError(
"compile-time probe".to_string(),
)))
}
}
#[test]
fn traits() {
fn assert_completion_model<M: CompletionModel>() {}
assert_completion_model::<NonCloneModel>();
assert_completion_model::<std::sync::Arc<NonCloneModel>>();
let _ = || {
let handle = ModelHandle::new(NonCloneModel);
let named = ModelHandle::named("probe", NonCloneModel);
let via_arc = std::sync::Arc::new(NonCloneModel).completion_request("go");
let builder = crate::AgentBuilder::new(NonCloneModel);
(handle, named, via_arc, builder)
};
}
}