Skip to main content

alopex_cli/commands/
hnsw.rs

1//! HNSW Command - HNSW index management
2//!
3//! Supports: create, stats, drop
4
5use std::io::Write;
6use std::path::PathBuf;
7
8use alopex_embedded::{Database, HnswConfig, Metric};
9use serde::{Deserialize, Serialize};
10
11use crate::batch::BatchMode;
12use crate::cli::{DistanceMetric, HnswCommand, OutputFormat};
13use crate::client::http::{ClientError, HttpClient};
14use crate::error::Result;
15use crate::models::{Column, DataType, Row, Value};
16use crate::output::formatter::Formatter;
17use crate::output::RowCollector;
18use crate::streaming::StreamingWriter;
19use crate::tui::admin::{AdminBackend, AdminContext, AdminTarget, AuthCapabilities};
20use crate::tui::renderer::render_output;
21
22/// Default M parameter (max connections per node)
23const DEFAULT_M: usize = 16;
24/// Default ef_construction parameter
25const DEFAULT_EF_CONSTRUCTION: usize = 200;
26
27#[derive(Debug, Serialize)]
28struct RemoteHnswCreateRequest {
29    index: String,
30    dim: usize,
31    metric: String,
32}
33
34#[derive(Debug, Serialize)]
35struct RemoteHnswDropRequest {
36    index: String,
37}
38
39#[derive(Debug, Serialize)]
40struct RemoteHnswStatsRequest {
41    index: String,
42}
43
44#[derive(Debug, Deserialize)]
45struct RemoteHnswStatsResponse {
46    stats: RemoteHnswStats,
47}
48
49#[derive(Debug, Deserialize)]
50struct RemoteHnswStats {
51    node_count: u64,
52    deleted_count: u64,
53    memory_bytes: u64,
54    avg_edges_per_node: f64,
55}
56
57#[derive(Debug, Deserialize)]
58struct RemoteHnswStatusResponse {
59    success: bool,
60}
61
62/// Execute an HNSW command.
63///
64/// # Arguments
65///
66/// * `db` - The database instance.
67/// * `cmd` - The HNSW subcommand to execute.
68/// * `writer` - The streaming writer for output.
69pub fn execute<W: Write>(
70    db: &Database,
71    cmd: HnswCommand,
72    writer: &mut StreamingWriter<W>,
73) -> Result<()> {
74    match cmd {
75        HnswCommand::Create { name, dim, metric } => execute_create(db, &name, dim, metric, writer),
76        HnswCommand::Stats { name } => execute_stats(db, &name, writer),
77        HnswCommand::Drop { name } => execute_drop(db, &name, writer),
78    }
79}
80
81#[allow(clippy::too_many_arguments)]
82pub fn execute_tui(
83    db: &Database,
84    cmd: HnswCommand,
85    batch_mode: &BatchMode,
86    output_format: OutputFormat,
87    columns: Vec<Column>,
88    limit: Option<usize>,
89    quiet: bool,
90    connection_label: impl Into<String>,
91    data_dir: Option<PathBuf>,
92) -> Result<()> {
93    let connection_label = connection_label.into();
94    let context_message = Some(hnsw_command_context(&cmd));
95    let admin_label = connection_label.clone();
96    let admin_data_dir = data_dir.clone();
97    let admin_launcher: Option<Box<dyn FnMut() -> Result<()> + '_>> = Some(Box::new(move || {
98        let connection_label = admin_label.clone();
99        let data_dir = admin_data_dir.clone();
100        crate::tui::admin::run_admin_ui(AdminContext {
101            connection_label,
102            auth: AuthCapabilities::full(),
103            backend: AdminBackend::Local {
104                db,
105                batch_mode,
106                output_format,
107                limit,
108                quiet,
109                data_dir,
110            },
111            initial_target: Some(AdminTarget::Hnsw),
112        })
113    }));
114    let collector = RowCollector::new();
115    let formatter = Box::new(collector.formatter());
116    let mut sink = std::io::sink();
117    let mut writer =
118        StreamingWriter::new(&mut sink, formatter, columns.clone(), limit).with_quiet(quiet);
119    execute(db, cmd, &mut writer)?;
120    let warning = collector.truncation_warning();
121    render_output(
122        columns,
123        collector.rows(),
124        connection_label,
125        context_message,
126        true,
127        warning,
128        output_format,
129        admin_launcher,
130    )
131}
132
133/// Execute an HNSW command against a remote server.
134pub async fn execute_remote_with_formatter<W: Write>(
135    client: &HttpClient,
136    cmd: &HnswCommand,
137    writer: &mut W,
138    formatter: Box<dyn Formatter>,
139    limit: Option<usize>,
140    quiet: bool,
141) -> Result<()> {
142    match cmd {
143        HnswCommand::Create { name, dim, metric } => {
144            execute_remote_create(client, name, *dim, *metric, writer, formatter, limit, quiet)
145                .await
146        }
147        HnswCommand::Stats { name } => {
148            execute_remote_stats(client, name, writer, formatter, limit, quiet).await
149        }
150        HnswCommand::Drop { name } => {
151            execute_remote_drop(client, name, writer, formatter, limit, quiet).await
152        }
153    }
154}
155
156#[allow(clippy::too_many_arguments)]
157pub async fn execute_remote_tui<'a>(
158    client: &HttpClient,
159    cmd: &HnswCommand,
160    columns: Vec<Column>,
161    output_format: OutputFormat,
162    limit: Option<usize>,
163    quiet: bool,
164    connection_label: impl Into<String>,
165    admin_launcher: Option<Box<dyn FnMut() -> Result<()> + 'a>>,
166) -> Result<()> {
167    let collector = RowCollector::new();
168    let formatter = Box::new(collector.formatter());
169    let mut sink = std::io::sink();
170    execute_remote_with_formatter(client, cmd, &mut sink, formatter, limit, quiet).await?;
171    let warning = collector.truncation_warning();
172    render_output(
173        columns,
174        collector.rows(),
175        connection_label,
176        Some(hnsw_command_context(cmd)),
177        true,
178        warning,
179        output_format,
180        admin_launcher,
181    )
182}
183
184/// Convert CLI distance metric to embedded Metric.
185fn to_embedded_metric(metric: DistanceMetric) -> Metric {
186    match metric {
187        DistanceMetric::Cosine => Metric::Cosine,
188        DistanceMetric::L2 => Metric::L2,
189        DistanceMetric::Ip => Metric::InnerProduct,
190    }
191}
192
193fn metric_to_string(metric: DistanceMetric) -> String {
194    match metric {
195        DistanceMetric::Cosine => "cosine".to_string(),
196        DistanceMetric::L2 => "l2".to_string(),
197        DistanceMetric::Ip => "ip".to_string(),
198    }
199}
200
201#[allow(clippy::too_many_arguments)]
202async fn execute_remote_create<W: Write>(
203    client: &HttpClient,
204    name: &str,
205    dim: usize,
206    metric: DistanceMetric,
207    writer: &mut W,
208    formatter: Box<dyn Formatter>,
209    limit: Option<usize>,
210    quiet: bool,
211) -> Result<()> {
212    let request = RemoteHnswCreateRequest {
213        index: name.to_string(),
214        dim,
215        metric: metric_to_string(metric),
216    };
217    let response: RemoteHnswStatusResponse = client
218        .post_json("hnsw/create", &request)
219        .await
220        .map_err(map_client_error)?;
221    if response.success {
222        if quiet {
223            return Ok(());
224        }
225        let columns = hnsw_status_columns();
226        let mut streaming_writer =
227            StreamingWriter::new(writer, formatter, columns, limit).with_quiet(quiet);
228        streaming_writer.prepare(Some(1))?;
229        let row = Row::new(vec![
230            Value::Text("OK".to_string()),
231            Value::Text(format!("Created HNSW index: {}", name)),
232        ]);
233        streaming_writer.write_row(row)?;
234        streaming_writer.finish()
235    } else {
236        Err(crate::error::CliError::InvalidArgument(
237            "Failed to create HNSW index".to_string(),
238        ))
239    }
240}
241
242#[allow(clippy::too_many_arguments)]
243async fn execute_remote_stats<W: Write>(
244    client: &HttpClient,
245    name: &str,
246    writer: &mut W,
247    formatter: Box<dyn Formatter>,
248    limit: Option<usize>,
249    quiet: bool,
250) -> Result<()> {
251    let request = RemoteHnswStatsRequest {
252        index: name.to_string(),
253    };
254    let response: RemoteHnswStatsResponse = client
255        .post_json("hnsw/stats", &request)
256        .await
257        .map_err(map_client_error)?;
258    let stats = response.stats;
259    let columns = hnsw_stats_columns();
260    let mut streaming_writer =
261        StreamingWriter::new(writer, formatter, columns, limit).with_quiet(quiet);
262    streaming_writer.prepare(Some(4))?;
263    let stats_rows = vec![
264        ("node_count", Value::Int(stats.node_count as i64)),
265        ("deleted_count", Value::Int(stats.deleted_count as i64)),
266        ("memory_bytes", Value::Int(stats.memory_bytes as i64)),
267        ("avg_edges_per_node", Value::Float(stats.avg_edges_per_node)),
268    ];
269    for (key, value) in stats_rows {
270        let row = Row::new(vec![Value::Text(key.to_string()), value]);
271        streaming_writer.write_row(row)?;
272    }
273    streaming_writer.finish()
274}
275
276#[allow(clippy::too_many_arguments)]
277async fn execute_remote_drop<W: Write>(
278    client: &HttpClient,
279    name: &str,
280    writer: &mut W,
281    formatter: Box<dyn Formatter>,
282    limit: Option<usize>,
283    quiet: bool,
284) -> Result<()> {
285    let request = RemoteHnswDropRequest {
286        index: name.to_string(),
287    };
288    let response: RemoteHnswStatusResponse = client
289        .post_json("hnsw/drop", &request)
290        .await
291        .map_err(map_client_error)?;
292    if response.success {
293        if quiet {
294            return Ok(());
295        }
296        let columns = hnsw_status_columns();
297        let mut streaming_writer =
298            StreamingWriter::new(writer, formatter, columns, limit).with_quiet(quiet);
299        streaming_writer.prepare(Some(1))?;
300        let row = Row::new(vec![
301            Value::Text("OK".to_string()),
302            Value::Text(format!("Dropped HNSW index: {}", name)),
303        ]);
304        streaming_writer.write_row(row)?;
305        streaming_writer.finish()
306    } else {
307        Err(crate::error::CliError::InvalidArgument(
308            "Failed to drop HNSW index".to_string(),
309        ))
310    }
311}
312
313fn map_client_error(err: ClientError) -> crate::error::CliError {
314    match err {
315        ClientError::Request { source, .. } => {
316            crate::error::CliError::ServerConnection(format!("request failed: {source}"))
317        }
318        ClientError::InvalidUrl(message) => crate::error::CliError::InvalidArgument(message),
319        ClientError::Build(message) => crate::error::CliError::InvalidArgument(message),
320        ClientError::Auth(err) => crate::error::CliError::InvalidArgument(err.to_string()),
321        ClientError::HttpStatus { status, body } => crate::error::CliError::InvalidArgument(
322            format!("Server error: HTTP {} - {}", status.as_u16(), body),
323        ),
324    }
325}
326
327fn hnsw_command_context(cmd: &HnswCommand) -> String {
328    match cmd {
329        HnswCommand::Create { name, dim, metric } => format!(
330            "hnsw create {name} --dim {dim} --metric {}",
331            metric_to_string(*metric)
332        ),
333        HnswCommand::Stats { name } => format!("hnsw stats {name}"),
334        HnswCommand::Drop { name } => format!("hnsw drop {name}"),
335    }
336}
337
338/// Execute an HNSW create command.
339fn execute_create<W: Write>(
340    db: &Database,
341    name: &str,
342    dim: usize,
343    metric: DistanceMetric,
344    writer: &mut StreamingWriter<W>,
345) -> Result<()> {
346    let config = HnswConfig {
347        dimension: dim,
348        metric: to_embedded_metric(metric),
349        m: DEFAULT_M,
350        ef_construction: DEFAULT_EF_CONSTRUCTION,
351    };
352
353    db.create_hnsw_index(name, config)?;
354
355    // Suppress status output in quiet mode
356    if !writer.is_quiet() {
357        writer.prepare(Some(1))?;
358        let row = Row::new(vec![
359            Value::Text("OK".to_string()),
360            Value::Text(format!("Created HNSW index: {}", name)),
361        ]);
362        writer.write_row(row)?;
363        writer.finish()?;
364    }
365
366    Ok(())
367}
368
369/// Execute an HNSW stats command.
370fn execute_stats<W: Write>(
371    db: &Database,
372    name: &str,
373    writer: &mut StreamingWriter<W>,
374) -> Result<()> {
375    let stats = db.get_hnsw_stats(name)?;
376
377    // Output stats as rows
378    writer.prepare(Some(4))?;
379
380    // Output each stat as a row
381    let stats_rows = vec![
382        ("node_count", Value::Int(stats.node_count as i64)),
383        ("deleted_count", Value::Int(stats.deleted_count as i64)),
384        ("memory_bytes", Value::Int(stats.memory_bytes as i64)),
385        ("avg_edges_per_node", Value::Float(stats.avg_edges_per_node)),
386    ];
387
388    for (key, value) in stats_rows {
389        let row = Row::new(vec![Value::Text(key.to_string()), value]);
390        writer.write_row(row)?;
391    }
392
393    writer.finish()?;
394    Ok(())
395}
396
397/// Execute an HNSW drop command.
398fn execute_drop<W: Write>(
399    db: &Database,
400    name: &str,
401    writer: &mut StreamingWriter<W>,
402) -> Result<()> {
403    db.drop_hnsw_index(name)?;
404
405    // Suppress status output in quiet mode
406    if !writer.is_quiet() {
407        writer.prepare(Some(1))?;
408        let row = Row::new(vec![
409            Value::Text("OK".to_string()),
410            Value::Text(format!("Dropped HNSW index: {}", name)),
411        ]);
412        writer.write_row(row)?;
413        writer.finish()?;
414    }
415
416    Ok(())
417}
418
419/// Create columns for HNSW stats output.
420///
421/// Note: value column is Text because stats include both integer and float values.
422pub fn hnsw_stats_columns() -> Vec<Column> {
423    vec![
424        Column::new("property", DataType::Text),
425        Column::new("value", DataType::Text),
426    ]
427}
428
429/// Create columns for HNSW status output.
430pub fn hnsw_status_columns() -> Vec<Column> {
431    vec![
432        Column::new("status", DataType::Text),
433        Column::new("message", DataType::Text),
434    ]
435}
436
437#[cfg(test)]
438mod tests {
439    use super::*;
440    use crate::output::jsonl::JsonlFormatter;
441
442    fn create_test_db() -> Database {
443        Database::open_in_memory().unwrap()
444    }
445
446    fn create_stats_writer(output: &mut Vec<u8>) -> StreamingWriter<&mut Vec<u8>> {
447        let formatter = Box::new(JsonlFormatter::new());
448        let columns = hnsw_stats_columns();
449        StreamingWriter::new(output, formatter, columns, None)
450    }
451
452    fn create_status_writer(output: &mut Vec<u8>) -> StreamingWriter<&mut Vec<u8>> {
453        let formatter = Box::new(JsonlFormatter::new());
454        let columns = hnsw_status_columns();
455        StreamingWriter::new(output, formatter, columns, None)
456    }
457
458    #[test]
459    fn test_create_hnsw_index() {
460        let db = create_test_db();
461
462        let mut output = Vec::new();
463        {
464            let mut writer = create_status_writer(&mut output);
465            execute_create(&db, "test_index", 128, DistanceMetric::Cosine, &mut writer).unwrap();
466        }
467
468        let result = String::from_utf8(output).unwrap();
469        assert!(result.contains("OK"));
470        assert!(result.contains("Created HNSW index"));
471    }
472
473    #[test]
474    fn test_create_hnsw_index_l2() {
475        let db = create_test_db();
476
477        let mut output = Vec::new();
478        {
479            let mut writer = create_status_writer(&mut output);
480            execute_create(&db, "l2_index", 64, DistanceMetric::L2, &mut writer).unwrap();
481        }
482
483        let result = String::from_utf8(output).unwrap();
484        assert!(result.contains("OK"));
485        assert!(result.contains("Created HNSW index: l2_index"));
486    }
487
488    #[test]
489    fn test_get_hnsw_stats() {
490        let db = create_test_db();
491
492        // Create index first
493        {
494            let mut output = Vec::new();
495            let mut writer = create_status_writer(&mut output);
496            execute_create(&db, "stats_test", 64, DistanceMetric::Cosine, &mut writer).unwrap();
497        }
498
499        // Get stats
500        let mut output = Vec::new();
501        {
502            let mut writer = create_stats_writer(&mut output);
503            execute_stats(&db, "stats_test", &mut writer).unwrap();
504        }
505
506        let result = String::from_utf8(output).unwrap();
507        assert!(result.contains("node_count"));
508        assert!(result.contains("memory_bytes"));
509    }
510
511    #[test]
512    fn test_drop_hnsw_index() {
513        let db = create_test_db();
514
515        // Create index first
516        {
517            let mut output = Vec::new();
518            let mut writer = create_status_writer(&mut output);
519            execute_create(&db, "drop_test", 32, DistanceMetric::Cosine, &mut writer).unwrap();
520        }
521
522        // Drop index
523        let mut output = Vec::new();
524        {
525            let mut writer = create_status_writer(&mut output);
526            execute_drop(&db, "drop_test", &mut writer).unwrap();
527        }
528
529        let result = String::from_utf8(output).unwrap();
530        assert!(result.contains("OK"));
531        assert!(result.contains("Dropped HNSW index"));
532    }
533
534    #[test]
535    fn test_create_and_query_index() {
536        let db = create_test_db();
537
538        // Create index
539        {
540            let mut output = Vec::new();
541            let mut writer = create_status_writer(&mut output);
542            execute_create(&db, "query_test", 3, DistanceMetric::Cosine, &mut writer).unwrap();
543        }
544
545        // Stats should show 0 vectors initially
546        let mut output = Vec::new();
547        {
548            let mut writer = create_stats_writer(&mut output);
549            execute_stats(&db, "query_test", &mut writer).unwrap();
550        }
551
552        let result = String::from_utf8(output).unwrap();
553        assert!(result.contains("node_count"));
554    }
555}