Skip to main content

rivetkit_core/actor/
state.rs

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