kcode-k1-chat-thread-actions 0.5.0

In-memory Kmap actions and session state for K1 chat threads
Documentation
use std::collections::{HashMap, HashSet};
use std::sync::Arc;

pub use kcode_k1_access_kmap::{AccessContext, AccessPolicy, ProfileId};
use kcode_k1_access_kmap::{
    AccessId, ConnectionSpec, ConnectionTier, K1AccessKmap, LoadedNode, Measurement,
    MeasurementImportance, OpenMode, TxId,
};
use serde::Deserialize;
use serde::de::DeserializeOwned;
use serde_json::Value;

const NO_AUTHORIZATION: &str = "Kmap authorization is not active";
const DIFFERENT_AUTHORIZATION: &str =
    "different Kmap access context, profile ID, or policy is already active";
const UNKNOWN_TOOL: &str = "unknown Ktool";

pub struct ChatThreadActions {
    kmap: Arc<K1AccessKmap>,
    authorization: Option<Authorization>,
    session: SessionState,
}

#[derive(Clone, Eq, PartialEq)]
struct Authorization {
    context: AccessContext,
    profile_id: ProfileId,
    policy: AccessPolicy,
}

impl ChatThreadActions {
    pub fn new(kmap: Arc<K1AccessKmap>) -> Self {
        Self {
            kmap,
            authorization: None,
            session: SessionState::default(),
        }
    }

    pub fn bind_authorization(
        &mut self,
        context: AccessContext,
        profile_id: ProfileId,
        policy: AccessPolicy,
    ) -> Result<(), String> {
        bind_once(
            &mut self.authorization,
            Authorization {
                context,
                profile_id,
                policy,
            },
        )
    }

    pub fn clear_authorization(&mut self) {
        self.authorization = None;
    }

    pub fn launch(&mut self, name: &str, arguments: &str) -> Result<String, String> {
        match name {
            "CurrentTime" => launch_current_time(arguments),
            "KmapCreateNode" => {
                let parsed: CreateArguments = parse(name, arguments)?;
                let connections = navigation_connections(&parsed.connections, name)?;
                self.create(parsed, connections)
            }
            "KmapOpenNode" => {
                let parsed: OpenArguments = parse(name, arguments)?;
                let node = parse_id(&parsed.node_id, name)?;
                if !parsed.budget.is_finite() || parsed.budget < 0.0 {
                    return Err(invalid(name));
                }
                self.open_node(node, parsed.budget)
            }
            "KmapUpdateNode" => {
                let parsed: UpdateArguments = parse(name, arguments)?;
                let node = parse_id(&parsed.node_id, name)?;
                let connections = navigation_connections(&parsed.connections, name)?;
                self.update(node, parsed, connections)
            }
            "KmapPenalizeNodes" => {
                let parsed: PenalizeArguments = parse(name, arguments)?;
                let nodes = parse_ids(&parsed.node_ids, name)?;
                self.penalize(nodes)
            }
            "KmapConnectNodes" => {
                parse_empty(name, arguments)?;
                self.connect()
            }
            _ => Err(UNKNOWN_TOOL.to_owned()),
        }
    }

    fn authorization(&self) -> Result<Authorization, String> {
        self.authorization
            .clone()
            .ok_or_else(|| NO_AUTHORIZATION.to_owned())
    }

    fn create(
        &mut self,
        parsed: CreateArguments,
        connections: Vec<ConnectionSpec>,
    ) -> Result<String, String> {
        let authorization = self.authorization()?;
        let revision = self.kmap.create_node(
            &authorization.context,
            authorization.profile_id,
            authorization.policy,
            parsed.title,
            parsed.navigation_hint,
            parsed.narrative,
            connections,
        )?;
        let node = revision.access_id();
        self.session.add_loaded(node);
        Ok(format!(
            "Node successfully created with id {}",
            render_id(node)
        ))
    }

    fn open_node(&mut self, node: AccessId, budget: f64) -> Result<String, String> {
        let authorization = self.authorization()?;
        let result =
            self.kmap
                .open_node(&authorization.context, node, budget, 1.0, OpenMode::Full)?;
        Ok(self.session.record_open(result.nodes))
    }

    fn update(
        &mut self,
        node: AccessId,
        parsed: UpdateArguments,
        connections: Vec<ConnectionSpec>,
    ) -> Result<String, String> {
        let authorization = self.authorization()?;
        self.kmap.update_node(
            &authorization.context,
            node,
            Some(parsed.title),
            Some(parsed.navigation_hint),
            Some(parsed.narrative),
            connections,
        )?;
        self.session.add_loaded(node);
        Ok("success".to_owned())
    }

    fn penalize(&mut self, supplied: Vec<AccessId>) -> Result<String, String> {
        let authorization = self.authorization()?;
        let nodes = self.session.new_penalties(supplied);
        let measurements = self.session.measurements(&nodes);
        if !measurements.is_empty() {
            self.kmap
                .apply_measurements(&authorization.context, measurements)?;
        }
        self.session.commit_penalties(nodes);
        Ok("success".to_owned())
    }

