use crate::agent::Agent;
use crate::context::HookContext;
use crate::error::CompletionError;
use crate::hook::ToolCallDecision;
use crate::message::Message;
use crate::provider::{CompletionModel, CompletionRequest, CompletionResponse};
use crate::telemetry::Metrics;
use std::time::Instant;
pub struct AgentCompletion<'a, M: CompletionModel> {
agent: &'a Agent<M>,
message: Message,
history: Vec<Message>,
temperature: Option<f64>,
max_tokens: Option<u64>,
additional_params: Option<serde_json::Value>,
}
impl<'a, M: CompletionModel> AgentCompletion<'a, M> {
pub fn new(agent: &'a Agent<M>, message: Message) -> Self {
Self {
agent,
message,
history: Vec::new(),
temperature: None,
max_tokens: None,
additional_params: None,
}
}
pub fn history(mut self, history: &[Message]) -> Self {
self.history = history.to_vec();
self
}
pub fn add_history(mut self, message: Message) -> Self {
self.history.push(message);
self
}
pub fn temperature(mut self, temp: f64) -> Self {
self.temperature = Some(temp);
self
}
pub fn max_tokens(mut self, max: u64) -> Self {
self.max_tokens = Some(max);
self
}
pub fn additional_params(mut self, params: serde_json::Value) -> Self {
self.additional_params = Some(params);
self
}
pub async fn send(self) -> Result<CompletionResponse<M::Response>, CompletionError> {
let request_start = Instant::now();
let metrics = Metrics::global();
let mut ctx = HookContext::new_with_uuid();
let message = self
.agent
.hook_chain
.execute_pre_completion(self.message, &self.history, &mut ctx)
.await?;
let tools = if self.agent.has_tools() {
let defs = self.agent.tool_definitions(&message.text()).await;
self.agent
.hook_chain
.execute_filter_tools(defs, &mut ctx)
.await?
} else {
vec![]
};
let mut messages = self.history.clone();
messages.push(message);
if let Some(optimizer) = &self.agent.optimizer {
messages = optimizer
.optimize(messages, &self.agent.optimization_config)
.await;
}
let request = CompletionRequest {
preamble: self.agent.preamble.clone(),
messages: messages.clone(),
tools: tools.clone(),
temperature: self.temperature,
max_tokens: self.max_tokens,
additional_params: self.additional_params.clone(),
};
let mut response = match self.agent.model.completion(request).await {
Ok(response) => {
metrics.record_completion_request(
self.agent.model.provider(),
self.agent.model.model_id(),
true,
);
response
}
Err(e) => {
metrics.record_completion_request(
self.agent.model.provider(),
self.agent.model.model_id(),
false,
);
return Err(e);
}
};
self.agent
.hook_chain
.execute_on_assistant_message(&response.message, &mut ctx)
.await?;
while response.has_tool_calls() {
messages.push(response.message.clone());
for tool_call in response.tool_calls() {
let tool_name = &tool_call.function.name;
let tool_args_str = &tool_call.function.arguments;
let args: serde_json::Value =
serde_json::from_str(tool_args_str).unwrap_or_else(|_| serde_json::json!({}));
let decision = self
.agent
.hook_chain
.execute_pre_tool_call(tool_name, args, &mut ctx)
.await?;
let tool_result = match decision {
ToolCallDecision::Block(reason) => {
format!("Tool call blocked: {}", reason)
}
ToolCallDecision::Proceed(modified_args) => {
let result = if let Some(tool) = self.agent.find_tool(tool_name) {
let args_str = serde_json::to_string(&modified_args)?;
match tool.call(args_str).await {
Ok(output) => output,
Err(e) => format!("Tool execution error: {}", e),
}
} else {
format!("Tool '{}' not found", tool_name)
};
self.agent
.hook_chain
.execute_post_tool_call(tool_name, result, &mut ctx)
.await?
}
};
let tool_result_message = Message::tool_result(&tool_call.id, tool_result);
messages.push(tool_result_message);
}
if let Some(optimizer) = &self.agent.optimizer {
messages = optimizer
.optimize(messages, &self.agent.optimization_config)
.await;
}
let request = CompletionRequest {
preamble: self.agent.preamble.clone(),
messages: messages.clone(),
tools: tools.clone(),
temperature: self.temperature,
max_tokens: self.max_tokens,
additional_params: self.additional_params.clone(),
};
response = match self.agent.model.completion(request).await {
Ok(response) => {
metrics.record_completion_request(
self.agent.model.provider(),
self.agent.model.model_id(),
true,
);
response
}
Err(e) => {
metrics.record_completion_request(
self.agent.model.provider(),
self.agent.model.model_id(),
false,
);
return Err(e);
}
};
self.agent
.hook_chain
.execute_on_assistant_message(&response.message, &mut ctx)
.await?;
}
let response = self
.agent
.hook_chain
.execute_post_completion(response, &mut ctx)
.await?;
metrics.record_request_latency(request_start.elapsed(), &[]);
Ok(response)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::agent::AgentBuilder;
use crate::error::CompletionError;
use crate::provider::Usage;
#[derive(Clone)]
struct MockModel;
impl CompletionModel for MockModel {
type Response = serde_json::Value;
async fn completion(
&self,
request: CompletionRequest,
) -> Result<CompletionResponse<Self::Response>, CompletionError> {
let last_message = request
.messages
.last()
.map(|m| m.text())
.unwrap_or_default();
Ok(CompletionResponse::new(
Message::assistant(format!("Temp={:?}: {}", request.temperature, last_message)),
Usage::new(10, 20),
serde_json::json!({}),
))
}
fn model_id(&self) -> &str {
"mock-model"
}
fn provider(&self) -> &str {
"mock"
}
}
#[tokio::test]
async fn test_completion_builder() {
let agent = AgentBuilder::new(MockModel).build();
let response = agent
.completion("Hello")
.temperature(0.7)
.max_tokens(100)
.send()
.await
.unwrap();
assert!(response.content().contains("Temp=Some(0.7)"));
}
#[tokio::test]
async fn test_completion_with_history() {
let agent = AgentBuilder::new(MockModel).build();
let response = agent
.completion("Follow up question")
.history(&[
Message::user("First message"),
Message::assistant("First response"),
])
.send()
.await
.unwrap();
assert!(response.content().contains("Follow up question"));
}
}