Skip to main content

rig_core/driver/
dyn_model.rs

1//! A model erased to its operation: what a consumer stores when it holds
2//! any model of one operation without naming its wire and transport. A
3//! [`DynModel`] runs the same driver as the [`Model`] it was made from.
4
5use std::fmt;
6use std::sync::Arc;
7
8use super::{Model, Transport};
9use crate::error::ProviderError;
10use crate::observe::AdapterContext;
11use crate::streaming::Streamed;
12use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
13use crate::wire::{Capabilities, Descriptor, Mode, Operation, Wire};
14
15/// Object-safe mirror of the calls a [`Model`] answers, with the wire and
16/// transport fixed. Private: the only way to reach it is through
17/// [`DynModel`], which re-exposes the public surface.
18pub(crate) trait ErasedModel<Op: Operation>: WasmCompatSend + WasmCompatSync {
19    fn describe(&self) -> Descriptor<'_>;
20
21    fn open(
22        &self,
23        request: Op::Request,
24        mode: Mode,
25        observation: Option<AdapterContext>,
26    ) -> Result<Streamed<Op>, ProviderError>;
27}
28
29impl<W, T> ErasedModel<W::Op> for Model<W, T>
30where
31    W: Wire,
32    T: Transport<W>,
33{
34    fn describe(&self) -> Descriptor<'_> {
35        self.wire.describe()
36    }
37
38    fn open(
39        &self,
40        request: <W::Op as Operation>::Request,
41        mode: Mode,
42        observation: Option<AdapterContext>,
43    ) -> Result<Streamed<W::Op>, ProviderError> {
44        Model::open(self, request, mode, observation)
45    }
46}
47
48/// A model of one operation with its wire and transport erased. Clones share
49/// the model. Built with [`Model::erase`] or `From<Model<W, T>>`, so a
50/// consumer takes `impl Into<DynModel<Op>>` and accepts either.
51///
52/// Every call runs the driver the concrete model runs: spans, request ids,
53/// the operation's fold and error enrichment are the same.
54///
55/// A model used by one consumer is passed as is; a model shared by several
56/// is erased once and the handle is cloned. A model the bus serves is a
57/// `ModelHandle` in the agent runtime instead: its calls are dispatched,
58/// recorded and observed by the bus.
59///
60/// ```no_run
61/// use rig_core::embeddings::EmbeddingsBuilder;
62/// use rig_core::vector_store::in_memory_store::InMemoryVectorStore;
63/// use rig_core::{Model, providers::openai::{self, OpenAI}};
64///
65/// # async fn example(http: rig_core::http_client::DynHttpClient) -> Result<(), Box<dyn std::error::Error>> {
66/// let model = OpenAI::from_env()?.with_http(http).embedding(openai::TEXT_EMBEDDING_3_SMALL, None).erase();
67/// let embeddings = EmbeddingsBuilder::new(model.clone())
68///     .documents(["a document".to_owned()])?
69///     .build()
70///     .await?;
71/// let index = InMemoryVectorStore::from_documents(embeddings).index(model);
72/// # let _ = index;
73/// # Ok(())
74/// # }
75/// ```
76pub struct DynModel<Op: Operation> {
77    inner: Arc<dyn ErasedModel<Op>>,
78}
79
80impl<W, T> Model<W, T>
81where
82    W: Wire,
83    T: Transport<W>,
84{
85    /// Erase this model to its operation.
86    pub fn erase(self) -> DynModel<W::Op> {
87        DynModel {
88            inner: Arc::new(self),
89        }
90    }
91}
92
93impl<W, T> From<Model<W, T>> for DynModel<W::Op>
94where
95    W: Wire,
96    T: Transport<W>,
97{
98    fn from(model: Model<W, T>) -> Self {
99        model.erase()
100    }
101}
102
103impl<Op: Operation> Clone for DynModel<Op> {
104    fn clone(&self) -> Self {
105        Self {
106            inner: Arc::clone(&self.inner),
107        }
108    }
109}
110
111impl<Op: Operation> fmt::Debug for DynModel<Op> {
112    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
113        f.debug_struct("DynModel")
114            .field("name", &self.name())
115            .field("id", &self.id())
116            .finish()
117    }
118}
119
120impl<Op: Operation> DynModel<Op> {
121    /// The wire's provider descriptor name (`"anthropic"`).
122    pub fn name(&self) -> &str {
123        self.inner.describe().name
124    }
125
126    /// The model id the wire addresses, when the operation addresses one.
127    pub fn id(&self) -> Option<&str> {
128        self.inner.describe().model
129    }
130
131    /// What a runtime accounts for about this model.
132    pub fn capabilities(&self) -> Capabilities {
133        self.inner.describe().capabilities
134    }
135
136    /// Send `request` and fold the whole reply into the operation's
137    /// response; [`Model::call`] with the model erased. The call owns what
138    /// it sends, so it can be spawned.
139    pub fn call(
140        &self,
141        request: impl Into<Op::Request>,
142    ) -> impl Future<Output = Result<Op::Response, ProviderError>> + WasmCompatSend + 'static {
143        self.finished(request.into(), None)
144    }
145
146    /// [`Self::call`], with the attempt observed under `observation`.
147    pub fn call_observed(
148        &self,
149        request: impl Into<Op::Request>,
150        observation: AdapterContext,
151    ) -> impl Future<Output = Result<Op::Response, ProviderError>> + WasmCompatSend + 'static {
152        self.finished(request.into(), Some(observation))
153    }
154
155    /// The call opens when first polled, as [`Model::call`] does.
156    fn finished(
157        &self,
158        request: Op::Request,
159        observation: Option<AdapterContext>,
160    ) -> impl Future<Output = Result<Op::Response, ProviderError>> + WasmCompatSend + 'static {
161        let inner = self.inner.clone();
162        async move {
163            inner
164                .open(request, Mode::Unary, observation)?
165                .finish()
166                .await
167        }
168    }
169
170    /// Open a streamed reply; [`Model::stream`] with the model erased.
171    pub fn stream(&self, request: impl Into<Op::Request>) -> Result<Streamed<Op>, ProviderError> {
172        self.inner.open(request.into(), Mode::Streaming, None)
173    }
174
175    /// [`Self::stream`], with the attempt observed under `observation`.
176    pub fn stream_observed(
177        &self,
178        request: impl Into<Op::Request>,
179        observation: AdapterContext,
180    ) -> Result<Streamed<Op>, ProviderError> {
181        self.inner
182            .open(request.into(), Mode::Streaming, Some(observation))
183    }
184}
185
186#[cfg(test)]
187mod tests;