1use 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
22const DEFAULT_M: usize = 16;
24const 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
62pub 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
133pub 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
184fn 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
338fn 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 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
369fn 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 writer.prepare(Some(4))?;
379
380 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
397fn 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 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
419pub fn hnsw_stats_columns() -> Vec<Column> {
423 vec![
424 Column::new("property", DataType::Text),
425 Column::new("value", DataType::Text),
426 ]
427}
428
429pub 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 {
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 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 {
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 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 {
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 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}