pub mod pipeline;
#[cfg(feature = "fastembed-reranker")]
pub mod fastembed_reranker;
pub use pipeline::{DEFAULT_MIN_CANDIDATES, DEFAULT_TOP_K_RERANK, rerank_step};
#[cfg(feature = "fastembed-reranker")]
pub use fastembed_reranker::FastEmbedReranker;
use std::fmt;
use std::fs;
use std::path::{Path, PathBuf};
use asupersync::Cx;
use asupersync::sync::{LockError, Mutex};
use ort::session::Session;
use serde::Deserialize;
use tokenizers::Tokenizer;
use tracing::instrument;
use frankensearch_core::error::{SearchError, SearchResult};
use frankensearch_core::traits::{RerankDocument, RerankScore, Reranker, SearchFuture};
pub const DEFAULT_MODEL_NAME: &str = "flashrank";
pub const DEFAULT_MAX_LENGTH: usize = 512;
const INFERENCE_BATCH_SIZE: usize = 32;
const MODEL_ONNX_SUBDIR: &str = "onnx/model.onnx";
const MODEL_ONNX_LEGACY: &str = "model.onnx";
const TOKENIZER_JSON: &str = "tokenizer.json";
const CONFIG_JSON: &str = "config.json";
const TOKENIZER_CONFIG_JSON: &str = "tokenizer_config.json";
const OUTPUT_TENSOR_CANDIDATES: [&str; 3] = ["logits", "output", "sentence_embedding"];
#[derive(Debug, Deserialize)]
struct ModelConfig {
#[serde(default)]
pad_token_id: u32,
}
#[derive(Debug, Deserialize)]
struct TokenizerConfig {
model_max_length: Option<usize>,
}
pub struct FlashRankReranker {
session: Mutex<Session>,
tokenizer: Tokenizer,
input_names: Vec<String>,
max_length: usize,
name: String,
model_dir: PathBuf,
}
impl fmt::Debug for FlashRankReranker {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("FlashRankReranker")
.field("name", &self.name)
.field("input_names", &self.input_names)
.field("max_length", &self.max_length)
.field("model_dir", &self.model_dir)
.finish_non_exhaustive()
}
}
impl FlashRankReranker {
#[instrument(skip_all, fields(model_dir = %model_dir.as_ref().display()))]
pub fn load(model_dir: impl AsRef<Path>) -> SearchResult<Self> {
Self::load_with_config(model_dir, DEFAULT_MODEL_NAME, DEFAULT_MAX_LENGTH)
}
pub fn load_with_config(
model_dir: impl AsRef<Path>,
name: &str,
max_length: usize,
) -> SearchResult<Self> {
let model_dir = resolve_model_dir(model_dir.as_ref(), name)?;
let model_file =
select_model_file(&model_dir).ok_or_else(|| SearchError::ModelNotFound {
name: format!("{name} (missing {MODEL_ONNX_SUBDIR} or {MODEL_ONNX_LEGACY})"),
})?;
let tokenizer_path = model_dir.join(TOKENIZER_JSON);
if !tokenizer_path.is_file() {
return Err(SearchError::ModelNotFound {
name: format!(
"{name} (missing {TOKENIZER_JSON} in {})",
model_dir.display()
),
});
}
let config_path = model_dir.join(CONFIG_JSON);
let pad_token_id = if config_path.is_file() {
let content =
fs::read_to_string(&config_path).map_err(|e| SearchError::ModelLoadFailed {
path: config_path.clone(),
source: format!("failed to read config.json: {e}").into(),
})?;
let config: ModelConfig =
serde_json::from_str(&content).map_err(|e| SearchError::ModelLoadFailed {
path: config_path,
source: format!("failed to parse config.json: {e}").into(),
})?;
config.pad_token_id
} else {
0
};
let tokenizer_config_path = model_dir.join(TOKENIZER_CONFIG_JSON);
let model_max_len = if tokenizer_config_path.is_file() {
let content = fs::read_to_string(&tokenizer_config_path).map_err(|e| {
SearchError::ModelLoadFailed {
path: tokenizer_config_path.clone(),
source: format!("failed to read tokenizer_config.json: {e}").into(),
}
})?;
let config: TokenizerConfig =
serde_json::from_str(&content).map_err(|e| SearchError::ModelLoadFailed {
path: tokenizer_config_path,
source: format!("failed to parse tokenizer_config.json: {e}").into(),
})?;
config.model_max_length.unwrap_or(max_length)
} else {
max_length
};
let session_builder = Session::builder().map_err(|e| SearchError::ModelLoadFailed {
path: model_file.clone(),
source: format!("ONNX session creation failed: {e}").into(),
})?;
let session_builder = session_builder
.with_optimization_level(ort::session::builder::GraphOptimizationLevel::Level3)
.map_err(|e| SearchError::ModelLoadFailed {
path: model_file.clone(),
source: format!("ONNX session creation failed: {e}").into(),
})?;
let session_builder = session_builder
.with_intra_threads(num_cpus())
.map_err(|e| SearchError::ModelLoadFailed {
path: model_file.clone(),
source: format!("ONNX session creation failed: {e}").into(),
})?;
let session = session_builder.commit_from_file(&model_file).map_err(|e| {
SearchError::ModelLoadFailed {
path: model_file.clone(),
source: format!("ONNX session creation failed: {e}").into(),
}
})?;
let input_names: Vec<String> = session
.inputs()
.iter()
.map(|i| i.name().to_owned())
.collect();
let mut tokenizer =
Tokenizer::from_file(&tokenizer_path).map_err(|e| SearchError::ModelLoadFailed {
path: tokenizer_path,
source: format!("tokenizer load failed: {e}").into(),
})?;
let final_max_len = model_max_len.min(max_length);
tokenizer
.with_truncation(Some(tokenizers::TruncationParams {
max_length: final_max_len,
..Default::default()
}))
.map_err(|e| SearchError::ModelLoadFailed {
path: model_dir.join(TOKENIZER_JSON),
source: format!("failed to enable truncation: {e}").into(),
})?;
tokenizer.with_padding(Some(tokenizers::PaddingParams {
pad_id: pad_token_id,
pad_token: "[PAD]".to_owned(),
..Default::default()
}));
tracing::info!(
model = %name,
input_names = ?input_names,
max_length = final_max_len,
pad_token_id,
model_dir = %model_dir.display(),
"FlashRank reranker model loaded"
);
Ok(Self {
session: Mutex::new(session),
tokenizer,
input_names,
max_length: final_max_len,
name: name.to_owned(),
model_dir,
})
}
fn tokenize_pair(&self, query: &str, document: &str) -> SearchResult<TokenizedPair> {
let encoding = self
.tokenizer
.encode((query, document), true)
.map_err(|e| SearchError::RerankFailed {
model: self.name.clone(),
source: format!("tokenization failed: {e}").into(),
})?;
let mut input_ids: Vec<i64> = encoding.get_ids().iter().map(|&id| i64::from(id)).collect();
let mut attention_mask: Vec<i64> = encoding
.get_attention_mask()
.iter()
.map(|&m| i64::from(m))
.collect();
let mut token_type_ids: Vec<i64> = encoding
.get_type_ids()
.iter()
.map(|&t| i64::from(t))
.collect();
if input_ids.len() > self.max_length {
input_ids.truncate(self.max_length);
attention_mask.truncate(self.max_length);
token_type_ids.truncate(self.max_length);
}
Ok(TokenizedPair {
input_ids,
attention_mask,
token_type_ids,
})
}
#[allow(clippy::cast_possible_wrap)]
fn infer_batch(
session: &mut Session,
input_names: &[String],
pairs: &[TokenizedPair],
model_name: &str,
) -> SearchResult<Vec<(f32, Option<f32>)>> {
if pairs.is_empty() {
return Ok(Vec::new());
}
let batch_size = pairs.len();
let seq_len = pairs.iter().map(|p| p.input_ids.len()).max().unwrap_or(0);
let flat_len =
batch_size
.checked_mul(seq_len)
.ok_or_else(|| SearchError::RerankFailed {
model: model_name.to_owned(),
source: format!("tensor size overflow: {batch_size} * {seq_len}").into(),
})?;
let mut flat_input_ids = vec![0_i64; flat_len];
let mut flat_attention_mask = vec![0_i64; flat_len];
let mut flat_token_type_ids = vec![0_i64; flat_len];
for (i, pair) in pairs.iter().enumerate() {
let offset = i * seq_len;
flat_input_ids[offset..offset + pair.input_ids.len()].copy_from_slice(&pair.input_ids);
flat_attention_mask[offset..offset + pair.attention_mask.len()]
.copy_from_slice(&pair.attention_mask);
flat_token_type_ids[offset..offset + pair.token_type_ids.len()]
.copy_from_slice(&pair.token_type_ids);
}
let batch_i64 = i64::try_from(batch_size).map_err(|_| SearchError::RerankFailed {
model: model_name.to_owned(),
source: format!("batch_size {batch_size} exceeds i64::MAX").into(),
})?;
let seq_i64 = i64::try_from(seq_len).map_err(|_| SearchError::RerankFailed {
model: model_name.to_owned(),
source: format!("seq_len {seq_len} exceeds i64::MAX").into(),
})?;
let shape = [batch_i64, seq_i64];
let mut inputs: Vec<(&str, ort::value::Value<ort::value::TensorValueType<i64>>)> =
Vec::with_capacity(3);
if input_names.iter().any(|n| n == "input_ids") {
let tensor = ort::value::Tensor::from_array((shape, flat_input_ids))
.map_err(|e| rerank_ort_error(model_name, "input_ids tensor", &e))?;
inputs.push(("input_ids", tensor));
}
if input_names.iter().any(|n| n == "attention_mask") {
let tensor = ort::value::Tensor::from_array((shape, flat_attention_mask))
.map_err(|e| rerank_ort_error(model_name, "attention_mask tensor", &e))?;
inputs.push(("attention_mask", tensor));
}
if input_names.iter().any(|n| n == "token_type_ids") {
let tensor = ort::value::Tensor::from_array((shape, flat_token_type_ids))
.map_err(|e| rerank_ort_error(model_name, "token_type_ids tensor", &e))?;
inputs.push(("token_type_ids", tensor));
}
let outputs = session.run(inputs).map_err(|e| SearchError::RerankFailed {
model: model_name.to_owned(),
source: format!("ONNX inference failed: {e}").into(),
})?;
let scores = extract_scores(&outputs, model_name, batch_size)?;
Ok(scores)
}
#[must_use]
pub fn model_dir(&self) -> &Path {
&self.model_dir
}
}
impl Reranker for FlashRankReranker {
fn rerank<'a>(
&'a self,
cx: &'a Cx,
query: &'a str,
documents: &'a [RerankDocument],
) -> SearchFuture<'a, Vec<RerankScore>> {
Box::pin(async move {
if documents.is_empty() {
return Ok(Vec::new());
}
let mut all_pairs = Vec::with_capacity(documents.len());
for doc in documents {
all_pairs.push(self.tokenize_pair(query, &doc.text)?);
}
let mut session = self
.session
.lock(cx)
.await
.map_err(|err| map_lock_error(&self.name, err))?;
let mut all_scores = Vec::with_capacity(documents.len());
for chunk in all_pairs.chunks(INFERENCE_BATCH_SIZE) {
let batch_scores =
Self::infer_batch(&mut session, &self.input_names, chunk, &self.name)?;
all_scores.extend(batch_scores);
}
if all_scores.len() != documents.len() {
return Err(SearchError::RerankFailed {
model: self.name.clone(),
source: format!(
"inference returned {} scores for {} documents",
all_scores.len(),
documents.len()
)
.into(),
});
}
let mut results = Vec::with_capacity(documents.len());
for (rank, doc) in documents.iter().enumerate() {
let Some(&(score, raw_logit)) = all_scores.get(rank) else {
return Err(SearchError::RerankFailed {
model: self.name.clone(),
source: format!(
"missing score at rank {rank} ({} scores for {} documents)",
all_scores.len(),
documents.len()
)
.into(),
});
};
results.push(RerankScore {
doc_id: doc.doc_id.clone(),
score,
original_rank: rank,
raw_logit,
});
}
results.sort_by(|a, b| {
sanitize_score(b.score)
.total_cmp(&sanitize_score(a.score))
.then_with(|| a.doc_id.cmp(&b.doc_id))
});
Ok(results)
})
}
fn id(&self) -> &str {
&self.name
}
fn model_name(&self) -> &str {
&self.name
}
fn max_length(&self) -> usize {
self.max_length
}
}
struct TokenizedPair {
input_ids: Vec<i64>,
attention_mask: Vec<i64>,
token_type_ids: Vec<i64>,
}
#[inline]
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
#[inline]
const fn sanitize_score(score: f32) -> f32 {
if score.is_finite() {
score
} else {
f32::NEG_INFINITY
}
}
#[allow(clippy::cast_sign_loss)]
fn extract_scores(
outputs: &ort::session::SessionOutputs<'_>,
model_name: &str,
batch_size: usize,
) -> SearchResult<Vec<(f32, Option<f32>)>> {
for name in &OUTPUT_TENSOR_CANDIDATES {
if let Some(value) = outputs.get(*name)
&& let Ok((shape, data)) = value.try_extract_tensor::<f32>()
{
return extract_scores_from_raw((&**shape, data), batch_size, model_name);
}
}
if let Some((_, value)) = outputs.iter().next()
&& let Ok((shape, data)) = value.try_extract_tensor::<f32>()
{
return extract_scores_from_raw((&**shape, data), batch_size, model_name);
}
Err(SearchError::RerankFailed {
model: model_name.to_owned(),
source: "no extractable output tensor found in ONNX session output".into(),
})
}
fn extract_scores_from_raw(
(shape, data): (&[i64], &[f32]),
batch_size: usize,
model_name: &str,
) -> SearchResult<Vec<(f32, Option<f32>)>> {
let dim_eq = |n: &i64| usize::try_from(*n).is_ok_and(|v| v == batch_size);
match shape {
[n, 1] if dim_eq(n) && data.len() >= batch_size => Ok(data[..batch_size]
.iter()
.map(|&logit| (sigmoid(logit), Some(logit)))
.collect()),
[n, 2] if dim_eq(n) && data.len() == batch_size * 2 => Ok(data
.chunks_exact(2)
.map(|pair| {
let (l0, l1) = (pair[0], pair[1]);
let delta = l1 - l0;
(sigmoid(delta), Some(delta))
})
.collect()),
[n] if dim_eq(n) => Ok(data
.iter()
.map(|&logit| (sigmoid(logit), Some(logit)))
.collect()),
[1, n] if dim_eq(n) => Ok(data
.iter()
.map(|&logit| (sigmoid(logit), Some(logit)))
.collect()),
_ => {
if data.len() >= batch_size {
Ok(data[..batch_size]
.iter()
.map(|&logit| (sigmoid(logit), Some(logit)))
.collect())
} else {
Err(SearchError::RerankFailed {
model: model_name.to_owned(),
source: format!(
"unexpected output tensor shape {shape:?} for batch size {batch_size}"
)
.into(),
})
}
}
}
}
fn map_lock_error(model: &str, error: LockError) -> SearchError {
match error {
LockError::Cancelled => SearchError::Cancelled {
phase: "rerank".to_owned(),
reason: "mutex lock cancelled".to_owned(),
},
LockError::Poisoned => SearchError::RerankFailed {
model: model.to_owned(),
source: "reranker mutex poisoned".into(),
},
other => SearchError::RerankFailed {
model: model.to_owned(),
source: std::io::Error::other(format!("reranker mutex lock failed: {other}")).into(),
},
}
}
#[cfg(test)]
mod score_tests {
use super::*;
fn assert_close(actual: f32, expected: f32) {
let delta = (actual - expected).abs();
assert!(
delta <= 1e-4,
"expected {expected}, got {actual} (delta {delta})"
);
}
#[test]
fn extract_scores_batch_one_column_preserves_logits() {
let shape = [2_i64, 1];
let data = [0.0_f32, 1.0];
let scores = extract_scores_from_raw((&shape, &data), 2, "test").unwrap();
assert_eq!(scores.len(), 2);
assert_close(scores[0].0, sigmoid(0.0));
assert_close(scores[1].0, sigmoid(1.0));
assert_eq!(scores[0].1, Some(0.0));
assert_eq!(scores[1].1, Some(1.0));
}
#[test]
fn extract_scores_batch_two_column_uses_logit_delta() {
let shape = [2_i64, 2];
let data = [1.0_f32, 2.0, 0.5, -0.5];
let scores = extract_scores_from_raw((&shape, &data), 2, "test").unwrap();
assert_eq!(scores.len(), 2);
let delta0 = 2.0 - 1.0;
let delta1 = -0.5 - 0.5;
assert_close(scores[0].0, sigmoid(delta0));
assert_close(scores[1].0, sigmoid(delta1));
assert_eq!(scores[0].1, Some(delta0));
assert_eq!(scores[1].1, Some(delta1));
}
#[test]
fn extract_scores_flat_shape_preserves_logits() {
let shape = [2_i64];
let data = [0.25_f32, -0.75];
let scores = extract_scores_from_raw((&shape, &data), 2, "test").unwrap();
assert_eq!(scores.len(), 2);
assert_close(scores[0].0, sigmoid(0.25));
assert_close(scores[1].0, sigmoid(-0.75));
assert_eq!(scores[0].1, Some(0.25));
assert_eq!(scores[1].1, Some(-0.75));
}
#[test]
fn extract_scores_transposed_shape_preserves_logits() {
let shape = [1_i64, 2];
let data = [0.8_f32, -0.2];
let scores = extract_scores_from_raw((&shape, &data), 2, "test").unwrap();
assert_eq!(scores.len(), 2);
assert_close(scores[0].0, sigmoid(0.8));
assert_close(scores[1].0, sigmoid(-0.2));
assert_eq!(scores[0].1, Some(0.8));
assert_eq!(scores[1].1, Some(-0.2));
}
}
fn rerank_ort_error(model: &str, context: &str, error: &ort::Error) -> SearchError {
SearchError::RerankFailed {
model: model.to_owned(),
source: format!("{context}: {error}").into(),
}
}
fn resolve_model_dir(base_dir: &Path, model_name: &str) -> SearchResult<PathBuf> {
if has_required_files(base_dir) {
return Ok(base_dir.to_path_buf());
}
if model_name.contains("..") || model_name.starts_with('/') || model_name.starts_with('\\') {
return Err(SearchError::ModelNotFound {
name: format!("{model_name} (unsafe model name)"),
});
}
let nested = base_dir.join(model_name);
if has_required_files(&nested) {
return Ok(nested);
}
Err(SearchError::ModelNotFound {
name: format!(
"{model_name} (missing required files in {} or {})",
base_dir.display(),
nested.display()
),
})
}
fn select_model_file(model_dir: &Path) -> Option<PathBuf> {
let modern = model_dir.join(MODEL_ONNX_SUBDIR);
if modern.is_file() {
return Some(modern);
}
let legacy = model_dir.join(MODEL_ONNX_LEGACY);
if legacy.is_file() {
return Some(legacy);
}
None
}
fn has_required_files(dir: &Path) -> bool {
select_model_file(dir).is_some() && dir.join(TOKENIZER_JSON).is_file()
}
#[must_use]
pub fn find_model_dir(model_name: &str) -> Option<PathBuf> {
if model_name.contains("..") || model_name.starts_with('/') || model_name.starts_with('\\') {
return None;
}
let mut candidates = Vec::new();
if let Ok(dir) = std::env::var("FRANKENSEARCH_MODEL_DIR") {
let base = PathBuf::from(dir);
candidates.push(base.join(model_name));
candidates.push(base);
}
if let Some(cache_dir) = dirs::cache_dir() {
candidates.push(cache_dir.join("frankensearch/models").join(model_name));
candidates.push(cache_dir.join("flashrank").join(model_name));
}
if let Some(data_dir) = dirs::data_local_dir() {
candidates.push(data_dir.join("frankensearch/models").join(model_name));
}
candidates.into_iter().find(|dir| has_required_files(dir))
}
fn num_cpus() -> usize {
std::thread::available_parallelism().map_or(4, |n| n.get().min(8))
}
#[cfg(test)]
mod tests {
use std::fs;
use super::*;
#[test]
fn sigmoid_zero_is_half() {
assert!((sigmoid(0.0) - 0.5).abs() < 1e-6);
}
#[test]
fn sigmoid_large_positive_approaches_one() {
assert!(sigmoid(10.0) > 0.99);
assert!(sigmoid(100.0) > 0.999);
}
#[test]
fn sigmoid_large_negative_approaches_zero() {
assert!(sigmoid(-10.0) < 0.01);
assert!(sigmoid(-100.0) < 0.001);
}
#[test]
fn sigmoid_output_range() {
for x in [-100.0, -10.0, -1.0, 0.0, 1.0, 10.0, 100.0] {
let s = sigmoid(x);
assert!((0.0..=1.0).contains(&s), "sigmoid({x}) = {s} out of [0,1]");
}
}
#[test]
fn sigmoid_symmetry() {
for x in [0.5, 1.0, 2.5, 5.0] {
let sum = sigmoid(x) + sigmoid(-x);
assert!(
(sum - 1.0).abs() < 1e-6,
"sigmoid({x}) + sigmoid(-{x}) = {sum}"
);
}
}
#[test]
fn sanitize_finite_scores() {
assert!((sanitize_score(0.5) - 0.5).abs() < f32::EPSILON);
assert!((sanitize_score(1.0) - 1.0).abs() < f32::EPSILON);
assert!((sanitize_score(-1.0) - (-1.0)).abs() < f32::EPSILON);
}
#[test]
fn sanitize_nan_to_neg_infinity() {
assert_eq!(
sanitize_score(f32::NAN).to_bits(),
f32::NEG_INFINITY.to_bits()
);
}
#[test]
fn sanitize_infinity_to_neg_infinity() {
assert_eq!(
sanitize_score(f32::INFINITY).to_bits(),
f32::NEG_INFINITY.to_bits()
);
assert_eq!(
sanitize_score(f32::NEG_INFINITY).to_bits(),
f32::NEG_INFINITY.to_bits()
);
}
#[test]
fn has_required_files_with_modern_layout() {
let temp = tempfile::tempdir().unwrap();
create_stub_model(temp.path(), true);
assert!(has_required_files(temp.path()));
}
#[test]
fn has_required_files_with_legacy_layout() {
let temp = tempfile::tempdir().unwrap();
create_stub_model(temp.path(), false);
assert!(has_required_files(temp.path()));
}
#[test]
fn has_required_files_missing_tokenizer() {
let temp = tempfile::tempdir().unwrap();
fs::create_dir_all(temp.path().join("onnx")).unwrap();
fs::write(temp.path().join("onnx/model.onnx"), b"stub").unwrap();
assert!(!has_required_files(temp.path()));
}
#[test]
fn resolve_model_dir_direct_path() {
let temp = tempfile::tempdir().unwrap();
create_stub_model(temp.path(), true);
let resolved = resolve_model_dir(temp.path(), "flashrank").unwrap();
assert_eq!(resolved, temp.path());
}
#[test]
fn resolve_model_dir_nested_path() {
let temp = tempfile::tempdir().unwrap();
let child = temp.path().join("flashrank");
fs::create_dir_all(&child).unwrap();
create_stub_model(&child, true);
let resolved = resolve_model_dir(temp.path(), "flashrank").unwrap();
assert_eq!(resolved, child);
}
#[test]
fn resolve_model_dir_missing() {
let temp = tempfile::tempdir().unwrap();
let err = resolve_model_dir(temp.path(), "flashrank").unwrap_err();
assert!(matches!(err, SearchError::ModelNotFound { .. }));
}
#[test]
fn select_model_file_prefers_modern() {
let temp = tempfile::tempdir().unwrap();
create_stub_model(temp.path(), true);
fs::write(temp.path().join("model.onnx"), b"legacy").unwrap();
let selected = select_model_file(temp.path()).unwrap();
assert!(selected.ends_with(MODEL_ONNX_SUBDIR));
}
#[test]
fn lock_cancelled_maps_to_search_cancelled() {
let err = map_lock_error("flashrank", LockError::Cancelled);
assert!(matches!(err, SearchError::Cancelled { .. }));
}
#[test]
fn lock_poisoned_maps_to_rerank_failed() {
let err = map_lock_error("flashrank", LockError::Poisoned);
assert!(matches!(err, SearchError::RerankFailed { .. }));
}
#[test]
fn lock_polled_after_completion_maps_to_rerank_failed() {
let err = map_lock_error("flashrank", LockError::PolledAfterCompletion);
match err {
SearchError::RerankFailed { source, .. } => {
assert!(
source
.to_string()
.contains("future reused after completion")
);
}
other => panic!("expected rerank failure, got {other:?}"),
}
}
#[test]
fn num_cpus_returns_reasonable_value() {
let n = num_cpus();
assert!((1..=8).contains(&n));
}
#[test]
fn find_model_dir_rejects_dotdot_traversal() {
assert!(find_model_dir("../../etc").is_none());
assert!(find_model_dir("foo/../bar").is_none());
}
#[test]
fn find_model_dir_rejects_absolute_path() {
assert!(find_model_dir("/etc/passwd").is_none());
}
#[test]
fn find_model_dir_rejects_backslash_prefix() {
assert!(find_model_dir("\\Windows\\System32").is_none());
}
fn create_stub_model(dir: &Path, use_onnx_subdir: bool) {
if use_onnx_subdir {
fs::create_dir_all(dir.join("onnx")).unwrap();
fs::write(dir.join("onnx/model.onnx"), b"stub-onnx").unwrap();
} else {
fs::write(dir.join("model.onnx"), b"stub-onnx").unwrap();
}
fs::write(dir.join("tokenizer.json"), "{}").unwrap();
}
}