    fn connect(&self) -> Result<String, String> {
        let authorization = self.authorization()?;
        let eligible = self.session.eligible();
        for source in &eligible {
            let node = self.kmap.get_node(&authorization.context, *source)?;
            let existing = node
                .connections
                .into_iter()
                .map(|connection| connection.target)
                .collect::<HashSet<_>>();
            let missing = eligible
                .iter()
                .copied()
                .filter(|target| target != source && !existing.contains(target))
                .map(|target| ConnectionSpec {
                    target,
                    tier: ConnectionTier::Automated,
                })
                .collect::<Vec<_>>();
            if !missing.is_empty() {
                self.kmap.update_node(
                    &authorization.context,
                    *source,
                    None,
                    None,
                    None,
                    missing,
                )?;
            }
        }
        Ok("success".to_owned())
    }
}

#[derive(Default)]
struct SessionState {
    loaded_order: Vec<AccessId>,
    loaded: HashSet<AccessId>,
    provenance: HashMap<AccessId, AccessId>,
    returned_previews: HashSet<AccessId>,
    returned_narratives: HashSet<AccessId>,
    penalized: HashSet<AccessId>,
}

impl SessionState {
    fn add_loaded(&mut self, node: AccessId) {
        if self.loaded.insert(node) {
            self.loaded_order.push(node);
        }
    }

    fn record_open(&mut self, nodes: Vec<LoadedNode>) -> String {
        let mut blocks = Vec::new();
        for node in nodes {
            self.add_loaded(node.access_id);
            if let Some(source) = node.source {
                self.provenance.entry(node.access_id).or_insert(source);
            }
            let preview = self.returned_previews.insert(node.access_id);
            let narrative =
                node.narrative.is_some() && self.returned_narratives.insert(node.access_id);
            if preview {
                let mut block = format!(
                    "Node ID: {}\nTitle: {}\nNavigation Hint: {}",
                    render_id(node.access_id),
                    node.title,
                    node.navigation_hint
                );
                if let (true, Some(text)) = (narrative, node.narrative.as_deref()) {
                    block.push_str("\nNarrative: ");
                    block.push_str(text);
                }
                blocks.push(block);
            } else if let (true, Some(text)) = (narrative, node.narrative.as_deref()) {
                blocks.push(format!(
                    "Node ID: {}\nNarrative: {text}",
                    render_id(node.access_id)
                ));
            }
        }
        if blocks.is_empty() {
            "No new Kmap node components.".to_owned()
        } else {
            blocks.join("\n\n")
        }
    }

    fn new_penalties(&self, supplied: Vec<AccessId>) -> Vec<AccessId> {
        let mut seen = HashSet::new();
        supplied
            .into_iter()
            .filter(|node| {
                seen.insert(*node) && self.loaded.contains(node) && !self.penalized.contains(node)
            })
            .collect()
    }

    fn measurements(&self, nodes: &[AccessId]) -> Vec<Measurement> {
        nodes
            .iter()
            .filter_map(|target| {
                self.provenance.get(target).map(|source| Measurement {
                    source: *source,
                    target: *target,
                    useful: false,
                    importance: MeasurementImportance::NonCritical,
                })
            })
            .collect()
    }

    fn commit_penalties(&mut self, nodes: Vec<AccessId>) {
        self.penalized.extend(nodes);
    }

    fn eligible(&self) -> Vec<AccessId> {
        self.loaded_order
            .iter()
            .copied()
            .filter(|node| !self.penalized.contains(node))
            .collect()
    }
}

fn bind_once<T: Eq>(slot: &mut Option<T>, value: T) -> Result<(), String> {
    match slot {
        Some(active) if active == &value => Ok(()),
        Some(_) => Err(DIFFERENT_AUTHORIZATION.to_owned()),
        None => {
            *slot = Some(value);
            Ok(())
        }
    }
}

fn launch_current_time(arguments: &str) -> Result<String, String> {
    parse_empty("CurrentTime", arguments)?;
    Ok(kcode_k1_ktool_current_time::current_time())
}

fn parse<T: DeserializeOwned>(name: &str, arguments: &str) -> Result<T, String> {
    serde_json::from_str(arguments).map_err(|_| invalid(name))
}

fn parse_empty(name: &str, arguments: &str) -> Result<(), String> {
    match serde_json::from_str::<Value>(arguments) {
        Ok(Value::Object(fields)) if fields.is_empty() => Ok(()),
        _ => Err(invalid(name)),
    }
}

fn parse_id(value: &str, name: &str) -> Result<AccessId, String> {
    value
        .parse::<TxId>()
        .map(AccessId::new)
        .map_err(|_| invalid(name))
}

fn parse_ids(values: &[String], name: &str) -> Result<Vec<AccessId>, String> {
    values.iter().map(|value| parse_id(value, name)).collect()
}

fn navigation_connections(values: &[String], name: &str) -> Result<Vec<ConnectionSpec>, String> {
    Ok(parse_ids(values, name)?
        .into_iter()
        .map(|target| ConnectionSpec {
            target,
            tier: ConnectionTier::Navigation,
        })
        .collect())
}

fn render_id(id: AccessId) -> String {
    id.txid().to_string()
}

fn invalid(name: &str) -> String {
    format!("invalid {name} arguments")
}

#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct CreateArguments {
    title: String,
    navigation_hint: String,
    narrative: String,
    connections: Vec<String>,
}

#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct OpenArguments {
    node_id: String,
    budget: f64,
}

#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct UpdateArguments {
    node_id: String,
    title: String,
    navigation_hint: String,
    narrative: String,
    connections: Vec<String>,
}

#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct PenalizeArguments {
    node_ids: Vec<String>,
}

#[cfg(test)]
mod tests;