1use std::sync::Arc;
2use std::sync::atomic::Ordering;
3use std::time::Duration;
4
5use crate::time::Instant as StdInstant;
6use crate::time::sleep;
7
8use anyhow::{Context, Result};
9use rivetkit_actor_persist::{generated::v4 as persist_v4, versioned as persist_versioned};
10#[cfg(not(feature = "wasm-runtime"))]
11use tokio::runtime::Handle;
12use tokio::sync::{mpsc, oneshot};
13use tokio::task::JoinHandle;
14#[cfg(test)]
15use tokio::time::timeout;
16use tracing::Instrument;
17
18use crate::actor::connection::{PersistedConnection, decode_persisted_connection};
19use crate::actor::context::ActorContext;
20use crate::actor::internal_storage;
21use crate::actor::lifecycle_hooks::Reply;
22use crate::actor::messages::{StateDelta, WorkflowKvWrite};
23use crate::actor::persist::decode_latest_with_embedded_version;
24#[cfg(test)]
25use crate::actor::persist::encode_latest_with_embedded_version;
26use crate::actor::task::LifecycleEvent;
27use crate::actor::task_types::StateMutationReason;
28use crate::error::ActorRuntime;
29#[cfg(feature = "wasm-runtime")]
30use crate::runtime::RuntimeSpawner;
31use crate::types::SaveStateOpts;
32
33#[cfg(test)]
34const LAST_PUSHED_ALARM_VERSION: u16 = 1;
35
36pub type PersistedScheduleEvent = persist_v4::ScheduleEvent;
37pub type PersistedActor = persist_v4::Actor;
38
39#[cfg(test)]
40pub(crate) fn encode_persisted_actor(actor: &PersistedActor) -> Result<Vec<u8>> {
41 encode_latest_with_embedded_version::<persist_versioned::Actor>(
42 actor.clone(),
43 rivetkit_actor_persist::CURRENT_VERSION,
44 "persisted actor",
45 )
46}
47
48pub(crate) fn decode_persisted_actor(payload: &[u8]) -> Result<PersistedActor> {
49 let actor = decode_latest_with_embedded_version::<persist_versioned::Actor>(
50 payload,
51 "persisted actor",
52 )?;
53 Ok(actor)
54}
55
56#[cfg(test)]
57pub(crate) fn encode_last_pushed_alarm(alarm_ts: Option<i64>) -> Result<Vec<u8>> {
58 encode_latest_with_embedded_version::<persist_versioned::LastPushedAlarm>(
59 alarm_ts,
60 LAST_PUSHED_ALARM_VERSION,
61 "last pushed alarm",
62 )
63}
64
65pub(crate) fn decode_last_pushed_alarm(payload: &[u8]) -> Result<Option<i64>> {
66 decode_latest_with_embedded_version::<persist_versioned::LastPushedAlarm>(
67 payload,
68 "last pushed alarm",
69 )
70}
71
72#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
73pub struct RequestSaveOpts {
74 pub immediate: bool,
75 pub max_wait_ms: Option<u32>,
76}
77
78pub(super) struct PendingSave {
79 scheduled_at: StdInstant,
80 handle: JoinHandle<()>,
81}
82
83pub struct OnStateChangeGuard {
84 ctx: Option<ActorContext>,
85}
86
87impl OnStateChangeGuard {
88 fn new(ctx: ActorContext) -> Self {
89 ctx.on_state_change_started();
90 Self { ctx: Some(ctx) }
91 }
92}
93
94impl Drop for OnStateChangeGuard {
95 fn drop(&mut self) {
96 if let Some(ctx) = self.ctx.take() {
97 ctx.on_state_change_finished();
98 }
99 }
100}
101
102impl ActorContext {
103 pub fn state(&self) -> Vec<u8> {
104 self.0.current_state.read().clone()
105 }
106
107 pub(crate) async fn persist_state(&self, opts: SaveStateOpts) -> Result<()> {
108 if !self.is_dirty() {
109 return Ok(());
110 }
111
112 let result = if opts.immediate {
113 self.clear_pending_save();
114 self.persist_if_dirty().await
115 } else {
116 let delay = self.compute_save_delay(None);
117 if !delay.is_zero() {
118 sleep(delay).await;
119 }
120 self.persist_if_dirty().await
121 };
122 result?;
123 self.record_state_updated();
124 Ok(())
125 }
126
127 pub fn set_state_initial(&self, state: Vec<u8>) {
130 self.set_initial_state(state);
131 }
132
133 pub fn request_save(&self, opts: RequestSaveOpts) {
141 #[cfg(target_arch = "wasm32")]
142 {
143 self.request_save_best_effort(opts);
144 }
145
146 #[cfg(not(target_arch = "wasm32"))]
147 if let Err(error) = self.request_save_with_revision(opts) {
148 tracing::warn!(?error, "failed to request actor state save");
149 }
150 }
151
152 #[cfg(target_arch = "wasm32")]
153 fn request_save_best_effort(&self, opts: RequestSaveOpts) {
154 let immediate = opts.immediate;
155 let _save_request_revision =
156 self.0.save_request_revision.fetch_add(1, Ordering::SeqCst) + 1;
157 self.notify_request_save_hooks(opts);
158 let already_requested = self.0.save_requested.swap(true, Ordering::SeqCst);
159 let immediate_already_requested = if immediate {
160 self.0.save_requested_immediate.swap(true, Ordering::SeqCst)
161 } else {
162 self.0.save_requested_immediate.load(Ordering::SeqCst)
163 };
164
165 if let Some(max_wait_ms) = opts.max_wait_ms {
166 let deadline = StdInstant::now() + Duration::from_millis(u64::from(max_wait_ms));
167 let mut requested_deadline = self.0.save_requested_within_deadline.lock();
168 *requested_deadline = Some(match *requested_deadline {
169 Some(existing) => existing.min(deadline),
170 None => deadline,
171 });
172 }
173
174 let Some(sender) = self.lifecycle_event_sender() else {
175 return;
176 };
177
178 if opts.max_wait_ms.is_none()
179 && already_requested
180 && (!immediate || immediate_already_requested)
181 {
182 return;
183 }
184
185 let _ = sender.send(LifecycleEvent::SaveRequested { immediate });
186 }
187
188 pub async fn request_save_and_wait(&self, opts: RequestSaveOpts) -> Result<()> {
189 let save_request_revision = self.request_save_with_revision(opts)?;
190 self.wait_for_save_request(save_request_revision).await;
191 Ok(())
192 }
193
194 pub async fn save_state(&self, deltas: Vec<StateDelta>) -> Result<()> {
195 let save_request_revision = self.save_request_revision();
196 self.save_state_with_revision(deltas, save_request_revision)
197 .await
198 }
199
200 pub async fn save_state_and_workflow_batch(
203 &self,
204 workflow_writes: Vec<WorkflowKvWrite>,
205 ) -> Result<()> {
206 let Some(sender) = self.lifecycle_event_sender() else {
207 return Err(ActorRuntime::NotConfigured {
208 component: "lifecycle events".to_owned(),
209 }
210 .build());
211 };
212 let (reply_tx, reply_rx) = oneshot::channel();
213 sender
214 .send(LifecycleEvent::WorkflowFlushRequested {
215 writes: workflow_writes,
216 reply: Reply::from(reply_tx),
217 })
218 .map_err(|_| {
219 ActorRuntime::NotConfigured {
220 component: "lifecycle events".to_owned(),
221 }
222 .build()
223 })?;
224 reply_rx
225 .await
226 .context("receive workflow flush lifecycle reply")?
227 }
228
229 #[cfg(test)]
232 pub(crate) async fn commit_serialized_state_and_workflow_batch(
233 &self,
234 deltas: Vec<StateDelta>,
235 workflow_writes: Vec<WorkflowKvWrite>,
236 ) -> Result<()> {
237 let save_request_revision = self.save_request_revision();
238 self.save_state_and_workflow_batch_with_revision(
239 deltas,
240 workflow_writes,
241 save_request_revision,
242 )
243 .await
244 }
245
246 pub(crate) fn request_save_with_revision(&self, opts: RequestSaveOpts) -> Result<u64> {
247 let immediate = opts.immediate;
248 let save_request_revision = self.0.save_request_revision.fetch_add(1, Ordering::SeqCst) + 1;
249 self.notify_request_save_hooks(opts);
250 let already_requested = self.0.save_requested.swap(true, Ordering::SeqCst);
251 let immediate_already_requested = if immediate {
252 self.0.save_requested_immediate.swap(true, Ordering::SeqCst)
253 } else {
254 self.0.save_requested_immediate.load(Ordering::SeqCst)
255 };
256
257 if let Some(max_wait_ms) = opts.max_wait_ms {
258 let deadline = StdInstant::now() + Duration::from_millis(u64::from(max_wait_ms));
259 let mut requested_deadline = self.0.save_requested_within_deadline.lock();
260 *requested_deadline = Some(match *requested_deadline {
261 Some(existing) => existing.min(deadline),
262 None => deadline,
263 });
264 }
265
266 let Some(sender) = self.lifecycle_event_sender() else {
267 return Err(ActorRuntime::NotConfigured {
268 component: "lifecycle events".to_owned(),
269 }
270 .build());
271 };
272
273 if opts.max_wait_ms.is_none()
274 && already_requested
275 && (!immediate || immediate_already_requested)
276 {
277 return Ok(save_request_revision);
278 }
279
280 sender
281 .send(LifecycleEvent::SaveRequested { immediate })
282 .map(|()| save_request_revision)
283 .map_err(|_| {
284 ActorRuntime::NotConfigured {
285 component: "lifecycle events".to_owned(),
286 }
287 .build()
288 })
289 }
290
291 pub(crate) async fn wait_for_save_request(&self, save_request_revision: u64) {
292 loop {
293 if self.0.save_completed_revision.load(Ordering::SeqCst) >= save_request_revision {
294 return;
295 }
296
297 self.0.save_completion.notified().await;
298 }
299 }
300
301 pub(crate) fn save_requested(&self) -> bool {
302 self.0.save_requested.load(Ordering::SeqCst)
303 }
304
305 pub(crate) fn save_requested_immediate(&self) -> bool {
306 self.0.save_requested_immediate.load(Ordering::SeqCst)
307 }
308
309 pub(crate) fn save_deadline(&self, immediate: bool) -> StdInstant {
310 self.compute_save_deadline(immediate)
311 }
312
313 pub(crate) fn compute_save_deadline(&self, immediate: bool) -> StdInstant {
314 if immediate || self.save_requested_immediate() {
315 return StdInstant::now();
316 }
317
318 let throttled_deadline = StdInstant::now() + self.compute_save_delay(None);
319 let requested_deadline = *self.0.save_requested_within_deadline.lock();
320
321 match requested_deadline {
322 Some(requested_deadline) => throttled_deadline.min(requested_deadline),
323 None => throttled_deadline,
324 }
325 }
326
327 pub(crate) fn save_request_revision(&self) -> u64 {
328 self.0.save_request_revision.load(Ordering::SeqCst)
329 }
330
331 pub(crate) async fn apply_state_deltas(
332 &self,
333 deltas: Vec<StateDelta>,
334 save_request_revision: u64,
335 ) -> Result<()> {
336 self.apply_state_deltas_inner(deltas, None, save_request_revision)
337 .await
338 }
339
340 pub(crate) async fn apply_state_deltas_and_workflow(
341 &self,
342 deltas: Vec<StateDelta>,
343 workflow_writes: Vec<WorkflowKvWrite>,
344 save_request_revision: u64,
345 ) -> Result<()> {
346 self.apply_state_deltas_inner(deltas, Some(workflow_writes), save_request_revision)
347 .await
348 }
349
350 async fn apply_state_deltas_inner(
351 &self,
352 deltas: Vec<StateDelta>,
353 workflow_writes: Option<Vec<WorkflowKvWrite>>,
354 save_request_revision: u64,
355 ) -> Result<()> {
356 let delta_count = deltas.len();
357 let delta_bytes: usize = deltas.iter().map(StateDelta::payload_len).sum();
358 let workflow_write_count = workflow_writes.as_ref().map_or(0, Vec::len);
359 let current_revision = self.0.state_revision.load(Ordering::SeqCst);
360 tracing::debug!(
361 delta_count,
362 delta_bytes,
363 state_revision = current_revision,
364 save_request_revision,
365 "applying actor state deltas"
366 );
367 self.clear_pending_save();
368
369 if deltas.is_empty() && workflow_write_count == 0 {
370 self.mark_save_request_completed(save_request_revision);
371 self.finish_save_request(save_request_revision);
372 tracing::debug!(
373 delta_count,
374 state_revision = current_revision,
375 save_request_revision,
376 "actor state deltas applied without kv write"
377 );
378 return Ok(());
379 }
380
381 let (
382 next_state,
383 actor_to_persist,
384 connections_to_persist,
385 connections_to_delete,
386 revision,
387 _write_guard,
388 ) = {
389 let _save_guard = self.0.save_guard.lock().await;
390 let revision = self.0.state_revision.load(Ordering::SeqCst);
391 let mut persisted = self.persisted();
392 let mut next_state = None;
393 let mut actor_to_persist = None;
394 let mut connections_to_persist: Vec<PersistedConnection> = Vec::new();
395 let mut connections_to_delete = Vec::new();
396
397 for delta in deltas {
398 match delta {
399 StateDelta::ActorState(bytes) => {
400 next_state = Some(bytes.clone());
401 persisted.state = bytes;
402 }
403 StateDelta::ConnHibernation { conn: _, bytes } => {
404 connections_to_persist.push(
405 decode_persisted_connection(&bytes)
406 .context("decode hibernatable connection state delta")?,
407 );
408 }
409 StateDelta::ConnHibernationRemoved(conn) => {
410 connections_to_delete.push(conn);
411 }
412 }
413 }
414
415 if next_state.is_some() {
416 actor_to_persist = Some(persisted.clone());
417 }
418
419 (
420 next_state,
421 actor_to_persist,
422 connections_to_persist,
423 connections_to_delete,
424 revision,
425 self.begin_write(),
426 )
427 };
428
429 if let Some(workflow_writes) = workflow_writes.as_deref() {
430 internal_storage::persist_actor_core_connections_and_workflow(
431 self.sql(),
432 actor_to_persist.as_ref(),
433 &connections_to_persist,
434 &connections_to_delete,
435 workflow_writes,
436 )
437 .await
438 .context("atomically persist actor state, connection deltas, and workflow kv")?;
439 } else if actor_to_persist.is_some()
440 || !connections_to_persist.is_empty()
441 || !connections_to_delete.is_empty()
442 {
443 internal_storage::persist_actor_core_and_connections(
444 self.sql(),
445 actor_to_persist.as_ref(),
446 &connections_to_persist,
447 &connections_to_delete,
448 )
449 .await
450 .context("persist actor state and connection deltas to sqlite")?;
451 }
452
453 if let Some(state) = next_state {
454 self.0.persisted.write().state = state.clone();
455 *self.0.current_state.write() = state;
456 }
457 for connection in &connections_to_persist {
458 if let Some(handle) = self.connection(&connection.id) {
459 handle.set_state_initial(connection.state.clone());
460 }
461 }
462
463 *self.0.last_save_at.lock() = Some(StdInstant::now());
464
465 if self.0.state_revision.load(Ordering::SeqCst) == revision {
466 self.0.state_dirty.store(false, Ordering::SeqCst);
467 }
468
469 self.mark_save_request_completed(save_request_revision);
470 self.finish_save_request(save_request_revision);
471 tracing::debug!(
472 delta_count,
473 delta_bytes,
474 workflow_write_count,
475 state_revision = self.0.state_revision.load(Ordering::SeqCst),
476 save_request_revision,
477 "actor state deltas applied"
478 );
479 Ok(())
480 }
481
482 pub(crate) async fn wait_for_pending_writes(&self) {
483 loop {
484 if let Some(handle) = self.take_tracked_persist() {
485 let _ = handle.await;
486 continue;
487 }
488
489 let save_guard = self.0.save_guard.lock().await;
490 if self.has_tracked_persist() {
491 drop(save_guard);
492 continue;
493 }
494
495 if self.0.in_flight_state_writes.load(Ordering::SeqCst) == 0 {
496 return;
497 }
498 drop(save_guard);
499
500 self.wait_for_in_flight_writes().await;
501 }
502 }
503
504 pub(crate) async fn wait_for_pending_state_writes(&self) {
505 self.wait_for_pending_writes().await;
506 }
507
508 pub fn begin_on_state_change(&self) -> OnStateChangeGuard {
509 OnStateChangeGuard::new(self.clone())
510 }
511
512 pub fn on_state_change_started(&self) {
513 self.0
514 .on_state_change_in_flight
515 .fetch_add(1, Ordering::SeqCst);
516 self.0.sleep.work.keep_awake.increment();
517 self.reset_sleep_timer();
518 }
519
520 pub fn on_state_change_finished(&self) {
521 let previous = self.0.on_state_change_in_flight.fetch_update(
522 Ordering::SeqCst,
523 Ordering::SeqCst,
524 |count| count.checked_sub(1),
525 );
526
527 match previous {
528 Ok(1) => {
529 self.0.sleep.work.keep_awake.decrement();
530 self.0.on_state_change_idle.notify_waiters();
531 self.reset_sleep_timer();
532 }
533 Ok(_) => {
534 self.0.sleep.work.keep_awake.decrement();
535 self.reset_sleep_timer();
536 }
537 Err(_) => {
538 tracing::warn!(
539 actor_id = %self.actor_id(),
540 "on_state_change finished without a matching start"
541 );
542 }
543 }
544 }
545
546 #[cfg(test)]
547 #[allow(dead_code)]
548 pub(crate) async fn wait_for_on_state_change_idle(&self, timeout_duration: Duration) -> bool {
549 if self.0.on_state_change_in_flight.load(Ordering::SeqCst) == 0 {
550 return true;
551 }
552
553 timeout(timeout_duration, async {
554 loop {
555 let idle = self.0.on_state_change_idle.notified();
556 tokio::pin!(idle);
557 idle.as_mut().enable();
558
559 if self.0.on_state_change_in_flight.load(Ordering::SeqCst) == 0 {
560 return;
561 }
562
563 idle.await;
564 }
565 })
566 .await
567 .is_ok()
568 }
569
570 pub fn persisted(&self) -> PersistedActor {
571 self.0.persisted.read().clone()
572 }
573
574 pub fn load_persisted(&self, persisted: PersistedActor) {
575 let state = persisted.state.clone();
576 *self.0.persisted.write() = persisted;
577 *self.0.current_state.write() = state;
578 self.0.state_dirty.store(false, Ordering::SeqCst);
579 self.finish_save_request(self.save_request_revision());
580 self.0
581 .metrics
582 .inc_state_mutation(StateMutationReason::InternalReplace);
583 }
584
585 pub(crate) fn load_last_pushed_alarm(&self, alarm_ts: Option<i64>) {
586 *self.0.last_pushed_alarm.write() = alarm_ts;
587 }
588
589 pub(crate) fn last_pushed_alarm(&self) -> Option<i64> {
590 *self.0.last_pushed_alarm.read()
591 }
592
593 pub(crate) async fn persist_last_pushed_alarm(&self, alarm_ts: Option<i64>) -> Result<()> {
594 internal_storage::persist_last_pushed_alarm(self.sql(), alarm_ts)
595 .await
596 .context("persist last pushed alarm to sqlite")?;
597 self.load_last_pushed_alarm(alarm_ts);
598 Ok(())
599 }
600
601 pub(crate) fn set_initial_state(&self, state: Vec<u8>) {
602 *self.0.current_state.write() = state.clone();
603 self.0.persisted.write().state = state;
604 self.0.state_dirty.store(true, Ordering::SeqCst);
605 self.0.state_revision.fetch_add(1, Ordering::SeqCst);
606 }
607
608 pub fn scheduled_events(&self) -> Vec<PersistedScheduleEvent> {
609 self.0.persisted.read().scheduled_events.clone()
610 }
611
612 pub fn set_scheduled_events(&self, scheduled_events: Vec<PersistedScheduleEvent>) {
613 self.0.persisted.write().scheduled_events = scheduled_events;
614 self.0
615 .metrics
616 .inc_state_mutation(StateMutationReason::ScheduledEventsUpdate);
617 self.mark_dirty();
618 self.schedule_save(None);
619 }
620
621 pub fn set_input(&self, input: Option<Vec<u8>>) {
622 self.0.persisted.write().input = input;
623 self.0
624 .metrics
625 .inc_state_mutation(StateMutationReason::InputSet);
626 self.mark_dirty();
627 self.schedule_save(None);
628 }
629
630 pub fn input(&self) -> Option<Vec<u8>> {
631 self.0.persisted.read().input.clone()
632 }
633
634 pub fn set_has_initialized(&self, has_initialized: bool) {
635 {
636 let mut persisted = self.0.persisted.write();
637 if persisted.has_initialized == has_initialized {
638 return;
639 }
640 persisted.has_initialized = has_initialized;
641 }
642 self.0
643 .metrics
644 .inc_state_mutation(StateMutationReason::HasInitialized);
645 self.mark_dirty();
646 self.schedule_save(None);
647 }
648
649 pub fn has_initialized(&self) -> bool {
650 self.0.persisted.read().has_initialized
651 }
652
653 pub fn flush_on_shutdown(&self) {
654 self.persist_now_tracked("shutdown_flush");
655 }
656
657 pub fn on_request_save(&self, hook: Box<dyn Fn(RequestSaveOpts) + Send + Sync>) {
658 self.0.request_save_hooks.write().push(Arc::from(hook));
659 }
660
661 fn is_dirty(&self) -> bool {
662 self.0.state_dirty.load(Ordering::SeqCst)
663 }
664
665 fn mark_dirty(&self) {
666 self.0.state_dirty.store(true, Ordering::SeqCst);
667 self.0.state_revision.fetch_add(1, Ordering::SeqCst);
668 }
669
670 fn lifecycle_event_sender(&self) -> Option<mpsc::UnboundedSender<LifecycleEvent>> {
671 self.0.lifecycle_events.read().clone()
672 }
673
674 fn compute_save_delay(&self, max_wait: Option<Duration>) -> Duration {
675 let elapsed = self
676 .0
677 .last_save_at
678 .lock()
679 .map(|instant| instant.elapsed())
680 .unwrap_or_default();
681
682 throttled_save_delay(self.0.state_save_interval, elapsed, max_wait)
683 }
684
685 fn schedule_save(&self, max_wait: Option<Duration>) {
686 if !self.is_dirty() {
687 return;
688 }
689
690 let delay = self.compute_save_delay(max_wait);
691 let scheduled_at = StdInstant::now() + delay;
692
693 let mut pending_save = self.0.pending_save.lock();
694
695 if let Some(existing) = pending_save.as_ref() {
696 if existing.scheduled_at <= scheduled_at {
697 return;
698 }
699
700 existing.handle.abort();
701 }
702
703 let state = self.clone();
704 let task = async move {
708 if !delay.is_zero() {
709 sleep(delay).await;
710 }
711
712 state.take_pending_save();
713
714 if let Err(error) = state.persist_if_dirty().await {
715 tracing::error!(?error, "failed to persist actor state");
716 }
717 }
718 .in_current_span();
719
720 #[cfg(not(feature = "wasm-runtime"))]
721 let handle = {
722 let Ok(tokio_handle) = Handle::try_current() else {
723 return;
724 };
725 tokio_handle.spawn(task)
726 };
727
728 #[cfg(feature = "wasm-runtime")]
729 let handle = RuntimeSpawner::spawn(task);
730
731 *pending_save = Some(PendingSave {
732 scheduled_at,
733 handle,
734 });
735 }
736
737 fn clear_pending_save(&self) {
738 if let Some(pending_save) = self.take_pending_save() {
739 pending_save.handle.abort();
740 }
741 }
742
743 pub(crate) fn persist_now_tracked(&self, description: &'static str) {
744 self.clear_pending_save();
745
746 let state = self.clone();
747 let mut tracked_persist = self.0.tracked_persist.lock();
748 let previous = tracked_persist.take();
749 let task = async move {
750 if let Some(previous) = previous {
751 let _ = previous.await;
752 }
753
754 if let Err(error) = state.persist_state(SaveStateOpts { immediate: true }).await {
755 tracing::error!(?error, description, "failed to persist actor state");
756 }
757 }
758 .in_current_span();
759
760 #[cfg(not(feature = "wasm-runtime"))]
761 let handle = {
762 let Ok(tokio_handle) = Handle::try_current() else {
763 tracing::warn!(
764 description,
765 "skipping tracked actor state persistence without runtime"
766 );
767 return;
768 };
769 tokio_handle.spawn(task)
770 };
771
772 #[cfg(feature = "wasm-runtime")]
773 let handle = RuntimeSpawner::spawn(task);
774
775 *tracked_persist = Some(handle);
776 }
777
778 fn take_pending_save(&self) -> Option<PendingSave> {
779 self.0.pending_save.lock().take()
780 }
781
782 fn take_tracked_persist(&self) -> Option<JoinHandle<()>> {
783 self.0.tracked_persist.lock().take()
784 }
785
786 fn has_tracked_persist(&self) -> bool {
787 self.0.tracked_persist.lock().is_some()
788 }
789
790 #[cfg(test)]
791 pub(crate) fn tracked_persist_pending(&self) -> bool {
792 self.has_tracked_persist()
793 }
794
795 async fn persist_if_dirty(&self) -> Result<()> {
796 if !self.is_dirty() {
797 return Ok(());
798 }
799
800 let (revision, actor_to_persist, _write_guard) = {
801 let _save_guard = self.0.save_guard.lock().await;
802 if !self.is_dirty() {
803 return Ok(());
804 }
805
806 let revision = self.0.state_revision.load(Ordering::SeqCst);
807 let persisted = self.persisted();
808 (revision, persisted, self.begin_write())
809 };
810
811 internal_storage::persist_actor_snapshot(self.sql(), &actor_to_persist)
812 .await
813 .context("persist actor state to sqlite")?;
814
815 *self.0.last_save_at.lock() = Some(StdInstant::now());
816
817 if self.0.state_revision.load(Ordering::SeqCst) == revision {
818 self.0.state_dirty.store(false, Ordering::SeqCst);
819 }
820
821 Ok(())
822 }
823
824 fn begin_write(&self) -> InFlightWrite {
825 self.0.in_flight_state_writes.fetch_add(1, Ordering::SeqCst);
826 InFlightWrite { ctx: self.clone() }
827 }
828
829 async fn wait_for_in_flight_writes(&self) {
830 loop {
831 if self.0.in_flight_state_writes.load(Ordering::SeqCst) == 0 {
832 return;
833 }
834 self.0.state_write_completion.notified().await;
835 }
836 }
837
838 fn finish_save_request(&self, save_request_revision: u64) {
839 if self.0.save_request_revision.load(Ordering::SeqCst) == save_request_revision {
840 self.0.save_requested.store(false, Ordering::SeqCst);
841 self.0
842 .save_requested_immediate
843 .store(false, Ordering::SeqCst);
844 *self.0.save_requested_within_deadline.lock() = None;
845 }
846 }
847
848 fn mark_save_request_completed(&self, save_request_revision: u64) {
849 self.0
850 .save_completed_revision
851 .fetch_max(save_request_revision, Ordering::SeqCst);
852 self.0.save_completion.notify_waiters();
853 }
854
855 fn notify_request_save_hooks(&self, opts: RequestSaveOpts) {
856 let hooks = self.0.request_save_hooks.read().clone();
857 for hook in hooks {
858 hook(opts);
859 }
860 }
861}
862
863struct InFlightWrite {
864 ctx: ActorContext,
865}
866
867impl Drop for InFlightWrite {
868 fn drop(&mut self) {
869 if self
870 .ctx
871 .0
872 .in_flight_state_writes
873 .fetch_sub(1, Ordering::SeqCst)
874 == 1
875 {
876 self.ctx.0.state_write_completion.notify_waiters();
877 self.ctx.0.state_write_completion.notify_one();
878 }
879 }
880}
881
882fn throttled_save_delay(
883 save_interval: Duration,
884 time_since_last_save: Duration,
885 max_wait: Option<Duration>,
886) -> Duration {
887 let save_delay = save_interval.saturating_sub(time_since_last_save);
888 if let Some(max_wait) = max_wait {
889 save_delay.min(max_wait)
890 } else {
891 save_delay
892 }
893}
894
895#[cfg(test)]
897#[path = "../../tests/state.rs"]
898mod tests;