Skip to content
Merged
26 changes: 26 additions & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

3 changes: 3 additions & 0 deletions crates/goose/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,9 @@ aws-sdk-bedrockruntime = "1.72.0"
# For GCP Vertex AI provider auth
jsonwebtoken = "9.3.1"

# Added blake3 hashing library as a dependency
blake3 = "1.5"

[target.'cfg(target_os = "windows")'.dependencies]
winapi = { version = "0.3", features = ["wincred"] }

Expand Down
2 changes: 2 additions & 0 deletions crates/goose/src/agents/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ mod capabilities;
pub mod extension;
mod factory;
mod permission_judge;
mod permission_store;
mod reference;
mod truncate;

Expand All @@ -11,3 +12,4 @@ pub use capabilities::Capabilities;
pub use extension::ExtensionConfig;
pub use factory::{register_agent, AgentFactory};
pub use permission_judge::detect_read_only_tools;
pub use permission_store::ToolPermissionStore;
295 changes: 295 additions & 0 deletions crates/goose/src/agents/permission_store.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,295 @@
use crate::message::ToolRequest;
use anyhow::Result;
use blake3::Hasher;
use chrono::Utc;
use etcetera::{choose_app_strategy, AppStrategy};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::time::Duration;
use std::{fs::File, path::PathBuf};

#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct ToolPermissionRecord {
tool_name: String,
allowed: bool,
context_hash: String, // Hash of the tool's arguments/context to differentiate similar calls
#[serde(skip_serializing_if = "Option::is_none")] // Don't serialize if None
readable_context: Option<String>, // Add this field
timestamp: i64,
expiry: Option<i64>, // Optional expiry timestamp
}

#[derive(Debug, Serialize, Deserialize)]
pub struct ToolPermissionStore {
permissions: HashMap<String, Vec<ToolPermissionRecord>>,
version: u32, // For future schema migrations
#[serde(skip)] // Don't serialize this field
permissions_dir: PathBuf,
}

impl Default for ToolPermissionStore {
fn default() -> Self {
Self::new()
}
}

impl ToolPermissionStore {
pub fn new() -> Self {
let permissions_dir = choose_app_strategy(crate::config::APP_STRATEGY.clone())
.map(|strategy| strategy.config_dir())
.unwrap_or_else(|_| PathBuf::from(".config/goose"));

Self {
permissions: HashMap::new(),
version: 1,
permissions_dir,
}
}

pub fn load() -> Result<Self> {
let store = Self::new();
let file_path = store.permissions_dir.join("tool_permissions.json");

if !file_path.exists() {
return Ok(store);
}

let file = File::open(file_path)?;
let mut permissions: ToolPermissionStore = serde_json::from_reader(file)?;
permissions.permissions_dir = store.permissions_dir;

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

do we need to clean up the expired entries?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

yes, good call


// Clean up expired entries on load
permissions.cleanup_expired()?;

Ok(permissions)
}

pub fn save(&self) -> anyhow::Result<()> {
std::fs::create_dir_all(&self.permissions_dir)?;

let path = self.permissions_dir.join("tool_permissions.json");
let temp_path = path.with_extension("tmp");

// Write complete content to temporary file
let content = serde_json::to_string_pretty(self)?;
std::fs::write(&temp_path, &content)?;

// Atomically rename temp file to target file
std::fs::rename(temp_path, path)?;

Ok(())
}

pub fn check_permission(&self, tool_request: &ToolRequest) -> Option<bool> {
let context_hash = self.hash_tool_context(tool_request);
let tool_call = tool_request.tool_call.as_ref().unwrap();
let key = format!("{}:{}", tool_call.name, context_hash);

self.permissions.get(&key).and_then(|records| {
records
.iter()
.filter(|record| record.expiry.is_none_or(|exp| exp > Utc::now().timestamp()))
.last()
.map(|record| record.allowed)
})
}

pub fn record_permission(
&mut self,
tool_request: &ToolRequest,
allowed: bool,
expiry_duration: Option<Duration>,
) -> anyhow::Result<()> {
let context_hash = self.hash_tool_context(tool_request);
let tool_call = tool_request.tool_call.as_ref().unwrap();
let key = format!("{}:{}", tool_call.name, context_hash);

let record = ToolPermissionRecord {
tool_name: tool_call.name.clone(),
allowed,
context_hash,
readable_context: Some(tool_request.to_readable_string()),
timestamp: Utc::now().timestamp(),
expiry: expiry_duration.map(|d| Utc::now().timestamp() + d.as_secs() as i64),
};

self.permissions.entry(key).or_default().push(record);

self.save()?;
Ok(())
}

fn hash_tool_context(&self, tool_request: &ToolRequest) -> String {
// Create a hash of the tool's arguments to differentiate similar calls
// This helps identify when the same tool is being used in a different context
let mut hasher = Hasher::new();
hasher.update(
serde_json::to_string(&tool_request.tool_call.as_ref().unwrap().arguments)

@yingjiehe-xyz yingjiehe-xyz Mar 7, 2025

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

do we want to hash on some argument params? I am wondering whether we will have low hit rate if we hash all, including argument param and value, like write file, the name can be different, but they are super similar

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Similar thought - I was thinking to hash by tool name only at first, but this would lump all bash commands in one hash.

Any ideas about how to do the hash at the right level of granularity?

@yingjiehe-xyz yingjiehe-xyz Mar 7, 2025

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

yeah, kind of tricky, introducing some cosine similarity check may be overkilled. I am wondering whether we can have a set of normalized keys? for example in bash case, we have "command":write, "file_path": xxx, and we maintain a set containing "file_path", if it is in the normalized key set, like "file_path", we replace the value with "normalized_value" before hashing. Even with such case, we cannot handle all argument keys, we can just start with bash? how do you think?

.unwrap_or_default()
.as_bytes(),
);
hasher.finalize().to_hex().to_string()
}

pub fn cleanup_expired(&mut self) -> anyhow::Result<()> {
let now = Utc::now().timestamp();
let mut changed = false;

self.permissions.retain(|_, records| {
records.retain(|record| record.expiry.map_or(true, |exp| exp > now));
changed = changed || records.is_empty();
!records.is_empty()
});

if changed {
self.save()?;
}
Ok(())
}
}

