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