use std::collections::{HashMap, HashSet};
use std::sync::Arc;
pub use kcode_k1_access_kmap::{AccessContext, AccessPolicy};
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 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,
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,
policy: AccessPolicy,
) -> Result<(), String> {
bind_once(&mut self.authorization, Authorization { context, 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.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;