use super::prompt_request::PromptRequest;
use crate::{
completion::{
Chat, Completion, CompletionError, CompletionModel, CompletionRequestBuilder, Document,
Message, Prompt, PromptError,
},
streaming::{StreamingChat, StreamingCompletion, StreamingCompletionResponse, StreamingPrompt},
tool::ToolSet,
vector_store::VectorStoreError,
};
use futures::{stream, StreamExt, TryStreamExt};
use std::collections::HashMap;
pub struct Agent<M: CompletionModel> {
pub model: M,
pub preamble: String,
pub static_context: Vec<Document>,
pub static_tools: Vec<String>,
pub temperature: Option<f64>,
pub max_tokens: Option<u64>,
pub additional_params: Option<serde_json::Value>,
pub dynamic_context: Vec<(usize, Box<dyn crate::vector_store::VectorStoreIndexDyn>)>,
pub dynamic_tools: Vec<(usize, Box<dyn crate::vector_store::VectorStoreIndexDyn>)>,
pub tools: ToolSet,
}
impl<M: CompletionModel> Completion<M> for Agent<M> {
async fn completion(
&self,
prompt: impl Into<Message> + Send,
chat_history: Vec<Message>,
) -> Result<CompletionRequestBuilder<M>, CompletionError> {
let prompt = prompt.into();
let rag_text = prompt.rag_text();
let rag_text = rag_text.or_else(|| {
chat_history
.iter()
.rev()
.find_map(|message| message.rag_text())
});
let completion_request = self
.model
.completion_request(prompt)
.preamble(self.preamble.clone())
.messages(chat_history)
.temperature_opt(self.temperature)
.max_tokens_opt(self.max_tokens)
.additional_params_opt(self.additional_params.clone())
.documents(self.static_context.clone());
let agent = match &rag_text {
Some(text) => {
let dynamic_context = stream::iter(self.dynamic_context.iter())
.then(|(num_sample, index)| async {
Ok::<_, VectorStoreError>(
index
.top_n(text, *num_sample)
.await?
.into_iter()
.map(|(_, id, doc)| {
let text = serde_json::to_string_pretty(&doc)
.unwrap_or_else(|_| doc.to_string());
Document {
id,
text,
additional_props: HashMap::new(),
}
})
.collect::<Vec<_>>(),
)
})
.try_fold(vec![], |mut acc, docs| async {
acc.extend(docs);
Ok(acc)
})
.await
.map_err(|e| CompletionError::RequestError(Box::new(e)))?;
let dynamic_tools = stream::iter(self.dynamic_tools.iter())
.then(|(num_sample, index)| async {
Ok::<_, VectorStoreError>(
index
.top_n_ids(text, *num_sample)
.await?
.into_iter()
.map(|(_, id)| id)
.collect::<Vec<_>>(),
)
})
.try_fold(vec![], |mut acc, docs| async {
for doc in docs {
if let Some(tool) = self.tools.get(&doc) {
acc.push(tool.definition(text.into()).await)
} else {
tracing::warn!("Tool implementation not found in toolset: {}", doc);
}
}
Ok(acc)
})
.await
.map_err(|e| CompletionError::RequestError(Box::new(e)))?;
let static_tools = stream::iter(self.static_tools.iter())
.filter_map(|toolname| async move {
if let Some(tool) = self.tools.get(toolname) {
Some(tool.definition(text.into()).await)
} else {
tracing::warn!(
"Tool implementation not found in toolset: {}",
toolname
);
None
}
})
.collect::<Vec<_>>()
.await;
completion_request
.documents(dynamic_context)
.tools([static_tools.clone(), dynamic_tools].concat())
}
None => {
let static_tools = stream::iter(self.static_tools.iter())
.filter_map(|toolname| async move {
if let Some(tool) = self.tools.get(toolname) {
Some(tool.definition("".into()).await)
} else {
tracing::warn!(
"Tool implementation not found in toolset: {}",
toolname
);
None
}
})
.collect::<Vec<_>>()
.await;
completion_request.tools(static_tools)
}
};
Ok(agent)
}
}
#[allow(refining_impl_trait)]
impl<M: CompletionModel> Prompt for Agent<M> {
fn prompt(&self, prompt: impl Into<Message> + Send) -> PromptRequest<M> {
PromptRequest::new(self, prompt)
}
}
#[allow(refining_impl_trait)]
impl<M: CompletionModel> Prompt for &Agent<M> {
fn prompt(&self, prompt: impl Into<Message> + Send) -> PromptRequest<M> {
PromptRequest::new(*self, prompt)
}
}
#[allow(refining_impl_trait)]
impl<M: CompletionModel> Chat for Agent<M> {
async fn chat(
&self,
prompt: impl Into<Message> + Send,
chat_history: Vec<Message>,
) -> Result<String, PromptError> {
let mut cloned_history = chat_history.clone();
PromptRequest::new(self, prompt)
.with_history(&mut cloned_history)
.await
}
}
impl<M: CompletionModel> StreamingCompletion<M> for Agent<M> {
async fn stream_completion(
&self,
prompt: impl Into<Message> + Send,
chat_history: Vec<Message>,
) -> Result<CompletionRequestBuilder<M>, CompletionError> {
self.completion(prompt, chat_history).await
}
}
impl<M: CompletionModel> StreamingPrompt<M::StreamingResponse> for Agent<M> {
async fn stream_prompt(
&self,
prompt: impl Into<Message> + Send,
) -> Result<StreamingCompletionResponse<M::StreamingResponse>, CompletionError> {
self.stream_chat(prompt, vec![]).await
}
}
impl<M: CompletionModel> StreamingChat<M::StreamingResponse> for Agent<M> {
async fn stream_chat(
&self,
prompt: impl Into<Message> + Send,
chat_history: Vec<Message>,
) -> Result<StreamingCompletionResponse<M::StreamingResponse>, CompletionError> {
self.stream_completion(prompt, chat_history)
.await?
.stream()
.await
}
}