1use std::time::Duration;
2
3use serde::{Deserialize, Serialize};
4
5pub const DEFAULT_SHELL_TIMEOUT_SECS: u32 = 120;
7
8pub const DEFAULT_SHELL_TIMEOUT_MAX_SECS: u32 = 600;
10
11pub const LONGEST_SHELL_TIMEOUT_SECS: u32 = 86_400;
13
14#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
16pub enum ShellTimeoutsError {
17 #[error("timeouts must be whole seconds from 1 to {LONGEST_SHELL_TIMEOUT_SECS}")]
18 OutOfRange,
19 #[error("the default timeout ({default} s) must not exceed the maximum ({max} s)")]
20 DefaultAboveMax { default: u32, max: u32 },
21}
22
23#[derive(Serialize, Deserialize)]
24struct RawShellTimeouts {
25 default_secs: u32,
26 max_secs: u32,
27}
28
29#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
31#[serde(try_from = "RawShellTimeouts", into = "RawShellTimeouts")]
32pub struct ShellTimeouts {
33 default_secs: u32,
34 max_secs: u32,
35}
36
37impl ShellTimeouts {
38 pub const DEFAULT: Self = Self {
39 default_secs: DEFAULT_SHELL_TIMEOUT_SECS,
40 max_secs: DEFAULT_SHELL_TIMEOUT_MAX_SECS,
41 };
42
43 pub fn new(default_secs: u32, max_secs: u32) -> Result<Self, ShellTimeoutsError> {
49 let range = 1..=LONGEST_SHELL_TIMEOUT_SECS;
50 if !range.contains(&default_secs) || !range.contains(&max_secs) {
51 Err(ShellTimeoutsError::OutOfRange)
52 } else if default_secs > max_secs {
53 Err(ShellTimeoutsError::DefaultAboveMax {
54 default: default_secs,
55 max: max_secs,
56 })
57 } else {
58 Ok(Self {
59 default_secs,
60 max_secs,
61 })
62 }
63 }
64
65 #[must_use]
67 pub fn effective(self, requested_secs: Option<u64>) -> Duration {
68 let secs = requested_secs
69 .unwrap_or(self.default_secs.into())
70 .clamp(1, self.max_secs.into());
71 Duration::from_secs(secs)
72 }
73
74 #[must_use]
75 pub fn max(self) -> Duration {
76 Duration::from_secs(self.max_secs.into())
77 }
78}
79
80impl Default for ShellTimeouts {
81 fn default() -> Self {
82 Self::DEFAULT
83 }
84}
85
86impl TryFrom<RawShellTimeouts> for ShellTimeouts {
87 type Error = ShellTimeoutsError;
88
89 fn try_from(raw: RawShellTimeouts) -> Result<Self, Self::Error> {
90 Self::new(raw.default_secs, raw.max_secs)
91 }
92}
93
94impl From<ShellTimeouts> for RawShellTimeouts {
95 fn from(timeouts: ShellTimeouts) -> Self {
96 Self {
97 default_secs: timeouts.default_secs,
98 max_secs: timeouts.max_secs,
99 }
100 }
101}
102
103#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
105pub struct ShellRequest {
106 pub command: String,
107 pub timeout_secs: Option<u64>,
109}
110
111#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
113#[serde(tag = "kind", rename_all = "snake_case")]
114pub enum ShellOutcome {
115 Exited {
116 code: i32,
117 },
118 Signaled {
119 signal: i32,
120 },
121 Cancelled,
124 TimedOut {
125 after_secs: u64,
126 },
127}
128
129#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
131pub struct ShellReply {
132 pub outcome: ShellOutcome,
133 pub duration_ms: u64,
134 pub stdout: String,
136 pub stderr: String,
138}
139
140#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
142pub struct SetCwdRequest {
143 pub path: String,
144}
145
146#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
148pub struct SetCwdReply {
149 pub cwd: String,
151}
152
153#[cfg(test)]
154mod tests {
155 use super::*;
156
157 #[test]
158 fn agent_timeouts_are_held_between_one_second_and_the_maximum() {
159 let timeouts = ShellTimeouts::new(120, 600).unwrap();
160 let secs = |requested| timeouts.effective(requested).as_secs();
161 assert_eq!(secs(None), 120);
162 assert_eq!(secs(Some(30)), 30);
163 assert_eq!(secs(Some(0)), 1);
164 assert_eq!(secs(Some(100_000)), 600);
165 }
166
167 #[test]
168 fn timeout_settings_must_be_in_range_and_the_default_within_the_maximum() {
169 assert!(ShellTimeouts::new(600, 600).is_ok());
170 assert_eq!(
171 ShellTimeouts::new(0, 600),
172 Err(ShellTimeoutsError::OutOfRange)
173 );
174 assert_eq!(
175 ShellTimeouts::new(1, LONGEST_SHELL_TIMEOUT_SECS + 1),
176 Err(ShellTimeoutsError::OutOfRange)
177 );
178 assert_eq!(
179 ShellTimeouts::new(601, 600),
180 Err(ShellTimeoutsError::DefaultAboveMax {
181 default: 601,
182 max: 600
183 })
184 );
185 let bad = serde_json::json!({ "default_secs": 700, "max_secs": 600 });
186 assert!(serde_json::from_value::<ShellTimeouts>(bad).is_err());
187 }
188}