1use std::io::{Read, Write}; 2use std::net::TcpListener; 3use std::thread; 4 5use whiskers_core::{Image, Model, Speaker, Turn}; 6use whiskers_gateway::Gateway; 7 8/// Serves one canned response and returns the request it received. 9fn serve(status: &str, body: &str) -> (String, thread::JoinHandle<String>) { 10 let listener = TcpListener::bind("127.0.0.1:0").unwrap(); 11 let url = format!("http://{}", listener.local_addr().unwrap()); 12 let (status, body) = (status.to_owned(), body.to_owned()); 13 let h = thread::spawn(move || { 14 let (mut s, _) = listener.accept().unwrap(); 15 // Headers and body can arrive in separate packets: read until the 16 // announced body length is complete. 17 let mut req = Vec::new(); 18 let mut buf = [0u8; 4096]; 19 loop { 20 let n = s.read(&mut buf).unwrap(); 21 req.extend_from_slice(&buf[..n]); 22 let text = String::from_utf8_lossy(&req); 23 if let Some((head, body)) = text.split_once("\r\n\r\n") { 24 let want = head 25 .lines() 26 .find_map(|l| l.to_ascii_lowercase().strip_prefix("content-length:").map(|v| v.trim().parse::<usize>().unwrap())) 27 .unwrap_or(0); 28 if body.len() >= want { 29 break; 30 } 31 } 32 assert!(n > 0, "connection closed early"); 33 } 34 let req = String::from_utf8_lossy(&req).into_owned(); 35 let resp = format!( 36 "HTTP/1.1 {status}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}", 37 body.len() 38 ); 39 s.write_all(resp.as_bytes()).unwrap(); 40 req 41 }); 42 (url, h) 43} 44 45fn turns() -> Vec<Turn> { 46 vec![ 47 Turn::said(Speaker::Whiskers, "stale"), 48 Turn::said(Speaker::Child, "hi"), 49 ] 50} 51 52#[test] 53fn keeps_only_text_blocks_and_starts_on_a_user_turn() { 54 let body = r#"{"content":[{"type":"thinking","thinking":"","signature":"x"},{"type":"text","text":"Meow!"}]}"#; 55 let (url, h) = serve("200 OK", body); 56 let out = Gateway::new(url, "m", 100).complete("sys", &turns()).unwrap(); 57 assert_eq!(out, "Meow!"); 58 let req = h.join().unwrap(); 59 assert!(req.contains("POST /v1/messages")); 60 assert!(!req.contains("stale"), "a leading assistant turn must be dropped"); 61 assert!(req.contains(r#""system":"sys""#)); 62} 63 64#[test] 65fn a_gateway_error_is_an_error() { 66 let (url, _h) = serve("500 Internal Server Error", r#"{"error":"boom"}"#); 67 assert!(Gateway::new(url, "m", 100).complete("s", &turns()).is_err()); 68} 69 70#[test] 71fn an_answer_with_no_text_is_an_error() { 72 let (url, _h) = serve("200 OK", r#"{"content":[{"type":"thinking","thinking":""}]}"#); 73 assert!(Gateway::new(url, "m", 100).complete("s", &turns()).is_err()); 74} 75 76#[test] 77fn pictures_go_as_image_blocks_before_the_words() { 78 let (url, h) = serve("200 OK", r#"{"content":[{"type":"text","text":"A bunny!"}]}"#); 79 let turns = vec![Turn { 80 speaker: Speaker::Child, 81 text: "what is this".into(), 82 pictures: vec![Image { media_type: "image/jpeg".into(), bytes: vec![1, 2, 3] }], 83 }]; 84 assert_eq!(Gateway::new(url, "m", 100).complete("s", &turns).unwrap(), "A bunny!"); 85 let req = h.join().unwrap(); 86 assert!(req.contains(r#""type":"image""#) && req.contains(r#""media_type":"image/jpeg""#)); 87 assert!(req.contains("AQID"), "base64 of the bytes"); 88 assert!(req.find(r#""type":"image""#) < req.find(r#""type":"text""#)); 89}