Skip to main content

kcode_k1_kmap_loader/
lib.rs

1use std::collections::{HashMap, HashSet, hash_map::Entry};
2
3use kcode_k1_kmap_format::ConnectionTier;
4pub use kcode_k1_kmap_format::{Node, NodeId};
5use kcode_k1_kmap_selection::score;
6
7pub const PREVIEW_COST: f64 = 0.3;
8pub const NARRATIVE_COST: f64 = 1.0;
9
10const PREVIEW_TENTHS: u64 = 3;
11const NARRATIVE_TENTHS: u64 = 10;
12
13type CandidateFilter<'a> = dyn FnMut(&[NodeId]) -> Result<Vec<NodeId>, String> + 'a;
14type Ticket = (usize, f64, u64);
15
16#[derive(Clone, Debug, PartialEq)]
17pub struct LoadedNode {
18    pub node_id: NodeId,
19    pub source: Option<NodeId>,
20    pub title: String,
21    pub navigation_hint: String,
22    pub narrative: Option<String>,
23}
24
25#[derive(Clone, Debug, PartialEq)]
26pub struct OpenResult {
27    pub nodes: Vec<LoadedNode>,
28    pub automatic_attention_spent: f64,
29}
30
31#[derive(Clone, Copy, Debug, Eq, PartialEq)]
32pub enum OpenMode {
33    Full,
34    NavigationOnly,
35}
36
37pub fn open_node(
38    node_id: NodeId,
39    budget: f64,
40    temperature: f64,
41    mode: OpenMode,
42    load_node: impl FnMut(NodeId) -> Result<Option<Node>, String>,
43    candidate_filter: impl FnMut(&[NodeId]) -> Result<Vec<NodeId>, String>,
44) -> Result<OpenResult, String> {
45    open_node_with_random(
46        node_id,
47        budget,
48        temperature,
49        mode,
50        load_node,
51        candidate_filter,
52        kcode_k1_kmap_selection::os_random_unit,
53    )
54}
55
56fn open_node_with_random(
57    node_id: NodeId,
58    budget: f64,
59    temperature: f64,
60    mode: OpenMode,
61    mut load_node: impl FnMut(NodeId) -> Result<Option<Node>, String>,
62    mut candidate_filter: impl FnMut(&[NodeId]) -> Result<Vec<NodeId>, String>,
63    mut random: impl FnMut() -> Result<f64, String>,
64) -> Result<OpenResult, String> {
65    if !budget.is_finite() || budget < 0.0 {
66        return Err("budget must be finite and nonnegative".to_owned());
67    }
68    if !temperature.is_finite() || temperature < 0.0 {
69        return Err("temperature must be finite and nonnegative".to_owned());
70    }
71    let mut engine = Engine {
72        budget,
73        temperature,
74        load_node: &mut load_node,
75        candidate_filter: &mut candidate_filter,
76        random: &mut random,
77        decisions: HashMap::from([(node_id, true)]),
78    };
79    let root = engine.required(node_id)?;
80    match mode {
81        OpenMode::Full => engine.full(node_id, root),
82        OpenMode::NavigationOnly => engine.navigation_only(node_id, root),
83    }
84}
85
86struct Engine<'a> {
87    budget: f64,
88    temperature: f64,
89    load_node: &'a mut dyn FnMut(NodeId) -> Result<Option<Node>, String>,
90    candidate_filter: &'a mut CandidateFilter<'a>,
91    random: &'a mut dyn FnMut() -> Result<f64, String>,
92    decisions: HashMap<NodeId, bool>,
93}
94
95impl Engine<'_> {
96    fn full(&mut self, node_id: NodeId, root: Node) -> Result<OpenResult, String> {
97        let mut outputs = vec![loaded(node_id, None, &root, true)];
98        let mut states = HashMap::from([(node_id, NodeState::Opened)]);
99        let root_targets = self.nav(node_id, &root, 1.0, &states)?;
100        for occurrence in root_targets {
101            let target = occurrence.target;
102            preview(
103                target,
104                occurrence.source,
105                self.required(target)?,
106                &mut outputs,
107                &mut states,
108            );
109        }
110        let mut frontier = self.edges(node_id, &root, 1.0, &states)?;
111        let mut spent = 0_u64;
112
113        loop {
114            let mut candidates = Vec::new();
115            let mut opening_costs = HashMap::new();
116            for (index, occurrence) in frontier.iter().enumerate() {
117                if matches!(states.get(&occurrence.target), Some(NodeState::Opened)) {
118                    continue;
119                }
120                if !occurrence.strength.is_finite() || occurrence.strength <= 0.0 {
121                    continue;
122                }
123                let cost = match states.get(&occurrence.target) {
124                    None => PREVIEW_TENTHS,
125                    Some(NodeState::Previewed(node)) => {
126                        if let Some(cost) = opening_costs.get(&occurrence.target) {
127                            *cost
128                        } else {
129                            let previews = self
130                                .nav(occurrence.target, node, occurrence.strength, &states)?
131                                .len();
132                            let cost =
133                                attention(preview_cost(previews)?.checked_add(NARRATIVE_TENTHS))?;
134                            let _ = opening_costs.insert(occurrence.target, cost);
135                            cost
136                        }
137                    }
138                    Some(NodeState::Opened) => continue,
139                };
140                if affordable(spent, cost, self.budget)? {
141                    candidates.push((index, occurrence.strength, cost));
142                }
143            }
144            if candidates.is_empty() {
145                break;
146            }
147
148            let choice = self.choose(&candidates)?;
149            let (selected, _, cost) = candidates[choice];
150            let occurrence = frontier[selected].clone();
151            if let Some(NodeState::Previewed(node)) = states.get(&occurrence.target).cloned() {
152                let _ = states.insert(occurrence.target, NodeState::Opened);
153                let guarantees =
154                    self.nav(occurrence.target, &node, occurrence.strength, &states)?;
155                outputs
156                    .iter_mut()
157                    .find(|node| node.node_id == occurrence.target)
158                    .ok_or_else(|| "previewed Kmap node had no output".to_owned())?
159                    .narrative = Some(node.narrative.clone());
160                frontier.retain(|entry| entry.target != occurrence.target);
161                for guarantee in guarantees {
162                    let target = guarantee.target;
163                    preview(
164                        target,
165                        guarantee.source,
166                        self.required(target)?,
167                        &mut outputs,
168                        &mut states,
169                    );
170                }
171                frontier.extend(self.edges(
172                    occurrence.target,
173                    &node,
174                    occurrence.strength,
175                    &states,
176                )?);
177            } else {
178                preview(
179                    occurrence.target,
180                    occurrence.source,
181                    self.required(occurrence.target)?,
182                    &mut outputs,
183                    &mut states,
184                );
185            }
186            spent = attention(spent.checked_add(cost))?;
187        }
188        Ok(result(outputs, spent))
189    }
190
191    fn navigation_only(&mut self, node_id: NodeId, root: Node) -> Result<OpenResult, String> {
192        let mut outputs = vec![loaded(node_id, None, &root, false)];
193        let mut states = HashMap::from([(node_id, NodeState::Opened)]);
194        let root_targets = self.nav(node_id, &root, 1.0, &states)?;
195        let root_cost = preview_cost(root_targets.len())?;
196        if !affordable(0, root_cost, self.budget)? {
197            return Ok(result(outputs, 0));
198        }
199        let mut root_nodes = Vec::with_capacity(root_targets.len());
200        for occurrence in root_targets {
201            let node = self.required(occurrence.target)?;
202            outputs.push(loaded(
203                occurrence.target,
204                Some(occurrence.source),
205                &node,
206                false,
207            ));
208            let _ = states.insert(occurrence.target, NodeState::Opened);
209            root_nodes.push((occurrence.target, node, occurrence.strength));
210        }
211
212        let mut spent = root_cost;
213        let mut frontier = self.edges(node_id, &root, 1.0, &states)?;
214        for (node_id, node, strength) in &root_nodes {
215            frontier.extend(self.edges(*node_id, node, *strength, &states)?);
216        }
217        loop {
218            let candidates: Vec<Ticket> = frontier
219                .iter()
220                .enumerate()
221                .filter(|(_, occurrence)| !states.contains_key(&occurrence.target))
222                .filter_map(|(index, occurrence)| {
223                    (occurrence.strength.is_finite() && occurrence.strength > 0.0).then_some((
224                        index,
225                        occurrence.strength,
226                        0,
227                    ))
228                })
229                .collect();
230            if candidates.is_empty() {
231                break;
232            }
233
234            let choice = self.choose(&candidates)?;
235            let occurrence = frontier[candidates[choice].0].clone();
236            let node = self.required(occurrence.target)?;
237            let _ = states.insert(occurrence.target, NodeState::Opened);
238            let children = self.nav(occurrence.target, &node, occurrence.strength, &states)?;
239            let cost = preview_cost(attention(children.len().checked_add(1))?)?;
240            if !affordable(spent, cost, self.budget)? {
241                break;
242            }
243            outputs.push(loaded(
244                occurrence.target,
245                Some(occurrence.source),
246                &node,
247                false,
248            ));
249            let mut child_nodes = Vec::with_capacity(children.len());
250            for child in children {
251                let node = self.required(child.target)?;
252                let _ = states.insert(child.target, NodeState::Opened);
253                outputs.push(loaded(child.target, Some(child.source), &node, false));
254                child_nodes.push((child.target, node, child.strength));
255            }
256            let mut additions =
257                self.edges(occurrence.target, &node, occurrence.strength, &states)?;
258            for (node_id, child, strength) in &child_nodes {
259                additions.extend(self.edges(*node_id, child, *strength, &states)?);
260            }
261            frontier.retain(|entry| !states.contains_key(&entry.target));
262            frontier.extend(additions);
263            spent = attention(spent.checked_add(cost))?;
264        }
265        Ok(result(outputs, spent))
266    }
267
268    fn choose(&mut self, candidates: &[Ticket]) -> Result<usize, String> {
269        kcode_k1_kmap_selection::choose(candidates, self.temperature, &mut self.random)
270    }
271
272    fn required(&mut self, node_id: NodeId) -> Result<Node, String> {
273        (self.load_node)(node_id)?
274            .ok_or_else(|| format!("authorized Kmap target {node_id:?} is missing"))
275    }
276
277    fn resolve_candidates(&mut self, node: &Node) -> Result<(), String> {
278        let mut expected = HashSet::with_capacity(node.connections.len());
279        let requested: Vec<NodeId> = node
280            .connections
281            .iter()
282            .map(|connection| connection.target)
283            .filter(|target| !self.decisions.contains_key(target) && expected.insert(*target))
284            .collect();
285        if requested.is_empty() {
286            return Ok(());
287        }
288        let returned = (self.candidate_filter)(&requested)
289            .map_err(|error| format!("Kmap candidate filter failed: {error}"))?;
290        let mut allowed = HashSet::with_capacity(returned.len());
291        for target in returned {
292            if !expected.contains(&target) {
293                return Err(format!(
294                    "Kmap candidate filter returned unrequested node {target:?}"
295                ));
296            }
297            if !allowed.insert(target) {
298                return Err(format!(
299                    "Kmap candidate filter returned duplicate node {target:?}"
300                ));
301            }
302        }
303        self.decisions.extend(
304            requested
305                .into_iter()
306                .map(|target| (target, allowed.contains(&target))),
307        );
308        Ok(())
309    }
310
311    fn nav(
312        &mut self,
313        source: NodeId,
314        node: &Node,
315        inherited_strength: f64,
316        states: &HashMap<NodeId, NodeState>,
317    ) -> Result<Vec<Occurrence>, String> {
318        self.resolve_candidates(node)?;
319        Ok(node
320            .connections
321            .iter()
322            .filter(|connection| {
323                !states.contains_key(&connection.target)
324                    && self.decisions.get(&connection.target) == Some(&true)
325                    && connection.tier == ConnectionTier::Navigation
326            })
327            .map(|connection| Occurrence {
328                source,
329                target: connection.target,
330                strength: score(connection.weight.value, inherited_strength),
331            })
332            .collect())
333    }
334
335    fn edges(
336        &mut self,
337        source: NodeId,
338        node: &Node,
339        inherited_strength: f64,
340        states: &HashMap<NodeId, NodeState>,
341    ) -> Result<Vec<Occurrence>, String> {
342        self.resolve_candidates(node)?;
343        Ok(node
344            .connections
345            .iter()
346            .filter(|connection| {
347                !matches!(states.get(&connection.target), Some(NodeState::Opened))
348                    && self.decisions.get(&connection.target) == Some(&true)
349            })
350            .map(|connection| Occurrence {
351                source,
352                target: connection.target,
353                strength: score(connection.weight.value, inherited_strength),
354            })
355            .collect())
356    }
357}
358
359#[derive(Clone)]
360enum NodeState {
361    Previewed(Node),
362    Opened,
363}
364
365#[derive(Clone)]
366struct Occurrence {
367    source: NodeId,
368    target: NodeId,
369    strength: f64,
370}
371
372fn loaded(node_id: NodeId, source: Option<NodeId>, node: &Node, opened: bool) -> LoadedNode {
373    LoadedNode {
374        node_id,
375        source,
376        title: node.title.clone(),
377        navigation_hint: node.navigation_hint.clone(),
378        narrative: opened.then(|| node.narrative.clone()),
379    }
380}
381
382fn preview(
383    node_id: NodeId,
384    source: NodeId,
385    node: Node,
386    outputs: &mut Vec<LoadedNode>,
387    states: &mut HashMap<NodeId, NodeState>,
388) {
389    if let Entry::Vacant(entry) = states.entry(node_id) {
390        outputs.push(loaded(node_id, Some(source), &node, false));
391        entry.insert(NodeState::Previewed(node));
392    }
393}
394
395fn result(nodes: Vec<LoadedNode>, spent: u64) -> OpenResult {
396    OpenResult {
397        nodes,
398        automatic_attention_spent: spent as f64 / 10.0,
399    }
400}
401
402fn attention<T>(value: Option<T>) -> Result<T, String> {
403    value.ok_or_else(|| "Kmap attention cost overflow".to_owned())
404}
405
406fn preview_cost(count: usize) -> Result<u64, String> {
407    let count = attention(u64::try_from(count).ok())?;
408    attention(count.checked_mul(PREVIEW_TENTHS))
409}
410
411fn affordable(spent: u64, cost: u64, budget: f64) -> Result<bool, String> {
412    let total = attention(spent.checked_add(cost))? as f64 / 10.0;
413    let tolerance = 8.0 * f64::EPSILON * total.abs().max(budget.abs()).max(1.0);
414    Ok(total <= budget || total - budget <= tolerance)
415}
416
417#[cfg(test)]
418mod tests {
419    use std::{
420        cell::{Cell, RefCell},
421        collections::HashMap,
422    };
423
424    use kcode_k1_kmap_format::{Connection, ConnectionTier};
425
426    use super::{Node, NodeId, OpenMode, OpenResult, open_node_with_random};
427
428    fn id(value: u64) -> NodeId {
429        let mut bytes = [0; 12];
430        bytes[..8].copy_from_slice(&value.to_le_bytes());
431        NodeId(bytes)
432    }
433
434    fn node(connections: Vec<Connection>) -> Node {
435        Node {
436            title: String::new(),
437            navigation_hint: String::new(),
438            narrative: String::new(),
439            connections,
440        }
441    }
442
443    fn edge(target: NodeId, tier: ConnectionTier, weight: f64) -> Connection {
444        let mut connection = Connection::new(target, tier);
445        connection.weight.value = weight;
446        connection
447    }
448
449    fn filtered(returned: Result<Vec<NodeId>, String>) -> Result<OpenResult, String> {
450        let root = id(0);
451        let root_node = node(vec![Connection::new(id(1), ConnectionTier::Automated)]);
452        open_node_with_random(
453            root,
454            1.0,
455            1.0,
456            OpenMode::NavigationOnly,
457            |node_id| Ok((node_id == root).then(|| root_node.clone())),
458            move |_| returned.clone(),
459            || panic!("RNG invoked for rejected or denied candidates"),
460        )
461    }
462
463    #[test]
464    fn batches_large_denial_before_reads_and_randomness() {
465        let root = id(0);
466        let root_node = node(
467            (1..=10_000)
468                .map(|value| Connection::new(id(value), ConnectionTier::Automated))
469                .collect(),
470        );
471        let filters = Cell::new(0);
472        let result = open_node_with_random(
473            root,
474            10_000.0,
475            1.0,
476            OpenMode::NavigationOnly,
477            |node_id| Ok((node_id == root).then(|| root_node.clone())),
478            |targets| {
479                filters.set(filters.get() + 1);
480                assert_eq!(targets.len(), 10_000);
481                assert_eq!((targets[0], targets[9_999]), (id(1), id(10_000)));
482                Ok(Vec::new())
483            },
484            || panic!("RNG invoked after every candidate was denied"),
485        )
486        .unwrap();
487        assert_eq!((result.nodes.len(), result.nodes[0].source), (1, None));
488        assert_eq!(filters.get(), 1);
489    }
490
491    #[test]
492    fn propagates_and_validates_filter_results() {
493        assert_eq!(
494            filtered(Err("Access unavailable".to_owned())).unwrap_err(),
495            "Kmap candidate filter failed: Access unavailable"
496        );
497        assert!(
498            filtered(Ok(vec![id(1), id(1)]))
499                .unwrap_err()
500                .contains("returned duplicate node")
501        );
502        assert!(
503            filtered(Ok(vec![id(2)]))
504                .unwrap_err()
505                .contains("returned unrequested node")
506        );
507        assert_eq!(filtered(Ok(Vec::new())).unwrap().nodes.len(), 1);
508    }
509
510    #[test]
511    fn memoizes_visibility_across_repeated_paths() {
512        let (root, a, f, x, b, c, e, d) = (id(0), id(1), id(2), id(3), id(4), id(5), id(6), id(7));
513        let nodes = HashMap::from([
514            (
515                root,
516                node(vec![
517                    edge(a, ConnectionTier::Navigation, 0.5),
518                    edge(f, ConnectionTier::Navigation, 0.1),
519                    edge(x, ConnectionTier::Automated, 0.15),
520                ]),
521            ),
522            (a, node(vec![edge(b, ConnectionTier::Automated, 0.4)])),
523            (f, node(vec![edge(b, ConnectionTier::Automated, 0.9)])),
524            (
525                b,
526                node(vec![
527                    edge(e, ConnectionTier::Automated, 0.8),
528                    edge(c, ConnectionTier::Navigation, 0.5),
529                ]),
530            ),
531            (c, node(vec![edge(d, ConnectionTier::Automated, 1.0)])),
532            (x, node(Vec::new())),
533            (e, node(Vec::new())),
534            (d, node(Vec::new())),
535        ]);
536        let batches = RefCell::new(Vec::new());
537        let result = open_node_with_random(
538            root,
539            2.1,
540            0.0,
541            OpenMode::NavigationOnly,
542            |node_id| Ok(nodes.get(&node_id).cloned()),
543            |targets| {
544                batches.borrow_mut().push(targets.to_vec());
545                Ok(targets.to_vec())
546            },
547            || Ok(0.0),
548        )
549        .unwrap();
550        let ids = result
551            .nodes
552            .iter()
553            .map(|node| node.node_id)
554            .collect::<Vec<_>>();
555        let sources = result
556            .nodes
557            .iter()
558            .map(|node| node.source)
559            .collect::<Vec<_>>();
560        assert_eq!(ids, vec![root, a, f, b, c, e, x, d]);
561        assert_eq!(
562            sources,
563            vec![
564                None,
565                Some(root),
566                Some(root),
567                Some(a),
568                Some(b),
569                Some(b),
570                Some(root),
571                Some(c)
572            ]
573        );
574        assert_eq!(
575            batches.into_inner(),
576            vec![vec![a, f, x], vec![b], vec![e, c], vec![d]]
577        );
578    }
579}