use std::fmt;
use lgwks_std::json::{Deserialize, Serialize};
use lgwks_std::similarity::{Cosine, CosineError};
use crate::language::{Alias, LanguageResolver, decide};
use crate::session::{
AnswerDomain, DegradedReason, MatchTier, PolicyVersion, Provenance, Question, Resolution,
Resolver, Verdict, by_score_descending,
};
pub trait Embedder {
type Error: std::error::Error + 'static;
fn identity(&self) -> &EmbedderIdentity;
fn embed(&self, text: &str) -> Result<Vec<f32>, Self::Error>;
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, Self::Error> {
texts.iter().map(|text| self.embed(text)).collect()
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(crate = "lgwks_std::json::serde", deny_unknown_fields)]
#[non_exhaustive]
pub struct EmbedderIdentity {
name: String,
digest: String,
dimension: usize,
}
impl EmbedderIdentity {
pub fn new(
name: impl Into<String>,
digest: impl Into<String>,
dimension: usize,
) -> Result<Self, SemanticError> {
let name = name.into();
let digest = digest.into();
if name.is_empty() {
let refusal = Err(SemanticError::UnnamedEmbedder);
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "new: returning an error to the caller");
return refusal;
}
if digest.is_empty() {
let refusal = Err(SemanticError::UndigestedEmbedder { name });
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "new: returning an error to the caller");
return refusal;
}
if dimension == 0 {
let refusal = Err(SemanticError::ZeroDimension { name });
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "new: returning an error to the caller");
return refusal;
}
Ok(Self {
name,
digest,
dimension,
})
}
#[must_use]
pub fn name(&self) -> &str {
&self.name
}
#[must_use]
pub fn digest(&self) -> &str {
&self.digest
}
#[must_use]
pub const fn dimension(&self) -> usize {
self.dimension
}
}
impl fmt::Display for EmbedderIdentity {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
formatter,
"{} (digest {}, {} dimensions)",
self.name, self.digest, self.dimension
)
}
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub enum SemanticError {
UnnamedEmbedder,
UndigestedEmbedder {
name: String,
},
ZeroDimension {
name: String,
},
NonFiniteThreshold,
InvalidThreshold {
threshold: f64,
},
NonFiniteMargin,
InvalidMargin {
margin: f64,
},
}
impl fmt::Display for SemanticError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match *self {
Self::UnnamedEmbedder => {
formatter.write_str("a semantic embedder must be named to be recorded")
}
Self::UndigestedEmbedder { ref name } => write!(
formatter,
"semantic embedder {name} has no weight digest, so a run using it cannot be reproduced"
),
Self::ZeroDimension { ref name } => write!(
formatter,
"semantic embedder {name} declares a zero-width vector"
),
Self::InvalidThreshold { threshold } => write!(
formatter,
"semantic threshold {threshold} must be within [0, 1]"
),
Self::NonFiniteThreshold => {
formatter.write_str("the semantic threshold must be a finite number")
}
Self::NonFiniteMargin => {
formatter.write_str("the semantic margin must be a finite number")
}
Self::InvalidMargin { margin } => {
write!(formatter, "semantic margin {margin} must be within [0, 1]")
}
}
}
}
impl std::error::Error for SemanticError {}
const SEMANTIC_LABEL: &str = "semantic";
#[derive(Debug, Clone, Copy, PartialEq)]
#[non_exhaustive]
pub struct SemanticPolicy {
threshold: f64,
margin: f64,
}
impl SemanticPolicy {
pub const DEFAULT: Self = Self {
threshold: 0.72,
margin: 0.05,
};
pub fn new(threshold: f64, margin: f64) -> Result<Self, SemanticError> {
if !threshold.is_finite() {
let refusal = Err(SemanticError::NonFiniteThreshold);
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "new: returning an error to the caller");
return refusal;
}
if !(0.0..=1.0).contains(&threshold) {
let refusal = Err(SemanticError::InvalidThreshold { threshold });
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "new: returning an error to the caller");
return refusal;
}
if !margin.is_finite() {
let refusal = Err(SemanticError::NonFiniteMargin);
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "new: returning an error to the caller");
return refusal;
}
if !(0.0..=1.0).contains(&margin) {
let refusal = Err(SemanticError::InvalidMargin { margin });
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "new: returning an error to the caller");
return refusal;
}
Ok(Self { threshold, margin })
}
#[must_use]
pub const fn threshold(&self) -> f64 {
self.threshold
}
#[must_use]
pub const fn margin(&self) -> f64 {
self.margin
}
}
impl Default for SemanticPolicy {
fn default() -> Self {
Self::DEFAULT
}
}
#[non_exhaustive]
pub struct SemanticResolver<E> {
lexicon: LanguageResolver,
embedder: E,
policy: SemanticPolicy,
}
impl<E> core::fmt::Debug for SemanticResolver<E> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("SemanticResolver")
.field("lexicon", &self.lexicon)
.field("policy", &self.policy)
.finish_non_exhaustive()
}
}
impl<E: Embedder> SemanticResolver<E> {
#[must_use]
pub fn new(embedder: E) -> Self {
Self {
lexicon: LanguageResolver::new(),
embedder,
policy: SemanticPolicy::DEFAULT,
}
}
#[must_use]
pub fn with_policy(embedder: E, policy: SemanticPolicy) -> Self {
Self {
lexicon: LanguageResolver::new(),
embedder,
policy,
}
}
#[must_use]
pub fn embedder_identity(&self) -> &EmbedderIdentity {
self.embedder.identity()
}
#[must_use]
pub const fn policy(&self) -> SemanticPolicy {
self.policy
}
#[must_use]
pub fn policy_version(&self) -> PolicyVersion {
PolicyVersion::new(
SEMANTIC_LABEL,
&[self.policy.threshold(), self.policy.margin()],
)
}
pub fn learn(&mut self, question: &str, utterance: &str, option: &str) -> Option<Alias> {
self.lexicon.learn(question, utterance, option)
}
pub fn forget(&mut self, question: &str, utterance: &str) -> Option<Alias> {
self.lexicon.forget(question, utterance)
}
#[must_use]
pub fn learned(&self) -> usize {
self.lexicon.learned()
}
fn embed(&self, text: &str) -> Result<Vec<f32>, DegradedReason> {
match self.embedder.embed(text) {
Ok(vector) if vector.len() == self.embedder.identity().dimension() => Ok(vector),
Ok(_) | Err(_) => Err(DegradedReason::EmbedderUnavailable),
}
}
fn degrade(reason: CosineError) -> DegradedReason {
match reason {
CosineError::DimensionMismatch { .. } => DegradedReason::EmbedderUnavailable,
_ => DegradedReason::UnmeasurableEmbedding,
}
}
fn ensure_measurable(metric: &Cosine, vector: &[f32]) -> Result<(), DegradedReason> {
match metric.try_score(vector, vector) {
Ok(_) => Ok(()),
Err(reason) => Err(Self::degrade(reason)),
}
}
fn score_semantically(
&self,
utterance: &str,
options: &[String],
) -> Result<Vec<(usize, MatchTier, f64)>, DegradedReason> {
let metric = Cosine::new();
let target = self.embed(utterance)?;
Self::ensure_measurable(&metric, &target)?;
let mut scored: Vec<(usize, MatchTier, f64)> = Vec::with_capacity(options.len());
for (index, option) in options.iter().enumerate() {
let candidate = self.embed(option)?;
match metric.try_score(&target, &candidate) {
Ok(score) => scored.push((index, MatchTier::Semantic, score)),
Err(reason) => {
let refusal = Err(Self::degrade(reason));
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "score_semantically: returning an error to the caller");
return refusal;
}
}
}
scored.sort_by(by_score_descending);
Ok(scored)
}
}
impl<E: Embedder> Resolver for SemanticResolver<E> {
fn resolve(&self, utterance: &str, question: &Question<'_>) -> Verdict {
let deterministic = self.lexicon.decide_for(utterance, question);
if question.domain() == AnswerDomain::Integer {
return Verdict::new(
deterministic,
Provenance::without_model(self.lexicon.policy_version()),
);
}
if !matches!(deterministic, Resolution::Absent { .. }) {
return Verdict::new(
deterministic,
Provenance::without_model(self.lexicon.policy_version()),
);
}
let provenance =
Provenance::with_model(self.policy_version(), self.embedder.identity().clone());
match self.score_semantically(utterance, question.options()) {
Ok(scored) => Verdict::new(
decide(&scored, self.policy.threshold(), self.policy.margin()),
provenance,
),
Err(reason) => Verdict::new(Resolution::Degraded { reason }, provenance),
}
}
}
#[cfg(test)]
mod tests {
use super::{
Alias, Embedder, EmbedderIdentity, SemanticError, SemanticPolicy, SemanticResolver,
};
use crate::session::{AnswerDomain, DegradedReason, MatchTier, Question, Resolution, Resolver};
use std::cell::Cell;
use std::collections::BTreeMap;
use std::fmt;
use std::rc::Rc;
#[derive(Debug)]
struct StubFailure;
impl fmt::Display for StubFailure {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("stub embedder failure")
}
}
impl std::error::Error for StubFailure {}
struct StubEmbedder {
identity: EmbedderIdentity,
vectors: BTreeMap<String, Vec<f32>>,
fallback: Vec<f32>,
calls: Rc<Cell<usize>>,
failing: bool,
}
impl Embedder for StubEmbedder {
type Error = StubFailure;
fn identity(&self) -> &EmbedderIdentity {
&self.identity
}
fn embed(&self, text: &str) -> Result<Vec<f32>, Self::Error> {
self.calls.set(self.calls.get().saturating_add(1));
if self.failing {
let refusal = Err(StubFailure);
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "embed: returning an error to the caller");
return refusal;
}
match self.vectors.get(text) {
Some(vector) => Ok(vector.clone()),
None => Ok(self.fallback.clone()),
}
}
}
const DIMENSION: usize = 2;
fn the_field() -> Vec<String> {
options(&["Repeat last order", "Cancel"])
}
fn resolve_with(embedder: StubEmbedder) -> crate::session::Verdict {
SemanticResolver::new(embedder).resolve("the usual", &ask(&the_field()))
}
fn near_tie_vectors() -> Vec<(&'static str, Vec<f32>)> {
vec![
("the usual", vec![1.0, 0.0]),
(
"Repeat last order",
vec![0.73, (1.0_f32 - 0.73_f32.powi(2)).sqrt()],
),
("Cancel", vec![0.71, (1.0_f32 - 0.71_f32.powi(2)).sqrt()]),
]
}
fn stub(
vectors: Vec<(&str, Vec<f32>)>,
fallback: Vec<f32>,
failing: bool,
) -> Result<(StubEmbedder, Rc<Cell<usize>>), SemanticError> {
let calls = Rc::new(Cell::new(0));
let embedder = StubEmbedder {
identity: EmbedderIdentity::new("stub-model", "digest-0123456789", DIMENSION)?,
vectors: vectors
.into_iter()
.map(|(text, vector)| (text.to_owned(), vector))
.collect(),
fallback,
calls: Rc::clone(&calls),
failing,
};
Ok((embedder, calls))
}
fn options(names: &[&str]) -> Vec<String> {
names.iter().map(|name| (*name).to_owned()).collect()
}
fn ask(options: &[String]) -> Question<'_> {
Question::new("ask", options)
}
fn paraphrase_vectors() -> Vec<(&'static str, Vec<f32>)> {
vec![
("the usual", vec![1.0, 0.0]),
("Repeat last order", vec![1.0, 0.1]),
("Cancel", vec![0.0, 1.0]),
("Later", vec![-1.0, 0.0]),
]
}
fn assert_vector(actual: &[f32], expected: &[f32]) {
assert_eq!(
actual.len(),
expected.len(),
"vector length changed: {actual:?} vs {expected:?}"
);
for (index, (actual_component, expected_component)) in
actual.iter().zip(expected).enumerate()
{
assert!(
(actual_component - expected_component).abs() < 1e-6,
"component {index} was {actual_component}, expected {expected_component}"
);
}
}
#[test]
fn an_exact_lexicon_match_never_consults_the_model() -> Result<(), Box<dyn std::error::Error>> {
let (embedder, calls) = stub(paraphrase_vectors(), vec![0.0, 1.0], false)?;
let resolver = SemanticResolver::new(embedder);
let resolution = resolver
.resolve("yes", &ask(&options(&["yes", "no"])))
.into_resolution();
assert_eq!(
resolution,
Resolution::Resolved {
index: 0,
tier: MatchTier::Exact,
score: 1.0,
lead: 1.0,
},
"the lexicon's verdict is returned unchanged"
);
assert_eq!(
calls.get(),
0,
"an exact match must not reach the embedder at all"
);
Ok(())
}
#[test]
fn a_lexicon_tie_is_not_rescored_by_the_model() -> Result<(), Box<dyn std::error::Error>> {
let (embedder, calls) = stub(paraphrase_vectors(), vec![0.0, 1.0], false)?;
let resolver = SemanticResolver::new(embedder);
let resolution = resolver
.resolve("order", &ask(&options(&["order", "order"])))
.into_resolution();
assert!(
matches!(resolution, Resolution::Ambiguous { .. }),
"the lexicon's ambiguity survives, got {resolution:?}"
);
assert_eq!(
calls.get(),
0,
"an ambiguous verdict must not reach the model"
);
Ok(())
}
#[test]
fn a_phrase_the_lexicon_cannot_reach_resolves_semantically()
-> Result<(), Box<dyn std::error::Error>> {
let (embedder, calls) = stub(paraphrase_vectors(), vec![0.0, 1.0], false)?;
let resolver = SemanticResolver::new(embedder);
let options = options(&["Repeat last order", "Cancel"]);
let lexicon = crate::language::LanguageResolver::new();
assert!(
matches!(
lexicon
.resolve("the usual", &ask(&options))
.into_resolution(),
Resolution::Absent { .. }
),
"\"the usual\" must not be reachable by letters alone, or this test proves nothing"
);
let resolution = resolver
.resolve("the usual", &ask(&options))
.into_resolution();
assert!(
matches!(
resolution,
Resolution::Resolved {
index: 0,
tier: MatchTier::Semantic,
..
}
),
"expected a semantic resolution of option 0, got {resolution:?}"
);
assert!(
calls.get() > 0,
"the model is what resolved it, so it must have been consulted"
);
Ok(())
}
#[test]
fn two_equally_close_options_are_ambiguous_rather_than_guessed()
-> Result<(), Box<dyn std::error::Error>> {
let (embedder, _calls) = stub(paraphrase_vectors(), vec![0.0, 1.0], false)?;
let resolver = SemanticResolver::new(embedder);
let resolution = resolver
.resolve(
"the usual",
&ask(&options(&["Repeat last order", "Repeat last order"])),
)
.into_resolution();
assert!(
matches!(&resolution, Resolution::Ambiguous { tied, .. } if *tied == vec![0, 1]),
"both options are still in play, got {resolution:?}"
);
Ok(())
}
#[test]
fn a_runner_up_below_threshold_still_counts_against_the_margin()
-> Result<(), Box<dyn std::error::Error>> {
let (embedder, _) = stub(near_tie_vectors(), vec![0.0, 1.0], false)?;
let verdict = resolve_with(embedder);
assert!(
matches!(verdict.resolution(), Resolution::Ambiguous { tied, .. } if *tied == vec![0, 1]),
"the near tie must be reported as one, got {verdict:?}"
);
Ok(())
}
#[test]
fn a_below_threshold_runner_up_is_reported_in_the_lead()
-> Result<(), Box<dyn std::error::Error>> {
let (embedder, _) = stub(near_tie_vectors(), vec![0.0, 1.0], false)?;
let policy = SemanticPolicy::new(0.72, 0.01)?;
let verdict = SemanticResolver::with_policy(embedder, policy)
.resolve("the usual", &ask(&the_field()));
assert!(
matches!(
verdict.resolution(),
Resolution::Resolved {
index: 0,
score,
lead,
..
} if (*score - 0.73).abs() < 1e-6 && (*lead - 0.02).abs() < 1e-6
),
"the lead must be the measured gap over the runner-up, got {verdict:?}"
);
Ok(())
}
#[test]
fn an_option_below_the_threshold_is_not_a_candidate() -> Result<(), Box<dyn std::error::Error>>
{
let (embedder, calls) = stub(paraphrase_vectors(), vec![0.0, 1.0], false)?;
let resolver = SemanticResolver::new(embedder);
let resolution = resolver
.resolve("the usual", &ask(&options(&["Cancel", "Later"])))
.into_resolution();
assert!(
matches!(resolution, Resolution::Absent { .. }),
"expected the semantic tier to stay silent, got {resolution:?}"
);
assert!(calls.get() > 0, "the tier did look, and found nothing");
Ok(())
}
#[test]
fn a_broken_competitor_cannot_manufacture_semantic_confidence()
-> Result<(), Box<dyn std::error::Error>> {
let (embedder, _) = stub(
vec![
("the usual", vec![1.0, 0.0]),
("Repeat last order", vec![1.0, 0.0]),
("Cancel", vec![0.0, 0.0]),
],
vec![0.0, 1.0],
false,
)?;
let verdict = resolve_with(embedder);
assert_eq!(
verdict.into_resolution(),
Resolution::Degraded {
reason: DegradedReason::UnmeasurableEmbedding,
},
"an unmeasured competitor cannot leave the field intact"
);
Ok(())
}
#[test]
fn a_degenerate_utterance_degrades_rather_than_finding_nothing()
-> Result<(), Box<dyn std::error::Error>> {
let options = options(&["Repeat last order", "Cancel"]);
let lexicon = crate::language::LanguageResolver::new();
for phrase in ["the usual", "a second phrasing"] {
assert!(
matches!(
lexicon.resolve(phrase, &ask(&options)).into_resolution(),
Resolution::Absent { .. }
),
"{phrase:?} must not be reachable by the lexicon, or this proves nothing"
);
}
for broken in [
vec![0.0, 0.0],
vec![f32::NAN, 0.0],
vec![f32::INFINITY, 0.0],
] {
let (embedder, _) = stub(
vec![
("the usual", broken.clone()),
("a second phrasing", vec![1.0, 0.1]),
("Repeat last order", vec![1.0, 0.1]),
("Cancel", vec![0.0, 1.0]),
],
vec![0.0, 1.0],
false,
)?;
let resolver = SemanticResolver::new(embedder);
let verdict = resolver.resolve("the usual", &ask(&options));
assert_eq!(
verdict.into_resolution(),
Resolution::Degraded {
reason: DegradedReason::UnmeasurableEmbedding,
},
"a {broken:?} utterance vector is not an absence of meaning"
);
let healthy = resolver.resolve("a second phrasing", &ask(&options));
assert!(
matches!(
healthy.resolution(),
Resolution::Resolved {
index: 0,
tier: MatchTier::Semantic,
score,
lead,
} if (*score - 1.0).abs() < 1e-6 && *lead > 0.8
),
"a valid utterance still resolves after a degraded one, got {healthy:?}"
);
}
Ok(())
}
#[test]
fn an_invalid_embedding_at_either_end_of_the_field_degrades()
-> Result<(), Box<dyn std::error::Error>> {
let cases = [
("Repeat last order", vec![0.0, 0.0]),
("Later", vec![f32::INFINITY, 0.0]),
];
for (broken_option, broken_vector) in cases {
let (embedder, _) = stub(
vec![
("the usual", vec![1.0, 0.0]),
("Repeat last order", vec![1.0, 0.1]),
("Cancel", vec![0.0, 1.0]),
("Later", vec![-1.0, 0.0]),
(broken_option, broken_vector),
],
vec![0.0, 1.0],
false,
)?;
let verdict = SemanticResolver::new(embedder).resolve(
"the usual",
&ask(&options(&["Repeat last order", "Cancel", "Later"])),
);
assert_eq!(
verdict.into_resolution(),
Resolution::Degraded {
reason: DegradedReason::UnmeasurableEmbedding,
},
"an invalid embedding for {broken_option} must stop the decision"
);
}
Ok(())
}
#[test]
fn a_field_of_unmeasurable_vectors_is_degraded_not_absent()
-> Result<(), Box<dyn std::error::Error>> {
let (embedder, _) = stub(
vec![
("the usual", vec![1.0, 0.0]),
("Repeat last order", vec![0.0, 0.0]),
("Cancel", vec![f32::NAN, 0.0]),
],
vec![0.0, 1.0],
false,
)?;
let verdict = resolve_with(embedder);
assert_eq!(
verdict.into_resolution(),
Resolution::Degraded {
reason: DegradedReason::UnmeasurableEmbedding,
},
"no comparison was made, so no best score exists"
);
Ok(())
}
#[test]
fn an_unmeasurable_embedding_is_named_apart_from_a_missing_embedder() {
assert_ne!(
DegradedReason::UnmeasurableEmbedding,
DegradedReason::EmbedderUnavailable,
"the two causes must not collapse into one verdict"
);
assert_ne!(
DegradedReason::UnmeasurableEmbedding.to_string(),
DegradedReason::EmbedderUnavailable.to_string(),
"and must not render identically in the transcript"
);
}
#[test]
fn a_failing_embedder_is_degraded_rather_than_absent() -> Result<(), Box<dyn std::error::Error>>
{
let (embedder, _calls) = stub(paraphrase_vectors(), vec![0.0, 1.0], true)?;
let resolution = resolve_with(embedder).into_resolution();
assert_eq!(
resolution,
Resolution::Degraded {
reason: DegradedReason::EmbedderUnavailable,
},
"a failed embedder is not a person who said something unclear"
);
Ok(())
}
#[test]
fn a_vector_of_the_wrong_length_is_degraded_rather_than_scored()
-> Result<(), Box<dyn std::error::Error>> {
let (embedder, _calls) = stub(
vec![("the usual", vec![1.0, 0.0])],
vec![0.0, 1.0, 2.0],
false,
)?;
let resolution = resolve_with(embedder).into_resolution();
assert_eq!(
resolution,
Resolution::Degraded {
reason: DegradedReason::EmbedderUnavailable,
},
"a provider that contradicts its own declared width is broken, not silent"
);
Ok(())
}
#[test]
fn a_learned_alias_moves_a_phrase_off_the_model() -> Result<(), Box<dyn std::error::Error>> {
let (embedder, calls) = stub(paraphrase_vectors(), vec![0.0, 1.0], false)?;
let mut resolver = SemanticResolver::new(embedder);
let options = options(&["Repeat last order", "Cancel"]);
assert_eq!(
resolver.learn("ask", "the usual", "Repeat last order"),
None,
"the first binding for a phrase has no predecessor"
);
let resolution = resolver
.resolve("the usual", &ask(&options))
.into_resolution();
assert!(
matches!(
resolution,
Resolution::Resolved {
index: 0,
tier: MatchTier::Exact,
score,
lead,
} if (score - 1.0).abs() < 1e-9 && (lead - score).abs() < 1e-9
),
"a confirmed phrase resolves exactly from then on. The lead is the whole \
score and that is the measurement, not a gap left unmeasured: an alias \
puts exactly one candidate in the exact tier, and the lead is held over \
the runner-up in the tier the question is answered in, so a lone exact \
winner has none to lead. A below-tier competitor does not count against \
it — see `the_lead_is_measured_against_the_highest_other_measured_score` \
for the measurement itself. Got {resolution:?}"
);
assert_eq!(
calls.get(),
0,
"and a confirmed correction leaves the model out of it entirely"
);
assert_eq!(resolver.learned(), 1);
assert_eq!(
resolver
.forget("ask", "the usual")
.as_ref()
.map(Alias::option),
Some("Repeat last order")
);
assert_eq!(resolver.learned(), 0);
Ok(())
}
#[test]
fn a_superseded_alias_is_not_handed_to_the_model() -> Result<(), Box<dyn std::error::Error>> {
let (embedder, calls) = stub(paraphrase_vectors(), vec![0.0, 1.0], false)?;
let mut resolver = SemanticResolver::new(embedder);
assert_eq!(
resolver.learn("ask", "the usual", "Repeat last order"),
None
);
let edited = options(&["Delete account", "Keep account"]);
let resolution = resolver
.resolve("the usual", &ask(&edited))
.into_resolution();
assert_eq!(
resolution,
Resolution::StaleAlias {
question: String::from("ask"),
option: String::from("Repeat last order"),
},
"the semantic tier must return a stale verdict verbatim"
);
assert_eq!(
calls.get(),
0,
"and must not spend the model trying to legitimate it"
);
Ok(())
}
#[test]
fn a_numeric_question_never_reaches_the_model() -> Result<(), Box<dyn std::error::Error>> {
let (embedder, calls) = stub(
vec![
("-5", vec![1.0, 0.0]),
("5", vec![1.0, 0.1]),
("10", vec![0.0, 1.0]),
],
vec![1.0, 0.0],
false,
)?;
let resolver = SemanticResolver::new(embedder);
let options = options(&["5", "10"]);
let question = ask(&options).with_domain(AnswerDomain::Integer);
assert_eq!(
resolver.resolve("-5", &question).into_resolution(),
Resolution::Absent { best_score: 0.0 },
"a value the question does not offer is absent, not the number it resembles"
);
assert_eq!(
calls.get(),
0,
"no tier may spend the model on relating two numbers"
);
assert_eq!(
resolver.resolve("5", &question).into_resolution(),
Resolution::Resolved {
index: 0,
tier: MatchTier::Exact,
score: 1.0,
lead: 1.0,
},
"and the values that are offered resolve exactly, by value"
);
assert_eq!(
calls.get(),
0,
"with the model still untouched after a successful numeric answer"
);
Ok(())
}
#[test]
fn identity_refuses_a_record_that_could_not_be_reproduced() {
assert_eq!(
EmbedderIdentity::new("", "digest", 4),
Err(SemanticError::UnnamedEmbedder)
);
assert_eq!(
EmbedderIdentity::new("model", "", 4),
Err(SemanticError::UndigestedEmbedder {
name: String::from("model")
})
);
assert_eq!(
EmbedderIdentity::new("model", "digest", 0),
Err(SemanticError::ZeroDimension {
name: String::from("model")
})
);
let identity = EmbedderIdentity::new("model", "digest", 4);
assert!(
identity.is_ok(),
"a named, digested, sized model is accepted"
);
}
#[test]
fn policy_refuses_a_threshold_outside_the_contract_interval() {
assert_eq!(
SemanticPolicy::new(f64::NAN, 0.05),
Err(SemanticError::NonFiniteThreshold)
);
assert_eq!(
SemanticPolicy::new(f64::INFINITY, 0.05),
Err(SemanticError::NonFiniteThreshold)
);
assert_eq!(
SemanticPolicy::new(1.5, 0.05),
Err(SemanticError::InvalidThreshold { threshold: 1.5 })
);
assert_eq!(
SemanticPolicy::new(0.7, f64::NAN),
Err(SemanticError::NonFiniteMargin)
);
assert_eq!(
SemanticPolicy::new(0.7, 2.0),
Err(SemanticError::InvalidMargin { margin: 2.0 })
);
assert!(SemanticPolicy::new(0.7, 0.05).is_ok());
assert!(
(SemanticPolicy::DEFAULT.threshold() - 0.72).abs() < 1e-12,
"the shipped threshold is the declared one"
);
assert_eq!(SemanticPolicy::default(), SemanticPolicy::DEFAULT);
}
#[test]
fn the_default_embed_batch_preserves_input_order() -> Result<(), Box<dyn std::error::Error>> {
let (embedder, calls) = stub(paraphrase_vectors(), vec![0.0, 1.0], false)?;
let batch = embedder
.embed_batch(&["Cancel", "the usual"])
.map_err(|_| "the stub was configured not to fail")?;
let expected = [vec![0.0_f32, 1.0], vec![1.0, 0.0]];
assert_eq!(batch.len(), expected.len(), "batch length changed");
for (actual, want) in batch.iter().zip(expected.iter()) {
assert_vector(actual, want);
}
assert_eq!(calls.get(), 2, "the default embeds one text at a time");
Ok(())
}
}