1use async_trait::async_trait;
8use serde_json::json;
9use tracing::info;
10
11use super::policy::{Check, StageSet, Verdict};
12use super::{ConnectionContext, IdentifierContext};
13use crate::script_hook::{ScriptError, ScriptHook, ScriptStdin};
14
15#[derive(Debug, Clone)]
17pub struct Settings {
18 pub script_path: String,
19 pub timeout_ms: u64,
20 pub pass_stdin: bool,
21 pub args: Vec<String>,
22}
23
24impl Default for Settings {
25 fn default() -> Self {
26 Self {
27 script_path: String::new(),
28 timeout_ms: 5000,
29 pass_stdin: true,
30 args: Vec::new(),
31 }
32 }
33}
34
35#[derive(Debug)]
37pub struct CustomScriptFilter {
38 hook: ScriptHook,
39 pass_stdin: bool,
40 check_name: String,
42}
43
44impl CustomScriptFilter {
45 pub fn from_settings(name: &str, settings: &Settings) -> anyhow::Result<Self> {
47 let Some(hook) =
48 ScriptHook::new(&settings.script_path, &settings.args, settings.timeout_ms)
49 else {
50 anyhow::bail!(
51 "filter.check.{name}.script_path is empty; provide a path to an \
52 executable script or drop the check"
53 );
54 };
55
56 info!(
57 event = "filter_custom_loaded",
58 outcome = "success",
59 check = name,
60 script_path = %hook.path().display(),
61 timeout_ms = settings.timeout_ms,
62 pass_stdin = settings.pass_stdin,
63 args = ?settings.args,
64 );
65
66 Ok(Self {
67 hook,
68 pass_stdin: settings.pass_stdin,
69 check_name: name.to_string(),
70 })
71 }
72
73 async fn run_script(&self, envs: &[(&str, &str)], payload: &serde_json::Value) -> Verdict {
82 let stdin = if self.pass_stdin {
83 ScriptStdin::Json(payload)
84 } else {
85 ScriptStdin::Null
86 };
87
88 let outcome = match self.hook.run(envs, stdin).await {
89 Ok(outcome) => outcome,
90 Err(
91 error @ (ScriptError::Spawn { .. }
92 | ScriptError::Serialize(_)
93 | ScriptError::Wait(_)
94 | ScriptError::Timeout(_)),
95 ) => return Verdict::Undecided(format!("custom filter {error}")),
96 };
97
98 if outcome.output.status.success() {
99 Verdict::Pass
100 } else {
101 Verdict::Fail(ScriptHook::detail(&outcome, "custom filter script"))
102 }
103 }
104}
105
106#[async_trait]
107impl Check for CustomScriptFilter {
108 fn kind(&self) -> &'static str {
109 "custom"
110 }
111
112 fn stages(&self) -> StageSet {
113 StageSet::both()
114 }
115
116 async fn check_connection(&self, context: &ConnectionContext<'_>) -> Verdict {
117 let client_ip_str = context
118 .client_ip
119 .map(|ip| super::canonical(ip).to_string())
120 .unwrap_or_default();
121 let envs = [
122 ("ACME_FILTER_HOOK", "connection"),
123 ("ACME_FILTER_CHECK_NAME", self.check_name.as_str()),
124 ("ACME_FILTER_CLIENT_IP", client_ip_str.as_str()),
125 ("ACME_FILTER_METHOD", context.method.as_str()),
126 ("ACME_FILTER_PATH", context.path),
127 ];
128
129 let payload = json!({
130 "hook": "connection",
131 "check": self.check_name,
132 "client_ip": if client_ip_str.is_empty() { None } else { Some(&client_ip_str) },
133 "method": context.method.as_str(),
134 "path": context.path,
135 });
136
137 self.run_script(&envs, &payload).await
138 }
139
140 async fn check_identifiers(&self, context: &IdentifierContext<'_>) -> Verdict {
141 let client_ip_str = context
142 .client_ip
143 .map(|ip| super::canonical(ip).to_string())
144 .unwrap_or_default();
145 let identifiers_vec: Vec<String> = context
146 .identifiers
147 .iter()
148 .map(|identifier| identifier.value.clone())
149 .collect();
150 let identifiers_str = identifiers_vec.join(",");
151
152 let envs = [
153 ("ACME_FILTER_HOOK", "identifiers"),
154 ("ACME_FILTER_CHECK_NAME", self.check_name.as_str()),
155 ("ACME_FILTER_CLIENT_IP", client_ip_str.as_str()),
156 ("ACME_FILTER_ACCOUNT_ID", context.account_id),
157 ("ACME_FILTER_STAGE", context.stage.as_str()),
158 ("ACME_FILTER_IDENTIFIERS", identifiers_str.as_str()),
159 ];
160
161 let payload = json!({
162 "hook": "identifiers",
163 "check": self.check_name,
164 "client_ip": if client_ip_str.is_empty() { None } else { Some(&client_ip_str) },
165 "account_id": context.account_id,
166 "stage": context.stage.as_str(),
167 "identifiers": context.identifiers,
168 });
169
170 self.run_script(&envs, &payload).await
171 }
172}
173
174#[cfg(test)]
175mod tests {
176 use super::*;
177 use crate::filter::IdentifierStage;
178 use crate::sqlite::order::Identifier;
179 use crate::testutil::TempDir;
180 use axum::http::Method;
181 use std::time::Duration;
182
183 fn write_script(dir: &TempDir, name: &str, body: &str) -> Settings {
189 let script_path = crate::testutil::write_script(dir, name, body);
190 Settings {
191 script_path: script_path.to_str().unwrap().to_string(),
192 ..Default::default()
193 }
194 }
195
196 #[test]
197 fn missing_script_path_bails() {
198 let cfg = Settings {
199 script_path: " ".to_string(),
200 ..Default::default()
201 };
202 assert!(CustomScriptFilter::from_settings("hook", &cfg).is_err());
203 }
204
205 #[tokio::test]
210 async fn a_script_that_cannot_be_spawned_is_internal_not_denied() {
211 let filter = CustomScriptFilter::from_settings(
212 "hook",
213 &Settings {
214 script_path: "/nonexistent/filter.sh".to_string(),
215 ..Default::default()
216 },
217 )
218 .unwrap();
219
220 let ctx = ConnectionContext {
221 client_ip: "127.0.0.1".parse().ok(),
222 method: &Method::GET,
223 path: "/newOrder",
224 };
225 match filter.check_connection(&ctx).await {
226 Verdict::Undecided(detail) => {
227 assert!(
228 detail.contains("failed to spawn script")
229 && detail.contains("/nonexistent/filter.sh"),
230 "{detail}"
231 )
232 }
233 other => panic!("expected Undecided, got {other:?}"),
234 }
235
236 let identifiers = vec![Identifier::dns("example.com")];
237 let ctx = IdentifierContext {
238 client_ip: "127.0.0.1".parse().ok(),
239 account_id: "acc_1",
240 stage: IdentifierStage::NewOrder,
241 identifiers: &identifiers,
242
243 eab: None,
244 };
245 assert!(matches!(
246 filter.check_identifiers(&ctx).await,
247 Verdict::Undecided(_)
248 ));
249 }
250
251 #[tokio::test]
252 async fn passing_script_allows() {
253 let dir = TempDir::new("filter-custom");
254 let cfg = write_script(&dir, "pass.sh", "#!/bin/sh\nexit 0\n");
255 let filter = CustomScriptFilter::from_settings("hook", &cfg).unwrap();
256
257 let ctx = ConnectionContext {
258 client_ip: "127.0.0.1".parse().ok(),
259 method: &Method::GET,
260 path: "/health",
261 };
262 assert_eq!(filter.check_connection(&ctx).await, Verdict::Pass);
263 }
264
265 #[tokio::test]
266 async fn failing_script_denies() {
267 let dir = TempDir::new("filter-custom");
268 let cfg = write_script(
269 &dir,
270 "fail.sh",
271 "#!/bin/sh\necho \"custom denial\"\nexit 1\n",
272 );
273 let filter = CustomScriptFilter::from_settings("hook", &cfg).unwrap();
274
275 let ctx = ConnectionContext {
276 client_ip: "127.0.0.1".parse().ok(),
277 method: &Method::POST,
278 path: "/acme/new-order",
279 };
280 let res = filter.check_connection(&ctx).await;
281 match res {
282 Verdict::Fail(detail) => {
292 assert!(detail.starts_with("custom denial"), "{detail}");
293 }
294 other => panic!("expected Fail, got {other:?}"),
295 }
296 }
297
298 #[tokio::test]
299 async fn script_receives_env_and_stdin() {
300 let dir = TempDir::new("filter-custom");
301 let script_content = r#"#!/bin/sh
302if [ "$ACME_FILTER_HOOK" != "identifiers" ]; then
303 echo "wrong hook: $ACME_FILTER_HOOK"
304 exit 1
305fi
306if [ "$ACME_FILTER_IDENTIFIERS" != "example.com" ]; then
307 echo "wrong identifiers: $ACME_FILTER_IDENTIFIERS"
308 exit 1
309fi
310exit 0
311"#;
312 let cfg = write_script(&dir, "check_env.sh", script_content);
313 let filter = CustomScriptFilter::from_settings("hook", &cfg).unwrap();
314
315 let identifiers = vec![Identifier::dns("example.com")];
316 let ctx = IdentifierContext {
317 client_ip: "10.0.0.1".parse().ok(),
318 account_id: "acc_123",
319 stage: IdentifierStage::NewOrder,
320 identifiers: &identifiers,
321
322 eab: None,
323 };
324
325 assert_eq!(filter.check_identifiers(&ctx).await, Verdict::Pass);
326 }
327
328 #[tokio::test]
329 async fn script_timeout_returns_internal() {
330 let dir = TempDir::new("filter-custom");
331 let cfg = Settings {
332 timeout_ms: 100,
333 ..write_script(&dir, "sleep.sh", "#!/bin/sh\nsleep 2\nexit 0\n")
334 };
335 let filter = CustomScriptFilter::from_settings("hook", &cfg).unwrap();
336
337 let ctx = ConnectionContext {
338 client_ip: None,
339 method: &Method::GET,
340 path: "/health",
341 };
342
343 let res = filter.check_connection(&ctx).await;
344 match res {
345 Verdict::Undecided(detail) => assert!(detail.contains("timed out")),
346 other => panic!("expected Undecided on timeout, got {other:?}"),
347 }
348 }
349
350 #[tokio::test]
359 async fn the_script_does_not_inherit_the_server_environment() {
360 assert!(
361 std::env::var_os("CARGO_MANIFEST_DIR").is_some(),
362 "the canary must exist in the parent, otherwise the test proves nothing"
363 );
364
365 let dir = TempDir::new("filter-custom");
366 let cfg = write_script(
367 &dir,
368 "env_leak.sh",
369 r#"#!/bin/sh
370if [ -n "$CARGO_MANIFEST_DIR" ]; then
371 echo "inherited CARGO_MANIFEST_DIR=$CARGO_MANIFEST_DIR"
372 exit 1
373fi
374# The minimal PATH, however, must be provided.
375if [ -z "$PATH" ]; then
376 echo "no PATH"
377 exit 1
378fi
379# And the documented filter variables too.
380if [ "$ACME_FILTER_HOOK" != "connection" ]; then
381 echo "missing ACME_FILTER_HOOK"
382 exit 1
383fi
384exit 0
385"#,
386 );
387 let filter = CustomScriptFilter::from_settings("hook", &cfg).unwrap();
388
389 let ctx = ConnectionContext {
390 client_ip: "127.0.0.1".parse().ok(),
391 method: &Method::GET,
392 path: "/newOrder",
393 };
394 assert_eq!(filter.check_connection(&ctx).await, Verdict::Pass);
395 }
396
397 #[tokio::test]
401 async fn a_timed_out_script_is_killed_rather_than_left_running() {
402 let dir = TempDir::new("filter-custom");
403 let marker = dir.path().join("survived");
404 let cfg = Settings {
405 timeout_ms: 50,
406 ..write_script(
407 &dir,
408 "slow.sh",
409 &format!("#!/bin/sh\nsleep 1\ntouch {}\n", marker.to_str().unwrap()),
410 )
411 };
412 let filter = CustomScriptFilter::from_settings("hook", &cfg).unwrap();
413
414 let ctx = ConnectionContext {
415 client_ip: None,
416 method: &Method::GET,
417 path: "/newOrder",
418 };
419 match filter.check_connection(&ctx).await {
420 Verdict::Undecided(detail) => assert!(detail.contains("timed out")),
421 other => panic!("expected Undecided on timeout, got {other:?}"),
422 }
423
424 tokio::time::sleep(Duration::from_millis(1_800)).await;
427 assert!(
428 !marker.exists(),
429 "the script survived the timeout and continued executing"
430 );
431 }
432}