whiskers.git / crates / whiskersd / src / model.rs
model.rsannotatedmodel.rssource121 lines · 5.0 KB · raw

The model, through the operator's gateway. The tablet talks to one service and the gateway itself (which may have no login of its own) is never exposed on the private network the tablet uses. What may pass, and what it costs, is the service's policy; this forwards the vetted body and brings back what the gateway said.

6use std::time::Duration;
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 {

WHISKERS_MODEL_URL is the gateway's address and has no default: there is no address that is right for everyone, and a wrong one would send the child's words somewhere unintended. 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    }
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}