use crate::CliError;
use serde::{Deserialize, Serialize};
use std::fs;
use std::path::{Path, PathBuf};
const SKILL_TEMPLATE: &str = include_str!("../skills/mushroom/SKILL.md");
const CURSOR_RULES_TEMPLATE: &str = include_str!("../skills/mushroom/cursor-rules.mdc");
const DB_PATH_PLACEHOLDER: &str = "{{DB_PATH}}";
const BIN_PLACEHOLDER: &str = "{{BIN}}";
const SERVER_NAME: &str = "mushroomdb";
const BIN_NAME: &str = "mushroomdb";
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum BinaryLocation {
OnPath,
CopyFrom(PathBuf),
}
pub fn detect_binary_location() -> BinaryLocation {
if bin_on_path() {
return BinaryLocation::OnPath;
}
match std::env::current_exe() {
Ok(exe) => BinaryLocation::CopyFrom(exe),
Err(_) => BinaryLocation::OnPath,
}
}
fn bin_on_path() -> bool {
let Some(path) = std::env::var_os("PATH") else {
return false;
};
std::env::split_paths(&path).any(|dir| dir.join(BIN_NAME).is_file())
}
fn stable_bin_path(home: &Path) -> PathBuf {
home.join(".mushroomdb").join("bin").join(BIN_NAME)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Platform {
ClaudeCode,
Cursor,
All,
}
impl Platform {
pub fn parse(s: &str) -> Result<Self, String> {
match s {
"claude-code" => Ok(Platform::ClaudeCode),
"cursor" => Ok(Platform::Cursor),
"all" => Ok(Platform::All),
other => Err(format!(
"--platform must be claude-code | cursor | all, got: {other}"
)),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct InstallOpts {
pub platform: Option<Platform>,
pub project: bool,
pub db: Option<PathBuf>,
}
impl InstallOpts {
pub fn default_db(&self, project_root: &Path, home: &Path) -> PathBuf {
if self.project {
project_root.join("mushroom-memory")
} else {
home.join(".mushroomdb").join("memory")
}
}
}
#[derive(Serialize, Deserialize, Default, Debug)]
struct Manifest {
files: Vec<PathBuf>,
mcp_keys: Vec<ManagedMcpKey>,
#[serde(default)]
hooks: Vec<ManagedHook>,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
struct ManagedMcpKey {
file: PathBuf,
server: String,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)]
struct ManagedHook {
file: PathBuf,
event: String,
command: String,
}
const HOOK_EVENT: &str = "UserPromptSubmit";
const HOOK_TIMEOUT_SECS: u64 = 5;
fn sh_quote(s: &str) -> String {
format!("'{}'", s.replace('\'', r"'\''"))
}
fn recall_hook_command(bin_cmd: &str, db_str: &str) -> String {
format!("{} recall {}", sh_quote(bin_cmd), sh_quote(db_str))
}
fn hook_entry(command: &str) -> serde_json::Value {
serde_json::json!({ "hooks": [ { "type": "command", "command": command, "timeout": HOOK_TIMEOUT_SECS } ] })
}
fn settings_has_hook(root: &serde_json::Value, event: &str, command: &str) -> bool {
root["hooks"][event]
.as_array()
.map(|groups| {
groups.iter().any(|g| {
g["hooks"]
.as_array()
.map(|hs| hs.iter().any(|h| h["command"] == command))
.unwrap_or(false)
})
})
.unwrap_or(false)
}
fn merge_hook_entry(
settings_file: &Path,
command: &str,
manifest: &mut Manifest,
) -> Result<(), CliError> {
let mut root: serde_json::Value = if settings_file.exists() {
let raw = fs::read_to_string(settings_file)
.map_err(|e| CliError(format!("cannot read {}: {e}", settings_file.display())))?;
serde_json::from_str(&raw)
.map_err(|e| CliError(format!("invalid JSON in {}: {e}", settings_file.display())))?
} else {
serde_json::json!({})
};
if !root.is_object() {
return Err(CliError(format!(
"{} is not a JSON object at its top level — refusing to add a hook",
settings_file.display()
)));
}
if settings_has_hook(&root, HOOK_EVENT, command) {
return Ok(());
}
match root.get("hooks") {
None => root["hooks"] = serde_json::json!({}),
Some(v) if v.is_object() => {}
Some(_) => {
return Err(CliError(format!(
"{}: \"hooks\" is not a JSON object — refusing to overwrite it",
settings_file.display()
)));
}
}
match root["hooks"].get(HOOK_EVENT) {
None => root["hooks"][HOOK_EVENT] = serde_json::json!([]),
Some(v) if v.is_array() => {}
Some(_) => {
return Err(CliError(format!(
"{}: \"hooks.{HOOK_EVENT}\" is not a JSON array — refusing to overwrite it",
settings_file.display()
)));
}
}
root["hooks"][HOOK_EVENT]
.as_array_mut()
.unwrap()
.push(hook_entry(command));
let parent = settings_file.parent().unwrap_or(Path::new("."));
fs::create_dir_all(parent)
.map_err(|e| CliError(format!("cannot create {}: {e}", parent.display())))?;
let json = serde_json::to_string_pretty(&root)
.map_err(|e| CliError(format!("cannot serialize settings: {e}")))?;
fs::write(settings_file, json)
.map_err(|e| CliError(format!("cannot write {}: {e}", settings_file.display())))?;
manifest.hooks.push(ManagedHook {
file: settings_file.to_path_buf(),
event: HOOK_EVENT.into(),
command: command.into(),
});
Ok(())
}
fn remove_hook_entry(settings_file: &Path, event: &str, command: &str) -> Result<(), CliError> {
if !settings_file.exists() {
return Ok(());
}
let raw = fs::read_to_string(settings_file)
.map_err(|e| CliError(format!("cannot read {}: {e}", settings_file.display())))?;
let mut root: serde_json::Value = serde_json::from_str(&raw).map_err(|e| {
CliError(format!(
"corrupt settings json at {}: {e}",
settings_file.display()
))
})?;
let Some(mut groups) = root
.get("hooks")
.and_then(|h| h.get(event))
.and_then(|g| g.as_array())
.cloned()
else {
return Ok(());
};
for g in groups.iter_mut() {
if let Some(hs) = g["hooks"].as_array_mut() {
hs.retain(|h| h["command"] != command);
}
}
groups.retain(|g| {
g["hooks"]
.as_array()
.map(|hs| !hs.is_empty())
.unwrap_or(true)
});
let before = root.clone();
if groups.is_empty() {
root["hooks"].as_object_mut().unwrap().remove(event);
} else {
root["hooks"][event] = serde_json::Value::Array(groups);
}
if root == before {
return Ok(());
}
let json = serde_json::to_string_pretty(&root)
.map_err(|e| CliError(format!("cannot serialize settings: {e}")))?;
fs::write(settings_file, json)
.map_err(|e| CliError(format!("cannot write {}: {e}", settings_file.display())))?;
Ok(())
}
pub fn run_install(
project_root: &Path,
home: &Path,
opts: &InstallOpts,
) -> Result<String, CliError> {
run_install_with(project_root, home, opts, &detect_binary_location())
}
pub fn run_install_with(
project_root: &Path,
home: &Path,
opts: &InstallOpts,
bin: &BinaryLocation,
) -> Result<String, CliError> {
let db = opts
.db
.clone()
.unwrap_or_else(|| opts.default_db(project_root, home));
let db_str = db.to_string_lossy();
let resolved = resolve_platform(project_root, home, opts.platform.as_ref())?;
let platforms = expand_platform(&resolved);
for plat in &platforms {
preflight_check(project_root, home, plat, opts.project, &db_str)?;
}
let manifest_path = manifest_path(project_root, home, opts.project, &platforms);
let existing = load_manifest(&manifest_path);
let mut manifest = Manifest::default();
let bin_cmd = match bin {
BinaryLocation::OnPath => BIN_NAME.to_string(),
BinaryLocation::CopyFrom(src) => {
let dest = stable_bin_path(home);
copy_binary(src, &dest, &mut manifest)?;
dest.to_string_lossy().into_owned()
}
};
for plat in &platforms {
let step = install_platform(
project_root,
home,
plat,
opts.project,
&db_str,
&bin_cmd,
&mut manifest,
);
if let Err(e) = step {
let anything_written = !manifest.files.is_empty()
|| !manifest.mcp_keys.is_empty()
|| !manifest.hooks.is_empty();
if anything_written {
let merged = union_manifests(load_manifest(&manifest_path), &manifest);
let _ = write_manifest(&manifest_path, &merged);
}
return Err(e);
}
}
let anything_written =
!manifest.files.is_empty() || !manifest.mcp_keys.is_empty() || !manifest.hooks.is_empty();
if anything_written {
let merged = union_manifests(existing, &manifest);
write_manifest(&manifest_path, &merged)?;
}
let mut out = format!("mushroomdb installed ({} platform(s))\n", platforms.len());
for f in &manifest.files {
out.push_str(&format!(" wrote {}\n", f.display()));
}
for k in &manifest.mcp_keys {
out.push_str(&format!(
" added mcpServers.{} in {}\n",
k.server,
k.file.display()
));
}
for h in &manifest.hooks {
out.push_str(&format!(
" added {} hook in {}\n",
h.event,
h.file.display()
));
}
if anything_written {
out.push_str(&format!(" manifest {}\n", manifest_path.display()));
out.push_str(&format!(
" mcp command {bin_cmd}\n restart your assistant to connect the MCP server\n"
));
} else {
out.push_str(" (already installed — no changes)\n");
}
Ok(out)
}
fn copy_binary(src: &Path, dest: &Path, manifest: &mut Manifest) -> Result<(), CliError> {
let bytes = fs::read(src)
.map_err(|e| CliError(format!("cannot read binary {}: {e}", src.display())))?;
if fs::read(dest).map(|cur| cur == bytes).unwrap_or(false) {
return Ok(());
}
let parent = dest.parent().unwrap_or(Path::new("."));
fs::create_dir_all(parent)
.map_err(|e| CliError(format!("cannot create {}: {e}", parent.display())))?;
let tmp = parent.join(format!(".{BIN_NAME}.tmp-{}", std::process::id()));
fs::write(&tmp, &bytes)
.map_err(|e| CliError(format!("cannot write {}: {e}", tmp.display())))?;
let finish = || -> Result<(), CliError> {
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
fs::set_permissions(&tmp, fs::Permissions::from_mode(0o755))
.map_err(|e| CliError(format!("cannot chmod {}: {e}", tmp.display())))?;
}
fs::rename(&tmp, dest)
.map_err(|e| CliError(format!("cannot move binary into {}: {e}", dest.display())))
};
if let Err(e) = finish() {
let _ = fs::remove_file(&tmp); return Err(e);
}
manifest.files.push(dest.to_path_buf());
Ok(())
}
pub fn run_uninstall(
project_root: &Path,
home: &Path,
opts: &InstallOpts,
) -> Result<String, CliError> {
let resolved = resolve_platform(project_root, home, opts.platform.as_ref())?;
let platforms = expand_platform(&resolved);
let manifest_path = manifest_path(project_root, home, opts.project, &platforms);
if !manifest_path.exists() {
return Err(CliError(format!(
"no install manifest found at {} — nothing to uninstall",
manifest_path.display()
)));
}
let raw = fs::read_to_string(&manifest_path)
.map_err(|e| CliError(format!("cannot read manifest: {e}")))?;
let manifest: Manifest =
serde_json::from_str(&raw).map_err(|e| CliError(format!("corrupt manifest: {e}")))?;
let mut removed = Vec::new();
for key in &manifest.mcp_keys {
if key.file.exists() {
remove_mcp_key(&key.file, &key.server)?;
removed.push(format!(
"removed mcpServers.{} from {}",
key.server,
key.file.display()
));
}
}
for h in &manifest.hooks {
if h.file.exists() {
remove_hook_entry(&h.file, &h.event, &h.command)?;
removed.push(format!(
"removed {} hook from {}",
h.event,
h.file.display()
));
}
}
for f in &manifest.files {
if f.exists() {
fs::remove_file(f)
.map_err(|e| CliError(format!("cannot remove {}: {e}", f.display())))?;
removed.push(format!("removed {}", f.display()));
}
}
if manifest_path.exists() {
fs::remove_file(&manifest_path)
.map_err(|e| CliError(format!("cannot remove manifest: {e}")))?;
}
let mut out = "mushroomdb uninstalled\n".to_string();
for line in &removed {
out.push_str(&format!(" {line}\n"));
}
Ok(out)
}
fn resolve_platform(
project_root: &Path,
home: &Path,
requested: Option<&Platform>,
) -> Result<Platform, CliError> {
if let Some(p) = requested {
return Ok(p.clone());
}
let has_claude = home.join(".claude").exists() || project_root.join(".claude").exists();
let has_cursor = project_root.join(".cursor").exists() || home.join(".cursor").exists();
match (has_claude, has_cursor) {
(true, true) => Ok(Platform::All),
(true, false) => Ok(Platform::ClaudeCode),
(false, true) => Ok(Platform::Cursor),
(false, false) => Err(CliError(
"cannot auto-detect platform: neither ~/.claude nor .cursor/ found.\n\
Pass --platform claude-code, --platform cursor, or --platform all."
.to_string(),
)),
}
}
fn expand_platform(p: &Platform) -> Vec<Platform> {
match p {
Platform::All => vec![Platform::ClaudeCode, Platform::Cursor],
Platform::ClaudeCode => vec![Platform::ClaudeCode],
Platform::Cursor => vec![Platform::Cursor],
}
}
fn preflight_check(
project_root: &Path,
home: &Path,
platform: &Platform,
project_scope: bool,
db_str: &str,
) -> Result<(), CliError> {
match platform {
Platform::ClaudeCode => {
let mcp_file = if project_scope {
project_root.join(".mcp.json")
} else {
home.join(".claude.json")
};
check_mcp_conflict(&mcp_file, db_str)?;
}
Platform::Cursor => {
let mcp_file = if project_scope {
project_root.join(".cursor").join("mcp.json")
} else {
home.join(".cursor").join("mcp.json")
};
check_mcp_conflict(&mcp_file, db_str)?;
}
Platform::All => unreachable!("expand_platform never produces All"),
}
Ok(())
}
fn check_mcp_conflict(mcp_file: &Path, db_str: &str) -> Result<(), CliError> {
if !mcp_file.exists() {
return Ok(());
}
let raw = fs::read_to_string(mcp_file)
.map_err(|e| CliError(format!("cannot read {}: {e}", mcp_file.display())))?;
let v: serde_json::Value = serde_json::from_str(&raw)
.map_err(|e| CliError(format!("invalid JSON in {}: {e}", mcp_file.display())))?;
let existing = &v["mcpServers"][SERVER_NAME];
if existing.is_null() {
return Ok(()); }
let existing_db = existing["args"]
.get(1)
.and_then(|v| v.as_str())
.unwrap_or("");
if existing_db == db_str {
return Ok(()); }
Err(CliError(format!(
"conflict: {} already has mcpServers.mushroomdb pointing to {:?}\n\
To update it, run `mushroomdb uninstall` first, then re-install.\n\
Or manually edit {} and remove the existing mushroomdb entry.",
mcp_file.display(),
existing_db,
mcp_file.display()
)))
}
fn install_platform(
project_root: &Path,
home: &Path,
platform: &Platform,
project_scope: bool,
db_str: &str,
bin_cmd: &str,
manifest: &mut Manifest,
) -> Result<(), CliError> {
match platform {
Platform::ClaudeCode => {
install_claude_code(project_root, home, project_scope, db_str, bin_cmd, manifest)
}
Platform::Cursor => {
install_cursor(project_root, home, project_scope, db_str, bin_cmd, manifest)
}
Platform::All => unreachable!("expand_platform never produces All"),
}
}
fn render_template(template: &str, db_str: &str, bin_cmd: &str) -> String {
template
.replace(DB_PATH_PLACEHOLDER, db_str)
.replace(BIN_PLACEHOLDER, bin_cmd)
}
fn install_claude_code(
project_root: &Path,
home: &Path,
project_scope: bool,
db_str: &str,
bin_cmd: &str,
manifest: &mut Manifest,
) -> Result<(), CliError> {
let skill_content = render_template(SKILL_TEMPLATE, db_str, bin_cmd);
let skill_dir = if project_scope {
project_root.join(".claude").join("skills").join("mushroom")
} else {
home.join(".claude").join("skills").join("mushroom")
};
let skill_file = skill_dir.join("SKILL.md");
if !file_matches(&skill_file, &skill_content) {
fs::create_dir_all(&skill_dir)
.map_err(|e| CliError(format!("cannot create {}: {e}", skill_dir.display())))?;
fs::write(&skill_file, &skill_content)
.map_err(|e| CliError(format!("cannot write {}: {e}", skill_file.display())))?;
manifest.files.push(skill_file);
}
let mcp_file = if project_scope {
project_root.join(".mcp.json")
} else {
home.join(".claude.json")
};
merge_mcp_entry(&mcp_file, db_str, bin_cmd, manifest)?;
let settings_file = if project_scope {
project_root.join(".claude").join("settings.json")
} else {
home.join(".claude").join("settings.json")
};
merge_hook_entry(
&settings_file,
&recall_hook_command(bin_cmd, db_str),
manifest,
)?;
Ok(())
}
fn install_cursor(
project_root: &Path,
home: &Path,
project_scope: bool,
db_str: &str,
bin_cmd: &str,
manifest: &mut Manifest,
) -> Result<(), CliError> {
let rules_content = render_template(CURSOR_RULES_TEMPLATE, db_str, bin_cmd);
let rules_dir = if project_scope {
project_root.join(".cursor").join("rules")
} else {
home.join(".cursor").join("rules")
};
let rules_file = rules_dir.join("mushroom.mdc");
if !file_matches(&rules_file, &rules_content) {
fs::create_dir_all(&rules_dir)
.map_err(|e| CliError(format!("cannot create {}: {e}", rules_dir.display())))?;
fs::write(&rules_file, &rules_content)
.map_err(|e| CliError(format!("cannot write {}: {e}", rules_file.display())))?;
manifest.files.push(rules_file);
}
let mcp_file = if project_scope {
project_root.join(".cursor").join("mcp.json")
} else {
home.join(".cursor").join("mcp.json")
};
merge_mcp_entry(&mcp_file, db_str, bin_cmd, manifest)?;
Ok(())
}
fn merge_mcp_entry(
mcp_file: &Path,
db_str: &str,
bin_cmd: &str,
manifest: &mut Manifest,
) -> Result<(), CliError> {
let mut root: serde_json::Value = if mcp_file.exists() {
let raw = fs::read_to_string(mcp_file)
.map_err(|e| CliError(format!("cannot read {}: {e}", mcp_file.display())))?;
serde_json::from_str(&raw)
.map_err(|e| CliError(format!("invalid JSON in {}: {e}", mcp_file.display())))?
} else {
serde_json::json!({})
};
if !root["mcpServers"].is_object() {
root["mcpServers"] = serde_json::json!({});
}
let desired = mcp_server_entry(db_str, bin_cmd);
let existing = &root["mcpServers"][SERVER_NAME];
if existing == &desired {
return Ok(()); }
root["mcpServers"][SERVER_NAME] = desired;
let parent = mcp_file.parent().unwrap_or(Path::new("."));
fs::create_dir_all(parent)
.map_err(|e| CliError(format!("cannot create {}: {e}", parent.display())))?;
let json = serde_json::to_string_pretty(&root)
.map_err(|e| CliError(format!("cannot serialize mcp json: {e}")))?;
fs::write(mcp_file, json)
.map_err(|e| CliError(format!("cannot write {}: {e}", mcp_file.display())))?;
manifest.mcp_keys.push(ManagedMcpKey {
file: mcp_file.to_path_buf(),
server: SERVER_NAME.to_string(),
});
Ok(())
}
fn remove_mcp_key(mcp_file: &Path, server: &str) -> Result<(), CliError> {
if !mcp_file.exists() {
return Ok(());
}
let raw = fs::read_to_string(mcp_file)
.map_err(|e| CliError(format!("cannot read {}: {e}", mcp_file.display())))?;
let mut root: serde_json::Value = serde_json::from_str(&raw)
.map_err(|e| CliError(format!("corrupt mcp json at {}: {e}", mcp_file.display())))?;
if let Some(servers) = root["mcpServers"].as_object_mut() {
servers.remove(server);
}
let json = serde_json::to_string_pretty(&root)
.map_err(|e| CliError(format!("cannot serialize mcp json: {e}")))?;
fs::write(mcp_file, json)
.map_err(|e| CliError(format!("cannot write {}: {e}", mcp_file.display())))?;
Ok(())
}
fn mcp_server_entry(db_str: &str, bin_cmd: &str) -> serde_json::Value {
serde_json::json!({
"command": bin_cmd,
"args": ["mcp", db_str]
})
}
fn manifest_path(
project_root: &Path,
home: &Path,
project_scope: bool,
platforms: &[Platform],
) -> PathBuf {
if !project_scope {
return home.join(".mushroomdb").join("install-manifest.json");
}
if platforms.contains(&Platform::ClaudeCode) {
project_root
.join(".claude")
.join("skills")
.join("mushroom")
.join(".install-manifest.json")
} else {
project_root.join(".cursor").join(".install-manifest.json")
}
}
fn load_manifest(path: &Path) -> Manifest {
let raw = match fs::read_to_string(path) {
Ok(s) => s,
Err(_) => return Manifest::default(),
};
serde_json::from_str(&raw).unwrap_or_default()
}
fn union_manifests(mut existing: Manifest, this_run: &Manifest) -> Manifest {
for f in &this_run.files {
if !existing.files.contains(f) {
existing.files.push(f.clone());
}
}
for k in &this_run.mcp_keys {
let already = existing
.mcp_keys
.iter()
.any(|e| e.file == k.file && e.server == k.server);
if !already {
existing.mcp_keys.push(k.clone());
}
}
for h in &this_run.hooks {
if !existing.hooks.contains(h) {
existing.hooks.push(h.clone());
}
}
existing
}
fn write_manifest(path: &Path, manifest: &Manifest) -> Result<(), CliError> {
let parent = path.parent().unwrap_or(Path::new("."));
fs::create_dir_all(parent).map_err(|e| {
CliError(format!(
"cannot create manifest dir {}: {e}",
parent.display()
))
})?;
let json = serde_json::to_string_pretty(manifest)
.map_err(|e| CliError(format!("cannot serialize manifest: {e}")))?;
fs::write(path, json)
.map_err(|e| CliError(format!("cannot write manifest {}: {e}", path.display())))?;
Ok(())
}
fn file_matches(path: &Path, expected: &str) -> bool {
fs::read_to_string(path)
.map(|s| s == expected)
.unwrap_or(false)
}