Skip to main content

kcode_k1_access_kmap/
lib.rs

1use std::collections::HashMap;
2use std::hash::Hash;
3use std::sync::Arc;
4
5use kcode_k1_access::{K1Access, RequestPrincipal, SubsystemId, Target};
6use kcode_k1_access_profiles::K1AccessProfiles;
7use kcode_k1_kmap::{
8    ConnectionMeasurement, ConnectionSpec as RawConnectionSpec, K1Kmap,
9    LoadedNode as RawLoadedNode, NodeId,
10};
11
12pub use kcode_k1_access::{AccessCheck, AccessId, AccessRevision, ModelId, UserId};
13pub use kcode_k1_access_profiles::ProfileSelection as AccessProfile;
14pub use kcode_k1_kmap::{ConnectionTier, MeasurementImportance, OpenMode, TxId, Weight};
15
16const SUBSYSTEM_NAME: &str = "k1-kmap";
17const NODE_UNAVAILABLE: &str = "node unavailable";
18
19#[derive(Clone, Copy, Debug, Eq, PartialEq)]
20pub struct ConnectionSpec {
21    pub target: AccessId,
22    pub tier: ConnectionTier,
23}
24
25#[derive(Clone, Debug, PartialEq)]
26pub struct Connection {
27    pub target: AccessId,
28    pub tier: ConnectionTier,
29    pub weight: Weight,
30}
31
32#[derive(Clone, Copy, Debug, Eq, PartialEq)]
33pub struct Measurement {
34    pub source: AccessId,
35    pub target: AccessId,
36    pub useful: bool,
37    pub importance: MeasurementImportance,
38}
39
40#[derive(Clone, Debug, PartialEq)]
41pub struct Node {
42    pub access_id: AccessId,
43    pub title: String,
44    pub navigation_hint: String,
45    pub narrative: String,
46    pub connections: Vec<Connection>,
47}
48
49#[derive(Clone, Debug, PartialEq)]
50pub struct LoadedNode {
51    pub access_id: AccessId,
52    pub title: String,
53    pub navigation_hint: String,
54    pub narrative: Option<String>,
55}
56
57#[derive(Clone, Debug, PartialEq)]
58pub struct OpenResult {
59    pub nodes: Vec<LoadedNode>,
60    pub automatic_attention_spent: f64,
61}
62
63pub struct K1AccessKmap {
64    access: Arc<K1Access>,
65    profiles: Arc<K1AccessProfiles>,
66    kmap: Arc<K1Kmap>,
67}
68
69impl K1AccessKmap {
70    pub fn open(
71        access: Arc<K1Access>,
72        profiles: Arc<K1AccessProfiles>,
73        kmap: Arc<K1Kmap>,
74    ) -> Result<Self, String> {
75        SubsystemId::from_str(SUBSYSTEM_NAME)
76            .map_err(|error| with_context("validate kmap subsystem", error))?;
77        Ok(Self {
78            access,
79            profiles,
80            kmap,
81        })
82    }
83
84    #[allow(clippy::too_many_arguments)]
85    pub fn create_node(
86        &self,
87        user: UserId,
88        model: ModelId,
89        access_profile: AccessProfile,
90        title: String,
91        navigation_hint: String,
92        narrative: String,
93        connections: Vec<ConnectionSpec>,
94    ) -> Result<AccessRevision, String> {
95        let resolved = self
96            .profiles
97            .resolve(RequestPrincipal::new(user, model), access_profile)
98            .map_err(|error| with_context("resolve access profile", error))?;
99        let access_ids = connections
100            .iter()
101            .map(|connection| connection.target)
102            .collect::<Vec<_>>();
103        let raw_ids = self.authorize_nodes(
104            user,
105            model,
106            &access_ids,
107            &vec![false; access_ids.len()],
108            "authorize initial connections",
109        )?;
110        let raw_connections = connections
111            .into_iter()
112            .zip(raw_ids)
113            .map(|(connection, target)| RawConnectionSpec {
114                target,
115                tier: connection.tier,
116            })
117            .collect();
118        let revision = self
119            .kmap
120            .create_node(title, navigation_hint, narrative, raw_connections)
121            .map_err(|error| with_context("create kmap node", error))?;
122        self.access
123            .create(
124                target_from_node(NodeId::from(revision)),
125                resolved.into_authorizations(),
126            )
127            .map_err(|error| {
128                with_context(
129                    "create node access after raw creation (possible inaccessible orphan)",
130                    error,
131                )
132            })
133    }
134
135    pub fn get_node(&self, user: UserId, model: ModelId, node: AccessId) -> Result<Node, String> {
136        let raw_id = self
137            .authorize_nodes(user, model, &[node], &[false], "authorize node read")?
138            .into_iter()
139            .next()
140            .ok_or_else(inconsistent)?;
141        let raw = self
142            .kmap
143            .get_node(raw_id)
144            .map_err(|error| with_context("get kmap node", error))?
145            .ok_or_else(unavailable)?;
146        let targets = raw
147            .connections
148            .iter()
149            .map(|connection| target_from_node(connection.target))
150            .collect::<Vec<_>>();
151        let resolved = self
152            .access
153            .resolve_visible_targets(user, model, &targets, subsystem())
154            .map_err(|error| with_context("resolve visible connections", error))?;
155        if resolved.len() != raw.connections.len() {
156            return Err(inconsistent());
157        }
158        let connections = raw
159            .connections
160            .into_iter()
161            .zip(resolved)
162            .filter_map(|(connection, access_id)| {
163                access_id.map(|target| Connection {
164                    target,
165                    tier: connection.tier,
166                    weight: connection.weight,
167                })
168            })
169            .collect();
170        Ok(Node {
171            access_id: node,
172            title: raw.title,
173            navigation_hint: raw.navigation_hint,
174            narrative: raw.narrative,
175            connections,
176        })
177    }
178
179    pub fn open_node(
180        &self,
181        user: UserId,
182        model: ModelId,
183        root: AccessId,
184        budget: f64,
185        temperature: f64,
186        mode: OpenMode,
187    ) -> Result<OpenResult, String> {
188        let raw_root = self
189            .authorize_nodes(user, model, &[root], &[false], "authorize root node")?
190            .into_iter()
191            .next()
192            .ok_or_else(inconsistent)?;
193        let mut visible = HashMap::from([(raw_root, root)]);
194        let raw = self
195            .kmap
196            .open_node(
197                raw_root,
198                budget,
199                temperature,
200                mode,
201                |candidates: &[NodeId]| {
202                    let targets = candidates
203                        .iter()
204                        .copied()
205                        .map(target_from_node)
206                        .collect::<Vec<_>>();
207                    let resolved = self
208                        .access
209                        .resolve_visible_targets(user, model, &targets, subsystem())
210                        .map_err(|error| with_context("filter open candidates", error))?;
211                    if resolved.len() != candidates.len() {
212                        return Err(inconsistent());
213                    }
214                    let mut allowed = Vec::new();
215                    for (candidate, access_id) in candidates.iter().copied().zip(resolved) {
216                        if let Some(access_id) = access_id {
217                            visible.insert(candidate, access_id);
218                            allowed.push(candidate);
219                        }
220                    }
221                    Ok(allowed)
222                },
223            )
224            .map_err(|error| with_context("open kmap node", error))?;
225        let nodes = raw
226            .nodes
227            .into_iter()
228            .map(|loaded: RawLoadedNode| {
229                let access_id = visible.get(&loaded.node_id).copied().ok_or_else(|| {
230                    with_context(
231                        "open kmap node",
232                        "dependency returned a loaded node without an authorized mapping",
233                    )
234                })?;
235                Ok(LoadedNode {
236                    access_id,
237                    title: loaded.title,
238                    navigation_hint: loaded.navigation_hint,
239                    narrative: loaded.narrative,
240                })
241            })
242            .collect::<Result<Vec<_>, String>>()?;
243        Ok(OpenResult {
244            nodes,
245            automatic_attention_spent: raw.automatic_attention_spent,
246        })
247    }
248
249    #[allow(clippy::too_many_arguments)]
250    pub fn update_node(
251        &self,
252        user: UserId,
253        model: ModelId,
254        node: AccessId,
255        title: Option<String>,
256        navigation_hint: Option<String>,
257        narrative: Option<String>,
258        connection_updates: Vec<ConnectionSpec>,
259    ) -> Result<TxId, String> {
260        let mut access_ids = Vec::with_capacity(connection_updates.len() + 1);
261        access_ids.push(node);
262        access_ids.extend(connection_updates.iter().map(|update| update.target));
263        let mut roles = vec![false; access_ids.len()];
264        roles[0] = true;
265        let raw_ids =
266            self.authorize_nodes(user, model, &access_ids, &roles, "authorize node update")?;
267        let (source, targets) = raw_ids.split_first().ok_or_else(inconsistent)?;
268        let updates = connection_updates
269            .into_iter()
270            .zip(targets.iter().copied())
271            .map(|(update, target)| RawConnectionSpec {
272                target,
273                tier: update.tier,
274            })
275            .collect();
276        self.kmap
277            .update_node(*source, title, navigation_hint, narrative, updates)
278            .map_err(|error| with_context("update kmap node", error))
279    }
280
281    pub fn apply_measurements(
282        &self,
283        user: UserId,
284        model: ModelId,
285        measurements: Vec<Measurement>,
286    ) -> Result<TxId, String> {
287        let roles = merged_roles(
288            measurements
289                .iter()
290                .flat_map(|measurement| [(measurement.source, true), (measurement.target, false)]),
291        );
292        let access_ids = roles.iter().map(|(id, _)| *id).collect::<Vec<_>>();
293        let required = roles.iter().map(|(_, manage)| *manage).collect::<Vec<_>>();
294        let raw_ids = self.authorize_nodes(
295            user,
296            model,
297            &access_ids,
298            &required,
299            "authorize measurements",
300        )?;
301        let raw_by_access = access_ids
302            .into_iter()
303            .zip(raw_ids)
304            .collect::<HashMap<_, _>>();
305        let raw_measurements = measurements
306            .into_iter()
307            .map(|measurement| {
308                let source = raw_by_access
309                    .get(&measurement.source)
310                    .copied()
311                    .ok_or_else(inconsistent)?;
312                let target = raw_by_access
313                    .get(&measurement.target)
314                    .copied()
315                    .ok_or_else(inconsistent)?;
316                Ok(ConnectionMeasurement::new(
317                    source,
318                    target,
319                    measurement.useful,
320                    measurement.importance,
321                ))
322            })
323            .collect::<Result<Vec<_>, String>>()?;
324        self.kmap
325            .apply_measurements(raw_measurements)
326            .map_err(|error| with_context("apply kmap measurements", error))
327    }
328
329    fn authorize_nodes(
330        &self,
331        user: UserId,
332        model: ModelId,
333        access_ids: &[AccessId],
334        require_manage: &[bool],
335        operation: &str,
336    ) -> Result<Vec<NodeId>, String> {
337        if access_ids.len() != require_manage.len() {
338            return Err(inconsistent());
339        }
340        let checks = self
341            .access
342            .check_many(user, model, access_ids, subsystem())
343            .map_err(|error| with_context(operation, error))?;
344        if checks.len() != access_ids.len() {
345            return Err(inconsistent());
346        }
347        checks
348            .iter()
349            .zip(require_manage.iter().copied())
350            .map(|(check, manage)| checked_node_id(check, manage))
351            .collect()
352    }
353}
354
355fn subsystem() -> SubsystemId {
356    SubsystemId::from_str(SUBSYSTEM_NAME).expect("validated subsystem literal")
357}
358
359fn target_from_node(node_id: NodeId) -> Target {
360    Target::new(subsystem(), node_id.0.to_vec())
361}
362
363fn node_id_from_target(target: &Target) -> Result<NodeId, String> {
364    if target.subsystem() != subsystem() {
365        return Err(unavailable());
366    }
367    let bytes = target.object_id().try_into().map_err(|_| unavailable())?;
368    Ok(NodeId(bytes))
369}
370
371fn checked_node_id(check: &AccessCheck, require_manage: bool) -> Result<NodeId, String> {
372    if !check.can_view() || (require_manage && !check.can_manage()) {
373        return Err(unavailable());
374    }
375    node_id_from_target(check.target().ok_or_else(unavailable)?)
376}
377
378fn merged_roles<K>(uses: impl IntoIterator<Item = (K, bool)>) -> Vec<(K, bool)>
379where
380    K: Copy + Eq + Hash,
381{
382    let mut positions: HashMap<K, usize> = HashMap::new();
383    let mut roles: Vec<(K, bool)> = Vec::new();
384    for (key, manage) in uses {
385        if let Some(position) = positions.get(&key).copied() {
386            roles[position].1 |= manage;
387        } else {
388            positions.insert(key, roles.len());
389            roles.push((key, manage));
390        }
391    }
392    roles
393}
394
395fn unavailable() -> String {
396    NODE_UNAVAILABLE.to_owned()
397}
398
399fn inconsistent() -> String {
400    "dependency returned an inconsistent positional response".to_owned()
401}
402
403fn with_context(operation: &str, error: impl std::fmt::Display) -> String {
404    format!("{operation}: {error}")
405}
406
407#[cfg(test)]
408mod tests {
409    use super::*;
410
411    type Open =
412        fn(Arc<K1Access>, Arc<K1AccessProfiles>, Arc<K1Kmap>) -> Result<K1AccessKmap, String>;
413    type Create = fn(
414        &K1AccessKmap,
415        UserId,
416        ModelId,
417        AccessProfile,
418        String,
419        String,
420        String,
421        Vec<ConnectionSpec>,
422    ) -> Result<AccessRevision, String>;
423    type OpenNode = fn(
424        &K1AccessKmap,
425        UserId,
426        ModelId,
427        AccessId,
428        f64,
429        f64,
430        OpenMode,
431    ) -> Result<OpenResult, String>;
432    type Update = fn(
433        &K1AccessKmap,
434        UserId,
435        ModelId,
436        AccessId,
437        Option<String>,
438        Option<String>,
439        Option<String>,
440        Vec<ConnectionSpec>,
441    ) -> Result<TxId, String>;
442
443    #[test]
444    fn target_conversion_is_exact_and_safe() {
445        let node = NodeId([7; 12]);
446        let target = target_from_node(node);
447        assert_eq!(target.object_id(), &[7; 12]);
448        assert_eq!(node_id_from_target(&target), Ok(node));
449        let malformed = Target::new(subsystem(), vec![7; 11]);
450        assert_eq!(node_id_from_target(&malformed), Err(unavailable()));
451    }
452
453    #[test]
454    fn measurement_roles_deduplicate_and_upgrade() {
455        let roles = merged_roles([(1_u8, false), (2, true), (1, true), (2, false)]);
456        assert_eq!(roles, vec![(1, true), (2, true)]);
457    }
458
459    #[test]
460    fn approved_public_signatures_compile() {
461        let _: Open = K1AccessKmap::open;
462        let _: Create = K1AccessKmap::create_node;
463        let _: fn(&K1AccessKmap, UserId, ModelId, AccessId) -> Result<Node, String> =
464            K1AccessKmap::get_node;
465        let _: OpenNode = K1AccessKmap::open_node;
466        let _: Update = K1AccessKmap::update_node;
467        let _: fn(&K1AccessKmap, UserId, ModelId, Vec<Measurement>) -> Result<TxId, String> =
468            K1AccessKmap::apply_measurements;
469    }
470}