Skip to main content

kcode_k1_chat_thread_actions/
lib.rs

1use std::collections::{HashMap, HashSet};
2use std::sync::Arc;
3
4pub use kcode_k1_access_kmap::{AccessContext, AccessPolicy};
5use kcode_k1_access_kmap::{
6    AccessId, ConnectionSpec, ConnectionTier, K1AccessKmap, LoadedNode, Measurement,
7    MeasurementImportance, OpenMode, TxId,
8};
9use serde::Deserialize;
10use serde::de::DeserializeOwned;
11use serde_json::Value;
12
13const NO_AUTHORIZATION: &str = "Kmap authorization is not active";
14const DIFFERENT_AUTHORIZATION: &str = "different Kmap access context or policy is already active";
15const UNKNOWN_TOOL: &str = "unknown Ktool";
16
17pub struct ChatThreadActions {
18    kmap: Arc<K1AccessKmap>,
19    authorization: Option<Authorization>,
20    session: SessionState,
21}
22
23#[derive(Clone, Eq, PartialEq)]
24struct Authorization {
25    context: AccessContext,
26    policy: AccessPolicy,
27}
28
29impl ChatThreadActions {
30    pub fn new(kmap: Arc<K1AccessKmap>) -> Self {
31        Self {
32            kmap,
33            authorization: None,
34            session: SessionState::default(),
35        }
36    }
37
38    pub fn bind_authorization(
39        &mut self,
40        context: AccessContext,
41        policy: AccessPolicy,
42    ) -> Result<(), String> {
43        bind_once(&mut self.authorization, Authorization { context, policy })
44    }
45
46    pub fn clear_authorization(&mut self) {
47        self.authorization = None;
48    }
49
50    pub fn launch(&mut self, name: &str, arguments: &str) -> Result<String, String> {
51        match name {
52            "CurrentTime" => launch_current_time(arguments),
53            "KmapCreateNode" => {
54                let parsed: CreateArguments = parse(name, arguments)?;
55                let connections = navigation_connections(&parsed.connections, name)?;
56                self.create(parsed, connections)
57            }
58            "KmapOpenNode" => {
59                let parsed: OpenArguments = parse(name, arguments)?;
60                let node = parse_id(&parsed.node_id, name)?;
61                if !parsed.budget.is_finite() || parsed.budget < 0.0 {
62                    return Err(invalid(name));
63                }
64                self.open_node(node, parsed.budget)
65            }
66            "KmapUpdateNode" => {
67                let parsed: UpdateArguments = parse(name, arguments)?;
68                let node = parse_id(&parsed.node_id, name)?;
69                let connections = navigation_connections(&parsed.connections, name)?;
70                self.update(node, parsed, connections)
71            }
72            "KmapPenalizeNodes" => {
73                let parsed: PenalizeArguments = parse(name, arguments)?;
74                let nodes = parse_ids(&parsed.node_ids, name)?;
75                self.penalize(nodes)
76            }
77            "KmapConnectNodes" => {
78                parse_empty(name, arguments)?;
79                self.connect()
80            }
81            _ => Err(UNKNOWN_TOOL.to_owned()),
82        }
83    }
84
85    fn authorization(&self) -> Result<Authorization, String> {
86        self.authorization
87            .clone()
88            .ok_or_else(|| NO_AUTHORIZATION.to_owned())
89    }
90
91    fn create(
92        &mut self,
93        parsed: CreateArguments,
94        connections: Vec<ConnectionSpec>,
95    ) -> Result<String, String> {
96        let authorization = self.authorization()?;
97        let revision = self.kmap.create_node(
98            &authorization.context,
99            authorization.policy,
100            parsed.title,
101            parsed.navigation_hint,
102            parsed.narrative,
103            connections,
104        )?;
105        let node = revision.access_id();
106        self.session.add_loaded(node);
107        Ok(format!(
108            "Node successfully created with id {}",
109            render_id(node)
110        ))
111    }
112
113    fn open_node(&mut self, node: AccessId, budget: f64) -> Result<String, String> {
114        let authorization = self.authorization()?;
115        let result =
116            self.kmap
117                .open_node(&authorization.context, node, budget, 1.0, OpenMode::Full)?;
118        Ok(self.session.record_open(result.nodes))
119    }
120
121    fn update(
122        &mut self,
123        node: AccessId,
124        parsed: UpdateArguments,
125        connections: Vec<ConnectionSpec>,
126    ) -> Result<String, String> {
127        let authorization = self.authorization()?;
128        self.kmap.update_node(
129            &authorization.context,
130            node,
131            Some(parsed.title),
132            Some(parsed.navigation_hint),
133            Some(parsed.narrative),
134            connections,
135        )?;
136        self.session.add_loaded(node);
137        Ok("success".to_owned())
138    }
139
140    fn penalize(&mut self, supplied: Vec<AccessId>) -> Result<String, String> {
141        let authorization = self.authorization()?;
142        let nodes = self.session.new_penalties(supplied);
143        let measurements = self.session.measurements(&nodes);
144        if !measurements.is_empty() {
145            self.kmap
146                .apply_measurements(&authorization.context, measurements)?;
147        }
148        self.session.commit_penalties(nodes);
149        Ok("success".to_owned())
150    }
151
152    fn connect(&self) -> Result<String, String> {
153        let authorization = self.authorization()?;
154        let eligible = self.session.eligible();
155        for source in &eligible {
156            let node = self.kmap.get_node(&authorization.context, *source)?;
157            let existing = node
158                .connections
159                .into_iter()
160                .map(|connection| connection.target)
161                .collect::<HashSet<_>>();
162            let missing = eligible
163                .iter()
164                .copied()
165                .filter(|target| target != source && !existing.contains(target))
166                .map(|target| ConnectionSpec {
167                    target,
168                    tier: ConnectionTier::Automated,
169                })
170                .collect::<Vec<_>>();
171            if !missing.is_empty() {
172                self.kmap.update_node(
173                    &authorization.context,
174                    *source,
175                    None,
176                    None,
177                    None,
178                    missing,
179                )?;
180            }
181        }
182        Ok("success".to_owned())
183    }
184}
185
186#[derive(Default)]
187struct SessionState {
188    loaded_order: Vec<AccessId>,
189    loaded: HashSet<AccessId>,
190    provenance: HashMap<AccessId, AccessId>,
191    returned_previews: HashSet<AccessId>,
192    returned_narratives: HashSet<AccessId>,
193    penalized: HashSet<AccessId>,
194}
195
196impl SessionState {
197    fn add_loaded(&mut self, node: AccessId) {
198        if self.loaded.insert(node) {
199            self.loaded_order.push(node);
200        }
201    }
202
203    fn record_open(&mut self, nodes: Vec<LoadedNode>) -> String {
204        let mut blocks = Vec::new();
205        for node in nodes {
206            self.add_loaded(node.access_id);
207            if let Some(source) = node.source {
208                self.provenance.entry(node.access_id).or_insert(source);
209            }
210            let preview = self.returned_previews.insert(node.access_id);
211            let narrative =
212                node.narrative.is_some() && self.returned_narratives.insert(node.access_id);
213            if preview {
214                let mut block = format!(
215                    "Node ID: {}\nTitle: {}\nNavigation Hint: {}",
216                    render_id(node.access_id),
217                    node.title,
218                    node.navigation_hint
219                );
220                if let (true, Some(text)) = (narrative, node.narrative.as_deref()) {
221                    block.push_str("\nNarrative: ");
222                    block.push_str(text);
223                }
224                blocks.push(block);
225            } else if let (true, Some(text)) = (narrative, node.narrative.as_deref()) {
226                blocks.push(format!(
227                    "Node ID: {}\nNarrative: {text}",
228                    render_id(node.access_id)
229                ));
230            }
231        }
232        if blocks.is_empty() {
233            "No new Kmap node components.".to_owned()
234        } else {
235            blocks.join("\n\n")
236        }
237    }
238
239    fn new_penalties(&self, supplied: Vec<AccessId>) -> Vec<AccessId> {
240        let mut seen = HashSet::new();
241        supplied
242            .into_iter()
243            .filter(|node| {
244                seen.insert(*node) && self.loaded.contains(node) && !self.penalized.contains(node)
245            })
246            .collect()
247    }
248
249    fn measurements(&self, nodes: &[AccessId]) -> Vec<Measurement> {
250        nodes
251            .iter()
252            .filter_map(|target| {
253                self.provenance.get(target).map(|source| Measurement {
254                    source: *source,
255                    target: *target,
256                    useful: false,
257                    importance: MeasurementImportance::NonCritical,
258                })
259            })
260            .collect()
261    }
262
263    fn commit_penalties(&mut self, nodes: Vec<AccessId>) {
264        self.penalized.extend(nodes);
265    }
266
267    fn eligible(&self) -> Vec<AccessId> {
268        self.loaded_order
269            .iter()
270            .copied()
271            .filter(|node| !self.penalized.contains(node))
272            .collect()
273    }
274}
275
276fn bind_once<T: Eq>(slot: &mut Option<T>, value: T) -> Result<(), String> {
277    match slot {
278        Some(active) if active == &value => Ok(()),
279        Some(_) => Err(DIFFERENT_AUTHORIZATION.to_owned()),
280        None => {
281            *slot = Some(value);
282            Ok(())
283        }
284    }
285}
286
287fn launch_current_time(arguments: &str) -> Result<String, String> {
288    parse_empty("CurrentTime", arguments)?;
289    Ok(kcode_k1_ktool_current_time::current_time())
290}
291
292fn parse<T: DeserializeOwned>(name: &str, arguments: &str) -> Result<T, String> {
293    serde_json::from_str(arguments).map_err(|_| invalid(name))
294}
295
296fn parse_empty(name: &str, arguments: &str) -> Result<(), String> {
297    match serde_json::from_str::<Value>(arguments) {
298        Ok(Value::Object(fields)) if fields.is_empty() => Ok(()),
299        _ => Err(invalid(name)),
300    }
301}
302
303fn parse_id(value: &str, name: &str) -> Result<AccessId, String> {
304    value
305        .parse::<TxId>()
306        .map(AccessId::new)
307        .map_err(|_| invalid(name))
308}
309
310fn parse_ids(values: &[String], name: &str) -> Result<Vec<AccessId>, String> {
311    values.iter().map(|value| parse_id(value, name)).collect()
312}
313
314fn navigation_connections(values: &[String], name: &str) -> Result<Vec<ConnectionSpec>, String> {
315    Ok(parse_ids(values, name)?
316        .into_iter()
317        .map(|target| ConnectionSpec {
318            target,
319            tier: ConnectionTier::Navigation,
320        })
321        .collect())
322}
323
324fn render_id(id: AccessId) -> String {
325    id.txid().to_string()
326}
327
328fn invalid(name: &str) -> String {
329    format!("invalid {name} arguments")
330}
331
332#[derive(Deserialize)]
333#[serde(deny_unknown_fields)]
334struct CreateArguments {
335    title: String,
336    navigation_hint: String,
337    narrative: String,
338    connections: Vec<String>,
339}
340
341#[derive(Deserialize)]
342#[serde(deny_unknown_fields)]
343struct OpenArguments {
344    node_id: String,
345    budget: f64,
346}
347
348#[derive(Deserialize)]
349#[serde(deny_unknown_fields)]
350struct UpdateArguments {
351    node_id: String,
352    title: String,
353    navigation_hint: String,
354    narrative: String,
355    connections: Vec<String>,
356}
357
358#[derive(Deserialize)]
359#[serde(deny_unknown_fields)]
360struct PenalizeArguments {
361    node_ids: Vec<String>,
362}
363
364#[cfg(test)]
365mod tests;