Files
mcp-memory/server/src/handlers/git.rs
T

217 lines
7.9 KiB
Rust

use crate::router::McpTool;
use crate::state::MemoryState;
use crate::tools::GetActiveWorktreeContextTool;
use async_trait::async_trait;
use serde_json::{Value, json};
use std::env;
use std::sync::Arc;
pub struct GetActiveWorktreeContextHandler;
#[async_trait]
impl McpTool for GetActiveWorktreeContextHandler {
fn name(&self) -> &'static str {
"get_active_worktree_context"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<GetActiveWorktreeContextTool>(
"get_active_worktree_context",
"Get the active worktree context, including branch name, modified files, and a truncated git diff.",
)
}
async fn execute(
&self,
_args: Value,
_state: Arc<MemoryState>,
) -> crate::error::Result<String> {
let result =
tokio::task::spawn_blocking(move || -> crate::error::Result<serde_json::Value> {
let cwd = env::current_dir().map_err(|e| e.to_string())?;
let repo = git2::Repository::discover(&cwd).map_err(|e| {
crate::error::AppError::Internal(format!("Not in a git repository: {}", e))
})?;
let mut branch_name = String::new();
if let Ok(head) = repo.head()
&& let Some(name) = head.shorthand()
{
branch_name = name.to_string();
}
let mut opts = git2::DiffOptions::new();
let mut diff = None;
// Try to diff against HEAD
if let Ok(tree) = repo.head().and_then(|h| h.peel_to_tree()) {
diff = repo
.diff_tree_to_workdir_with_index(Some(&tree), Some(&mut opts))
.ok();
}
let mut files = Vec::new();
let mut diff_output = String::new();
if let Some(diff) = diff {
let _ = diff.print(git2::DiffFormat::Patch, |_delta, _hunk, line| {
match line.origin() {
'+' | '-' | ' ' => diff_output.push(line.origin()),
_ => {}
}
let content = std::str::from_utf8(line.content()).unwrap_or("");
diff_output.push_str(content);
true
});
for delta in diff.deltas() {
if let Some(path) = delta.new_file().path() {
files.push(path.to_string_lossy().into_owned());
}
}
}
// Truncate diff output if it's too large to save tokens
if diff_output.len() > 10000 {
let valid_boundary = diff_output.floor_char_boundary(10000);
diff_output.truncate(valid_boundary);
diff_output.push_str("\n... [Diff truncated due to size]");
}
Ok(json!({
"branch": branch_name,
"modified_files": files,
"diff": diff_output
}))
})
.await
.map_err(|e| crate::error::AppError::Internal(format!("Task panic: {}", e)))??;
Ok::<String, crate::error::AppError>(serde_json::to_string_pretty(&result)?)
}
}
pub struct QueryGitDiffsHandler;
#[async_trait]
impl McpTool for QueryGitDiffsHandler {
fn name(&self) -> &'static str {
"query_git_diffs"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<crate::tools::QueryGitDiffsTool>(
"query_git_diffs",
"Query recent git commit history, diffs, and change ledger entries.",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String> {
let req: crate::tools::QueryGitDiffsTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let limit = req.limit.unwrap_or(5);
let q = req.query.to_lowercase();
let mut matches = Vec::new();
state.code.ledger.read_with(|ledger| {
for change in ledger {
if change.file_path.to_lowercase().contains(&q)
|| change.description.to_lowercase().contains(&q)
|| change.git_commit.as_ref().is_some_and(|c| c.contains(&q))
{
matches.push(json!({
"file_path": change.file_path,
"description": change.description,
"commit": change.git_commit,
"branch": change.git_branch,
"timestamp": change.timestamp,
}));
if matches.len() >= limit {
break;
}
}
}
});
if matches.len() < limit {
let remaining = limit - matches.len();
let git_matches = tokio::task::spawn_blocking(move || {
let mut results = Vec::new();
let cwd = env::current_dir().unwrap_or_default();
if let Ok(repo) = git2::Repository::discover(&cwd)
&& let Ok(mut revwalk) = repo.revwalk() {
let _ = revwalk.push_head();
let mut count = 0;
for oid in revwalk.flatten() {
if count >= remaining {
break;
}
if let Ok(commit) = repo.find_commit(oid) {
let summary = commit.summary().unwrap_or("");
if summary.to_lowercase().contains(&q) {
count += 1;
results.push(json!({
"commit_id": oid.to_string(),
"author": commit.author().name().unwrap_or("unknown"),
"message": summary,
"timestamp": commit.time().seconds(),
}));
}
}
}
}
results
})
.await
.unwrap_or_default();
matches.extend(git_matches);
}
Ok(serde_json::to_string_pretty(&matches)?)
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use std::sync::Arc;
use tempfile::tempdir;
#[tokio::test]
async fn test_get_active_worktree_context() {
let dir = tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let handler = GetActiveWorktreeContextHandler;
let result = handler
.execute(json!({}), state)
.await
.map_err(|e| format!("Failed to get worktree context: {}", e))
.unwrap();
let parsed: serde_json::Value = serde_json::from_str(&result).unwrap();
assert!(parsed.get("branch").is_some());
assert!(parsed.get("modified_files").is_some());
assert!(parsed.get("diff").is_some());
}
#[tokio::test]
async fn test_get_active_worktree_context_empty_git_repo() {
let dir = tempfile::tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let handler = GetActiveWorktreeContextHandler;
let result = handler
.execute(serde_json::json!({}), state)
.await
.map_err(|e| format!("Failed to get worktree context: {}", e))
.unwrap();
let parsed: serde_json::Value = serde_json::from_str(&result).unwrap();
assert!(parsed.get("branch").is_some() || parsed.is_object());
}
}