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,
|
token_counter: &AsyncTokenCounter,
|
||||||
messages: &[Message],
|
messages: &[Message],
|
||||||
) -> Vec<usize> {
|
) -> 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
|
messages
|
||||||
.iter()
|
.iter()
|
||||||
|
.filter(|m| m.is_agent_visible())
|
||||||
.map(|msg| token_counter.count_chat_tokens("", std::slice::from_ref(msg), &[]))
|
.map(|msg| token_counter.count_chat_tokens("", std::slice::from_ref(msg), &[]))
|
||||||
.collect()
|
.collect()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
use crate::conversation::message::{Message, MessageContent};
|
use crate::conversation::message::{Message, MessageContent, MessageMetadata};
|
||||||
use rmcp::model::Role;
|
use rmcp::model::Role;
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use std::collections::HashSet;
|
use std::collections::HashSet;
|
||||||
@@ -103,6 +103,25 @@ impl Conversation {
|
|||||||
self.0.clear();
|
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> {
|
fn validate(self) -> Result<Self, InvalidConversation> {
|
||||||
let (_messages, issues) = fix_messages(self.0.clone());
|
let (_messages, issues) = fix_messages(self.0.clone());
|
||||||
if !issues.is_empty() {
|
if !issues.is_empty() {
|
||||||
|
|||||||
@@ -330,7 +330,6 @@ pub trait Provider: Send + Sync {
|
|||||||
) -> Result<(Message, ProviderUsage), ProviderError>;
|
) -> Result<(Message, ProviderUsage), ProviderError>;
|
||||||
|
|
||||||
// Default implementation: use the provider's configured model
|
// Default implementation: use the provider's configured model
|
||||||
// This method filters messages to only include agent_visible ones
|
|
||||||
async fn complete(
|
async fn complete(
|
||||||
&self,
|
&self,
|
||||||
system: &str,
|
system: &str,
|
||||||
@@ -338,20 +337,11 @@ pub trait Provider: Send + Sync {
|
|||||||
tools: &[Tool],
|
tools: &[Tool],
|
||||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||||
let model_config = self.get_model_config();
|
let model_config = self.get_model_config();
|
||||||
|
self.complete_with_model(&model_config, system, messages, tools)
|
||||||
// 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)
|
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if a fast model is configured, otherwise fall back to regular model
|
// 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(
|
async fn complete_fast(
|
||||||
&self,
|
&self,
|
||||||
system: &str,
|
system: &str,
|
||||||
@@ -361,15 +351,8 @@ pub trait Provider: Send + Sync {
|
|||||||
let model_config = self.get_model_config();
|
let model_config = self.get_model_config();
|
||||||
let fast_config = model_config.use_fast_model();
|
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
|
match self
|
||||||
.complete_with_model(&fast_config, system, &agent_visible_messages, tools)
|
.complete_with_model(&fast_config, system, messages, tools)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
Ok(result) => Ok(result),
|
Ok(result) => Ok(result),
|
||||||
@@ -381,7 +364,7 @@ pub trait Provider: Send + Sync {
|
|||||||
e,
|
e,
|
||||||
model_config.model_name
|
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
|
.await
|
||||||
} else {
|
} else {
|
||||||
Err(e)
|
Err(e)
|
||||||
|
|||||||
@@ -124,6 +124,7 @@ impl BedrockProvider {
|
|||||||
.set_messages(Some(
|
.set_messages(Some(
|
||||||
messages
|
messages
|
||||||
.iter()
|
.iter()
|
||||||
|
.filter(|m| m.is_agent_visible())
|
||||||
.map(to_bedrock_message)
|
.map(to_bedrock_message)
|
||||||
.collect::<Result<_>>()?,
|
.collect::<Result<_>>()?,
|
||||||
));
|
));
|
||||||
|
|||||||
@@ -129,7 +129,7 @@ impl ClaudeCodeProvider {
|
|||||||
fn messages_to_claude_format(&self, _system: &str, messages: &[Message]) -> Result<Value> {
|
fn messages_to_claude_format(&self, _system: &str, messages: &[Message]) -> Result<Value> {
|
||||||
let mut claude_messages = Vec::new();
|
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 {
|
let role = match message.role {
|
||||||
Role::User => "user",
|
Role::User => "user",
|
||||||
Role::Assistant => "assistant",
|
Role::Assistant => "assistant",
|
||||||
|
|||||||
@@ -133,7 +133,7 @@ impl CursorAgentProvider {
|
|||||||
full_prompt.push_str("\n\n");
|
full_prompt.push_str("\n\n");
|
||||||
|
|
||||||
// Add conversation history
|
// Add conversation history
|
||||||
for message in messages {
|
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||||
let role_prefix = match message.role {
|
let role_prefix = match message.role {
|
||||||
Role::User => "Human: ",
|
Role::User => "Human: ",
|
||||||
Role::Assistant => "Assistant: ",
|
Role::Assistant => "Assistant: ",
|
||||||
|
|||||||
@@ -31,8 +31,7 @@ const DATA_FIELD: &str = "data";
|
|||||||
pub fn format_messages(messages: &[Message]) -> Vec<Value> {
|
pub fn format_messages(messages: &[Message]) -> Vec<Value> {
|
||||||
let mut anthropic_messages = Vec::new();
|
let mut anthropic_messages = Vec::new();
|
||||||
|
|
||||||
// Convert messages to Anthropic format
|
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||||
for message in messages {
|
|
||||||
let role = match message.role {
|
let role = match message.role {
|
||||||
Role::User => USER_ROLE,
|
Role::User => USER_ROLE,
|
||||||
Role::Assistant => ASSISTANT_ROLE,
|
Role::Assistant => ASSISTANT_ROLE,
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ struct DatabricksMessage {
|
|||||||
/// even though the message structure is otherwise following openai, the enum switches this
|
/// even though the message structure is otherwise following openai, the enum switches this
|
||||||
fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<DatabricksMessage> {
|
fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<DatabricksMessage> {
|
||||||
let mut result = Vec::new();
|
let mut result = Vec::new();
|
||||||
for message in messages {
|
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||||
let mut converted = DatabricksMessage {
|
let mut converted = DatabricksMessage {
|
||||||
content: Value::Null,
|
content: Value::Null,
|
||||||
role: match message.role {
|
role: match message.role {
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ use std::ops::Deref;
|
|||||||
pub fn format_messages(messages: &[Message]) -> Vec<Value> {
|
pub fn format_messages(messages: &[Message]) -> Vec<Value> {
|
||||||
messages
|
messages
|
||||||
.iter()
|
.iter()
|
||||||
|
.filter(|m| m.is_agent_visible())
|
||||||
.filter(|message| {
|
.filter(|message| {
|
||||||
message
|
message
|
||||||
.content
|
.content
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ struct StreamingChunk {
|
|||||||
/// even though the message structure is otherwise following openai, the enum switches this
|
/// even though the message structure is otherwise following openai, the enum switches this
|
||||||
pub fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<Value> {
|
pub fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<Value> {
|
||||||
let mut messages_spec = Vec::new();
|
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!({
|
let mut converted = json!({
|
||||||
"role": message.role
|
"role": message.role
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ pub fn format_messages(messages: &[Message]) -> Vec<Value> {
|
|||||||
let mut snowflake_messages = Vec::new();
|
let mut snowflake_messages = Vec::new();
|
||||||
|
|
||||||
// Convert messages to Snowflake format
|
// 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 {
|
let role = match message.role {
|
||||||
Role::User => "user",
|
Role::User => "user",
|
||||||
Role::Assistant => "assistant",
|
Role::Assistant => "assistant",
|
||||||
|
|||||||
@@ -140,7 +140,7 @@ impl GeminiCliProvider {
|
|||||||
full_prompt.push_str("\n\n");
|
full_prompt.push_str("\n\n");
|
||||||
|
|
||||||
// Add conversation history
|
// Add conversation history
|
||||||
for message in messages {
|
for message in messages.iter().filter(|m| m.is_agent_visible()) {
|
||||||
let role_prefix = match message.role {
|
let role_prefix = match message.role {
|
||||||
Role::User => "Human: ",
|
Role::User => "Human: ",
|
||||||
Role::Assistant => "Assistant: ",
|
Role::Assistant => "Assistant: ",
|
||||||
|
|||||||
Reference in New Issue
Block a user