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::{Mutex as AsyncMutex, OwnedMutexGuard, 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::sqlite::{BindParam, ExecuteResult, SqliteTransaction};
32use crate::types::SaveStateOpts;
33
34#[cfg(test)]
35const LAST_PUSHED_ALARM_VERSION: u16 = 1;
36
37pub type PersistedScheduleEvent = persist_v4::ScheduleEvent;
38pub type PersistedActor = persist_v4::Actor;
39
40#[cfg(test)]
41pub(crate) fn encode_persisted_actor(actor: &PersistedActor) -> Result<Vec<u8>> {
42	encode_latest_with_embedded_version::<persist_versioned::Actor>(
43		actor.clone(),
44		rivetkit_actor_persist::CURRENT_VERSION,
45		"persisted actor",
46	)
47}
48
49pub(crate) fn decode_persisted_actor(payload: &[u8]) -> Result<PersistedActor> {
50	let actor = decode_latest_with_embedded_version::<persist_versioned::Actor>(
51		payload,
52		"persisted actor",
53	)?;
54	Ok(actor)
55}
56
57#[cfg(test)]
58pub(crate) fn encode_last_pushed_alarm(alarm_ts: Option<i64>) -> Result<Vec<u8>> {
59	encode_latest_with_embedded_version::<persist_versioned::LastPushedAlarm>(
60		alarm_ts,
61		LAST_PUSHED_ALARM_VERSION,
62		"last pushed alarm",
63	)
64}
65
66pub(crate) fn decode_last_pushed_alarm(payload: &[u8]) -> Result<Option<i64>> {
67	decode_latest_with_embedded_version::<persist_versioned::LastPushedAlarm>(
68		payload,
69		"last pushed alarm",
70	)
71}
72
73#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
74pub struct RequestSaveOpts {
75	pub immediate: bool,
76	pub max_wait_ms: Option<u32>,
77}
78
79pub(super) struct PendingSave {
80	scheduled_at: StdInstant,
81	handle: JoinHandle<()>,
82}
83
84pub struct OnStateChangeGuard {
85	ctx: Option<ActorContext>,
86}
87/// A SQLite transaction that owns the actor state-save exclusion until it is
88/// committed or rolled back.
89#[derive(Clone)]
90pub struct ActorStateTransaction {
91	owner: Arc<AsyncMutex<ActorStateTransactionOwner>>,
92}
93
94struct ActorStateTransactionOwner {
95	ctx: ActorContext,
96	transaction: SqliteTransaction,
97	save_guard: Option<OwnedMutexGuard<()>>,
98	write_guard: Option<InFlightWrite>,
99	finalized: bool,
100	committed: bool,
101	save_request_revision: u64,
102	statement_count: usize,
103	bind_payload_bytes: usize,
104	rollback_connection_states: Vec<(String, Vec<u8>)>,
105}
106
107impl ActorStateTransactionOwner {
108	fn restore_connection_states(&self) {
109		for (conn_id, state) in &self.rollback_connection_states {
110			if let Some(conn) = self.ctx.connection(conn_id) {
111				conn.set_state(state.clone());
112			}
113		}
114	}
115
116	fn advance_epoch(&self) {
117		self.ctx
118			.0
119			.state_transaction_epoch
120			.fetch_add(1, Ordering::SeqCst);
121	}
122}
123
124impl Drop for ActorStateTransactionOwner {
125	fn drop(&mut self) {
126		if !self.finalized {
127			self.restore_connection_states();
128			self.advance_epoch();
129			let transaction = self.transaction.clone();
130			#[cfg(not(feature = "wasm-runtime"))]
131			if let Ok(handle) = Handle::try_current() {
132				handle.spawn(async move {
133					if let Err(error) = transaction.rollback().await {
134						tracing::debug!(?error, "dropped actor state transaction rollback failed");
135					}
136				});
137			}
138			#[cfg(feature = "wasm-runtime")]
139			wasm_bindgen_futures::spawn_local(async move {
140				if let Err(error) = transaction.rollback().await {
141					tracing::debug!(?error, "dropped actor state transaction rollback failed");
142				}
143			});
144		}
145		if !self.committed {
146			self.ctx.schedule_save(None);
147		}
148	}
149}
150
151impl ActorStateTransaction {
152	pub async fn execute(
153		&self,
154		sql: impl Into<String>,
155		params: Option<Vec<BindParam>>,
156	) -> Result<ExecuteResult> {
157		let mut owner = self.owner.lock().await;
158		if owner.finalized {
159			return Err(anyhow::anyhow!(
160				"actor state transaction is already finalized"
161			));
162		}
163		let sql = sql.into();
164		if is_transaction_terminator(&sql) {
165			return Err(anyhow::anyhow!(
166				"cannot execute transaction terminator inside an actor state transaction; use the transaction commit or rollback method"
167			));
168		}
169		let payload_bytes = match params.as_deref() {
170			Some(params) => params
171				.iter()
172				.map(internal_storage::bind_param_payload_len)
173				.fold(0usize, usize::saturating_add),
174			None => sql.len(),
175		};
176		let result = owner.transaction.execute(sql, params).await?;
177		owner.statement_count = owner.statement_count.saturating_add(1);
178		owner.bind_payload_bytes = owner.bind_payload_bytes.saturating_add(payload_bytes);
179		Ok(result)
180	}
181
182	pub async fn commit(&self, deltas: Vec<StateDelta>) -> Result<()> {
183		let mut owner = self.owner.lock().await;
184		if owner.finalized {
185			return Err(anyhow::anyhow!(
186				"actor state transaction is already finalized"
187			));
188		}
189
190		let save_request_revision = owner.save_request_revision;
191		let transaction = owner.transaction.clone();
192		let statement_count = owner.statement_count;
193		let bind_payload_bytes = owner.bind_payload_bytes;
194		let result = owner
195			.ctx
196			.commit_state_transaction(
197				&transaction,
198				deltas,
199				save_request_revision,
200				statement_count,
201				bind_payload_bytes,
202			)
203			.await;
204		owner.finalized = true;
205		owner.committed = result.is_ok();
206		if result.is_err() {
207			owner.restore_connection_states();
208		}
209		owner.advance_epoch();
210		owner.write_guard.take();
211		owner.save_guard.take();
212		if result.is_err() {
213			owner.ctx.schedule_save(None);
214		}
215		result
216	}
217
218	pub async fn rollback(&self) -> Result<()> {
219		let mut owner = self.owner.lock().await;
220		if owner.finalized {
221			return Err(anyhow::anyhow!(
222				"actor state transaction is already finalized"
223			));
224		}
225		let result = owner.transaction.rollback().await;
226		owner.finalized = true;
227		owner.restore_connection_states();
228		owner.advance_epoch();
229		owner.write_guard.take();
230		owner.save_guard.take();
231		owner.ctx.schedule_save(None);
232		result
233	}
234}
235
236fn is_transaction_terminator(sql: &str) -> bool {
237	let mut offset = 0;
238	let Some(first) = next_sql_keyword(sql, &mut offset) else {
239		return false;
240	};
241	if first.eq_ignore_ascii_case("COMMIT") || first.eq_ignore_ascii_case("END") {
242		return true;
243	}
244	if !first.eq_ignore_ascii_case("ROLLBACK") {
245		return false;
246	}
247
248	let mut next = next_sql_keyword(sql, &mut offset);
249	if next.is_some_and(|keyword| keyword.eq_ignore_ascii_case("TRANSACTION")) {
250		next = next_sql_keyword(sql, &mut offset);
251	}
252	!next.is_some_and(|keyword| keyword.eq_ignore_ascii_case("TO"))
253}
254
255fn next_sql_keyword<'a>(sql: &'a str, offset: &mut usize) -> Option<&'a str> {
256	let bytes = sql.as_bytes();
257	loop {
258		while bytes.get(*offset).is_some_and(u8::is_ascii_whitespace) {
259			*offset += 1;
260		}
261		if bytes.get(*offset..*offset + 2) == Some(b"--") {
262			*offset += 2;
263			while bytes.get(*offset).is_some_and(|byte| *byte != b'\n') {
264				*offset += 1;
265			}
266			continue;
267		}
268		if bytes.get(*offset..*offset + 2) == Some(b"/*") {
269			*offset += 2;
270			while bytes.get(*offset..*offset + 2) != Some(b"*/") {
271				bytes.get(*offset)?;
272				*offset += 1;
273			}
274			*offset += 2;
275			continue;
276		}
277		break;
278	}
279
280	let start = *offset;
281	while bytes.get(*offset).is_some_and(u8::is_ascii_alphabetic) {
282		*offset += 1;
283	}
284	(start != *offset).then(|| &sql[start..*offset])
285}
286
287impl OnStateChangeGuard {
288	fn new(ctx: ActorContext) -> Self {
289		ctx.on_state_change_started();
290		Self { ctx: Some(ctx) }
291	}
292}
293
294impl Drop for OnStateChangeGuard {
295	fn drop(&mut self) {
296		if let Some(ctx) = self.ctx.take() {
297			ctx.on_state_change_finished();
298		}
299	}
300}
301
302impl ActorContext {
303	pub async fn begin_state_transaction(
304		&self,
305		timeout: Option<Duration>,
306	) -> Result<ActorStateTransaction> {
307		self.clear_pending_save();
308		let save_guard = Arc::clone(&self.0.save_guard).lock_owned().await;
309		self.0
310			.state_transaction_epoch
311			.fetch_add(1, Ordering::SeqCst);
312		self.wait_for_in_flight_writes().await;
313		let rollback_connection_states = self
314			.iter_connections()
315			.filter(|conn| conn.is_hibernatable())
316			.map(|conn| (conn.id().to_owned(), conn.state()))
317			.collect();
318		let transaction = match self.sql().begin_transaction(timeout).await {
319			Ok(transaction) => transaction,
320			Err(error) => {
321				self.0
322					.state_transaction_epoch
323					.fetch_add(1, Ordering::SeqCst);
324				drop(save_guard);
325				self.schedule_save(None);
326				return Err(error);
327			}
328		};
329		let write_guard = self.begin_write();
330		let save_request_revision = self.save_request_revision();
331		Ok(ActorStateTransaction {
332			owner: Arc::new(AsyncMutex::new(ActorStateTransactionOwner {
333				ctx: self.clone(),
334				transaction,
335				save_guard: Some(save_guard),
336				write_guard: Some(write_guard),
337				finalized: false,
338				committed: false,
339				save_request_revision,
340				statement_count: 0,
341				bind_payload_bytes: 0,
342				rollback_connection_states,
343			})),
344		})
345	}
346	pub fn state(&self) -> Vec<u8> {
347		self.0.current_state.read().clone()
348	}
349
350	pub(crate) async fn persist_state(&self, opts: SaveStateOpts) -> Result<()> {
351		if !self.is_dirty() {
352			return Ok(());
353		}
354
355		let result = if opts.immediate {
356			self.clear_pending_save();
357			self.persist_if_dirty().await
358		} else {
359			let delay = self.compute_save_delay(None);
360			if !delay.is_zero() {
361				sleep(delay).await;
362			}
363			self.persist_if_dirty().await
364		};
365		result?;
366		self.record_state_updated();
367		Ok(())
368	}
369
370	/// Foreign-runtime bootstrap hook for installing the actor state snapshot
371	/// before the actor starts handling lifecycle/dispatch work.
372	pub fn set_state_initial(&self, state: Vec<u8>) {
373		self.set_initial_state(state);
374	}
375
376	/// Fire-and-forget save request helper.
377	///
378	/// If the lifecycle event inbox is unavailable, this only logs a warning and
379	/// returns. That `warn!` is the sole failure signal for this path; callers do
380	/// not receive a `Result`. Call
381	/// [`Self::request_save_and_wait`] when the caller must observe
382	/// save-request delivery failures.
383	pub fn request_save(&self, opts: RequestSaveOpts) {
384		#[cfg(target_arch = "wasm32")]
385		{
386			self.request_save_best_effort(opts);
387		}
388
389		#[cfg(not(target_arch = "wasm32"))]
390		if let Err(error) = self.request_save_with_revision(opts) {
391			tracing::warn!(?error, "failed to request actor state save");
392		}
393	}
394
395	#[cfg(target_arch = "wasm32")]
396	fn request_save_best_effort(&self, opts: RequestSaveOpts) {
397		let immediate = opts.immediate;
398		let _save_request_revision =
399			self.0.save_request_revision.fetch_add(1, Ordering::SeqCst) + 1;
400		self.notify_request_save_hooks(opts);
401		let already_requested = self.0.save_requested.swap(true, Ordering::SeqCst);
402		let immediate_already_requested = if immediate {
403			self.0.save_requested_immediate.swap(true, Ordering::SeqCst)
404		} else {
405			self.0.save_requested_immediate.load(Ordering::SeqCst)
406		};
407
408		if let Some(max_wait_ms) = opts.max_wait_ms {
409			let deadline = StdInstant::now() + Duration::from_millis(u64::from(max_wait_ms));
410			let mut requested_deadline = self.0.save_requested_within_deadline.lock();
411			*requested_deadline = Some(match *requested_deadline {
412				Some(existing) => existing.min(deadline),
413				None => deadline,
414			});
415		}
416
417		let Some(sender) = self.lifecycle_event_sender() else {
418			return;
419		};
420
421		if opts.max_wait_ms.is_none()
422			&& already_requested
423			&& (!immediate || immediate_already_requested)
424		{
425			return;
426		}
427
428		let _ = sender.send(LifecycleEvent::SaveRequested { immediate });
429	}
430
431	pub async fn request_save_and_wait(&self, opts: RequestSaveOpts) -> Result<()> {
432		let save_request_revision = self.request_save_with_revision(opts)?;
433		self.wait_for_save_request(save_request_revision).await;
434		Ok(())
435	}
436
437	pub async fn save_state(&self, deltas: Vec<StateDelta>) -> Result<()> {
438		let save_request_revision = self.save_request_revision();
439		self.save_state_with_revision(deltas, save_request_revision)
440			.await
441	}
442
443	/// Requests one logical workflow-engine flush. The actor lifecycle owns
444	/// serializing actor state and committing both sides in one SQLite transaction.
445	pub async fn save_state_and_workflow_batch(
446		&self,
447		workflow_writes: Vec<WorkflowKvWrite>,
448	) -> Result<()> {
449		let Some(sender) = self.lifecycle_event_sender() else {
450			return Err(ActorRuntime::NotConfigured {
451				component: "lifecycle events".to_owned(),
452			}
453			.build());
454		};
455		let (reply_tx, reply_rx) = oneshot::channel();
456		sender
457			.send(LifecycleEvent::WorkflowFlushRequested {
458				writes: workflow_writes,
459				reply: Reply::from(reply_tx),
460			})
461			.map_err(|_| {
462				ActorRuntime::NotConfigured {
463					component: "lifecycle events".to_owned(),
464				}
465				.build()
466			})?;
467		reply_rx
468			.await
469			.context("receive workflow flush lifecycle reply")?
470	}
471
472	/// Commits an already serialized snapshot and workflow flush atomically for
473	/// storage-level fault tests. Runtime bridges must use the lifecycle-owned API.
474	#[cfg(test)]
475	pub(crate) async fn commit_serialized_state_and_workflow_batch(
476		&self,
477		deltas: Vec<StateDelta>,
478		workflow_writes: Vec<WorkflowKvWrite>,
479	) -> Result<()> {
480		let save_request_revision = self.save_request_revision();
481		self.save_state_and_workflow_batch_with_revision(
482			deltas,
483			workflow_writes,
484			save_request_revision,
485		)
486		.await
487	}
488
489	pub(crate) fn request_save_with_revision(&self, opts: RequestSaveOpts) -> Result<u64> {
490		let immediate = opts.immediate;
491		let save_request_revision = self.0.save_request_revision.fetch_add(1, Ordering::SeqCst) + 1;
492		self.notify_request_save_hooks(opts);
493		let already_requested = self.0.save_requested.swap(true, Ordering::SeqCst);
494		let immediate_already_requested = if immediate {
495			self.0.save_requested_immediate.swap(true, Ordering::SeqCst)
496		} else {
497			self.0.save_requested_immediate.load(Ordering::SeqCst)
498		};
499
500		if let Some(max_wait_ms) = opts.max_wait_ms {
501			let deadline = StdInstant::now() + Duration::from_millis(u64::from(max_wait_ms));
502			let mut requested_deadline = self.0.save_requested_within_deadline.lock();
503			*requested_deadline = Some(match *requested_deadline {
504				Some(existing) => existing.min(deadline),
505				None => deadline,
506			});
507		}
508
509		let Some(sender) = self.lifecycle_event_sender() else {
510			return Err(ActorRuntime::NotConfigured {
511				component: "lifecycle events".to_owned(),
512			}
513			.build());
514		};
515
516		if opts.max_wait_ms.is_none()
517			&& already_requested
518			&& (!immediate || immediate_already_requested)
519		{
520			return Ok(save_request_revision);
521		}
522
523		sender
524			.send(LifecycleEvent::SaveRequested { immediate })
525			.map(|()| save_request_revision)
526			.map_err(|_| {
527				ActorRuntime::NotConfigured {
528					component: "lifecycle events".to_owned(),
529				}
530				.build()
531			})
532	}
533
534	pub(crate) async fn wait_for_save_request(&self, save_request_revision: u64) {
535		loop {
536			if self.0.save_completed_revision.load(Ordering::SeqCst) >= save_request_revision {
537				return;
538			}
539
540			self.0.save_completion.notified().await;
541		}
542	}
543
544	pub(crate) fn save_requested(&self) -> bool {
545		self.0.save_requested.load(Ordering::SeqCst)
546	}
547
548	pub(crate) fn save_requested_immediate(&self) -> bool {
549		self.0.save_requested_immediate.load(Ordering::SeqCst)
550	}
551
552	pub(crate) fn save_deadline(&self, immediate: bool) -> StdInstant {
553		self.compute_save_deadline(immediate)
554	}
555
556	pub(crate) fn compute_save_deadline(&self, immediate: bool) -> StdInstant {
557		if immediate || self.save_requested_immediate() {
558			return StdInstant::now();
559		}
560
561		let throttled_deadline = StdInstant::now() + self.compute_save_delay(None);
562		let requested_deadline = *self.0.save_requested_within_deadline.lock();
563
564		match requested_deadline {
565			Some(requested_deadline) => throttled_deadline.min(requested_deadline),
566			None => throttled_deadline,
567		}
568	}
569
570	pub(crate) fn save_request_revision(&self) -> u64 {
571		self.0.save_request_revision.load(Ordering::SeqCst)
572	}
573
574	pub(crate) async fn apply_state_deltas(
575		&self,
576		deltas: Vec<StateDelta>,
577		save_request_revision: u64,
578	) -> Result<()> {
579		self.apply_state_deltas_inner(deltas, None, save_request_revision, None)
580			.await
581			.map(|_| ())
582	}
583
584	pub(super) async fn apply_state_deltas_inner(
585		&self,
586		deltas: Vec<StateDelta>,
587		workflow_writes: Option<Vec<WorkflowKvWrite>>,
588		save_request_revision: u64,
589		expected_state_transaction_epoch: Option<u64>,
590	) -> Result<bool> {
591		let delta_count = deltas.len();
592		let delta_bytes: usize = deltas.iter().map(StateDelta::payload_len).sum();
593		let workflow_write_count = workflow_writes.as_ref().map_or(0, Vec::len);
594		let current_revision = self.0.state_revision.load(Ordering::SeqCst);
595		tracing::debug!(
596			delta_count,
597			delta_bytes,
598			state_revision = current_revision,
599			save_request_revision,
600			"applying actor state deltas"
601		);
602		self.clear_pending_save();
603
604		if deltas.is_empty() && workflow_write_count == 0 {
605			self.mark_save_request_completed(save_request_revision);
606			self.finish_save_request(save_request_revision);
607			tracing::debug!(
608				delta_count,
609				state_revision = current_revision,
610				save_request_revision,
611				"actor state deltas applied without kv write"
612			);
613			return Ok(true);
614		}
615
616		let prepared = {
617			let _save_guard = self.0.save_guard.lock().await;
618			if expected_state_transaction_epoch.is_some_and(|expected| {
619				self.0.state_transaction_epoch.load(Ordering::SeqCst) != expected
620			}) {
621				None
622			} else {
623				let revision = self.0.state_revision.load(Ordering::SeqCst);
624				let mut persisted = self.persisted();
625				let mut next_state = None;
626				let mut actor_to_persist = None;
627				let mut connections_to_persist: Vec<PersistedConnection> = Vec::new();
628				let mut connections_to_delete = Vec::new();
629
630				for delta in deltas {
631					match delta {
632						StateDelta::ActorState(bytes) => {
633							next_state = Some(bytes.clone());
634							persisted.state = bytes;
635						}
636						StateDelta::ConnHibernation { conn: _, bytes } => {
637							connections_to_persist.push(
638								decode_persisted_connection(&bytes)
639									.context("decode hibernatable connection state delta")?,
640							);
641						}
642						StateDelta::ConnHibernationRemoved(conn) => {
643							connections_to_delete.push(conn);
644						}
645					}
646				}
647
648				if next_state.is_some() {
649					actor_to_persist = Some(persisted.clone());
650				}
651
652				Some((
653					next_state,
654					actor_to_persist,
655					connections_to_persist,
656					connections_to_delete,
657					revision,
658					self.begin_write(),
659				))
660			}
661		};
662		let Some((
663			next_state,
664			actor_to_persist,
665			connections_to_persist,
666			connections_to_delete,
667			revision,
668			_write_guard,
669		)) = prepared
670		else {
671			tracing::debug!(
672				save_request_revision,
673				"discarding state serialized across a state transaction boundary"
674			);
675			return Ok(false);
676		};
677
678		if let Some(workflow_writes) = workflow_writes.as_deref() {
679			internal_storage::persist_actor_core_connections_and_workflow(
680				self.sql(),
681				actor_to_persist.as_ref(),
682				&connections_to_persist,
683				&connections_to_delete,
684				workflow_writes,
685			)
686			.await
687			.context("atomically persist actor state, connection deltas, and workflow kv")?;
688		} else if actor_to_persist.is_some()
689			|| !connections_to_persist.is_empty()
690			|| !connections_to_delete.is_empty()
691		{
692			internal_storage::persist_actor_core_and_connections(
693				self.sql(),
694				actor_to_persist.as_ref(),
695				&connections_to_persist,
696				&connections_to_delete,
697			)
698			.await
699			.context("persist actor state and connection deltas to sqlite")?;
700		}
701
702		if let Some(state) = next_state {
703			self.0.persisted.write().state = state.clone();
704			*self.0.current_state.write() = state;
705		}
706		for connection in &connections_to_persist {
707			if let Some(handle) = self.connection(&connection.id) {
708				handle.set_state_initial(connection.state.clone());
709			}
710		}
711
712		*self.0.last_save_at.lock() = Some(StdInstant::now());
713
714		if self.0.state_revision.load(Ordering::SeqCst) == revision {
715			self.0.state_dirty.store(false, Ordering::SeqCst);
716		}
717
718		self.mark_save_request_completed(save_request_revision);
719		self.finish_save_request(save_request_revision);
720		tracing::debug!(
721			delta_count,
722			delta_bytes,
723			workflow_write_count,
724			state_revision = self.0.state_revision.load(Ordering::SeqCst),
725			save_request_revision,
726			"actor state deltas applied"
727		);
728		Ok(true)
729	}
730	async fn commit_state_transaction(
731		&self,
732		transaction: &SqliteTransaction,
733		deltas: Vec<StateDelta>,
734		save_request_revision: u64,
735		statement_count: usize,
736		bind_payload_bytes: usize,
737	) -> Result<()> {
738		let (deltas, pending_hibernation_changes) = match self.prepare_state_deltas(deltas) {
739			Ok(prepared) => prepared,
740			Err(error) => {
741				let _ = transaction.rollback().await;
742				return Err(error);
743			}
744		};
745		let commit_result = async {
746			let revision = self.0.state_revision.load(Ordering::SeqCst);
747			let mut persisted = self.persisted();
748			let mut next_state = None;
749			let mut actor_to_persist = None;
750			let mut connections_to_persist: Vec<PersistedConnection> = Vec::new();
751			let mut connections_to_delete = Vec::new();
752
753			for delta in deltas {
754				match delta {
755					StateDelta::ActorState(bytes) => {
756						next_state = Some(bytes.clone());
757						persisted.state = bytes;
758					}
759					StateDelta::ConnHibernation { conn: _, bytes } => {
760						connections_to_persist.push(
761							decode_persisted_connection(&bytes)
762								.context("decode hibernatable connection state delta")?,
763						);
764					}
765					StateDelta::ConnHibernationRemoved(conn) => {
766						connections_to_delete.push(conn);
767					}
768				}
769			}
770
771			if next_state.is_some() {
772				actor_to_persist = Some(persisted);
773			}
774			let statements = internal_storage::build_actor_core_and_connection_statements(
775				actor_to_persist.as_ref(),
776				&connections_to_persist,
777				&connections_to_delete,
778			)?;
779			internal_storage::validate_atomic_state_transaction_budget(
780				statement_count.saturating_add(statements.len()),
781				bind_payload_bytes
782					.saturating_add(internal_storage::statement_bind_payload_len(&statements)),
783			)?;
784			for statement in statements {
785				transaction
786					.execute(statement.sql, statement.params)
787					.await
788					.context("persist actor state inside sqlite transaction")?;
789			}
790			transaction
791				.commit()
792				.await
793				.context("commit sqlite transaction with actor state")?;
794
795			if let Some(state) = next_state {
796				self.0.persisted.write().state = state.clone();
797				*self.0.current_state.write() = state;
798			}
799			for connection in &connections_to_persist {
800				if let Some(handle) = self.connection(&connection.id) {
801					handle.set_state_initial(connection.state.clone());
802				}
803			}
804			*self.0.last_save_at.lock() = Some(StdInstant::now());
805			if self.0.state_revision.load(Ordering::SeqCst) == revision {
806				self.0.state_dirty.store(false, Ordering::SeqCst);
807			}
808			self.mark_save_request_completed(save_request_revision);
809			self.finish_save_request(save_request_revision);
810			self.record_state_updated();
811			Ok(())
812		}
813		.await;
814
815		if let Err(error) = commit_result {
816			self.restore_pending_hibernation_changes(pending_hibernation_changes);
817			let _ = transaction.rollback().await;
818			return Err(error);
819		}
820		Ok(())
821	}
822
823	pub(crate) async fn wait_for_pending_writes(&self) {
824		loop {
825			if let Some(handle) = self.take_tracked_persist() {
826				let _ = handle.await;
827				continue;
828			}
829
830			let save_guard = self.0.save_guard.lock().await;
831			if self.has_tracked_persist() {
832				drop(save_guard);
833				continue;
834			}
835
836			if self.0.in_flight_state_writes.load(Ordering::SeqCst) == 0 {
837				return;
838			}
839			drop(save_guard);
840
841			self.wait_for_in_flight_writes().await;
842		}
843	}
844
845	pub(crate) async fn wait_for_pending_state_writes(&self) {
846		self.wait_for_pending_writes().await;
847	}
848
849	pub fn begin_on_state_change(&self) -> OnStateChangeGuard {
850		OnStateChangeGuard::new(self.clone())
851	}
852
853	pub fn on_state_change_started(&self) {
854		self.0
855			.on_state_change_in_flight
856			.fetch_add(1, Ordering::SeqCst);
857		self.0.sleep.work.keep_awake.increment();
858		self.reset_sleep_timer();
859	}
860
861	pub fn on_state_change_finished(&self) {
862		let previous = self.0.on_state_change_in_flight.fetch_update(
863			Ordering::SeqCst,
864			Ordering::SeqCst,
865			|count| count.checked_sub(1),
866		);
867
868		match previous {
869			Ok(1) => {
870				self.0.sleep.work.keep_awake.decrement();
871				self.0.on_state_change_idle.notify_waiters();
872				self.reset_sleep_timer();
873			}
874			Ok(_) => {
875				self.0.sleep.work.keep_awake.decrement();
876				self.reset_sleep_timer();
877			}
878			Err(_) => {
879				tracing::warn!(
880					actor_id = %self.actor_id(),
881					"on_state_change finished without a matching start"
882				);
883			}
884		}
885	}
886
887	#[cfg(test)]
888	#[allow(dead_code)]
889	pub(crate) async fn wait_for_on_state_change_idle(&self, timeout_duration: Duration) -> bool {
890		if self.0.on_state_change_in_flight.load(Ordering::SeqCst) == 0 {
891			return true;
892		}
893
894		timeout(timeout_duration, async {
895			loop {
896				let idle = self.0.on_state_change_idle.notified();
897				tokio::pin!(idle);
898				idle.as_mut().enable();
899
900				if self.0.on_state_change_in_flight.load(Ordering::SeqCst) == 0 {
901					return;
902				}
903
904				idle.await;
905			}
906		})
907		.await
908		.is_ok()
909	}
910
911	pub fn persisted(&self) -> PersistedActor {
912		self.0.persisted.read().clone()
913	}
914
915	pub fn load_persisted(&self, persisted: PersistedActor) {
916		let state = persisted.state.clone();
917		*self.0.persisted.write() = persisted;
918		*self.0.current_state.write() = state;
919		self.0.state_dirty.store(false, Ordering::SeqCst);
920		self.finish_save_request(self.save_request_revision());
921		self.0
922			.metrics
923			.inc_state_mutation(StateMutationReason::InternalReplace);
924	}
925
926	pub(crate) fn load_last_pushed_alarm(&self, alarm_ts: Option<i64>) {
927		*self.0.last_pushed_alarm.write() = alarm_ts;
928	}
929
930	pub(crate) fn last_pushed_alarm(&self) -> Option<i64> {
931		*self.0.last_pushed_alarm.read()
932	}
933
934	pub(crate) async fn persist_last_pushed_alarm(&self, alarm_ts: Option<i64>) -> Result<()> {
935		internal_storage::persist_last_pushed_alarm(self.sql(), alarm_ts)
936			.await
937			.context("persist last pushed alarm to sqlite")?;
938		self.load_last_pushed_alarm(alarm_ts);
939		Ok(())
940	}
941
942	pub(crate) fn load_run_wake_at(&self, wake_at: Option<i64>) {
943		*self.0.run_wake_at.write() = wake_at;
944	}
945
946	pub(crate) fn run_wake_at(&self) -> Option<i64> {
947		*self.0.run_wake_at.read()
948	}
949
950	pub(crate) async fn persist_run_wake_at(&self, wake_at: Option<i64>) -> Result<()> {
951		internal_storage::persist_run_wake_at(self.sql(), wake_at)
952			.await
953			.context("persist run wake deadline to sqlite")?;
954		self.load_run_wake_at(wake_at);
955		Ok(())
956	}
957
958	pub(crate) fn set_initial_state(&self, state: Vec<u8>) {
959		*self.0.current_state.write() = state.clone();
960		self.0.persisted.write().state = state;
961		self.0.state_dirty.store(true, Ordering::SeqCst);
962		self.0.state_revision.fetch_add(1, Ordering::SeqCst);
963	}
964
965	pub fn scheduled_events(&self) -> Vec<PersistedScheduleEvent> {
966		self.0.persisted.read().scheduled_events.clone()
967	}
968
969	pub fn set_scheduled_events(&self, scheduled_events: Vec<PersistedScheduleEvent>) {
970		self.0.persisted.write().scheduled_events = scheduled_events;
971		self.0
972			.metrics
973			.inc_state_mutation(StateMutationReason::ScheduledEventsUpdate);
974		self.mark_dirty();
975		self.schedule_save(None);
976	}
977
978	pub fn set_input(&self, input: Option<Vec<u8>>) {
979		self.0.persisted.write().input = input;
980		self.0
981			.metrics
982			.inc_state_mutation(StateMutationReason::InputSet);
983		self.mark_dirty();
984		self.schedule_save(None);
985	}
986
987	pub fn input(&self) -> Option<Vec<u8>> {
988		self.0.persisted.read().input.clone()
989	}
990
991	pub fn set_has_initialized(&self, has_initialized: bool) {
992		{
993			let mut persisted = self.0.persisted.write();
994			if persisted.has_initialized == has_initialized {
995				return;
996			}
997			persisted.has_initialized = has_initialized;
998		}
999		self.0
1000			.metrics
1001			.inc_state_mutation(StateMutationReason::HasInitialized);
1002		self.mark_dirty();
1003		self.schedule_save(None);
1004	}
1005
1006	pub fn has_initialized(&self) -> bool {
1007		self.0.persisted.read().has_initialized
1008	}
1009
1010	pub fn flush_on_shutdown(&self) {
1011		self.persist_now_tracked("shutdown_flush");
1012	}
1013
1014	pub fn on_request_save(&self, hook: Box<dyn Fn(RequestSaveOpts) + Send + Sync>) {
1015		self.0.request_save_hooks.write().push(Arc::from(hook));
1016	}
1017
1018	fn is_dirty(&self) -> bool {
1019		self.0.state_dirty.load(Ordering::SeqCst)
1020	}
1021
1022	fn mark_dirty(&self) {
1023		self.0.state_dirty.store(true, Ordering::SeqCst);
1024		self.0.state_revision.fetch_add(1, Ordering::SeqCst);
1025	}
1026
1027	fn lifecycle_event_sender(&self) -> Option<mpsc::UnboundedSender<LifecycleEvent>> {
1028		self.0.lifecycle_events.read().clone()
1029	}
1030
1031	fn compute_save_delay(&self, max_wait: Option<Duration>) -> Duration {
1032		let elapsed = self
1033			.0
1034			.last_save_at
1035			.lock()
1036			.map(|instant| instant.elapsed())
1037			.unwrap_or_default();
1038
1039		throttled_save_delay(self.0.state_save_interval, elapsed, max_wait)
1040	}
1041
1042	fn schedule_save(&self, max_wait: Option<Duration>) {
1043		if !self.is_dirty() {
1044			return;
1045		}
1046
1047		let delay = self.compute_save_delay(max_wait);
1048		let scheduled_at = StdInstant::now() + delay;
1049
1050		let mut pending_save = self.0.pending_save.lock();
1051
1052		if let Some(existing) = pending_save.as_ref() {
1053			if existing.scheduled_at <= scheduled_at {
1054				return;
1055			}
1056
1057			existing.handle.abort();
1058		}
1059
1060		let state = self.clone();
1061		// Intentionally detached but abortable: pending delayed saves are
1062		// retained in `pending_save`, replaced by newer saves, and awaited at
1063		// shutdown through the state save guard.
1064		let task = async move {
1065			if !delay.is_zero() {
1066				sleep(delay).await;
1067			}
1068
1069			state.take_pending_save();
1070
1071			if let Err(error) = state.persist_if_dirty().await {
1072				tracing::error!(?error, "failed to persist actor state");
1073			}
1074		}
1075		.in_current_span();
1076
1077		#[cfg(not(feature = "wasm-runtime"))]
1078		let handle = {
1079			let Ok(tokio_handle) = Handle::try_current() else {
1080				return;
1081			};
1082			tokio_handle.spawn(task)
1083		};
1084
1085		#[cfg(feature = "wasm-runtime")]
1086		let handle = RuntimeSpawner::spawn(task);
1087
1088		*pending_save = Some(PendingSave {
1089			scheduled_at,
1090			handle,
1091		});
1092	}
1093
1094	fn clear_pending_save(&self) {
1095		if let Some(pending_save) = self.take_pending_save() {
1096			pending_save.handle.abort();
1097		}
1098	}
1099
1100	pub(crate) fn persist_now_tracked(&self, description: &'static str) {
1101		self.clear_pending_save();
1102
1103		let state = self.clone();
1104		let mut tracked_persist = self.0.tracked_persist.lock();
1105		let previous = tracked_persist.take();
1106		let task = async move {
1107			if let Some(previous) = previous {
1108				let _ = previous.await;
1109			}
1110
1111			if let Err(error) = state.persist_state(SaveStateOpts { immediate: true }).await {
1112				tracing::error!(?error, description, "failed to persist actor state");
1113			}
1114		}
1115		.in_current_span();
1116
1117		#[cfg(not(feature = "wasm-runtime"))]
1118		let handle = {
1119			let Ok(tokio_handle) = Handle::try_current() else {
1120				tracing::warn!(
1121					description,
1122					"skipping tracked actor state persistence without runtime"
1123				);
1124				return;
1125			};
1126			tokio_handle.spawn(task)
1127		};
1128
1129		#[cfg(feature = "wasm-runtime")]
1130		let handle = RuntimeSpawner::spawn(task);
1131
1132		*tracked_persist = Some(handle);
1133	}
1134
1135	fn take_pending_save(&self) -> Option<PendingSave> {
1136		self.0.pending_save.lock().take()
1137	}
1138
1139	fn take_tracked_persist(&self) -> Option<JoinHandle<()>> {
1140		self.0.tracked_persist.lock().take()
1141	}
1142
1143	fn has_tracked_persist(&self) -> bool {
1144		self.0.tracked_persist.lock().is_some()
1145	}
1146
1147	#[cfg(test)]
1148	pub(crate) fn tracked_persist_pending(&self) -> bool {
1149		self.has_tracked_persist()
1150	}
1151
1152	async fn persist_if_dirty(&self) -> Result<()> {
1153		if !self.is_dirty() {
1154			return Ok(());
1155		}
1156
1157		let (revision, actor_to_persist, _write_guard) = {
1158			let _save_guard = self.0.save_guard.lock().await;
1159			if !self.is_dirty() {
1160				return Ok(());
1161			}
1162
1163			let revision = self.0.state_revision.load(Ordering::SeqCst);
1164			let persisted = self.persisted();
1165			(revision, persisted, self.begin_write())
1166		};
1167
1168		internal_storage::persist_actor_snapshot(self.sql(), &actor_to_persist)
1169			.await
1170			.context("persist actor state to sqlite")?;
1171
1172		*self.0.last_save_at.lock() = Some(StdInstant::now());
1173
1174		if self.0.state_revision.load(Ordering::SeqCst) == revision {
1175			self.0.state_dirty.store(false, Ordering::SeqCst);
1176		}
1177
1178		Ok(())
1179	}
1180
1181	fn begin_write(&self) -> InFlightWrite {
1182		self.0.in_flight_state_writes.fetch_add(1, Ordering::SeqCst);
1183		InFlightWrite { ctx: self.clone() }
1184	}
1185
1186	async fn wait_for_in_flight_writes(&self) {
1187		loop {
1188			if self.0.in_flight_state_writes.load(Ordering::SeqCst) == 0 {
1189				return;
1190			}
1191			self.0.state_write_completion.notified().await;
1192		}
1193	}
1194
1195	fn finish_save_request(&self, save_request_revision: u64) {
1196		if self.0.save_request_revision.load(Ordering::SeqCst) == save_request_revision {
1197			self.0.save_requested.store(false, Ordering::SeqCst);
1198			self.0
1199				.save_requested_immediate
1200				.store(false, Ordering::SeqCst);
1201			*self.0.save_requested_within_deadline.lock() = None;
1202		}
1203	}
1204
1205	fn mark_save_request_completed(&self, save_request_revision: u64) {
1206		self.0
1207			.save_completed_revision
1208			.fetch_max(save_request_revision, Ordering::SeqCst);
1209		self.0.save_completion.notify_waiters();
1210	}
1211
1212	fn notify_request_save_hooks(&self, opts: RequestSaveOpts) {
1213		let hooks = self.0.request_save_hooks.read().clone();
1214		for hook in hooks {
1215			hook(opts);
1216		}
1217	}
1218}
1219
1220struct InFlightWrite {
1221	ctx: ActorContext,
1222}
1223
1224impl Drop for InFlightWrite {
1225	fn drop(&mut self) {
1226		if self
1227			.ctx
1228			.0
1229			.in_flight_state_writes
1230			.fetch_sub(1, Ordering::SeqCst)
1231			== 1
1232		{
1233			self.ctx.0.state_write_completion.notify_waiters();
1234			self.ctx.0.state_write_completion.notify_one();
1235		}
1236	}
1237}
1238
1239fn throttled_save_delay(
1240	save_interval: Duration,
1241	time_since_last_save: Duration,
1242	max_wait: Option<Duration>,
1243) -> Duration {
1244	let save_delay = save_interval.saturating_sub(time_since_last_save);
1245	if let Some(max_wait) = max_wait {
1246		save_delay.min(max_wait)
1247	} else {
1248		save_delay
1249	}
1250}
1251
1252// Test shim keeps moved tests in crate-root tests/ with private-module access.
1253#[cfg(test)]
1254#[path = "../../tests/state.rs"]
1255mod tests;