use anyhow::Result;
use serde::{Deserialize, Serialize};
use smallvec::SmallVec;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbeddingConfig {
pub model_type: String,
pub model_path: String,
pub embedding_dim: usize,
pub max_tokens: usize,
pub batch_size: usize,
pub device: String,
pub use_simd: bool,
pub memory_pool_mb: usize,
}
impl Default for EmbeddingConfig {
fn default() -> Self {
Self {
model_type: "sentence-transformers".to_string(),
model_path: "all-MiniLM-L6-v2".to_string(),
embedding_dim: 384,
max_tokens: 512,
batch_size: 32,
device: "cpu".to_string(),
use_simd: true,
memory_pool_mb: 512,
}
}
}
#[derive(Debug, Clone)]
pub struct CodeEmbedding {
pub vector: SmallVec<[f32; 768]>,
pub dim: usize,
pub norm: f32,
pub metadata: EmbeddingMetadata,
}
impl CodeEmbedding {
pub fn new(vector: Vec<f32>, metadata: EmbeddingMetadata) -> Self {
let dim = vector.len();
let norm = vector.iter().map(|&x| x * x).sum::<f32>().sqrt();
Self {
vector: SmallVec::from_vec(vector),
dim,
norm,
metadata,
}
}
pub fn cosine_similarity(&self, other: &Self) -> f32 {
if self.dim != other.dim {
return 0.0;
}
if self.norm == 0.0 || other.norm == 0.0 {
return 0.0;
}
let dot_product: f32 = self.vector.iter()
.zip(other.vector.iter())
.map(|(a, b)| a * b)
.sum();
dot_product / (self.norm * other.norm)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbeddingMetadata {
pub doc_id: String,
pub file_path: Option<String>,
pub location: Option<(usize, usize)>,
pub language: Option<String>,
pub encoded_at: u64,
pub model_version: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CacheStats {
pub hit_count: u64,
pub miss_count: u64,
pub hit_rate: f64,
pub cache_size: usize,
pub evictions: u64,
}
pub struct SemanticEncoder {
config: EmbeddingConfig,
cache_state: Arc<RwLock<EncoderCacheState>>,
}
#[derive(Debug)]
struct EncoderCacheState {
token_cache: HashMap<String, CodeEmbedding>,
cache_hits: u64,
cache_misses: u64,
cache_evictions: u64,
}
impl SemanticEncoder {
pub async fn new(config: EmbeddingConfig) -> Result<Self> {
let cache_state = EncoderCacheState {
token_cache: HashMap::new(),
cache_hits: 0,
cache_misses: 0,
cache_evictions: 0,
};
Ok(Self {
config,
cache_state: Arc::new(RwLock::new(cache_state)),
})
}
pub async fn encode_query(&self, query: &str) -> Result<CodeEmbedding> {
{
let cache = self.cache_state.read().await;
if let Some(cached) = cache.token_cache.get(query) {
drop(cache);
let mut cache = self.cache_state.write().await;
cache.cache_hits += 1;
if let Some(cached) = cache.token_cache.get(query) {
return Ok(cached.clone());
}
}
}
{
let mut cache = self.cache_state.write().await;
cache.cache_misses += 1;
}
let normalized_query = query.trim().to_lowercase();
let tokens: Vec<&str> = normalized_query.split_whitespace().collect();
let mut vector = vec![0.0; self.config.embedding_dim];
for (i, token) in tokens.iter().enumerate() {
let token_hash = self.hash_token(token);
for j in 0..self.config.embedding_dim {
let idx = (token_hash + j) % self.config.embedding_dim;
vector[idx] += 1.0 / (i + 1) as f32; }
}
let norm: f32 = vector.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for v in &mut vector {
*v /= norm;
}
}
let metadata = EmbeddingMetadata {
doc_id: format!("query_{}", self.hash_token(query)),
file_path: None,
location: None,
language: Some("query".to_string()),
encoded_at: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
model_version: format!("{}_{}", self.config.model_type, self.config.model_path),
};
let embedding = CodeEmbedding::new(vector, metadata);
{
let mut cache = self.cache_state.write().await;
Self::manage_cache_size(&mut cache);
cache.token_cache.insert(query.to_string(), embedding.clone());
}
Ok(embedding)
}
pub async fn encode_code(&self, code: &str) -> Result<CodeEmbedding> {
{
let cache = self.cache_state.read().await;
if let Some(cached) = cache.token_cache.get(code) {
drop(cache);
let mut cache = self.cache_state.write().await;
cache.cache_hits += 1;
if let Some(cached) = cache.token_cache.get(code) {
return Ok(cached.clone());
}
}
}
{
let mut cache = self.cache_state.write().await;
cache.cache_misses += 1;
}
let normalized_code = code.trim();
let tokens: Vec<&str> = normalized_code
.split(|c: char| c.is_whitespace() || "(){}[],.;:".contains(c))
.filter(|s| !s.is_empty())
.collect();
let mut vector = vec![0.0; self.config.embedding_dim];
let has_function = normalized_code.contains("fn ") || normalized_code.contains("function");
let has_class = normalized_code.contains("class ") || normalized_code.contains("struct ");
let has_import = normalized_code.contains("import ") || normalized_code.contains("use ");
if has_function { vector[0] += 2.0; }
if has_class { vector[1] += 2.0; }
if has_import { vector[2] += 1.5; }
for (i, token) in tokens.iter().enumerate() {
let token_hash = self.hash_token(token);
let weight = if self.is_keyword(token) { 2.0 } else { 1.0 };
for j in 0..self.config.embedding_dim {
let idx = (token_hash + j * 3) % self.config.embedding_dim;
vector[idx] += weight / (i + 1) as f32; }
}
let norm: f32 = vector.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for v in &mut vector {
*v /= norm;
}
}
let metadata = EmbeddingMetadata {
doc_id: format!("code_{}", self.hash_token(code)),
file_path: None,
location: None,
language: self.detect_language(normalized_code),
encoded_at: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
model_version: format!("{}_{}", self.config.model_type, self.config.model_path),
};
let embedding = CodeEmbedding::new(vector, metadata);
{
let mut cache = self.cache_state.write().await;
Self::manage_cache_size(&mut cache);
cache.token_cache.insert(code.to_string(), embedding.clone());
}
Ok(embedding)
}
fn manage_cache_size(cache_state: &mut EncoderCacheState) {
const MAX_CACHE_SIZE: usize = 10000;
if cache_state.token_cache.len() >= MAX_CACHE_SIZE {
let keys_to_remove: Vec<_> = cache_state.token_cache.keys().take(MAX_CACHE_SIZE / 4).cloned().collect();
for key in keys_to_remove {
cache_state.token_cache.remove(&key);
cache_state.cache_evictions += 1;
}
}
}
pub async fn get_cache_stats(&self) -> CacheStats {
let cache = self.cache_state.read().await;
let total_requests = cache.cache_hits + cache.cache_misses;
let hit_rate = if total_requests > 0 {
cache.cache_hits as f64 / total_requests as f64
} else {
0.0
};
CacheStats {
hit_count: cache.cache_hits,
miss_count: cache.cache_misses,
hit_rate,
cache_size: cache.token_cache.len(),
evictions: cache.cache_evictions,
}
}
pub async fn health_check(&self) -> Result<()> {
if self.config.embedding_dim == 0 {
return Err(anyhow::anyhow!("Invalid embedding dimension: 0"));
}
if self.config.max_tokens == 0 {
return Err(anyhow::anyhow!("Invalid max tokens: 0"));
}
let test_query = "test health check";
let result = self.encode_query(test_query).await?;
if result.vector.is_empty() {
return Err(anyhow::anyhow!("Health check failed: empty embedding vector"));
}
if result.dim != self.config.embedding_dim {
return Err(anyhow::anyhow!(
"Health check failed: dimension mismatch {} vs {}",
result.dim,
self.config.embedding_dim
));
}
Ok(())
}
fn hash_token(&self, token: &str) -> usize {
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
let mut hasher = DefaultHasher::new();
token.hash(&mut hasher);
hasher.finish() as usize
}
fn is_keyword(&self, token: &str) -> bool {
matches!(token.to_lowercase().as_str(),
"fn" | "function" | "class" | "struct" | "enum" | "impl" | "trait" |
"if" | "else" | "while" | "for" | "loop" | "match" | "return" |
"let" | "mut" | "const" | "static" | "pub" | "use" | "mod" |
"async" | "await" | "try" | "catch" | "throw" | "import" | "export" |
"var" | "const" | "let" | "def" | "lambda" | "yield" | "with"
)
}
fn detect_language(&self, code: &str) -> Option<String> {
if code.contains("fn ") && code.contains("->") {
Some("rust".to_string())
} else if code.contains("function ") || code.contains("const ") || code.contains("=>") {
Some("javascript".to_string())
} else if code.contains("def ") && code.contains(":") {
Some("python".to_string())
} else if code.contains("class ") && code.contains("{") {
Some("java".to_string())
} else if code.contains("#include") || code.contains("int main") {
Some("c".to_string())
} else {
Some("unknown".to_string())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_embedding_config_default() {
let config = EmbeddingConfig::default();
assert_eq!(config.model_type, "sentence-transformers");
assert_eq!(config.model_path, "all-MiniLM-L6-v2");
assert_eq!(config.embedding_dim, 384);
assert_eq!(config.max_tokens, 512);
assert_eq!(config.batch_size, 32);
assert_eq!(config.device, "cpu");
assert!(config.use_simd);
assert_eq!(config.memory_pool_mb, 512);
}
#[test]
fn test_code_embedding_creation() {
let metadata = EmbeddingMetadata {
doc_id: "test".to_string(),
file_path: None,
location: None,
language: None,
encoded_at: 0,
model_version: "test".to_string(),
};
let vector = vec![0.5; 384];
let embedding = CodeEmbedding::new(vector.clone(), metadata);
assert_eq!(embedding.vector.len(), 384);
assert_eq!(embedding.dim, 384);
assert!(embedding.norm > 0.0);
}
#[test]
fn test_embedding_metadata_creation() {
let metadata = EmbeddingMetadata {
doc_id: "test-doc".to_string(),
file_path: Some("/path/to/file.rs".to_string()),
location: Some((10, 20)),
language: Some("rust".to_string()),
encoded_at: 12345,
model_version: "1.0.0".to_string(),
};
assert_eq!(metadata.doc_id, "test-doc");
assert_eq!(metadata.file_path, Some("/path/to/file.rs".to_string()));
assert_eq!(metadata.location, Some((10, 20)));
assert_eq!(metadata.language, Some("rust".to_string()));
assert_eq!(metadata.encoded_at, 12345);
assert_eq!(metadata.model_version, "1.0.0");
}
#[tokio::test]
async fn test_semantic_encoder_creation() {
let config = EmbeddingConfig::default();
let encoder_result = SemanticEncoder::new(config).await;
assert!(encoder_result.is_ok());
let encoder = encoder_result.unwrap();
assert_eq!(encoder.config.embedding_dim, 384);
}
#[tokio::test]
async fn test_query_encoding() {
let config = EmbeddingConfig::default();
let encoder = SemanticEncoder::new(config).await.unwrap();
let query = "fn main() { println!(\"Hello\"); }";
let result = encoder.encode_query(query).await;
assert!(result.is_ok());
let embedding = result.unwrap();
assert_eq!(embedding.dim, 384);
assert_eq!(embedding.vector.len(), 384);
assert!(embedding.norm > 0.0);
}
#[tokio::test]
async fn test_code_encoding() {
let config = EmbeddingConfig::default();
let encoder = SemanticEncoder::new(config).await.unwrap();
let code = "fn main() { println!(\"Hello\"); }";
let result = encoder.encode_code(code).await;
assert!(result.is_ok());
let embedding = result.unwrap();
assert_eq!(embedding.dim, 384);
assert_eq!(embedding.vector.len(), 384);
assert!(embedding.norm > 0.0);
}
#[tokio::test]
async fn test_health_check() {
let config = EmbeddingConfig::default();
let encoder = SemanticEncoder::new(config).await.unwrap();
let result = encoder.health_check().await;
assert!(result.is_ok());
}
#[test]
fn test_cosine_similarity_real() {
let metadata1 = EmbeddingMetadata {
doc_id: "test1".to_string(),
file_path: None,
location: None,
language: None,
encoded_at: 0,
model_version: "test".to_string(),
};
let metadata2 = EmbeddingMetadata {
doc_id: "test2".to_string(),
file_path: None,
location: None,
language: None,
encoded_at: 0,
model_version: "test".to_string(),
};
let vector1 = vec![1.0, 0.0, 0.0];
let vector2 = vec![1.0, 0.0, 0.0];
let embedding1 = CodeEmbedding::new(vector1, metadata1.clone());
let embedding2 = CodeEmbedding::new(vector2, metadata2.clone());
let similarity = embedding1.cosine_similarity(&embedding2);
assert!((similarity - 1.0).abs() < 0.001);
let vector3 = vec![1.0, 0.0, 0.0];
let vector4 = vec![0.0, 1.0, 0.0];
let embedding3 = CodeEmbedding::new(vector3, metadata1);
let embedding4 = CodeEmbedding::new(vector4, metadata2);
let similarity2 = embedding3.cosine_similarity(&embedding4);
assert!((similarity2 - 0.0).abs() < 0.001); }
}