Nest TODO State in session data (#4361)

Co-authored-by: Alex Hancock <alexhancock@block.xyz>
This commit is contained in:
David Katz
2025-08-28 14:11:29 -04:00
committed by GitHub
parent 7879445c4b
commit ef57f5062d
11 changed files with 294 additions and 45 deletions
+173
View File
@@ -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"}))
);
}
}
+2
View File
@@ -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};
+11 -7
View File
@@ -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(),
}
}
}