1use std::collections::hash_map::RandomState;
31use std::hash::{BuildHasher, Hasher};
32use std::io::Write;
33use std::path::{Path, PathBuf};
34use std::sync::atomic::{AtomicU64, Ordering};
35use std::sync::Mutex;
36use std::time::{SystemTime, UNIX_EPOCH};
37
38use crate::admission::lock;
39
40#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
42pub enum ExecutionState {
43 Running,
45 Exited,
47 Unknown,
51}
52
53impl ExecutionState {
54 pub fn label(self) -> &'static str {
55 match self {
56 Self::Running => "running",
57 Self::Exited => "exited",
58 Self::Unknown => "unknown",
59 }
60 }
61 fn parse(s: &str) -> Option<Self> {
62 match s {
63 "running" => Some(Self::Running),
64 "exited" => Some(Self::Exited),
65 "unknown" => Some(Self::Unknown),
66 _ => None,
67 }
68 }
69}
70
71#[derive(Debug, Clone, PartialEq, Eq, Hash)]
74pub struct ExecutionIdentity(String);
75
76impl ExecutionIdentity {
77 pub fn generate() -> Self {
79 static COUNTER: AtomicU64 = AtomicU64::new(0);
80 let nanos = SystemTime::now()
81 .duration_since(UNIX_EPOCH)
82 .map(|d| d.as_nanos())
83 .unwrap_or(0);
84 let n = COUNTER.fetch_add(1, Ordering::Relaxed);
85 let word = |salt: u64| {
86 let mut h = RandomState::new().build_hasher();
87 h.write_u128(nanos);
88 h.write_u32(std::process::id());
89 h.write_u64(n);
90 h.write_u64(salt);
91 h.finish()
92 };
93 Self(format!("{:016x}{:016x}", word(1), word(2)))
94 }
95 pub fn new(id: impl Into<String>) -> Result<Self, StoreError> {
97 let id = id.into();
98 let ok = !id.is_empty()
99 && id.len() <= 128
100 && id
101 .bytes()
102 .all(|b| b.is_ascii_alphanumeric() || b"._:-".contains(&b));
103 if ok {
104 Ok(Self(id))
105 } else {
106 Err(StoreError::Invalid(format!(
107 "execution identity {id:?} must be 1-128 characters of [A-Za-z0-9._:-]"
108 )))
109 }
110 }
111 pub fn as_str(&self) -> &str {
112 &self.0
113 }
114}
115
116impl std::fmt::Display for ExecutionIdentity {
117 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
118 f.write_str(&self.0)
119 }
120}
121
122#[derive(Debug, Clone, PartialEq, Eq)]
124pub struct LeaseRecord {
125 pub name: String,
126 pub execution_id: ExecutionIdentity,
127 pub execution_state: ExecutionState,
128 pub epoch: u64,
130 pub holder_epoch: Option<u64>,
132 pub next_input_sequence: u64,
133 pub acked_input_sequence: u64,
135 pub unknown_input: Option<(u64, u64)>,
139}
140
141impl LeaseRecord {
142 fn validate(&self) -> Result<(), StoreError> {
143 let bad = |m: &str| Err(StoreError::Corrupt(m.into()));
144 if self.name.contains(['\n', '\r']) {
145 return bad("lease name contains a line break");
146 }
147 if self.acked_input_sequence > self.next_input_sequence {
148 return bad("acknowledged sequence is ahead of the next sequence");
149 }
150 if let Some(h) = self.holder_epoch {
151 if h != self.epoch || h == 0 {
152 return bad("holder epoch does not match the current epoch");
153 }
154 }
155 if let Some(r) = self.unknown_input {
156 if r != (self.acked_input_sequence, self.next_input_sequence) || r.0 >= r.1 {
157 return bad("unknown input range does not match the unacknowledged range");
158 }
159 }
160 Ok(())
161 }
162}
163
164#[derive(Debug, Clone, PartialEq, Eq)]
166pub enum StoreError {
167 Io(String),
168 Corrupt(String),
170 Invalid(String),
172}
173
174impl std::fmt::Display for StoreError {
175 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
176 match self {
177 Self::Io(m) => write!(f, "lease store i/o: {m}"),
178 Self::Corrupt(m) => write!(f, "lease store record is corrupt: {m}"),
179 Self::Invalid(m) => write!(f, "lease store: {m}"),
180 }
181 }
182}
183
184impl std::error::Error for StoreError {}
185
186pub trait LeaseStore: Send + Sync {
190 fn load(&self) -> Result<Option<LeaseRecord>, StoreError>;
192 fn save(&self, record: &LeaseRecord) -> Result<(), StoreError>;
193}
194
195#[derive(Debug, Default)]
197pub struct MemoryLeaseStore(Mutex<Option<LeaseRecord>>);
198
199impl MemoryLeaseStore {
200 pub fn new() -> Self {
201 Self::default()
202 }
203 pub fn record(&self) -> Option<LeaseRecord> {
204 lock(&self.0).clone()
205 }
206}
207
208impl LeaseStore for MemoryLeaseStore {
209 fn load(&self) -> Result<Option<LeaseRecord>, StoreError> {
210 Ok(lock(&self.0).clone())
211 }
212 fn save(&self, record: &LeaseRecord) -> Result<(), StoreError> {
213 record.validate()?;
214 *lock(&self.0) = Some(record.clone());
215 Ok(())
216 }
217}
218
219const HEADER: &str = "rightkit-control-lease 1";
220
221#[derive(Debug, Clone)]
230pub struct FileLeaseStore {
231 path: PathBuf,
232}
233
234impl FileLeaseStore {
235 pub fn new(path: impl Into<PathBuf>) -> Self {
236 Self { path: path.into() }
237 }
238 pub fn path(&self) -> &Path {
239 &self.path
240 }
241 pub fn temp_path(&self) -> PathBuf {
243 let mut p = self.path.clone().into_os_string();
244 p.push(".tmp");
245 PathBuf::from(p)
246 }
247}
248
249impl LeaseStore for FileLeaseStore {
250 fn load(&self) -> Result<Option<LeaseRecord>, StoreError> {
251 match std::fs::read_to_string(&self.path) {
252 Ok(text) => decode(&text).map(Some),
253 Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(None),
254 Err(e) if e.kind() == std::io::ErrorKind::InvalidData => {
255 Err(StoreError::Corrupt("record is not UTF-8".into()))
256 }
257 Err(e) => Err(StoreError::Io(format!("{}: {e}", self.path.display()))),
258 }
259 }
260
261 fn save(&self, record: &LeaseRecord) -> Result<(), StoreError> {
262 record.validate().map_err(|e| match e {
263 StoreError::Corrupt(m) => StoreError::Invalid(m),
264 other => other,
265 })?;
266 let body = encode(record);
267 let tmp = self.temp_path();
268 let io = |what: &str, e: std::io::Error| StoreError::Io(format!("{what}: {e}"));
269 if let Some(dir) = self.path.parent().filter(|d| !d.as_os_str().is_empty()) {
270 std::fs::create_dir_all(dir).map_err(|e| io("create lease directory", e))?;
271 }
272 let mut opts = std::fs::OpenOptions::new();
273 opts.write(true).create(true).truncate(true);
274 #[cfg(unix)]
275 {
276 use std::os::unix::fs::OpenOptionsExt;
277 opts.mode(0o600);
278 }
279 let mut f = opts.open(&tmp).map_err(|e| io("open temp record", e))?;
280 f.write_all(body.as_bytes())
281 .map_err(|e| io("write temp record", e))?;
282 f.sync_all().map_err(|e| io("fsync temp record", e))?;
283 drop(f);
284 std::fs::rename(&tmp, &self.path).map_err(|e| io("rename temp record", e))?;
285 #[cfg(unix)]
286 if let Some(dir) = self.path.parent().filter(|d| !d.as_os_str().is_empty()) {
287 std::fs::File::open(dir)
288 .and_then(|d| d.sync_all())
289 .map_err(|e| io("fsync lease directory", e))?;
290 }
291 Ok(())
292 }
293}
294
295fn fnv1a(text: &str) -> u64 {
296 text.bytes().fold(0xcbf2_9ce4_8422_2325, |h, b| {
297 (h ^ u64::from(b)).wrapping_mul(0x0000_0100_0000_01b3)
298 })
299}
300
301fn opt(v: Option<u64>) -> String {
302 v.map_or_else(|| "-".into(), |n| n.to_string())
303}
304
305fn encode(r: &LeaseRecord) -> String {
306 let unknown = r
307 .unknown_input
308 .map_or_else(|| "-".into(), |(a, b)| format!("{a}..{b}"));
309 let body = format!(
310 "{HEADER}\nname={}\nexecution_id={}\nexecution_state={}\nepoch={}\nholder_epoch={}\nnext_input_sequence={}\nacked_input_sequence={}\nunknown_input={unknown}\n",
311 r.name,
312 r.execution_id,
313 r.execution_state.label(),
314 r.epoch,
315 opt(r.holder_epoch),
316 r.next_input_sequence,
317 r.acked_input_sequence,
318 );
319 let sum = fnv1a(&body);
320 format!("{body}checksum={sum:016x}\n")
321}
322
323fn decode(text: &str) -> Result<LeaseRecord, StoreError> {
324 let corrupt = |m: &str| StoreError::Corrupt(m.to_string());
325 let split = text
326 .rfind("checksum=")
327 .ok_or_else(|| corrupt("missing checksum"))?;
328 let (body, tail) = text.split_at(split);
329 let sum = tail
330 .strip_prefix("checksum=")
331 .and_then(|t| t.strip_suffix('\n'))
332 .and_then(|t| u64::from_str_radix(t, 16).ok())
333 .ok_or_else(|| corrupt("malformed checksum"))?;
334 if sum != fnv1a(body) {
335 return Err(corrupt("checksum mismatch"));
336 }
337 let mut lines = body.lines();
338 if lines.next() != Some(HEADER) {
339 return Err(corrupt("unknown header or version"));
340 }
341 let mut field = |key: &str| -> Result<String, StoreError> {
342 lines
343 .next()
344 .and_then(|l| l.strip_prefix(key))
345 .and_then(|l| l.strip_prefix('='))
346 .map(str::to_string)
347 .ok_or_else(|| StoreError::Corrupt(format!("missing field {key}")))
348 };
349 let num = |v: String, key: &str| {
350 v.parse::<u64>()
351 .map_err(|_| StoreError::Corrupt(format!("field {key} is not a number")))
352 };
353 let opt_num = |v: String, key: &str| {
354 if v == "-" {
355 Ok(None)
356 } else {
357 num(v, key).map(Some)
358 }
359 };
360 let name = field("name")?;
361 let execution_id = ExecutionIdentity::new(field("execution_id")?)
362 .map_err(|_| corrupt("invalid execution identity"))?;
363 let execution_state = ExecutionState::parse(&field("execution_state")?)
364 .ok_or_else(|| corrupt("invalid execution state"))?;
365 let epoch = num(field("epoch")?, "epoch")?;
366 let holder_epoch = opt_num(field("holder_epoch")?, "holder_epoch")?;
367 let next_input_sequence = num(field("next_input_sequence")?, "next_input_sequence")?;
368 let acked_input_sequence = num(field("acked_input_sequence")?, "acked_input_sequence")?;
369 let unknown = field("unknown_input")?;
370 let unknown_input = if unknown == "-" {
371 None
372 } else {
373 let (a, b) = unknown
374 .split_once("..")
375 .ok_or_else(|| corrupt("malformed unknown input range"))?;
376 Some((
377 num(a.into(), "unknown_input")?,
378 num(b.into(), "unknown_input")?,
379 ))
380 };
381 if lines.next().is_some() {
382 return Err(corrupt("unexpected trailing fields"));
383 }
384 let record = LeaseRecord {
385 name,
386 execution_id,
387 execution_state,
388 epoch,
389 holder_epoch,
390 next_input_sequence,
391 acked_input_sequence,
392 unknown_input,
393 };
394 record.validate()?;
395 Ok(record)
396}
397
398#[cfg(test)]
399mod tests {
400 use super::*;
401
402 fn sample() -> LeaseRecord {
403 LeaseRecord {
404 name: "pty:1".into(),
405 execution_id: ExecutionIdentity::generate(),
406 execution_state: ExecutionState::Running,
407 epoch: 3,
408 holder_epoch: Some(3),
409 next_input_sequence: 7,
410 acked_input_sequence: 5,
411 unknown_input: Some((5, 7)),
412 }
413 }
414
415 #[test]
416 fn round_trip_and_checksum() {
417 let r = sample();
418 let text = encode(&r);
419 assert_eq!(decode(&text).unwrap(), r);
420 let tampered = text.replace("epoch=3\n", "epoch=1\n");
421 assert!(matches!(decode(&tampered), Err(StoreError::Corrupt(_))));
422 let truncated = &text[..text.len() / 2];
423 assert!(matches!(decode(truncated), Err(StoreError::Corrupt(_))));
424 }
425
426 #[test]
427 fn identities_are_unique_and_validated() {
428 assert_ne!(ExecutionIdentity::generate(), ExecutionIdentity::generate());
429 assert_eq!(ExecutionIdentity::generate().as_str().len(), 32);
430 assert!(ExecutionIdentity::new("job:42-a").is_ok());
431 assert!(ExecutionIdentity::new("bad id\n").is_err());
432 assert!(ExecutionIdentity::new("").is_err());
433 }
434}