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(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
203fn 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
215fn publish(dispatch: &Dispatch, context: crate::tool::ToolContext) {
219 if let Some(published) = dispatch.scope::<crate::tool::PublishedContext>() {
220 published.publish(context);
221 }
222}
223
224pub struct ToolAdapter<T> {
226 tool: T,
227 embedding: Option<ToolEmbeddingDescriptor>,
228}
229
230impl<T: Tool> ToolAdapter<T> {
231 pub fn new(tool: T) -> Self {
233 Self {
234 tool,
235 embedding: None,
236 }
237 }
238
239 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 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 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
307pub 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
333pub 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 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 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
404pub struct MemoryAdapter<M> {
406 memory: M,
407 label: Option<String>,
408}
409
410impl<M> MemoryAdapter<M> {
411 pub fn new(memory: M) -> Self {
413 Self {
414 memory,
415 label: None,
416 }
417 }
418
419 pub fn labelled(label: impl Into<String>, memory: M) -> Self {
423 Self {
424 memory,
425 label: Some(label.into()),
426 }
427 }
428
429 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
489pub struct RetrieveAdapter<I> {
493 index: I,
494 label: Option<String>,
495}
496
497impl<I> RetrieveAdapter<I> {
498 pub fn new(index: I) -> Self {
500 Self { index, label: None }
501 }
502
503 pub fn labelled(label: impl Into<String>, index: I) -> Self {
507 Self {
508 index,
509 label: Some(label.into()),
510 }
511 }
512
513 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}