use super::batch::{adaptive_batch_for_dim, entity_cache_key, entity_embed_cache};
use super::*;
use std::sync::Arc;
#[test]
fn f32_to_bytes_roundtrip() {
let input = vec![0.0_f32, 1.5, -2.25, f32::MIN, f32::MAX];
let bytes = f32_to_bytes(&input);
assert_eq!(bytes.len(), input.len() * 4);
let out = bytes_to_f32(&bytes);
assert_eq!(out, input);
}
#[test]
fn embedding_dim_matches_constants_source() {
assert_eq!(embedding_dim(), crate::constants::embedding_dim());
}
#[test]
fn effective_permits_clamps_to_bounds() {
assert!(effective_permits(0) >= 1);
assert!(effective_permits(1000) <= 32);
}
#[test]
fn adaptive_batch_dim64_keeps_calibrated_sizes() {
assert_eq!(adaptive_batch_for_dim(CHUNK_EMBED_BATCH_SIZE, 64), 8);
assert_eq!(adaptive_batch_for_dim(ENTITY_EMBED_BATCH_SIZE, 64), 25);
}
#[test]
fn adaptive_batch_dim384_shrinks() {
assert_eq!(adaptive_batch_for_dim(CHUNK_EMBED_BATCH_SIZE, 384), 1);
assert_eq!(adaptive_batch_for_dim(ENTITY_EMBED_BATCH_SIZE, 384), 4);
}
#[test]
fn adaptive_batch_intermediate_dims() {
assert_eq!(adaptive_batch_for_dim(8, 128), 4);
assert_eq!(adaptive_batch_for_dim(8, 256), 2);
}
#[test]
fn adaptive_batch_small_dim_clamps_to_base() {
assert_eq!(adaptive_batch_for_dim(8, 8), 8);
}
#[test]
fn adaptive_batch_total_function() {
assert_eq!(adaptive_batch_for_dim(8, 4096), 1);
assert_eq!(adaptive_batch_for_dim(8, 0), 8);
assert_eq!(adaptive_batch_for_dim(0, 64), 1);
}
#[test]
#[serial_test::serial(env)]
fn adaptive_wrappers_follow_active_dim() {
crate::constants::set_active_embedding_dim(384);
let chunk = chunk_embed_batch_size();
let entity = entity_embed_batch_size();
crate::constants::set_active_embedding_dim(crate::constants::DEFAULT_EMBEDDING_DIM);
assert_eq!(chunk, 1, "384-dim chunk batch must shrink to 1 (G44)");
assert_eq!(entity, 4, "384-dim entity batch must shrink to 4 (G44)");
}
#[test]
#[serial_test::serial(env)]
fn retired_embedding_dim_env_is_inert() {
crate::constants::set_active_embedding_dim(crate::constants::DEFAULT_EMBEDDING_DIM);
let before = chunk_embed_batch_size();
std::env::set_var("SQLITE_GRAPHRAG_EMBEDDING_DIM", "384");
let during = chunk_embed_batch_size();
std::env::remove_var("SQLITE_GRAPHRAG_EMBEDDING_DIM");
assert_eq!(
before, during,
"SQLITE_GRAPHRAG_EMBEDDING_DIM must not change the active dim"
);
}
#[test]
fn embedding_error_kind_classify_oauth_message() {
assert_eq!(
EmbeddingErrorKind::classify("OAuth token expired for claude"),
EmbeddingErrorKind::OAuth,
);
assert_eq!(
EmbeddingErrorKind::classify("oauth authentication failed"),
EmbeddingErrorKind::OAuth,
);
}
#[test]
fn embedding_error_kind_classify_quota_message() {
assert_eq!(
EmbeddingErrorKind::classify("quota exhausted on backend"),
EmbeddingErrorKind::Quota,
);
assert_eq!(
EmbeddingErrorKind::classify("Usage quota limit reached"),
EmbeddingErrorKind::Quota,
);
}
#[test]
fn embedding_error_kind_classify_slot_exhausted_message() {
assert_eq!(
EmbeddingErrorKind::classify("slot exhausted: failed to acquire LLM slot after backoff"),
EmbeddingErrorKind::SlotExhausted,
);
}
#[test]
fn embedding_error_kind_classify_zero_dimension_message() {
assert_eq!(
EmbeddingErrorKind::classify("embedding returned dim=zero"),
EmbeddingErrorKind::ZeroDimension,
);
assert_eq!(
EmbeddingErrorKind::classify("got zero-dim vector from LLM"),
EmbeddingErrorKind::ZeroDimension,
);
}
#[test]
fn embedding_error_kind_classify_unknown_fallback() {
assert_eq!(
EmbeddingErrorKind::classify("unrelated subprocess error"),
EmbeddingErrorKind::Unknown,
);
assert_eq!(
EmbeddingErrorKind::classify("rate limit hit"),
EmbeddingErrorKind::Unknown,
);
assert_eq!(EmbeddingErrorKind::OAuth.code(), "oauth");
assert_eq!(EmbeddingErrorKind::Quota.code(), "quota");
assert_eq!(EmbeddingErrorKind::SlotExhausted.code(), "slot-exhausted");
assert_eq!(
EmbeddingErrorKind::BackendMismatch.code(),
"backend-mismatch"
);
assert_eq!(EmbeddingErrorKind::ZeroDimension.code(), "zero-dimension");
assert_eq!(EmbeddingErrorKind::Unknown.code(), "unknown");
}
#[test]
fn fallback_reason_display_does_not_panic() {
let _ = FallbackReason::EmbeddingFailed("rate limit".into()).to_string();
let _ = FallbackReason::Cancelled.to_string();
let _ = FallbackReason::Timeout {
operation: "embed_query".into(),
duration_secs: 30,
}
.to_string();
}
#[test]
fn fallback_reason_is_partial_eq() {
assert_eq!(
FallbackReason::EmbeddingFailed("a".into()),
FallbackReason::EmbeddingFailed("a".into())
);
assert_eq!(FallbackReason::Cancelled, FallbackReason::Cancelled);
assert_ne!(
FallbackReason::EmbeddingFailed("a".into()),
FallbackReason::EmbeddingFailed("b".into())
);
assert_ne!(
FallbackReason::Cancelled,
FallbackReason::Timeout {
operation: "x".into(),
duration_secs: 1
}
);
}
#[test]
fn fallback_reason_timeout_preserves_fields() {
let r = FallbackReason::Timeout {
operation: "embed_query_local".into(),
duration_secs: 300,
};
match r {
FallbackReason::Timeout {
operation,
duration_secs,
} => {
assert_eq!(operation, "embed_query_local");
assert_eq!(duration_secs, 300);
}
other => panic!("expected Timeout, got {other:?}"),
}
}
#[test]
fn embedding_errors_map_to_the_fallback_reason_that_names_them() {
let reason = classify_embedding_error(crate::errors::AppError::Embedding(
"models dir /nonexistent not found".to_string(),
));
match reason {
FallbackReason::EmbeddingFailed(msg) => assert!(
msg.contains("/nonexistent"),
"the original error must survive for triage, got {msg:?}"
),
other => panic!("expected EmbeddingFailed, got {other:?}"),
}
let reason = classify_embedding_error(crate::errors::AppError::Timeout {
operation: "embed_query_local".to_string(),
duration_secs: 300,
});
match reason {
FallbackReason::Timeout {
operation,
duration_secs,
} => {
assert_eq!(operation, "embed_query_local");
assert_eq!(duration_secs, 300);
}
other => panic!("expected Timeout, got {other:?}"),
}
for (message, expected) in [
("embedding returned dim=zero", "dim_zero"),
("operation cancelled by signal", "cancelled"),
] {
let reason =
classify_embedding_error(crate::errors::AppError::Embedding(message.to_string()));
assert_eq!(
reason.reason_code(),
expected,
"{message:?} must classify as {expected}"
);
}
}
#[test]
fn g56_entity_cache_key_is_stable_and_distinct() {
let k1 = entity_cache_key("codex:default", "sqlite-graphrag");
let k2 = entity_cache_key("codex:default", "sqlite-graphrag");
let k3 = entity_cache_key("codex:default", "claude-code");
let k4 = entity_cache_key("claude:default", "sqlite-graphrag");
assert_eq!(k1, k2, "same model+text must hash identically");
assert_ne!(k1, k3, "different text must hash differently");
assert_ne!(k1, k4, "different model must hash differently");
}
#[test]
fn g56_entity_embed_cache_stats_hit_rate() {
let zero = EmbedCacheStats::default();
assert_eq!(zero.hit_rate(), 0.0);
let half = EmbedCacheStats {
requested: 4,
hits: 2,
misses: 2,
};
assert!((half.hit_rate() - 0.5).abs() < 1e-9);
let all = EmbedCacheStats {
requested: 7,
hits: 7,
misses: 0,
};
assert!((all.hit_rate() - 1.0).abs() < 1e-9);
}
#[test]
fn g56_entity_embed_cache_populates_and_hits() {
let cache = entity_embed_cache();
let model = "test-model";
let text = "sqlite-graphrag";
let key = entity_cache_key(model, text);
let stored = Arc::new(vec![0.42_f32; crate::constants::embedding_dim()]);
cache.lock().insert(key, Arc::clone(&stored));
let guard = cache.lock();
let hit = guard.get(&key).expect("cache must return stored value");
assert_eq!(hit.len(), crate::constants::embedding_dim());
assert!((hit[0] - 0.42).abs() < 1e-6);
}
#[test]
fn p1_openrouter_chain_ignores_llm_backend_none() {
use crate::cli::{EmbeddingBackendChoice, LlmBackendChoice};
let chain = EmbeddingBackendChoice::Openrouter.to_chain(LlmBackendChoice::None);
assert_eq!(
chain,
vec![LlmBackendKind::OpenRouter],
"openrouter embedding must not be silenced by --llm-backend none"
);
let none_chain = LlmBackendChoice::None.to_chain();
assert_eq!(none_chain, vec![LlmBackendKind::None]);
}
#[test]
fn g56_empty_texts_short_circuits_with_zero_stats() {
let stats = EmbedCacheStats::default();
assert_eq!(stats.requested, 0);
assert_eq!(stats.hits, 0);
assert_eq!(stats.misses, 0);
assert_eq!(stats.hit_rate(), 0.0);
}