Skip to main content

remem/rules/
hook.rs

1use std::path::Path;
2
3use anyhow::{bail, Context, Result};
4use serde::Deserialize;
5use serde_json::{json, Value};
6use sha2::{Digest, Sha256};
7
8use super::{
9    artifact_path_for_project, evaluate_artifact_file_with_codes, EvaluationDiagnosticCode,
10    EvaluationInput, EvaluationVerdict, RuleMatch,
11};
12
13#[derive(Debug, Deserialize)]
14struct PreToolUsePayload {
15    session_id: Option<String>,
16    cwd: String,
17    hook_event_name: String,
18    tool_name: String,
19    tool_input: Value,
20}
21
22#[derive(Debug, Clone, PartialEq)]
23pub struct RuleHookEvaluation {
24    pub session_id: Option<String>,
25    pub output: Option<Value>,
26    pub diagnostics: Vec<String>,
27}
28
29#[derive(Debug, Clone, PartialEq)]
30pub(crate) struct DetailedRuleHookEvaluation {
31    pub evaluation: RuleHookEvaluation,
32    pub project: Option<String>,
33    pub diagnostic_codes: Vec<EvaluationDiagnosticCode>,
34}
35
36pub fn session_id_hint(raw: &str) -> Option<String> {
37    serde_json::from_str::<Value>(raw)
38        .ok()?
39        .get("session_id")?
40        .as_str()
41        .map(str::to_string)
42}
43
44pub(crate) fn project_hint(raw: &str) -> Option<String> {
45    let cwd = serde_json::from_str::<Value>(raw)
46        .ok()?
47        .get("cwd")?
48        .as_str()?
49        .to_string();
50    Some(crate::db::project_from_cwd(&cwd))
51}
52
53pub fn evaluate_pre_tool_use(
54    raw: &str,
55    host: Option<&str>,
56    data_dir: &Path,
57    enabled: bool,
58) -> Result<RuleHookEvaluation> {
59    evaluate_pre_tool_use_with_diagnostics(raw, host, data_dir, enabled)
60        .map(|detailed| detailed.evaluation)
61}
62
63pub(crate) fn evaluate_pre_tool_use_with_diagnostics(
64    raw: &str,
65    host: Option<&str>,
66    data_dir: &Path,
67    enabled: bool,
68) -> Result<DetailedRuleHookEvaluation> {
69    let host = host
70        .map(crate::runtime_config::normalize_host)
71        .unwrap_or_else(|| "unknown".to_string());
72    if host != crate::runtime_config::CLAUDE_HOST {
73        bail!("compiled command-rule enforcement is unsupported for host '{host}'");
74    }
75
76    if !enabled {
77        return Ok(DetailedRuleHookEvaluation {
78            evaluation: RuleHookEvaluation {
79                session_id: session_id_hint(raw),
80                output: None,
81                diagnostics: Vec::new(),
82            },
83            project: None,
84            diagnostic_codes: Vec::new(),
85        });
86    }
87
88    let payload: PreToolUsePayload =
89        serde_json::from_str(raw).context("parse Claude PreToolUse hook input")?;
90    if payload.hook_event_name != "PreToolUse" {
91        bail!(
92            "rules eval expected hook_event_name=PreToolUse, got '{}'",
93            payload.hook_event_name
94        );
95    }
96    if payload.tool_name != "Bash" {
97        bail!(
98            "rules eval expected tool_name=Bash, got '{}'",
99            payload.tool_name
100        );
101    }
102    let command = payload
103        .tool_input
104        .get("command")
105        .and_then(Value::as_str)
106        .filter(|command| !command.trim().is_empty())
107        .context("Claude PreToolUse Bash input is missing tool_input.command")?;
108    let project = crate::db::project_from_cwd(&payload.cwd);
109    let coded_outcome = evaluate_artifact_file_with_codes(
110        artifact_path_for_project(data_dir, &project),
111        &EvaluationInput {
112            command: command.to_string(),
113        },
114    );
115    let diagnostic_codes = coded_outcome.diagnostic_codes;
116    let outcome = coded_outcome.outcome;
117    let diagnostics = outcome
118        .diagnostics
119        .into_iter()
120        .map(|diagnostic| sanitize_diagnostic(&diagnostic.message))
121        .collect::<Vec<_>>();
122    if !diagnostics.is_empty() {
123        return Ok(DetailedRuleHookEvaluation {
124            evaluation: RuleHookEvaluation {
125                session_id: payload.session_id,
126                output: None,
127                diagnostics,
128            },
129            project: Some(project),
130            diagnostic_codes,
131        });
132    }
133
134    Ok(DetailedRuleHookEvaluation {
135        evaluation: RuleHookEvaluation {
136            session_id: payload.session_id,
137            output: render_hook_output(outcome.verdict, &outcome.matches),
138            diagnostics,
139        },
140        project: Some(project),
141        diagnostic_codes,
142    })
143}
144
145fn render_hook_output(verdict: EvaluationVerdict, matches: &[RuleMatch]) -> Option<Value> {
146    match verdict {
147        EvaluationVerdict::Allow => None,
148        EvaluationVerdict::Warn => {
149            let message = static_match_message("warning", matches);
150            Some(json!({
151                "systemMessage": message,
152                "hookSpecificOutput": {
153                    "hookEventName": "PreToolUse",
154                    "additionalContext": message,
155                }
156            }))
157        }
158        EvaluationVerdict::Block => {
159            let message = static_match_message("blocked", matches);
160            Some(json!({
161                "systemMessage": message,
162                "hookSpecificOutput": {
163                    "hookEventName": "PreToolUse",
164                    "permissionDecision": "deny",
165                    "permissionDecisionReason": message,
166                }
167            }))
168        }
169    }
170}
171
172fn static_match_message(disposition: &str, matches: &[RuleMatch]) -> String {
173    let mut source_ids = matches
174        .iter()
175        .map(|matched| matched.source_memory_id)
176        .collect::<Vec<_>>();
177    source_ids.sort_unstable();
178    source_ids.dedup();
179    let sources = source_ids
180        .iter()
181        .map(ToString::to_string)
182        .collect::<Vec<_>>()
183        .join(", ");
184    format!(
185        "remem compiled preference rule {disposition} for source memory(s) {sources}; inspect with `remem rules list`."
186    )
187}
188
189fn sanitize_diagnostic(message: &str) -> String {
190    let single_line = message
191        .chars()
192        .map(|ch| if ch.is_control() { ' ' } else { ch })
193        .collect::<String>();
194    crate::db::truncate_str(&single_line, 1000).to_string()
195}
196
197pub fn log_evaluation_error_once(data_dir: &Path, session_id: Option<&str>, message: &str) {
198    log_evaluation_error_once_with_diagnostic(
199        data_dir,
200        session_id,
201        None,
202        &[EvaluationDiagnosticCode::HookInput],
203        message,
204    );
205}
206
207pub(crate) fn log_evaluation_error_once_with_diagnostic(
208    data_dir: &Path,
209    session_id: Option<&str>,
210    project: Option<&str>,
211    codes: &[EvaluationDiagnosticCode],
212    message: &str,
213) {
214    let Some(session_key) = session_id.filter(|session_id| !session_id.trim().is_empty()) else {
215        crate::log::error("rules-eval", &sanitize_diagnostic(message));
216        return;
217    };
218    let marker_dir = super::evaluation_marker_dir(data_dir);
219    if let Err(error) = std::fs::create_dir_all(&marker_dir) {
220        crate::log::error(
221            "rules-eval",
222            &format!(
223                "could not create evaluation diagnostic marker directory: {error}; {}",
224                sanitize_diagnostic(message)
225            ),
226        );
227        return;
228    }
229    let digest = Sha256::digest(session_key.as_bytes());
230    let marker = marker_dir.join(format!("{digest:x}"));
231    match super::upsert_evaluation_error_record(&marker, data_dir, project, codes) {
232        Ok(true) => crate::log::error("rules-eval", &sanitize_diagnostic(message)),
233        Ok(false) => {}
234        Err(error) => crate::log::error(
235            "rules-eval",
236            &format!(
237                "could not publish evaluation diagnostic marker: {error:#}; {}",
238                sanitize_diagnostic(message)
239            ),
240        ),
241    }
242}
243
244#[cfg(test)]
245mod tests;