1use std::io::{Read, Write};
2use std::net::TcpListener;
3use std::thread;
4
5use whiskers_core::{Age, Direction, Embedder, Guard, RefusalKind, Ranker, Verdict};
6use whiskers_guard::{CheckReply, EmbedReply, RankReply, RemoteEmbedder, RemoteGuard, RemoteRanker, RemoteSync};
7
8fn serve(status: &str, body: String) -> String {
9    let listener = TcpListener::bind("127.0.0.1:0").unwrap();
10    let url = format!("http://{}", listener.local_addr().unwrap());
11    let status = status.to_owned();
12    thread::spawn(move || {
13        let (mut s, _) = listener.accept().unwrap();
14        let mut buf = [0u8; 8192];
15        let _ = s.read(&mut buf);
16        let resp = format!(
17            "HTTP/1.1 {status}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
18            body.len()
19        );
20        s.write_all(resp.as_bytes()).unwrap();
21    });
22    url
23}
24
25#[test]
26fn a_verdict_comes_back_as_a_verdict() {
27    let v = Verdict::Refuse { reason: "hurt_or_unsafe".into(), kind: RefusalKind::NeedsAGrownUp };
28    let url = serve("200 OK", serde_json::to_string(&CheckReply::Verdict(v.clone())).unwrap());
29    assert_eq!(RemoteGuard::new(url).check(Direction::FromChild, Age::YOUNGEST, "x").unwrap(), v);
30}
31
32#[test]
33fn unavailable_is_an_error_never_an_allow() {
34    let url = serve("200 OK", serde_json::to_string(&CheckReply::Unavailable("no key".into())).unwrap());
35    assert!(RemoteGuard::new(url).check(Direction::ToChild, Age::YOUNGEST, "x").is_err());
36}
37
38#[test]
39fn a_server_error_or_garbage_is_an_error() {
40    let url = serve("500 Internal Server Error", "{}".into());
41    assert!(RemoteGuard::new(url).check(Direction::ToChild, Age::YOUNGEST, "x").is_err());
42    let url = serve("200 OK", "not json".into());
43    assert!(RemoteGuard::new(url).check(Direction::ToChild, Age::YOUNGEST, "x").is_err());
44}
45
46#[test]
47fn nothing_listening_is_an_error() {
48    assert!(RemoteGuard::new("http://127.0.0.1:1").check(Direction::ToChild, Age::YOUNGEST, "x").is_err());
49}
50
51#[test]
52fn vectors_come_back_one_per_text_and_a_miscount_is_an_error() {
53    let url = serve("200 OK", serde_json::to_string(&EmbedReply::Vectors(vec![vec![1.0, 2.0], vec![3.0, 4.0]])).unwrap());
54    let got = RemoteEmbedder::new(url).embed(&["a".into(), "b".into()]).unwrap();
55    assert_eq!(got, vec![vec![1.0, 2.0], vec![3.0, 4.0]]);
56    let url = serve("200 OK", serde_json::to_string(&EmbedReply::Vectors(vec![vec![1.0]])).unwrap());
57    assert!(RemoteEmbedder::new(url).embed(&["a".into(), "b".into()]).is_err());
58    let url = serve("200 OK", serde_json::to_string(&EmbedReply::Unavailable("down".into())).unwrap());
59    assert!(RemoteEmbedder::new(url).embed(&["a".into()]).is_err());
60    assert!(RemoteEmbedder::new("http://127.0.0.1:1").embed(&["a".into()]).is_err());
61}
62
63#[test]
64fn rank_probabilities_come_back_one_per_candidate() {
65    let url = serve("200 OK", serde_json::to_string(&RankReply::Probabilities(vec![0.7, 0.3])).unwrap());
66    assert_eq!(RemoteRanker::new(url).rank("q", &["a".into(), "b".into()]).unwrap(), vec![0.7, 0.3]);
67    let url = serve("200 OK", serde_json::to_string(&RankReply::Probabilities(vec![1.0])).unwrap());
68    assert!(RemoteRanker::new(url).rank("q", &["a".into(), "b".into()]).is_err());
69    let url = serve("200 OK", serde_json::to_string(&RankReply::Unavailable("throttled".into())).unwrap());
70    assert!(RemoteRanker::new(url).rank("q", &["a".into(), "b".into()]).is_err());
71}
72
73#[test]
74fn a_500_from_chat_sync_is_an_error_naming_the_status_never_a_copy_to_adopt() {
75    let mine = whiskers_core::ChatState { version: 3, ..Default::default() };
76    let url = serve("500 Internal Server Error", String::new());
77    assert_eq!(RemoteSync::new(url).chat(&mine), Err("service answered 500 Internal Server Error".to_owned()));
78    // And a 200 is still the service's copy.
79    let url = serve("200 OK", serde_json::to_string(&whiskers_core::ChatState { version: 5, ..Default::default() }).unwrap());
80    assert_eq!(RemoteSync::new(url).chat(&mine).unwrap().version, 5);
81}