Skip to main content

rig_core/serve/
adapters.rs

1//! Effect handlers adapting model, tool, memory, and retrieval traits.
2//! Each adapter translates its effect family into trait calls and returns
3//! outcomes or streaming events.
4//!
5//! ```
6//! use rig_core::{memory::InMemoryConversationMemory, serve::{Serve, adapters::MemoryAdapter}};
7//!
8//! let handler = MemoryAdapter::new(InMemoryConversationMemory::new());
9//! assert_eq!(handler.descriptor().key.as_str(), "memory");
10//! ```
11
12use crate::{
13    completion::ModelRef,
14    driver::DynModel,
15    effect::{
16        EffectFamily, EffectKind, EmbedInputs, EmbedModality, EmbedOutputs, FamilyDescriptor,
17        HandlerDescriptor, HandlerKey, MemoryOp, MemoryOutcome, Outcome, RetrieveQuery,
18        RetrievedDocuments, ToolEmbeddingDescriptor,
19    },
20    error::{ErrorKind, ErrorReport},
21    memory::ConversationMemory,
22    operation::{Completion, Embedding, Rerank},
23    tool::{ErasedTool, Tool, ToolEmbedding},
24    vector_store::{VectorStoreError, VectorStoreIndex, request::DynamicSearchFilter},
25    wasm_compat::{WasmBoxedFuture, WasmCompatSend, WasmCompatSync},
26    wire::Operation,
27};
28
29use super::{Dispatch, Reply, Serve};
30use crate::effect::family;
31
32fn wrong_family(handler: EffectFamily, kind: &EffectKind) -> ErrorReport {
33    ErrorReport::new(
34        ErrorKind::HandlerUnavailable,
35        format!(
36            "a {handler} handler cannot serve a `{}` effect",
37            kind.name()
38        ),
39    )
40}
41
42/// A model as a handler under `label`. Completion, text-embedding and
43/// rerank models serve their effect families. The model is erased to its
44/// operation, so one adapter type serves every wire and transport.
45pub struct ModelAdapter<Op: Operation> {
46    label: ModelRef,
47    model: DynModel<Op>,
48}
49
50impl<Op: Operation> ModelAdapter<Op> {
51    /// Wrap `model` under `label`.
52    pub fn new(label: impl Into<ModelRef>, model: impl Into<DynModel<Op>>) -> Self {
53        Self {
54            label: label.into(),
55            model: model.into(),
56        }
57    }
58}
59
60/// Unary and streaming completions both route here; the descriptor carries
61/// the model's label and capability snapshot.
62impl Serve for ModelAdapter<Completion> {
63    type Family = family::Completion;
64
65    fn descriptor(&self) -> HandlerDescriptor {
66        HandlerDescriptor {
67            key: crate::effect::model_key(self.label.as_str()),
68            family: FamilyDescriptor::Completion {
69                model: self.label.clone(),
70                capabilities: self.model.capabilities().completion,
71            },
72            layers: Vec::new(),
73        }
74    }
75
76    async fn serve(&self, kind: EffectKind, dispatch: Dispatch) -> Reply {
77        let model = &self.model;
78        let context = dispatch.adapter_context();
79        match kind {
80            EffectKind::Completion {
81                request,
82                stream: false,
83            } => {
84                let result = match context {
85                    Some(context) => model.call_observed(request, context).await,
86                    None => model.call(request).await,
87                };
88                Reply::Outcome(result.map(Outcome::Completion).map_err(ErrorReport::from))
89            }
90            EffectKind::Completion {
91                request,
92                stream: true,
93            } => {
94                let opened = match context {
95                    Some(context) => model.stream_observed(request, context),
96                    None => model.stream(request),
97                };
98                match opened {
99                    Ok(stream) => Reply::Stream(stream.into_relay()),
100                    Err(error) => Reply::Outcome(Err(ErrorReport::from(error))),
101                }
102            }
103            other @ (EffectKind::ToolCall { .. }
104            | EffectKind::Embed { .. }
105            | EffectKind::Memory { .. }
106            | EffectKind::Retrieve { .. }
107            | EffectKind::Rerank { .. }
108            | EffectKind::Custom { .. }) => {
109                Reply::Outcome(Err(wrong_family(EffectFamily::Completion, &other)))
110            }
111        }
112    }
113}
114
115/// A text embedding model embeds texts and refuses images.
116impl Serve for ModelAdapter<Embedding> {
117    type Family = family::Embed;
118
119    fn descriptor(&self) -> HandlerDescriptor {
120        let capabilities = self.model.capabilities();
121        HandlerDescriptor {
122            key: crate::effect::embed_key(self.label.as_str()),
123            family: FamilyDescriptor::Embed {
124                model: self.label.to_string(),
125                dims: Some(capabilities.ndims),
126                max_documents: capabilities.max_documents,
127                modality: EmbedModality::Text,
128            },
129            layers: Vec::new(),
130        }
131    }
132
133    async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
134        let model = &self.model;
135        match kind {
136            EffectKind::Embed {
137                inputs: EmbedInputs::Texts(texts),
138            } => Reply::Outcome(
139                model
140                    .call(texts)
141                    .await
142                    .map(|response| Outcome::Embeddings(EmbedOutputs::Texts(response)))
143                    .map_err(ErrorReport::from),
144            ),
145            EffectKind::Embed {
146                inputs: EmbedInputs::Images(_),
147            } => Reply::Outcome(Err(ErrorReport::new(
148                ErrorKind::HandlerUnavailable,
149                "a text embedding handler cannot embed images",
150            ))),
151            other @ (EffectKind::Completion { .. }
152            | EffectKind::ToolCall { .. }
153            | EffectKind::Memory { .. }
154            | EffectKind::Retrieve { .. }
155            | EffectKind::Rerank { .. }
156            | EffectKind::Custom { .. }) => {
157                Reply::Outcome(Err(wrong_family(EffectFamily::Embed, &other)))
158            }
159        }
160    }
161}
162
163/// A reranking model orders documents for a query.
164impl Serve for ModelAdapter<Rerank> {
165    type Family = family::Rerank;
166
167    fn descriptor(&self) -> HandlerDescriptor {
168        HandlerDescriptor {
169            key: crate::effect::rerank_key(self.label.as_str()),
170            family: FamilyDescriptor::Rerank {
171                model: self.label.to_string(),
172                max_documents: self.model.capabilities().max_documents,
173            },
174            layers: Vec::new(),
175        }
176    }
177
178    async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
179        let model = &self.model;
180        match kind {
181            EffectKind::Rerank { request } => Reply::Outcome(
182                model
183                    .call(request)
184                    .await
185                    .map(Outcome::Reranked)
186                    .map_err(ErrorReport::from),
187            ),
188            other @ (EffectKind::Completion { .. }
189            | EffectKind::ToolCall { .. }
190            | EffectKind::Embed { .. }
191            | EffectKind::Memory { .. }
192            | EffectKind::Retrieve { .. }
193            | EffectKind::Custom { .. }) => {
194                Reply::Outcome(Err(wrong_family(EffectFamily::Rerank, &other)))
195            }
196        }
197    }
198}
199
200/// The context a tool call runs with: the driver's inbound values from
201/// the dispatch's scope (`ToolContext`, as `for_dispatch`), else empty, with
202/// every scope of the dispatch attached so the tool reaches its runtime by
203/// type for the length of the call.
204fn dispatch_context(dispatch: &Dispatch) -> crate::tool::ToolContext {
205    dispatch
206        .scope::<crate::tool::ToolContext>()
207        .map(|inbound| inbound.for_dispatch())
208        .unwrap_or_default()
209        .with_scopes(dispatch.scopes())
210}
211
212/// Hand what the tool published back beside the dispatch, when the driver
213/// asked for it ([`PublishedContext`](crate::tool::PublishedContext) in
214/// the dispatch's scope); the result carries data only.
215fn publish(dispatch: &Dispatch, context: crate::tool::ToolContext) {
216    if let Some(published) = dispatch.scope::<crate::tool::PublishedContext>() {
217        published.publish(context);
218    }
219}
220
221/// A [`Tool`] as a handler, keyed by its name.
222pub struct ToolAdapter<T> {
223    tool: T,
224    embedding: Option<ToolEmbeddingDescriptor>,
225}
226
227impl<T: Tool> ToolAdapter<T> {
228    /// Wrap a static tool.
229    pub fn new(tool: T) -> Self {
230        Self {
231            tool,
232            embedding: None,
233        }
234    }
235
236    /// Wraps a tool with embedding descriptions and serialized reconstruction
237    /// context. Returns an error if context serialization fails.
238    pub fn retrievable(tool: T) -> Result<Self, serde_json::Error>
239    where
240        T: ToolEmbedding,
241    {
242        let embedding = ToolEmbeddingDescriptor {
243            context: serde_json::to_value(tool.context())?,
244            embedding_docs: tool.embedding_docs(),
245        };
246        Ok(Self {
247            tool,
248            embedding: Some(embedding),
249        })
250    }
251
252    /// The wrapped tool.
253    pub fn tool(&self) -> &T {
254        &self.tool
255    }
256}
257
258impl<T> Serve for ToolAdapter<T>
259where
260    T: Tool + 'static,
261{
262    type Family = family::Tool;
263
264    fn descriptor(&self) -> HandlerDescriptor {
265        HandlerDescriptor {
266            key: crate::effect::tool_key(T::NAME),
267            family: FamilyDescriptor::Tool {
268                name: T::NAME.to_owned(),
269                description: self.tool.description(),
270                parameters: self.tool.parameters(),
271                embedding: self.embedding.clone(),
272            },
273            layers: Vec::new(),
274        }
275    }
276
277    async fn serve(&self, kind: EffectKind, dispatch: Dispatch) -> Reply {
278        match kind {
279            EffectKind::ToolCall { name, .. } if name != T::NAME => {
280                // Key routing must not invoke a different tool than the bound target.
281                Reply::Outcome(Err(ErrorReport::new(
282                    ErrorKind::Internal,
283                    format!("tool handler `{}` asked to run `{name}`", T::NAME),
284                )))
285            }
286            EffectKind::ToolCall { args, .. } => {
287                let mut context = dispatch_context(&dispatch);
288                let result = ErasedTool::execute(&self.tool, args, &mut context).await;
289                publish(&dispatch, context);
290                Reply::Outcome(Ok(Outcome::ToolResult { result }))
291            }
292            other @ (EffectKind::Completion { .. }
293            | EffectKind::Embed { .. }
294            | EffectKind::Memory { .. }
295            | EffectKind::Retrieve { .. }
296            | EffectKind::Rerank { .. }
297            | EffectKind::Custom { .. }) => {
298                Reply::Outcome(Err(wrong_family(EffectFamily::Tool, &other)))
299            }
300        }
301    }
302}
303
304/// Contextual tool callback accepting JSON arguments and returning canonical
305/// output or a tool execution error.
306pub trait ToolCallback:
307    for<'a> Fn(
308        &'a mut crate::tool::ToolContext,
309        serde_json::Value,
310    ) -> WasmBoxedFuture<
311        'a,
312        Result<crate::tool::ToolOutput, crate::tool::ToolExecutionError>,
313    > + WasmCompatSend
314    + WasmCompatSync
315{
316}
317
318impl<F> ToolCallback for F where
319    F: for<'a> Fn(
320            &'a mut crate::tool::ToolContext,
321            serde_json::Value,
322        ) -> WasmBoxedFuture<
323            'a,
324            Result<crate::tool::ToolOutput, crate::tool::ToolExecutionError>,
325        > + WasmCompatSend
326        + WasmCompatSync
327{
328}
329
330/// A runtime-defined tool handler with a name, argument schema, and callback.
331pub struct ToolFn<F> {
332    name: String,
333    description: String,
334    parameters: serde_json::Value,
335    callback: F,
336}
337
338impl<F: ToolCallback> ToolFn<F> {
339    /// Build a runtime-defined tool.
340    pub fn new(
341        name: impl Into<String>,
342        description: impl Into<String>,
343        parameters: serde_json::Value,
344        callback: F,
345    ) -> Self {
346        Self {
347            name: name.into(),
348            description: description.into(),
349            parameters,
350            callback,
351        }
352    }
353
354    /// The tool's name.
355    pub fn name(&self) -> &str {
356        &self.name
357    }
358}
359
360impl<F> Serve for ToolFn<F>
361where
362    F: ToolCallback + 'static,
363{
364    type Family = family::Tool;
365
366    fn descriptor(&self) -> HandlerDescriptor {
367        HandlerDescriptor {
368            key: crate::effect::tool_key(&self.name),
369            family: FamilyDescriptor::Tool {
370                name: self.name.clone(),
371                description: self.description.clone(),
372                parameters: self.parameters.clone(),
373                embedding: None,
374            },
375            layers: Vec::new(),
376        }
377    }
378
379    async fn serve(&self, kind: EffectKind, dispatch: Dispatch) -> Reply {
380        match kind {
381            EffectKind::ToolCall { args, .. } => {
382                let mut context = dispatch_context(&dispatch);
383                let result =
384                    crate::tool::contextual::execute_callback(&self.callback, args, &mut context)
385                        .await;
386                publish(&dispatch, context);
387                Reply::Outcome(Ok(Outcome::ToolResult { result }))
388            }
389            other @ (EffectKind::Completion { .. }
390            | EffectKind::Embed { .. }
391            | EffectKind::Memory { .. }
392            | EffectKind::Retrieve { .. }
393            | EffectKind::Rerank { .. }
394            | EffectKind::Custom { .. }) => {
395                Reply::Outcome(Err(wrong_family(EffectFamily::Tool, &other)))
396            }
397        }
398    }
399}
400
401/// A [`ConversationMemory`] as a handler.
402pub struct MemoryAdapter<M> {
403    memory: M,
404    label: Option<String>,
405}
406
407impl<M> MemoryAdapter<M> {
408    /// Wrap `memory` as an agent's own memory, under the bare `memory` key.
409    pub fn new(memory: M) -> Self {
410        Self {
411            memory,
412            label: None,
413        }
414    }
415
416    /// Wrap `memory` under `memory:<label>`, so a host can serve several
417    /// backends (per tenant, per store) on one bus and dispatch to each by
418    /// [`memory_key`](crate::effect::memory_key).
419    pub fn labelled(label: impl Into<String>, memory: M) -> Self {
420        Self {
421            memory,
422            label: Some(label.into()),
423        }
424    }
425
426    /// The wrapped backend.
427    pub fn memory(&self) -> &M {
428        &self.memory
429    }
430}
431
432impl<M> Serve for MemoryAdapter<M>
433where
434    M: ConversationMemory + 'static,
435{
436    type Family = family::Memory;
437
438    fn descriptor(&self) -> HandlerDescriptor {
439        HandlerDescriptor {
440            key: self
441                .label
442                .as_deref()
443                .map_or_else(|| HandlerKey::from("memory"), crate::effect::memory_key),
444            family: FamilyDescriptor::Memory {},
445            layers: Vec::new(),
446        }
447    }
448
449    async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
450        match kind {
451            EffectKind::Memory { op } => {
452                let outcome = match op {
453                    MemoryOp::Load { conversation } => self
454                        .memory
455                        .load(&conversation)
456                        .await
457                        .map(|messages| Outcome::Memory(MemoryOutcome::Loaded { messages })),
458                    MemoryOp::Append {
459                        conversation,
460                        messages,
461                    } => self
462                        .memory
463                        .append(&conversation, messages)
464                        .await
465                        .map(|()| Outcome::Memory(MemoryOutcome::Appended)),
466                    MemoryOp::Clear { conversation } => self
467                        .memory
468                        .clear(&conversation)
469                        .await
470                        .map(|()| Outcome::Memory(MemoryOutcome::Cleared)),
471                };
472                Reply::Outcome(outcome.map_err(ErrorReport::from))
473            }
474            other @ (EffectKind::Completion { .. }
475            | EffectKind::ToolCall { .. }
476            | EffectKind::Embed { .. }
477            | EffectKind::Retrieve { .. }
478            | EffectKind::Rerank { .. }
479            | EffectKind::Custom { .. }) => {
480                Reply::Outcome(Err(wrong_family(EffectFamily::Memory, &other)))
481            }
482        }
483    }
484}
485
486/// A [`VectorStoreIndex`] as a handler. The index's filter type is rebuilt
487/// from the dynamic filter on the wire; documents come back as JSON, and the
488/// typed view deserialises on the client side.
489pub struct RetrieveAdapter<I> {
490    index: I,
491    label: Option<String>,
492}
493
494impl<I> RetrieveAdapter<I> {
495    /// Wrap `index` as an agent's own index, under the bare `retrieve` key.
496    pub fn new(index: I) -> Self {
497        Self { index, label: None }
498    }
499
500    /// Wrap `index` under `retrieve:<label>`, so a host can serve several
501    /// indexes on one bus and dispatch to each by
502    /// [`retrieve_key`](crate::effect::retrieve_key).
503    pub fn labelled(label: impl Into<String>, index: I) -> Self {
504        Self {
505            index,
506            label: Some(label.into()),
507        }
508    }
509
510    /// The wrapped index.
511    pub fn index(&self) -> &I {
512        &self.index
513    }
514}
515
516impl<I, F> Serve for RetrieveAdapter<I>
517where
518    I: VectorStoreIndex<Filter = F> + 'static,
519    F: DynamicSearchFilter + WasmCompatSend + WasmCompatSync + 'static,
520{
521    type Family = family::Retrieve;
522
523    fn descriptor(&self) -> HandlerDescriptor {
524        HandlerDescriptor {
525            key: self
526                .label
527                .as_deref()
528                .map_or_else(|| HandlerKey::from("retrieve"), crate::effect::retrieve_key),
529            family: FamilyDescriptor::Retrieve {},
530            layers: Vec::new(),
531        }
532    }
533
534    async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
535        match kind {
536            EffectKind::Retrieve { query } => {
537                let outcome = match query {
538                    RetrieveQuery::TopN { req } => {
539                        match req.try_map_filter(F::from_dynamic_filter) {
540                            Ok(req) => self
541                                .index
542                                .top_n::<serde_json::Value>(req)
543                                .await
544                                .map(|results| {
545                                    Outcome::Documents(RetrievedDocuments::Scored(
546                                        results
547                                            .into_iter()
548                                            .map(|result| {
549                                                (
550                                                    result.score,
551                                                    result.id,
552                                                    F::normalize_dynamic_document(result.document),
553                                                )
554                                            })
555                                            .collect(),
556                                    ))
557                                })
558                                .map_err(ErrorReport::from),
559                            Err(error) => Err(ErrorReport::from(VectorStoreError::from(error))),
560                        }
561                    }
562                    RetrieveQuery::TopNIds { req } => {
563                        match req.try_map_filter(F::from_dynamic_filter) {
564                            Ok(req) => self
565                                .index
566                                .top_n_ids(req)
567                                .await
568                                .map(|results| {
569                                    Outcome::Documents(RetrievedDocuments::Ids(
570                                        results
571                                            .into_iter()
572                                            .map(|result| (result.score, result.id))
573                                            .collect(),
574                                    ))
575                                })
576                                .map_err(ErrorReport::from),
577                            Err(error) => Err(ErrorReport::from(VectorStoreError::from(error))),
578                        }
579                    }
580                };
581                Reply::Outcome(outcome)
582            }
583            other @ (EffectKind::Completion { .. }
584            | EffectKind::ToolCall { .. }
585            | EffectKind::Embed { .. }
586            | EffectKind::Memory { .. }
587            | EffectKind::Rerank { .. }
588            | EffectKind::Custom { .. }) => {
589                Reply::Outcome(Err(wrong_family(EffectFamily::Retrieve, &other)))
590            }
591        }
592    }
593}