use crate::model::capabilities::ModelCapabilities;
use anyhow::Result;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ModelRole {
Embedder,
Reranker,
Llm,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "lowercase")]
pub enum ModelSource {
Hf { repo: String, files: Vec<String> },
Local { dir: String },
Url { base: String, files: Vec<String> },
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Pooling {
Mean,
Cls,
LateInteraction,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum QuantKind {
None,
Int8Symmetric,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ScoreStrategyKind {
ClassifierHead,
YesNoLogit,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct EmbedderSpec {
pub dims: usize,
pub max_context: usize,
#[serde(default)]
pub query_prefix: String,
#[serde(default)]
pub doc_prefix: String,
pub pooling: Pooling,
#[serde(default = "default_true")]
pub normalize: bool,
pub quant: QuantKind,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct RerankerSpec {
pub score_strategy: ScoreStrategyKind,
#[serde(default = "default_reranker_max_context")]
pub max_context: usize,
#[serde(default)]
pub prompt_prefix: String,
#[serde(default)]
pub prompt_middle: String,
#[serde(default)]
pub prompt_suffix: String,
#[serde(default)]
pub yes_token_id: Option<usize>,
#[serde(default)]
pub no_token_id: Option<usize>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct LlmSpec {
#[serde(default)]
pub provider: String,
pub model: String,
#[serde(default)]
pub endpoint: String,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum RoleData {
Embedder(EmbedderSpec),
Reranker(RerankerSpec),
Llm(LlmSpec),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ModelSpec {
pub id: String,
pub role: ModelRole,
pub source: ModelSource,
#[serde(default)]
pub capabilities: ModelCapabilities,
#[serde(flatten)]
pub role_data: RoleData,
}
impl ModelSpec {
pub fn validate(&self) -> Result<()> {
anyhow::ensure!(!self.id.trim().is_empty(), "model spec has an empty `id`");
let role_ok = matches!(
(self.role, &self.role_data),
(ModelRole::Embedder, RoleData::Embedder(_))
| (ModelRole::Reranker, RoleData::Reranker(_))
| (ModelRole::Llm, RoleData::Llm(_))
);
anyhow::ensure!(
role_ok,
"model `{}`: `role` {:?} disagrees with its role data",
self.id,
self.role
);
match &self.source {
ModelSource::Hf { repo, files } => {
anyhow::ensure!(
!repo.trim().is_empty(),
"model `{}`: empty hf `repo`",
self.id
);
anyhow::ensure!(
!files.is_empty(),
"model `{}`: hf source lists no `files`",
self.id
);
}
ModelSource::Local { dir } => {
anyhow::ensure!(
!dir.trim().is_empty(),
"model `{}`: empty local `dir`",
self.id
);
}
ModelSource::Url { base, files } => {
anyhow::ensure!(
!base.trim().is_empty(),
"model `{}`: empty url `base`",
self.id
);
anyhow::ensure!(
!files.is_empty(),
"model `{}`: url source lists no `files`",
self.id
);
}
}
if let RoleData::Embedder(e) = &self.role_data {
anyhow::ensure!(
e.dims > 0,
"model `{}`: embedder `dims` must be > 0",
self.id
);
anyhow::ensure!(
e.max_context > 0,
"model `{}`: embedder `max_context` must be > 0",
self.id
);
}
if let RoleData::Reranker(r) = &self.role_data
&& matches!(r.score_strategy, ScoreStrategyKind::YesNoLogit)
{
anyhow::ensure!(
r.yes_token_id.is_some() && r.no_token_id.is_some(),
"model `{}`: yes_no_logit reranker needs both `yes_token_id` and `no_token_id`",
self.id
);
}
Ok(())
}
}
pub struct EmbedderFingerprint;
impl EmbedderFingerprint {
#[must_use]
pub fn compute(id: &str, spec: &EmbedderSpec, dense_context: bool) -> String {
let pooling = match spec.pooling {
Pooling::Mean => "mean",
Pooling::Cls => "cls",
Pooling::LateInteraction => "late_interaction",
};
let quant = match spec.quant {
QuantKind::None => "none",
QuantKind::Int8Symmetric => "int8_symmetric",
};
let ctx = u8::from(dense_context);
let canonical = format!(
"id={id};dims={};pooling={pooling};quant={quant};norm={};qpre={};dpre={};ctx={ctx}",
spec.dims, spec.normalize, spec.query_prefix, spec.doc_prefix
);
let hash = xxhash_rust::xxh64::xxh64(canonical.as_bytes(), 0);
format!("{hash:016x}")
}
}
fn default_true() -> bool {
true
}
fn default_reranker_max_context() -> usize {
512
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn role_round_trips_through_toml() {
for (role, s) in [
(ModelRole::Embedder, "embedder"),
(ModelRole::Reranker, "reranker"),
(ModelRole::Llm, "llm"),
] {
let doc = format!("role = \"{s}\"\n");
let parsed: RoleHolder = toml::from_str(&doc).unwrap();
assert_eq!(parsed.role, role);
}
}
#[test]
fn source_hf_round_trips() {
let doc = r#"
kind = "hf"
repo = "owner/Model-onnx-int8"
files = ["model_int8.onnx", "tokenizer.json"]
"#;
let src: ModelSource = toml::from_str(doc).unwrap();
match src {
ModelSource::Hf { repo, files } => {
assert_eq!(repo, "owner/Model-onnx-int8");
assert_eq!(files, vec!["model_int8.onnx", "tokenizer.json"]);
}
other => panic!("expected Hf, got {other:?}"),
}
}
#[test]
fn embedder_spec_carries_every_nuance() {
let e = EmbedderSpec {
dims: 768,
max_context: 8192,
query_prefix: "Represent this query for searching relevant code: ".to_string(),
doc_prefix: String::new(),
pooling: Pooling::Cls,
normalize: true,
quant: QuantKind::Int8Symmetric,
};
assert_eq!(e.dims, 768);
assert_eq!(e.pooling, Pooling::Cls);
assert_eq!(e.quant, QuantKind::Int8Symmetric);
assert!(e.query_prefix.ends_with(' '));
}
#[test]
fn fingerprint_is_stable_and_sensitive() {
let base = EmbedderSpec {
dims: 768,
max_context: 8192,
query_prefix: "Q: ".to_string(),
doc_prefix: String::new(),
pooling: Pooling::Cls,
normalize: true,
quant: QuantKind::Int8Symmetric,
};
let fp1 = EmbedderFingerprint::compute("coderank-137m", &base, false);
let fp2 = EmbedderFingerprint::compute("coderank-137m", &base, false);
assert_eq!(fp1, fp2, "same id+spec → same fingerprint (deterministic)");
assert!(!fp1.is_empty());
let mut diff_dims = base.clone();
diff_dims.dims = 384;
assert_ne!(
fp1,
EmbedderFingerprint::compute("coderank-137m", &diff_dims, false)
);
let mut diff_pool = base.clone();
diff_pool.pooling = Pooling::Mean;
assert_ne!(
fp1,
EmbedderFingerprint::compute("coderank-137m", &diff_pool, false)
);
let mut diff_quant = base.clone();
diff_quant.quant = QuantKind::None;
assert_ne!(
fp1,
EmbedderFingerprint::compute("coderank-137m", &diff_quant, false)
);
assert_ne!(
fp1,
EmbedderFingerprint::compute("other-embedder", &base, false)
);
}
#[test]
fn fingerprint_is_sensitive_to_dense_context() {
let base = EmbedderSpec {
dims: 768,
max_context: 8192,
query_prefix: "Q: ".to_string(),
doc_prefix: String::new(),
pooling: Pooling::Cls,
normalize: true,
quant: QuantKind::Int8Symmetric,
};
let off = EmbedderFingerprint::compute("coderank-137m", &base, false);
let on = EmbedderFingerprint::compute("coderank-137m", &base, true);
assert_ne!(
off, on,
"dense_context must be part of the fingerprint (it changes embedded text)"
);
assert_eq!(
on,
EmbedderFingerprint::compute("coderank-137m", &base, true)
);
}
#[test]
fn fingerprint_is_sensitive_to_doc_prefix() {
let base = EmbedderSpec {
dims: 768,
max_context: 8192,
query_prefix: "Q: ".to_string(),
doc_prefix: String::new(),
pooling: Pooling::Cls,
normalize: true,
quant: QuantKind::Int8Symmetric,
};
let mut diff_doc = base.clone();
diff_doc.doc_prefix = "passage: ".to_string();
assert_ne!(
EmbedderFingerprint::compute("coderank-137m", &base, false),
EmbedderFingerprint::compute("coderank-137m", &diff_doc, false),
"doc_prefix must be part of the fingerprint"
);
}
#[derive(Deserialize)]
struct RoleHolder {
role: ModelRole,
}
}