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(crate::operation::RerankRequest {
184                        query: request.query,
185                        documents: request.documents,
186                    })
187                    .await
188                    .map(Outcome::Reranked)
189                    .map_err(ErrorReport::from),
190            ),
191            other @ (EffectKind::Completion { .. }
192            | EffectKind::ToolCall { .. }
193            | EffectKind::Embed { .. }
194            | EffectKind::Memory { .. }
195            | EffectKind::Retrieve { .. }
196            | EffectKind::Custom { .. }) => {
197                Reply::Outcome(Err(wrong_family(EffectFamily::Rerank, &other)))
198            }
199        }
200    }
201}
202
203/// The context a tool call runs with: the driver's inbound values from
204/// the dispatch's scope (`ToolContext`, as `for_dispatch`), else empty, with
205/// every scope of the dispatch attached so the tool reaches its runtime by
206/// type for the length of the call.
207fn dispatch_context(dispatch: &Dispatch) -> crate::tool::ToolContext {
208    dispatch
209        .scope::<crate::tool::ToolContext>()
210        .map(|inbound| inbound.for_dispatch())
211        .unwrap_or_default()
212        .with_scopes(dispatch.scopes())
213}
214
215/// Hand what the tool published back beside the dispatch, when the driver
216/// asked for it ([`PublishedContext`](crate::tool::PublishedContext) in
217/// the dispatch's scope); the result carries data only.
218fn publish(dispatch: &Dispatch, context: crate::tool::ToolContext) {
219    if let Some(published) = dispatch.scope::<crate::tool::PublishedContext>() {
220        published.publish(context);
221    }
222}
223
224/// A [`Tool`] as a handler, keyed by its name.
225pub struct ToolAdapter<T> {
226    tool: T,
227    embedding: Option<ToolEmbeddingDescriptor>,
228}
229
230impl<T: Tool> ToolAdapter<T> {
231    /// Wrap a static tool.
232    pub fn new(tool: T) -> Self {
233        Self {
234            tool,
235            embedding: None,
236        }
237    }
238
239    /// Wraps a tool with embedding descriptions and serialized reconstruction
240    /// context. Returns an error if context serialization fails.
241    pub fn retrievable(tool: T) -> Result<Self, serde_json::Error>
242    where
243        T: ToolEmbedding,
244    {
245        let embedding = ToolEmbeddingDescriptor {
246            context: serde_json::to_value(tool.context())?,
247            embedding_docs: tool.embedding_docs(),
248        };
249        Ok(Self {
250            tool,
251            embedding: Some(embedding),
252        })
253    }
254
255    /// The wrapped tool.
256    pub fn tool(&self) -> &T {
257        &self.tool
258    }
259}
260
261impl<T> Serve for ToolAdapter<T>
262where
263    T: Tool + 'static,
264{
265    type Family = family::Tool;
266
267    fn descriptor(&self) -> HandlerDescriptor {
268        HandlerDescriptor {
269            key: crate::effect::tool_key(T::NAME),
270            family: FamilyDescriptor::Tool {
271                name: T::NAME.to_owned(),
272                description: self.tool.description(),
273                parameters: self.tool.parameters(),
274                embedding: self.embedding.clone(),
275            },
276            layers: Vec::new(),
277        }
278    }
279
280    async fn serve(&self, kind: EffectKind, dispatch: Dispatch) -> Reply {
281        match kind {
282            EffectKind::ToolCall { name, .. } if name != T::NAME => {
283                // Key routing must not invoke a different tool than the bound target.
284                Reply::Outcome(Err(ErrorReport::new(
285                    ErrorKind::Internal,
286                    format!("tool handler `{}` asked to run `{name}`", T::NAME),
287                )))
288            }
289            EffectKind::ToolCall { args, .. } => {
290                let mut context = dispatch_context(&dispatch);
291                let result = ErasedTool::execute(&self.tool, args, &mut context).await;
292                publish(&dispatch, context);
293                Reply::Outcome(Ok(Outcome::ToolResult { result }))
294            }
295            other @ (EffectKind::Completion { .. }
296            | EffectKind::Embed { .. }
297            | EffectKind::Memory { .. }
298            | EffectKind::Retrieve { .. }
299            | EffectKind::Rerank { .. }
300            | EffectKind::Custom { .. }) => {
301                Reply::Outcome(Err(wrong_family(EffectFamily::Tool, &other)))
302            }
303        }
304    }
305}
306
307/// Contextual tool callback accepting JSON arguments and returning canonical
308/// output or a tool execution error.
309pub trait ToolCallback:
310    for<'a> Fn(
311        &'a mut crate::tool::ToolContext,
312        serde_json::Value,
313    ) -> WasmBoxedFuture<
314        'a,
315        Result<crate::tool::ToolOutput, crate::tool::ToolExecutionError>,
316    > + WasmCompatSend
317    + WasmCompatSync
318{
319}
320
321impl<F> ToolCallback for F where
322    F: for<'a> Fn(
323            &'a mut crate::tool::ToolContext,
324            serde_json::Value,
325        ) -> WasmBoxedFuture<
326            'a,
327            Result<crate::tool::ToolOutput, crate::tool::ToolExecutionError>,
328        > + WasmCompatSend
329        + WasmCompatSync
330{
331}
332
333/// A runtime-defined tool handler with a name, argument schema, and callback.
334pub struct ToolFn<F> {
335    name: String,
336    description: String,
337    parameters: serde_json::Value,
338    callback: F,
339}
340
341impl<F: ToolCallback> ToolFn<F> {
342    /// Build a runtime-defined tool.
343    pub fn new(
344        name: impl Into<String>,
345        description: impl Into<String>,
346        parameters: serde_json::Value,
347        callback: F,
348    ) -> Self {
349        Self {
350            name: name.into(),
351            description: description.into(),
352            parameters,
353            callback,
354        }
355    }
356
357    /// The tool's name.
358    pub fn name(&self) -> &str {
359        &self.name
360    }
361}
362
363impl<F> Serve for ToolFn<F>
364where
365    F: ToolCallback + 'static,
366{
367    type Family = family::Tool;
368
369    fn descriptor(&self) -> HandlerDescriptor {
370        HandlerDescriptor {
371            key: crate::effect::tool_key(&self.name),
372            family: FamilyDescriptor::Tool {
373                name: self.name.clone(),
374                description: self.description.clone(),
375                parameters: self.parameters.clone(),
376                embedding: None,
377            },
378            layers: Vec::new(),
379        }
380    }
381
382    async fn serve(&self, kind: EffectKind, dispatch: Dispatch) -> Reply {
383        match kind {
384            EffectKind::ToolCall { args, .. } => {
385                let mut context = dispatch_context(&dispatch);
386                let result =
387                    crate::tool::contextual::execute_callback(&self.callback, args, &mut context)
388                        .await;
389                publish(&dispatch, context);
390                Reply::Outcome(Ok(Outcome::ToolResult { result }))
391            }
392            other @ (EffectKind::Completion { .. }
393            | EffectKind::Embed { .. }
394            | EffectKind::Memory { .. }
395            | EffectKind::Retrieve { .. }
396            | EffectKind::Rerank { .. }
397            | EffectKind::Custom { .. }) => {
398                Reply::Outcome(Err(wrong_family(EffectFamily::Tool, &other)))
399            }
400        }
401    }
402}
403
404/// A [`ConversationMemory`] as a handler.
405pub struct MemoryAdapter<M> {
406    memory: M,
407    label: Option<String>,
408}
409
410impl<M> MemoryAdapter<M> {
411    /// Wrap `memory` as an agent's own memory, under the bare `memory` key.
412    pub fn new(memory: M) -> Self {
413        Self {
414            memory,
415            label: None,
416        }
417    }
418
419    /// Wrap `memory` under `memory:<label>`, so a host can serve several
420    /// backends (per tenant, per store) on one bus and dispatch to each by
421    /// [`memory_key`](crate::effect::memory_key).
422    pub fn labelled(label: impl Into<String>, memory: M) -> Self {
423        Self {
424            memory,
425            label: Some(label.into()),
426        }
427    }
428
429    /// The wrapped backend.
430    pub fn memory(&self) -> &M {
431        &self.memory
432    }
433}
434
435impl<M> Serve for MemoryAdapter<M>
436where
437    M: ConversationMemory + 'static,
438{
439    type Family = family::Memory;
440
441    fn descriptor(&self) -> HandlerDescriptor {
442        HandlerDescriptor {
443            key: self
444                .label
445                .as_deref()
446                .map_or_else(|| HandlerKey::from("memory"), crate::effect::memory_key),
447            family: FamilyDescriptor::Memory {},
448            layers: Vec::new(),
449        }
450    }
451
452    async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
453        match kind {
454            EffectKind::Memory { op } => {
455                let outcome = match op {
456                    MemoryOp::Load { conversation } => self
457                        .memory
458                        .load(&conversation)
459                        .await
460                        .map(|messages| Outcome::Memory(MemoryOutcome::Loaded { messages })),
461                    MemoryOp::Append {
462                        conversation,
463                        messages,
464                    } => self
465                        .memory
466                        .append(&conversation, messages)
467                        .await
468                        .map(|()| Outcome::Memory(MemoryOutcome::Appended)),
469                    MemoryOp::Clear { conversation } => self
470                        .memory
471                        .clear(&conversation)
472                        .await
473                        .map(|()| Outcome::Memory(MemoryOutcome::Cleared)),
474                };
475                Reply::Outcome(outcome.map_err(ErrorReport::from))
476            }
477            other @ (EffectKind::Completion { .. }
478            | EffectKind::ToolCall { .. }
479            | EffectKind::Embed { .. }
480            | EffectKind::Retrieve { .. }
481            | EffectKind::Rerank { .. }
482            | EffectKind::Custom { .. }) => {
483                Reply::Outcome(Err(wrong_family(EffectFamily::Memory, &other)))
484            }
485        }
486    }
487}
488
489/// A [`VectorStoreIndex`] as a handler. The index's filter type is rebuilt
490/// from the dynamic filter on the wire; documents come back as JSON, and the
491/// typed view deserialises on the client side.
492pub struct RetrieveAdapter<I> {
493    index: I,
494    label: Option<String>,
495}
496
497impl<I> RetrieveAdapter<I> {
498    /// Wrap `index` as an agent's own index, under the bare `retrieve` key.
499    pub fn new(index: I) -> Self {
500        Self { index, label: None }
501    }
502
503    /// Wrap `index` under `retrieve:<label>`, so a host can serve several
504    /// indexes on one bus and dispatch to each by
505    /// [`retrieve_key`](crate::effect::retrieve_key).
506    pub fn labelled(label: impl Into<String>, index: I) -> Self {
507        Self {
508            index,
509            label: Some(label.into()),
510        }
511    }
512
513    /// The wrapped index.
514    pub fn index(&self) -> &I {
515        &self.index
516    }
517}
518
519impl<I, F> Serve for RetrieveAdapter<I>
520where
521    I: VectorStoreIndex<Filter = F> + 'static,
522    F: DynamicSearchFilter + WasmCompatSend + WasmCompatSync + 'static,
523{
524    type Family = family::Retrieve;
525
526    fn descriptor(&self) -> HandlerDescriptor {
527        HandlerDescriptor {
528            key: self
529                .label
530                .as_deref()
531                .map_or_else(|| HandlerKey::from("retrieve"), crate::effect::retrieve_key),
532            family: FamilyDescriptor::Retrieve {},
533            layers: Vec::new(),
534        }
535    }
536
537    async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
538        match kind {
539            EffectKind::Retrieve { query } => {
540                let outcome = match query {
541                    RetrieveQuery::TopN { req } => {
542                        match req.try_map_filter(F::from_dynamic_filter) {
543                            Ok(req) => self
544                                .index
545                                .top_n::<serde_json::Value>(req)
546                                .await
547                                .map(|results| {
548                                    Outcome::Documents(RetrievedDocuments::Scored(
549                                        results
550                                            .into_iter()
551                                            .map(|(score, id, doc)| {
552                                                (score, id, F::normalize_dynamic_document(doc))
553                                            })
554                                            .collect(),
555                                    ))
556                                })
557                                .map_err(ErrorReport::from),
558                            Err(error) => Err(ErrorReport::from(VectorStoreError::from(error))),
559                        }
560                    }
561                    RetrieveQuery::TopNIds { req } => {
562                        match req.try_map_filter(F::from_dynamic_filter) {
563                            Ok(req) => self
564                                .index
565                                .top_n_ids(req)
566                                .await
567                                .map(|results| Outcome::Documents(RetrievedDocuments::Ids(results)))
568                                .map_err(ErrorReport::from),
569                            Err(error) => Err(ErrorReport::from(VectorStoreError::from(error))),
570                        }
571                    }
572                };
573                Reply::Outcome(outcome)
574            }
575            other @ (EffectKind::Completion { .. }
576            | EffectKind::ToolCall { .. }
577            | EffectKind::Embed { .. }
578            | EffectKind::Memory { .. }
579            | EffectKind::Rerank { .. }
580            | EffectKind::Custom { .. }) => {
581                Reply::Outcome(Err(wrong_family(EffectFamily::Retrieve, &other)))
582            }
583        }
584    }
585}