Skip to main content

ic_canister_kit/functions/
record.rs

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/// 记录 id
12#[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    /// 取出内部数据
23    pub fn into_inner(&self) -> u64 {
24        self.0
25    }
26    /// 下一个 id
27    pub fn next(self) -> Self {
28        Self(self.0 + 1)
29    }
30}
31
32/// 查询
33pub trait Searchable<Record> {
34    /// 查询
35    fn test(&self, record: &Record) -> bool;
36}
37
38/// 可以记录的操作
39pub trait Recordable<Record, RecordTopic, Search: Searchable<Record>> {
40    // 查询
41    /// 查询所有
42    fn record_find_all(&self) -> &[Record];
43
44    // 修改
45    /// 插入记录
46    fn record_push(&mut self, caller: CallerId, topic: RecordTopic, content: String) -> RecordId;
47    /// 更新记录
48    fn record_update(&mut self, record_id: RecordId, result: String);
49    /// 按 id 批量删除记录,返回实际删除的记录数量
50    fn record_delete(&mut self, ids: &HashSet<RecordId>) -> u64;
51
52    /// 分页查询
53    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
67// ================== 简单实现 ==================
68
69/// 记录功能简单实现
70pub 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    /// 记录主题
83    pub type RecordTopic = u8;
84
85    /// 每条记录
86    #[derive(CandidType, Serialize, Deserialize, Debug, Clone)]
87    pub struct Record {
88        /// 记录 id
89        pub id: RecordId,
90        /// 创建时间戳 纳秒
91        pub created: TimestampNanos,
92        /// 调用人
93        pub caller: CallerId,
94        /// 记录主题
95        pub topic: RecordTopic,
96        /// 记录内容
97        pub content: String,
98        /// 完成时间与执行结果
99        #[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    /// 记录检索
111    #[derive(CandidType, Serialize, Deserialize, Debug, Clone)]
112    pub struct RecordSearch {
113        /// id 范围过滤,依次为包含下界和包含上界
114        #[serde(alias = "id")]
115        pub id_range: Option<(Option<RecordId>, Option<RecordId>)>,
116        /// 创建时间纳秒范围过滤,依次为包含下界和包含上界
117        #[serde(alias = "created")]
118        pub created_at_nanos_range: Option<(Option<TimestampNanos>, Option<TimestampNanos>)>,
119        /// 调用人过滤
120        pub caller: Option<HashSet<CallerId>>,
121        /// 主题过滤
122        pub topic: Option<HashSet<RecordTopic>>,
123        /// 内容过滤
124        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    /// 持久化的记录对象
176    #[derive(CandidType, Serialize, Deserialize, Debug, Clone)]
177    pub struct Records {
178        /// 最多保留的记录条数
179        #[serde(alias = "max")]
180        pub retention_limit: u64,
181        /// 因保留上限被累计淘汰的记录数量
182        #[serde(alias = "removed")]
183        pub retention_evicted_count: u64,
184        /// 下一个未使用的 id
185        pub next_id: RecordId,
186        /// 当前保留的记录列表
187        pub records: Vec<Record>,
188    }
189
190    impl Default for Records {
191        fn default() -> Self {
192            Self {
193                retention_limit: 1024 * 64, // 假设一条占用 1KB 则最大 64MB 记录
194                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        // 查询
245
246        // 查询所有 正序
247        fn record_find_all(&self) -> &[Record] {
248            &self.records
249        }
250
251        // 修改
252        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        /// 更新记录
257        fn record_update(&mut self, record_id: RecordId, result: String) {
258            self.update_at(record_id, result, crate::times::now());
259        }
260
261        // 删除
262        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    /// 记录检索
274    #[derive(CandidType, Serialize, Deserialize, Debug, Clone)]
275    pub struct RecordSearchArg {
276        /// id 范围过滤,依次为包含下界和包含上界
277        #[serde(alias = "id")]
278        pub id_range: Option<(Option<u64>, Option<u64>)>,
279        /// 创建时间纳秒范围过滤,依次为包含下界和包含上界
280        #[serde(alias = "created")]
281        pub created_at_nanos_range: Option<(Option<u64>, Option<u64>)>,
282        /// 调用人过滤
283        pub caller: Option<HashSet<CallerId>>,
284        /// 主题过滤
285        pub topic: Option<HashSet<String>>,
286        /// 内容过滤
287        pub content: Option<String>,
288    }
289
290    impl RecordSearchArg {
291        /// 参数转变
292        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(&current);
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(&current_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}