1use 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
42pub struct ModelAdapter<Op: Operation> {
46 label: ModelRef,
47 model: DynModel<Op>,
48}
49
50impl<Op: Operation> ModelAdapter<Op> {
51 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
60impl 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
115impl 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
163impl 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
200fn 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
212fn publish(dispatch: &Dispatch, context: crate::tool::ToolContext) {
216 if let Some(published) = dispatch.scope::<crate::tool::PublishedContext>() {
217 published.publish(context);
218 }
219}
220
221pub struct ToolAdapter<T> {
223 tool: T,
224 embedding: Option<ToolEmbeddingDescriptor>,
225}
226
227impl<T: Tool> ToolAdapter<T> {
228 pub fn new(tool: T) -> Self {
230 Self {
231 tool,
232 embedding: None,
233 }
234 }
235
236 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 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 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
304pub 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
330pub 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 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 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
401pub struct MemoryAdapter<M> {
403 memory: M,
404 label: Option<String>,
405}
406
407impl<M> MemoryAdapter<M> {
408 pub fn new(memory: M) -> Self {
410 Self {
411 memory,
412 label: None,
413 }
414 }
415
416 pub fn labelled(label: impl Into<String>, memory: M) -> Self {
420 Self {
421 memory,
422 label: Some(label.into()),
423 }
424 }
425
426 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
486pub struct RetrieveAdapter<I> {
490 index: I,
491 label: Option<String>,
492}
493
494impl<I> RetrieveAdapter<I> {
495 pub fn new(index: I) -> Self {
497 Self { index, label: None }
498 }
499
500 pub fn labelled(label: impl Into<String>, index: I) -> Self {
504 Self {
505 index,
506 label: Some(label.into()),
507 }
508 }
509
510 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}