stub.rsannotatedstub.rssource89 lines · 3.4 KB · raw
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}