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 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 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}