Skip to main content

kcode_k1_chat_testkit/
lib.rs

1use std::{
2    collections::VecDeque,
3    future::Future,
4    sync::{
5        Arc, Mutex,
6        atomic::{AtomicUsize, Ordering},
7    },
8    time::Duration,
9};
10
11use kcode_k1_chat_core::{
12    Call, ChatError, ChatEvent, ChatView, CompactFuture, CompactRequest, Inference, Llm, LlmError,
13    LlmFuture, LlmThread, PendingAction, Runtime, ToolMode, ToolOutput, ToolRequest, ToolStart,
14    Updates, WorkerRequest, WorkerStart,
15};
16use tokio::sync::{Notify, mpsc};
17
18pub trait Candidate {
19    type Chat: Send + Sync + 'static;
20    fn open(
21        runtime: Arc<dyn Runtime>,
22        initial_primary: String,
23        llm: Arc<dyn Llm>,
24    ) -> (Self::Chat, mpsc::UnboundedReceiver<ChatEvent>);
25    fn append(
26        chat: &Self::Chat,
27        text: String,
28    ) -> impl Future<Output = Result<(), ChatError>> + Send;
29    fn restart(chat: &Self::Chat) -> impl Future<Output = Result<(), ChatError>> + Send;
30    fn view(chat: &Self::Chat) -> impl Future<Output = Result<ChatView, ChatError>> + Send;
31    fn finalize(
32        chat: Self::Chat,
33    ) -> impl Future<Output = Result<ChatView, ChatError>> + Send + 'static;
34}
35
36struct Reply {
37    gate: Option<Arc<Notify>>,
38    result: Result<Inference, LlmError>,
39}
40
41struct ScriptedInner {
42    starts: AtomicUsize,
43    replies: Mutex<VecDeque<Reply>>,
44    deltas: Mutex<Vec<(usize, String)>>,
45}
46
47#[derive(Clone)]
48struct Scripted(Arc<ScriptedInner>);
49
50impl Scripted {
51    fn new(replies: Vec<Reply>) -> Arc<Self> {
52        Arc::new(Self(Arc::new(ScriptedInner {
53            starts: AtomicUsize::new(0),
54            replies: Mutex::new(replies.into()),
55            deltas: Mutex::new(Vec::new()),
56        })))
57    }
58
59    fn deltas(&self) -> Vec<(usize, String)> {
60        self.0.deltas.lock().unwrap().clone()
61    }
62
63    fn starts(&self) -> usize {
64        self.0.starts.load(Ordering::SeqCst)
65    }
66}
67
68struct ScriptedThread {
69    owner: Scripted,
70    id: usize,
71}
72
73impl Llm for Scripted {
74    fn start(&self) -> Box<dyn LlmThread> {
75        let id = self.0.starts.fetch_add(1, Ordering::SeqCst) + 1;
76        Box::new(ScriptedThread {
77            owner: self.clone(),
78            id,
79        })
80    }
81}
82
83impl LlmThread for ScriptedThread {
84    fn infer<'a>(&'a mut self, delta: &'a str) -> LlmFuture<'a> {
85        self.owner
86            .0
87            .deltas
88            .lock()
89            .unwrap()
90            .push((self.id, delta.to_owned()));
91        let reply = self
92            .owner
93            .0
94            .replies
95            .lock()
96            .unwrap()
97            .pop_front()
98            .expect("scripted LLM reply exhausted");
99        Box::pin(async move {
100            if let Some(gate) = reply.gate {
101                gate.notified().await;
102            }
103            reply.result
104        })
105    }
106}
107
108struct ToolPlan {
109    mode: ToolMode,
110    queued: String,
111    result: String,
112    gate: Option<Arc<Notify>>,
113    activity: Option<String>,
114}
115
116#[derive(Default)]
117struct FixtureRuntime {
118    tools: Mutex<VecDeque<(String, ToolPlan)>>,
119}
120
121impl FixtureRuntime {
122    fn install(&self, name: &str, plan: ToolPlan) {
123        self.tools
124            .lock()
125            .unwrap()
126            .push_back((name.to_owned(), plan));
127    }
128}
129
130impl Runtime for FixtureRuntime {
131    fn start_tool(&self, request: ToolRequest, updates: Updates) -> Result<ToolStart, String> {
132        let mut tools = self.tools.lock().unwrap();
133        let index = tools
134            .iter()
135            .position(|(name, _)| name == &request.name)
136            .ok_or_else(|| format!("missing tool plan: {}", request.name))?;
137        let (_, plan) = tools.remove(index).expect("tool plan index disappeared");
138        Ok(ToolStart {
139            mode: plan.mode,
140            queued: plan.queued,
141            future: Box::pin(async move {
142                if let Some(activity) = plan.activity {
143                    let _ = updates.activity(activity).send();
144                }
145                if let Some(gate) = plan.gate {
146                    gate.notified().await;
147                }
148                ToolOutput {
149                    text: plan.result,
150                    cost_cents: Default::default(),
151                }
152            }),
153        })
154    }
155
156    fn start_worker(&self, _request: &WorkerRequest) -> Result<WorkerStart, String> {
157        Err("workers are unavailable in this verifier".to_owned())
158    }
159
160    fn compact(&self, _request: CompactRequest, _primary: String) -> CompactFuture {
161        Box::pin(async { Err("compaction is unavailable in this verifier".to_owned()) })
162    }
163}
164
165fn success(text: &str, calls: Vec<Call>) -> Reply {
166    Reply {
167        gate: None,
168        result: Ok(Inference {
169            text: text.to_owned(),
170            calls,
171            continue_inference: false,
172        }),
173    }
174}
175
176fn gated(gate: Arc<Notify>, text: &str, calls: Vec<Call>) -> Reply {
177    Reply {
178        gate: Some(gate),
179        result: success(text, calls).result,
180    }
181}
182
183fn transient(text: &str) -> Reply {
184    Reply {
185        gate: None,
186        result: Err(LlmError::Transient(text.to_owned())),
187    }
188}
189
190fn permanent(text: &str) -> Reply {
191    Reply {
192        gate: None,
193        result: Err(LlmError::Permanent(text.to_owned())),
194    }
195}
196
197fn tool(name: &str) -> Call {
198    Call::Tool(ToolRequest {
199        name: name.to_owned(),
200        input: String::new(),
201    })
202}
203
204fn plan(
205    mode: ToolMode,
206    queued: &str,
207    result: &str,
208    gate: Option<Arc<Notify>>,
209    activity: Option<&str>,
210) -> ToolPlan {
211    ToolPlan {
212        mode,
213        queued: queued.to_owned(),
214        result: result.to_owned(),
215        gate,
216        activity: activity.map(str::to_owned),
217    }
218}
219
220async fn settle() {
221    for _ in 0..30 {
222        tokio::task::yield_now().await;
223    }
224}
225
226pub fn verify_initial_primary_delta_output_order_and_no_self_trigger<C: Candidate>() {
227    let runtime = tokio::runtime::Builder::new_current_thread()
228        .enable_time()
229        .start_paused(true)
230        .build()
231        .expect("failed to build verifier runtime");
232    runtime.block_on(async {
233        let gate = Arc::new(Notify::new());
234        let llm = Scripted::new(vec![
235            gated(gate.clone(), "o", Vec::new()),
236            success("", Vec::new()),
237            success("", Vec::new()),
238        ]);
239        let (chat, mut events) = C::open(
240            Arc::new(FixtureRuntime::default()),
241            "i".to_owned(),
242            llm.clone(),
243        );
244
245        settle().await;
246        assert!(llm.deltas().is_empty());
247        assert_eq!(C::view(&chat).await.unwrap().primary, "i");
248        assert_eq!(C::append(&chat, String::new()).await, Err(ChatError::Empty));
249
250        C::append(&chat, "u".to_owned()).await.unwrap();
251        settle().await;
252        assert_eq!(llm.deltas(), vec![(1, "iu".to_owned())]);
253
254        C::append(&chat, "a".to_owned()).await.unwrap();
255        assert_eq!(C::view(&chat).await.unwrap().pending, "a");
256        gate.notify_one();
257        settle().await;
258
259        assert_eq!(llm.deltas()[1], (1, "a".to_owned()));
260        assert_eq!(C::view(&chat).await.unwrap().primary, "iuoa");
261        assert_eq!(events.try_recv(), Ok(ChatEvent::Text("o".to_owned())));
262        settle().await;
263        assert_eq!(llm.deltas().len(), 2);
264
265        C::append(&chat, "v".to_owned()).await.unwrap();
266        settle().await;
267        assert_eq!(llm.deltas()[2], (1, "v".to_owned()));
268        assert_eq!(C::view(&chat).await.unwrap().primary, "iuoav");
269        assert_eq!(C::finalize(chat).await.unwrap().primary, "iuoav");
270        assert_eq!(events.recv().await, None);
271    });
272}
273
274pub fn verify_retry_schedule_stall_pending_and_fresh_restart<C: Candidate>() {
275    let runtime = tokio::runtime::Builder::new_current_thread()
276        .enable_time()
277        .start_paused(true)
278        .build()
279        .expect("failed to build verifier runtime");
280    runtime.block_on(async {
281        let mut replies = (1..=5)
282            .map(|number| transient(&number.to_string()))
283            .collect::<Vec<_>>();
284        replies.push(success("", Vec::new()));
285        let llm = Scripted::new(replies);
286        let (chat, mut events) = C::open(
287            Arc::new(FixtureRuntime::default()),
288            "i".to_owned(),
289            llm.clone(),
290        );
291
292        C::append(&chat, "u".to_owned()).await.unwrap();
293        settle().await;
294        assert_eq!(
295            C::view(&chat).await.unwrap().actions,
296            vec![PendingAction::Inference { attempt: 1 }]
297        );
298
299        for (wait, count) in [(10, 2), (20, 3), (40, 4), (80, 5)] {
300            tokio::time::advance(Duration::from_secs(wait - 1)).await;
301            settle().await;
302            assert_eq!(llm.deltas().len(), count - 1);
303            assert_eq!(
304                C::view(&chat).await.unwrap().actions,
305                vec![PendingAction::Inference {
306                    attempt: (count - 1) as u8
307                }]
308            );
309            tokio::time::advance(Duration::from_secs(1)).await;
310            settle().await;
311            assert_eq!(llm.deltas().len(), count);
312            if count < 5 {
313                assert_eq!(
314                    C::view(&chat).await.unwrap().actions,
315                    vec![PendingAction::Inference {
316                        attempt: count as u8
317                    }]
318                );
319            }
320        }
321
322        assert_eq!(llm.deltas(), vec![(1, "iu".to_owned()); 5]);
323        assert_eq!(events.try_recv(), Ok(ChatEvent::Stalled("5".to_owned())));
324        C::append(&chat, "later".to_owned()).await.unwrap();
325        assert_eq!(C::view(&chat).await.unwrap().pending, "later");
326        C::restart(&chat).await.unwrap();
327        settle().await;
328        assert_eq!(llm.starts(), 2);
329        assert_eq!(llm.deltas().last(), Some(&(2, "iulater".to_owned())));
330        assert_eq!(C::finalize(chat).await.unwrap().primary, "iulater");
331        assert_eq!(events.recv().await, None);
332    });
333}
334
335pub fn verify_blocked_threads_are_independent<C: Candidate>() {
336    let runtime = tokio::runtime::Builder::new_current_thread()
337        .enable_time()
338        .start_paused(true)
339        .build()
340        .expect("failed to build verifier runtime");
341    runtime.block_on(async {
342        let gate = Arc::new(Notify::new());
343        let blocked = Scripted::new(vec![gated(gate.clone(), "", Vec::new())]);
344        let free = Scripted::new(vec![success("x", Vec::new())]);
345        let fixtures: Arc<dyn Runtime> = Arc::new(FixtureRuntime::default());
346        let (first, mut first_events) = C::open(fixtures.clone(), String::new(), blocked);
347        let (second, mut second_events) = C::open(fixtures, String::new(), free);
348
349        C::append(&first, "a".to_owned()).await.unwrap();
350        C::append(&second, "b".to_owned()).await.unwrap();
351        settle().await;
352        assert_eq!(C::view(&second).await.unwrap().primary, "bx");
353        assert_eq!(
354            C::view(&first).await.unwrap().actions,
355            vec![PendingAction::Inference { attempt: 1 }]
356        );
357
358        gate.notify_one();
359        settle().await;
360        assert_eq!(C::finalize(first).await.unwrap().primary, "a");
361        assert_eq!(C::finalize(second).await.unwrap().primary, "bx");
362        assert_eq!(first_events.recv().await, None);
363        assert_eq!(
364            second_events.recv().await,
365            Some(ChatEvent::Text("x".to_owned()))
366        );
367        assert_eq!(second_events.recv().await, None);
368    });
369}
370
371pub fn verify_finalize_waits_closes_and_stalled_finalize_returns<C: Candidate>() {
372    let runtime = tokio::runtime::Builder::new_current_thread()
373        .enable_time()
374        .start_paused(true)
375        .build()
376        .expect("failed to build verifier runtime");
377    runtime.block_on(async {
378        let gate = Arc::new(Notify::new());
379        let fixtures = Arc::new(FixtureRuntime::default());
380        fixtures.install(
381            "q",
382            plan(ToolMode::Queued, "q", "r", Some(gate.clone()), None),
383        );
384        let llm = Scripted::new(vec![
385            success("", vec![tool("q")]),
386            success("", Vec::new()),
387            success("", Vec::new()),
388        ]);
389        let (chat, mut events) = C::open(fixtures, String::new(), llm);
390
391        C::append(&chat, "u".to_owned()).await.unwrap();
392        settle().await;
393        assert_eq!(C::view(&chat).await.unwrap().primary, "uq");
394        let task = tokio::spawn(C::finalize(chat));
395        settle().await;
396        assert!(!task.is_finished());
397
398        gate.notify_one();
399        settle().await;
400        assert_eq!(task.await.unwrap().unwrap().primary, "uqr");
401        assert_eq!(events.recv().await, None);
402
403        let llm = Scripted::new(vec![permanent("stop")]);
404        let (stalled, mut stalled_events) =
405            C::open(Arc::new(FixtureRuntime::default()), String::new(), llm);
406        C::append(&stalled, "u".to_owned()).await.unwrap();
407        settle().await;
408        let view = tokio::time::timeout(Duration::from_secs(1), C::finalize(stalled))
409            .await
410            .expect("stalled finalization timed out")
411            .expect("stalled finalization failed");
412        assert_eq!(view.primary, "u");
413        assert_eq!(
414            stalled_events.try_recv(),
415            Ok(ChatEvent::Stalled("stop".to_owned()))
416        );
417        assert_eq!(stalled_events.recv().await, None);
418    });
419}
420
421pub fn verify_dropped_text_and_activity_receivers_stall_cleanly<C: Candidate>() {
422    let runtime = tokio::runtime::Builder::new_current_thread()
423        .enable_time()
424        .start_paused(true)
425        .build()
426        .expect("failed to build verifier runtime");
427    runtime.block_on(async {
428        let llm = Scripted::new(vec![success("text", Vec::new()), success("", Vec::new())]);
429        let (chat, events) = C::open(
430            Arc::new(FixtureRuntime::default()),
431            "i".to_owned(),
432            llm.clone(),
433        );
434        drop(events);
435
436        C::append(&chat, "u".to_owned()).await.unwrap();
437        settle().await;
438        C::restart(&chat).await.unwrap();
439        settle().await;
440        assert_eq!(llm.starts(), 2);
441        assert_eq!(llm.deltas()[1], (2, "iutext".to_owned()));
442        assert_eq!(C::finalize(chat).await.unwrap().primary, "iutext");
443
444        let fixtures = Arc::new(FixtureRuntime::default());
445        fixtures.install(
446            "status",
447            plan(ToolMode::Fast, "unused", "R", None, Some("status")),
448        );
449        let llm = Scripted::new(vec![
450            success("", vec![tool("status")]),
451            success("", Vec::new()),
452        ]);
453        let (chat, events) = C::open(fixtures, String::new(), llm.clone());
454        drop(events);
455
456        C::append(&chat, "u".to_owned()).await.unwrap();
457        settle().await;
458        let view = C::view(&chat).await.unwrap();
459        assert!(view.actions.is_empty());
460        assert_eq!(view.pending, "R");
461        C::restart(&chat).await.unwrap();
462        settle().await;
463        assert_eq!(llm.starts(), 2);
464        assert_eq!(llm.deltas()[1], (2, "uR".to_owned()));
465        assert_eq!(C::finalize(chat).await.unwrap().primary, "uR");
466    });
467}