mirror of
https://github.com/aaif-goose/goose.git
synced 2026-07-03 14:10:03 +02:00
Add filtering for agentVisible: false messages on streaming providers (#4847)
This commit is contained in:
@@ -43,10 +43,9 @@ pub fn get_messages_token_counts_async(
|
||||
token_counter: &AsyncTokenCounter,
|
||||
messages: &[Message],
|
||||
) -> Vec<usize> {
|
||||
// Calculate current token count of each message, use count_chat_tokens to ensure we
|
||||
// capture the full content of the message, include ToolRequests and ToolResponses
|
||||
messages
|
||||
.iter()
|
||||
.filter(|m| m.is_agent_visible())
|
||||
.map(|msg| token_counter.count_chat_tokens("", std::slice::from_ref(msg), &[]))
|
||||
.collect()
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::conversation::message::{Message, MessageContent, MessageMetadata};
|
||||
use rmcp::model::Role;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashSet;
|
||||
@@ -102,6 +102,25 @@ impl Conversation {
|
||||
self.0.clear();
|
||||
}
|
||||
|
||||
pub fn filtered_messages<F>(&self, filter: F) -> Vec<Message>
|
||||
where
|
||||
F: Fn(&MessageMetadata) -> bool,
|
||||
{
|
||||
self.0
|
||||
.iter()
|
||||
.filter(|msg| filter(&msg.metadata))
|
||||
.cloned()
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn agent_visible_messages(&self) -> Vec<Message> {
|
||||
self.filtered_messages(|meta| meta.agent_visible)
|
||||
}
|
||||
|
||||
pub fn user_visible_messages(&self) -> Vec<Message> {
|
||||
self.filtered_messages(|meta| meta.user_visible)
|
||||
}
|
||||
|
||||
fn validate(self) -> Result<Self, InvalidConversation> {
|
||||
let (_messages, issues) = fix_messages(self.0.clone());
|
||||
if !issues.is_empty() {
|
||||
|
||||
@@ -328,7 +328,6 @@ pub trait Provider: Send + Sync {
|
||||
) -> Result<(Message, ProviderUsage), ProviderError>;
|
||||
|
||||
// Default implementation: use the provider's configured model
|
||||
// This method filters messages to only include agent_visible ones
|
||||
async fn complete(
|
||||
&self,
|
||||
system: &str,
|
||||
@@ -336,20 +335,11 @@ pub trait Provider: Send + Sync {
|
||||
tools: &[Tool],
|
||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||
let model_config = self.get_model_config();
|
||||
|
||||
// Filter messages to only include agent_visible ones
|
||||
let agent_visible_messages: Vec<Message> = messages
|
||||
.iter()
|
||||
.filter(|m| m.is_agent_visible())
|
||||
.cloned()
|
||||
.collect();
|
||||
|
||||
self.complete_with_model(&model_config, system, &agent_visible_messages, tools)
|
||||
self.complete_with_model(&model_config, system, messages, tools)
|
||||
.await
|
||||
}
|
||||
|
||||
// Check if a fast model is configured, otherwise fall back to regular model
|
||||
// This method filters messages to only include agent_visible ones
|
||||
async fn complete_fast(
|
||||
&self,
|
||||
system: &str,
|
||||
@@ -359,15 +349,8 @@ pub trait Provider: Send + Sync {
|
||||
let model_config = self.get_model_config();
|
||||
let fast_config = model_config.use_fast_model();
|
||||
|
||||
// Filter messages to only include agent_visible ones
|
||||
let agent_visible_messages: Vec<Message> = messages
|
||||
.iter()
|
||||
.filter(|m| m.is_agent_visible())
|
||||
.cloned()
|
||||
.collect();
|
||||
|
||||
match self
|
||||
.complete_with_model(&fast_config, system, &agent_visible_messages, tools)
|
||||
.complete_with_model(&fast_config, system, messages, tools)
|
||||
.await
|
||||
{
|
||||
Ok(result) => Ok(result),
|
||||
@@ -379,7 +362,7 @@ pub trait Provider: Send + Sync {
|
||||
e,
|
||||
model_config.model_name
|
||||
);
|
||||
self.complete_with_model(&model_config, system, &agent_visible_messages, tools)
|
||||
self.complete_with_model(&model_config, system, messages, tools)
|
||||
.await
|
||||
} else {
|
||||
Err(e)
|
||||
|
||||
@@ -124,6 +124,7 @@ impl BedrockProvider {
|
||||
.set_messages(Some(
|
||||
messages
|
||||
.iter()
|
||||
.filter(|m| m.is_agent_visible())
|
||||
.map(to_bedrock_message)
|
||||
.collect::<Result<_>>()?,
|
||||
));
|
||||
|
||||
@@ -129,7 +129,7 @@ impl ClaudeCodeProvider {
|
||||
fn messages_to_claude_format(&self, _system: &str, messages: &[Message]) -> Result<Value> {
|
||||
let mut claude_messages = Vec::new();
|
||||
|
||||
for message in messages {
|
||||
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||
let role = match message.role {
|
||||
Role::User => "user",
|
||||
Role::Assistant => "assistant",
|
||||
|
||||
@@ -133,7 +133,7 @@ impl CursorAgentProvider {
|
||||
full_prompt.push_str("\n\n");
|
||||
|
||||
// Add conversation history
|
||||
for message in messages {
|
||||
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||
let role_prefix = match message.role {
|
||||
Role::User => "Human: ",
|
||||
Role::Assistant => "Assistant: ",
|
||||
|
||||
@@ -32,8 +32,7 @@ const DATA_FIELD: &str = "data";
|
||||
pub fn format_messages(messages: &[Message]) -> Vec<Value> {
|
||||
let mut anthropic_messages = Vec::new();
|
||||
|
||||
// Convert messages to Anthropic format
|
||||
for message in messages {
|
||||
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||
let role = match message.role {
|
||||
Role::User => USER_ROLE,
|
||||
Role::Assistant => ASSISTANT_ROLE,
|
||||
|
||||
@@ -29,7 +29,7 @@ struct DatabricksMessage {
|
||||
/// even though the message structure is otherwise following openai, the enum switches this
|
||||
fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<DatabricksMessage> {
|
||||
let mut result = Vec::new();
|
||||
for message in messages {
|
||||
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||
let mut converted = DatabricksMessage {
|
||||
content: Value::Null,
|
||||
role: match message.role {
|
||||
|
||||
@@ -16,6 +16,7 @@ use std::ops::Deref;
|
||||
pub fn format_messages(messages: &[Message]) -> Vec<Value> {
|
||||
messages
|
||||
.iter()
|
||||
.filter(|m| m.is_agent_visible())
|
||||
.filter(|message| {
|
||||
message
|
||||
.content
|
||||
|
||||
@@ -59,7 +59,7 @@ struct StreamingChunk {
|
||||
/// even though the message structure is otherwise following openai, the enum switches this
|
||||
pub fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<Value> {
|
||||
let mut messages_spec = Vec::new();
|
||||
for message in messages {
|
||||
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||
let mut converted = json!({
|
||||
"role": message.role
|
||||
});
|
||||
|
||||
@@ -13,7 +13,7 @@ pub fn format_messages(messages: &[Message]) -> Vec<Value> {
|
||||
let mut snowflake_messages = Vec::new();
|
||||
|
||||
// Convert messages to Snowflake format
|
||||
for message in messages {
|
||||
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||
let role = match message.role {
|
||||
Role::User => "user",
|
||||
Role::Assistant => "assistant",
|
||||
|
||||
@@ -140,7 +140,7 @@ impl GeminiCliProvider {
|
||||
full_prompt.push_str("\n\n");
|
||||
|
||||
// Add conversation history
|
||||
for message in messages {
|
||||
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||
let role_prefix = match message.role {
|
||||
Role::User => "Human: ",
|
||||
Role::Assistant => "Assistant: ",
|
||||
|
||||
Reference in New Issue
Block a user