kcode_k1_chat_thread_actions/
lib.rs1use 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;