1use std::collections::HashSet;
2
3use candid::CandidType;
4use serde::{Deserialize, Serialize};
5
6use crate::{
7 identity::CallerId,
8 types::{PageData, QueryPage, QueryPageError},
9};
10
11#[derive(CandidType, Serialize, Deserialize, Debug, Clone, Copy, Default, PartialEq, Eq, PartialOrd, Ord, Hash)]
13pub struct RecordId(u64);
14
15impl From<u64> for RecordId {
16 fn from(value: u64) -> Self {
17 Self(value)
18 }
19}
20
21impl RecordId {
22 pub fn into_inner(&self) -> u64 {
24 self.0
25 }
26 pub fn next(self) -> Self {
28 Self(self.0 + 1)
29 }
30}
31
32pub trait Searchable<Record> {
34 fn test(&self, record: &Record) -> bool;
36}
37
38pub trait Recordable<Record, RecordTopic, Search: Searchable<Record>> {
40 fn record_find_all(&self) -> &[Record];
43
44 fn record_push(&mut self, caller: CallerId, topic: RecordTopic, content: String) -> RecordId;
47 fn record_update(&mut self, record_id: RecordId, result: String);
49 fn record_delete(&mut self, ids: &HashSet<RecordId>) -> u64;
51
52 fn record_find_by_page(
54 &self,
55 page: &QueryPage,
56 max_page_size: u32,
57 search: &Option<Search>,
58 ) -> Result<PageData<&Record>, QueryPageError> {
59 let list = self.record_find_all();
60 if let Some(search) = search {
61 return page.query_desc_by_list_and_filter(list, max_page_size, |item| search.test(item));
62 }
63 page.query_desc_by_list(list, max_page_size)
64 }
65}
66
67pub mod basic {
71 use std::collections::HashSet;
72
73 use candid::CandidType;
74 use serde::{Deserialize, Serialize};
75
76 use crate::{
77 functions::types::{RecordId, Recordable, Searchable},
78 identity::CallerId,
79 types::TimestampNanos,
80 };
81
82 pub type RecordTopic = u8;
84
85 #[derive(CandidType, Serialize, Deserialize, Debug, Clone)]
87 pub struct Record {
88 pub id: RecordId,
90 pub created: TimestampNanos,
92 pub caller: CallerId,
94 pub topic: RecordTopic,
96 pub content: String,
98 #[serde(alias = "done")]
100 pub completion: Option<(TimestampNanos, String)>,
101 }
102
103 impl Record {
104 #[inline]
105 fn same(&self, id: &RecordId) -> bool {
106 self.id == *id
107 }
108 }
109
110 #[derive(CandidType, Serialize, Deserialize, Debug, Clone)]
112 pub struct RecordSearch {
113 #[serde(alias = "id")]
115 pub id_range: Option<(Option<RecordId>, Option<RecordId>)>,
116 #[serde(alias = "created")]
118 pub created_at_nanos_range: Option<(Option<TimestampNanos>, Option<TimestampNanos>)>,
119 pub caller: Option<HashSet<CallerId>>,
121 pub topic: Option<HashSet<RecordTopic>>,
123 pub content: Option<String>,
125 }
126
127 impl Searchable<Record> for RecordSearch {
128 #[allow(unused)]
129 #[inline]
130 fn test(&self, record: &Record) -> bool {
131 if let Some((id_min, id_max)) = &self.id_range {
132 if let Some(id_min) = &id_min
133 && record.id < *id_min
134 {
135 return false;
136 }
137 if let Some(id_max) = &id_max
138 && *id_max < record.id
139 {
140 return false;
141 }
142 }
143 if let Some(created) = self.created_at_nanos_range {
144 let (created_min, created_max) = created;
145 if let Some(created_min) = created_min
146 && record.created < created_min
147 {
148 return false;
149 }
150 if let Some(created_max) = created_max
151 && created_max < record.created
152 {
153 return false;
154 }
155 }
156 if let Some(caller) = &self.caller
157 && !caller.contains(&record.caller)
158 {
159 return false;
160 }
161 if let Some(topic) = &self.topic
162 && !topic.contains(&record.topic)
163 {
164 return false;
165 }
166 if let Some(content) = &self.content
167 && !record.content.contains(content)
168 {
169 return false;
170 }
171 true
172 }
173 }
174
175 #[derive(CandidType, Serialize, Deserialize, Debug, Clone)]
177 pub struct Records {
178 #[serde(alias = "max")]
180 pub retention_limit: u64,
181 #[serde(alias = "removed")]
183 pub retention_evicted_count: u64,
184 pub next_id: RecordId,
186 pub records: Vec<Record>,
188 }
189
190 impl Default for Records {
191 fn default() -> Self {
192 Self {
193 retention_limit: 1024 * 64, retention_evicted_count: Default::default(),
195 next_id: Default::default(),
196 records: Default::default(),
197 }
198 }
199 }
200
201 impl Records {
202 fn push_at(
203 &mut self,
204 caller: CallerId,
205 topic: RecordTopic,
206 content: String,
207 created: TimestampNanos,
208 ) -> RecordId {
209 let id = self.next_id;
210 self.next_id = self.next_id.next();
211
212 if self.retention_limit == 0 {
213 self.retention_evicted_count = self.retention_evicted_count.saturating_add(1);
214 return id;
215 }
216
217 let retention_limit = usize::try_from(self.retention_limit).unwrap_or(usize::MAX);
218 if retention_limit <= self.records.len() {
219 let remove_count = self.records.len() - retention_limit + 1;
220 self.records.drain(..remove_count);
221 self.retention_evicted_count = self.retention_evicted_count.saturating_add(remove_count as u64);
222 }
223
224 self.records.push(Record {
225 id,
226 created,
227 caller,
228 topic,
229 content,
230 completion: None,
231 });
232
233 id
234 }
235
236 fn update_at(&mut self, record_id: RecordId, result: String, completed_at: TimestampNanos) {
237 if let Some(item) = self.records.iter_mut().rev().find(|item| item.same(&record_id)) {
238 item.completion = Some((completed_at, result));
239 }
240 }
241 }
242
243 impl Recordable<Record, RecordTopic, RecordSearch> for Records {
244 fn record_find_all(&self) -> &[Record] {
248 &self.records
249 }
250
251 fn record_push(&mut self, caller: CallerId, topic: RecordTopic, content: String) -> RecordId {
253 self.push_at(caller, topic, content, crate::times::now())
254 }
255
256 fn record_update(&mut self, record_id: RecordId, result: String) {
258 self.update_at(record_id, result, crate::times::now());
259 }
260
261 fn record_delete(&mut self, ids: &HashSet<RecordId>) -> u64 {
263 if ids.is_empty() {
264 return 0;
265 }
266
267 let before = self.records.len();
268 self.records.retain(|record| !ids.contains(&record.id));
269 u64::try_from(before - self.records.len()).unwrap_or(u64::MAX)
270 }
271 }
272
273 #[derive(CandidType, Serialize, Deserialize, Debug, Clone)]
275 pub struct RecordSearchArg {
276 #[serde(alias = "id")]
278 pub id_range: Option<(Option<u64>, Option<u64>)>,
279 #[serde(alias = "created")]
281 pub created_at_nanos_range: Option<(Option<u64>, Option<u64>)>,
282 pub caller: Option<HashSet<CallerId>>,
284 pub topic: Option<HashSet<String>>,
286 pub content: Option<String>,
288 }
289
290 impl RecordSearchArg {
291 pub fn into<E, F: Fn(&str) -> Result<RecordTopic, E>>(self, f: F) -> Result<RecordSearch, E> {
293 Ok(RecordSearch {
294 id_range: self.id_range.map(|(a, b)| (a.map(|a| a.into()), b.map(|b| b.into()))),
295 created_at_nanos_range: self
296 .created_at_nanos_range
297 .map(|(a, b)| (a.map(|a| (a as i128).into()), b.map(|b| (b as i128).into()))),
298 caller: self.caller,
299 topic: self
300 .topic
301 .map(|topic| topic.iter().map(|t| f(t)).collect::<Result<HashSet<_>, _>>())
302 .transpose()?,
303 content: self.content,
304 })
305 }
306 }
307
308 #[cfg(test)]
309 mod tests {
310 use std::collections::HashSet;
311
312 use candid::Principal;
313 use ciborium::value::Value;
314 use serde::Serialize;
315
316 use super::{RecordSearchArg, Records};
317 use crate::{
318 functions::{record::RecordId, types::Recordable},
319 types::TimestampNanos,
320 };
321
322 #[derive(Serialize)]
323 struct LegacyRecord {
324 id: RecordId,
325 created: TimestampNanos,
326 caller: Principal,
327 topic: u8,
328 content: String,
329 done: Option<(TimestampNanos, String)>,
330 }
331
332 #[derive(Serialize)]
333 struct LegacyRecords {
334 max: u64,
335 removed: u64,
336 next_id: RecordId,
337 records: Vec<LegacyRecord>,
338 }
339
340 #[derive(Serialize)]
341 struct LegacyRecordSearchArg {
342 id: Option<(Option<u64>, Option<u64>)>,
343 created: Option<(Option<u64>, Option<u64>)>,
344 caller: Option<HashSet<Principal>>,
345 topic: Option<HashSet<String>>,
346 content: Option<String>,
347 }
348
349 fn map_keys(value: &Value) -> Vec<&str> {
350 let Value::Map(entries) = value else {
351 panic!("expected a CBOR map")
352 };
353 entries
354 .iter()
355 .filter_map(|(key, _)| match key {
356 Value::Text(key) => Some(key.as_str()),
357 _ => None,
358 })
359 .collect()
360 }
361
362 fn push(records: &mut Records, value: &str, time: i128) -> RecordId {
363 records.push_at(Principal::anonymous(), 1, value.to_string(), TimestampNanos::from(time))
364 }
365
366 #[test]
367 fn zero_capacity_discards_records_without_panicking() {
368 let mut records = Records {
369 retention_limit: 0,
370 ..Default::default()
371 };
372
373 let id = push(&mut records, "discarded", 1);
374 records.update_at(id, "ignored".to_string(), TimestampNanos::from(2));
375
376 assert!(records.records.is_empty());
377 assert_eq!(records.retention_evicted_count, 1);
378 assert_eq!(records.next_id.into_inner(), 1);
379 }
380
381 #[test]
382 fn enforces_capacity_after_capacity_is_reduced() {
383 let mut records = Records {
384 retention_limit: 4,
385 ..Default::default()
386 };
387 push(&mut records, "zero", 0);
388 push(&mut records, "one", 1);
389 push(&mut records, "two", 2);
390 records.retention_limit = 2;
391
392 let newest = push(&mut records, "three", 3);
393
394 assert_eq!(records.records.len(), 2);
395 assert_eq!(records.records[0].content, "two");
396 assert_eq!(records.records[1].id, newest);
397 assert_eq!(records.retention_evicted_count, 2);
398 }
399
400 #[test]
401 fn updates_existing_record_and_ignores_missing_record() {
402 let mut records = Records::default();
403 records.update_at(RecordId::from(99), "missing".to_string(), TimestampNanos::from(1));
404
405 let id = push(&mut records, "created", 2);
406 records.update_at(id, "done".to_string(), TimestampNanos::from(3));
407
408 assert_eq!(
409 records.records[0].completion.as_ref().map(|(_, value)| value.as_str()),
410 Some("done")
411 );
412 }
413
414 #[test]
415 fn deletes_only_requested_ids_and_is_safe_to_retry() {
416 let mut records = Records::default();
417 let first = push(&mut records, "first", 1);
418 let second = push(&mut records, "second", 2);
419 let third = push(&mut records, "third", 3);
420 let ids = HashSet::from([first, third, RecordId::from(99)]);
421
422 assert_eq!(records.record_delete(&ids), 2);
423 assert_eq!(records.records.len(), 1);
424 assert_eq!(records.records[0].id, second);
425 assert_eq!(records.next_id.into_inner(), 3);
426 assert_eq!(records.retention_evicted_count, 0);
427
428 assert_eq!(records.record_delete(&ids), 0);
429 assert_eq!(records.records.len(), 1);
430 }
431
432 #[test]
433 fn deserializes_legacy_aliases_and_serializes_current_names() {
434 let legacy = LegacyRecords {
435 max: 42,
436 removed: 3,
437 next_id: RecordId::from(7),
438 records: vec![LegacyRecord {
439 id: RecordId::from(6),
440 created: TimestampNanos::from(1),
441 caller: Principal::anonymous(),
442 topic: 1,
443 content: "legacy".to_string(),
444 done: Some((TimestampNanos::from(2), "ok".to_string())),
445 }],
446 };
447
448 let mut cbor = Vec::new();
449 ciborium::ser::into_writer(&legacy, &mut cbor).unwrap();
450 let decoded: Records = ciborium::de::from_reader(cbor.as_slice()).unwrap();
451 assert_eq!(decoded.retention_limit, 42);
452 assert_eq!(decoded.retention_evicted_count, 3);
453 assert_eq!(decoded.records[0].completion.as_ref().unwrap().1, "ok");
454
455 let mut current_cbor = Vec::new();
456 ciborium::ser::into_writer(&decoded, &mut current_cbor).unwrap();
457 let current: Value = ciborium::de::from_reader(current_cbor.as_slice()).unwrap();
458 let current_keys = map_keys(¤t);
459 assert!(current_keys.contains(&"retention_limit"));
460 assert!(current_keys.contains(&"retention_evicted_count"));
461 assert!(!current_keys.contains(&"max"));
462 assert!(!current_keys.contains(&"removed"));
463
464 let Value::Map(entries) = current else { unreachable!() };
465 let records = entries
466 .iter()
467 .find_map(|(key, value)| (key == &Value::Text("records".to_string())).then_some(value))
468 .unwrap();
469 let Value::Array(records) = records else {
470 panic!("expected records to be a CBOR array")
471 };
472 let record_keys = map_keys(&records[0]);
473 assert!(record_keys.contains(&"completion"));
474 assert!(!record_keys.contains(&"done"));
475
476 let legacy_search = LegacyRecordSearchArg {
477 id: Some((Some(1), Some(2))),
478 created: Some((Some(3), Some(4))),
479 caller: None,
480 topic: None,
481 content: None,
482 };
483 let mut search_cbor = Vec::new();
484 ciborium::ser::into_writer(&legacy_search, &mut search_cbor).unwrap();
485 let search: RecordSearchArg = ciborium::de::from_reader(search_cbor.as_slice()).unwrap();
486 assert_eq!(search.id_range, Some((Some(1), Some(2))));
487 assert_eq!(search.created_at_nanos_range, Some((Some(3), Some(4))));
488 let mut current_search_cbor = Vec::new();
489 ciborium::ser::into_writer(&search, &mut current_search_cbor).unwrap();
490 let current_search: Value = ciborium::de::from_reader(current_search_cbor.as_slice()).unwrap();
491 let search_keys = map_keys(¤t_search);
492 assert!(search_keys.contains(&"id_range"));
493 assert!(search_keys.contains(&"created_at_nanos_range"));
494 assert!(!search_keys.contains(&"id"));
495 assert!(!search_keys.contains(&"created"));
496 }
497 }
498}