1use std::collections::HashSet;
15use std::io::Write as _;
16use std::path::Path;
17
18use serde::Deserialize;
19
20use crate::core::bm25_index::{BM25Index, CodeChunk};
21
22#[derive(Debug, Clone)]
23pub struct PgvectorConfig {
24 pub url: String,
26 pub timeout_secs: u64,
28 pub table_prefix: String,
30}
31
32impl PgvectorConfig {
33 pub fn from_env() -> Result<Self, String> {
34 let url = std::env::var("LEANCTX_PGVECTOR_URL")
35 .map_err(|_| "LEANCTX_PGVECTOR_URL is required for pgvector backend".to_string())?;
36 let url = url.trim().to_string();
37 if url.is_empty() {
38 return Err("LEANCTX_PGVECTOR_URL is required for pgvector backend".to_string());
39 }
40
41 let timeout_secs = std::env::var("LEANCTX_PGVECTOR_TIMEOUT_SECS")
42 .ok()
43 .and_then(|v| v.trim().parse::<u64>().ok())
44 .filter(|v| *v > 0)
45 .unwrap_or(10);
46
47 let table_prefix = std::env::var("LEANCTX_PGVECTOR_TABLE_PREFIX")
48 .ok()
49 .map(|v| v.trim().to_string())
50 .filter(|v| !v.is_empty())
51 .unwrap_or_else(|| "lctx_code_".to_string());
52 validate_pg_identifier_prefix(&table_prefix)?;
53
54 Ok(Self {
55 url,
56 timeout_secs,
57 table_prefix,
58 })
59 }
60}
61
62#[derive(Debug, Clone)]
63pub struct PgvectorStore {
64 cfg: PgvectorConfig,
65}
66
67#[derive(Debug, Clone)]
68pub struct PgvectorHit {
69 pub score: f32,
70 pub file_path: String,
71 pub symbol_name: String,
72 pub kind: crate::core::bm25_index::ChunkKind,
73 pub start_line: usize,
74 pub end_line: usize,
75}
76
77#[derive(Debug, Deserialize)]
78struct PgRow {
79 score: f32,
80 file_path: String,
81 symbol_name: String,
82 kind: String,
83 start_line: usize,
84 end_line: usize,
85}
86
87impl PgvectorStore {
88 pub fn from_env() -> Result<Self, String> {
89 let cfg = PgvectorConfig::from_env()?;
90 Ok(Self { cfg })
91 }
92
93 pub fn table_name(&self, root: &Path, dimensions: usize) -> Result<String, String> {
96 let ns = crate::core::index_namespace::namespace_hash(root);
97 let name = format!("{}{}_d{}", self.cfg.table_prefix, ns, dimensions);
98 if name.len() > 63 {
99 return Err(format!(
100 "pgvector table name exceeds PostgreSQL's 63-byte identifier limit: {name}"
101 ));
102 }
103 Ok(name)
104 }
105
106 pub fn ensure_table(&self, table: &str, dimensions: usize) -> Result<bool, String> {
108 let existed = self.table_exists(table)?;
109 if existed {
110 return Ok(false);
111 }
112 let sql = format!(
113 "CREATE EXTENSION IF NOT EXISTS vector;\n\
114 CREATE TABLE IF NOT EXISTS {table} (\n\
115 id BIGINT PRIMARY KEY,\n\
116 file_path TEXT NOT NULL,\n\
117 symbol_name TEXT NOT NULL,\n\
118 kind TEXT NOT NULL,\n\
119 start_line BIGINT NOT NULL,\n\
120 end_line BIGINT NOT NULL,\n\
121 embedding vector({dimensions}) NOT NULL\n\
122 );\n\
123 CREATE INDEX IF NOT EXISTS {table}_file_idx ON {table} (file_path);"
124 );
125 self.run_sql(&sql)?;
126 Ok(true)
127 }
128
129 pub fn sync_index(
132 &self,
133 table: &str,
134 index: &BM25Index,
135 aligned_embeddings: &[Vec<f32>],
136 changed_files: &[String],
137 created_new: bool,
138 ) -> Result<(), String> {
139 if index.chunks.len() != aligned_embeddings.len() {
140 return Err("embedding alignment length mismatch".to_string());
141 }
142
143 if created_new {
144 return self.upsert_filtered(table, index, aligned_embeddings, None);
145 }
146
147 if changed_files.is_empty() {
148 return Ok(());
149 }
150
151 let mut unique: Vec<String> = changed_files.to_vec();
152 unique.sort();
153 unique.dedup();
154
155 for file in &unique {
156 self.delete_by_file(table, file)?;
157 }
158
159 let changed_set: HashSet<&str> = unique.iter().map(String::as_str).collect();
160 self.upsert_filtered(table, index, aligned_embeddings, Some(&changed_set))
161 }
162
163 pub fn search(
164 &self,
165 table: &str,
166 query_vec: &[f32],
167 limit: usize,
168 ) -> Result<Vec<PgvectorHit>, String> {
169 let vec_literal = vector_literal(query_vec);
170 let sql = format!(
172 "SELECT json_build_object(\
173 'score', 1 - (embedding <=> '{vec_literal}'::vector), \
174 'file_path', file_path, \
175 'symbol_name', symbol_name, \
176 'kind', kind, \
177 'start_line', start_line, \
178 'end_line', end_line\
179 )::text \
180 FROM {table} \
181 ORDER BY embedding <=> '{vec_literal}'::vector \
182 LIMIT {limit};"
183 );
184 let stdout = self.run_sql(&sql)?;
185
186 let mut out = Vec::new();
187 for line in stdout.lines() {
188 let line = line.trim();
189 if line.is_empty() {
190 continue;
191 }
192 let row: PgRow = serde_json::from_str(line)
193 .map_err(|e| format!("invalid pgvector row json: {e}"))?;
194 out.push(PgvectorHit {
195 score: row.score,
196 file_path: row.file_path,
197 symbol_name: row.symbol_name,
198 kind: crate::core::dense_backend::kind_from_str(&row.kind),
199 start_line: row.start_line,
200 end_line: row.end_line,
201 });
202 }
203 Ok(out)
204 }
205
206 fn table_exists(&self, table: &str) -> Result<bool, String> {
207 let literal = sql_string_literal(table)?;
208 let out = self.run_sql(&format!("SELECT to_regclass({literal}) IS NOT NULL;"))?;
209 Ok(out.trim() == "t")
210 }
211
212 fn upsert_filtered(
215 &self,
216 table: &str,
217 index: &BM25Index,
218 aligned_embeddings: &[Vec<f32>],
219 changed_set: Option<&HashSet<&str>>,
220 ) -> Result<(), String> {
221 let mut batch: Vec<String> = Vec::new();
222 for (i, chunk) in index.chunks.iter().enumerate() {
223 if let Some(set) = changed_set
224 && !set.contains(chunk.file_path.as_str())
225 {
226 continue;
227 }
228 let vec = aligned_embeddings
229 .get(i)
230 .ok_or_else(|| "embedding alignment missing".to_string())?;
231 batch.push(values_row_for_chunk(chunk, vec)?);
232 if batch.len() >= UPSERT_BATCH_ROWS {
233 self.upsert_rows(table, &batch)?;
234 batch.clear();
235 }
236 }
237 if !batch.is_empty() {
238 self.upsert_rows(table, &batch)?;
239 }
240 Ok(())
241 }
242
243 fn upsert_rows(&self, table: &str, rows: &[String]) -> Result<(), String> {
244 let sql = format!(
245 "INSERT INTO {table} (id, file_path, symbol_name, kind, start_line, end_line, embedding)\n\
246 VALUES\n{}\n\
247 ON CONFLICT (id) DO UPDATE SET\n\
248 file_path = EXCLUDED.file_path,\n\
249 symbol_name = EXCLUDED.symbol_name,\n\
250 kind = EXCLUDED.kind,\n\
251 start_line = EXCLUDED.start_line,\n\
252 end_line = EXCLUDED.end_line,\n\
253 embedding = EXCLUDED.embedding;",
254 rows.join(",\n")
255 );
256 self.run_sql(&sql).map(|_| ())
257 }
258
259 fn delete_by_file(&self, table: &str, file_path: &str) -> Result<(), String> {
260 let literal = sql_string_literal(file_path)?;
261 self.run_sql(&format!("DELETE FROM {table} WHERE file_path = {literal};"))
262 .map(|_| ())
263 }
264
265 fn run_sql(&self, sql: &str) -> Result<String, String> {
268 let mut tmp = tempfile::NamedTempFile::new()
269 .map_err(|e| format!("pgvector: temp file failed: {e}"))?;
270 tmp.write_all(sql.as_bytes())
271 .map_err(|e| format!("pgvector: temp write failed: {e}"))?;
272 tmp.flush()
273 .map_err(|e| format!("pgvector: temp flush failed: {e}"))?;
274
275 let output = std::process::Command::new("psql")
276 .arg(&self.cfg.url)
277 .args(["-X", "-q", "-v", "ON_ERROR_STOP=1", "-t", "-A", "-f"])
278 .arg(tmp.path())
279 .env("PGCONNECT_TIMEOUT", self.cfg.timeout_secs.to_string())
280 .output()
281 .map_err(|e| {
282 format!("pgvector: failed to run psql (is the PostgreSQL client installed?): {e}")
283 })?;
284
285 if !output.status.success() {
286 let stderr = String::from_utf8_lossy(&output.stderr);
287 return Err(format!("pgvector: psql error: {}", stderr.trim()));
288 }
289 Ok(String::from_utf8_lossy(&output.stdout).into_owned())
290 }
291}
292
293const UPSERT_BATCH_ROWS: usize = 256;
294
295fn values_row_for_chunk(chunk: &CodeChunk, vector: &[f32]) -> Result<String, String> {
297 let id = point_id_for_chunk(chunk) as i64; let file = sql_string_literal(&chunk.file_path)?;
299 let symbol = sql_string_literal(&chunk.symbol_name)?;
300 let kind = sql_string_literal(crate::core::dense_backend::kind_to_str(&chunk.kind))?;
301 Ok(format!(
302 "({id}, {file}, {symbol}, {kind}, {}, {}, '{}'::vector)",
303 chunk.start_line,
304 chunk.end_line,
305 vector_literal(vector)
306 ))
307}
308
309fn point_id_for_chunk(chunk: &CodeChunk) -> u64 {
312 use md5::{Digest, Md5};
313 let mut h = Md5::new();
314 h.update(chunk.file_path.as_bytes());
315 h.update(chunk.start_line.to_le_bytes());
316 h.update(chunk.end_line.to_le_bytes());
317 h.update(chunk.symbol_name.as_bytes());
318 h.update(crate::core::dense_backend::kind_to_str(&chunk.kind).as_bytes());
319 let out = h.finalize();
320 u64::from_le_bytes(out[0..8].try_into().unwrap_or([0u8; 8]))
321}
322
323fn vector_literal(vector: &[f32]) -> String {
325 let mut s = String::with_capacity(vector.len() * 10 + 2);
326 s.push('[');
327 for (i, v) in vector.iter().enumerate() {
328 if i > 0 {
329 s.push(',');
330 }
331 s.push_str(&format!("{v}"));
333 }
334 s.push(']');
335 s
336}
337
338fn sql_string_literal(s: &str) -> Result<String, String> {
351 if s.contains('\0') {
352 return Err("pgvector: NUL byte in string".to_string());
353 }
354 let escaped = s.replace('\\', "\\\\").replace('\'', "''");
355 Ok(format!("E'{escaped}'"))
356}
357
358fn validate_pg_identifier_prefix(name: &str) -> Result<(), String> {
361 let valid_start = name
362 .chars()
363 .next()
364 .is_some_and(|c| c.is_ascii_alphabetic() || c == '_');
365 let valid_rest = name
366 .chars()
367 .skip(1)
368 .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '$');
369 if name.is_empty() || name.len() > 40 || !valid_start || !valid_rest {
370 return Err(format!(
371 "Invalid LEANCTX_PGVECTOR_TABLE_PREFIX: {name:?} (allowed: [A-Za-z_][A-Za-z0-9_$]*, max 40 chars)"
372 ));
373 }
374 Ok(())
375}
376
377#[cfg(test)]
378mod tests {
379 use super::*;
380 use crate::core::bm25_index::ChunkKind;
381
382 fn chunk(file: &str, name: &str, start: usize, end: usize, kind: ChunkKind) -> CodeChunk {
383 CodeChunk {
384 file_path: file.to_string(),
385 symbol_name: name.to_string(),
386 kind,
387 start_line: start,
388 end_line: end,
389 content: "fn x() {}".to_string(),
390 tokens: vec![],
391 token_count: 0,
392 }
393 }
394
395 #[test]
396 fn point_id_matches_qdrant_scheme_and_is_stable() {
397 let c = chunk("src/main.rs", "main", 1, 10, ChunkKind::Function);
398 assert_eq!(point_id_for_chunk(&c), point_id_for_chunk(&c));
399 let c2 = chunk("src/main.rs", "main", 2, 10, ChunkKind::Function);
400 assert_ne!(point_id_for_chunk(&c), point_id_for_chunk(&c2));
401 }
402
403 #[test]
404 fn sql_string_literal_escapes_quotes() {
405 assert_eq!(sql_string_literal("a'b").unwrap(), "E'a''b'");
406 assert_eq!(sql_string_literal("plain").unwrap(), "E'plain'");
407 assert!(sql_string_literal("nul\0byte").is_err());
408 }
409
410 #[test]
411 fn sql_string_literal_escapes_backslashes() {
412 assert_eq!(sql_string_literal(r"a\b").unwrap(), r"E'a\\b'");
418 assert_eq!(sql_string_literal(r"trailing\").unwrap(), r"E'trailing\\'");
419 assert_eq!(sql_string_literal(r"a\'b").unwrap(), r"E'a\\''b'");
420 }
421
422 #[test]
423 fn vector_literal_is_bracketed_csv() {
424 assert_eq!(vector_literal(&[0.5, -1.0, 2.0]), "[0.5,-1,2]");
425 assert_eq!(vector_literal(&[]), "[]");
426 }
427
428 #[test]
429 fn values_row_contains_escaped_fields() {
430 let c = chunk("src/a'b.rs", "fn'x", 3, 9, ChunkKind::Method);
431 let row = values_row_for_chunk(&c, &[0.25, 0.75]).unwrap();
432 assert!(row.contains("'src/a''b.rs'"));
433 assert!(row.contains("'fn''x'"));
434 assert!(row.contains("'Method'"));
435 assert!(row.contains("'[0.25,0.75]'::vector"));
436 }
437
438 #[test]
439 fn table_prefix_validation() {
440 assert!(validate_pg_identifier_prefix("lctx_code_").is_ok());
441 assert!(validate_pg_identifier_prefix("with space").is_err());
442 assert!(validate_pg_identifier_prefix("1leading_digit").is_err());
443 assert!(validate_pg_identifier_prefix("drop;table").is_err());
444 assert!(validate_pg_identifier_prefix("").is_err());
445 }
446
447 #[test]
448 fn config_requires_url() {
449 let _env = crate::core::data_dir::test_env_lock();
450 crate::test_env::remove_var("LEANCTX_PGVECTOR_URL");
451 assert!(PgvectorConfig::from_env().is_err());
452
453 crate::test_env::set_var("LEANCTX_PGVECTOR_URL", "postgres://localhost/lctx");
454 crate::test_env::remove_var("LEANCTX_PGVECTOR_TABLE_PREFIX");
455 crate::test_env::remove_var("LEANCTX_PGVECTOR_TIMEOUT_SECS");
456 let cfg = PgvectorConfig::from_env().unwrap();
457 assert_eq!(cfg.url, "postgres://localhost/lctx");
458 assert_eq!(cfg.table_prefix, "lctx_code_");
459 assert_eq!(cfg.timeout_secs, 10);
460 crate::test_env::remove_var("LEANCTX_PGVECTOR_URL");
461 }
462
463 #[test]
467 #[ignore = "requires live PostgreSQL with pgvector extension (set LEANCTX_PGVECTOR_URL)"]
468 fn pgvector_e2e_round_trip() {
469 let store = PgvectorStore::from_env().expect("LEANCTX_PGVECTOR_URL must be set");
470 let table = "lctx_e2e_round_trip_d3".to_string();
471 let _ = store.run_sql(&format!("DROP TABLE IF EXISTS {table};"));
472
473 assert!(
475 store.ensure_table(&table, 3).unwrap(),
476 "table should be new"
477 );
478 assert!(
479 !store.ensure_table(&table, 3).unwrap(),
480 "second call sees it"
481 );
482
483 let mut index = BM25Index::new();
484 index
485 .chunks
486 .push(chunk("src/a.rs", "alpha", 1, 5, ChunkKind::Function));
487 index
488 .chunks
489 .push(chunk("src/b.rs", "beta", 10, 20, ChunkKind::Struct));
490 let embeddings = vec![vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0]];
491
492 store
493 .sync_index(&table, &index, &embeddings, &[], true)
494 .unwrap();
495
496 let hits = store.search(&table, &[1.0, 0.0, 0.0], 2).unwrap();
497 assert_eq!(hits.len(), 2);
498 assert_eq!(hits[0].file_path, "src/a.rs");
499 assert_eq!(hits[0].symbol_name, "alpha");
500 assert!(hits[0].score > 0.99, "cosine sim of identical vec ~ 1.0");
501 assert_eq!(hits[0].kind, ChunkKind::Function);
502
503 index.chunks[1] = chunk("src/b.rs", "beta", 30, 40, ChunkKind::Struct);
505 store
506 .sync_index(
507 &table,
508 &index,
509 &embeddings,
510 &["src/b.rs".to_string()],
511 false,
512 )
513 .unwrap();
514
515 let hits = store.search(&table, &[0.0, 1.0, 0.0], 2).unwrap();
516 assert_eq!(hits[0].file_path, "src/b.rs");
517 assert_eq!(hits[0].start_line, 30, "stale row was replaced");
518 assert_eq!(hits[0].end_line, 40);
519
520 index
522 .chunks
523 .push(chunk("src/it's.rs", "q'uote", 2, 3, ChunkKind::Method));
524 let embeddings = vec![
525 vec![1.0, 0.0, 0.0],
526 vec![0.0, 1.0, 0.0],
527 vec![0.0, 0.0, 1.0],
528 ];
529 store
530 .sync_index(
531 &table,
532 &index,
533 &embeddings,
534 &["src/it's.rs".to_string()],
535 false,
536 )
537 .unwrap();
538 let hits = store.search(&table, &[0.0, 0.0, 1.0], 1).unwrap();
539 assert_eq!(hits[0].file_path, "src/it's.rs");
540 assert_eq!(hits[0].symbol_name, "q'uote");
541
542 store.run_sql(&format!("DROP TABLE {table};")).unwrap();
543 }
544
545 #[test]
546 fn table_name_is_namespaced_and_bounded() {
547 let _env = crate::core::data_dir::test_env_lock();
548 crate::test_env::set_var("LEANCTX_PGVECTOR_URL", "postgres://localhost/lctx");
549 let store = PgvectorStore::from_env().unwrap();
550 let name = store
551 .table_name(Path::new("/tmp/some-project"), 384)
552 .unwrap();
553 assert!(name.starts_with("lctx_code_"));
554 assert!(name.ends_with("_d384"));
555 assert!(name.len() <= 63);
556 crate::test_env::remove_var("LEANCTX_PGVECTOR_URL");
557 }
558}