#[cfg(test)]
mod tests {
use crate::agents::permission_store::{ToolPermissionRecord, ToolPermissionStore};
use crate::message::ToolRequest;
use chrono::Utc;
use mcp_core::tool::ToolCall;
use std::time::Duration;

fn create_test_tool_request(name: &str, args: serde_json::Value) -> ToolRequest {
ToolRequest {
id: "test-id".to_string(),
tool_call: Ok(ToolCall {
name: name.to_string(),
arguments: args,
}),
}
}

#[test]
fn test_permission_store_basic() {
let mut store = ToolPermissionStore::new();
let tool_request =
create_test_tool_request("test_tool", serde_json::json!({"arg1": "value1"}));

// Initially no permission recorded
assert!(store.check_permission(&tool_request).is_none());

// Record a permission
store.record_permission(&tool_request, true, None).unwrap();

// Should now find the recorded permission
assert_eq!(store.check_permission(&tool_request), Some(true));
}

#[test]
fn test_permission_expiry() {
let mut store = ToolPermissionStore::new();
let tool_request =
create_test_tool_request("test_tool", serde_json::json!({"arg1": "value1"}));

// Record a permission that expires in 1 second
store
.record_permission(&tool_request, true, Some(Duration::from_secs(1)))
.unwrap();

// Should initially be allowed
assert_eq!(store.check_permission(&tool_request), Some(true));

// Manually set expiry to the past
if let Some(records) = store.permissions.get_mut(&format!(
"{}:{}",
tool_request.tool_call.as_ref().unwrap().name,
store.hash_tool_context(&tool_request)
)) {
if let Some(record) = records.last_mut() {
record.expiry = Some(Utc::now().timestamp() - 2);
}
}

// Should now be expired (no permission found)
assert!(store.check_permission(&tool_request).is_none());
}

#[test]
fn test_different_arguments() {
let mut store = ToolPermissionStore::new();

// Create two requests with same tool but different args
let request1 = create_test_tool_request("test_tool", serde_json::json!({"arg": "value1"}));
let request2 = create_test_tool_request("test_tool", serde_json::json!({"arg": "value2"}));

// Record permission for first request
store.record_permission(&request1, true, None).unwrap();

// Should only find permission for first request
assert_eq!(store.check_permission(&request1), Some(true));
assert!(store.check_permission(&request2).is_none());
}

#[test]
fn test_cleanup_expired() {
let mut store = ToolPermissionStore::new();
let tool_request =
create_test_tool_request("test_tool", serde_json::json!({"arg1": "value1"}));

// Compute hash and key first
let context_hash = store.hash_tool_context(&tool_request);
let key = format!(
"{}:{}",
tool_request.tool_call.as_ref().unwrap().name,
context_hash
);

// Add an expired permission
store
.permissions
.entry(key.clone())
.or_default()
.push(ToolPermissionRecord {
tool_name: "test_tool".to_string(),
allowed: true,
context_hash: context_hash.clone(),
readable_context: None,
timestamp: Utc::now().timestamp(),
expiry: Some(Utc::now().timestamp() - 1000),
});

// Add a valid permission
store
.record_permission(&tool_request, true, Some(Duration::from_secs(3600)))
.unwrap();

// Before cleanup - should have 2 records
assert_eq!(
store
.permissions
.get(&format!(
"{}:{}",
tool_request.tool_call.as_ref().unwrap().name,
store.hash_tool_context(&tool_request)
))
.map(|records| records.len()),
Some(2)
);

// Run cleanup
store.cleanup_expired().unwrap();

// After cleanup - should have 1 record
assert_eq!(
store
.permissions
.get(&format!(
"{}:{}",
tool_request.tool_call.as_ref().unwrap().name,
store.hash_tool_context(&tool_request)
))
.map(|records| records.len()),
Some(1)
);

// The remaining record should be valid
assert_eq!(store.check_permission(&tool_request), Some(true));
}
}
Loading