jev.rsannotatedjev.rssource115 lines · 4.8 KB · raw

/check, /embed and /rerank: Jev's decisions and the embedding model.

3use ::log::{debug, info, warn};
4use whiskers_core::Verdict;
5use whiskers_core::wire::{CheckReply, CheckRequest, EmbedReply, EmbedRequest, RankReply, RankRequest};
6use whiskers_ports::{Clock, EmbedBatch, Embedder, HasClock, HasEmbedder, HasJudge, Judge};
8use super::{Route, sealed};
9use crate::Response;

POST /check: Jev's decision on a message. Needs the judge and the clock.

12pub struct Check;

POST /embed: texts to vectors. Needs the embedder and the clock.

14pub struct Embed;

POST /rerank: which of the shortlisted memories fit the message. Needs the judge and the clock.

16pub struct Rerank;
18impl sealed::Sealed for Check {}
19impl sealed::Sealed for Embed {}
20impl sealed::Sealed for Rerank {}
21
22impl<C: HasJudge + HasClock> Route<C> for Check {
23    const PATHS: &'static [&'static str] = &["/check"];
24    async fn answer(&self, adapter: &C, _: &str, body: &str) -> Response {
25        let reply = match serde_json::from_str::<CheckRequest>(body) {
26            Ok(req) => judge_message(adapter, req).await,
27            Err(e) => {
28                warn!("/check: bad request: {}", super::kind_of(&e));
29                CheckReply::Unavailable(format!("bad request: {e}"))
30            }
31        };
32        Response::of(&reply)
33    }
34}
35
36async fn judge_message<C: HasJudge + HasClock>(adapter: &C, req: CheckRequest) -> CheckReply {
37    let age = req.age();
38    debug!("/check {:?} for age {}: {} chars (age sent = {})", req.direction, age.years(), req.text.len(), req.age.is_some());
39    let started = adapter.clock().now();
40    let answer = adapter.judge().check(req.direction, age, &req.text).await;
41    let ms = adapter.clock().now().get().saturating_sub(started.get());
42    match answer {
43        Ok(verdict) => {
44            info!("/check verdict {} in {ms} ms", if matches!(verdict, Verdict::Allow) { "allow" } else { "refuse" });
45            CheckReply::Verdict(verdict)
46        }
47        Err(e) => {
48            warn!("/check unavailable after {ms} ms: {e:?}");
49            CheckReply::Unavailable(e.to_string())
50        }
51    }
52}
53
54impl<C: HasEmbedder + HasClock> Route<C> for Embed {
55    const PATHS: &'static [&'static str] = &["/embed"];
56    async fn answer(&self, adapter: &C, _: &str, body: &str) -> Response {
57        let reply = match serde_json::from_str::<EmbedRequest>(body) {
58            Ok(req) => {
59                debug!("/embed: {} text(s)", req.texts.len());
60                embed_texts(adapter, req.texts).await
61            }
62            Err(e) => {
63                warn!("/embed: bad request: {}", super::kind_of(&e));
64                EmbedReply::Unavailable(format!("bad request: {e}"))
65            }
66        };
67        Response::of(&reply)
68    }
69}
70
71async fn embed_texts<C: HasEmbedder + HasClock>(adapter: &C, texts: Vec<String>) -> EmbedReply {
72    let batch = match EmbedBatch::new(texts) {
73        Ok(batch) => batch,
74        Err(e) => return EmbedReply::Unavailable(format!("bad request: {e}")),
75    };
76    let started = adapter.clock().now();
77    match adapter.embedder().embed(&batch).await {
78        Ok(vectors) => {
79            debug!("/embed: {} vector(s) in {} ms", vectors.len(), adapter.clock().now().get().saturating_sub(started.get()));
80            EmbedReply::Vectors(vectors)
81        }
82        Err(e) => {
83            warn!("/embed: the embedding model failed after {} ms: {e}", adapter.clock().now().get().saturating_sub(started.get()));
84            EmbedReply::Unavailable(e.to_string())
85        }
86    }
87}

Which of the shortlisted memories fit the message, as Jev sees it (one question, however many).

90impl<C: HasJudge + HasClock> Route<C> for Rerank {
91    const PATHS: &'static [&'static str] = &["/rerank"];
92    async fn answer(&self, adapter: &C, _: &str, body: &str) -> Response {
93        let reply = match serde_json::from_str::<RankRequest>(body) {
94            Ok(req) => {
95                debug!("/rerank: {} candidate(s), query of {} chars", req.candidates.len(), req.query.len());
96                let started = adapter.clock().now();
97                match adapter.judge().rerank(&req.query, &req.candidates).await {
98                    Ok(p) => {
99                        debug!("/rerank: {} probabilities in {} ms", p.len(), adapter.clock().now().get().saturating_sub(started.get()));
100                        RankReply::Probabilities(p)
101                    }
102                    Err(e) => {
103                        warn!("/rerank unavailable after {} ms: {e:?}", adapter.clock().now().get().saturating_sub(started.get()));
104                        RankReply::Unavailable(e.to_string())
105                    }
106                }
107            }
108            Err(e) => {
109                warn!("/rerank: bad request: {}", super::kind_of(&e));
110                RankReply::Unavailable(format!("bad request: {e}"))
111            }
112        };
113        Response::of(&reply)
114    }
115}