Skip to main content

vv_agent/runtime/stores/
redis.rs

1//! Redis checkpoint store.
2
3use std::sync::Mutex;
4use std::time::Duration;
5
6use redis::{Commands, Connection, Pipeline};
7use serde_json::Value;
8use sha2::{Digest, Sha256};
9
10use crate::checkpoint::{CheckpointError, CheckpointResult, ClaimMode, EventCursor};
11use crate::runtime::checkpoint_codec::{checkpoint_from_json, checkpoint_to_json};
12use crate::runtime::state::{
13    apply_claim, claim_candidate, prepare_ack, prepare_commit, prepare_event_delivery,
14    prepare_finalize, prepare_finalize_claimed, prepare_progress, prepare_suspend, Checkpoint,
15    CheckpointStore,
16};
17
18const KEY_PREFIX: &str = "vv-agent:checkpoint:";
19const LEASE_SUFFIX: &str = ":lease";
20const IO_TIMEOUT: Duration = Duration::from_secs(1);
21const TRANSACTION_MAX_ATTEMPTS: usize = 8;
22const MAX_EXTENSION_STATE_BYTES: u64 = crate::checkpoint::MAX_WIRE_INTEGER;
23
24pub struct RedisCheckpointStore {
25    connection: Mutex<Connection>,
26    redis_url: String,
27}
28
29impl std::fmt::Debug for RedisCheckpointStore {
30    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
31        formatter
32            .debug_struct("RedisCheckpointStore")
33            .field("redis_url", &self.redis_url)
34            .finish_non_exhaustive()
35    }
36}
37
38impl RedisCheckpointStore {
39    pub fn new(redis_url: impl AsRef<str>) -> CheckpointResult<Self> {
40        let redis_url = redis_url.as_ref().to_string();
41        let client = redis::Client::open(redis_url.as_str()).map_err(redis_error)?;
42        let connection = client
43            .get_connection_with_timeout(IO_TIMEOUT)
44            .map_err(redis_error)?;
45        connection
46            .set_read_timeout(Some(IO_TIMEOUT))
47            .map_err(redis_error)?;
48        connection
49            .set_write_timeout(Some(IO_TIMEOUT))
50            .map_err(redis_error)?;
51        Ok(Self {
52            connection: Mutex::new(connection),
53            redis_url,
54        })
55    }
56
57    pub fn data_key(checkpoint_key: &str) -> String {
58        let digest = Sha256::digest(checkpoint_key.as_bytes());
59        format!("{KEY_PREFIX}{digest:x}")
60    }
61
62    pub fn lease_key(checkpoint_key: &str) -> String {
63        format!("{}{LEASE_SUFFIX}", Self::data_key(checkpoint_key))
64    }
65
66    pub fn redis_url(&self) -> &str {
67        &self.redis_url
68    }
69
70    fn lock(&self) -> CheckpointResult<std::sync::MutexGuard<'_, Connection>> {
71        self.connection.lock().map_err(|_| {
72            CheckpointError::new(
73                "checkpoint_store_lock_poisoned",
74                "Redis store lock poisoned",
75            )
76        })
77    }
78
79    fn load_from_connection(
80        connection: &mut Connection,
81        data_key: &str,
82        lease_key: &str,
83    ) -> CheckpointResult<Option<Checkpoint>> {
84        for _ in 0..TRANSACTION_MAX_ATTEMPTS {
85            let Some(raw) = connection
86                .get::<_, Option<String>>(data_key)
87                .map_err(redis_error)?
88            else {
89                return Ok(None);
90            };
91            let lease = connection
92                .get::<_, Option<u64>>(lease_key)
93                .map_err(redis_error)?;
94            let raw_again = connection
95                .get::<_, Option<String>>(data_key)
96                .map_err(redis_error)?;
97            if raw_again.as_deref() != Some(raw.as_str()) {
98                continue;
99            }
100            return decode_storage(&raw, lease).map(Some);
101        }
102        Err(CheckpointError::new(
103            "checkpoint_store_read_conflict",
104            "Redis checkpoint load could not obtain a stable snapshot",
105        ))
106    }
107
108    fn transaction<T>(
109        &self,
110        data_key: &str,
111        lease_key: &str,
112        operation: impl Fn(&mut Connection, &mut Pipeline) -> CheckpointResult<Option<T>>,
113    ) -> CheckpointResult<T> {
114        let mut connection = self.lock()?;
115        for _ in 0..TRANSACTION_MAX_ATTEMPTS {
116            redis::cmd("WATCH")
117                .arg(data_key)
118                .arg(lease_key)
119                .query::<()>(&mut *connection)
120                .map_err(redis_error)?;
121            let mut pipeline = redis::pipe();
122            pipeline.atomic();
123            match operation(&mut connection, &mut pipeline)? {
124                None => {
125                    redis::cmd("UNWATCH")
126                        .query::<()>(&mut *connection)
127                        .map_err(redis_error)?;
128                    return Err(CheckpointError::new(
129                        "checkpoint_store_conflict",
130                        "checkpoint operation did not match its compare-and-set precondition",
131                    ));
132                }
133                Some(value) => match pipeline.query::<Option<()>>(&mut *connection) {
134                    Ok(Some(())) => {
135                        redis::cmd("UNWATCH")
136                            .query::<()>(&mut *connection)
137                            .map_err(redis_error)?;
138                        return Ok(value);
139                    }
140                    Ok(None) => continue,
141                    Err(error) => return Err(redis_error(error)),
142                },
143            }
144        }
145        Err(CheckpointError::new(
146            "checkpoint_store_transaction_retry_exhausted",
147            "Redis checkpoint transaction retry limit exceeded",
148        ))
149    }
150}
151
152impl CheckpointStore for RedisCheckpointStore {
153    fn create_checkpoint(&self, checkpoint: Checkpoint) -> CheckpointResult<bool> {
154        checkpoint.validate()?;
155        let data_key = Self::data_key(&checkpoint.checkpoint_key);
156        let lease_key = Self::lease_key(&checkpoint.checkpoint_key);
157        let payload = checkpoint_to_json(&checkpoint, MAX_EXTENSION_STATE_BYTES)?;
158        let mut connection = self.lock()?;
159        let created: bool = connection.set_nx(&data_key, payload).map_err(redis_error)?;
160        if created {
161            connection.del::<_, ()>(&lease_key).map_err(redis_error)?;
162        }
163        Ok(created)
164    }
165
166    fn load_checkpoint(&self, checkpoint_key: &str) -> CheckpointResult<Option<Checkpoint>> {
167        let data_key = Self::data_key(checkpoint_key);
168        let lease_key = Self::lease_key(checkpoint_key);
169        let mut connection = self.lock()?;
170        Self::load_from_connection(&mut connection, &data_key, &lease_key)
171    }
172
173    fn claim_checkpoint(
174        &self,
175        checkpoint_key: &str,
176        cycle_index: u64,
177        claim_token: &str,
178        lease_expires_at_ms: u64,
179        now_ms: u64,
180        claim_mode: ClaimMode,
181    ) -> CheckpointResult<Option<Checkpoint>> {
182        if claim_token.trim().is_empty() || lease_expires_at_ms <= now_ms {
183            return Err(CheckpointError::new(
184                "checkpoint_claim_invalid",
185                "claim token must be non-empty and lease must be in the future",
186            ));
187        }
188        let data_key = Self::data_key(checkpoint_key);
189        let lease_key = Self::lease_key(checkpoint_key);
190        let result = self.transaction(&data_key, &lease_key, |connection, pipeline| {
191            let Some(raw) = connection
192                .get::<_, Option<String>>(&data_key)
193                .map_err(redis_error)?
194            else {
195                return Ok(None);
196            };
197            let lease = connection
198                .get::<_, Option<u64>>(&lease_key)
199                .map_err(redis_error)?;
200            let current = decode_storage(&raw, lease)?;
201            if !claim_candidate(&current, cycle_index, now_ms, claim_mode)? {
202                return Ok(None);
203            }
204            let mut claimed = current;
205            apply_claim(
206                &mut claimed,
207                cycle_index,
208                claim_token,
209                lease_expires_at_ms,
210                claim_mode,
211            )?;
212            let payload = checkpoint_to_json(&claimed, MAX_EXTENSION_STATE_BYTES)?;
213            pipeline.set(&data_key, payload).ignore();
214            pipeline.set(&lease_key, lease_expires_at_ms).ignore();
215            Ok(Some(claimed))
216        });
217        match result {
218            Ok(value) => Ok(Some(value)),
219            Err(error) if error.code() == "checkpoint_store_conflict" => Ok(None),
220            Err(error) => Err(error),
221        }
222    }
223
224    fn progress_checkpoint(
225        &self,
226        checkpoint: Checkpoint,
227        claim_token: &str,
228        expected_revision: u64,
229    ) -> CheckpointResult<bool> {
230        self.replace_claimed(
231            checkpoint,
232            claim_token,
233            expected_revision,
234            ReplaceKind::Progress,
235        )
236    }
237
238    fn suspend_checkpoint(
239        &self,
240        checkpoint: Checkpoint,
241        claim_token: &str,
242        expected_revision: u64,
243    ) -> CheckpointResult<bool> {
244        self.replace_claimed(
245            checkpoint,
246            claim_token,
247            expected_revision,
248            ReplaceKind::Suspend,
249        )
250    }
251
252    fn commit_checkpoint(
253        &self,
254        checkpoint: Checkpoint,
255        claim_token: &str,
256        expected_revision: u64,
257    ) -> CheckpointResult<bool> {
258        self.replace_claimed(
259            checkpoint,
260            claim_token,
261            expected_revision,
262            ReplaceKind::Commit,
263        )
264    }
265
266    fn finalize_claimed_checkpoint(
267        &self,
268        checkpoint: Checkpoint,
269        claim_token: &str,
270        expected_revision: u64,
271    ) -> CheckpointResult<bool> {
272        self.replace_claimed(
273            checkpoint,
274            claim_token,
275            expected_revision,
276            ReplaceKind::FinalizeClaimed,
277        )
278    }
279
280    fn finalize_checkpoint(
281        &self,
282        checkpoint: Checkpoint,
283        expected_revision: u64,
284    ) -> CheckpointResult<bool> {
285        let data_key = Self::data_key(&checkpoint.checkpoint_key);
286        let lease_key = Self::lease_key(&checkpoint.checkpoint_key);
287        let result = self.transaction(&data_key, &lease_key, |connection, pipeline| {
288            let Some(raw) = connection
289                .get::<_, Option<String>>(&data_key)
290                .map_err(redis_error)?
291            else {
292                return Ok(None);
293            };
294            let current = decode_storage(
295                &raw,
296                connection
297                    .get::<_, Option<u64>>(&lease_key)
298                    .map_err(redis_error)?,
299            )?;
300            let Some(updated) = prepare_finalize(&current, checkpoint.clone(), expected_revision)?
301            else {
302                return Ok(None);
303            };
304            let payload = checkpoint_to_json(&updated, MAX_EXTENSION_STATE_BYTES)?;
305            pipeline.set(&data_key, payload).ignore();
306            pipeline.del(&lease_key).ignore();
307            Ok(Some(true))
308        });
309        match result {
310            Ok(value) => Ok(value),
311            Err(error) if error.code() == "checkpoint_store_conflict" => Ok(false),
312            Err(error) => Err(error),
313        }
314    }
315
316    fn renew_checkpoint_claim(
317        &self,
318        checkpoint_key: &str,
319        claim_token: &str,
320        lease_expires_at_ms: u64,
321        now_ms: u64,
322    ) -> CheckpointResult<bool> {
323        if claim_token.trim().is_empty() || lease_expires_at_ms <= now_ms {
324            return Err(CheckpointError::new(
325                "checkpoint_claim_invalid",
326                "claim token must be non-empty and lease must be in the future",
327            ));
328        }
329        let data_key = Self::data_key(checkpoint_key);
330        let lease_key = Self::lease_key(checkpoint_key);
331        let result = self.transaction(&data_key, &lease_key, |connection, pipeline| {
332            let Some(raw) = connection
333                .get::<_, Option<String>>(&data_key)
334                .map_err(redis_error)?
335            else {
336                return Ok(None);
337            };
338            let current_lease = connection
339                .get::<_, Option<u64>>(&lease_key)
340                .map_err(redis_error)?;
341            let current = decode_storage(&raw, current_lease)?;
342            if current.claim_token.as_deref() != Some(claim_token)
343                || current
344                    .lease_expires_at_ms
345                    .is_none_or(|expiry| expiry <= now_ms)
346            {
347                return Ok(None);
348            }
349            pipeline.set(&lease_key, lease_expires_at_ms).ignore();
350            Ok(Some(true))
351        });
352        match result {
353            Ok(value) => Ok(value),
354            Err(error) if error.code() == "checkpoint_store_conflict" => Ok(false),
355            Err(error) => Err(error),
356        }
357    }
358
359    fn acknowledge_terminal(
360        &self,
361        checkpoint_key: &str,
362        expected_revision: u64,
363    ) -> CheckpointResult<bool> {
364        let data_key = Self::data_key(checkpoint_key);
365        let lease_key = Self::lease_key(checkpoint_key);
366        let result = self.transaction(&data_key, &lease_key, |connection, pipeline| {
367            let Some(raw) = connection
368                .get::<_, Option<String>>(&data_key)
369                .map_err(redis_error)?
370            else {
371                return Ok(None);
372            };
373            let current = decode_storage(
374                &raw,
375                connection
376                    .get::<_, Option<u64>>(&lease_key)
377                    .map_err(redis_error)?,
378            )?;
379            let Some(updated) = prepare_ack(&current, expected_revision)? else {
380                return Ok(None);
381            };
382            let payload = checkpoint_to_json(&updated, MAX_EXTENSION_STATE_BYTES)?;
383            pipeline.set(&data_key, payload).ignore();
384            pipeline.del(&lease_key).ignore();
385            Ok(Some(true))
386        });
387        match result {
388            Ok(value) => Ok(value),
389            Err(error) if error.code() == "checkpoint_store_conflict" => Ok(false),
390            Err(error) => Err(error),
391        }
392    }
393
394    fn record_event_delivery(
395        &self,
396        checkpoint_key: &str,
397        claim_token: Option<&str>,
398        expected_revision: u64,
399        event_id: &str,
400        payload_digest: &str,
401        cursor: EventCursor,
402    ) -> CheckpointResult<bool> {
403        let data_key = Self::data_key(checkpoint_key);
404        let lease_key = Self::lease_key(checkpoint_key);
405        let result = self.transaction(&data_key, &lease_key, |connection, pipeline| {
406            let Some(raw) = connection
407                .get::<_, Option<String>>(&data_key)
408                .map_err(redis_error)?
409            else {
410                return Ok(None);
411            };
412            let current = decode_storage(
413                &raw,
414                connection
415                    .get::<_, Option<u64>>(&lease_key)
416                    .map_err(redis_error)?,
417            )?;
418            let Some(updated) = prepare_event_delivery(
419                &current,
420                claim_token,
421                expected_revision,
422                event_id,
423                payload_digest,
424                cursor.clone(),
425            )?
426            else {
427                return Ok(None);
428            };
429            let payload = checkpoint_to_json(&updated, MAX_EXTENSION_STATE_BYTES)?;
430            pipeline.set(&data_key, payload).ignore();
431            if updated.claim_token.is_none() {
432                pipeline.del(&lease_key).ignore();
433            }
434            Ok(Some(true))
435        });
436        match result {
437            Ok(value) => Ok(value),
438            Err(error) if error.code() == "checkpoint_store_conflict" => Ok(false),
439            Err(error) => Err(error),
440        }
441    }
442
443    fn delete_checkpoint(&self, checkpoint_key: &str) -> CheckpointResult<()> {
444        let mut connection = self.lock()?;
445        let data_key = Self::data_key(checkpoint_key);
446        let lease_key = Self::lease_key(checkpoint_key);
447        let keys = [data_key.as_str(), lease_key.as_str()];
448        let _: usize = connection.del(&keys).map_err(redis_error)?;
449        Ok(())
450    }
451
452    fn list_checkpoints(&self) -> CheckpointResult<Vec<String>> {
453        let mut connection = self.lock()?;
454        let keys = connection
455            .scan_match::<_, String>(format!("{KEY_PREFIX}*"))
456            .map_err(redis_error)?
457            .filter(|key| !key.ends_with(LEASE_SUFFIX))
458            .collect::<Vec<_>>();
459        let mut checkpoint_keys = Vec::new();
460        for key in keys {
461            let Some(raw) = connection
462                .get::<_, Option<String>>(&key)
463                .map_err(redis_error)?
464            else {
465                continue;
466            };
467            let checkpoint = decode_storage(
468                &raw,
469                connection
470                    .get::<_, Option<u64>>(format!("{key}{LEASE_SUFFIX}"))
471                    .map_err(redis_error)?,
472            )?;
473            checkpoint_keys.push(checkpoint.checkpoint_key);
474        }
475        checkpoint_keys.sort();
476        Ok(checkpoint_keys)
477    }
478}
479
480impl RedisCheckpointStore {
481    fn replace_claimed(
482        &self,
483        checkpoint: Checkpoint,
484        claim_token: &str,
485        expected_revision: u64,
486        kind: ReplaceKind,
487    ) -> CheckpointResult<bool> {
488        let data_key = Self::data_key(&checkpoint.checkpoint_key);
489        let lease_key = Self::lease_key(&checkpoint.checkpoint_key);
490        let result = self.transaction(&data_key, &lease_key, |connection, pipeline| {
491            let Some(raw) = connection
492                .get::<_, Option<String>>(&data_key)
493                .map_err(redis_error)?
494            else {
495                return Ok(None);
496            };
497            let current = decode_storage(
498                &raw,
499                connection
500                    .get::<_, Option<u64>>(&lease_key)
501                    .map_err(redis_error)?,
502            )?;
503            let updated = match kind {
504                ReplaceKind::Progress => {
505                    prepare_progress(&current, checkpoint.clone(), claim_token, expected_revision)?
506                }
507                ReplaceKind::Suspend => {
508                    prepare_suspend(&current, checkpoint.clone(), claim_token, expected_revision)?
509                }
510                ReplaceKind::Commit => {
511                    prepare_commit(&current, checkpoint.clone(), claim_token, expected_revision)?
512                }
513                ReplaceKind::FinalizeClaimed => prepare_finalize_claimed(
514                    &current,
515                    checkpoint.clone(),
516                    claim_token,
517                    expected_revision,
518                )?,
519            };
520            let Some(updated) = updated else {
521                return Ok(None);
522            };
523            let payload = checkpoint_to_json(&updated, MAX_EXTENSION_STATE_BYTES)?;
524            pipeline.set(&data_key, payload).ignore();
525            if updated.claim_token.is_none() {
526                pipeline.del(&lease_key).ignore();
527            }
528            Ok(Some(true))
529        });
530        match result {
531            Ok(value) => Ok(value),
532            Err(error) if error.code() == "checkpoint_store_conflict" => Ok(false),
533            Err(error) => Err(error),
534        }
535    }
536}
537
538#[derive(Clone, Copy)]
539enum ReplaceKind {
540    Progress,
541    Suspend,
542    Commit,
543    FinalizeClaimed,
544}
545
546fn decode_storage(raw: &str, lease: Option<u64>) -> CheckpointResult<Checkpoint> {
547    let mut value: Value = serde_json::from_str(raw)
548        .map_err(|error| CheckpointError::new("checkpoint_json_invalid", error.to_string()))?;
549    if let Some(object) = value.as_object_mut() {
550        object.insert(
551            "lease_expires_at_ms".to_string(),
552            lease.map_or(Value::Null, Value::from),
553        );
554    }
555    let payload = serde_json::to_string(&value)
556        .map_err(|error| CheckpointError::new("checkpoint_json_invalid", error.to_string()))?;
557    checkpoint_from_json(&payload, MAX_EXTENSION_STATE_BYTES)
558}
559
560fn redis_error(error: redis::RedisError) -> CheckpointError {
561    CheckpointError::new("checkpoint_store_redis", error.to_string())
562}