Files
tkmind_go/crates/goose/src/agents/reply_parts.rs
T

574 lines
20 KiB
Rust

use anyhow::Result;
use regex::Regex;
use std::sync::Arc;
use async_stream::try_stream;
use futures::stream::StreamExt;
use serde_json::{json, Value};
use tracing::debug;
use super::super::agents::Agent;
#[cfg(feature = "code-mode")]
use crate::agents::platform_extensions::code_execution;
use crate::conversation::message::{Message, MessageContent, ToolRequest};
use crate::conversation::Conversation;
#[cfg(test)]
use crate::providers::base::stream_from_single_message;
use crate::providers::base::{MessageStream, Provider, ProviderUsage};
use crate::providers::errors::ProviderError;
use crate::providers::toolshim::{
augment_message_with_tool_calls, convert_tool_messages_to_text,
modify_system_prompt_for_tool_json, OllamaInterpreter,
};
use rmcp::model::Tool;
async fn enhance_model_error(error: ProviderError, provider: &Arc<dyn Provider>) -> ProviderError {
let ProviderError::RequestFailed(ref msg) = error else {
return error;
};
let re = Regex::new(r"(?i)\b4\d{2}\b.*model|model.*\b4\d{2}\b").unwrap();
if !re.is_match(msg) {
return error;
}
let Ok(models) = provider.fetch_recommended_models().await else {
return error;
};
if models.is_empty() {
return error;
}
ProviderError::RequestFailed(format!(
"{}. Available models for this provider: {}",
msg,
models.join(", ")
))
}
fn coerce_value(s: &str, schema: &Value) -> Value {
let type_str = schema.get("type");
match type_str {
Some(Value::String(t)) => match t.as_str() {
"number" | "integer" => try_coerce_number(s),
"boolean" => try_coerce_boolean(s),
_ => Value::String(s.to_string()),
},
Some(Value::Array(types)) => {
// Try each type in order
for t in types {
if let Value::String(type_name) = t {
match type_name.as_str() {
"number" | "integer" if s.parse::<f64>().is_ok() => {
return try_coerce_number(s)
}
"boolean" if matches!(s.to_lowercase().as_str(), "true" | "false") => {
return try_coerce_boolean(s)
}
_ => continue,
}
}
}
Value::String(s.to_string())
}
_ => Value::String(s.to_string()),
}
}
fn try_coerce_number(s: &str) -> Value {
if let Ok(n) = s.parse::<f64>() {
if n.fract() == 0.0 && n >= i64::MIN as f64 && n <= i64::MAX as f64 {
json!(n as i64)
} else {
json!(n)
}
} else {
Value::String(s.to_string())
}
}
fn try_coerce_boolean(s: &str) -> Value {
match s.to_lowercase().as_str() {
"true" => json!(true),
"false" => json!(false),
_ => Value::String(s.to_string()),
}
}
fn coerce_tool_arguments(
arguments: Option<serde_json::Map<String, Value>>,
tool_schema: &Value,
) -> Option<serde_json::Map<String, Value>> {
let args = arguments?;
let properties = tool_schema.get("properties").and_then(|p| p.as_object())?;
let mut coerced = serde_json::Map::new();
for (key, value) in args.iter() {
let coerced_value =
if let (Value::String(s), Some(prop_schema)) = (value, properties.get(key)) {
coerce_value(s, prop_schema)
} else {
value.clone()
};
coerced.insert(key.clone(), coerced_value);
}
Some(coerced)
}
async fn toolshim_postprocess(
response: Message,
toolshim_tools: &[Tool],
) -> Result<Message, ProviderError> {
let interpreter = OllamaInterpreter::new().map_err(|e| {
ProviderError::ExecutionError(format!("Failed to create OllamaInterpreter: {}", e))
})?;
augment_message_with_tool_calls(&interpreter, response, toolshim_tools)
.await
.map_err(|e| ProviderError::ExecutionError(format!("Failed to augment message: {}", e)))
}
impl Agent {
pub async fn prepare_tools_and_prompt(
&self,
session_id: &str,
working_dir: &std::path::Path,
) -> Result<(Vec<Tool>, Vec<Tool>, String)> {
// Get tools from extension manager
let mut tools = self.list_tools(session_id, None).await;
// Add frontend tools
let frontend_tools = self.frontend_tools.lock().await;
for frontend_tool in frontend_tools.values() {
tools.push(frontend_tool.tool.clone());
}
#[cfg(feature = "code-mode")]
let code_execution_active = self
.extension_manager
.is_extension_enabled(code_execution::EXTENSION_NAME)
.await;
#[cfg(not(feature = "code-mode"))]
let code_execution_active = false;
if code_execution_active {
tools.retain(|tool| {
if let Some(owner) = crate::agents::extension_manager::get_tool_owner(tool) {
crate::agents::extension_manager::is_first_class_extension(&owner)
} else {
false
}
});
}
// Stable tool ordering is important for multi session prompt caching.
tools.sort_by(|a, b| a.name.cmp(&b.name));
// Prepare system prompt
let extensions_info = self
.extension_manager
.get_extensions_info(working_dir)
.await;
let (extension_count, tool_count) = self
.extension_manager
.get_extension_and_tool_counts(session_id)
.await;
// Get model name from provider
let provider = self.provider().await?;
let model_config = provider.get_model_config();
let prompt_manager = self.prompt_manager.lock().await;
let mut system_prompt = prompt_manager
.builder()
.with_extensions(extensions_info.into_iter())
.with_frontend_instructions(self.frontend_instructions.lock().await.clone())
.with_extension_and_tool_counts(extension_count, tool_count)
.with_code_execution_mode(code_execution_active)
.with_hints(working_dir)
.build();
// Handle toolshim if enabled
let mut toolshim_tools = vec![];
if model_config.toolshim {
// If tool interpretation is enabled, modify the system prompt
system_prompt = modify_system_prompt_for_tool_json(&system_prompt, &tools);
// Make a copy of tools before emptying
toolshim_tools = tools.clone();
// Empty the tools vector for provider completion
tools = vec![];
}
Ok((tools, toolshim_tools, system_prompt))
}
/// Stream a response from the LLM provider.
/// Handles toolshim transformations if needed
pub(crate) async fn stream_response_from_provider(
provider: Arc<dyn Provider>,
session_id: &str,
system_prompt: &str,
messages: &[Message],
tools: &[Tool],
toolshim_tools: &[Tool],
) -> Result<MessageStream, ProviderError> {
let config = provider.get_model_config();
let filtered_messages: Vec<Message> = messages
.iter()
.filter(|m| m.is_agent_visible())
.map(|m| m.agent_visible_content())
.collect();
// Convert tool messages to text if toolshim is enabled
let messages_for_provider = if config.toolshim {
convert_tool_messages_to_text(&filtered_messages)
} else {
Conversation::new_unvalidated(filtered_messages)
};
// Clone owned data to move into the async stream
let system_prompt = system_prompt.to_owned();
let tools = tools.to_owned();
let toolshim_tools = toolshim_tools.to_owned();
let provider = provider.clone();
// Capture errors during stream creation and return them as part of the stream
// so they can be handled by the existing error handling logic in the agent
let model_config = provider.get_model_config();
debug!("WAITING_LLM_STREAM_START");
let stream_result = provider
.stream(
&model_config,
session_id,
system_prompt.as_str(),
messages_for_provider.messages(),
&tools,
)
.await;
debug!("WAITING_LLM_STREAM_END");
// If there was an error creating the stream, return a stream that yields that error
let mut stream = match stream_result {
Ok(s) => s,
Err(e) => {
let enhanced_error = enhance_model_error(e, &provider).await;
// Return a stream that immediately yields the error
// This allows the error to be caught by existing error handling in agent.rs
return Ok(Box::pin(try_stream! {
yield Err(enhanced_error)?;
}));
}
};
Ok(Box::pin(try_stream! {
while let Some(result) = stream.next().await {
let (mut message, usage) = result?;
// Store the model information in the global store
if let Some(usage) = usage.as_ref() {
crate::providers::base::set_current_model(&usage.model);
}
// Post-process / structure the response only if tool interpretation is enabled
if message.is_some() && config.toolshim {
message = Some(toolshim_postprocess(message.unwrap(), &toolshim_tools).await?);
}
yield (message, usage);
}
}))
}
/// Categorize tool requests from the response into different types
/// Returns:
/// - frontend_requests: Tool requests that should be handled by the frontend
/// - other_requests: All other tool requests (including requests to enable extensions)
/// - filtered_message: The original message with frontend tool requests removed
pub(crate) async fn categorize_tool_requests(
&self,
response: &Message,
tools: &[Tool],
) -> (Vec<ToolRequest>, Vec<ToolRequest>, Message) {
// First collect all tool requests with coercion applied
let tool_requests: Vec<ToolRequest> = response
.content
.iter()
.filter_map(|content| {
if let MessageContent::ToolRequest(req) = content {
let mut coerced_req = req.clone();
if let Ok(ref mut tool_call) = coerced_req.tool_call {
if let Some(tool) = tools.iter().find(|t| t.name == tool_call.name) {
let schema_value = Value::Object(tool.input_schema.as_ref().clone());
tool_call.arguments =
coerce_tool_arguments(tool_call.arguments.clone(), &schema_value);
if let Some(ref meta) = tool.meta {
coerced_req.tool_meta = serde_json::to_value(meta).ok();
}
}
}
Some(coerced_req)
} else {
None
}
})
.collect();
// Create a filtered message with frontend tool requests removed
let mut filtered_content = Vec::new();
let mut tool_request_index = 0;
for content in &response.content {
match content {
MessageContent::ToolRequest(_) => {
if tool_request_index < tool_requests.len() {
let coerced_req = &tool_requests[tool_request_index];
tool_request_index += 1;
let should_include = if let Ok(tool_call) = &coerced_req.tool_call {
!self.is_frontend_tool(&tool_call.name).await
} else {
true
};
if should_include {
filtered_content.push(MessageContent::ToolRequest(coerced_req.clone()));
}
}
}
_ => {
filtered_content.push(content.clone());
}
}
}
let mut filtered_message =
Message::new(response.role.clone(), response.created, filtered_content);
// Preserve the ID if it exists
if let Some(id) = response.id.clone() {
filtered_message = filtered_message.with_id(id);
}
// Categorize tool requests
let mut frontend_requests = Vec::new();
let mut other_requests = Vec::new();
for request in tool_requests {
if let Ok(tool_call) = &request.tool_call {
if self.is_frontend_tool(&tool_call.name).await {
frontend_requests.push(request);
} else {
other_requests.push(request);
}
} else {
// If there's an error in the tool call, add it to other_requests
other_requests.push(request);
}
}
(frontend_requests, other_requests, filtered_message)
}
pub(crate) async fn update_session_metrics(
&self,
session_id: &str,
schedule_id: Option<String>,
usage: &ProviderUsage,
is_compaction_usage: bool,
) -> Result<()> {
let manager = self.config.session_manager.clone();
let session = manager.get_session(session_id, false).await?;
let accumulate = |a: Option<i32>, b: Option<i32>| -> Option<i32> {
match (a, b) {
(Some(x), Some(y)) => Some(x + y),
_ => a.or(b),
}
};
let accumulated_total =
accumulate(session.accumulated_total_tokens, usage.usage.total_tokens);
let accumulated_input =
accumulate(session.accumulated_input_tokens, usage.usage.input_tokens);
let accumulated_output =
accumulate(session.accumulated_output_tokens, usage.usage.output_tokens);
let (current_total, current_input, current_output) = if is_compaction_usage {
// After compaction: summary output becomes new input context
let new_input = usage.usage.output_tokens;
(new_input, new_input, None)
} else {
(
usage.usage.total_tokens,
usage.usage.input_tokens,
usage.usage.output_tokens,
)
};
manager
.update(session_id)
.schedule_id(schedule_id)
.total_tokens(current_total)
.input_tokens(current_input)
.output_tokens(current_output)
.accumulated_total_tokens(accumulated_total)
.accumulated_input_tokens(accumulated_input)
.accumulated_output_tokens(accumulated_output)
.apply()
.await?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::conversation::message::Message;
use crate::model::ModelConfig;
use crate::providers::base::{Provider, ProviderUsage, Usage};
use crate::providers::errors::ProviderError;
use crate::session::session_manager::SessionType;
use async_trait::async_trait;
use rmcp::object;
#[derive(Clone)]
struct MockProvider {
model_config: ModelConfig,
}
#[async_trait]
impl Provider for MockProvider {
fn get_name(&self) -> &str {
"mock"
}
fn get_model_config(&self) -> ModelConfig {
self.model_config.clone()
}
async fn stream(
&self,
_model_config: &ModelConfig,
_session_id: &str,
_system: &str,
_messages: &[Message],
_tools: &[Tool],
) -> Result<MessageStream, ProviderError> {
let message = Message::assistant().with_text("ok");
let usage = ProviderUsage::new("mock".to_string(), Usage::default());
Ok(stream_from_single_message(message, usage))
}
}
#[tokio::test]
async fn prepare_tools_returns_sorted_tools_including_frontend() -> anyhow::Result<()> {
let agent = crate::agents::Agent::new();
let session = agent
.config
.session_manager
.create_session(
std::env::current_dir().unwrap(),
"test-prepare-tools".to_string(),
SessionType::Hidden,
)
.await?;
let model_config = ModelConfig::new("test-model").unwrap();
let provider = std::sync::Arc::new(MockProvider { model_config });
agent.update_provider(provider, &session.id).await?;
// Add unsorted frontend tools
let frontend_tools = vec![
Tool::new(
"frontend__z_tool".to_string(),
"Z tool".to_string(),
object!({ "type": "object", "properties": { } }),
),
Tool::new(
"frontend__a_tool".to_string(),
"A tool".to_string(),
object!({ "type": "object", "properties": { } }),
),
];
agent
.add_extension(
crate::agents::extension::ExtensionConfig::Frontend {
name: "frontend".to_string(),
description: "desc".to_string(),
tools: frontend_tools,
instructions: None,
bundled: None,
available_tools: vec![],
},
&session.id,
)
.await
.unwrap();
let (tools, _toolshim_tools, _system_prompt) = agent
.prepare_tools_and_prompt(&session.id, session.working_dir.as_path())
.await?;
let names: Vec<String> = tools.iter().map(|t| t.name.clone().into_owned()).collect();
assert!(names.iter().any(|n| n == "frontend__a_tool"));
assert!(names.iter().any(|n| n == "frontend__z_tool"));
// Verify the names are sorted ascending
let mut sorted = names.clone();
sorted.sort();
assert_eq!(names, sorted);
Ok(())
}
#[tokio::test]
async fn test_stream_error_propagation() {
use futures::StreamExt;
type StreamItem = Result<(Option<Message>, Option<ProviderUsage>), ProviderError>;
let stream = futures::stream::iter(vec![
Ok((Some(Message::assistant().with_text("chunk1")), None)),
Ok((Some(Message::assistant().with_text("chunk2")), None)),
Err(ProviderError::RequestFailed(
"simulated stream error".to_string(),
)),
] as Vec<StreamItem>);
let mut pinned = Box::pin(stream);
let mut results = Vec::new();
let mut error_seen = false;
while let Some(result) = pinned.next().await {
match result {
Ok((message, _usage)) => {
if let Some(msg) = message {
results.push(msg.as_concat_text());
}
}
Err(_e) => {
error_seen = true;
break;
}
}
}
assert_eq!(results.len(), 2);
assert_eq!(results[0], "chunk1");
assert_eq!(results[1], "chunk2");
assert!(
error_seen,
"Error should have been propagated, not silently ignored"
);
}
}