This commit is contained in:
pavel 2026-02-11 17:45:04 +01:00
commit beab32894b
3 changed files with 128 additions and 132 deletions

View file

@ -13,7 +13,7 @@ pub struct Agent {
url: String,
zen_api_key: Option<String>,
tavily_api_key: Option<String>,
messages: Vec<Message>,
pub messages: Vec<Message>,
tools: Option<Vec<Tool>>,
logs: String,
answer: Option<String>,
@ -50,10 +50,19 @@ impl Agent {
},
];
Self::with_messages(db, zen_api_key, tavily_api_key, messages)
}
pub fn with_messages(
db: DatabaseConnection,
zen_api_key: Option<String>,
tavily_api_key: Option<String>,
messages: Vec<Message>,
) -> AppResult<Self> {
let tools = Some(tools::get_tools());
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(60))
.timeout(std::time::Duration::from_secs(120))
.build()
.map_err(|e| AppError::Internal(format!("Failed to build HTTP client: {}", e)))?;
@ -105,6 +114,35 @@ impl Agent {
turns, current_role
));
let assistant_message = self.execute_turn().await?;
if let Some(tool_calls) = &assistant_message.tool_calls {
if tool_calls.iter().any(|tc| tc.function.name == "answer") {
finished = true;
}
}
if self.answer.is_some() {
finished = true;
}
}
self.log("\n--- Execution Finished ---");
Ok((self.logs.clone(), self.answer.clone()))
}
pub async fn execute_turn(&mut self) -> AppResult<Message> {
let max_sub_turns = 10;
let mut sub_turns = 0;
loop {
sub_turns += 1;
if sub_turns > max_sub_turns {
return Err(AppError::Internal(
"Interaction cycle turn limit exceeded".into(),
));
}
let chat_response = self.call_llm().await?;
let assistant_message = chat_response
.choices
@ -121,20 +159,24 @@ impl Agent {
}
}
if let Some(tool_calls) = assistant_message.tool_calls {
if let Some(tool_calls) = &assistant_message.tool_calls {
let mut is_final_cycle = false;
let mut final_answer = None;
for tool_call in tool_calls {
self.log(&format!("Calling tool: {}", tool_call.function.name));
let (tool_message, is_final, tool_answer) =
tools::handle_tool_call(&tool_call, &self.tavily_api_key, &self.db)
tools::handle_tool_call(tool_call, &self.tavily_api_key, &self.db)
.await
.map_err(|e| {
AppError::Internal(format!("Tool execution failed: {}", e))
})?;
if let Some(ans) = tool_answer {
self.answer = Some(ans);
self.log("Task marked as finished by tool.");
self.answer = Some(ans.clone());
final_answer = Some(ans);
self.log("Interaction marked as finished by tool.");
}
if let Some(content) = &tool_message.content {
@ -143,20 +185,24 @@ impl Agent {
self.messages.push(tool_message);
if is_final {
finished = true;
is_final_cycle = true;
}
}
} else if assistant_message.content.is_some() {
// If assistant just talked without tools, we might be stuck or finished.
// But typically we expect a 'finish' tool call.
self.log("Assistant responded without tool calls.");
// For now we continue unless the assistant explicitly uses a tool to finish,
// or we could add heuristic here if needed.
}
}
self.log("\n--- Execution Finished ---");
Ok((self.logs.clone(), self.answer.clone()))
if is_final_cycle {
return Ok(Message {
role: "assistant".to_string(),
content: final_answer.or(assistant_message.content),
tool_calls: None,
tool_call_id: None,
});
}
continue;
}
return Ok(assistant_message);
}
}
async fn call_llm(&self) -> AppResult<ChatResponse> {
@ -172,23 +218,46 @@ impl Agent {
request_builder = request_builder.header("Authorization", format!("Bearer {}", key));
}
let response = request_builder.send().await.map_err(AppError::Network)?;
let start = std::time::Instant::now();
let response = request_builder.send().await.map_err(|e| {
let duration = start.elapsed();
let is_timeout = e.is_timeout();
tracing::error!(
"Network error after {:?} during LLM call (Timeout: {}): {:?}",
duration,
is_timeout,
e
);
AppError::Network(e)
})?;
let duration = start.elapsed();
tracing::info!("LLM request completed in {:?}", duration);
if !response.status().is_success() {
let status = response.status();
let error_text = response
let body_text = response
.text()
.await
.unwrap_or_else(|_| "Unknown error".into());
return Err(AppError::Internal(format!(
"API request failed: {} - {}",
status, error_text
)));
.unwrap_or_else(|_| "Unknown body".into());
let err = format!("API request failed: {} - {}", status, body_text);
tracing::error!("{}", err);
return Err(AppError::Internal(err));
}
response
.json()
.await
.map_err(|e| AppError::Internal(format!("Failed to parse LLM response: {}", e)))
let response_text = response.text().await.map_err(|e| {
let err = format!("Failed to read response text: {}", e);
tracing::error!("{}", err);
AppError::Internal(err)
})?;
serde_json::from_str(&response_text).map_err(|e| {
let err = format!(
"Failed to parse LLM response: {} | Raw Body: {}",
e, response_text
);
tracing::error!("{}", err);
AppError::Internal(err)
})
}
}