use crate::{
completion::ModelRef,
driver::DynModel,
effect::{
EffectFamily, EffectKind, EmbedInputs, EmbedModality, EmbedOutputs, FamilyDescriptor,
HandlerDescriptor, HandlerKey, MemoryOp, MemoryOutcome, Outcome, RetrieveQuery,
RetrievedDocuments, ToolEmbeddingDescriptor,
},
error::{ErrorKind, ErrorReport},
memory::ConversationMemory,
operation::{Completion, Embedding, Rerank},
tool::{ErasedTool, Tool, ToolEmbedding},
vector_store::{VectorStoreError, VectorStoreIndex, request::DynamicSearchFilter},
wasm_compat::{WasmBoxedFuture, WasmCompatSend, WasmCompatSync},
wire::Operation,
};
use super::{Dispatch, Reply, Serve};
use crate::effect::family;
fn wrong_family(handler: EffectFamily, kind: &EffectKind) -> ErrorReport {
ErrorReport::new(
ErrorKind::HandlerUnavailable,
format!(
"a {handler} handler cannot serve a `{}` effect",
kind.name()
),
)
}
pub struct ModelAdapter<Op: Operation> {
label: ModelRef,
model: DynModel<Op>,
}
impl<Op: Operation> ModelAdapter<Op> {
pub fn new(label: impl Into<ModelRef>, model: impl Into<DynModel<Op>>) -> Self {
Self {
label: label.into(),
model: model.into(),
}
}
}
impl Serve for ModelAdapter<Completion> {
type Family = family::Completion;
fn descriptor(&self) -> HandlerDescriptor {
HandlerDescriptor {
key: crate::effect::model_key(self.label.as_str()),
family: FamilyDescriptor::Completion {
model: self.label.clone(),
capabilities: self.model.capabilities().completion,
},
layers: Vec::new(),
}
}
async fn serve(&self, kind: EffectKind, dispatch: Dispatch) -> Reply {
let model = &self.model;
let context = dispatch.adapter_context();
match kind {
EffectKind::Completion {
request,
stream: false,
} => {
let result = match context {
Some(context) => model.call_observed(request, context).await,
None => model.call(request).await,
};
Reply::Outcome(result.map(Outcome::Completion).map_err(ErrorReport::from))
}
EffectKind::Completion {
request,
stream: true,
} => {
let opened = match context {
Some(context) => model.stream_observed(request, context),
None => model.stream(request),
};
match opened {
Ok(stream) => Reply::Stream(stream.into_relay()),
Err(error) => Reply::Outcome(Err(ErrorReport::from(error))),
}
}
other @ (EffectKind::ToolCall { .. }
| EffectKind::Embed { .. }
| EffectKind::Memory { .. }
| EffectKind::Retrieve { .. }
| EffectKind::Rerank { .. }
| EffectKind::Custom { .. }) => {
Reply::Outcome(Err(wrong_family(EffectFamily::Completion, &other)))
}
}
}
}
impl Serve for ModelAdapter<Embedding> {
type Family = family::Embed;
fn descriptor(&self) -> HandlerDescriptor {
let capabilities = self.model.capabilities();
HandlerDescriptor {
key: crate::effect::embed_key(self.label.as_str()),
family: FamilyDescriptor::Embed {
model: self.label.to_string(),
dims: Some(capabilities.ndims),
max_documents: capabilities.max_documents,
modality: EmbedModality::Text,
},
layers: Vec::new(),
}
}
async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
let model = &self.model;
match kind {
EffectKind::Embed {
inputs: EmbedInputs::Texts(texts),
} => Reply::Outcome(
model
.call(texts)
.await
.map(|response| Outcome::Embeddings(EmbedOutputs::Texts(response)))
.map_err(ErrorReport::from),
),
EffectKind::Embed {
inputs: EmbedInputs::Images(_),
} => Reply::Outcome(Err(ErrorReport::new(
ErrorKind::HandlerUnavailable,
"a text embedding handler cannot embed images",
))),
other @ (EffectKind::Completion { .. }
| EffectKind::ToolCall { .. }
| EffectKind::Memory { .. }
| EffectKind::Retrieve { .. }
| EffectKind::Rerank { .. }
| EffectKind::Custom { .. }) => {
Reply::Outcome(Err(wrong_family(EffectFamily::Embed, &other)))
}
}
}
}
impl Serve for ModelAdapter<Rerank> {
type Family = family::Rerank;
fn descriptor(&self) -> HandlerDescriptor {
HandlerDescriptor {
key: crate::effect::rerank_key(self.label.as_str()),
family: FamilyDescriptor::Rerank {
model: self.label.to_string(),
max_documents: self.model.capabilities().max_documents,
},
layers: Vec::new(),
}
}
async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
let model = &self.model;
match kind {
EffectKind::Rerank { request } => Reply::Outcome(
model
.call(request)
.await
.map(Outcome::Reranked)
.map_err(ErrorReport::from),
),
other @ (EffectKind::Completion { .. }
| EffectKind::ToolCall { .. }
| EffectKind::Embed { .. }
| EffectKind::Memory { .. }
| EffectKind::Retrieve { .. }
| EffectKind::Custom { .. }) => {
Reply::Outcome(Err(wrong_family(EffectFamily::Rerank, &other)))
}
}
}
}
fn dispatch_context(dispatch: &Dispatch) -> crate::tool::ToolContext {
dispatch
.scope::<crate::tool::ToolContext>()
.map(|inbound| inbound.for_dispatch())
.unwrap_or_default()
.with_scopes(dispatch.scopes())
}
fn publish(dispatch: &Dispatch, context: crate::tool::ToolContext) {
if let Some(published) = dispatch.scope::<crate::tool::PublishedContext>() {
published.publish(context);
}
}
pub struct ToolAdapter<T> {
tool: T,
embedding: Option<ToolEmbeddingDescriptor>,
}
impl<T: Tool> ToolAdapter<T> {
pub fn new(tool: T) -> Self {
Self {
tool,
embedding: None,
}
}
pub fn retrievable(tool: T) -> Result<Self, serde_json::Error>
where
T: ToolEmbedding,
{
let embedding = ToolEmbeddingDescriptor {
context: serde_json::to_value(tool.context())?,
embedding_docs: tool.embedding_docs(),
};
Ok(Self {
tool,
embedding: Some(embedding),
})
}
pub fn tool(&self) -> &T {
&self.tool
}
}
impl<T> Serve for ToolAdapter<T>
where
T: Tool + 'static,
{
type Family = family::Tool;
fn descriptor(&self) -> HandlerDescriptor {
HandlerDescriptor {
key: crate::effect::tool_key(T::NAME),
family: FamilyDescriptor::Tool {
name: T::NAME.to_owned(),
description: self.tool.description(),
parameters: self.tool.parameters(),
embedding: self.embedding.clone(),
},
layers: Vec::new(),
}
}
async fn serve(&self, kind: EffectKind, dispatch: Dispatch) -> Reply {
match kind {
EffectKind::ToolCall { name, .. } if name != T::NAME => {
Reply::Outcome(Err(ErrorReport::new(
ErrorKind::Internal,
format!("tool handler `{}` asked to run `{name}`", T::NAME),
)))
}
EffectKind::ToolCall { args, .. } => {
let mut context = dispatch_context(&dispatch);
let result = ErasedTool::execute(&self.tool, args, &mut context).await;
publish(&dispatch, context);
Reply::Outcome(Ok(Outcome::ToolResult { result }))
}
other @ (EffectKind::Completion { .. }
| EffectKind::Embed { .. }
| EffectKind::Memory { .. }
| EffectKind::Retrieve { .. }
| EffectKind::Rerank { .. }
| EffectKind::Custom { .. }) => {
Reply::Outcome(Err(wrong_family(EffectFamily::Tool, &other)))
}
}
}
}
pub trait ToolCallback:
for<'a> Fn(
&'a mut crate::tool::ToolContext,
serde_json::Value,
) -> WasmBoxedFuture<
'a,
Result<crate::tool::ToolOutput, crate::tool::ToolExecutionError>,
> + WasmCompatSend
+ WasmCompatSync
{
}
impl<F> ToolCallback for F where
F: for<'a> Fn(
&'a mut crate::tool::ToolContext,
serde_json::Value,
) -> WasmBoxedFuture<
'a,
Result<crate::tool::ToolOutput, crate::tool::ToolExecutionError>,
> + WasmCompatSend
+ WasmCompatSync
{
}
pub struct ToolFn<F> {
name: String,
description: String,
parameters: serde_json::Value,
callback: F,
}
impl<F: ToolCallback> ToolFn<F> {
pub fn new(
name: impl Into<String>,
description: impl Into<String>,
parameters: serde_json::Value,
callback: F,
) -> Self {
Self {
name: name.into(),
description: description.into(),
parameters,
callback,
}
}
pub fn name(&self) -> &str {
&self.name
}
}
impl<F> Serve for ToolFn<F>
where
F: ToolCallback + 'static,
{
type Family = family::Tool;
fn descriptor(&self) -> HandlerDescriptor {
HandlerDescriptor {
key: crate::effect::tool_key(&self.name),
family: FamilyDescriptor::Tool {
name: self.name.clone(),
description: self.description.clone(),
parameters: self.parameters.clone(),
embedding: None,
},
layers: Vec::new(),
}
}
async fn serve(&self, kind: EffectKind, dispatch: Dispatch) -> Reply {
match kind {
EffectKind::ToolCall { args, .. } => {
let mut context = dispatch_context(&dispatch);
let result =
crate::tool::contextual::execute_callback(&self.callback, args, &mut context)
.await;
publish(&dispatch, context);
Reply::Outcome(Ok(Outcome::ToolResult { result }))
}
other @ (EffectKind::Completion { .. }
| EffectKind::Embed { .. }
| EffectKind::Memory { .. }
| EffectKind::Retrieve { .. }
| EffectKind::Rerank { .. }
| EffectKind::Custom { .. }) => {
Reply::Outcome(Err(wrong_family(EffectFamily::Tool, &other)))
}
}
}
}
pub struct MemoryAdapter<M> {
memory: M,
label: Option<String>,
}
impl<M> MemoryAdapter<M> {
pub fn new(memory: M) -> Self {
Self {
memory,
label: None,
}
}
pub fn labelled(label: impl Into<String>, memory: M) -> Self {
Self {
memory,
label: Some(label.into()),
}
}
pub fn memory(&self) -> &M {
&self.memory
}
}
impl<M> Serve for MemoryAdapter<M>
where
M: ConversationMemory + 'static,
{
type Family = family::Memory;
fn descriptor(&self) -> HandlerDescriptor {
HandlerDescriptor {
key: self
.label
.as_deref()
.map_or_else(|| HandlerKey::from("memory"), crate::effect::memory_key),
family: FamilyDescriptor::Memory {},
layers: Vec::new(),
}
}
async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
match kind {
EffectKind::Memory { op } => {
let outcome = match op {
MemoryOp::Load { conversation } => self
.memory
.load(&conversation)
.await
.map(|messages| Outcome::Memory(MemoryOutcome::Loaded { messages })),
MemoryOp::Append {
conversation,
messages,
} => self
.memory
.append(&conversation, messages)
.await
.map(|()| Outcome::Memory(MemoryOutcome::Appended)),
MemoryOp::Clear { conversation } => self
.memory
.clear(&conversation)
.await
.map(|()| Outcome::Memory(MemoryOutcome::Cleared)),
};
Reply::Outcome(outcome.map_err(ErrorReport::from))
}
other @ (EffectKind::Completion { .. }
| EffectKind::ToolCall { .. }
| EffectKind::Embed { .. }
| EffectKind::Retrieve { .. }
| EffectKind::Rerank { .. }
| EffectKind::Custom { .. }) => {
Reply::Outcome(Err(wrong_family(EffectFamily::Memory, &other)))
}
}
}
}
pub struct RetrieveAdapter<I> {
index: I,
label: Option<String>,
}
impl<I> RetrieveAdapter<I> {
pub fn new(index: I) -> Self {
Self { index, label: None }
}
pub fn labelled(label: impl Into<String>, index: I) -> Self {
Self {
index,
label: Some(label.into()),
}
}
pub fn index(&self) -> &I {
&self.index
}
}
impl<I, F> Serve for RetrieveAdapter<I>
where
I: VectorStoreIndex<Filter = F> + 'static,
F: DynamicSearchFilter + WasmCompatSend + WasmCompatSync + 'static,
{
type Family = family::Retrieve;
fn descriptor(&self) -> HandlerDescriptor {
HandlerDescriptor {
key: self
.label
.as_deref()
.map_or_else(|| HandlerKey::from("retrieve"), crate::effect::retrieve_key),
family: FamilyDescriptor::Retrieve {},
layers: Vec::new(),
}
}
async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
match kind {
EffectKind::Retrieve { query } => {
let outcome = match query {
RetrieveQuery::TopN { req } => {
match req.try_map_filter(F::from_dynamic_filter) {
Ok(req) => self
.index
.top_n::<serde_json::Value>(req)
.await
.map(|results| {
Outcome::Documents(RetrievedDocuments::Scored(
results
.into_iter()
.map(|result| {
(
result.score,
result.id,
F::normalize_dynamic_document(result.document),
)
})
.collect(),
))
})
.map_err(ErrorReport::from),
Err(error) => Err(ErrorReport::from(VectorStoreError::from(error))),
}
}
RetrieveQuery::TopNIds { req } => {
match req.try_map_filter(F::from_dynamic_filter) {
Ok(req) => self
.index
.top_n_ids(req)
.await
.map(|results| {
Outcome::Documents(RetrievedDocuments::Ids(
results
.into_iter()
.map(|result| (result.score, result.id))
.collect(),
))
})
.map_err(ErrorReport::from),
Err(error) => Err(ErrorReport::from(VectorStoreError::from(error))),
}
}
};
Reply::Outcome(outcome)
}
other @ (EffectKind::Completion { .. }
| EffectKind::ToolCall { .. }
| EffectKind::Embed { .. }
| EffectKind::Memory { .. }
| EffectKind::Rerank { .. }
| EffectKind::Custom { .. }) => {
Reply::Outcome(Err(wrong_family(EffectFamily::Retrieve, &other)))
}
}
}
}