chore: use a Conversation type (#3735)
This commit is contained in:
+24
-13
@@ -5,7 +5,8 @@ use std::sync::Arc;
|
||||
use anyhow::Result;
|
||||
use futures::StreamExt;
|
||||
use goose::agents::{Agent, AgentEvent};
|
||||
use goose::message::Message;
|
||||
use goose::conversation::message::Message;
|
||||
use goose::conversation::Conversation;
|
||||
use goose::model::ModelConfig;
|
||||
use goose::providers::base::Provider;
|
||||
use goose::providers::{
|
||||
@@ -118,7 +119,7 @@ async fn run_truncate_test(
|
||||
agent.update_provider(provider).await?;
|
||||
let repeat_count = context_window + 10_000;
|
||||
let large_message_content = "hello ".repeat(repeat_count);
|
||||
let messages = vec![
|
||||
let messages = Conversation::new(vec![
|
||||
Message::user().with_text("hi there. what is 2 + 2?"),
|
||||
Message::assistant().with_text("hey! I think it's 4."),
|
||||
Message::user().with_text(&large_message_content),
|
||||
@@ -128,9 +129,10 @@ async fn run_truncate_test(
|
||||
Message::user().with_text(
|
||||
"did I ask you what's 2+2 in this message history? just respond with 'yes' or 'no'",
|
||||
),
|
||||
];
|
||||
])
|
||||
.unwrap();
|
||||
|
||||
let reply_stream = agent.reply(&messages, None, None).await?;
|
||||
let reply_stream = agent.reply(messages, None, None).await?;
|
||||
tokio::pin!(reply_stream);
|
||||
|
||||
let mut responses = Vec::new();
|
||||
@@ -166,11 +168,11 @@ async fn run_truncate_test(
|
||||
assert_eq!(responses[0].content.len(), 1);
|
||||
|
||||
match responses[0].content[0] {
|
||||
goose::message::MessageContent::Text(ref text_content) => {
|
||||
goose::conversation::message::MessageContent::Text(ref text_content) => {
|
||||
assert!(text_content.text.to_lowercase().contains("no"));
|
||||
assert!(!text_content.text.to_lowercase().contains("yes"));
|
||||
}
|
||||
goose::message::MessageContent::ContextLengthExceeded(_) => {
|
||||
goose::conversation::message::MessageContent::ContextLengthExceeded(_) => {
|
||||
// This is an acceptable outcome for providers that don't truncate themselves
|
||||
// and correctly report that the context length was exceeded.
|
||||
println!(
|
||||
@@ -546,12 +548,14 @@ mod final_output_tool_tests {
|
||||
use goose::agents::final_output_tool::{
|
||||
FINAL_OUTPUT_CONTINUATION_MESSAGE, FINAL_OUTPUT_TOOL_NAME,
|
||||
};
|
||||
use goose::conversation::Conversation;
|
||||
use goose::providers::base::MessageStream;
|
||||
use goose::recipe::Response;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_final_output_assistant_message_in_reply() -> Result<()> {
|
||||
use async_trait::async_trait;
|
||||
use goose::conversation::message::Message;
|
||||
use goose::model::ModelConfig;
|
||||
use goose::providers::base::{Provider, ProviderUsage, Usage};
|
||||
use goose::providers::errors::ProviderError;
|
||||
@@ -626,7 +630,7 @@ mod final_output_tool_tests {
|
||||
);
|
||||
|
||||
// Simulate the reply stream continuing after the final output tool call.
|
||||
let reply_stream = agent.reply(&vec![], None, None).await?;
|
||||
let reply_stream = agent.reply(Conversation::empty(), None, None).await?;
|
||||
tokio::pin!(reply_stream);
|
||||
|
||||
let mut responses = Vec::new();
|
||||
@@ -652,6 +656,7 @@ mod final_output_tool_tests {
|
||||
#[tokio::test]
|
||||
async fn test_when_final_output_not_called_in_reply() -> Result<()> {
|
||||
use async_trait::async_trait;
|
||||
use goose::conversation::message::Message;
|
||||
use goose::model::ModelConfig;
|
||||
use goose::providers::base::{Provider, ProviderUsage};
|
||||
use goose::providers::errors::ProviderError;
|
||||
@@ -723,7 +728,7 @@ mod final_output_tool_tests {
|
||||
agent.add_final_output_tool(response).await;
|
||||
|
||||
// Simulate the reply stream being called.
|
||||
let reply_stream = agent.reply(&vec![], None, None).await?;
|
||||
let reply_stream = agent.reply(Conversation::empty(), None, None).await?;
|
||||
tokio::pin!(reply_stream);
|
||||
|
||||
let mut responses = Vec::new();
|
||||
@@ -773,6 +778,8 @@ mod retry_tests {
|
||||
use super::*;
|
||||
use async_trait::async_trait;
|
||||
use goose::agents::types::{RetryConfig, SessionConfig, SuccessCheck};
|
||||
use goose::conversation::message::Message;
|
||||
use goose::conversation::Conversation;
|
||||
use goose::model::ModelConfig;
|
||||
use goose::providers::base::{Provider, ProviderUsage, Usage};
|
||||
use goose::providers::errors::ProviderError;
|
||||
@@ -855,10 +862,11 @@ mod retry_tests {
|
||||
retry_config: Some(retry_config),
|
||||
};
|
||||
|
||||
let initial_messages = vec![Message::user().with_text("Complete this task")];
|
||||
let conversation =
|
||||
Conversation::new(vec![Message::user().with_text("Complete this task")]).unwrap();
|
||||
|
||||
let reply_stream = agent
|
||||
.reply(&initial_messages, Some(session_config), None)
|
||||
.reply(conversation, Some(session_config), None)
|
||||
.await?;
|
||||
tokio::pin!(reply_stream);
|
||||
|
||||
@@ -952,7 +960,8 @@ mod retry_tests {
|
||||
mod max_turns_tests {
|
||||
use super::*;
|
||||
use async_trait::async_trait;
|
||||
use goose::message::MessageContent;
|
||||
use goose::conversation::message::{Message, MessageContent};
|
||||
use goose::conversation::Conversation;
|
||||
use goose::model::ModelConfig;
|
||||
use goose::providers::base::{Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use goose::providers::errors::ProviderError;
|
||||
@@ -1021,9 +1030,11 @@ mod max_turns_tests {
|
||||
max_turns: Some(1),
|
||||
retry_config: None,
|
||||
};
|
||||
let messages = vec![Message::user().with_text("Hello")];
|
||||
let conversation = Conversation::new(vec![Message::user().with_text("Hello")]).unwrap();
|
||||
|
||||
let reply_stream = agent.reply(&messages, Some(session_config), None).await?;
|
||||
let reply_stream = agent
|
||||
.reply(conversation, Some(session_config), None)
|
||||
.await?;
|
||||
tokio::pin!(reply_stream);
|
||||
|
||||
let mut responses = Vec::new();
|
||||
|
||||
@@ -782,10 +782,10 @@ async fn test_schedule_tool_session_content_action_with_real_session() {
|
||||
|
||||
// Create test metadata and messages
|
||||
let metadata = create_test_session_metadata(2, "/tmp");
|
||||
let messages = vec![
|
||||
goose::message::Message::user().with_text("Hello"),
|
||||
goose::message::Message::assistant().with_text("Hi there!"),
|
||||
];
|
||||
let messages = goose::conversation::Conversation::new_unvalidated(vec![
|
||||
goose::conversation::message::Message::user().with_text("Hello"),
|
||||
goose::conversation::message::Message::assistant().with_text("Hi there!"),
|
||||
]);
|
||||
|
||||
// Save the session file
|
||||
goose::session::storage::save_messages_with_metadata(&session_path, &metadata, &messages)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use anyhow::Result;
|
||||
use dotenvy::dotenv;
|
||||
use goose::message::{Message, MessageContent};
|
||||
use goose::conversation::message::{Message, MessageContent};
|
||||
use goose::providers::base::Provider;
|
||||
use goose::providers::errors::ProviderError;
|
||||
use goose::providers::{
|
||||
@@ -257,6 +257,7 @@ impl ProviderTester {
|
||||
|
||||
async fn test_image_content_support(&self) -> Result<()> {
|
||||
use base64::{engine::general_purpose::STANDARD as BASE64, Engine as _};
|
||||
use goose::conversation::message::Message;
|
||||
use std::fs;
|
||||
|
||||
// Try to read the test image
|
||||
|
||||
Reference in New Issue
Block a user