Skip to main content

awaken_stores/memory/
pending.rs

1use std::collections::HashMap;
2use std::collections::HashSet;
3
4use async_trait::async_trait;
5use awaken_server_contract::contract::message::{
6    DeliveryBoundary, DeliveryMode, Message, MessageRecord, PendingMessageRecord,
7    pending_queue_revision, select_pending_for_freeze, select_pending_for_freeze_for_run,
8};
9use awaken_server_contract::contract::storage::{
10    RunRecord, StorageError, checkpoint_parent_thread_id, message_append,
11};
12use awaken_server_contract::thread::Thread;
13
14use crate::PendingMessageStore;
15use crate::pending_message_store::{
16    validate_pending_message_record, validate_pending_message_records,
17};
18
19use super::validate_thread_hierarchy_map;
20use super::{InMemoryStore, current_millis};
21
22fn normalize_pending_positions(pending: &mut [PendingMessageRecord]) {
23    for (index, record) in pending.iter_mut().enumerate() {
24        record.position = index as u64 + 1;
25    }
26}
27
28fn pending_not_found(thread_id: &str, pending_id: &str) -> StorageError {
29    StorageError::NotFound(format!(
30        "pending message '{pending_id}' in thread '{thread_id}'"
31    ))
32}
33
34fn already_consumed(pending_id: &str) -> StorageError {
35    StorageError::Validation(format!(
36        "pending message '{pending_id}' is already consumed"
37    ))
38}
39
40fn duplicate_pending_id(pending_id: &str) -> StorageError {
41    StorageError::Validation(format!("pending message '{pending_id}' already exists"))
42}
43
44fn selected_pending_ids(
45    pending: &[PendingMessageRecord],
46    selected_indexes: &[usize],
47) -> Vec<String> {
48    selected_indexes
49        .iter()
50        .map(|index| pending[*index].pending_id.clone())
51        .collect()
52}
53
54impl InMemoryStore {
55    async fn committed_message_exists(
56        &self,
57        thread_id: &str,
58        message_id: &str,
59    ) -> Result<bool, StorageError> {
60        let guard = self.messages.read().await;
61        Ok(guard.get(thread_id).is_some_and(|messages| {
62            messages
63                .iter()
64                .any(|message| message.id.as_deref() == Some(message_id))
65        }))
66    }
67}
68
69#[async_trait]
70impl PendingMessageStore for InMemoryStore {
71    async fn load_pending_message_records(
72        &self,
73        thread_id: &str,
74    ) -> Result<Vec<PendingMessageRecord>, StorageError> {
75        let guard = self.pending_messages.read().await;
76        let records = guard.get(thread_id).cloned().unwrap_or_default();
77        validate_pending_message_records(&records)?;
78        Ok(records)
79    }
80
81    async fn list_threads_with_pending_messages(
82        &self,
83        limit: usize,
84        after: Option<&str>,
85    ) -> Result<Vec<String>, StorageError> {
86        let guard = self.pending_messages.read().await;
87        let mut ids: Vec<String> = guard
88            .iter()
89            .filter(|(_, records)| !records.is_empty())
90            .map(|(thread_id, _)| thread_id.clone())
91            .filter(|thread_id| after.is_none_or(|cursor| thread_id.as_str() > cursor))
92            .collect();
93        ids.sort();
94        if limit > 0 {
95            ids.truncate(limit);
96        }
97        Ok(ids)
98    }
99
100    async fn append_pending_message_records(
101        &self,
102        thread_id: &str,
103        messages: &[Message],
104        delivery_mode: DeliveryMode,
105    ) -> Result<Vec<PendingMessageRecord>, StorageError> {
106        let now = current_millis() / 1000;
107        let committed_guard = self.messages.read().await;
108        let mut guard = self.pending_messages.write().await;
109        let pending = guard.entry(thread_id.to_owned()).or_default();
110        let start_position = pending.len() as u64 + 1;
111        let records = messages
112            .iter()
113            .cloned()
114            .enumerate()
115            .map(|(index, message)| {
116                let mut record = PendingMessageRecord::from_message(
117                    thread_id.to_owned(),
118                    start_position + index as u64,
119                    message,
120                    delivery_mode.clone(),
121                );
122                record.created_at = Some(now);
123                record.updated_at = Some(now);
124                record
125            })
126            .collect::<Vec<_>>();
127        let mut seen = pending
128            .iter()
129            .map(|record| record.pending_id.as_str())
130            .collect::<HashSet<_>>();
131        for record in &records {
132            validate_pending_message_record(record)?;
133            if !seen.insert(record.pending_id.as_str()) {
134                return Err(duplicate_pending_id(&record.pending_id));
135            }
136            if committed_guard.get(thread_id).is_some_and(|committed| {
137                committed
138                    .iter()
139                    .any(|message| message.id.as_deref() == Some(record.pending_id.as_str()))
140            }) {
141                return Err(already_consumed(&record.pending_id));
142            }
143        }
144        pending.extend(records.iter().cloned());
145        Ok(records)
146    }
147
148    async fn update_pending_message_record_checked(
149        &self,
150        thread_id: &str,
151        pending_id: &str,
152        expected_revision: Option<u64>,
153        mut message: Message,
154    ) -> Result<PendingMessageRecord, StorageError> {
155        let mut guard = self.pending_messages.write().await;
156        if let Some(pending) = guard.get_mut(thread_id)
157            && let Some(record) = pending
158                .iter_mut()
159                .find(|record| record.pending_id == pending_id)
160        {
161            if let Some(expected) = expected_revision
162                && record.revision != expected
163            {
164                return Err(StorageError::VersionConflict {
165                    expected,
166                    actual: record.revision,
167                });
168            }
169            match message.id.as_deref() {
170                Some(message_id) if message_id != pending_id => {
171                    return Err(StorageError::Validation(format!(
172                        "pending message '{pending_id}' cannot change message id to '{message_id}'"
173                    )));
174                }
175                Some(_) => {}
176                None => message.id = Some(pending_id.to_owned()),
177            }
178            let mut updated = record.clone();
179            updated.message = message;
180            validate_pending_message_record(&updated)?;
181            record.message = updated.message;
182            record.revision += 1;
183            record.updated_at = Some(current_millis() / 1000);
184            return Ok(record.clone());
185        }
186        drop(guard);
187        if self.committed_message_exists(thread_id, pending_id).await? {
188            return Err(already_consumed(pending_id));
189        }
190        Err(pending_not_found(thread_id, pending_id))
191    }
192
193    async fn retract_pending_message_record_checked(
194        &self,
195        thread_id: &str,
196        pending_id: &str,
197        expected_revision: Option<u64>,
198    ) -> Result<PendingMessageRecord, StorageError> {
199        let mut guard = self.pending_messages.write().await;
200        if let Some(pending) = guard.get_mut(thread_id)
201            && let Some(index) = pending
202                .iter()
203                .position(|record| record.pending_id == pending_id)
204        {
205            if let Some(expected) = expected_revision
206                && pending[index].revision != expected
207            {
208                return Err(StorageError::VersionConflict {
209                    expected,
210                    actual: pending[index].revision,
211                });
212            }
213            let removed = pending.remove(index);
214            normalize_pending_positions(pending);
215            return Ok(removed);
216        }
217        drop(guard);
218        if self.committed_message_exists(thread_id, pending_id).await? {
219            return Err(already_consumed(pending_id));
220        }
221        Err(pending_not_found(thread_id, pending_id))
222    }
223
224    async fn reorder_pending_message_records_checked(
225        &self,
226        thread_id: &str,
227        expected_queue_revision: Option<u64>,
228        ordered_pending_ids: &[String],
229    ) -> Result<Vec<PendingMessageRecord>, StorageError> {
230        let mut guard = self.pending_messages.write().await;
231        let Some(pending) = guard.get_mut(thread_id) else {
232            drop(guard);
233            for pending_id in ordered_pending_ids {
234                if self.committed_message_exists(thread_id, pending_id).await? {
235                    return Err(already_consumed(pending_id));
236                }
237            }
238            return Err(StorageError::NotFound(thread_id.to_owned()));
239        };
240        let actual_queue_revision = pending_queue_revision(pending);
241        if let Some(expected) = expected_queue_revision
242            && expected != actual_queue_revision
243        {
244            return Err(StorageError::VersionConflict {
245                expected,
246                actual: actual_queue_revision,
247            });
248        }
249        let pending_ids = pending
250            .iter()
251            .map(|record| record.pending_id.as_str())
252            .collect::<HashSet<_>>();
253        let consumed_id = ordered_pending_ids
254            .iter()
255            .find(|pending_id| !pending_ids.contains(pending_id.as_str()))
256            .cloned();
257        if let Some(pending_id) = consumed_id {
258            drop(guard);
259            if self
260                .committed_message_exists(thread_id, &pending_id)
261                .await?
262            {
263                return Err(already_consumed(&pending_id));
264            }
265            return Err(StorageError::NotFound(pending_id));
266        }
267        if pending.len() != ordered_pending_ids.len() {
268            return Err(StorageError::VersionConflict {
269                expected: ordered_pending_ids.len() as u64,
270                actual: pending.len() as u64,
271            });
272        }
273        let mut by_id = pending
274            .iter()
275            .cloned()
276            .map(|record| (record.pending_id.clone(), record))
277            .collect::<HashMap<_, _>>();
278        let mut reordered = Vec::with_capacity(ordered_pending_ids.len());
279        for pending_id in ordered_pending_ids {
280            let record = by_id
281                .remove(pending_id)
282                .ok_or_else(|| StorageError::NotFound(pending_id.clone()))?;
283            reordered.push(record);
284        }
285        if !by_id.is_empty() {
286            return Err(StorageError::Validation(format!(
287                "reorder for thread '{thread_id}' omitted pending ids"
288            )));
289        }
290        let now = current_millis() / 1000;
291        normalize_pending_positions(&mut reordered);
292        for record in &mut reordered {
293            record.revision += 1;
294            record.updated_at = Some(now);
295        }
296        *pending = reordered.clone();
297        Ok(reordered)
298    }
299
300    async fn freeze_pending_message_records(
301        &self,
302        thread_id: &str,
303        boundary: DeliveryBoundary,
304        expected_message_version: Option<u64>,
305    ) -> Result<Vec<MessageRecord>, StorageError> {
306        let mut messages_guard = self.messages.write().await;
307        let mut pending_guard = self.pending_messages.write().await;
308        let committed = messages_guard.entry(thread_id.to_owned()).or_default();
309        let actual = committed.len() as u64;
310        if let Some(expected) = expected_message_version
311            && expected != actual
312        {
313            return Err(StorageError::VersionConflict { expected, actual });
314        }
315        let Some(pending) = pending_guard.get_mut(thread_id) else {
316            return Ok(Vec::new());
317        };
318        let selected_indexes = select_pending_for_freeze(pending, boundary);
319        if selected_indexes.is_empty() {
320            return Ok(Vec::new());
321        }
322        let selected_messages = selected_indexes
323            .iter()
324            .map(|index| pending[*index].message.clone())
325            .collect::<Vec<_>>();
326        message_append::validate_append_only_delta(committed, &selected_messages)?;
327
328        let mut selected = Vec::with_capacity(selected_indexes.len());
329        for index in selected_indexes.iter().rev() {
330            selected.push(pending.remove(*index));
331        }
332        selected.reverse();
333        normalize_pending_positions(pending);
334        let start_seq = committed.len() as u64 + 1;
335        let appended = selected
336            .into_iter()
337            .enumerate()
338            .map(|(index, record)| {
339                let message = record.message;
340                committed.push(message.clone());
341                MessageRecord::from_message(thread_id.to_owned(), start_seq + index as u64, message)
342            })
343            .collect();
344        Ok(appended)
345    }
346
347    async fn freeze_pending_message_records_with_run(
348        &self,
349        thread_id: &str,
350        boundary: DeliveryBoundary,
351        expected_message_version: Option<u64>,
352        expected_pending_ids: &[String],
353        run: &RunRecord,
354    ) -> Result<Vec<MessageRecord>, StorageError> {
355        run.validate_for_persist()?;
356        let now = current_millis();
357        let mut thread_guard = self.threads.write().await;
358        let existing_thread = thread_guard.get(thread_id).cloned();
359        validate_thread_hierarchy_map(
360            &thread_guard,
361            thread_id,
362            checkpoint_parent_thread_id(existing_thread.as_ref(), run),
363        )?;
364        let mut messages_guard = self.messages.write().await;
365        let mut pending_guard = self.pending_messages.write().await;
366        let mut run_guard = self.runs.write().await;
367        let actual = messages_guard
368            .get(thread_id)
369            .map(|messages| messages.len() as u64)
370            .unwrap_or(0);
371        if let Some(expected) = expected_message_version
372            && expected != actual
373        {
374            return Err(StorageError::VersionConflict { expected, actual });
375        }
376        let pending = pending_guard.entry(thread_id.to_owned()).or_default();
377        let selected_indexes =
378            select_pending_for_freeze_for_run(pending, boundary, Some(&run.run_id));
379        let selected_ids = selected_pending_ids(pending, &selected_indexes);
380        if selected_ids != expected_pending_ids {
381            return Err(StorageError::PendingSelectionConflict {
382                expected_ids: expected_pending_ids.to_vec(),
383                actual_ids: selected_ids,
384            });
385        }
386        let committed = messages_guard.entry(thread_id.to_owned()).or_default();
387        let selected_messages = selected_indexes
388            .iter()
389            .map(|index| pending[*index].message.clone())
390            .collect::<Vec<_>>();
391        message_append::validate_append_only_delta(committed, &selected_messages)?;
392
393        let mut selected = Vec::with_capacity(selected_indexes.len());
394        for index in selected_indexes.iter().rev() {
395            selected.push(pending.remove(*index));
396        }
397        selected.reverse();
398        normalize_pending_positions(pending);
399        let start_seq = committed.len() as u64 + 1;
400        let appended = selected
401            .into_iter()
402            .enumerate()
403            .map(|(index, record)| {
404                let message = record.message;
405                committed.push(message.clone());
406                MessageRecord::from_message(thread_id.to_owned(), start_seq + index as u64, message)
407            })
408            .collect::<Vec<_>>();
409        let mut thread = existing_thread.unwrap_or_else(|| Thread::with_id(thread_id));
410        thread.touch(now);
411        thread.apply_run_projection(run);
412        thread.normalize_lineage();
413        thread_guard.insert(thread_id.to_owned(), thread);
414        run_guard.insert(run.run_id.clone(), run.clone());
415        self.run_insertion
416            .write()
417            .await
418            .insert(run.run_id.clone(), self.next_run_seq());
419        Ok(appended)
420    }
421
422    async fn append_and_freeze_pending_message_records_with_run(
423        &self,
424        thread_id: &str,
425        new_messages: &[Message],
426        append_delivery_mode: DeliveryMode,
427        boundary: DeliveryBoundary,
428        expected_message_version: Option<u64>,
429        expected_pending_ids: &[String],
430        run: &RunRecord,
431    ) -> Result<Vec<MessageRecord>, StorageError> {
432        run.validate_for_persist()?;
433        let now = current_millis();
434        let mut thread_guard = self.threads.write().await;
435        let existing_thread = thread_guard.get(thread_id).cloned();
436        validate_thread_hierarchy_map(
437            &thread_guard,
438            thread_id,
439            checkpoint_parent_thread_id(existing_thread.as_ref(), run),
440        )?;
441        let mut messages_guard = self.messages.write().await;
442        let mut pending_guard = self.pending_messages.write().await;
443        let mut run_guard = self.runs.write().await;
444        let actual = messages_guard
445            .get(thread_id)
446            .map(|messages| messages.len() as u64)
447            .unwrap_or(0);
448        if let Some(expected) = expected_message_version
449            && expected != actual
450        {
451            return Err(StorageError::VersionConflict { expected, actual });
452        }
453        let pending = pending_guard.entry(thread_id.to_owned()).or_default();
454        // Append the new messages to pending *under the same held locks* as the
455        // freeze below, so the two are one atomic boundary (ADR-0042 D7).
456        let append_now = now / 1000;
457        let start_position = pending.len() as u64 + 1;
458        let mut seen = pending
459            .iter()
460            .map(|record| record.pending_id.clone())
461            .collect::<HashSet<_>>();
462        let mut appended_pending = Vec::with_capacity(new_messages.len());
463        for (index, message) in new_messages.iter().cloned().enumerate() {
464            let mut record = PendingMessageRecord::from_message(
465                thread_id.to_owned(),
466                start_position + index as u64,
467                message,
468                append_delivery_mode.clone(),
469            );
470            record.created_at = Some(append_now);
471            record.updated_at = Some(append_now);
472            validate_pending_message_record(&record)?;
473            if !seen.insert(record.pending_id.clone()) {
474                return Err(duplicate_pending_id(&record.pending_id));
475            }
476            if messages_guard.get(thread_id).is_some_and(|committed| {
477                committed
478                    .iter()
479                    .any(|message| message.id.as_deref() == Some(record.pending_id.as_str()))
480            }) {
481                return Err(already_consumed(&record.pending_id));
482            }
483            appended_pending.push(record);
484        }
485        pending.extend(appended_pending);
486
487        // Freeze the caller's selection over the now-appended pending. Mirrors
488        // `freeze_pending_message_records_with_run`, but with append folded into
489        // the same locked region.
490        let selected_indexes =
491            select_pending_for_freeze_for_run(pending, boundary, Some(&run.run_id));
492        let selected_ids = selected_pending_ids(pending, &selected_indexes);
493        if selected_ids != expected_pending_ids {
494            return Err(StorageError::PendingSelectionConflict {
495                expected_ids: expected_pending_ids.to_vec(),
496                actual_ids: selected_ids,
497            });
498        }
499        let committed = messages_guard.entry(thread_id.to_owned()).or_default();
500        let selected_messages = selected_indexes
501            .iter()
502            .map(|index| pending[*index].message.clone())
503            .collect::<Vec<_>>();
504        message_append::validate_append_only_delta(committed, &selected_messages)?;
505        let mut selected = Vec::with_capacity(selected_indexes.len());
506        for index in selected_indexes.iter().rev() {
507            selected.push(pending.remove(*index));
508        }
509        selected.reverse();
510        normalize_pending_positions(pending);
511        let start_seq = committed.len() as u64 + 1;
512        let appended = selected
513            .into_iter()
514            .enumerate()
515            .map(|(index, record)| {
516                let message = record.message;
517                committed.push(message.clone());
518                MessageRecord::from_message(thread_id.to_owned(), start_seq + index as u64, message)
519            })
520            .collect::<Vec<_>>();
521        let mut thread = existing_thread.unwrap_or_else(|| Thread::with_id(thread_id));
522        thread.touch(now);
523        thread.apply_run_projection(run);
524        thread.normalize_lineage();
525        thread_guard.insert(thread_id.to_owned(), thread);
526        run_guard.insert(run.run_id.clone(), run.clone());
527        self.run_insertion
528            .write()
529            .await
530            .insert(run.run_id.clone(), self.next_run_seq());
531        Ok(appended)
532    }
533}