whiskers.git / crates / whiskersd / src / model.rs
model.rsannotatedmodel.rssource121 lines · 5.0 KB · raw
1//! The model, through the operator's gateway. The tablet talks to one service and the gateway itself
2//! (which may have no login of its own) is never exposed on the private network the tablet uses. What
3//! may pass, and what it costs, is the service's policy; this forwards the vetted body and brings back
4//! what the gateway said.
5
6use std::time::Duration;
7
8use log::{info, trace, warn};
9use whiskers_ports::{Diagnostic, ModelName, Reply, ThinkError, ThinkRequest, Thinker};
10
11const DEFAULT_MODEL: &str = "claude-haiku-4-5-20251001";
12
13pub struct GatewayThinker {
14    upstream: String,
15    model: ModelName,
16    agent: ureq::Agent,
17}
18
19impl GatewayThinker {
20    /// `WHISKERS_MODEL_URL` is the gateway's address and has no default: there is no address that is
21    /// right for everyone, and a wrong one would send the child's words somewhere unintended.
22    /// `WHISKERS_MODEL` names the one model allowed.
23    pub fn from_env() -> Result<Self, String> {
24        let get = |k: &str| std::env::var(k).ok().map(|v| v.trim().to_owned()).filter(|v| !v.is_empty());
25        let url = get("WHISKERS_MODEL_URL").ok_or_else(|| {
26            "WHISKERS_MODEL_URL is not set: it is the address of the model gateway (an Anthropic Messages API), for example http://127.0.0.1:3456".to_owned()
27        })?;
28        let model = get("WHISKERS_MODEL").unwrap_or_else(|| DEFAULT_MODEL.to_owned());
29        let p = Self::new(url, &model);
30        info!("model proxy: model {}", p.model);
31        Ok(p)
32    }
33
34    pub fn new(upstream: String, model: &str) -> Self {
35        let agent = ureq::Agent::config_builder().timeout_global(Some(Duration::from_secs(60))).http_status_as_error(false).build().into();
36        Self {
37            upstream: upstream.trim_end_matches('/').to_owned(),
38            model: ModelName::new(model).expect("a model name that is not empty"),
39            agent,
40        }
41    }
42}
43
44impl Thinker for GatewayThinker {
45    fn serves(&self) -> &ModelName {
46        &self.model
47    }
48
49    async fn think(&self, request: &ThinkRequest) -> Result<Reply, ThinkError> {
50        trace!("model proxy: request of {} bytes", request.body().len());
51        let started = std::time::Instant::now();
52        let mut resp = self
53            .agent
54            .post(format!("{}/v1/messages", self.upstream))
55            .header("content-type", "application/json")
56            .header("anthropic-version", "2023-06-01")
57            .header("x-api-key", "gateway")
58            .send(request.body())
59            .map_err(|e| {
60                warn!("model gateway unreachable after {} ms: {e}", started.elapsed().as_millis());
61                ThinkError::Unreachable(Diagnostic::new(format!("gateway: {e}")))
62            })?;
63        let status = resp.status().as_u16();
64        let body = resp.body_mut().read_to_string().map_err(|e| {
65            warn!("model gateway reply unreadable (status {status}) after {} ms: {e}", started.elapsed().as_millis());
66            ThinkError::Unreadable(Diagnostic::new(format!("gateway: {e}")))
67        })?;
68        info!("model gateway answered {status} in {} ms, {} bytes", started.elapsed().as_millis(), body.len());
69        Ok(Reply { status, body })
70    }
71}
72
73#[cfg(test)]
74mod tests {
75    use std::io::{Read, Write};
76    use std::net::TcpListener;
77    use std::thread;
78
79    use super::*;
80    use whiskers_ports::run_ready;
81
82    fn upstream(status: &str, body: &str) -> String {
83        let l = TcpListener::bind("127.0.0.1:0").unwrap();
84        let url = format!("http://{}", l.local_addr().unwrap());
85        let (status, body) = (status.to_owned(), body.to_owned());
86        thread::spawn(move || {
87            let (mut s, _) = l.accept().unwrap();
88            let mut buf = [0u8; 65536];
89            let _ = s.read(&mut buf);
90            let r = format!("HTTP/1.1 {status}\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}", body.len());
91            s.write_all(r.as_bytes()).unwrap();
92        });
93        url
94    }
95
96    const GOOD: &str = r#"{"model":"m","max_tokens":300,"system":"s","messages":[{"role":"user","content":"hi"}]}"#;
97
98    fn vetted(t: &GatewayThinker) -> ThinkRequest {
99        ThinkRequest::vet(GOOD, t.serves()).unwrap()
100    }
101
102    #[test]
103    fn a_whiskers_request_is_forwarded_and_the_answer_comes_back_whole() {
104        let t = GatewayThinker::new(upstream("200 OK", r#"{"content":[]}"#), "m");
105        assert_eq!(run_ready(t.think(&vetted(&t))).unwrap(), Reply { status: 200, body: r#"{"content":[]}"#.to_owned() });
106    }
107
108    #[test]
109    fn the_gateways_own_errors_pass_through_with_their_status() {
110        let t = GatewayThinker::new(upstream("500 Internal Server Error", "boom"), "m");
111        assert_eq!(run_ready(t.think(&vetted(&t))).unwrap(), Reply { status: 500, body: "boom".to_owned() });
112    }
113
114    #[test]
115    fn an_unreachable_gateway_is_an_upstream_error_that_says_so() {
116        let t = GatewayThinker::new("http://127.0.0.1:1".into(), "m");
117        let e = run_ready(t.think(&vetted(&t))).unwrap_err();
118        assert!(matches!(e, ThinkError::Unreachable(_)));
119        assert!(e.to_string().starts_with("gateway: "), "{e}");
120    }
121}