mod_memoize/
lib.rs

1use config::epoch::{get_current_epoch, ConfigEpoch};
2use config::{any_err, from_lua_value, get_or_create_module, load_config, serialize_options};
3use dashmap::DashMap;
4use kumo_prometheus::declare_metric;
5use kumo_prometheus::prometheus::Counter;
6use lruttl::{ItemLookup, LruCacheWithTtl};
7use mlua::{
8    FromLua, Function, IntoLua, Lua, LuaSerdeExt, MetaMethod, MultiValue, UserData,
9    UserDataMethods, UserDataRef,
10};
11use serde::{Deserialize, Serialize};
12use std::collections::HashMap;
13use std::sync::{Arc, LazyLock};
14use std::time::Duration;
15
16/// Memoized is a helper type that allows native Rust types to be captured
17/// in memoization caches.
18/// Unfortunately, we cannot automatically make that work for all UserData
19/// that are exported to lua, but we can make it simple for a type to opt-in
20/// to that behavior.
21///
22/// When you impl UserData for your type, you can call
23/// `Memoized::impl_memoize(methods)` from your add_methods impl.
24/// That will add a metamethod to your UserData type that will clone your
25/// value and wrap it into a Memoized wrapper.
26///
27/// Since Clone is used, it is recommended that you use an Arc inside your
28/// type to avoid making large or expensive clones.
29#[derive(Clone, mlua::FromLua)]
30pub struct Memoized {
31    pub to_value: Arc<dyn Fn(&Lua) -> mlua::Result<mlua::Value> + Send + Sync>,
32}
33
34impl PartialEq for Memoized {
35    fn eq(&self, other: &Self) -> bool {
36        Arc::ptr_eq(&self.to_value, &other.to_value)
37    }
38}
39
40impl Memoized {
41    /// Call this from your `UserData::add_methods` implementation to
42    /// enable memoization for your UserData type
43    pub fn impl_memoize<T, M>(methods: &mut M)
44    where
45        T: UserData + Send + Sync + Clone + 'static,
46        M: UserDataMethods<T>,
47    {
48        methods.add_meta_method(
49            "__memoize",
50            move |_lua, this, _: ()| -> mlua::Result<Memoized> {
51                let this = this.clone();
52                Ok(Memoized {
53                    to_value: Arc::new(move |lua| this.clone().into_lua(lua)),
54                })
55            },
56        );
57    }
58}
59
60impl UserData for Memoized {}
61
62#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, Eq)]
63#[serde(deny_unknown_fields)]
64pub struct MemoizeParams {
65    #[serde(with = "duration_serde")]
66    pub ttl: Duration,
67    pub capacity: usize,
68    pub name: String,
69    #[serde(default)]
70    pub invalidate_with_epoch: bool,
71    #[serde(default)]
72    pub retry_on_populate_timeout: bool,
73    #[serde(default, with = "duration_serde")]
74    pub populate_timeout: Option<Duration>,
75    #[serde(default)]
76    pub allow_stale_reads: bool,
77    #[serde(default)]
78    pub detached: Option<bool>,
79}
80
81#[derive(Clone, Hash, Eq, PartialEq)]
82pub enum MapKey {
83    Integer(mlua::Integer),
84    String(Vec<u8>),
85}
86
87impl MapKey {
88    pub fn from_lua(v: mlua::Value) -> Option<Self> {
89        match v {
90            mlua::Value::String(s) => Some(Self::String(s.as_bytes().to_vec())),
91            mlua::Value::Integer(n) => Some(Self::Integer(n)),
92            _ => None,
93        }
94    }
95
96    pub fn as_lua(self, lua: &Lua) -> mlua::Result<mlua::Value> {
97        match self {
98            Self::Integer(j) => Ok(mlua::Value::Integer(j)),
99            Self::String(b) => Ok(mlua::Value::String(lua.create_string(b)?)),
100        }
101    }
102}
103
104#[derive(Clone, PartialEq)]
105pub enum CacheValue {
106    Table(Arc<HashMap<MapKey, CacheValue>>),
107    Json(serde_json::Value),
108    Memoized(Memoized),
109}
110
111impl std::fmt::Debug for CacheValue {
112    fn fmt(&self, fmt: &mut std::fmt::Formatter) -> std::fmt::Result {
113        fmt.debug_struct("CacheValue").finish()
114    }
115}
116
117impl FromLua for CacheValue {
118    fn from_lua(value: mlua::Value, lua: &Lua) -> mlua::Result<Self> {
119        match value {
120            mlua::Value::UserData(ud) => {
121                let mt = ud.metatable()?;
122                let func: Function = mt.get("__memoize")?;
123                let m: Memoized = func.call(mlua::Value::UserData(ud))?;
124                Ok(Self::Memoized(m))
125            }
126            mlua::Value::Table(tbl) => {
127                let mut map = HashMap::new();
128                for pair in tbl.pairs::<mlua::Value, mlua::Value>() {
129                    let (key, value) = pair?;
130                    let key = match key {
131                        mlua::Value::Integer(n) => MapKey::Integer(n),
132                        mlua::Value::String(n) => MapKey::String(n.as_bytes().to_vec()),
133                        _ => {
134                            return Err(anyhow::anyhow!(
135                                "table key {key:?} cannot be used as a key in a memoizable table"
136                            ))
137                            .map_err(any_err)
138                        }
139                    };
140                    let value = CacheValue::from_lua(value, lua)?;
141                    map.insert(key, value);
142                }
143                Ok(Self::Table(map.into()))
144            }
145            _ => Ok(Self::Json(from_lua_value(lua, value)?)),
146        }
147    }
148}
149
150impl IntoLua for CacheValue {
151    fn into_lua(self, lua: &Lua) -> mlua::Result<mlua::Value> {
152        self.as_lua(lua)
153    }
154}
155
156impl CacheValue {
157    pub fn as_lua(&self, lua: &Lua) -> mlua::Result<mlua::Value> {
158        match self {
159            Self::Json(j) => lua.to_value_with(j, serialize_options()),
160            Self::Memoized(m) => (m.to_value)(lua),
161            Self::Table(m) => Ok(mlua::Value::UserData(
162                lua.create_userdata(MemoizedTable::Shared(m.clone()))?,
163            )),
164        }
165    }
166}
167
168/// MemoizedTable is a helper type that is returned to represent
169/// cached table values.  We'll return the Shared variant by
170/// default as that presents the cheapest way to return the cached
171/// data--only a clone of the underlying Arc is required to return
172/// the value.
173///
174/// This type implements __index, __newindex, __len, and __pairs
175/// metamethods which allow reading and iterating the table.
176///
177/// Writing to the table via __newindex will "unshare" the table in
178/// a similar manner to the Cow type, creating a mutable copy of the top
179/// level of the table.
180enum MemoizedTable {
181    Shared(Arc<HashMap<MapKey, CacheValue>>),
182    Mut(HashMap<MapKey, CacheValue>),
183}
184
185impl MemoizedTable {
186    /// Get a reference to the table, facilitating get() and iter(),
187    /// regardless of whether we are Shared or Mut.
188    fn table(&self) -> &HashMap<MapKey, CacheValue> {
189        match self {
190            Self::Shared(s) => s,
191            Self::Mut(s) => s,
192        }
193    }
194
195    /// Transform Shared -> Mut
196    fn unshare(&mut self) -> &mut HashMap<MapKey, CacheValue> {
197        if let Self::Shared(t) = self {
198            *self = Self::Mut(t.iter().map(|(k, v)| (k.clone(), v.clone())).collect());
199        }
200
201        match self {
202            Self::Shared(_) => unreachable!(),
203            Self::Mut(map) => map,
204        }
205    }
206}
207
208impl UserData for MemoizedTable {
209    fn add_methods<M: UserDataMethods<Self>>(methods: &mut M) {
210        // Index allows reading fields of the table
211        methods.add_meta_method(MetaMethod::Index, move |lua, this, key: mlua::Value| {
212            match MapKey::from_lua(key) {
213                Some(key) => match this.table().get(&key) {
214                    Some(value) => value.as_lua(lua),
215                    None => Ok(mlua::Value::Nil),
216                },
217                None => Ok(mlua::Value::Nil),
218            }
219        });
220
221        // NewIndex allows writing fields of the table
222        methods.add_meta_method_mut(
223            MetaMethod::NewIndex,
224            move |lua, this, (key, value): (mlua::Value, mlua::Value)| match MapKey::from_lua(key) {
225                Some(key) => {
226                    let value = CacheValue::from_lua(value, lua)?;
227                    this.unshare().insert(key, value);
228                    Ok(())
229                }
230                None => Err(mlua::Error::external(
231                    "invalid key type while trying to call __newindex and assign a value",
232                )),
233            },
234        );
235        methods.add_meta_method(MetaMethod::Len, move |_lua, this, ()| {
236            Ok(this.table().len())
237        });
238
239        // Pairs iterates the keys of the table.
240        // We use add_meta_function rather than add_meta_method here
241        // because we need to return `this` as the "state" parameter
242        // for use in a generic-for statement
243        methods.add_meta_function(MetaMethod::Pairs, move |lua, this: mlua::Value| {
244            // Maintain our own local idea of the control variable,
245            // as it is much cheaper and simpler to iterate based
246            // on skipping than to keep comparing keys
247            let mut idx = 0;
248
249            let iter_func =
250                lua.create_function_mut(
251                    move |lua, (state, _control): (UserDataRef<MemoizedTable>, mlua::Value)| {
252                        match state.table().iter().nth(idx) {
253                            Some((key, value)) => {
254                                idx += 1;
255                                let key = key.clone().as_lua(lua)?;
256                                let value = value.as_lua(lua)?;
257                                Ok((key, value))
258                            }
259                            None => Ok((mlua::Value::Nil, mlua::Value::Nil)),
260                        }
261                    },
262                )?;
263
264            // Return the iterator, state and control values.
265            // The state and control will be passed back into iter_func
266            // as the for-loop iterates.
267            // Control is Nil here because we track our own idx
268            // value in the iter_func closure.
269            Ok((mlua::Value::Function(iter_func), this, mlua::Value::Nil))
270        });
271    }
272}
273
274#[derive(Clone, Debug)]
275enum CacheEntry {
276    Null,
277    Single(CacheValue),
278    Multi(Vec<CacheValue>),
279}
280
281impl CacheEntry {
282    fn to_value(&self, lua: &Lua) -> mlua::Result<mlua::Value> {
283        match self {
284            Self::Null => Ok(mlua::Value::Nil),
285            Self::Single(value) => value.as_lua(lua),
286            Self::Multi(values) => {
287                let mut result = vec![];
288                for v in values {
289                    result.push(v.as_lua(lua)?);
290                }
291                result.into_lua(lua)
292            }
293        }
294    }
295
296    fn from_multi_value(lua: &Lua, multi: MultiValue) -> mlua::Result<Self> {
297        let mut values = multi.into_vec();
298        if values.is_empty() {
299            Ok(Self::Null)
300        } else if values.len() == 1 {
301            Ok(Self::Single(CacheValue::from_lua(
302                values.pop().unwrap(),
303                lua,
304            )?))
305        } else {
306            let mut cvalues = vec![];
307            for v in values.into_iter() {
308                cvalues.push(CacheValue::from_lua(v, lua)?);
309            }
310            Ok(Self::Multi(cvalues))
311        }
312    }
313}
314
315struct MemoizeCache {
316    params: MemoizeParams,
317    cache: Arc<LruCacheWithTtl<CacheKey, CacheEntry>>,
318}
319
320static CACHES: LazyLock<DashMap<String, MemoizeCache>> = LazyLock::new(DashMap::new);
321
322type CacheKey = (Option<ConfigEpoch>, String);
323
324fn get_cache_by_name(
325    name: &str,
326) -> Option<(Arc<LruCacheWithTtl<CacheKey, CacheEntry>>, Duration, bool)> {
327    CACHES.get(name).map(|item| {
328        (
329            item.cache.clone(),
330            item.params.ttl,
331            item.params.invalidate_with_epoch,
332        )
333    })
334}
335
336declare_metric! {
337/// How many times a memoize cache lookup was initiated for a given cache.
338///
339/// Redundant with the newer [lruttl_lookup_count](lruttl_lookup_count.md) metric.
340static CACHE_LOOKUP: CounterVec(
341        "memoize_cache_lookup_count",
342        &["cache_name"]);
343}
344
345declare_metric! {
346/// How many times a memoize cache lookup was a hit for a given cache.
347///
348/// Redundant with the newer [lruttl_hit_count](lruttl_hit_count.md) metric.
349static CACHE_HIT: CounterVec(
350        "memoize_cache_hit_count",
351        &["cache_name"]);
352}
353
354declare_metric! {
355/// How many times a memoize cache lookup was a miss for a given cache
356///
357/// Redundant with the newer [lruttl_miss_count](lruttl_miss_count.md) metric.
358static CACHE_MISS: CounterVec(
359        "memoize_cache_miss_count",
360        &["cache_name"]);
361}
362
363declare_metric! {
364/// How many times a memoize cache lookup resulted in performing the work to populate the entry
365///
366/// Redundant with the newer [lruttl_populated_count](lruttl_populated_count.md) metric.
367static CACHE_POPULATED: CounterVec(
368        "memoize_cache_populated_count",
369        &["cache_name"]);
370}
371
372/// Returns the call arguments as a JSON array with one entry per argument,
373/// preserving the argument count, unlike the collapsed JSON value that
374/// `multi_value_to_json_value` builds for the cache key.
375fn multi_value_to_json_args(lua: &Lua, multi: MultiValue) -> mlua::Result<Vec<serde_json::Value>> {
376    multi
377        .into_vec()
378        .into_iter()
379        .map(|v| from_lua_value(lua, v))
380        .collect()
381}
382
383/// Returns the name of the event handler the current call is running inside,
384/// or None when the call is running at top-level policy scope.
385fn calling_event_handler(lua: &Lua) -> Option<String> {
386    lua.globals().get::<String>("_KUMO_CURRENT_EVENT").ok()
387}
388
389fn multi_value_to_json_value(lua: &Lua, multi: MultiValue) -> mlua::Result<serde_json::Value> {
390    let mut values = multi.into_vec();
391    if values.is_empty() {
392        Ok(serde_json::Value::Null)
393    } else if values.len() == 1 {
394        from_lua_value(lua, values.pop().unwrap())
395    } else {
396        let mut jvalues = vec![];
397        for v in values.into_iter() {
398            jvalues.push(from_lua_value(lua, v)?);
399        }
400        Ok(serde_json::Value::Array(jvalues))
401    }
402}
403
404/// Looks up `key`, populating it on the calling task on a miss. Has no
405/// timeout of its own: a slow populate runs for as long as it takes and its
406/// result is still cached. If the caller is cancelled while the populate is
407/// running, the populate is aborted and nothing is cached.
408async fn populate_inline(
409    cache: &LruCacheWithTtl<CacheKey, CacheEntry>,
410    key: &CacheKey,
411    ttl: Duration,
412    lua: &Lua,
413    func: Function,
414    params: &MultiValue,
415    populate_counter: &Counter,
416) -> Result<ItemLookup<CacheEntry>, Arc<anyhow::Error>> {
417    cache
418        .get_or_try_insert(key, |_| ttl, async {
419            tracing::trace!("populate {key:?}");
420            populate_counter.inc();
421            let result: MultiValue = func.call_async(params.clone()).await?;
422            CacheEntry::from_multi_value(lua, result)
423        })
424        .await
425}
426
427/// Looks up `key`, populating it on a miss. The populate runs on a task of its
428/// own: if the caller's own task is later dropped, such as an HTTP handler
429/// whose client disconnected, the populate keeps running to completion and
430/// other callers waiting on the same `key` still get its result. `args` must be
431/// the JSON form of the call arguments (see `multi_value_to_json_args`), and
432/// `registry_name` must name the populate function in the Lua registry. The
433/// populate is bounded by `populate_timeout`. Once it elapses, the entry is
434/// cached as failed and this call returns that failure as an error, even
435/// though the populate task itself keeps running to completion.
436async fn populate_detached(
437    cache: &LruCacheWithTtl<CacheKey, CacheEntry>,
438    key: &CacheKey,
439    ttl: Duration,
440    cache_name: String,
441    registry_name: String,
442    args: Vec<serde_json::Value>,
443    populate_counter: Counter,
444) -> Result<ItemLookup<CacheEntry>, Arc<anyhow::Error>> {
445    // populate_timeout reuses the sema timeout of the cache, set from
446    // MemoizeParams::populate_timeout. get_or_try_insert_detached hard-cancels
447    // the populate once it elapses: it caches a Failed entry with a 60-second
448    // TTL in place of a result, and the next lookup spawns a new populate.
449    let populate_timeout = cache.get_sema_timeout();
450    let make_fut = move || {
451        let registry_name = registry_name.clone();
452        let args = args.clone();
453        let populate_counter = populate_counter.clone();
454        let cache_name = cache_name.clone();
455        async move {
456            let config = load_config().await?;
457            let entry = {
458                let lua = config.lua()?;
459                // The function called here can be the newer policy's version
460                // if a reload re-ran the kumo.memoize call before this point.
461                // We still cache its result under the key built from
462                // epoch_at_start, not the newer epoch: we want a result
463                // produced under an older policy never mistaken for one that
464                // reflects the current policy, which is what tagging it with
465                // the newer epoch would do.
466                let func: Function = lua.named_registry_value(&registry_name).map_err(|_| {
467                    anyhow::anyhow!(
468                        "memoize populate function for cache {cache_name} is not registered in \
469                         a freshly loaded config context. This usually means the kumo.memoize \
470                         call does not run at top-level policy scope, where a config reload \
471                         would re-run it and re-register the function; it can also happen if a \
472                         policy reload removed the kumo.memoize call"
473                    )
474                })?;
475                let mut arg_vec = Vec::with_capacity(args.len());
476                for a in &args {
477                    arg_vec.push(lua.to_value_with(a, serialize_options())?);
478                }
479                populate_counter.inc();
480                let result: MultiValue = func.call_async(MultiValue::from_vec(arg_vec)).await?;
481                CacheEntry::from_multi_value(lua, result)?
482            };
483            config.put();
484            Ok::<_, anyhow::Error>(entry)
485        }
486    };
487    cache
488        .get_or_try_insert_detached(key, move |_| ttl, make_fut, populate_timeout)
489        .await
490}
491
492pub fn register(lua: &Lua) -> anyhow::Result<()> {
493    let kumo_mod = get_or_create_module(lua, "kumo")?;
494
495    kumo_mod.set(
496        "memoize",
497        lua.create_function(move |lua, (func, params): (mlua::Function, mlua::Value)| {
498            let params: MemoizeParams = from_lua_value(lua, params)?;
499
500            let cache_name = params.name.to_string();
501
502            if !lruttl::is_name_available(&cache_name) {
503                return Err(mlua::Error::external(format!(
504                    "cannot use name `{cache_name}` for a memoize cache, \
505                    as it collides with a built-in cache. \
506                    Suggestion: prefix your cache name with `user.` to \
507                    avoid conflicts with current and future caches."
508                )));
509            }
510
511            CACHES.remove_if(&params.name, |_k, item| {
512                let changed = item.params != params;
513                if changed {
514                    tracing::trace!("memoize parameters changed, replacing old cache {params:?}");
515                }
516                changed
517            });
518            CACHES.entry(cache_name.to_string()).or_insert_with(|| {
519                let cache = LruCacheWithTtl::new(cache_name.clone(), params.capacity);
520                if let Some(duration) = params.populate_timeout {
521                    cache.set_sema_timeout(duration);
522                }
523                cache.set_allow_stale_reads(params.allow_stale_reads);
524
525                MemoizeCache {
526                    params: params.clone(),
527                    cache: Arc::new(cache),
528                }
529            });
530
531            let lookup_counter = CACHE_LOOKUP
532                .get_metric_with_label_values(&[&cache_name])
533                .map_err(any_err)?;
534            let hit_counter = CACHE_HIT
535                .get_metric_with_label_values(&[&cache_name])
536                .map_err(any_err)?;
537            let miss_counter = CACHE_MISS
538                .get_metric_with_label_values(&[&cache_name])
539                .map_err(any_err)?;
540            let populate_counter = CACHE_POPULATED
541                .get_metric_with_label_values(&[&cache_name])
542                .map_err(any_err)?;
543            let retry_on_populate_timeout = params.retry_on_populate_timeout;
544            let allow_stale_reads = params.allow_stale_reads;
545            let detached = match params.detached {
546                Some(true) => {
547                    if let Some(event) = calling_event_handler(lua) {
548                        return Err(mlua::Error::external(format!(
549                            "kumo.memoize cache `{cache_name}` sets `detached = true`, but is \
550                            being called from within the `{event}` event handler. A detached \
551                            populate reloads the policy to re-establish the populate function, \
552                            and reloading runs only top-level policy code, not event handlers. \
553                            Move this kumo.memoize call to top-level policy scope, or set \
554                            `detached = false` if the populate is fast enough to run inline."
555                        )));
556                    }
557                    true
558                }
559                Some(false) => false,
560                None => calling_event_handler(lua).is_none(),
561            };
562
563            let registry_name = format!("kumo-memoize-fn.{cache_name}");
564            lua.set_named_registry_value(&registry_name, func.clone())?;
565
566            let func_ref = lua.create_registry_value(func)?;
567
568            lua.create_async_function(move |lua, params: MultiValue| {
569                let cache_name = cache_name.clone();
570                let registry_name = registry_name.clone();
571                let func = lua.registry_value::<mlua::Function>(&func_ref);
572                let lookup_counter = lookup_counter.clone();
573                let hit_counter = hit_counter.clone();
574                let miss_counter = miss_counter.clone();
575                let populate_counter = populate_counter.clone();
576                async move {
577                    lookup_counter.inc();
578                    let key = multi_value_to_json_value(&lua, params.clone())?;
579
580                    let func = func?;
581
582                    let mut last_failure = None;
583
584                    for _attempt in 0..3 {
585                        // We use the epoch from the start of the lookup as part
586                        // of the cache key. If the epoch changes while we are in
587                        // the middle of computing this value then subsequent calls
588                        // through to the cached function will see the newer epoch
589                        // and encounter a cache miss. This prevents a race condition
590                        // poisoning the cache with a stale value during an epoch
591                        // bump. The caller will still observe the stale value, so
592                        // ultimately should have some accommodation for detecting
593                        // the epoch change and retrying their call through here,
594                        // if it is important to not see a stale value.
595                        let epoch_at_start = get_current_epoch();
596
597                        let (cache, ttl, invalidate_with_epoch) = get_cache_by_name(&cache_name)
598                            .ok_or_else(|| anyhow::anyhow!("cache is somehow undefined!?"))
599                            .map_err(any_err)?;
600
601                        let epoch_key = if invalidate_with_epoch && !allow_stale_reads {
602                            Some(epoch_at_start)
603                        } else {
604                            None
605                        };
606                        let key = serde_json::to_string(&key).map_err(any_err)?;
607                        let key = (epoch_key, key);
608
609                        let value_result = if detached {
610                            let args = multi_value_to_json_args(&lua, params.clone())?;
611                            populate_detached(
612                                &cache,
613                                &key,
614                                ttl,
615                                cache_name.clone(),
616                                registry_name.clone(),
617                                args,
618                                populate_counter.clone(),
619                            )
620                            .await
621                        } else {
622                            populate_inline(
623                                &cache,
624                                &key,
625                                ttl,
626                                &lua,
627                                func.clone(),
628                                &params,
629                                &populate_counter,
630                            )
631                            .await
632                        };
633
634                        match value_result {
635                            Ok(lookup) => {
636                                if lookup.is_fresh {
637                                    miss_counter.inc();
638                                } else {
639                                    hit_counter.inc();
640                                }
641                                return lookup.item.to_value(&lua);
642                            }
643                            Err(err) => {
644                                tracing::error!("{cache_name} {key:?} failed: {err:#}");
645                                let error = format!("{err:#}");
646                                if !retry_on_populate_timeout {
647                                    return Err(mlua::Error::external(error));
648                                }
649                                last_failure.replace(error);
650                            }
651                        }
652                    }
653
654                    Err(mlua::Error::external(
655                        last_failure.expect("last_failure to always be set in loop above"),
656                    ))
657                }
658            })
659        })?,
660    )?;
661
662    Ok(())
663}
664
665#[cfg(test)]
666mod test {
667    use super::*;
668    use mlua::UserDataMethods;
669    use std::sync::atomic::{AtomicUsize, Ordering};
670
671    #[tokio::test]
672    async fn test_memoize() {
673        let lua = Lua::new();
674        register(&lua).unwrap();
675
676        let call_count = Arc::new(AtomicUsize::new(0));
677
678        let globals = lua.globals();
679        let counter = Arc::clone(&call_count);
680        globals
681            .set(
682                "do_thing",
683                lua.create_function(move |_lua, _: ()| {
684                    let count = counter.fetch_add(1, Ordering::SeqCst);
685                    Ok(count)
686                })
687                .unwrap(),
688            )
689            .unwrap();
690
691        let result: usize = lua
692            .load(
693                r#"
694            local kumo = require 'kumo';
695            -- make cached_do_thing a global for use in the expiry test below
696            cached_do_thing = kumo.memoize(do_thing, {
697                ttl = "1s",
698                capacity = 4,
699                name = "test_memoize_do_thing",
700                -- bare test Lua has no policy for load_config to reload, so
701                -- exercise the inline populate
702                detached = false,
703            })
704            return cached_do_thing() + cached_do_thing() + cached_do_thing()
705        "#,
706            )
707            .eval_async()
708            .await
709            .unwrap();
710
711        assert_eq!(result, 0);
712        assert_eq!(call_count.load(Ordering::SeqCst), 1);
713
714        // And confirm that expiry works
715        tokio::time::sleep(tokio::time::Duration::from_secs(2)).await;
716
717        let result: usize = lua
718            .load(
719                r#"
720            return cached_do_thing()
721        "#,
722            )
723            .eval()
724            .unwrap();
725
726        assert_eq!(result, 1);
727        assert_eq!(call_count.load(Ordering::SeqCst), 2);
728    }
729
730    #[tokio::test]
731    async fn test_memoize_rust() {
732        let lua = Lua::new();
733        register(&lua).unwrap();
734
735        let call_count = Arc::new(AtomicUsize::new(0));
736
737        #[derive(Clone)]
738        struct Foo {
739            value: usize,
740        }
741
742        impl UserData for Foo {
743            fn add_methods<M: UserDataMethods<Self>>(methods: &mut M) {
744                Memoized::impl_memoize(methods);
745                methods.add_method("get_value", move |_lua, this, _: ()| Ok(this.value));
746            }
747        }
748
749        let globals = lua.globals();
750        let counter = Arc::clone(&call_count);
751        globals
752            .set(
753                "make_foo",
754                lua.create_function(move |_lua, _: ()| {
755                    let count = counter.fetch_add(1, Ordering::SeqCst);
756                    Ok(Foo { value: count })
757                })
758                .unwrap(),
759            )
760            .unwrap();
761
762        let result: usize = lua
763            .load(
764                r#"
765            local kumo = require 'kumo';
766            local cached_make_foo = kumo.memoize(make_foo, {
767                ttl = "1s",
768                capacity = 4,
769                name = "test_memoize_make_foo",
770                -- bare test Lua has no policy for load_config to reload, so
771                -- exercise the inline populate
772                detached = false,
773            })
774            return cached_make_foo():get_value() +
775                   cached_make_foo():get_value() +
776                   cached_make_foo():get_value()
777        "#,
778            )
779            .eval()
780            .unwrap();
781
782        assert_eq!(result, 0);
783        assert_eq!(call_count.load(Ordering::SeqCst), 1);
784    }
785
786    #[tokio::test]
787    #[test_log::test]
788    async fn test_memoize_blocked() {
789        use std::sync::Mutex;
790        use tokio::sync::Notify;
791
792        let call_count = Arc::new(AtomicUsize::new(0));
793        let notify = Arc::new(Notify::new());
794        let shared_notify = Arc::new(Mutex::new(Some(Arc::clone(&notify))));
795
796        async fn setup_lua(
797            call_count: &Arc<AtomicUsize>,
798            notify: Arc<Mutex<Option<Arc<Notify>>>>,
799        ) -> Lua {
800            let lua = Lua::new();
801            register(&lua).unwrap();
802            let globals = lua.globals();
803            let counter = Arc::clone(&call_count);
804
805            fn take_notify(n: &Arc<Mutex<Option<Arc<Notify>>>>) -> Option<Arc<Notify>> {
806                n.lock().unwrap().take()
807            }
808
809            globals
810                .set(
811                    "do_thing",
812                    lua.create_async_function(move |_lua, _: ()| {
813                        let counter = counter.clone();
814                        let notify = notify.clone();
815                        async move {
816                            eprintln!("do_thing called!");
817                            match dbg!(take_notify(&notify)) {
818                                Some(notify) => {
819                                    eprintln!("do_thing: wait for notify");
820                                    notify.notified().await;
821                                    eprintln!("notified!");
822                                }
823                                None => {
824                                    eprintln!("do_thing: sleeping");
825                                    tokio::time::sleep(Duration::from_secs(1)).await;
826                                    eprintln!("do_thing: slept");
827                                }
828                            };
829                            eprintln!("do_thing: increment");
830                            let count = counter.fetch_add(1, Ordering::SeqCst);
831                            Ok(count)
832                        }
833                    })
834                    .unwrap(),
835                )
836                .unwrap();
837
838            let init = r#"
839            local kumo = require 'kumo';
840            -- make cached_do_thing a global for use in the expiry test below
841            cached_do_thing = kumo.memoize(do_thing, {
842                ttl = "1s",
843                capacity = 4,
844                name = "test_memoize_do_thing",
845                populate_timeout = "2s",
846                -- bare test Lua has no policy for load_config to reload, so
847                -- exercise the inline populate
848                detached = false,
849            })
850        "#;
851
852            let () = lua.load(init).eval_async().await.unwrap();
853            lua
854        }
855
856        let lua = setup_lua(&call_count, shared_notify.clone()).await;
857        async fn do_thing(lua: Lua) -> mlua::Result<usize> {
858            lua.load("return cached_do_thing()").eval_async().await
859        }
860
861        // Set up a future that will get far enough to own the lookup,
862        // but that won't complete until we notify it to do so.
863        eprintln!("spawning first call to do_thing");
864        let first_future = tokio::spawn(do_thing(lua));
865        // Let it progress to await on the notifier
866        tokio::task::yield_now().await;
867
868        // Now setup a second call; we expect this one to time out
869        // because the first one owns the lookup
870        let lua = setup_lua(&call_count, shared_notify.clone()).await;
871        eprintln!("second call to do_thing");
872
873        let res = do_thing(lua).await;
874        eprintln!("second_future is done!");
875        let error = res.unwrap_err().to_string();
876        assert!(
877            error.contains(
878                "timed out after 2s on semaphore acquire while waiting for cache to populate"
879            ),
880            "error: {error}"
881        );
882
883        // The third call will succeed because we've de-coupled the
884        // original lookup from the herd protection
885        eprintln!("third call to do_thing");
886        let lua = setup_lua(&call_count, shared_notify.clone()).await;
887        let result = do_thing(lua).await.unwrap();
888        assert_eq!(result, 0);
889
890        // Now wake up the original future and verify what it returns
891        notify.notify_one();
892        assert_eq!(first_future.await.unwrap().unwrap(), 1);
893    }
894}