Skip to main content

polyhook_core/
detect.rs

1use crate::types::CallerKind;
2
3/// Detect which agent is calling the hook.
4///
5/// Priority:
6/// 1. `POLYHOOK_CALLER` env var (explicit override)
7/// 2. Agent-specific env vars
8/// 3. Heuristics on the raw stdin JSON shape
9/// 4. `Unknown`
10pub fn detect_caller(stdin: &serde_json::Value) -> CallerKind {
11    // 1. Explicit override via env var
12    if let Ok(val) = std::env::var("POLYHOOK_CALLER") {
13        match val.to_lowercase().as_str() {
14            "claude-code" | "claudecode" => return CallerKind::ClaudeCode,
15            "cursor" => return CallerKind::Cursor,
16            "windsurf" => return CallerKind::Windsurf,
17            "cline" => return CallerKind::Cline,
18            "amp" => return CallerKind::Amp,
19            _ => {}
20        }
21    }
22
23    // 2. Agent-specific env vars
24    if std::env::var("CLAUDE_CODE_VERSION").is_ok() {
25        return CallerKind::ClaudeCode;
26    }
27    if std::env::var("CURSOR_SESSION_ID").is_ok() {
28        return CallerKind::Cursor;
29    }
30    if std::env::var("WINDSURF_SESSION_ID").is_ok() {
31        return CallerKind::Windsurf;
32    }
33    if std::env::var("CLINE_SESSION_ID").is_ok() {
34        return CallerKind::Cline;
35    }
36    if std::env::var("AMP_SESSION_ID").is_ok() {
37        return CallerKind::Amp;
38    }
39
40    // 3. JSON shape heuristics
41    if let Some(obj) = stdin.as_object() {
42        let has = |key: &str| obj.contains_key(key);
43
44        if has("tool_name") && has("tool_input") {
45            return CallerKind::ClaudeCode;
46        }
47        if has("type") && has("toolCall") {
48            return CallerKind::Cursor;
49        }
50        if has("event") && has("parameters") {
51            return CallerKind::Windsurf;
52        }
53        // Cline uses toolName (not toolCall)
54        if has("type") && has("toolName") && !has("toolCall") {
55            return CallerKind::Cline;
56        }
57        if has("kind") {
58            return CallerKind::Amp;
59        }
60    }
61
62    CallerKind::Unknown
63}
64
65// ---------------------------------------------------------------------------
66// Tests
67// ---------------------------------------------------------------------------
68
69#[cfg(test)]
70mod tests {
71    use super::detect_caller;
72    use crate::CallerKind;
73
74    const AGENT_ENV_VARS: &[&str] = &[
75        "POLYHOOK_CALLER",
76        "CLAUDE_CODE_VERSION",
77        "CURSOR_SESSION_ID",
78        "WINDSURF_SESSION_ID",
79        "CLINE_SESSION_ID",
80        "AMP_SESSION_ID",
81    ];
82
83    fn with_clean_env<F: FnOnce()>(f: F) {
84        let vars: Vec<(&str, Option<&str>)> = AGENT_ENV_VARS.iter().map(|k| (*k, None)).collect();
85        temp_env::with_vars(vars, f);
86    }
87
88    #[test]
89    fn claude_code_version_env_var_detected() {
90        let val = serde_json::json!({});
91        with_clean_env(|| {
92            temp_env::with_var("CLAUDE_CODE_VERSION", Some("1.0.0"), || {
93                assert_eq!(detect_caller(&val), CallerKind::ClaudeCode);
94            });
95        });
96    }
97
98    #[test]
99    fn cursor_session_id_env_var_detected() {
100        let val = serde_json::json!({});
101        with_clean_env(|| {
102            temp_env::with_var("CURSOR_SESSION_ID", Some("cursor-sess-abc"), || {
103                assert_eq!(detect_caller(&val), CallerKind::Cursor);
104            });
105        });
106    }
107
108    #[test]
109    fn windsurf_session_id_env_var_detected() {
110        let val = serde_json::json!({});
111        with_clean_env(|| {
112            temp_env::with_var("WINDSURF_SESSION_ID", Some("ws-sess-xyz"), || {
113                assert_eq!(detect_caller(&val), CallerKind::Windsurf);
114            });
115        });
116    }
117
118    #[test]
119    fn cline_session_id_env_var_detected() {
120        let val = serde_json::json!({});
121        with_clean_env(|| {
122            temp_env::with_var("CLINE_SESSION_ID", Some("cline-sess-999"), || {
123                assert_eq!(detect_caller(&val), CallerKind::Cline);
124            });
125        });
126    }
127
128    #[test]
129    fn amp_session_id_env_var_detected() {
130        let val = serde_json::json!({});
131        with_clean_env(|| {
132            temp_env::with_var("AMP_SESSION_ID", Some("amp-sess-000"), || {
133                assert_eq!(detect_caller(&val), CallerKind::Amp);
134            });
135        });
136    }
137
138    #[test]
139    fn polyhook_caller_garbage_falls_through_to_heuristics() {
140        let val = serde_json::json!({"tool_name": "Bash", "tool_input": {}, "session_id": "s1"});
141        with_clean_env(|| {
142            temp_env::with_var("POLYHOOK_CALLER", Some("garbage_value_xyz"), || {
143                assert_eq!(detect_caller(&val), CallerKind::ClaudeCode);
144            });
145        });
146    }
147
148    #[test]
149    fn polyhook_caller_garbage_with_no_heuristic_match_returns_unknown() {
150        let val = serde_json::json!({"some_random_key": "some_value"});
151        with_clean_env(|| {
152            temp_env::with_var("POLYHOOK_CALLER", Some("not_a_known_caller"), || {
153                assert_eq!(detect_caller(&val), CallerKind::Unknown);
154            });
155        });
156    }
157
158    #[test]
159    fn unknown_json_shape_returns_unknown() {
160        let val = serde_json::json!({"foo": "bar"});
161        with_clean_env(|| {
162            assert_eq!(detect_caller(&val), CallerKind::Unknown);
163        });
164    }
165
166    #[test]
167    fn empty_object_returns_unknown() {
168        let val = serde_json::json!({});
169        with_clean_env(|| {
170            assert_eq!(detect_caller(&val), CallerKind::Unknown);
171        });
172    }
173
174    #[test]
175    fn non_object_json_returns_unknown() {
176        let val = serde_json::json!(["tool_name", "tool_input"]);
177        with_clean_env(|| {
178            assert_eq!(detect_caller(&val), CallerKind::Unknown);
179        });
180    }
181
182    #[test]
183    fn polyhook_caller_claudecode_alias_detected() {
184        let val = serde_json::json!({});
185        with_clean_env(|| {
186            temp_env::with_var("POLYHOOK_CALLER", Some("claudecode"), || {
187                assert_eq!(detect_caller(&val), CallerKind::ClaudeCode);
188            });
189        });
190    }
191}