use std::collections::HashMap;
use std::future::Future;
use std::path::PathBuf;
use std::sync::Mutex;
use std::time::{Duration, Instant};
use indexmap::IndexMap;
use crate::answer::{ChoiceAnswer, NoulAnswer, ScoreAnswer};
use crate::backends::fake::FakeBackend;
use crate::error::{BackendError, DecodeError, Error, PolicyError};
use crate::ids::QuestionId;
use crate::policy::{Fail, Policy};
use crate::question::{ChoiceLabels, Question, ScoreLabels};
use crate::state::State;
use crate::usage::UsageFn;
use crate::verdict::{Decision, UnsureReason, UntypedDecision};
use crate::wire::{
self, Usage, WireAnswer, WireQuestion, WireRequest, WireResponse, renormalize_probabilities,
};
pub const TYPESAFE_ORIGIN: &str = "https://api.typesafe.ai";
#[derive(Debug)]
pub struct Evaluated {
pub wire: WireResponse,
pub meta: IndexMap<String, AnswerMeta>,
pub backend_id: String,
}
pub use crate::verdict::{AnswerMeta, CascadeHop};
pub trait Backend: Send + Sync {
fn id(&self) -> &str;
fn evaluate(
&self,
req: WireRequest,
deadline: Instant,
) -> impl Future<Output = Result<Evaluated, BackendError>> + Send;
fn replace_api_key(&self, key: Option<String>) -> Result<Option<String>, Error> {
let _ = key;
Ok(None)
}
}
pub enum AnyBackend {
Fake(FakeBackend),
#[cfg(feature = "http")]
Http(crate::backends::http::HttpBackend),
#[cfg(feature = "http")]
CascadeHttp(
Box<
crate::backends::cascade::Cascaded<
crate::backends::http::HttpBackend,
crate::backends::http::HttpBackend,
>,
>,
),
}
impl Backend for AnyBackend {
fn id(&self) -> &str {
match self {
AnyBackend::Fake(backend) => backend.id(),
#[cfg(feature = "http")]
AnyBackend::Http(backend) => backend.id(),
#[cfg(feature = "http")]
AnyBackend::CascadeHttp(backend) => backend.id(),
}
}
async fn evaluate(
&self,
req: WireRequest,
deadline: Instant,
) -> Result<Evaluated, BackendError> {
match self {
AnyBackend::Fake(backend) => backend.evaluate(req, deadline).await,
#[cfg(feature = "http")]
AnyBackend::Http(backend) => backend.evaluate(req, deadline).await,
#[cfg(feature = "http")]
AnyBackend::CascadeHttp(backend) => backend.evaluate(req, deadline).await,
}
}
fn replace_api_key(&self, key: Option<String>) -> Result<Option<String>, Error> {
match self {
AnyBackend::Fake(backend) => backend.replace_api_key(key),
#[cfg(feature = "http")]
AnyBackend::Http(backend) => backend.replace_api_key(key),
#[cfg(feature = "http")]
AnyBackend::CascadeHttp(backend) => backend.replace_api_key(key),
}
}
}
pub(crate) struct GateCache {
cap: usize,
ttl: Duration,
entries: HashMap<u64, (Instant, crate::verdict::Verdict)>,
}
impl GateCache {
pub(crate) fn new(cap: usize) -> Self {
Self {
cap,
ttl: Duration::from_secs(30),
entries: HashMap::new(),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct ClientConfig {
pub backend: String,
pub shadow: bool,
pub model: Option<String>,
pub timeout_ms: Option<String>,
pub policy: Option<String>,
pub base_url: Option<String>,
pub api_key: Option<String>,
pub typesafe_key: Option<String>,
pub allow_private_http: bool,
pub cascade: Option<String>,
pub log_path: Option<PathBuf>,
pub cache_capacity: Option<usize>,
}
#[derive(Clone, Copy, Default)]
pub struct CallChoice<'a> {
pub model: Option<&'a str>,
pub policy: Option<&'a Policy>,
pub timeout: Option<Duration>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ClientStatus {
pub backend: String,
pub model: String,
pub policy: String,
pub log: bool,
pub cache: bool,
}
pub struct Client<B: Backend> {
pub(crate) backend: B,
pub(crate) policy: Option<Policy>,
pub(crate) on_usage: Option<UsageFn>,
pub(crate) timeout: Duration,
pub(crate) model: String,
shadow_override: Option<bool>,
fail_override: Option<Fail>,
pub(crate) log_path: Option<PathBuf>,
pub(crate) log_lock: Mutex<()>,
pub(crate) cache: Option<Mutex<GateCache>>,
}
impl<B: Backend> Client<B> {
pub fn new(backend: B) -> Self {
Self {
backend,
policy: None,
on_usage: None,
timeout: Duration::from_millis(2000),
model: "jev-latest".to_string(),
shadow_override: None,
fail_override: None,
log_path: None,
log_lock: Mutex::new(()),
cache: None,
}
}
pub fn fail(mut self, fail: Fail) -> Self {
self.fail_override = Some(fail);
self
}
pub(crate) fn fail_for(&self, policy: Option<&Policy>) -> Fail {
if let Some(fail) = self.fail_override {
return fail;
}
policy.map(|policy| policy.fail).unwrap_or(Fail::Closed)
}
pub fn shadow(mut self, on: bool) -> Self {
self.shadow_override = Some(on);
self
}
pub(crate) fn shadow_on(&self, policy_shadow: bool) -> bool {
self.shadow_override.unwrap_or(policy_shadow)
}
pub fn policy(mut self, policy: Policy) -> Self {
self.policy = Some(policy);
self
}
pub fn on_usage(mut self, on_usage: UsageFn) -> Self {
self.on_usage = Some(on_usage);
self
}
pub fn timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout;
self
}
pub fn model(mut self, id: impl AsRef<str>) -> Result<Self, Error> {
let id = id.as_ref().trim();
if id.is_empty() {
return Err(Error::Policy(PolicyError::Config(
"SNAPIF_MODEL must not be blank".to_string(),
)));
}
self.model = id.to_string();
Ok(self)
}
pub fn backend(&self) -> &B {
&self.backend
}
pub fn status(&self) -> ClientStatus {
ClientStatus {
backend: self.backend.id().to_string(),
model: self.model.clone(),
policy: self.policy.as_ref().map(policy_label).unwrap_or_default(),
log: self.log_path.is_some(),
cache: self.cache.is_some(),
}
}
pub fn replace_api_key(&self, key: Option<String>) -> Result<(), Error> {
self.backend.replace_api_key(key).map(|_| ())
}
pub fn battery_id(&self) -> Option<&str> {
self.policy.as_ref().map(|policy| policy.battery.0.as_str())
}
pub(crate) fn cache_get(&self, key: u64) -> Option<crate::verdict::Verdict> {
let cache = self.cache.as_ref()?;
let mut guard = cache.lock().ok()?;
let expired = guard
.entries
.get(&key)
.is_some_and(|(stored, _)| stored.elapsed() > guard.ttl);
if expired {
guard.entries.remove(&key);
return None;
}
guard.entries.get(&key).map(|(_, verdict)| verdict.clone())
}
pub(crate) fn cache_put(&self, key: u64, verdict: crate::verdict::Verdict) {
let Some(cache) = &self.cache else {
return;
};
let Ok(mut guard) = cache.lock() else {
return;
};
if guard.entries.len() >= guard.cap
&& !guard.entries.contains_key(&key)
&& let Some(old) = guard.entries.keys().next().copied()
{
guard.entries.remove(&old);
}
guard.entries.insert(key, (Instant::now(), verdict));
}
pub async fn ask(&self, state: State, questions: Vec<Question>) -> Result<AskOut, Error> {
self.ask_with(state, questions, CallChoice::default()).await
}
pub async fn ask_with(
&self,
state: State,
questions: Vec<Question>,
choice: CallChoice<'_>,
) -> Result<AskOut, Error> {
let policy = choice
.policy
.or(self.policy.as_ref())
.ok_or_else(|| Error::Policy(PolicyError::Invariant("policy".to_string())))?;
let model = match choice.model.map(str::trim).filter(|text| !text.is_empty()) {
Some(model) => model.to_string(),
None => self.model.clone(),
};
policy.ensure_checked()?;
let mut request = WireRequest {
model: model.clone(),
state: state.to_wire(None),
questions: questions.iter().map(wire_question).collect(),
};
let encoded = wire::encode(&request)?;
if encoded.truncated_untrusted {
request = wire::decode_request(&encoded.body)?;
}
let timeout = choice.timeout.unwrap_or(self.timeout);
let deadline = Instant::now() + timeout;
let evaluated = self
.backend
.evaluate(request.clone(), deadline)
.await
.map_err(|err| map_backend(err, timeout))?;
if let Some(on_usage) = &self.on_usage {
on_usage(evaluated.wire.usage);
}
wire::check_response(&request.questions, &evaluated.wire)?;
let Evaluated {
wire,
mut meta,
backend_id,
} = evaluated;
let mut decisions = IndexMap::new();
let mut scores = IndexMap::new();
for (id, question) in &request.questions {
let Some(answer) = wire.answers.get(id) else {
continue;
};
record_prob_sum(&mut meta, id, answer);
scores.insert(id.clone(), answer_score(answer));
let key = QuestionId::new(id);
let decision = untyped(policy, question, answer).map_err(Error::Decode)?;
decisions.insert(key, decision);
}
let (pack, pack_version) = match policy.shipped_id.clone() {
Some(id) => {
let version = Policy::shipped_pack_version(&id);
(id, version)
}
None => (String::new(), 0),
};
Ok(AskOut {
decisions,
scores,
usage: wire.usage,
backend_id,
meta,
truncated_untrusted: encoded.truncated_untrusted,
pack,
pack_version,
model,
})
}
}
impl<B: Backend> std::fmt::Debug for Client<B> {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let status = self.status();
formatter
.debug_struct("Client")
.field("backend", &status.backend)
.field("model", &status.model)
.field("policy", &status.policy)
.field("log", &status.log)
.field("cache", &status.cache)
.finish()
}
}
#[derive(Debug, Default)]
struct BackendEnv {
#[cfg(feature = "http")]
typesafe_key: Option<String>,
#[cfg(feature = "http")]
snapif_key: Option<String>,
#[cfg(feature = "http")]
base_url: Option<String>,
#[cfg(feature = "http")]
allow_private_http: bool,
#[cfg_attr(not(feature = "http"), allow(dead_code))]
cascade: Option<String>,
model: Option<String>,
timeout_ms: Option<String>,
policy: Option<String>,
}
impl Client<AnyBackend> {
pub fn from_env() -> Result<Self, Error> {
let cache_capacity = match std::env::var("SNAPIF_CACHE").ok().as_deref() {
None => None,
Some(raw) => {
let raw = raw.trim();
if raw.is_empty() {
None
} else {
Some(raw.parse::<usize>().map_err(|_| {
Error::Policy(PolicyError::Config(format!(
"SNAPIF_CACHE must be an integer, got {raw}"
)))
})?)
}
}
};
let config = ClientConfig {
backend: std::env::var("SNAPIF_BACKEND").unwrap_or_default(),
shadow: matches!(
std::env::var("SNAPIF_SHADOW").ok().as_deref(),
Some("1" | "true" | "TRUE" | "True")
),
model: std::env::var("SNAPIF_MODEL").ok(),
timeout_ms: std::env::var("SNAPIF_TIMEOUT_MS").ok(),
policy: std::env::var("SNAPIF_POLICY").ok(),
#[cfg(feature = "http")]
base_url: std::env::var("SNAPIF_BASE_URL").ok(),
#[cfg(not(feature = "http"))]
base_url: None,
#[cfg(feature = "http")]
api_key: std::env::var("SNAPIF_API_KEY").ok(),
#[cfg(not(feature = "http"))]
api_key: None,
#[cfg(feature = "http")]
typesafe_key: std::env::var("TYPESAFE_API_KEY").ok(),
#[cfg(not(feature = "http"))]
typesafe_key: None,
#[cfg(feature = "http")]
allow_private_http: matches!(
std::env::var("SNAPIF_ALLOW_PRIVATE_HTTP").ok().as_deref(),
Some("1" | "true" | "TRUE" | "True")
),
#[cfg(not(feature = "http"))]
allow_private_http: false,
cascade: std::env::var("SNAPIF_CASCADE_BASE_URL").ok(),
log_path: std::env::var("SNAPIF_LOG")
.ok()
.map(|raw| raw.trim().to_string())
.filter(|raw| !raw.is_empty())
.map(PathBuf::from),
cache_capacity,
};
Self::from_config(&config)
}
pub fn from_config(config: &ClientConfig) -> Result<Self, Error> {
let env = BackendEnv {
#[cfg(feature = "http")]
typesafe_key: config.typesafe_key.clone(),
#[cfg(feature = "http")]
snapif_key: config.api_key.clone(),
#[cfg(feature = "http")]
base_url: config.base_url.clone(),
#[cfg(feature = "http")]
allow_private_http: config.allow_private_http,
cascade: config.cascade.clone(),
model: config.model.clone(),
timeout_ms: config.timeout_ms.clone(),
policy: config.policy.clone(),
};
let name = match config.backend.trim() {
"" => None,
other => Some(other),
};
let mut client = Self::from_parts(name, config.shadow, &env)?;
client.log_path = config.log_path.clone();
if let Some(cap) = config.cache_capacity.filter(|cap| *cap > 0) {
client.cache = Some(Mutex::new(GateCache::new(cap)));
}
Ok(client)
}
pub fn replace_fake(&mut self, backend: FakeBackend) -> Result<(), Error> {
match &mut self.backend {
AnyBackend::Fake(slot) => {
*slot = backend;
Ok(())
}
#[cfg(feature = "http")]
_ => Err(Error::Policy(PolicyError::Config(
"replace_fake requires a fake backend".to_string(),
))),
}
}
#[cfg_attr(not(feature = "http"), allow(unused_variables))]
fn from_parts(name: Option<&str>, shadow: bool, env: &BackendEnv) -> Result<Self, Error> {
let name = match name.map(str::trim) {
None | Some("") => None,
Some(text) => Some(text),
};
let policy = env_policy(env)?;
let client = match name {
None => {
return Err(Error::Policy(PolicyError::BackendName(
"SNAPIF_BACKEND must be fake, typesafe, or compatible".to_string(),
)));
}
Some("fake") => Self::new(AnyBackend::Fake(FakeBackend::new())).policy(policy),
#[cfg(feature = "http")]
Some(name @ ("typesafe" | "compatible")) => http_client(name, env, policy)?,
#[cfg(not(feature = "http"))]
Some(name @ ("typesafe" | "compatible")) => {
return Err(Error::Policy(PolicyError::BackendName(format!(
"SNAPIF_BACKEND {name} needs the http feature"
))));
}
Some(other) => {
return Err(Error::Policy(PolicyError::BackendName(format!(
"unknown SNAPIF_BACKEND {other}; expected fake, typesafe, or compatible"
))));
}
};
let client = apply_runtime(client, env)?;
Ok(if shadow { client.shadow(true) } else { client })
}
}
fn policy_label(policy: &Policy) -> String {
if let Some(id) = &policy.shipped_id {
return id.clone();
}
match &policy.source_path {
Some(path) => format!("{} {path}", policy.battery.0),
None => policy.battery.0.clone(),
}
}
fn env_policy(env: &BackendEnv) -> Result<Policy, Error> {
match nonempty(env.policy.as_deref()) {
Some(spec) => Policy::load(spec),
None => Ok(Policy::shipped("tool-gate")?),
}
}
fn apply_runtime(
mut client: Client<AnyBackend>,
env: &BackendEnv,
) -> Result<Client<AnyBackend>, Error> {
if let Some(raw) = nonempty(env.timeout_ms.as_deref()) {
let ms: u64 = raw.parse().map_err(|_| {
Error::Policy(PolicyError::Config(format!(
"SNAPIF_TIMEOUT_MS must be an integer, got {raw}"
)))
})?;
client = client.timeout(Duration::from_millis(ms));
}
if let Some(raw) = env.model.as_deref() {
let model = raw.trim();
if model.is_empty() {
return Err(Error::Policy(PolicyError::Config(
"SNAPIF_MODEL must not be blank".to_string(),
)));
}
client.model = model.to_string();
}
Ok(client)
}
fn nonempty(value: Option<&str>) -> Option<&str> {
value.map(str::trim).filter(|text| !text.is_empty())
}
#[cfg(feature = "http")]
fn http_client(name: &str, env: &BackendEnv, policy: Policy) -> Result<Client<AnyBackend>, Error> {
let backend = if let Some(raw) = nonempty(env.cascade.as_deref()) {
let first = compatible_backend(
raw,
env.snapif_key.clone(),
"SNAPIF_CASCADE_BASE_URL",
env.allow_private_http,
)?;
let fallback = match name {
"typesafe" => crate::backends::http::HttpBackend::typesafe(
env.typesafe_key.clone().unwrap_or_default(),
)?,
"compatible" => {
let Some(base) = env.base_url.as_deref().filter(|value| !value.is_empty()) else {
return Err(Error::Policy(PolicyError::Config(
"SNAPIF_BASE_URL is required".to_string(),
)));
};
compatible_backend(
base,
env.snapif_key.clone(),
"SNAPIF_BASE_URL",
env.allow_private_http,
)?
}
_ => return Err(Error::Policy(PolicyError::Invariant(name.to_string()))),
};
let rule = crate::backends::cascade::CascadeRule::new(policy.cascade_min);
AnyBackend::CascadeHttp(Box::new(crate::backends::cascade::Cascaded::new(
first, fallback, rule,
)))
} else {
match name {
"typesafe" => AnyBackend::Http(crate::backends::http::HttpBackend::typesafe(
env.typesafe_key.clone().unwrap_or_default(),
)?),
"compatible" => {
let Some(base) = env.base_url.as_deref().filter(|value| !value.is_empty()) else {
return Err(Error::Policy(PolicyError::Config(
"SNAPIF_BASE_URL is required".to_string(),
)));
};
AnyBackend::Http(compatible_backend(
base,
env.snapif_key.clone(),
"SNAPIF_BASE_URL",
env.allow_private_http,
)?)
}
_ => return Err(Error::Policy(PolicyError::Invariant(name.to_string()))),
}
};
Ok(Client::new(backend).policy(policy))
}
#[cfg(feature = "http")]
fn compatible_backend(
raw: &str,
key: Option<String>,
invariant: &str,
allow_private_http: bool,
) -> Result<crate::backends::http::HttpBackend, Error> {
let url = url::Url::parse(raw)
.map_err(|_| Error::Policy(PolicyError::Config(format!("{invariant} must be a URL"))))?;
if allow_private_http {
crate::backends::http::HttpBackend::compatible_private(url, key)
} else {
crate::backends::http::HttpBackend::compatible(url, key)
}
}
#[derive(Debug)]
#[non_exhaustive]
pub struct AskOut {
pub decisions: IndexMap<QuestionId, UntypedDecision>,
pub scores: IndexMap<String, f64>,
pub usage: Usage,
pub backend_id: String,
pub meta: IndexMap<String, AnswerMeta>,
pub truncated_untrusted: bool,
pub pack: String,
pub pack_version: u32,
pub model: String,
}
impl AskOut {
pub fn choice<T: ChoiceLabels>(&self, id: &QuestionId) -> Result<Decision<T>, DecodeError> {
match self.decisions.get(id) {
None => Err(DecodeError::MissingAnswer { key: id.clone() }),
Some(UntypedDecision::Choice(Decision::Known(label))) => match T::from_label(label) {
Some(value) => Ok(Decision::Known(value)),
None => Err(DecodeError::UnknownLabel {
key: id.clone(),
label: label.clone(),
}),
},
Some(UntypedDecision::Choice(Decision::Unsure { reason, guess })) => {
Ok(Decision::Unsure {
reason: reason.clone(),
guess: guess.as_deref().and_then(T::from_label),
})
}
Some(_) => Err(DecodeError::TypeMismatch { key: id.clone() }),
}
}
pub fn score<T: ScoreLabels>(&self, id: &QuestionId) -> Result<Decision<T>, DecodeError> {
match self.decisions.get(id) {
None => Err(DecodeError::MissingAnswer { key: id.clone() }),
Some(UntypedDecision::Score(Decision::Known(score))) => {
score_label::<T>(*score, id.clone())
}
Some(UntypedDecision::Score(Decision::Unsure { reason, guess })) => {
Ok(Decision::Unsure {
reason: reason.clone(),
guess: guess.and_then(|score| score_index(score).and_then(T::from_index)),
})
}
Some(_) => Err(DecodeError::TypeMismatch { key: id.clone() }),
}
}
pub fn noul(&self, id: &QuestionId) -> Result<Decision<bool>, DecodeError> {
match self.decisions.get(id) {
None => Err(DecodeError::MissingAnswer { key: id.clone() }),
Some(UntypedDecision::Noul(decision)) => Ok(decision.clone()),
Some(_) => Err(DecodeError::TypeMismatch { key: id.clone() }),
}
}
}
fn score_label<T: ScoreLabels>(score: f64, id: QuestionId) -> Result<Decision<T>, DecodeError> {
match score_index(score).and_then(T::from_index) {
Some(value) => Ok(Decision::Known(value)),
None => Err(DecodeError::OutOfRange { key: id }),
}
}
fn score_index(score: f64) -> Option<usize> {
if score.is_finite() && score >= 0.0 {
Some(score.round() as usize)
} else {
None
}
}
fn map_backend(err: BackendError, timeout: Duration) -> Error {
match err {
BackendError::Timeout => Error::Timeout(timeout),
BackendError::RateLimit => Error::RateLimit,
BackendError::Overloaded => Error::Overloaded,
BackendError::Auth => Error::Auth("authentication failed (HTTP 401)".to_string()),
BackendError::Rejected { status, body } => Error::Rejected { status, body },
other => Error::Backend(other.to_string()),
}
}
pub(crate) fn wire_question(question: &Question) -> (String, WireQuestion) {
match question {
Question::Choice(choice) => (
choice.id.to_string(),
WireQuestion::Choice {
instructions: choice.instructions.clone(),
criteria: choice.criteria.clone(),
},
),
Question::Score(score) => (
score.id.to_string(),
WireQuestion::Score {
instructions: score.instructions.clone(),
criteria: score.criteria.clone(),
},
),
Question::Noul(noul) => (
noul.id.to_string(),
WireQuestion::Noul {
instructions: noul.instructions.clone(),
criteria: noul.criteria.clone(),
},
),
}
}
fn answer_score(answer: &WireAnswer) -> f64 {
match answer {
WireAnswer::Choice { confidence, .. } => *confidence,
WireAnswer::Score { score, .. } => *score,
WireAnswer::Noul { noul } => *noul,
}
}
pub(crate) fn record_prob_sum(
meta: &mut IndexMap<String, AnswerMeta>,
id: &str,
answer: &WireAnswer,
) {
let probabilities = match answer {
WireAnswer::Choice { probabilities, .. } | WireAnswer::Score { probabilities, .. } => {
probabilities
}
WireAnswer::Noul { .. } => return,
};
let (_, original_sum) = renormalize_probabilities(probabilities);
if (original_sum - 1.0).abs() > 1e-6 {
meta.entry(id.to_string()).or_default().original_prob_sum = Some(original_sum);
}
}
fn untyped(
policy: &Policy,
question: &WireQuestion,
answer: &WireAnswer,
) -> Result<UntypedDecision, DecodeError> {
match (question, answer) {
(
WireQuestion::Choice { .. },
WireAnswer::Choice {
choice,
probabilities,
confidence,
},
) => {
let decoded = ChoiceAnswer {
label: choice.clone(),
confidence: *confidence,
probabilities: probabilities.clone(),
};
let signal = decoded.signal(policy.choice.signal);
let floor = policy.choice.escalate_below;
let decision = if signal < floor {
Decision::Unsure {
reason: UnsureReason::BelowFloor {
confidence: signal,
floor,
},
guess: Some(decoded.label),
}
} else {
Decision::Known(decoded.label)
};
Ok(UntypedDecision::Choice(decision))
}
(
WireQuestion::Score { .. },
WireAnswer::Score {
score,
probabilities,
confidence,
..
},
) => {
let decoded = ScoreAnswer {
score: *score,
confidence: *confidence,
probabilities: probabilities.clone(),
};
let signal = decoded.signal(policy.choice.signal);
let floor = policy.choice.escalate_below;
let decision = if signal < floor {
Decision::Unsure {
reason: UnsureReason::BelowFloor {
confidence: signal,
floor,
},
guess: Some(decoded.score),
}
} else {
Decision::Known(decoded.score)
};
Ok(UntypedDecision::Score(decision))
}
(WireQuestion::Noul { .. }, WireAnswer::Noul { noul }) => {
let decision = NoulAnswer { p: *noul }.decide(&policy.noul, None);
Ok(UntypedDecision::Noul(decision))
}
_ => Err(DecodeError::TypeMismatch {
key: QuestionId::new("answer"),
}),
}
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Condvar, Mutex};
use std::time::{Duration, Instant};
use super::{AnyBackend, Backend, Client, ClientConfig};
use crate::error::{BackendError, Error, PolicyError};
use crate::wire::WireRequest;
#[test]
fn from_name_selects_without_env() {
let env = super::BackendEnv::default();
let config = ClientConfig {
backend: "fake".to_string(),
timeout_ms: Some("nope".to_string()),
..ClientConfig::default()
};
let bad = match Client::<AnyBackend>::from_config(&config) {
Err(err) => err,
Ok(_) => panic!("timeout"),
};
assert!(bad.to_string().contains("SNAPIF_TIMEOUT_MS"), "{bad}");
let ok = Client::<AnyBackend>::from_config(&ClientConfig {
backend: "fake".to_string(),
..ClientConfig::default()
})
.expect("fake");
assert_eq!(ok.backend().id(), "fake");
let Err(unset) = Client::<AnyBackend>::from_parts(None, false, &env) else {
panic!("unset must be policy");
};
assert!(matches!(
unset,
Error::Policy(PolicyError::BackendName(ref message))
if message.contains("SNAPIF_BACKEND")
&& message.contains("fake")
&& message.contains("typesafe")
&& message.contains("compatible")
));
let spaced =
Client::<AnyBackend>::from_parts(Some(" fake "), false, &env).expect("trimmed fake");
assert_eq!(spaced.backend().id(), "fake");
let shown = unset.to_string();
assert!(!shown.contains("threshold invariant"), "{shown}");
let Err(unknown_name) = Client::<AnyBackend>::from_parts(Some("laya"), false, &env) else {
panic!("unknown backend");
};
let unknown_text = unknown_name.to_string();
assert!(
unknown_text.contains("unknown SNAPIF_BACKEND laya"),
"{unknown_text}"
);
assert!(unknown_text.contains("fake"), "{unknown_text}");
assert!(
!unknown_text.contains("threshold invariant"),
"{unknown_text}"
);
let client = Client::<AnyBackend>::from_parts(Some("fake"), false, &env).expect("fake");
assert_eq!(client.backend().id(), "fake");
let shadowed = Client::<AnyBackend>::from_parts(Some("fake"), true, &env).expect("shadow");
assert_eq!(shadowed.shadow_override, Some(true));
#[cfg(not(feature = "http"))]
{
let Err(unknown) = Client::<AnyBackend>::from_parts(Some("typesafe"), false, &env)
else {
panic!("typesafe must be policy");
};
assert!(matches!(
unknown,
Error::Policy(PolicyError::BackendName(ref message))
if message.contains("typesafe") && message.contains("http feature")
));
assert!(!unknown.to_string().contains("threshold invariant"));
let cascade_env = super::BackendEnv {
cascade: Some("http://127.0.0.1:9".to_string()),
..super::BackendEnv::default()
};
let err = Client::<AnyBackend>::from_parts(Some("typesafe"), false, &cascade_env);
assert!(matches!(err, Err(Error::Policy(_))));
}
#[cfg(feature = "http")]
{
let Err(missing_key) = Client::<AnyBackend>::from_parts(Some("typesafe"), false, &env)
else {
panic!("typesafe without a key must be auth");
};
assert!(matches!(missing_key, Error::Auth(_)));
let selected = Client::<AnyBackend>::from_parts(
Some("typesafe"),
false,
&super::BackendEnv {
typesafe_key: Some("secret".to_string()),
..super::BackendEnv::default()
},
)
.expect("typesafe");
assert_eq!(selected.backend().id(), "typesafe");
let cascaded = Client::<AnyBackend>::from_parts(
Some("typesafe"),
false,
&super::BackendEnv {
typesafe_key: Some("secret".to_string()),
cascade: Some("http://127.0.0.1:9".to_string()),
..super::BackendEnv::default()
},
)
.expect("cascade");
assert_eq!(cascaded.backend().id(), "cascade");
}
let tuned = Client::<AnyBackend>::from_parts(
Some("fake"),
false,
&super::BackendEnv {
model: Some("custom-model".to_string()),
timeout_ms: Some("1500".to_string()),
..super::BackendEnv::default()
},
)
.expect("runtime");
assert_eq!(tuned.model, "custom-model");
assert_eq!(tuned.timeout, std::time::Duration::from_millis(1500));
let bad_timeout = Client::<AnyBackend>::from_parts(
Some("fake"),
false,
&super::BackendEnv {
timeout_ms: Some("nope".to_string()),
..super::BackendEnv::default()
},
);
assert!(matches!(bad_timeout, Err(Error::Policy(_))));
let bad_policy = Client::<AnyBackend>::from_parts(
Some("fake"),
false,
&super::BackendEnv {
policy: Some("missing-policy".to_string()),
..super::BackendEnv::default()
},
);
assert!(matches!(bad_policy, Err(Error::Policy(_))));
}
#[test]
fn status_reads_the_resolved_model_and_hides_the_key() {
let client = Client::<AnyBackend>::from_config(&ClientConfig {
backend: "fake".to_string(),
model: Some("jev-1.12".to_string()),
policy: Some("tool-gate".to_string()),
api_key: Some("super-secret-key".to_string()),
log_path: Some(std::path::PathBuf::from("/tmp/snapif.log")),
cache_capacity: Some(2),
..ClientConfig::default()
})
.expect("config");
let status = client.status();
assert_eq!(status.backend, "fake");
assert_eq!(status.model, "jev-1.12");
assert_eq!(status.policy, "tool-gate");
assert!(status.log);
assert!(status.cache);
let shown = format!("{client:?} {status:?}");
assert!(!shown.contains("super-secret-key"), "{shown}");
let plain = Client::<AnyBackend>::from_config(&ClientConfig {
backend: "fake".to_string(),
..ClientConfig::default()
})
.expect("default");
assert_eq!(plain.status().model, "jev-latest");
assert_eq!(plain.status().policy, "tool-gate");
let raw = include_str!("../policies/tool-gate.toml");
let file_policy = Client::new(crate::backends::fake::FakeBackend::new())
.policy(crate::policy::Policy::from_toml_str(raw).expect("toml"));
assert_eq!(file_policy.status().policy, "tool-gate");
}
#[test]
fn status_names_a_file_policy_without_the_key() {
let dir = std::env::temp_dir().join(format!("snapif-status-{}", std::process::id()));
let _ = std::fs::create_dir_all(&dir);
let path = dir.join("desk.toml");
std::fs::write(
&path,
r#"
schema_version = 1
battery = "desk-pack"
[choice]
escalate_below = 0.8
review_below = 1.0
[default_action]
review = 0.8
when_unsure = "review_guess"
class = "read"
"#,
)
.expect("toml");
let client = Client::<AnyBackend>::from_config(&ClientConfig {
backend: "fake".to_string(),
policy: Some(path.display().to_string()),
api_key: Some("super-secret-key".to_string()),
..ClientConfig::default()
})
.expect("config");
let status = client.status();
assert!(status.policy.contains("desk-pack"), "{}", status.policy);
assert!(
status.policy.contains(&path.display().to_string()),
"{}",
status.policy
);
let shown = format!("{client:?} {status:?}");
assert!(!shown.contains("super-secret-key"), "{shown}");
assert!(!shown.contains("when_unsure"), "{shown}");
let _ = std::fs::remove_dir_all(&dir);
}
#[cfg(feature = "http")]
#[test]
fn replace_fake_on_typesafe_does_not_blame_the_env() {
let mut client = Client::<AnyBackend>::from_config(&ClientConfig {
backend: "typesafe".to_string(),
typesafe_key: Some("sekrit-key-xyz".to_string()),
..ClientConfig::default()
})
.expect("typesafe");
let err = client
.replace_fake(crate::backends::fake::FakeBackend::new())
.expect_err("not fake");
let text = err.to_string();
assert!(
text.contains("replace_fake requires a fake backend"),
"{text}"
);
assert!(!text.contains("SNAPIF_BACKEND"), "{text}");
assert!(!text.contains("script"), "{text}");
assert!(!text.contains("sekrit-key-xyz"), "{text}");
}
#[test]
fn model_env_rejects_blank_and_keeps_the_default() {
let default =
Client::<AnyBackend>::from_parts(Some("fake"), false, &super::BackendEnv::default())
.expect("fake");
assert_eq!(default.model, "jev-latest");
for blank in ["", " ", "\t"] {
let err = Client::<AnyBackend>::from_parts(
Some("fake"),
false,
&super::BackendEnv {
model: Some(blank.to_string()),
..super::BackendEnv::default()
},
);
let Err(err) = err else {
panic!("blank model must fail");
};
assert_eq!(
err.to_string(),
"policy: SNAPIF_MODEL must not be blank",
"{blank:?} {err}"
);
}
let trimmed = Client::<AnyBackend>::from_parts(
Some("fake"),
false,
&super::BackendEnv {
model: Some(" local-model ".to_string()),
..super::BackendEnv::default()
},
)
.expect("trimmed");
assert_eq!(trimmed.model, "local-model");
}
#[cfg(feature = "http")]
#[test]
fn private_http_flag_is_off_unless_set() {
let off = Client::<AnyBackend>::from_parts(
Some("compatible"),
false,
&super::BackendEnv {
base_url: Some("http://10.0.0.1".to_string()),
snapif_key: Some("snapif-key".to_string()),
..super::BackendEnv::default()
},
);
assert!(matches!(off, Err(Error::Policy(_))));
let on = Client::<AnyBackend>::from_parts(
Some("compatible"),
false,
&super::BackendEnv {
base_url: Some("http://10.0.0.1".to_string()),
snapif_key: Some("snapif-key".to_string()),
allow_private_http: true,
..super::BackendEnv::default()
},
);
let Ok(on) = on else {
panic!("private http");
};
assert_eq!(on.backend().id(), "compatible");
}
#[cfg(feature = "http")]
#[test]
fn private_http_flag_covers_the_cascade_origin() {
let denied = Client::<AnyBackend>::from_parts(
Some("compatible"),
false,
&super::BackendEnv {
base_url: Some("http://10.0.0.1".into()),
cascade: Some("http://192.168.1.50".into()),
snapif_key: Some("snapif-key".into()),
allow_private_http: false,
..super::BackendEnv::default()
},
);
assert!(matches!(denied, Err(Error::Policy(_))));
let allowed = Client::<AnyBackend>::from_parts(
Some("compatible"),
false,
&super::BackendEnv {
base_url: Some("http://10.0.0.1".into()),
cascade: Some("http://192.168.1.50".into()),
snapif_key: Some("snapif-key".into()),
allow_private_http: true,
..super::BackendEnv::default()
},
);
let Ok(allowed) = allowed else {
panic!("private http cascade");
};
assert_eq!(allowed.backend().id(), "cascade");
}
#[test]
fn auth_error_reaches_the_client_text() {
let err = super::map_backend(
crate::error::BackendError::Auth,
std::time::Duration::from_secs(1),
);
assert_eq!(err.to_string(), "auth: authentication failed (HTTP 401)");
}
struct HoldKey {
key: Mutex<String>,
seen: Mutex<Vec<String>>,
entered: Arc<(Mutex<bool>, Condvar)>,
release: Arc<(Mutex<bool>, Condvar)>,
}
impl Backend for HoldKey {
fn id(&self) -> &str {
"hold"
}
fn replace_api_key(&self, key: Option<String>) -> Result<Option<String>, Error> {
let Some(key) = key.filter(|value| !value.is_empty()) else {
return Err(Error::Auth("SNAPIF_API_KEY".to_string()));
};
let mut guard = self.key.lock().expect("key");
let previous = Some(guard.clone());
*guard = key;
Ok(previous)
}
async fn evaluate(
&self,
_req: WireRequest,
_deadline: std::time::Instant,
) -> Result<crate::backend::Evaluated, BackendError> {
let key = self.key.lock().expect("key").clone();
let n = {
let mut seen = self.seen.lock().expect("seen");
seen.push(key);
seen.len()
};
if n == 1 {
{
let (lock, cv) = &*self.entered;
*lock.lock().expect("entered") = true;
cv.notify_one();
}
let (lock, cv) = &*self.release;
let mut go = lock.lock().expect("release");
while !*go {
go = cv.wait(go).expect("wait");
}
}
Err(BackendError::Timeout)
}
}
#[test]
fn in_flight_evaluate_keeps_the_key_it_copied() {
let backend = HoldKey {
key: Mutex::new("old".to_string()),
seen: Mutex::new(Vec::new()),
entered: Arc::new((Mutex::new(false), Condvar::new())),
release: Arc::new((Mutex::new(false), Condvar::new())),
};
let entered = Arc::clone(&backend.entered);
let release = Arc::clone(&backend.release);
let client = Arc::new(Client::new(backend));
let first = Arc::clone(&client);
let handle = std::thread::spawn(move || {
let req = WireRequest {
model: "jev-latest".to_string(),
state: serde_json::json!({}),
questions: indexmap::IndexMap::new(),
};
pollster::block_on(
first
.backend()
.evaluate(req, Instant::now() + Duration::from_secs(2)),
)
});
{
let (lock, cv) = &*entered;
let mut ready = lock.lock().expect("entered");
while !*ready {
ready = cv.wait(ready).expect("wait");
}
}
client
.replace_api_key(Some("new".to_string()))
.expect("swap");
let req = WireRequest {
model: "jev-latest".to_string(),
state: serde_json::json!({}),
questions: indexmap::IndexMap::new(),
};
let _ = pollster::block_on(
client
.backend()
.evaluate(req, Instant::now() + Duration::from_secs(2)),
);
{
let (lock, cv) = &*release;
*lock.lock().expect("release") = true;
cv.notify_one();
}
let _ = handle.join().expect("join");
assert_eq!(
client.backend().seen.lock().expect("seen").as_slice(),
["old".to_string(), "new".to_string()]
);
}
}