use panproto_gat::Name;
use panproto_schema::Schema;
use thiserror::Error;
use super::{Anchor, StrategyTag, kinds_compatible};
#[derive(Debug, Error)]
pub enum EmbedError {
#[error("embedder failure: {0}")]
Failure(String),
#[error("embedding dimension mismatch: expected {expected}, got {actual}")]
DimensionMismatch {
expected: usize,
actual: usize,
},
}
pub trait Embedder {
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedError>;
fn dim(&self) -> usize;
}
#[must_use]
pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f64 {
if a.len() != b.len() {
return 0.0;
}
let mut dot: f64 = 0.0;
let mut na: f64 = 0.0;
let mut nb: f64 = 0.0;
for i in 0..a.len() {
let x = f64::from(a[i]);
let y = f64::from(b[i]);
dot += x * y;
na += x * x;
nb += y * y;
}
if na == 0.0 || nb == 0.0 {
return 0.0;
}
let cos = dot / (na.sqrt() * nb.sqrt());
cos.clamp(0.0, 1.0)
}
fn embedding_text(schema: &Schema, vertex_id: &Name) -> String {
let mut out = vertex_id.as_str().to_owned();
if let Some(cs) = schema.constraints.get(vertex_id) {
for c in cs {
if c.sort.as_str() == "description" && !c.value.is_empty() {
out.push(' ');
out.push_str(&c.value);
break;
}
}
}
out
}
pub fn embedding_anchors<E: Embedder>(
src: &Schema,
tgt: &Schema,
embedder: &E,
threshold: f64,
) -> Result<Vec<Anchor>, EmbedError> {
let mut src_ids: Vec<&Name> = src.vertices.keys().collect();
src_ids.sort_by(|a, b| a.as_str().cmp(b.as_str()));
let mut tgt_ids: Vec<&Name> = tgt.vertices.keys().collect();
tgt_ids.sort_by(|a, b| a.as_str().cmp(b.as_str()));
let expected_dim = embedder.dim();
let mut src_vecs: Vec<(&Name, Vec<f32>)> = Vec::with_capacity(src_ids.len());
for id in &src_ids {
let text = embedding_text(src, id);
let vec = embedder.embed(&text)?;
if vec.len() != expected_dim {
return Err(EmbedError::DimensionMismatch {
expected: expected_dim,
actual: vec.len(),
});
}
src_vecs.push((id, vec));
}
let mut tgt_vecs: Vec<(&Name, Vec<f32>)> = Vec::with_capacity(tgt_ids.len());
for id in &tgt_ids {
let text = embedding_text(tgt, id);
let vec = embedder.embed(&text)?;
if vec.len() != expected_dim {
return Err(EmbedError::DimensionMismatch {
expected: expected_dim,
actual: vec.len(),
});
}
tgt_vecs.push((id, vec));
}
let mut out = Vec::new();
for (src_id, src_vec) in &src_vecs {
let mut best: Option<(&Name, f64)> = None;
for (tgt_id, tgt_vec) in &tgt_vecs {
if !kinds_compatible(src, src_id, tgt, tgt_id) {
continue;
}
let score = cosine_similarity(src_vec, tgt_vec);
if best.as_ref().is_none_or(|(_, bs)| score > *bs) {
best = Some((tgt_id, score));
}
}
if let Some((tgt_id, score)) = best
&& score >= threshold
{
out.push(Anchor {
src: (*src_id).clone(),
tgt: (*tgt_id).clone(),
confidence: score,
strategy: StrategyTag::Llm,
explanation: format!(
"embedding cosine {:.3}: {} ↔ {}",
score,
src_id.as_str(),
tgt_id.as_str()
),
});
}
}
Ok(out)
}
#[derive(Clone, Debug)]
pub struct HashEmbedder {
pub dim: usize,
}
impl HashEmbedder {
#[must_use]
pub const fn new(dim: usize) -> Self {
Self { dim }
}
}
impl Embedder for HashEmbedder {
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedError> {
if self.dim == 0 {
return Err(EmbedError::Failure("zero-dimensional embedder".into()));
}
let mut v = vec![0.0f32; self.dim];
for tok in super::token_similarity::tokenize(text) {
let h = blake3::hash(tok.as_bytes());
let bytes = h.as_bytes();
let mut buf = [0u8; 8];
buf.copy_from_slice(&bytes[..8]);
let seed = u64::from_le_bytes(buf);
let dim_u64 = u64::try_from(self.dim).unwrap_or(u64::MAX);
let bucket = usize::try_from(seed % dim_u64).unwrap_or(0);
v[bucket] += 1.0;
}
Ok(v)
}
fn dim(&self) -> usize {
self.dim
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::float_cmp)]
mod tests {
use super::*;
use panproto_schema::{Protocol, SchemaBuilder};
#[test]
fn cosine_identical_vectors_is_one() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![1.0, 2.0, 3.0];
assert!((cosine_similarity(&a, &b) - 1.0).abs() < 1e-9);
}
#[test]
fn cosine_orthogonal_is_zero() {
let a = vec![1.0, 0.0];
let b = vec![0.0, 1.0];
assert!(cosine_similarity(&a, &b).abs() < 1e-9);
}
#[test]
fn cosine_mismatched_length_returns_zero() {
let a = vec![1.0, 0.0];
let b = vec![1.0, 0.0, 0.0];
assert_eq!(cosine_similarity(&a, &b), 0.0);
}
#[test]
fn cosine_zero_vector_is_zero() {
let a = vec![0.0, 0.0];
let b = vec![1.0, 1.0];
assert_eq!(cosine_similarity(&a, &b), 0.0);
}
fn proto() -> Protocol {
Protocol {
name: "t".into(),
schema_theory: "ThTest".into(),
instance_theory: "ThWType".into(),
edge_rules: vec![],
obj_kinds: vec!["string".into()],
constraint_sorts: vec![],
..Protocol::default()
}
}
#[test]
fn hash_embedder_end_to_end() {
let p = proto();
let src = SchemaBuilder::new(&p)
.vertex("a_shared_token", "string", None::<&str>)
.unwrap()
.vertex("completely_different", "string", None::<&str>)
.unwrap()
.build()
.unwrap();
let tgt = SchemaBuilder::new(&p)
.vertex("shared_a_token", "string", None::<&str>)
.unwrap()
.vertex("utterly_alien_words", "string", None::<&str>)
.unwrap()
.build()
.unwrap();
let embedder = HashEmbedder::new(64);
let anchors = embedding_anchors(&src, &tgt, &embedder, 0.5).unwrap();
assert!(
anchors
.iter()
.any(|a| a.src.as_str() == "a_shared_token" && a.tgt.as_str() == "shared_a_token"),
"expected a_shared_token ↔ shared_a_token anchor in {:?}",
anchors
.iter()
.map(|a| (a.src.as_str(), a.tgt.as_str(), a.confidence))
.collect::<Vec<_>>()
);
for anchor in &anchors {
assert_eq!(anchor.strategy, StrategyTag::Llm);
assert!(anchor.confidence >= 0.5);
}
}
#[test]
fn hash_embedder_zero_dim_errors() {
let e = HashEmbedder::new(0);
assert!(e.embed("anything").is_err());
}
#[test]
fn hash_embedder_emits_requested_dim() {
let e = HashEmbedder::new(16);
let v = e.embed("hello").unwrap();
assert_eq!(v.len(), 16);
assert_eq!(e.dim(), 16);
}
}