1use std::collections::HashMap;
2use std::future::Future;
3use std::path::PathBuf;
4use std::sync::Mutex;
5use std::time::{Duration, Instant};
6
7use indexmap::IndexMap;
8
9use crate::answer::{ChoiceAnswer, NoulAnswer, ScoreAnswer};
10use crate::backends::fake::FakeBackend;
11use crate::error::{BackendError, DecodeError, Error, PolicyError};
12use crate::ids::QuestionId;
13use crate::policy::{Fail, Policy};
14use crate::question::{ChoiceLabels, Question, ScoreLabels};
15use crate::state::State;
16use crate::usage::UsageFn;
17use crate::verdict::{Decision, UnsureReason, UntypedDecision};
18use crate::wire::{
19 self, Usage, WireAnswer, WireQuestion, WireRequest, WireResponse, renormalize_probabilities,
20};
21
22pub const TYPESAFE_ORIGIN: &str = "https://api.typesafe.ai";
24
25#[derive(Debug)]
26pub struct Evaluated {
27 pub wire: WireResponse,
28 pub meta: IndexMap<String, AnswerMeta>,
29 pub backend_id: String,
30}
31
32pub use crate::verdict::{AnswerMeta, CascadeHop};
33
34pub trait Backend: Send + Sync {
35 fn id(&self) -> &str;
36 fn evaluate(
37 &self,
38 req: WireRequest,
39 deadline: Instant,
40 ) -> impl Future<Output = Result<Evaluated, BackendError>> + Send;
41
42 fn replace_api_key(&self, key: Option<String>) -> Result<Option<String>, Error> {
50 let _ = key;
51 Ok(None)
52 }
53}
54
55pub enum AnyBackend {
56 Fake(FakeBackend),
57 #[cfg(feature = "http")]
58 Http(crate::backends::http::HttpBackend),
59 #[cfg(feature = "http")]
60 CascadeHttp(
61 Box<
62 crate::backends::cascade::Cascaded<
63 crate::backends::http::HttpBackend,
64 crate::backends::http::HttpBackend,
65 >,
66 >,
67 ),
68}
69
70impl Backend for AnyBackend {
71 fn id(&self) -> &str {
72 match self {
73 AnyBackend::Fake(backend) => backend.id(),
74 #[cfg(feature = "http")]
75 AnyBackend::Http(backend) => backend.id(),
76 #[cfg(feature = "http")]
77 AnyBackend::CascadeHttp(backend) => backend.id(),
78 }
79 }
80
81 async fn evaluate(
82 &self,
83 req: WireRequest,
84 deadline: Instant,
85 ) -> Result<Evaluated, BackendError> {
86 match self {
87 AnyBackend::Fake(backend) => backend.evaluate(req, deadline).await,
88 #[cfg(feature = "http")]
89 AnyBackend::Http(backend) => backend.evaluate(req, deadline).await,
90 #[cfg(feature = "http")]
91 AnyBackend::CascadeHttp(backend) => backend.evaluate(req, deadline).await,
92 }
93 }
94
95 fn replace_api_key(&self, key: Option<String>) -> Result<Option<String>, Error> {
96 match self {
97 AnyBackend::Fake(backend) => backend.replace_api_key(key),
98 #[cfg(feature = "http")]
99 AnyBackend::Http(backend) => backend.replace_api_key(key),
100 #[cfg(feature = "http")]
101 AnyBackend::CascadeHttp(backend) => backend.replace_api_key(key),
102 }
103 }
104}
105
106pub(crate) struct GateCache {
107 cap: usize,
108 ttl: Duration,
109 entries: HashMap<u64, (Instant, crate::verdict::Verdict)>,
110}
111
112impl GateCache {
113 pub(crate) fn new(cap: usize) -> Self {
114 Self {
115 cap,
116 ttl: Duration::from_secs(30),
117 entries: HashMap::new(),
118 }
119 }
120}
121
122#[derive(Debug, Clone, Default)]
124pub struct ClientConfig {
125 pub backend: String,
126 pub shadow: bool,
127 pub model: Option<String>,
128 pub timeout_ms: Option<String>,
129 pub policy: Option<String>,
130 pub base_url: Option<String>,
131 pub api_key: Option<String>,
132 pub typesafe_key: Option<String>,
133 pub allow_private_http: bool,
134 pub cascade: Option<String>,
135 pub log_path: Option<PathBuf>,
136 pub cache_capacity: Option<usize>,
138}
139
140#[derive(Clone, Copy, Default)]
145pub struct CallChoice<'a> {
146 pub model: Option<&'a str>,
147 pub policy: Option<&'a Policy>,
148 pub timeout: Option<Duration>,
149}
150
151#[derive(Debug, Clone, PartialEq, Eq)]
153pub struct ClientStatus {
154 pub backend: String,
155 pub model: String,
156 pub policy: String,
157 pub log: bool,
158 pub cache: bool,
159}
160
161pub struct Client<B: Backend> {
162 pub(crate) backend: B,
163 pub(crate) policy: Option<Policy>,
164 pub(crate) on_usage: Option<UsageFn>,
165 pub(crate) timeout: Duration,
166 pub(crate) model: String,
167 shadow_override: Option<bool>,
168 fail_override: Option<Fail>,
169 pub(crate) log_path: Option<PathBuf>,
170 pub(crate) log_lock: Mutex<()>,
171 pub(crate) cache: Option<Mutex<GateCache>>,
172}
173
174impl<B: Backend> Client<B> {
175 pub fn new(backend: B) -> Self {
176 Self {
177 backend,
178 policy: None,
179 on_usage: None,
180 timeout: Duration::from_millis(2000),
181 model: "jev-latest".to_string(),
182 shadow_override: None,
183 fail_override: None,
184 log_path: None,
185 log_lock: Mutex::new(()),
186 cache: None,
187 }
188 }
189
190 pub fn fail(mut self, fail: Fail) -> Self {
192 self.fail_override = Some(fail);
193 self
194 }
195
196 pub(crate) fn fail_for(&self, policy: Option<&Policy>) -> Fail {
198 if let Some(fail) = self.fail_override {
199 return fail;
200 }
201 policy.map(|policy| policy.fail).unwrap_or(Fail::Closed)
202 }
203
204 pub fn shadow(mut self, on: bool) -> Self {
206 self.shadow_override = Some(on);
207 self
208 }
209
210 pub(crate) fn shadow_on(&self, policy_shadow: bool) -> bool {
211 self.shadow_override.unwrap_or(policy_shadow)
212 }
213
214 pub fn policy(mut self, policy: Policy) -> Self {
215 self.policy = Some(policy);
216 self
217 }
218
219 pub fn on_usage(mut self, on_usage: UsageFn) -> Self {
220 self.on_usage = Some(on_usage);
221 self
222 }
223
224 pub fn timeout(mut self, timeout: Duration) -> Self {
225 self.timeout = timeout;
226 self
227 }
228
229 pub fn model(mut self, id: impl AsRef<str>) -> Result<Self, Error> {
231 let id = id.as_ref().trim();
232 if id.is_empty() {
233 return Err(Error::Policy(PolicyError::Config(
234 "SNAPIF_MODEL must not be blank".to_string(),
235 )));
236 }
237 self.model = id.to_string();
238 Ok(self)
239 }
240
241 pub fn backend(&self) -> &B {
242 &self.backend
243 }
244
245 pub fn status(&self) -> ClientStatus {
247 ClientStatus {
248 backend: self.backend.id().to_string(),
249 model: self.model.clone(),
250 policy: self.policy.as_ref().map(policy_label).unwrap_or_default(),
251 log: self.log_path.is_some(),
252 cache: self.cache.is_some(),
253 }
254 }
255
256 pub fn replace_api_key(&self, key: Option<String>) -> Result<(), Error> {
261 self.backend.replace_api_key(key).map(|_| ())
262 }
263
264 pub fn battery_id(&self) -> Option<&str> {
266 self.policy.as_ref().map(|policy| policy.battery.0.as_str())
267 }
268
269 pub(crate) fn cache_get(&self, key: u64) -> Option<crate::verdict::Verdict> {
270 let cache = self.cache.as_ref()?;
271 let mut guard = cache.lock().ok()?;
272 let expired = guard
273 .entries
274 .get(&key)
275 .is_some_and(|(stored, _)| stored.elapsed() > guard.ttl);
276 if expired {
277 guard.entries.remove(&key);
278 return None;
279 }
280 guard.entries.get(&key).map(|(_, verdict)| verdict.clone())
281 }
282
283 pub(crate) fn cache_put(&self, key: u64, verdict: crate::verdict::Verdict) {
284 let Some(cache) = &self.cache else {
285 return;
286 };
287 let Ok(mut guard) = cache.lock() else {
288 return;
289 };
290 if guard.entries.len() >= guard.cap
291 && !guard.entries.contains_key(&key)
292 && let Some(old) = guard.entries.keys().next().copied()
293 {
294 guard.entries.remove(&old);
295 }
296 guard.entries.insert(key, (Instant::now(), verdict));
297 }
298
299 pub async fn ask(&self, state: State, questions: Vec<Question>) -> Result<AskOut, Error> {
305 self.ask_with(state, questions, CallChoice::default()).await
306 }
307
308 pub async fn ask_with(
310 &self,
311 state: State,
312 questions: Vec<Question>,
313 choice: CallChoice<'_>,
314 ) -> Result<AskOut, Error> {
315 let policy = choice
316 .policy
317 .or(self.policy.as_ref())
318 .ok_or_else(|| Error::Policy(PolicyError::Invariant("policy".to_string())))?;
319 let model = match choice.model.map(str::trim).filter(|text| !text.is_empty()) {
320 Some(model) => model.to_string(),
321 None => self.model.clone(),
322 };
323 policy.ensure_checked()?;
324 let mut request = WireRequest {
325 model: model.clone(),
326 state: state.to_wire(None),
327 questions: questions.iter().map(wire_question).collect(),
328 };
329 let encoded = wire::encode(&request)?;
330 if encoded.truncated_untrusted {
331 request = wire::decode_request(&encoded.body)?;
332 }
333 let timeout = choice.timeout.unwrap_or(self.timeout);
334 let deadline = Instant::now() + timeout;
335 let evaluated = self
336 .backend
337 .evaluate(request.clone(), deadline)
338 .await
339 .map_err(|err| map_backend(err, timeout))?;
340 if let Some(on_usage) = &self.on_usage {
341 on_usage(evaluated.wire.usage);
342 }
343 wire::check_response(&request.questions, &evaluated.wire)?;
344 let Evaluated {
345 wire,
346 mut meta,
347 backend_id,
348 } = evaluated;
349 let mut decisions = IndexMap::new();
350 let mut scores = IndexMap::new();
351 for (id, question) in &request.questions {
352 let Some(answer) = wire.answers.get(id) else {
353 continue;
354 };
355 record_prob_sum(&mut meta, id, answer);
356 scores.insert(id.clone(), answer_score(answer));
357 let key = QuestionId::new(id);
358 let decision = untyped(policy, question, answer).map_err(Error::Decode)?;
359 decisions.insert(key, decision);
360 }
361 let (pack, pack_version) = match policy.shipped_id.clone() {
362 Some(id) => {
363 let version = Policy::shipped_pack_version(&id);
364 (id, version)
365 }
366 None => (String::new(), 0),
367 };
368 Ok(AskOut {
369 decisions,
370 scores,
371 usage: wire.usage,
372 backend_id,
373 meta,
374 truncated_untrusted: encoded.truncated_untrusted,
375 pack,
376 pack_version,
377 model,
378 })
379 }
380}
381
382impl<B: Backend> std::fmt::Debug for Client<B> {
383 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
384 let status = self.status();
385 formatter
386 .debug_struct("Client")
387 .field("backend", &status.backend)
388 .field("model", &status.model)
389 .field("policy", &status.policy)
390 .field("log", &status.log)
391 .field("cache", &status.cache)
392 .finish()
393 }
394}
395
396#[derive(Debug, Default)]
397struct BackendEnv {
398 #[cfg(feature = "http")]
399 typesafe_key: Option<String>,
400 #[cfg(feature = "http")]
401 snapif_key: Option<String>,
402 #[cfg(feature = "http")]
403 base_url: Option<String>,
404 #[cfg(feature = "http")]
405 allow_private_http: bool,
406 #[cfg_attr(not(feature = "http"), allow(dead_code))]
407 cascade: Option<String>,
408 model: Option<String>,
409 timeout_ms: Option<String>,
410 policy: Option<String>,
411}
412
413impl Client<AnyBackend> {
414 pub fn from_env() -> Result<Self, Error> {
435 let cache_capacity = match std::env::var("SNAPIF_CACHE").ok().as_deref() {
436 None => None,
437 Some(raw) => {
438 let raw = raw.trim();
439 if raw.is_empty() {
440 None
441 } else {
442 Some(raw.parse::<usize>().map_err(|_| {
443 Error::Policy(PolicyError::Config(format!(
444 "SNAPIF_CACHE must be an integer, got {raw}"
445 )))
446 })?)
447 }
448 }
449 };
450 let config = ClientConfig {
451 backend: std::env::var("SNAPIF_BACKEND").unwrap_or_default(),
452 shadow: matches!(
453 std::env::var("SNAPIF_SHADOW").ok().as_deref(),
454 Some("1" | "true" | "TRUE" | "True")
455 ),
456 model: std::env::var("SNAPIF_MODEL").ok(),
457 timeout_ms: std::env::var("SNAPIF_TIMEOUT_MS").ok(),
458 policy: std::env::var("SNAPIF_POLICY").ok(),
459 #[cfg(feature = "http")]
460 base_url: std::env::var("SNAPIF_BASE_URL").ok(),
461 #[cfg(not(feature = "http"))]
462 base_url: None,
463 #[cfg(feature = "http")]
464 api_key: std::env::var("SNAPIF_API_KEY").ok(),
465 #[cfg(not(feature = "http"))]
466 api_key: None,
467 #[cfg(feature = "http")]
468 typesafe_key: std::env::var("TYPESAFE_API_KEY").ok(),
469 #[cfg(not(feature = "http"))]
470 typesafe_key: None,
471 #[cfg(feature = "http")]
472 allow_private_http: matches!(
473 std::env::var("SNAPIF_ALLOW_PRIVATE_HTTP").ok().as_deref(),
474 Some("1" | "true" | "TRUE" | "True")
475 ),
476 #[cfg(not(feature = "http"))]
477 allow_private_http: false,
478 cascade: std::env::var("SNAPIF_CASCADE_BASE_URL").ok(),
479 log_path: std::env::var("SNAPIF_LOG")
480 .ok()
481 .map(|raw| raw.trim().to_string())
482 .filter(|raw| !raw.is_empty())
483 .map(PathBuf::from),
484 cache_capacity,
485 };
486 Self::from_config(&config)
487 }
488
489 pub fn from_config(config: &ClientConfig) -> Result<Self, Error> {
491 let env = BackendEnv {
492 #[cfg(feature = "http")]
493 typesafe_key: config.typesafe_key.clone(),
494 #[cfg(feature = "http")]
495 snapif_key: config.api_key.clone(),
496 #[cfg(feature = "http")]
497 base_url: config.base_url.clone(),
498 #[cfg(feature = "http")]
499 allow_private_http: config.allow_private_http,
500 cascade: config.cascade.clone(),
501 model: config.model.clone(),
502 timeout_ms: config.timeout_ms.clone(),
503 policy: config.policy.clone(),
504 };
505 let name = match config.backend.trim() {
506 "" => None,
507 other => Some(other),
508 };
509 let mut client = Self::from_parts(name, config.shadow, &env)?;
510 client.log_path = config.log_path.clone();
511 if let Some(cap) = config.cache_capacity.filter(|cap| *cap > 0) {
512 client.cache = Some(Mutex::new(GateCache::new(cap)));
513 }
514 Ok(client)
515 }
516
517 pub fn replace_fake(&mut self, backend: FakeBackend) -> Result<(), Error> {
519 match &mut self.backend {
520 AnyBackend::Fake(slot) => {
521 *slot = backend;
522 Ok(())
523 }
524 #[cfg(feature = "http")]
525 _ => Err(Error::Policy(PolicyError::Config(
526 "replace_fake requires a fake backend".to_string(),
527 ))),
528 }
529 }
530
531 #[cfg_attr(not(feature = "http"), allow(unused_variables))]
532 fn from_parts(name: Option<&str>, shadow: bool, env: &BackendEnv) -> Result<Self, Error> {
533 let name = match name.map(str::trim) {
534 None | Some("") => None,
535 Some(text) => Some(text),
536 };
537 let policy = env_policy(env)?;
538 let client = match name {
539 None => {
540 return Err(Error::Policy(PolicyError::BackendName(
541 "SNAPIF_BACKEND must be fake, typesafe, or compatible".to_string(),
542 )));
543 }
544 Some("fake") => Self::new(AnyBackend::Fake(FakeBackend::new())).policy(policy),
545 #[cfg(feature = "http")]
546 Some(name @ ("typesafe" | "compatible")) => http_client(name, env, policy)?,
547 #[cfg(not(feature = "http"))]
548 Some(name @ ("typesafe" | "compatible")) => {
549 return Err(Error::Policy(PolicyError::BackendName(format!(
550 "SNAPIF_BACKEND {name} needs the http feature"
551 ))));
552 }
553 Some(other) => {
554 return Err(Error::Policy(PolicyError::BackendName(format!(
555 "unknown SNAPIF_BACKEND {other}; expected fake, typesafe, or compatible"
556 ))));
557 }
558 };
559 let client = apply_runtime(client, env)?;
560 Ok(if shadow { client.shadow(true) } else { client })
561 }
562}
563
564fn policy_label(policy: &Policy) -> String {
566 if let Some(id) = &policy.shipped_id {
567 return id.clone();
568 }
569 match &policy.source_path {
570 Some(path) => format!("{} {path}", policy.battery.0),
571 None => policy.battery.0.clone(),
572 }
573}
574
575fn env_policy(env: &BackendEnv) -> Result<Policy, Error> {
576 match nonempty(env.policy.as_deref()) {
577 Some(spec) => Policy::load(spec),
578 None => Ok(Policy::shipped("tool-gate")?),
579 }
580}
581
582fn apply_runtime(
583 mut client: Client<AnyBackend>,
584 env: &BackendEnv,
585) -> Result<Client<AnyBackend>, Error> {
586 if let Some(raw) = nonempty(env.timeout_ms.as_deref()) {
587 let ms: u64 = raw.parse().map_err(|_| {
588 Error::Policy(PolicyError::Config(format!(
589 "SNAPIF_TIMEOUT_MS must be an integer, got {raw}"
590 )))
591 })?;
592 client = client.timeout(Duration::from_millis(ms));
593 }
594 if let Some(raw) = env.model.as_deref() {
595 let model = raw.trim();
596 if model.is_empty() {
597 return Err(Error::Policy(PolicyError::Config(
598 "SNAPIF_MODEL must not be blank".to_string(),
599 )));
600 }
601 client.model = model.to_string();
602 }
603 Ok(client)
604}
605
606fn nonempty(value: Option<&str>) -> Option<&str> {
607 value.map(str::trim).filter(|text| !text.is_empty())
608}
609
610#[cfg(feature = "http")]
611fn http_client(name: &str, env: &BackendEnv, policy: Policy) -> Result<Client<AnyBackend>, Error> {
612 let backend = if let Some(raw) = nonempty(env.cascade.as_deref()) {
613 let first = compatible_backend(
614 raw,
615 env.snapif_key.clone(),
616 "SNAPIF_CASCADE_BASE_URL",
617 env.allow_private_http,
618 )?;
619 let fallback = match name {
620 "typesafe" => crate::backends::http::HttpBackend::typesafe(
621 env.typesafe_key.clone().unwrap_or_default(),
622 )?,
623 "compatible" => {
624 let Some(base) = env.base_url.as_deref().filter(|value| !value.is_empty()) else {
625 return Err(Error::Policy(PolicyError::Config(
626 "SNAPIF_BASE_URL is required".to_string(),
627 )));
628 };
629 compatible_backend(
630 base,
631 env.snapif_key.clone(),
632 "SNAPIF_BASE_URL",
633 env.allow_private_http,
634 )?
635 }
636 _ => return Err(Error::Policy(PolicyError::Invariant(name.to_string()))),
637 };
638 let rule = crate::backends::cascade::CascadeRule::new(policy.cascade_min);
639 AnyBackend::CascadeHttp(Box::new(crate::backends::cascade::Cascaded::new(
640 first, fallback, rule,
641 )))
642 } else {
643 match name {
644 "typesafe" => AnyBackend::Http(crate::backends::http::HttpBackend::typesafe(
645 env.typesafe_key.clone().unwrap_or_default(),
646 )?),
647 "compatible" => {
648 let Some(base) = env.base_url.as_deref().filter(|value| !value.is_empty()) else {
649 return Err(Error::Policy(PolicyError::Config(
650 "SNAPIF_BASE_URL is required".to_string(),
651 )));
652 };
653 AnyBackend::Http(compatible_backend(
654 base,
655 env.snapif_key.clone(),
656 "SNAPIF_BASE_URL",
657 env.allow_private_http,
658 )?)
659 }
660 _ => return Err(Error::Policy(PolicyError::Invariant(name.to_string()))),
661 }
662 };
663 Ok(Client::new(backend).policy(policy))
664}
665
666#[cfg(feature = "http")]
667fn compatible_backend(
668 raw: &str,
669 key: Option<String>,
670 invariant: &str,
671 allow_private_http: bool,
672) -> Result<crate::backends::http::HttpBackend, Error> {
673 let url = url::Url::parse(raw)
674 .map_err(|_| Error::Policy(PolicyError::Config(format!("{invariant} must be a URL"))))?;
675 if allow_private_http {
676 crate::backends::http::HttpBackend::compatible_private(url, key)
677 } else {
678 crate::backends::http::HttpBackend::compatible(url, key)
679 }
680}
681
682#[derive(Debug)]
683#[non_exhaustive]
684pub struct AskOut {
685 pub decisions: IndexMap<QuestionId, UntypedDecision>,
686 pub scores: IndexMap<String, f64>,
688 pub usage: Usage,
689 pub backend_id: String,
690 pub meta: IndexMap<String, AnswerMeta>,
691 pub truncated_untrusted: bool,
693 pub pack: String,
695 pub pack_version: u32,
696 pub model: String,
698}
699
700impl AskOut {
701 pub fn choice<T: ChoiceLabels>(&self, id: &QuestionId) -> Result<Decision<T>, DecodeError> {
702 match self.decisions.get(id) {
703 None => Err(DecodeError::MissingAnswer { key: id.clone() }),
704 Some(UntypedDecision::Choice(Decision::Known(label))) => match T::from_label(label) {
705 Some(value) => Ok(Decision::Known(value)),
706 None => Err(DecodeError::UnknownLabel {
707 key: id.clone(),
708 label: label.clone(),
709 }),
710 },
711 Some(UntypedDecision::Choice(Decision::Unsure { reason, guess })) => {
712 Ok(Decision::Unsure {
713 reason: reason.clone(),
714 guess: guess.as_deref().and_then(T::from_label),
715 })
716 }
717 Some(_) => Err(DecodeError::TypeMismatch { key: id.clone() }),
718 }
719 }
720
721 pub fn score<T: ScoreLabels>(&self, id: &QuestionId) -> Result<Decision<T>, DecodeError> {
722 match self.decisions.get(id) {
723 None => Err(DecodeError::MissingAnswer { key: id.clone() }),
724 Some(UntypedDecision::Score(Decision::Known(score))) => {
725 score_label::<T>(*score, id.clone())
726 }
727 Some(UntypedDecision::Score(Decision::Unsure { reason, guess })) => {
728 Ok(Decision::Unsure {
729 reason: reason.clone(),
730 guess: guess.and_then(|score| score_index(score).and_then(T::from_index)),
731 })
732 }
733 Some(_) => Err(DecodeError::TypeMismatch { key: id.clone() }),
734 }
735 }
736
737 pub fn noul(&self, id: &QuestionId) -> Result<Decision<bool>, DecodeError> {
738 match self.decisions.get(id) {
739 None => Err(DecodeError::MissingAnswer { key: id.clone() }),
740 Some(UntypedDecision::Noul(decision)) => Ok(decision.clone()),
741 Some(_) => Err(DecodeError::TypeMismatch { key: id.clone() }),
742 }
743 }
744}
745
746fn score_label<T: ScoreLabels>(score: f64, id: QuestionId) -> Result<Decision<T>, DecodeError> {
747 match score_index(score).and_then(T::from_index) {
748 Some(value) => Ok(Decision::Known(value)),
749 None => Err(DecodeError::OutOfRange { key: id }),
750 }
751}
752
753fn score_index(score: f64) -> Option<usize> {
754 if score.is_finite() && score >= 0.0 {
755 Some(score.round() as usize)
756 } else {
757 None
758 }
759}
760
761fn map_backend(err: BackendError, timeout: Duration) -> Error {
762 match err {
763 BackendError::Timeout => Error::Timeout(timeout),
764 BackendError::RateLimit => Error::RateLimit,
765 BackendError::Overloaded => Error::Overloaded,
766 BackendError::Auth => Error::Auth("authentication failed (HTTP 401)".to_string()),
767 BackendError::Rejected { status, body } => Error::Rejected { status, body },
768 other => Error::Backend(other.to_string()),
769 }
770}
771
772pub(crate) fn wire_question(question: &Question) -> (String, WireQuestion) {
773 match question {
774 Question::Choice(choice) => (
775 choice.id.to_string(),
776 WireQuestion::Choice {
777 instructions: choice.instructions.clone(),
778 criteria: choice.criteria.clone(),
779 },
780 ),
781 Question::Score(score) => (
782 score.id.to_string(),
783 WireQuestion::Score {
784 instructions: score.instructions.clone(),
785 criteria: score.criteria.clone(),
786 },
787 ),
788 Question::Noul(noul) => (
789 noul.id.to_string(),
790 WireQuestion::Noul {
791 instructions: noul.instructions.clone(),
792 criteria: noul.criteria.clone(),
793 },
794 ),
795 }
796}
797
798fn answer_score(answer: &WireAnswer) -> f64 {
799 match answer {
800 WireAnswer::Choice { confidence, .. } => *confidence,
801 WireAnswer::Score { score, .. } => *score,
802 WireAnswer::Noul { noul } => *noul,
803 }
804}
805
806pub(crate) fn record_prob_sum(
807 meta: &mut IndexMap<String, AnswerMeta>,
808 id: &str,
809 answer: &WireAnswer,
810) {
811 let probabilities = match answer {
812 WireAnswer::Choice { probabilities, .. } | WireAnswer::Score { probabilities, .. } => {
813 probabilities
814 }
815 WireAnswer::Noul { .. } => return,
816 };
817 let (_, original_sum) = renormalize_probabilities(probabilities);
818 if (original_sum - 1.0).abs() > 1e-6 {
819 meta.entry(id.to_string()).or_default().original_prob_sum = Some(original_sum);
820 }
821}
822
823fn untyped(
824 policy: &Policy,
825 question: &WireQuestion,
826 answer: &WireAnswer,
827) -> Result<UntypedDecision, DecodeError> {
828 match (question, answer) {
829 (
830 WireQuestion::Choice { .. },
831 WireAnswer::Choice {
832 choice,
833 probabilities,
834 confidence,
835 },
836 ) => {
837 let decoded = ChoiceAnswer {
838 label: choice.clone(),
839 confidence: *confidence,
840 probabilities: probabilities.clone(),
841 };
842 let signal = decoded.signal(policy.choice.signal);
843 let floor = policy.choice.escalate_below;
844 let decision = if signal < floor {
845 Decision::Unsure {
846 reason: UnsureReason::BelowFloor {
847 confidence: signal,
848 floor,
849 },
850 guess: Some(decoded.label),
851 }
852 } else {
853 Decision::Known(decoded.label)
854 };
855 Ok(UntypedDecision::Choice(decision))
856 }
857 (
858 WireQuestion::Score { .. },
859 WireAnswer::Score {
860 score,
861 probabilities,
862 confidence,
863 ..
864 },
865 ) => {
866 let decoded = ScoreAnswer {
867 score: *score,
868 confidence: *confidence,
869 probabilities: probabilities.clone(),
870 };
871 let signal = decoded.signal(policy.choice.signal);
872 let floor = policy.choice.escalate_below;
873 let decision = if signal < floor {
874 Decision::Unsure {
875 reason: UnsureReason::BelowFloor {
876 confidence: signal,
877 floor,
878 },
879 guess: Some(decoded.score),
880 }
881 } else {
882 Decision::Known(decoded.score)
883 };
884 Ok(UntypedDecision::Score(decision))
885 }
886 (WireQuestion::Noul { .. }, WireAnswer::Noul { noul }) => {
887 let decision = NoulAnswer { p: *noul }.decide(&policy.noul, None);
888 Ok(UntypedDecision::Noul(decision))
889 }
890 _ => Err(DecodeError::TypeMismatch {
891 key: QuestionId::new("answer"),
892 }),
893 }
894}
895
896#[cfg(test)]
897mod tests {
898 use std::sync::{Arc, Condvar, Mutex};
899 use std::time::{Duration, Instant};
900
901 use super::{AnyBackend, Backend, Client, ClientConfig};
902 use crate::error::{BackendError, Error, PolicyError};
903 use crate::wire::WireRequest;
904
905 #[test]
906 fn from_name_selects_without_env() {
907 let env = super::BackendEnv::default();
908 let config = ClientConfig {
909 backend: "fake".to_string(),
910 timeout_ms: Some("nope".to_string()),
911 ..ClientConfig::default()
912 };
913 let bad = match Client::<AnyBackend>::from_config(&config) {
914 Err(err) => err,
915 Ok(_) => panic!("timeout"),
916 };
917 assert!(bad.to_string().contains("SNAPIF_TIMEOUT_MS"), "{bad}");
918 let ok = Client::<AnyBackend>::from_config(&ClientConfig {
919 backend: "fake".to_string(),
920 ..ClientConfig::default()
921 })
922 .expect("fake");
923 assert_eq!(ok.backend().id(), "fake");
924 let Err(unset) = Client::<AnyBackend>::from_parts(None, false, &env) else {
925 panic!("unset must be policy");
926 };
927 assert!(matches!(
928 unset,
929 Error::Policy(PolicyError::BackendName(ref message))
930 if message.contains("SNAPIF_BACKEND")
931 && message.contains("fake")
932 && message.contains("typesafe")
933 && message.contains("compatible")
934 ));
935 let spaced =
936 Client::<AnyBackend>::from_parts(Some(" fake "), false, &env).expect("trimmed fake");
937 assert_eq!(spaced.backend().id(), "fake");
938 let shown = unset.to_string();
939 assert!(!shown.contains("threshold invariant"), "{shown}");
940 let Err(unknown_name) = Client::<AnyBackend>::from_parts(Some("laya"), false, &env) else {
941 panic!("unknown backend");
942 };
943 let unknown_text = unknown_name.to_string();
944 assert!(
945 unknown_text.contains("unknown SNAPIF_BACKEND laya"),
946 "{unknown_text}"
947 );
948 assert!(unknown_text.contains("fake"), "{unknown_text}");
949 assert!(
950 !unknown_text.contains("threshold invariant"),
951 "{unknown_text}"
952 );
953
954 let client = Client::<AnyBackend>::from_parts(Some("fake"), false, &env).expect("fake");
955 assert_eq!(client.backend().id(), "fake");
956 let shadowed = Client::<AnyBackend>::from_parts(Some("fake"), true, &env).expect("shadow");
957 assert_eq!(shadowed.shadow_override, Some(true));
958
959 #[cfg(not(feature = "http"))]
960 {
961 let Err(unknown) = Client::<AnyBackend>::from_parts(Some("typesafe"), false, &env)
962 else {
963 panic!("typesafe must be policy");
964 };
965 assert!(matches!(
966 unknown,
967 Error::Policy(PolicyError::BackendName(ref message))
968 if message.contains("typesafe") && message.contains("http feature")
969 ));
970 assert!(!unknown.to_string().contains("threshold invariant"));
971 let cascade_env = super::BackendEnv {
972 cascade: Some("http://127.0.0.1:9".to_string()),
973 ..super::BackendEnv::default()
974 };
975 let err = Client::<AnyBackend>::from_parts(Some("typesafe"), false, &cascade_env);
976 assert!(matches!(err, Err(Error::Policy(_))));
977 }
978 #[cfg(feature = "http")]
979 {
980 let Err(missing_key) = Client::<AnyBackend>::from_parts(Some("typesafe"), false, &env)
981 else {
982 panic!("typesafe without a key must be auth");
983 };
984 assert!(matches!(missing_key, Error::Auth(_)));
985 let selected = Client::<AnyBackend>::from_parts(
986 Some("typesafe"),
987 false,
988 &super::BackendEnv {
989 typesafe_key: Some("secret".to_string()),
990 ..super::BackendEnv::default()
991 },
992 )
993 .expect("typesafe");
994 assert_eq!(selected.backend().id(), "typesafe");
995 let cascaded = Client::<AnyBackend>::from_parts(
996 Some("typesafe"),
997 false,
998 &super::BackendEnv {
999 typesafe_key: Some("secret".to_string()),
1000 cascade: Some("http://127.0.0.1:9".to_string()),
1001 ..super::BackendEnv::default()
1002 },
1003 )
1004 .expect("cascade");
1005 assert_eq!(cascaded.backend().id(), "cascade");
1006 }
1007
1008 let tuned = Client::<AnyBackend>::from_parts(
1009 Some("fake"),
1010 false,
1011 &super::BackendEnv {
1012 model: Some("custom-model".to_string()),
1013 timeout_ms: Some("1500".to_string()),
1014 ..super::BackendEnv::default()
1015 },
1016 )
1017 .expect("runtime");
1018 assert_eq!(tuned.model, "custom-model");
1019 assert_eq!(tuned.timeout, std::time::Duration::from_millis(1500));
1020
1021 let bad_timeout = Client::<AnyBackend>::from_parts(
1022 Some("fake"),
1023 false,
1024 &super::BackendEnv {
1025 timeout_ms: Some("nope".to_string()),
1026 ..super::BackendEnv::default()
1027 },
1028 );
1029 assert!(matches!(bad_timeout, Err(Error::Policy(_))));
1030 let bad_policy = Client::<AnyBackend>::from_parts(
1031 Some("fake"),
1032 false,
1033 &super::BackendEnv {
1034 policy: Some("missing-policy".to_string()),
1035 ..super::BackendEnv::default()
1036 },
1037 );
1038 assert!(matches!(bad_policy, Err(Error::Policy(_))));
1039 }
1040
1041 #[test]
1042 fn status_reads_the_resolved_model_and_hides_the_key() {
1043 let client = Client::<AnyBackend>::from_config(&ClientConfig {
1044 backend: "fake".to_string(),
1045 model: Some("jev-1.12".to_string()),
1046 policy: Some("tool-gate".to_string()),
1047 api_key: Some("super-secret-key".to_string()),
1048 log_path: Some(std::path::PathBuf::from("/tmp/snapif.log")),
1049 cache_capacity: Some(2),
1050 ..ClientConfig::default()
1051 })
1052 .expect("config");
1053 let status = client.status();
1054 assert_eq!(status.backend, "fake");
1055 assert_eq!(status.model, "jev-1.12");
1056 assert_eq!(status.policy, "tool-gate");
1057 assert!(status.log);
1058 assert!(status.cache);
1059 let shown = format!("{client:?} {status:?}");
1060 assert!(!shown.contains("super-secret-key"), "{shown}");
1061 let plain = Client::<AnyBackend>::from_config(&ClientConfig {
1062 backend: "fake".to_string(),
1063 ..ClientConfig::default()
1064 })
1065 .expect("default");
1066 assert_eq!(plain.status().model, "jev-latest");
1067 assert_eq!(plain.status().policy, "tool-gate");
1068 let raw = include_str!("../policies/tool-gate.toml");
1069 let file_policy = Client::new(crate::backends::fake::FakeBackend::new())
1070 .policy(crate::policy::Policy::from_toml_str(raw).expect("toml"));
1071 assert_eq!(file_policy.status().policy, "tool-gate");
1072 }
1073
1074 #[test]
1075 fn status_names_a_file_policy_without_the_key() {
1076 let dir = std::env::temp_dir().join(format!("snapif-status-{}", std::process::id()));
1077 let _ = std::fs::create_dir_all(&dir);
1078 let path = dir.join("desk.toml");
1079 std::fs::write(
1080 &path,
1081 r#"
1082schema_version = 1
1083battery = "desk-pack"
1084[choice]
1085escalate_below = 0.8
1086review_below = 1.0
1087[default_action]
1088review = 0.8
1089when_unsure = "review_guess"
1090class = "read"
1091"#,
1092 )
1093 .expect("toml");
1094 let client = Client::<AnyBackend>::from_config(&ClientConfig {
1095 backend: "fake".to_string(),
1096 policy: Some(path.display().to_string()),
1097 api_key: Some("super-secret-key".to_string()),
1098 ..ClientConfig::default()
1099 })
1100 .expect("config");
1101 let status = client.status();
1102 assert!(status.policy.contains("desk-pack"), "{}", status.policy);
1103 assert!(
1104 status.policy.contains(&path.display().to_string()),
1105 "{}",
1106 status.policy
1107 );
1108 let shown = format!("{client:?} {status:?}");
1109 assert!(!shown.contains("super-secret-key"), "{shown}");
1110 assert!(!shown.contains("when_unsure"), "{shown}");
1111 let _ = std::fs::remove_dir_all(&dir);
1112 }
1113
1114 #[cfg(feature = "http")]
1115 #[test]
1116 fn replace_fake_on_typesafe_does_not_blame_the_env() {
1117 let mut client = Client::<AnyBackend>::from_config(&ClientConfig {
1118 backend: "typesafe".to_string(),
1119 typesafe_key: Some("sekrit-key-xyz".to_string()),
1120 ..ClientConfig::default()
1121 })
1122 .expect("typesafe");
1123 let err = client
1124 .replace_fake(crate::backends::fake::FakeBackend::new())
1125 .expect_err("not fake");
1126 let text = err.to_string();
1127 assert!(
1128 text.contains("replace_fake requires a fake backend"),
1129 "{text}"
1130 );
1131 assert!(!text.contains("SNAPIF_BACKEND"), "{text}");
1132 assert!(!text.contains("script"), "{text}");
1133 assert!(!text.contains("sekrit-key-xyz"), "{text}");
1134 }
1135
1136 #[test]
1137 fn model_env_rejects_blank_and_keeps_the_default() {
1138 let default =
1139 Client::<AnyBackend>::from_parts(Some("fake"), false, &super::BackendEnv::default())
1140 .expect("fake");
1141 assert_eq!(default.model, "jev-latest");
1142
1143 for blank in ["", " ", "\t"] {
1144 let err = Client::<AnyBackend>::from_parts(
1145 Some("fake"),
1146 false,
1147 &super::BackendEnv {
1148 model: Some(blank.to_string()),
1149 ..super::BackendEnv::default()
1150 },
1151 );
1152 let Err(err) = err else {
1153 panic!("blank model must fail");
1154 };
1155 assert_eq!(
1156 err.to_string(),
1157 "policy: SNAPIF_MODEL must not be blank",
1158 "{blank:?} {err}"
1159 );
1160 }
1161
1162 let trimmed = Client::<AnyBackend>::from_parts(
1163 Some("fake"),
1164 false,
1165 &super::BackendEnv {
1166 model: Some(" local-model ".to_string()),
1167 ..super::BackendEnv::default()
1168 },
1169 )
1170 .expect("trimmed");
1171 assert_eq!(trimmed.model, "local-model");
1172 }
1173
1174 #[cfg(feature = "http")]
1175 #[test]
1176 fn private_http_flag_is_off_unless_set() {
1177 let off = Client::<AnyBackend>::from_parts(
1178 Some("compatible"),
1179 false,
1180 &super::BackendEnv {
1181 base_url: Some("http://10.0.0.1".to_string()),
1182 snapif_key: Some("snapif-key".to_string()),
1183 ..super::BackendEnv::default()
1184 },
1185 );
1186 assert!(matches!(off, Err(Error::Policy(_))));
1187
1188 let on = Client::<AnyBackend>::from_parts(
1189 Some("compatible"),
1190 false,
1191 &super::BackendEnv {
1192 base_url: Some("http://10.0.0.1".to_string()),
1193 snapif_key: Some("snapif-key".to_string()),
1194 allow_private_http: true,
1195 ..super::BackendEnv::default()
1196 },
1197 );
1198 let Ok(on) = on else {
1199 panic!("private http");
1200 };
1201 assert_eq!(on.backend().id(), "compatible");
1202 }
1203
1204 #[cfg(feature = "http")]
1205 #[test]
1206 fn private_http_flag_covers_the_cascade_origin() {
1207 let denied = Client::<AnyBackend>::from_parts(
1208 Some("compatible"),
1209 false,
1210 &super::BackendEnv {
1211 base_url: Some("http://10.0.0.1".into()),
1212 cascade: Some("http://192.168.1.50".into()),
1213 snapif_key: Some("snapif-key".into()),
1214 allow_private_http: false,
1215 ..super::BackendEnv::default()
1216 },
1217 );
1218 assert!(matches!(denied, Err(Error::Policy(_))));
1219
1220 let allowed = Client::<AnyBackend>::from_parts(
1221 Some("compatible"),
1222 false,
1223 &super::BackendEnv {
1224 base_url: Some("http://10.0.0.1".into()),
1225 cascade: Some("http://192.168.1.50".into()),
1226 snapif_key: Some("snapif-key".into()),
1227 allow_private_http: true,
1228 ..super::BackendEnv::default()
1229 },
1230 );
1231 let Ok(allowed) = allowed else {
1232 panic!("private http cascade");
1233 };
1234 assert_eq!(allowed.backend().id(), "cascade");
1235 }
1236
1237 #[test]
1238 fn auth_error_reaches_the_client_text() {
1239 let err = super::map_backend(
1240 crate::error::BackendError::Auth,
1241 std::time::Duration::from_secs(1),
1242 );
1243 assert_eq!(err.to_string(), "auth: authentication failed (HTTP 401)");
1244 }
1245
1246 struct HoldKey {
1247 key: Mutex<String>,
1248 seen: Mutex<Vec<String>>,
1249 entered: Arc<(Mutex<bool>, Condvar)>,
1250 release: Arc<(Mutex<bool>, Condvar)>,
1251 }
1252
1253 impl Backend for HoldKey {
1254 fn id(&self) -> &str {
1255 "hold"
1256 }
1257
1258 fn replace_api_key(&self, key: Option<String>) -> Result<Option<String>, Error> {
1259 let Some(key) = key.filter(|value| !value.is_empty()) else {
1260 return Err(Error::Auth("SNAPIF_API_KEY".to_string()));
1261 };
1262 let mut guard = self.key.lock().expect("key");
1263 let previous = Some(guard.clone());
1264 *guard = key;
1265 Ok(previous)
1266 }
1267
1268 async fn evaluate(
1269 &self,
1270 _req: WireRequest,
1271 _deadline: std::time::Instant,
1272 ) -> Result<crate::backend::Evaluated, BackendError> {
1273 let key = self.key.lock().expect("key").clone();
1274 let n = {
1275 let mut seen = self.seen.lock().expect("seen");
1276 seen.push(key);
1277 seen.len()
1278 };
1279 if n == 1 {
1280 {
1281 let (lock, cv) = &*self.entered;
1282 *lock.lock().expect("entered") = true;
1283 cv.notify_one();
1284 }
1285 let (lock, cv) = &*self.release;
1286 let mut go = lock.lock().expect("release");
1287 while !*go {
1288 go = cv.wait(go).expect("wait");
1289 }
1290 }
1291 Err(BackendError::Timeout)
1292 }
1293 }
1294
1295 #[test]
1296 fn in_flight_evaluate_keeps_the_key_it_copied() {
1297 let backend = HoldKey {
1298 key: Mutex::new("old".to_string()),
1299 seen: Mutex::new(Vec::new()),
1300 entered: Arc::new((Mutex::new(false), Condvar::new())),
1301 release: Arc::new((Mutex::new(false), Condvar::new())),
1302 };
1303 let entered = Arc::clone(&backend.entered);
1304 let release = Arc::clone(&backend.release);
1305 let client = Arc::new(Client::new(backend));
1306 let first = Arc::clone(&client);
1307 let handle = std::thread::spawn(move || {
1308 let req = WireRequest {
1309 model: "jev-latest".to_string(),
1310 state: serde_json::json!({}),
1311 questions: indexmap::IndexMap::new(),
1312 };
1313 pollster::block_on(
1314 first
1315 .backend()
1316 .evaluate(req, Instant::now() + Duration::from_secs(2)),
1317 )
1318 });
1319 {
1320 let (lock, cv) = &*entered;
1321 let mut ready = lock.lock().expect("entered");
1322 while !*ready {
1323 ready = cv.wait(ready).expect("wait");
1324 }
1325 }
1326 client
1327 .replace_api_key(Some("new".to_string()))
1328 .expect("swap");
1329 let req = WireRequest {
1330 model: "jev-latest".to_string(),
1331 state: serde_json::json!({}),
1332 questions: indexmap::IndexMap::new(),
1333 };
1334 let _ = pollster::block_on(
1335 client
1336 .backend()
1337 .evaluate(req, Instant::now() + Duration::from_secs(2)),
1338 );
1339 {
1340 let (lock, cv) = &*release;
1341 *lock.lock().expect("release") = true;
1342 cv.notify_one();
1343 }
1344 let _ = handle.join().expect("join");
1345 assert_eq!(
1346 client.backend().seen.lock().expect("seen").as_slice(),
1347 ["old".to_string(), "new".to_string()]
1348 );
1349 }
1350}