use std::collections::HashMap;
use std::hash::Hash;
use std::sync::Arc;
use kcode_k1_access::{K1Access, RequestPrincipal, SubsystemId, Target};
use kcode_k1_access_profiles::K1AccessProfiles;
use kcode_k1_kmap::{
ConnectionMeasurement, ConnectionSpec as RawConnectionSpec, K1Kmap,
LoadedNode as RawLoadedNode, NodeId,
};
pub use kcode_k1_access::{AccessCheck, AccessId, AccessRevision, ModelId, UserId};
pub use kcode_k1_access_profiles::ProfileSelection as AccessProfile;
pub use kcode_k1_kmap::{ConnectionTier, MeasurementImportance, OpenMode, TxId, Weight};
const SUBSYSTEM_NAME: &str = "k1-kmap";
const NODE_UNAVAILABLE: &str = "node unavailable";
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ConnectionSpec {
pub target: AccessId,
pub tier: ConnectionTier,
}
#[derive(Clone, Debug, PartialEq)]
pub struct Connection {
pub target: AccessId,
pub tier: ConnectionTier,
pub weight: Weight,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct Measurement {
pub source: AccessId,
pub target: AccessId,
pub useful: bool,
pub importance: MeasurementImportance,
}
#[derive(Clone, Debug, PartialEq)]
pub struct Node {
pub access_id: AccessId,
pub title: String,
pub navigation_hint: String,
pub narrative: String,
pub connections: Vec<Connection>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct LoadedNode {
pub access_id: AccessId,
pub title: String,
pub navigation_hint: String,
pub narrative: Option<String>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct OpenResult {
pub nodes: Vec<LoadedNode>,
pub automatic_attention_spent: f64,
}
pub struct K1AccessKmap {
access: Arc<K1Access>,
profiles: Arc<K1AccessProfiles>,
kmap: Arc<K1Kmap>,
}
impl K1AccessKmap {
pub fn open(
access: Arc<K1Access>,
profiles: Arc<K1AccessProfiles>,
kmap: Arc<K1Kmap>,
) -> Result<Self, String> {
SubsystemId::from_str(SUBSYSTEM_NAME)
.map_err(|error| with_context("validate kmap subsystem", error))?;
Ok(Self {
access,
profiles,
kmap,
})
}
#[allow(clippy::too_many_arguments)]
pub fn create_node(
&self,
user: UserId,
model: ModelId,
access_profile: AccessProfile,
title: String,
navigation_hint: String,
narrative: String,
connections: Vec<ConnectionSpec>,
) -> Result<AccessRevision, String> {
let resolved = self
.profiles
.resolve(RequestPrincipal::new(user, model), access_profile)
.map_err(|error| with_context("resolve access profile", error))?;
let access_ids = connections
.iter()
.map(|connection| connection.target)
.collect::<Vec<_>>();
let raw_ids = self.authorize_nodes(
user,
model,
&access_ids,
&vec![false; access_ids.len()],
"authorize initial connections",
)?;
let raw_connections = connections
.into_iter()
.zip(raw_ids)
.map(|(connection, target)| RawConnectionSpec {
target,
tier: connection.tier,
})
.collect();
let revision = self
.kmap
.create_node(title, navigation_hint, narrative, raw_connections)
.map_err(|error| with_context("create kmap node", error))?;
self.access
.create(
target_from_node(NodeId::from(revision)),
resolved.into_authorizations(),
)
.map_err(|error| {
with_context(
"create node access after raw creation (possible inaccessible orphan)",
error,
)
})
}
pub fn get_node(&self, user: UserId, model: ModelId, node: AccessId) -> Result<Node, String> {
let raw_id = self
.authorize_nodes(user, model, &[node], &[false], "authorize node read")?
.into_iter()
.next()
.ok_or_else(inconsistent)?;
let raw = self
.kmap
.get_node(raw_id)
.map_err(|error| with_context("get kmap node", error))?
.ok_or_else(unavailable)?;
let targets = raw
.connections
.iter()
.map(|connection| target_from_node(connection.target))
.collect::<Vec<_>>();
let resolved = self
.access
.resolve_visible_targets(user, model, &targets, subsystem())
.map_err(|error| with_context("resolve visible connections", error))?;
if resolved.len() != raw.connections.len() {
return Err(inconsistent());
}
let connections = raw
.connections
.into_iter()
.zip(resolved)
.filter_map(|(connection, access_id)| {
access_id.map(|target| Connection {
target,
tier: connection.tier,
weight: connection.weight,
})
})
.collect();
Ok(Node {
access_id: node,
title: raw.title,
navigation_hint: raw.navigation_hint,
narrative: raw.narrative,
connections,
})
}
pub fn open_node(
&self,
user: UserId,
model: ModelId,
root: AccessId,
budget: f64,
temperature: f64,
mode: OpenMode,
) -> Result<OpenResult, String> {
let raw_root = self
.authorize_nodes(user, model, &[root], &[false], "authorize root node")?
.into_iter()
.next()
.ok_or_else(inconsistent)?;
let mut visible = HashMap::from([(raw_root, root)]);
let raw = self
.kmap
.open_node(
raw_root,
budget,
temperature,
mode,
|candidates: &[NodeId]| {
let targets = candidates
.iter()
.copied()
.map(target_from_node)
.collect::<Vec<_>>();
let resolved = self
.access
.resolve_visible_targets(user, model, &targets, subsystem())
.map_err(|error| with_context("filter open candidates", error))?;
if resolved.len() != candidates.len() {
return Err(inconsistent());
}
let mut allowed = Vec::new();
for (candidate, access_id) in candidates.iter().copied().zip(resolved) {
if let Some(access_id) = access_id {
visible.insert(candidate, access_id);
allowed.push(candidate);
}
}
Ok(allowed)
},
)
.map_err(|error| with_context("open kmap node", error))?;
let nodes = raw
.nodes
.into_iter()
.map(|loaded: RawLoadedNode| {
let access_id = visible.get(&loaded.node_id).copied().ok_or_else(|| {
with_context(
"open kmap node",
"dependency returned a loaded node without an authorized mapping",
)
})?;
Ok(LoadedNode {
access_id,
title: loaded.title,
navigation_hint: loaded.navigation_hint,
narrative: loaded.narrative,
})
})
.collect::<Result<Vec<_>, String>>()?;
Ok(OpenResult {
nodes,
automatic_attention_spent: raw.automatic_attention_spent,
})
}
#[allow(clippy::too_many_arguments)]
pub fn update_node(
&self,
user: UserId,
model: ModelId,
node: AccessId,
title: Option<String>,
navigation_hint: Option<String>,
narrative: Option<String>,
connection_updates: Vec<ConnectionSpec>,
) -> Result<TxId, String> {
let mut access_ids = Vec::with_capacity(connection_updates.len() + 1);
access_ids.push(node);
access_ids.extend(connection_updates.iter().map(|update| update.target));
let mut roles = vec![false; access_ids.len()];
roles[0] = true;
let raw_ids =
self.authorize_nodes(user, model, &access_ids, &roles, "authorize node update")?;
let (source, targets) = raw_ids.split_first().ok_or_else(inconsistent)?;
let updates = connection_updates
.into_iter()
.zip(targets.iter().copied())
.map(|(update, target)| RawConnectionSpec {
target,
tier: update.tier,
})
.collect();
self.kmap
.update_node(*source, title, navigation_hint, narrative, updates)
.map_err(|error| with_context("update kmap node", error))
}
pub fn apply_measurements(
&self,
user: UserId,
model: ModelId,
measurements: Vec<Measurement>,
) -> Result<TxId, String> {
let roles = merged_roles(
measurements
.iter()
.flat_map(|measurement| [(measurement.source, true), (measurement.target, false)]),
);
let access_ids = roles.iter().map(|(id, _)| *id).collect::<Vec<_>>();
let required = roles.iter().map(|(_, manage)| *manage).collect::<Vec<_>>();
let raw_ids = self.authorize_nodes(
user,
model,
&access_ids,
&required,
"authorize measurements",
)?;
let raw_by_access = access_ids
.into_iter()
.zip(raw_ids)
.collect::<HashMap<_, _>>();
let raw_measurements = measurements
.into_iter()
.map(|measurement| {
let source = raw_by_access
.get(&measurement.source)
.copied()
.ok_or_else(inconsistent)?;
let target = raw_by_access
.get(&measurement.target)
.copied()
.ok_or_else(inconsistent)?;
Ok(ConnectionMeasurement::new(
source,
target,
measurement.useful,
measurement.importance,
))
})
.collect::<Result<Vec<_>, String>>()?;
self.kmap
.apply_measurements(raw_measurements)
.map_err(|error| with_context("apply kmap measurements", error))
}
fn authorize_nodes(
&self,
user: UserId,
model: ModelId,
access_ids: &[AccessId],
require_manage: &[bool],
operation: &str,
) -> Result<Vec<NodeId>, String> {
if access_ids.len() != require_manage.len() {
return Err(inconsistent());
}
let checks = self
.access
.check_many(user, model, access_ids, subsystem())
.map_err(|error| with_context(operation, error))?;
if checks.len() != access_ids.len() {
return Err(inconsistent());
}
checks
.iter()
.zip(require_manage.iter().copied())
.map(|(check, manage)| checked_node_id(check, manage))
.collect()
}
}
fn subsystem() -> SubsystemId {
SubsystemId::from_str(SUBSYSTEM_NAME).expect("validated subsystem literal")
}
fn target_from_node(node_id: NodeId) -> Target {
Target::new(subsystem(), node_id.0.to_vec())
}
fn node_id_from_target(target: &Target) -> Result<NodeId, String> {
if target.subsystem() != subsystem() {
return Err(unavailable());
}
let bytes = target.object_id().try_into().map_err(|_| unavailable())?;
Ok(NodeId(bytes))
}
fn checked_node_id(check: &AccessCheck, require_manage: bool) -> Result<NodeId, String> {
if !check.can_view() || (require_manage && !check.can_manage()) {
return Err(unavailable());
}
node_id_from_target(check.target().ok_or_else(unavailable)?)
}
fn merged_roles<K>(uses: impl IntoIterator<Item = (K, bool)>) -> Vec<(K, bool)>
where
K: Copy + Eq + Hash,
{
let mut positions: HashMap<K, usize> = HashMap::new();
let mut roles: Vec<(K, bool)> = Vec::new();
for (key, manage) in uses {
if let Some(position) = positions.get(&key).copied() {
roles[position].1 |= manage;
} else {
positions.insert(key, roles.len());
roles.push((key, manage));
}
}
roles
}
fn unavailable() -> String {
NODE_UNAVAILABLE.to_owned()
}
fn inconsistent() -> String {
"dependency returned an inconsistent positional response".to_owned()
}
fn with_context(operation: &str, error: impl std::fmt::Display) -> String {
format!("{operation}: {error}")
}
#[cfg(test)]
mod tests {
use super::*;
type Open =
fn(Arc<K1Access>, Arc<K1AccessProfiles>, Arc<K1Kmap>) -> Result<K1AccessKmap, String>;
type Create = fn(
&K1AccessKmap,
UserId,
ModelId,
AccessProfile,
String,
String,
String,
Vec<ConnectionSpec>,
) -> Result<AccessRevision, String>;
type OpenNode = fn(
&K1AccessKmap,
UserId,
ModelId,
AccessId,
f64,
f64,
OpenMode,
) -> Result<OpenResult, String>;
type Update = fn(
&K1AccessKmap,
UserId,
ModelId,
AccessId,
Option<String>,
Option<String>,
Option<String>,
Vec<ConnectionSpec>,
) -> Result<TxId, String>;
#[test]
fn target_conversion_is_exact_and_safe() {
let node = NodeId([7; 12]);
let target = target_from_node(node);
assert_eq!(target.object_id(), &[7; 12]);
assert_eq!(node_id_from_target(&target), Ok(node));
let malformed = Target::new(subsystem(), vec![7; 11]);
assert_eq!(node_id_from_target(&malformed), Err(unavailable()));
}
#[test]
fn measurement_roles_deduplicate_and_upgrade() {
let roles = merged_roles([(1_u8, false), (2, true), (1, true), (2, false)]);
assert_eq!(roles, vec![(1, true), (2, true)]);
}
#[test]
fn approved_public_signatures_compile() {
let _: Open = K1AccessKmap::open;
let _: Create = K1AccessKmap::create_node;
let _: fn(&K1AccessKmap, UserId, ModelId, AccessId) -> Result<Node, String> =
K1AccessKmap::get_node;
let _: OpenNode = K1AccessKmap::open_node;
let _: Update = K1AccessKmap::update_node;
let _: fn(&K1AccessKmap, UserId, ModelId, Vec<Measurement>) -> Result<TxId, String> =
K1AccessKmap::apply_measurements;
}
}