use std::env;
use std::time::Instant;
use ctx::embeddings::{self, EmbeddingProvider};
use ctx::error::Result;
use ctx::index;
use ctx::utils::{truncate_path, truncate_str};
pub fn run_embed(force: bool, verbose: bool, batch_size: usize, use_openai: bool) -> Result<()> {
let root = env::current_dir()?;
let db = index::open_database(&root)?;
let provider: Box<dyn EmbeddingProvider> = if use_openai {
use embeddings::openai::OpenAIProvider;
let p = OpenAIProvider::from_env().map_err(|_| {
"OPENAI_API_KEY environment variable not set.\n\
Set it with: export OPENAI_API_KEY=sk-..."
})?;
Box::new(p)
} else {
use embeddings::local::LocalProvider;
eprintln!("Initializing local embedding model (first run downloads ~90MB)...");
let p = LocalProvider::new()?;
Box::new(p)
};
if verbose {
println!(
"Using embedding provider: {} (dim={})",
provider.name(),
provider.dimension()
);
}
if force {
let deleted = db.delete_embeddings(provider.name(), None)?;
if verbose {
println!("Deleted {} existing embeddings", deleted);
}
}
let total_symbols = db.get_stats()?.symbols;
let existing_embeddings = db.count_embeddings()?;
if verbose {
println!("Total symbols: {}", total_symbols);
println!("Existing embeddings: {}", existing_embeddings);
}
if existing_embeddings >= total_symbols && !force {
println!("All symbols already have embeddings. Use --force to re-embed.");
return Ok(());
}
println!(
"Generating embeddings for {} symbols...",
total_symbols - existing_embeddings
);
let start = Instant::now();
let progress_callback = |done: usize, _total: usize| {
if verbose {
eprint!("\rEmbedded {} symbols...", done);
}
};
let embedded = embeddings::embed_missing_symbols(
&db,
provider.as_ref(),
batch_size,
Some(&progress_callback),
)?;
if verbose {
eprintln!();
}
let elapsed = start.elapsed();
println!(
"Embedded {} symbols in {:.2}s ({:.1} symbols/sec)",
embedded,
elapsed.as_secs_f64(),
embedded as f64 / elapsed.as_secs_f64()
);
Ok(())
}
pub fn run_embed_watch(verbose: bool, batch_size: usize, use_openai: bool) -> Result<()> {
use notify::RecursiveMode;
use notify_debouncer_mini::{new_debouncer, DebouncedEventKind};
use std::sync::mpsc::channel;
use std::time::Duration;
let root = env::current_dir()?;
let ctx_dir = root.join(".ctx");
let _db_path = ctx_dir.join("codebase.sqlite");
let provider: Box<dyn EmbeddingProvider> = if use_openai {
use embeddings::openai::OpenAIProvider;
let p = OpenAIProvider::from_env().map_err(|_| {
"OPENAI_API_KEY environment variable not set.\n\
Set it with: export OPENAI_API_KEY=sk-..."
})?;
Box::new(p)
} else {
use embeddings::local::LocalProvider;
eprintln!("Initializing local embedding model (first run downloads ~90MB)...");
let p = LocalProvider::new()?;
Box::new(p)
};
println!(
"Using embedding provider: {} (dim={})",
provider.name(),
provider.dimension()
);
{
let db = index::open_database(&root)?;
let total_symbols = db.get_stats()?.symbols;
let existing = db.count_embeddings()?;
if existing < total_symbols {
println!(
"Initial embedding: {} symbols missing embeddings...",
total_symbols - existing
);
let embedded =
embeddings::embed_missing_symbols(&db, provider.as_ref(), batch_size, None)?;
println!("Embedded {} symbols", embedded);
} else {
println!("All {} symbols already have embeddings", total_symbols);
}
}
let (tx, rx) = channel();
let mut debouncer = new_debouncer(Duration::from_secs(2), tx)
.map_err(|e| format!("Failed to create watcher: {}", e))?;
if ctx_dir.exists() {
debouncer
.watcher()
.watch(&ctx_dir, RecursiveMode::NonRecursive)
.map_err(|e| format!("Failed to watch .ctx directory: {}", e))?;
}
println!("\nWatching for index changes... (press Ctrl+C to stop)");
println!("Tip: Run 'ctx index --watch' in another terminal to auto-index file changes");
loop {
match rx.recv() {
Ok(Ok(events)) => {
let db_changed = events.iter().any(|e| {
e.kind == DebouncedEventKind::Any
&& e.path
.file_name()
.map(|n| n == "codebase.sqlite")
.unwrap_or(false)
});
if db_changed {
match index::open_database(&root) {
Ok(db) => {
let total = db.get_stats().map(|s| s.symbols).unwrap_or(0);
let existing = db.count_embeddings().unwrap_or(0);
if existing < total {
let missing = total - existing;
if verbose {
eprintln!("\nIndex updated: {} new symbols to embed", missing);
}
match embeddings::embed_missing_symbols(
&db,
provider.as_ref(),
batch_size,
None,
) {
Ok(embedded) => {
if embedded > 0 {
if verbose {
eprintln!("Embedded {} symbols", embedded);
} else {
eprint!("+{} ", embedded);
}
}
}
Err(e) => {
eprintln!("\nWarning: failed to embed: {}", e);
}
}
}
}
Err(e) => {
if verbose {
eprintln!("\nWarning: failed to open database: {}", e);
}
}
}
}
}
Ok(Err(e)) => {
eprintln!("\nWatch error: {:?}", e);
}
Err(e) => {
eprintln!("\nChannel error: {}", e);
break;
}
}
}
Ok(())
}
pub fn run_semantic(query: &str, limit: usize, output: &str, use_openai: bool) -> Result<()> {
let root = env::current_dir()?;
let db = index::open_database(&root)?;
let embedding_count = db.count_embeddings()?;
if embedding_count == 0 {
eprintln!("No embeddings found. Run 'ctx embed' first to generate embeddings.");
return Ok(());
}
let provider: Box<dyn EmbeddingProvider> = if use_openai {
use embeddings::openai::OpenAIProvider;
let p = OpenAIProvider::from_env().map_err(|_| {
"OPENAI_API_KEY environment variable not set.\n\
Set it with: export OPENAI_API_KEY=sk-..."
})?;
Box::new(p)
} else {
use embeddings::local::LocalProvider;
let p = LocalProvider::new()?;
Box::new(p)
};
let query_dim = provider.dimension();
if let Ok(metadata) = db.get_embedding_metadata() {
for (stored_provider, _model, stored_dim, count) in &metadata {
let stored_dim = *stored_dim as usize;
if stored_dim != query_dim {
eprintln!("Warning: Embedding dimension mismatch detected!");
eprintln!(
" Stored: {} embeddings from '{}' with dimension {}",
count, stored_provider, stored_dim
);
eprintln!(
" Query: Using '{}' with dimension {}",
provider.name(),
query_dim
);
eprintln!(
" Results may be inaccurate. Re-run 'ctx embed{}' to regenerate embeddings.",
if use_openai { " --openai" } else { "" }
);
eprintln!();
}
}
}
let query_embedding = provider.embed(query)?;
let results = embeddings::semantic_search(&db, &query_embedding, limit)?;
if results.is_empty() {
eprintln!("No results found for '{}'", query);
return Ok(());
}
if output == "json" {
let json_results: Vec<_> = results
.iter()
.map(|r| {
serde_json::json!({
"symbol_id": r.symbol_id,
"name": r.name,
"kind": r.kind,
"file": r.file_path,
"line": r.line,
"score": format!("{:.4}", r.score),
})
})
.collect();
println!("{}", serde_json::to_string_pretty(&json_results)?);
} else {
println!(
"Semantic search for '{}' ({} results):",
query,
results.len()
);
println!("{}", "-".repeat(80));
println!("{:<35} {:<10} {:<8} FILE", "SYMBOL", "KIND", "SCORE");
println!("{}", "-".repeat(80));
for result in &results {
let name = truncate_str(&result.name, 33);
let file = truncate_path(&result.file_path, 25);
let score_display = format!("{:.2}%", result.score * 100.0);
println!(
"{:<35} {:<10} {:<8} {}:{}",
name, result.kind, score_display, file, result.line
);
}
}
Ok(())
}