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