返回 CodeWhale
mock.rs
根目录 / crates / tui / src / llm_client / mock.rs
1 //! `MockLlmClient` — a queue-driven `LlmClient` implementation for tests.
2 //!
3 //! This client implements the [`LlmClient`](super::LlmClient) trait by replaying a
4 //! pre-loaded queue of canned responses (one per turn). It captures every
5 //! request the runtime sends so tests can assert on the outgoing payload —
6 //! e.g. confirming that prior `reasoning_content` is replayed in DeepSeek V4
7 //! thinking-mode tool-calling turns (V4 §5.1.1; the bug that broke
8 //! v0.4.9-v0.5.1).
9 //!
10 //! # Mocking strategy
11 //!
12 //! Tests mock at the **trait boundary** (`LlmClient`), never at the `reqwest`
13 //! HTTP layer. The trait is the durable abstraction — internal HTTP plumbing
14 //! changes frequently and is not part of the public engine contract.
15 //!
16 //! # Example
17 //!
18 //! ```ignore
19 //! use crate::llm_client::mock::{MockLlmClient, canned};
20 //! use crate::llm_client::LlmClient;
21 //!
22 //! // One canned turn that emits "hello world" as two text deltas, then
23 //! // finishes with stop_reason = "end_turn".
24 //! let turn = vec![
25 //! canned::message_start("msg_1"),
26 //! canned::text_delta(0, "hello "),
27 //! canned::text_delta(0, "world"),
28 //! canned::message_stop(),
29 //! ];
30 //!
31 //! let mock = MockLlmClient::new(vec![turn]);
32 //! let stream = mock.create_message_stream(/* ... */).await.unwrap();
33 //! // ... drain the stream, assert deltas ...
34 //! assert_eq!(mock.call_count(), 1);
35 //! assert_eq!(mock.captured_requests().len(), 1);
36 //! ```
37
38 // This module ships methods + builder helpers that integration tests rely on
39 // individually. Not every helper is exercised by unit tests — that's expected
40 // (the goal is a usable mock surface for downstream tests), so we silence
41 // per-item dead-code warnings at the module level.
42 #![allow(dead_code)]
43
44 use std::collections::VecDeque;
45 use std::pin::Pin;
46 use std::sync::Mutex;
47 use std::sync::atomic::{AtomicUsize, Ordering};
48
49 use anyhow::{Result, anyhow};
50 use async_stream::try_stream;
51 use futures_util::Stream;
52
53 use codewhale_models::{
54 ContentBlock, MessageDelta, MessageRequest, MessageResponse, StreamEvent, Usage,
55 };
56
57 use super::{LlmClient, StreamEventBox};
58 use codewhale_models::Role;
59
60 /// A pre-recorded "turn" the mock will replay on the next streaming call.
61 ///
62 /// `MessageStop` does *not* need to be the final element — the mock will
63 /// auto-emit one if missing, mirroring the real client's behaviour. Likewise
64 /// the mock does not require `MessageStart` to be present.
65 pub type CannedTurn = Vec<StreamEvent>;
66
67 /// A queued mock response step.
68 pub enum FauxStep {
69 Canned(CannedTurn),
70 /// Build a canned turn from the live outgoing request.
71 ///
72 /// Tests can assert DeepSeek V4's thinking-mode tool-call invariant here:
73 /// on the assistant turn that produced the previous tool call, the next
74 /// outgoing request must still carry `reasoning_content` (represented in
75 /// this model as a [`ContentBlock::Thinking`] block). If it is missing,
76 /// DeepSeek V4 returns HTTP 400 on the follow-up turn. This guards the
77 /// [v0.4.9-v0.5.1 regression range](https://github.com/codewhale-hq/CodeWhale/compare/v0.4.9...v0.5.1)
78 /// where that content was dropped.
79 Factory(Box<dyn Fn(&MessageRequest) -> CannedTurn + Send + Sync>),
80 /// Fail the request itself with this message, as a provider that refuses
81 /// it (a 401 for a bad key, say) does before any stream starts.
82 Error(String),
83 }
84
85 /// A queue-driven mock LLM client.
86 ///
87 /// The mock holds a FIFO queue of canned response turns. Each call to
88 /// [`LlmClient::create_message_stream`] dequeues the next turn and replays its
89 /// events as a stream. If the queue is exhausted, the call returns an error
90 /// — tests should ensure they push exactly as many turns as the runtime will
91 /// consume.
92 ///
93 /// The mock also captures the [`MessageRequest`] passed to every call so tests
94 /// can assert on the outgoing payload (e.g. that prior `reasoning_content` is
95 /// preserved across turns).
96 pub struct MockLlmClient {
97 canned: Mutex<VecDeque<FauxStep>>,
98 captured_requests: Mutex<Vec<MessageRequest>>,
99 calls: AtomicUsize,
100 provider_name: &'static str,
101 model: String,
102 /// If set, [`LlmClient::create_message`] returns this verbatim. Otherwise
103 /// it falls back to streaming + collection. Useful for non-streaming
104 /// compaction-style calls.
105 canned_messages: Mutex<VecDeque<MessageResponse>>,
106 }
107
108 impl MockLlmClient {
109 /// Construct a mock that will replay the given canned turns in order.
110 #[must_use]
111 pub fn new(canned: Vec<CannedTurn>) -> Self {
112 Self {
113 canned: Mutex::new(canned.into_iter().map(FauxStep::Canned).collect()),
114 captured_requests: Mutex::new(Vec::new()),
115 calls: AtomicUsize::new(0),
116 provider_name: "mock",
117 model: "mock-model".to_string(),
118 canned_messages: Mutex::new(VecDeque::new()),
119 }
120 }
121
122 /// Set the provider-name string returned by [`LlmClient::provider_name`].
123 #[must_use]
124 pub fn with_provider(mut self, name: &'static str) -> Self {
125 self.provider_name = name;
126 self
127 }
128
129 /// Set the model identifier returned by [`LlmClient::model`].
130 #[must_use]
131 pub fn with_model(mut self, model: impl Into<String>) -> Self {
132 self.model = model.into();
133 self
134 }
135
136 /// Push a canned turn onto the back of the queue.
137 pub fn push_turn(&self, turn: CannedTurn) {
138 self.canned
139 .lock()
140 .expect("MockLlmClient.canned mutex poisoned")
141 .push_back(FauxStep::Canned(turn));
142 }
143
144 /// Push a factory step onto the back of the queue.
145 ///
146 /// The closure receives the live outgoing [`MessageRequest`] before the
147 /// response stream is built, so assertions panic directly from the client
148 /// call rather than later while polling the returned stream.
149 pub fn push_factory<F>(&self, factory: F)
150 where
151 F: Fn(&MessageRequest) -> CannedTurn + Send + Sync + 'static,
152 {
153 self.canned
154 .lock()
155 .expect("MockLlmClient.canned mutex poisoned")
156 .push_back(FauxStep::Factory(Box::new(factory)));
157 }
158
159 /// Push a request failure onto the back of the queue.
160 pub fn push_error(&self, message: impl Into<String>) {
161 self.canned
162 .lock()
163 .expect("MockLlmClient.canned mutex poisoned")
164 .push_back(FauxStep::Error(message.into()));
165 }
166
167 /// Push a canned non-streaming `MessageResponse`. Consumed by
168 /// [`LlmClient::create_message`] (FIFO).
169 pub fn push_message_response(&self, response: MessageResponse) {
170 self.canned_messages
171 .lock()
172 .expect("MockLlmClient.canned_messages mutex poisoned")
173 .push_back(response);
174 }
175
176 /// Number of completed calls to either `create_message` or
177 /// `create_message_stream`.
178 #[must_use]
179 pub fn call_count(&self) -> usize {
180 self.calls.load(Ordering::SeqCst)
181 }
182
183 /// Number of canned turns still queued.
184 #[must_use]
185 pub fn remaining_turns(&self) -> usize {
186 self.canned
187 .lock()
188 .expect("MockLlmClient.canned mutex poisoned")
189 .len()
190 }
191
192 /// Snapshot of every request the mock has been asked to handle, in order.
193 #[must_use]
194 pub fn captured_requests(&self) -> Vec<MessageRequest> {
195 self.captured_requests
196 .lock()
197 .expect("MockLlmClient.captured_requests mutex poisoned")
198 .clone()
199 }
200
201 /// Convenience: return the most recently captured request, or `None` if
202 /// the mock has not been called yet.
203 #[must_use]
204 pub fn last_request(&self) -> Option<MessageRequest> {
205 self.captured_requests
206 .lock()
207 .expect("MockLlmClient.captured_requests mutex poisoned")
208 .last()
209 .cloned()
210 }
211
212 fn record_request(&self, request: &MessageRequest) {
213 self.captured_requests
214 .lock()
215 .expect("MockLlmClient.captured_requests mutex poisoned")
216 .push(request.clone());
217 self.calls.fetch_add(1, Ordering::SeqCst);
218 }
219
220 fn pop_step(&self) -> Option<FauxStep> {
221 self.canned
222 .lock()
223 .expect("MockLlmClient.canned mutex poisoned")
224 .pop_front()
225 }
226
227 fn turn_from_step(&self, step: FauxStep, request: &MessageRequest) -> Result<CannedTurn> {
228 match step {
229 FauxStep::Canned(turn) => Ok(turn),
230 FauxStep::Factory(factory) => Ok(factory(request)),
231 FauxStep::Error(message) => Err(anyhow!(message)),
232 }
233 }
234
235 fn pop_message(&self) -> Option<MessageResponse> {
236 self.canned_messages
237 .lock()
238 .expect("MockLlmClient.canned_messages mutex poisoned")
239 .pop_front()
240 }
241 }
242
243 impl LlmClient for MockLlmClient {
244 fn provider_name(&self) -> &'static str {
245 self.provider_name
246 }
247
248 fn model(&self) -> &str {
249 &self.model
250 }
251
252 async fn create_message(&self, request: MessageRequest) -> Result<MessageResponse> {
253 self.record_request(&request);
254
255 if let Some(canned) = self.pop_message() {
256 return Ok(canned);
257 }
258
259 // Fallback: synthesize a MessageResponse from the next streaming turn.
260 let Some(step) = self.pop_step() else {
261 return Err(anyhow!(
262 "MockLlmClient: create_message called but no canned response queued (request #{})",
263 self.calls.load(Ordering::SeqCst)
264 ));
265 };
266
267 let turn = self.turn_from_step(step, &request)?;
268 Ok(synthesize_message_response(turn, &self.model))
269 }
270
271 async fn create_message_stream(&self, request: MessageRequest) -> Result<StreamEventBox> {
272 self.record_request(&request);
273
274 let Some(step) = self.pop_step() else {
275 return Err(anyhow!(
276 "MockLlmClient: create_message_stream called but no canned turn queued (call #{})",
277 self.calls.load(Ordering::SeqCst)
278 ));
279 };
280
281 let turn = self.turn_from_step(step, &request)?;
282 Ok(stream_from_canned(turn))
283 }
284
285 async fn health_check(&self) -> Result<bool> {
286 Ok(true)
287 }
288 }
289
290 /// Wrap a canned event vector as a stream that yields each event in order and
291 /// auto-appends `MessageStop` if the trailing event is not already one.
292 fn stream_from_canned(turn: CannedTurn) -> StreamEventBox {
293 let s = try_stream! {
294 let has_stop = matches!(turn.last(), Some(StreamEvent::MessageStop));
295 for ev in turn {
296 yield ev;
297 }
298 if !has_stop {
299 yield StreamEvent::MessageStop;
300 }
301 };
302 Box::pin(s) as Pin<Box<dyn Stream<Item = Result<StreamEvent>> + Send + 'static>>
303 }
304
305 /// Best-effort: collapse a streaming turn into a non-streaming
306 /// `MessageResponse` by concatenating text deltas. Used only as a fallback
307 /// when callers `create_message` without a queued `MessageResponse`.
308 fn synthesize_message_response(turn: CannedTurn, model: &str) -> MessageResponse {
309 use codewhale_models::Delta;
310
311 let mut text = String::new();
312 let mut stop_reason: Option<String> = None;
313
314 for ev in turn {
315 match ev {
316 StreamEvent::ContentBlockDelta {
317 delta: Delta::TextDelta { text: t },
318 ..
319 } => text.push_str(&t),
320 StreamEvent::MessageDelta {
321 delta: MessageDelta {
322 stop_reason: sr, ..
323 },
324 ..
325 } => stop_reason = sr,
326 _ => {}
327 }
328 }
329
330 MessageResponse {
331 id: "mock_msg".to_string(),
332 r#type: "message".to_string(),
333 role: "assistant".to_string(),
334 content: vec![ContentBlock::Text {
335 text,
336 cache_control: None,
337 }],
338 model: model.to_string(),
339 stop_reason: stop_reason.or_else(|| Some("end_turn".to_string())),
340 stop_sequence: None,
341 container: None,
342 usage: Usage::default(),
343 }
344 }
345
346 /// Builders for common canned-event patterns. Re-exported so tests can build
347 /// realistic streams without wiring `StreamEvent` shapes by hand.
348 pub mod canned {
349 use serde_json::Value;
350
351 use codewhale_models::{
352 ContentBlockStart, Delta, MessageDelta, MessageResponse, StreamEvent, Usage,
353 };
354
355 /// `MessageStart` event with a synthetic message envelope.
356 #[must_use]
357 pub fn message_start(id: &str) -> StreamEvent {
358 StreamEvent::MessageStart {
359 message: MessageResponse {
360 id: id.to_string(),
361 r#type: "message".to_string(),
362 role: "assistant".to_string(),
363 content: vec![],
364 model: "mock-model".to_string(),
365 stop_reason: None,
366 stop_sequence: None,
367 container: None,
368 usage: Usage::default(),
369 },
370 }
371 }
372
373 /// Open a text content block at `index`.
374 #[must_use]
375 pub fn text_block_start(index: u32) -> StreamEvent {
376 StreamEvent::ContentBlockStart {
377 index,
378 content_block: ContentBlockStart::Text {
379 text: String::new(),
380 },
381 }
382 }
383
384 /// Append `text` to the content block at `index`.
385 #[must_use]
386 pub fn text_delta(index: u32, text: &str) -> StreamEvent {
387 StreamEvent::ContentBlockDelta {
388 index,
389 delta: Delta::TextDelta {
390 text: text.to_string(),
391 },
392 }
393 }
394
395 /// Append a thinking-content delta at `index`.
396 #[must_use]
397 pub fn thinking_delta(index: u32, thinking: &str) -> StreamEvent {
398 StreamEvent::ContentBlockDelta {
399 index,
400 delta: Delta::ThinkingDelta {
401 thinking: thinking.to_string(),
402 },
403 }
404 }
405
406 /// Open a tool_use content block at `index`.
407 #[must_use]
408 pub fn tool_use_block_start(index: u32, id: &str, name: &str) -> StreamEvent {
409 StreamEvent::ContentBlockStart {
410 index,
411 content_block: ContentBlockStart::ToolUse {
412 id: id.to_string(),
413 name: name.to_string(),
414 input: Value::Null,
415 caller: None,
416 thought_signature: None,
417 },
418 }
419 }
420
421 /// Stream partial JSON for a tool's input arguments.
422 #[must_use]
423 pub fn tool_input_delta(index: u32, partial_json: &str) -> StreamEvent {
424 StreamEvent::ContentBlockDelta {
425 index,
426 delta: Delta::InputJsonDelta {
427 partial_json: partial_json.to_string(),
428 },
429 }
430 }
431
432 /// Close the content block at `index`.
433 #[must_use]
434 pub fn block_stop(index: u32) -> StreamEvent {
435 StreamEvent::ContentBlockStop { index }
436 }
437
438 /// Emit a `message_delta` carrying `stop_reason` and optional `usage`.
439 #[must_use]
440 pub fn message_delta(stop_reason: &str, usage: Option<Usage>) -> StreamEvent {
441 StreamEvent::MessageDelta {
442 delta: MessageDelta {
443 stop_reason: Some(stop_reason.to_string()),
444 stop_sequence: None,
445 },
446 usage,
447 }
448 }
449
450 /// Final `message_stop` sentinel.
451 #[must_use]
452 pub fn message_stop() -> StreamEvent {
453 StreamEvent::MessageStop
454 }
455
456 /// Convenience: a complete "assistant emits this text" turn ending with
457 /// `stop_reason = "end_turn"`.
458 #[must_use]
459 pub fn simple_text_turn(text: &str) -> Vec<StreamEvent> {
460 vec![
461 message_start("mock_msg_1"),
462 text_block_start(0),
463 text_delta(0, text),
464 block_stop(0),
465 message_delta("end_turn", None),
466 message_stop(),
467 ]
468 }
469
470 /// Convenience: a turn that emits one assistant tool_call and stops.
471 #[must_use]
472 pub fn tool_call_turn(call_id: &str, tool_name: &str, args_json: &str) -> Vec<StreamEvent> {
473 vec![
474 message_start("mock_msg_tool"),
475 tool_use_block_start(0, call_id, tool_name),
476 tool_input_delta(0, args_json),
477 block_stop(0),
478 message_delta("tool_use", None),
479 message_stop(),
480 ]
481 }
482 }
483
484 // === Tests ===
485
486 #[cfg(test)]
487 mod tests {
488 use futures_util::StreamExt;
489
490 use super::*;
491 use crate::llm_client::LlmClient;
492 use codewhale_models::{Delta, Message, MessageRequest, StreamEvent};
493
494 fn empty_request() -> MessageRequest {
495 MessageRequest {
496 model: "mock-model".to_string(),
497 messages: vec![Message {
498 role: Role::User,
499 content: vec![],
500 }],
501 max_tokens: 1024,
502 system: None,
503 tools: None,
504 tool_choice: None,
505 metadata: None,
506 thinking: None,
507 reasoning_effort: None,
508 stream: Some(true),
509 temperature: None,
510 top_p: None,
511 }
512 }
513
514 #[tokio::test]
515 async fn replays_canned_turn_via_stream() {
516 let mock = MockLlmClient::new(vec![canned::simple_text_turn("hello world")]);
517
518 let mut stream = mock
519 .create_message_stream(empty_request())
520 .await
521 .expect("stream should open");
522
523 let mut text = String::new();
524 let mut saw_stop = false;
525 while let Some(ev) = stream.next().await {
526 match ev.expect("event") {
527 StreamEvent::ContentBlockDelta {
528 delta: Delta::TextDelta { text: t },
529 ..
530 } => text.push_str(&t),
531 StreamEvent::MessageStop => {
532 saw_stop = true;
533 break;
534 }
535 _ => {}
536 }
537 }
538
539 assert_eq!(text, "hello world");
540 assert!(saw_stop);
541 assert_eq!(mock.call_count(), 1);
542 assert_eq!(mock.captured_requests().len(), 1);
543 assert_eq!(mock.remaining_turns(), 0);
544 }
545
546 #[tokio::test]
547 async fn errors_when_queue_exhausted() {
548 let mock = MockLlmClient::new(Vec::new());
549 let result = mock.create_message_stream(empty_request()).await;
550 match result {
551 Ok(_) => panic!("should error on empty queue"),
552 Err(err) => assert!(format!("{err}").contains("no canned")),
553 }
554 }
555
556 #[tokio::test]
557 async fn captures_request_payload_for_assertions() {
558 let mock = MockLlmClient::new(vec![canned::simple_text_turn("ok")]);
559 let mut req = empty_request();
560 req.temperature = Some(0.42);
561 let _ = mock.create_message_stream(req).await.unwrap();
562
563 let captured = mock.last_request().expect("should have captured");
564 assert_eq!(captured.temperature, Some(0.42));
565 }
566
567 #[tokio::test]
568 async fn stream_auto_appends_message_stop() {
569 // Queue a turn missing MessageStop — mock should append one.
570 let turn = vec![canned::text_block_start(0), canned::text_delta(0, "x")];
571 let mock = MockLlmClient::new(vec![turn]);
572
573 let mut stream = mock.create_message_stream(empty_request()).await.unwrap();
574 let mut saw_stop = false;
575 while let Some(ev) = stream.next().await {
576 if matches!(ev.expect("event"), StreamEvent::MessageStop) {
577 saw_stop = true;
578 }
579 }
580 assert!(saw_stop, "auto MessageStop missing");
581 }
582
583 #[tokio::test]
584 async fn create_message_uses_canned_message_response_first() {
585 let mock = MockLlmClient::new(vec![canned::simple_text_turn("from stream")]);
586 mock.push_message_response(MessageResponse {
587 id: "preset".to_string(),
588 r#type: "message".to_string(),
589 role: "assistant".to_string(),
590 content: vec![ContentBlock::Text {
591 text: "from preset".to_string(),
592 cache_control: None,
593 }],
594 model: "mock-model".to_string(),
595 stop_reason: Some("end_turn".to_string()),
596 stop_sequence: None,
597 container: None,
598 usage: Usage::default(),
599 });
600
601 let resp = mock.create_message(empty_request()).await.unwrap();
602 assert_eq!(resp.id, "preset");
603 }
604
605 #[tokio::test]
606 async fn create_message_synthesizes_from_streaming_turn_when_no_message_queued() {
607 let mock = MockLlmClient::new(vec![canned::simple_text_turn("synthesized")]);
608 let resp = mock.create_message(empty_request()).await.unwrap();
609 let text = match &resp.content[0] {
610 ContentBlock::Text { text, .. } => text.clone(),
611 _ => panic!("expected text"),
612 };
613 assert_eq!(text, "synthesized");
614 assert_eq!(resp.stop_reason.as_deref(), Some("end_turn"));
615 }
616
617 #[tokio::test]
618 async fn create_message_synthesizes_from_factory_turn() {
619 let mock = MockLlmClient::new(Vec::new());
620 mock.push_factory(|request| {
621 assert_eq!(request.model, "mock-model");
622 canned::simple_text_turn("from factory")
623 });
624
625 let resp = mock.create_message(empty_request()).await.unwrap();
626 let text = match &resp.content[0] {
627 ContentBlock::Text { text, .. } => text.clone(),
628 _ => panic!("expected text"),
629 };
630 assert_eq!(text, "from factory");
631 }
632
633 #[tokio::test]
634 async fn provider_and_model_are_overridable() {
635 let mock = MockLlmClient::new(vec![canned::simple_text_turn("x")])
636 .with_provider("test-provider")
637 .with_model("test-model");
638 assert_eq!(mock.provider_name(), "test-provider");
639 assert_eq!(mock.model(), "test-model");
640 }
641
642 #[tokio::test]
643 async fn tool_call_turn_serializes_correctly() {
644 let mock = MockLlmClient::new(vec![canned::tool_call_turn(
645 "call_1",
646 "list_dir",
647 r#"{"path":"/tmp"}"#,
648 )]);
649 let mut stream = mock.create_message_stream(empty_request()).await.unwrap();
650
651 let mut saw_tool_use = false;
652 let mut json_seen = String::new();
653 while let Some(ev) = stream.next().await {
654 match ev.unwrap() {
655 StreamEvent::ContentBlockStart { content_block, .. } => {
656 use codewhale_models::ContentBlockStart;
657 if let ContentBlockStart::ToolUse { name, .. } = content_block {
658 assert_eq!(name, "list_dir");
659 saw_tool_use = true;
660 }
661 }
662 StreamEvent::ContentBlockDelta {
663 delta: Delta::InputJsonDelta { partial_json },
664 ..
665 } => json_seen.push_str(&partial_json),
666 _ => {}
667 }
668 }
669 assert!(saw_tool_use, "expected tool_use start event");
670 assert!(json_seen.contains("/tmp"));
671 }
672
673 #[tokio::test]
674 async fn multiple_turns_consumed_in_order() {
675 let mock = MockLlmClient::new(vec![
676 canned::simple_text_turn("turn-one"),
677 canned::simple_text_turn("turn-two"),
678 ]);
679 for expected in ["turn-one", "turn-two"] {
680 let mut stream = mock.create_message_stream(empty_request()).await.unwrap();
681 let mut text = String::new();
682 while let Some(ev) = stream.next().await {
683 if let StreamEvent::ContentBlockDelta {
684 delta: Delta::TextDelta { text: t },
685 ..
686 } = ev.unwrap()
687 {
688 text.push_str(&t);
689 }
690 }
691 assert_eq!(text, expected);
692 }
693 assert_eq!(mock.call_count(), 2);
694 assert_eq!(mock.remaining_turns(), 0);
695 }
696 }
697
697 lines RUST