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#[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 pub fn set_state_initial(&self, state: Vec<u8>) {
373 self.set_initial_state(state);
374 }
375
376 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 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 #[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 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#[cfg(test)]
1254#[path = "../../tests/state.rs"]
1255mod tests;