Nest TODO State in session data (#4361)
Co-authored-by: Alex Hancock <alexhancock@block.xyz>
This commit is contained in:
@@ -41,6 +41,7 @@ use crate::providers::errors::ProviderError;
|
||||
use crate::recipe::{Author, Recipe, Response, Settings, SubRecipe};
|
||||
use crate::scheduler_trait::SchedulerTrait;
|
||||
use crate::session;
|
||||
use crate::session::extension_data::ExtensionState;
|
||||
use crate::tool_monitor::{ToolCall, ToolMonitor};
|
||||
use crate::utils::is_token_cancelled;
|
||||
use mcp_core::ToolResult;
|
||||
@@ -494,7 +495,10 @@ impl Agent {
|
||||
let todo_content = if let Some(path) = session_file_path {
|
||||
session::storage::read_metadata(&path)
|
||||
.ok()
|
||||
.and_then(|m| m.todo_content)
|
||||
.and_then(|m| {
|
||||
session::TodoState::from_extension_data(&m.extension_data)
|
||||
.map(|state| state.content)
|
||||
})
|
||||
.unwrap_or_default()
|
||||
} else {
|
||||
String::new()
|
||||
@@ -531,7 +535,11 @@ impl Agent {
|
||||
match session::storage::get_path(session_config.id.clone()) {
|
||||
Ok(path) => match session::storage::read_metadata(&path) {
|
||||
Ok(mut metadata) => {
|
||||
metadata.todo_content = Some(content);
|
||||
let todo_state = session::TodoState::new(content);
|
||||
todo_state
|
||||
.to_extension_data(&mut metadata.extension_data)
|
||||
.ok();
|
||||
|
||||
let path_clone = path.clone();
|
||||
let metadata_clone = metadata.clone();
|
||||
let update_result = tokio::task::spawn(async move {
|
||||
|
||||
@@ -269,7 +269,7 @@ mod tests {
|
||||
accumulated_total_tokens: Some(100),
|
||||
accumulated_input_tokens: Some(50),
|
||||
accumulated_output_tokens: Some(50),
|
||||
todo_content: None,
|
||||
extension_data: crate::session::ExtensionData::new(),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1298,7 +1298,7 @@ async fn run_scheduled_job_internal(
|
||||
accumulated_total_tokens: None,
|
||||
accumulated_input_tokens: None,
|
||||
accumulated_output_tokens: None,
|
||||
todo_content: None,
|
||||
extension_data: crate::session::ExtensionData::new(),
|
||||
};
|
||||
if let Err(e_fb) = crate::session::storage::save_messages_with_metadata(
|
||||
&session_file_path,
|
||||
|
||||
@@ -0,0 +1,173 @@
|
||||
// Extension data management for sessions
|
||||
// Provides a simple way to store extension-specific data with versioned keys
|
||||
|
||||
use anyhow::Result;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use std::collections::HashMap;
|
||||
use utoipa::ToSchema;
|
||||
|
||||
/// Extension data containing all extension states
|
||||
/// Keys are in format "extension_name.version" (e.g., "todo.v0")
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default, ToSchema)]
|
||||
pub struct ExtensionData {
|
||||
#[serde(flatten)]
|
||||
pub extension_states: HashMap<String, Value>,
|
||||
}
|
||||
|
||||
impl ExtensionData {
|
||||
/// Create a new empty ExtensionData
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
extension_states: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get extension state for a specific extension and version
|
||||
pub fn get_extension_state(&self, extension_name: &str, version: &str) -> Option<&Value> {
|
||||
let key = format!("{}.{}", extension_name, version);
|
||||
self.extension_states.get(&key)
|
||||
}
|
||||
|
||||
/// Set extension state for a specific extension and version
|
||||
pub fn set_extension_state(&mut self, extension_name: &str, version: &str, state: Value) {
|
||||
let key = format!("{}.{}", extension_name, version);
|
||||
self.extension_states.insert(key, state);
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper trait for extension-specific state management
|
||||
pub trait ExtensionState: Sized + Serialize + for<'de> Deserialize<'de> {
|
||||
/// The name of the extension
|
||||
const EXTENSION_NAME: &'static str;
|
||||
|
||||
/// The version of the extension state format
|
||||
const VERSION: &'static str;
|
||||
|
||||
/// Convert from JSON value
|
||||
fn from_value(value: &Value) -> Result<Self> {
|
||||
serde_json::from_value(value.clone()).map_err(|e| {
|
||||
anyhow::anyhow!(
|
||||
"Failed to deserialize {} state: {}",
|
||||
Self::EXTENSION_NAME,
|
||||
e
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
/// Convert to JSON value
|
||||
fn to_value(&self) -> Result<Value> {
|
||||
serde_json::to_value(self).map_err(|e| {
|
||||
anyhow::anyhow!("Failed to serialize {} state: {}", Self::EXTENSION_NAME, e)
|
||||
})
|
||||
}
|
||||
|
||||
/// Get state from extension data
|
||||
fn from_extension_data(extension_data: &ExtensionData) -> Option<Self> {
|
||||
extension_data
|
||||
.get_extension_state(Self::EXTENSION_NAME, Self::VERSION)
|
||||
.and_then(|v| Self::from_value(v).ok())
|
||||
}
|
||||
|
||||
/// Save state to extension data
|
||||
fn to_extension_data(&self, extension_data: &mut ExtensionData) -> Result<()> {
|
||||
let value = self.to_value()?;
|
||||
extension_data.set_extension_state(Self::EXTENSION_NAME, Self::VERSION, value);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// TODO extension state implementation
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TodoState {
|
||||
pub content: String,
|
||||
}
|
||||
|
||||
impl ExtensionState for TodoState {
|
||||
const EXTENSION_NAME: &'static str = "todo";
|
||||
const VERSION: &'static str = "v0";
|
||||
}
|
||||
|
||||
impl TodoState {
|
||||
/// Create a new TODO state
|
||||
pub fn new(content: String) -> Self {
|
||||
Self { content }
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn test_extension_data_basic_operations() {
|
||||
let mut extension_data = ExtensionData::new();
|
||||
|
||||
// Test setting and getting extension state
|
||||
let todo_state = json!({"content": "- Task 1\n- Task 2"});
|
||||
extension_data.set_extension_state("todo", "v0", todo_state.clone());
|
||||
|
||||
assert_eq!(
|
||||
extension_data.get_extension_state("todo", "v0"),
|
||||
Some(&todo_state)
|
||||
);
|
||||
assert_eq!(extension_data.get_extension_state("todo", "v1"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_multiple_extension_states() {
|
||||
let mut extension_data = ExtensionData::new();
|
||||
|
||||
// Add multiple extension states
|
||||
extension_data.set_extension_state("todo", "v0", json!("TODO content"));
|
||||
extension_data.set_extension_state("memory", "v1", json!({"items": ["item1", "item2"]}));
|
||||
extension_data.set_extension_state("config", "v2", json!({"setting": true}));
|
||||
|
||||
// Check all states exist
|
||||
assert_eq!(extension_data.extension_states.len(), 3);
|
||||
assert!(extension_data.get_extension_state("todo", "v0").is_some());
|
||||
assert!(extension_data.get_extension_state("memory", "v1").is_some());
|
||||
assert!(extension_data.get_extension_state("config", "v2").is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_todo_state_trait() {
|
||||
let mut extension_data = ExtensionData::new();
|
||||
|
||||
// Create and save TODO state
|
||||
let todo = TodoState::new("- Task 1\n- Task 2".to_string());
|
||||
todo.to_extension_data(&mut extension_data).unwrap();
|
||||
|
||||
// Retrieve TODO state
|
||||
let retrieved = TodoState::from_extension_data(&extension_data);
|
||||
assert!(retrieved.is_some());
|
||||
assert_eq!(retrieved.unwrap().content, "- Task 1\n- Task 2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extension_data_serialization() {
|
||||
let mut extension_data = ExtensionData::new();
|
||||
extension_data.set_extension_state("todo", "v0", json!("TODO content"));
|
||||
extension_data.set_extension_state("memory", "v1", json!({"key": "value"}));
|
||||
|
||||
// Serialize to JSON
|
||||
let json = serde_json::to_value(&extension_data).unwrap();
|
||||
|
||||
// Check the structure
|
||||
assert!(json.is_object());
|
||||
assert_eq!(json.get("todo.v0"), Some(&json!("TODO content")));
|
||||
assert_eq!(json.get("memory.v1"), Some(&json!({"key": "value"})));
|
||||
|
||||
// Deserialize back
|
||||
let deserialized: ExtensionData = serde_json::from_value(json).unwrap();
|
||||
assert_eq!(
|
||||
deserialized.get_extension_state("todo", "v0"),
|
||||
Some(&json!("TODO content"))
|
||||
);
|
||||
assert_eq!(
|
||||
deserialized.get_extension_state("memory", "v1"),
|
||||
Some(&json!({"key": "value"}))
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
pub mod extension_data;
|
||||
pub mod info;
|
||||
pub mod storage;
|
||||
|
||||
@@ -9,4 +10,5 @@ pub use storage::{
|
||||
SessionMetadata,
|
||||
};
|
||||
|
||||
pub use extension_data::{ExtensionData, ExtensionState, TodoState};
|
||||
pub use info::{get_valid_sorted_sessions, SessionInfo};
|
||||
|
||||
@@ -8,6 +8,7 @@
|
||||
use crate::conversation::message::Message;
|
||||
use crate::conversation::Conversation;
|
||||
use crate::providers::base::Provider;
|
||||
use crate::session::extension_data::ExtensionData;
|
||||
use crate::utils::safe_truncate;
|
||||
use anyhow::Result;
|
||||
use chrono::Local;
|
||||
@@ -64,11 +65,13 @@ pub struct SessionMetadata {
|
||||
pub accumulated_input_tokens: Option<i32>,
|
||||
/// The number of output tokens used in the session. Accumulated across all messages.
|
||||
pub accumulated_output_tokens: Option<i32>,
|
||||
/// Session-scoped TODO list content
|
||||
pub todo_content: Option<String>,
|
||||
|
||||
/// Extension data containing extension states
|
||||
#[serde(default)]
|
||||
pub extension_data: ExtensionData,
|
||||
}
|
||||
|
||||
// Custom deserializer to handle old sessions without working_dir and todo_content
|
||||
// Custom deserializer to handle old sessions without working_dir
|
||||
impl<'de> Deserialize<'de> for SessionMetadata {
|
||||
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
||||
where
|
||||
@@ -78,7 +81,7 @@ impl<'de> Deserialize<'de> for SessionMetadata {
|
||||
struct Helper {
|
||||
description: String,
|
||||
message_count: usize,
|
||||
schedule_id: Option<String>, // For backward compatibility
|
||||
schedule_id: Option<String>,
|
||||
total_tokens: Option<i32>,
|
||||
input_tokens: Option<i32>,
|
||||
output_tokens: Option<i32>,
|
||||
@@ -86,7 +89,8 @@ impl<'de> Deserialize<'de> for SessionMetadata {
|
||||
accumulated_input_tokens: Option<i32>,
|
||||
accumulated_output_tokens: Option<i32>,
|
||||
working_dir: Option<PathBuf>,
|
||||
todo_content: Option<String>, // For backward compatibility
|
||||
#[serde(default)]
|
||||
extension_data: ExtensionData,
|
||||
}
|
||||
|
||||
let helper = Helper::deserialize(deserializer)?;
|
||||
@@ -108,7 +112,7 @@ impl<'de> Deserialize<'de> for SessionMetadata {
|
||||
accumulated_input_tokens: helper.accumulated_input_tokens,
|
||||
accumulated_output_tokens: helper.accumulated_output_tokens,
|
||||
working_dir,
|
||||
todo_content: helper.todo_content,
|
||||
extension_data: helper.extension_data,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -133,7 +137,7 @@ impl SessionMetadata {
|
||||
accumulated_total_tokens: None,
|
||||
accumulated_input_tokens: None,
|
||||
accumulated_output_tokens: None,
|
||||
todo_content: None,
|
||||
extension_data: ExtensionData::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user