mirror of
https://github.com/aaif-goose/goose.git
synced 2026-07-03 14:10:03 +02:00
WIP
This commit is contained in:
@@ -224,10 +224,7 @@ impl Agent {
|
||||
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);
|
||||
}
|
||||
filtered_message = filtered_message.with_id(response.id.clone());
|
||||
|
||||
// Categorize tool requests
|
||||
let mut frontend_requests = Vec::new();
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
use crate::conversation::tool_result_serde;
|
||||
use crate::mcp_utils::ToolResult;
|
||||
use crate::utils::sanitize_unicode_tags;
|
||||
use chrono::Utc;
|
||||
use rmcp::model::{
|
||||
AnnotateAble, CallToolRequestParam, Content, ImageContent, JsonObject, PromptMessage,
|
||||
@@ -9,9 +11,7 @@ use serde::{Deserialize, Deserializer, Serialize};
|
||||
use std::collections::HashSet;
|
||||
use std::fmt;
|
||||
use utoipa::ToSchema;
|
||||
|
||||
use crate::conversation::tool_result_serde;
|
||||
use crate::utils::sanitize_unicode_tags;
|
||||
use uuid::Uuid;
|
||||
|
||||
#[derive(ToSchema)]
|
||||
pub enum ToolCallResult<T> {
|
||||
@@ -450,10 +450,9 @@ impl MessageMetadata {
|
||||
}
|
||||
|
||||
#[derive(ToSchema, Clone, PartialEq, Serialize, Deserialize, Debug)]
|
||||
/// A message to or from an LLM
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct Message {
|
||||
pub id: Option<String>,
|
||||
pub id: String,
|
||||
pub role: Role,
|
||||
pub created: i64,
|
||||
#[serde(deserialize_with = "deserialize_sanitized_content")]
|
||||
@@ -462,15 +461,20 @@ pub struct Message {
|
||||
}
|
||||
|
||||
impl Message {
|
||||
pub fn msg_id() -> String {
|
||||
format!("msg_{}", Uuid::new_v4())
|
||||
}
|
||||
|
||||
pub fn new(role: Role, created: i64, content: Vec<MessageContent>) -> Self {
|
||||
Message {
|
||||
id: None,
|
||||
id: Message::msg_id(),
|
||||
role,
|
||||
created,
|
||||
content,
|
||||
metadata: MessageMetadata::default(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn debug(&self) -> String {
|
||||
format!("{:?}", self)
|
||||
}
|
||||
@@ -478,7 +482,7 @@ impl Message {
|
||||
/// Create a new user message with the current timestamp
|
||||
pub fn user() -> Self {
|
||||
Message {
|
||||
id: None,
|
||||
id: Message::msg_id(),
|
||||
role: Role::User,
|
||||
created: Utc::now().timestamp(),
|
||||
content: Vec::new(),
|
||||
@@ -489,7 +493,7 @@ impl Message {
|
||||
/// Create a new assistant message with the current timestamp
|
||||
pub fn assistant() -> Self {
|
||||
Message {
|
||||
id: None,
|
||||
id: Message::msg_id(),
|
||||
role: Role::Assistant,
|
||||
created: Utc::now().timestamp(),
|
||||
content: Vec::new(),
|
||||
@@ -498,7 +502,7 @@ impl Message {
|
||||
}
|
||||
|
||||
pub fn with_id<S: Into<String>>(mut self, id: S) -> Self {
|
||||
self.id = Some(id.into());
|
||||
self.id = id.into();
|
||||
self
|
||||
}
|
||||
|
||||
|
||||
@@ -42,11 +42,7 @@ impl Conversation {
|
||||
}
|
||||
|
||||
pub fn push(&mut self, message: Message) {
|
||||
if let Some(last) = self
|
||||
.0
|
||||
.last_mut()
|
||||
.filter(|m| m.id.is_some() && m.id == message.id)
|
||||
{
|
||||
if let Some(last) = self.0.last_mut().filter(|m| m.id == message.id) {
|
||||
match (last.content.last_mut(), message.content.last()) {
|
||||
(Some(MessageContent::Text(ref mut last)), Some(MessageContent::Text(new)))
|
||||
if message.content.len() == 1 =>
|
||||
|
||||
@@ -449,17 +449,16 @@ pub fn create_request(
|
||||
Ok(payload)
|
||||
}
|
||||
|
||||
/// Process streaming response from Anthropic's API
|
||||
pub fn response_to_streaming_message<S>(
|
||||
mut stream: S,
|
||||
) -> impl futures::Stream<
|
||||
Item = anyhow::Result<(
|
||||
Item = Result<(
|
||||
Option<Message>,
|
||||
Option<crate::providers::base::ProviderUsage>,
|
||||
)>,
|
||||
> + 'static
|
||||
where
|
||||
S: futures::Stream<Item = anyhow::Result<String>> + Unpin + Send + 'static,
|
||||
S: futures::Stream<Item = Result<String>> + Unpin + Send + 'static,
|
||||
{
|
||||
use async_stream::try_stream;
|
||||
use futures::StreamExt;
|
||||
@@ -483,19 +482,16 @@ where
|
||||
while let Some(line_result) = stream.next().await {
|
||||
let line = line_result?;
|
||||
|
||||
// Skip empty lines and non-data lines
|
||||
if line.trim().is_empty() || !line.starts_with("data: ") {
|
||||
continue;
|
||||
}
|
||||
|
||||
let data_part = line.strip_prefix("data: ").unwrap_or(&line);
|
||||
|
||||
// Handle end of stream
|
||||
if data_part.trim() == "[DONE]" {
|
||||
break;
|
||||
}
|
||||
|
||||
// Parse the JSON event
|
||||
let event: StreamingEvent = match serde_json::from_str(data_part) {
|
||||
Ok(event) => event,
|
||||
Err(e) => {
|
||||
@@ -555,7 +551,9 @@ where
|
||||
chrono::Utc::now().timestamp(),
|
||||
vec![MessageContent::text(text)],
|
||||
);
|
||||
message.id = message_id.clone();
|
||||
if let Some(msg_id) = message_id.clone() {
|
||||
message = message.with_id(msg_id);
|
||||
}
|
||||
yield (Some(message), None);
|
||||
}
|
||||
} else if delta.get("type") == Some(&json!("input_json_delta")) {
|
||||
@@ -593,7 +591,9 @@ where
|
||||
chrono::Utc::now().timestamp(),
|
||||
vec![MessageContent::tool_request(tool_id, Err(error))],
|
||||
);
|
||||
message.id = message_id.clone();
|
||||
if let Some(msg_id) = message_id.clone() {
|
||||
message = message.with_id(msg_id);
|
||||
}
|
||||
yield (Some(message), None);
|
||||
continue;
|
||||
}
|
||||
@@ -607,7 +607,9 @@ where
|
||||
chrono::Utc::now().timestamp(),
|
||||
vec![MessageContent::tool_request(tool_id, Ok(tool_call))],
|
||||
);
|
||||
message.id = message_id.clone();
|
||||
if let Some(msg_id) = message_id.clone() {
|
||||
message = message.with_id(msg_id);
|
||||
}
|
||||
yield (Some(message), None);
|
||||
}
|
||||
}
|
||||
@@ -678,9 +680,8 @@ where
|
||||
}
|
||||
}
|
||||
|
||||
// Yield final usage information if available
|
||||
if let Some(usage) = final_usage {
|
||||
yield (None, Some(usage));
|
||||
if let Some(usage) = &final_usage {
|
||||
yield (None, Some(usage.clone()));
|
||||
} else {
|
||||
tracing::debug!("🔍 Anthropic no final usage to yield");
|
||||
}
|
||||
|
||||
@@ -18,7 +18,7 @@ use tokio::sync::OnceCell;
|
||||
use tracing::{info, warn};
|
||||
use utoipa::ToSchema;
|
||||
|
||||
const CURRENT_SCHEMA_VERSION: i32 = 3;
|
||||
const CURRENT_SCHEMA_VERSION: i32 = 4;
|
||||
|
||||
static SESSION_STORAGE: OnceCell<Arc<SessionStorage>> = OnceCell::const_new();
|
||||
|
||||
@@ -620,6 +620,18 @@ impl SessionStorage {
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
}
|
||||
4 => {
|
||||
sqlx::query(
|
||||
r#"
|
||||
ALTER TABLE messages ADD COLUMN msg_id TEXT
|
||||
"#,
|
||||
)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
sqlx::query("CREATE INDEX idx_messages_msg_id ON messages(msg_id)")
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
}
|
||||
_ => {
|
||||
anyhow::bail!("Unknown migration version: {}", version);
|
||||
}
|
||||
@@ -784,15 +796,15 @@ impl SessionStorage {
|
||||
}
|
||||
|
||||
async fn get_conversation(&self, session_id: &str) -> Result<Conversation> {
|
||||
let rows = sqlx::query_as::<_, (String, String, i64, Option<String>)>(
|
||||
"SELECT role, content_json, created_timestamp, metadata_json FROM messages WHERE session_id = ? ORDER BY timestamp",
|
||||
let rows = sqlx::query_as::<_, (String, String, i64, Option<String>, Option<String>)>(
|
||||
"SELECT role, content_json, created_timestamp, metadata_json, msg_id FROM messages WHERE session_id = ? ORDER BY timestamp",
|
||||
)
|
||||
.bind(session_id)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
|
||||
let mut messages = Vec::new();
|
||||
for (role_str, content_json, created_timestamp, metadata_json) in rows {
|
||||
for (role_str, content_json, created_timestamp, metadata_json, msg_id) in rows {
|
||||
let role = match role_str.as_str() {
|
||||
"user" => Role::User,
|
||||
"assistant" => Role::Assistant,
|
||||
@@ -806,6 +818,9 @@ impl SessionStorage {
|
||||
|
||||
let mut message = Message::new(role, created_timestamp, content);
|
||||
message.metadata = metadata;
|
||||
if let Some(msg_id) = msg_id {
|
||||
message = message.with_id(msg_id);
|
||||
}
|
||||
messages.push(message);
|
||||
}
|
||||
|
||||
@@ -817,8 +832,8 @@ impl SessionStorage {
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO messages (session_id, role, content_json, created_timestamp, metadata_json)
|
||||
VALUES (?, ?, ?, ?, ?)
|
||||
INSERT INTO messages (session_id, role, content_json, created_timestamp, metadata_json, msg_id)
|
||||
VALUES (?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(session_id)
|
||||
@@ -826,6 +841,7 @@ impl SessionStorage {
|
||||
.bind(serde_json::to_string(&message.content)?)
|
||||
.bind(message.created)
|
||||
.bind(metadata_json)
|
||||
.bind(&message.id)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
|
||||
@@ -999,7 +1015,7 @@ mod tests {
|
||||
.add_message(
|
||||
&session.id,
|
||||
&Message {
|
||||
id: None,
|
||||
id: Message::msg_id(),
|
||||
role: Role::User,
|
||||
created: chrono::Utc::now().timestamp_millis(),
|
||||
content: vec![MessageContent::text("hello world")],
|
||||
@@ -1013,7 +1029,7 @@ mod tests {
|
||||
.add_message(
|
||||
&session.id,
|
||||
&Message {
|
||||
id: None,
|
||||
id: Message::msg_id(),
|
||||
role: Role::Assistant,
|
||||
created: chrono::Utc::now().timestamp_millis(),
|
||||
content: vec![MessageContent::text("sup world?")],
|
||||
@@ -1102,7 +1118,7 @@ mod tests {
|
||||
.add_message(
|
||||
&original.id,
|
||||
&Message {
|
||||
id: None,
|
||||
id: Message::msg_id(),
|
||||
role: Role::User,
|
||||
created: chrono::Utc::now().timestamp_millis(),
|
||||
content: vec![MessageContent::text(USER_MESSAGE)],
|
||||
@@ -1116,7 +1132,7 @@ mod tests {
|
||||
.add_message(
|
||||
&original.id,
|
||||
&Message {
|
||||
id: None,
|
||||
id: Message::msg_id(),
|
||||
role: Role::Assistant,
|
||||
created: chrono::Utc::now().timestamp_millis(),
|
||||
content: vec![MessageContent::text(ASSISTANT_MESSAGE)],
|
||||
|
||||
@@ -3128,8 +3128,8 @@
|
||||
},
|
||||
"Message": {
|
||||
"type": "object",
|
||||
"description": "A message to or from an LLM",
|
||||
"required": [
|
||||
"id",
|
||||
"role",
|
||||
"created",
|
||||
"content",
|
||||
@@ -3147,8 +3147,7 @@
|
||||
"format": "int64"
|
||||
},
|
||||
"id": {
|
||||
"type": "string",
|
||||
"nullable": true
|
||||
"type": "string"
|
||||
},
|
||||
"metadata": {
|
||||
"$ref": "#/components/schemas/MessageMetadata"
|
||||
|
||||
@@ -374,13 +374,10 @@ export type LoadedProvider = {
|
||||
is_editable: boolean;
|
||||
};
|
||||
|
||||
/**
|
||||
* A message to or from an LLM
|
||||
*/
|
||||
export type Message = {
|
||||
content: Array<MessageContent>;
|
||||
created: number;
|
||||
id?: string | null;
|
||||
id: string;
|
||||
metadata: MessageMetadata;
|
||||
role: Role;
|
||||
};
|
||||
|
||||
Reference in New Issue
Block a user