1use std::path::PathBuf;
4
5use clap::{Parser, Subcommand};
6use serde_json::json;
7
8pub use crate::backup::BackupCommands;
9pub use crate::daemon::DaemonCommands;
10pub use crate::hook::HookCommands;
11use crate::{AppraiseParams, RecallParams, RecordParams, Situation, APPRAISE_ADVISORY};
12
13fn default_db() -> PathBuf {
14 crate::paths::default_db_path()
15}
16
17#[derive(Parser)]
18#[command(name = "innate", version, about = "Self-growing knowledge layer")]
19pub struct Cli {
20 #[arg(long, global = true, env = "INNATE_DB")]
21 pub db: Option<PathBuf>,
22
23 #[command(subcommand)]
24 pub command: Commands,
25}
26
27#[derive(Subcommand)]
28pub enum Commands {
29 Recall {
31 query: String,
32 #[arg(long, default_value = "6000")]
33 budget: usize,
34 #[arg(long)]
35 top: Option<usize>,
36 #[arg(long, default_value = "text")]
37 format: String,
38 #[arg(long)]
39 include_sparks: bool,
40 #[arg(long, default_value = "false")]
42 expand_deps: String,
43 #[arg(long)]
45 allow_trim: bool,
46 #[arg(long, default_value = "off")]
48 refine_mode: String,
49 #[arg(long, default_value = "cli")]
51 source: String,
52 #[arg(long)]
55 min_score: Option<f64>,
56 #[arg(long)]
60 session: bool,
61 #[arg(long)]
64 rerank: bool,
65 },
66 Appraise {
69 #[arg(long, default_value = "")]
71 query: String,
72 #[arg(long)]
74 last_error: Option<String>,
75 #[arg(long)]
77 recent_actions: Option<String>,
78 #[arg(long)]
80 stage: Option<String>,
81 #[arg(long)]
83 file_context: Option<String>,
84 #[arg(long)]
86 candidate: Option<String>,
87 #[arg(long)]
88 top: Option<usize>,
89 #[arg(long)]
90 min_strength: Option<f64>,
91 #[arg(long, default_value = "cli")]
92 source: String,
93 #[arg(long, default_value = "json")]
94 format: String,
95 },
96 Record {
98 trace_id: String,
99 #[arg(long)]
100 query: Option<String>,
101 #[arg(long)]
102 outcome: Option<String>,
103 #[arg(long)]
105 used: Option<String>,
106 #[arg(long, default_value = "explicit")]
107 used_attribution: String,
108 #[arg(long)]
110 used_partial: bool,
111 #[arg(long)]
112 output: Option<String>,
113 #[arg(long)]
114 output_summary: Option<String>,
115 #[arg(long)]
116 nomination: Option<String>,
117 #[arg(long, default_value = "cli")]
118 source: String,
119 #[arg(long)]
121 feedback: Option<String>,
122 #[arg(long, default_value = "user")]
123 feedback_kind: String,
124 #[arg(long)]
125 feedback_actor: Option<String>,
126 #[arg(long)]
127 feedback_reason: Option<String>,
128 #[arg(long)]
129 task_state: Option<String>,
130 #[arg(long, default_value = "0")]
131 priority: i64,
132 #[arg(long)]
136 verdict_heeded: bool,
137 },
138 Add {
140 content: String,
141 #[arg(long, default_value = "note")]
142 kind: String,
143 #[arg(long)]
144 trigger: Option<String>,
145 #[arg(long)]
146 anti_trigger: Option<String>,
147 #[arg(long, default_value = "chat")]
148 source: String,
149 #[arg(long)]
150 skill_name: Option<String>,
151 #[arg(long = "depends-on")]
153 depends_on: Vec<String>,
154 #[arg(long, default_value = "hard")]
156 dep_kind: String,
157 },
158 Spark {
160 content: String,
161 #[arg(long)]
162 trigger: Option<String>,
163 },
164 Evolve {
166 #[arg(long, default_value = "manual")]
167 trigger: String,
168 #[arg(long)]
170 rebuild_embeddings: bool,
171 },
172 Inspect { id: Option<String> },
174 Approve { chunk_id: String },
176 Archive {
178 chunk_id: String,
179 #[arg(long, default_value = "stale")]
180 reason: String,
181 },
182 Invalidate {
184 chunk_id: String,
185 #[arg(long, default_value = "")]
186 reason: String,
187 },
188 Restore { chunk_id: String },
190 MatureSpark { spark_id: String, to: String },
192 PromoteSpark {
194 spark_id: String,
195 #[arg(long, default_value = "note")]
196 to: String,
197 },
198 DropSpark {
200 spark_id: String,
201 #[arg(long, default_value = "")]
202 reason: String,
203 },
204 Backup {
206 #[command(subcommand)]
207 action: BackupCommands,
208 },
209 Install,
211 Uninstall {
213 #[arg(long, short = 'y')]
215 yes: bool,
216 #[arg(long)]
218 purge_data: bool,
219 },
220 Migrate,
222 Metrics {
224 #[command(subcommand)]
225 action: MetricsAction,
226 },
227 Vacuum,
229 Export {
231 #[arg(long)]
233 out: Option<PathBuf>,
234 #[arg(long)]
236 include_archived: bool,
237 },
238 #[command(alias = "ingest")]
242 Import {
243 file: PathBuf,
245 #[arg(long, default_value = "doc")]
247 source: String,
248 },
249 RepairTraces {
252 #[arg(long)]
254 dry_run: bool,
255 },
256 RecallEval {
261 labels: PathBuf,
263 #[arg(long, default_value = "10")]
265 k: usize,
266 #[arg(long)]
269 save: bool,
270 },
271 Upgrade {
273 #[arg(long, value_name = "VERSION")]
275 version: Option<String>,
276 #[arg(long)]
278 check: bool,
279 },
280 Daemon {
282 #[command(subcommand)]
283 action: DaemonCommands,
284 },
285 Mcp,
287 Web {
289 #[arg(long, default_value = "127.0.0.1")]
291 bind: String,
292 #[arg(long, default_value_t = 8788)]
294 port: u16,
295 #[arg(long)]
297 no_token: bool,
298 #[arg(long)]
301 allow_remote: bool,
302 },
303 Hook {
305 #[command(subcommand)]
306 action: HookCommands,
307 },
308}
309
310#[derive(clap::Subcommand)]
311pub enum MetricsAction {
312 Snapshot,
314}
315
316pub fn run() -> anyhow::Result<()> {
317 let cli = Cli::parse();
318 crate::paths::ensure_layout();
321 let db_path = cli.db.unwrap_or_else(default_db);
322
323 if let Commands::Mcp = &cli.command {
324 return crate::mcp::run_server(db_path);
325 }
326
327 if let Commands::Install = &cli.command {
328 return crate::install::run_install();
329 }
330
331 if let Commands::Uninstall { yes, purge_data } = &cli.command {
332 return crate::install::run_uninstall(*yes, *purge_data);
333 }
334
335 if let Commands::Migrate = &cli.command {
336 let applied = crate::migrate::run_migrations(&db_path)?;
337 if applied.is_empty() {
338 println!(
339 "already at {} — nothing to do",
340 crate::migrate::target_version()
341 );
342 } else {
343 for step in &applied {
344 println!(" applied: {step}");
345 }
346 println!("migration complete");
347 }
348 return Ok(());
349 }
350
351 if let Commands::Daemon { action } = &cli.command {
352 return crate::daemon::run_command(action, &db_path);
353 }
354
355 if let Commands::Backup { action } = &cli.command {
356 return crate::backup::run_command(action, &db_path);
357 }
358
359 if let Commands::Upgrade { version, check } = &cli.command {
360 return crate::upgrade::run_upgrade(version.as_deref(), &db_path, *check);
361 }
362
363 if let Commands::Hook { action } = &cli.command {
364 return crate::hook::run_command(action, &db_path);
365 }
366
367 let kb = crate::open_kb(&db_path)?;
368
369 match cli.command {
370 Commands::Recall {
371 query,
372 budget,
373 top,
374 format,
375 include_sparks,
376 expand_deps,
377 allow_trim,
378 refine_mode,
379 source,
380 min_score,
381 session,
382 rerank,
383 } => {
384 let result = kb.recall(RecallParams {
385 query: &query,
386 budget,
387 trace: true,
388 include_sparks,
389 top,
390 source: &source,
391 expand_deps: &expand_deps,
392 allow_trim,
393 refine_mode: &refine_mode,
394 min_score,
395 session_only: session,
396 rerank,
397 })?;
398 match format.as_str() {
399 "json" => println!(
400 "{}",
401 serde_json::to_string_pretty(&json!({
402 "trace_id": result.trace_id,
403 "knowledge": result.knowledge,
404 "sparks": result.sparks,
405 "empty": result.empty,
406 }))?
407 ),
408 "prompt" => {
409 for chunk in &result.knowledge {
410 let content = chunk.get("content").and_then(|v| v.as_str()).unwrap_or("");
411 println!("{content}\n---");
412 }
413 println!("<!-- innate_trace_id: {} -->", result.trace_id);
415 println!(
416 "<!-- innate_selected: {} -->",
417 result
418 .knowledge
419 .iter()
420 .filter_map(|c| c.get("id").and_then(|v| v.as_str()))
421 .collect::<Vec<_>>()
422 .join(",")
423 );
424 }
425 _ => {
426 for chunk in &result.knowledge {
427 let id = chunk.get("id").and_then(|v| v.as_str()).unwrap_or("?");
428 let content = chunk.get("content").and_then(|v| v.as_str()).unwrap_or("");
429 let conf = chunk
430 .get("confidence")
431 .and_then(|v| v.as_f64())
432 .unwrap_or(0.5);
433 println!("[{id}] (conf={conf:.2})\n{content}\n");
434 }
435 if result.empty {
436 println!("(no results)");
437 }
438 }
439 }
440 }
441 Commands::Appraise {
442 query,
443 last_error,
444 recent_actions,
445 stage,
446 file_context,
447 candidate,
448 top,
449 min_strength,
450 source,
451 format,
452 } => {
453 let actions: Vec<String> = recent_actions
454 .as_deref()
455 .map(|raw| {
456 raw.split(',')
457 .map(str::trim)
458 .filter(|a| !a.is_empty())
459 .map(str::to_string)
460 .collect()
461 })
462 .unwrap_or_default();
463 let situation = Situation {
464 query: (!query.is_empty()).then_some(query.as_str()),
465 last_error: last_error.as_deref(),
466 recent_actions: &actions,
467 stage: stage.as_deref(),
468 file_context: file_context.as_deref(),
469 };
470 let verdict = kb.appraise(AppraiseParams {
471 situation,
472 candidate: candidate.as_deref(),
473 min_strength,
474 top,
475 trace: true,
476 source: &source,
477 })?;
478 match format.as_str() {
479 "text" => {
480 println!("ℹ {APPRAISE_ADVISORY}");
481 if verdict.abstained {
482 println!(
483 "ABSTAIN reason={:?} strength={:.3} trace_id={}",
484 verdict.abstain_reason, verdict.strength, verdict.trace_id
485 );
486 } else {
487 println!(
488 "valence={:?} tier={:?} strength={:.3} confidence={:.3} dispersion={:.3} trace_id={}",
489 verdict.valence, verdict.tier, verdict.strength,
490 verdict.confidence, verdict.dispersion, verdict.trace_id
491 );
492 }
493 for fp in &verdict.flagged_points {
494 println!(
495 " ⚠ [{}] {} (s={:.3})",
496 fp.chunk_id, fp.summary, fp.strength
497 );
498 }
499 }
500 _ => println!(
501 "{}",
502 serde_json::to_string_pretty(&json!({
503 "advisory": APPRAISE_ADVISORY,
504 "valence": verdict.valence,
505 "strength": verdict.strength,
506 "tier": verdict.tier,
507 "confidence": verdict.confidence,
508 "dispersion": verdict.dispersion,
509 "abstained": verdict.abstained,
510 "abstain_reason": verdict.abstain_reason,
511 "flagged_points": verdict.flagged_points,
512 "contributors": verdict.contributors,
513 "trace_id": verdict.trace_id,
514 }))?
515 ),
516 }
517 }
518 Commands::Record {
519 trace_id,
520 query,
521 outcome,
522 used,
523 used_attribution,
524 used_partial,
525 output,
526 output_summary,
527 nomination,
528 source,
529 feedback,
530 feedback_kind,
531 feedback_actor,
532 feedback_reason,
533 task_state,
534 priority,
535 verdict_heeded,
536 } => {
537 let used_ids = used.as_deref().map(|raw| {
538 raw.split(',')
539 .map(str::trim)
540 .filter(|id| !id.is_empty())
541 .map(str::to_string)
542 .collect::<Vec<_>>()
543 });
544 let used_ref = used_ids.as_deref();
545 let (fb_up, fb_down): (Option<Vec<String>>, Option<Vec<String>>) =
547 match feedback.as_deref() {
548 Some("up") if used_ids.as_ref().is_some_and(|ids| !ids.is_empty()) => {
549 (used_ids.clone(), None)
550 }
551 Some("down") if used_ids.as_ref().is_some_and(|ids| !ids.is_empty()) => {
552 (None, used_ids.clone())
553 }
554 Some("up") => (None, None), Some("down") => (None, None),
556 _ => (None, None),
557 };
558 let fb_up_ref = fb_up.as_deref();
559 let fb_down_ref = fb_down.as_deref();
560 kb.record(RecordParams {
561 trace_id: &trace_id,
562 query: query.as_deref(),
563 output: output.as_deref(),
564 output_summary: output_summary.as_deref(),
565 outcome: outcome.as_deref(),
566 used: used_ref,
567 used_attribution: &used_attribution,
568 used_complete: Some(!used_partial),
569 feedback_up: fb_up_ref,
570 feedback_down: fb_down_ref,
571 feedback_kind: &feedback_kind,
572 feedback_actor: feedback_actor.as_deref(),
573 feedback_reason: feedback_reason.as_deref(),
574 nomination: nomination.as_deref(),
575 priority,
576 task_state: task_state.as_deref(),
577 source: &source,
578 verdict_heeded,
579 })?;
580 println!("recorded");
581 }
582 Commands::Add {
583 content,
584 kind,
585 trigger,
586 anti_trigger,
587 source,
588 skill_name,
589 depends_on,
590 dep_kind,
591 } => {
592 let content = if kind == "skill" {
594 let p = std::path::Path::new(&content);
595 if p.exists() && p.is_file() {
596 std::fs::read_to_string(p).map_err(|e| {
597 anyhow::anyhow!("Failed to read skill file {}: {e}", p.display())
598 })?
599 } else {
600 content
601 }
602 } else {
603 content
604 };
605 let deps: Vec<(String, String)> = depends_on
606 .iter()
607 .map(|d| (d.clone(), dep_kind.clone()))
608 .collect();
609 let id = kb.add_with_deps(
610 &content,
611 &kind,
612 trigger.as_deref(),
613 anti_trigger.as_deref(),
614 &source,
615 skill_name.as_deref(),
616 &deps,
617 )?;
618 println!("{id}");
619 }
620 Commands::Spark { content, trigger } => {
621 let id = kb.spark(&content, trigger.as_deref(), None)?;
622 println!("{id}");
623 }
624 Commands::Evolve {
625 trigger,
626 rebuild_embeddings,
627 } => {
628 if rebuild_embeddings {
629 let rebuilt = kb.rebuild_embeddings()?;
630 let report = kb.evolve(&trigger)?;
631 println!(
632 "{}",
633 serde_json::to_string_pretty(&json!({
634 "rebuilt_embeddings": rebuilt,
635 "evolve": report
636 }))?
637 );
638 } else {
639 let report = kb.evolve(&trigger)?;
640 println!("{}", serde_json::to_string_pretty(&report)?);
641 }
642 }
643 Commands::Inspect { id } => match id.as_deref() {
644 None => {
645 let info = kb.inspect()?;
646 println!("{}", serde_json::to_string_pretty(&info)?);
647 }
648 Some(id) => {
649 let detail = kb.inspect_id(id)?;
650 println!("{}", serde_json::to_string_pretty(&detail)?);
651 }
652 },
653 Commands::Approve { chunk_id } => {
654 kb.approve(&chunk_id)?;
655 println!("approved");
656 }
657 Commands::Archive { chunk_id, reason } => {
658 kb.archive(&chunk_id, &reason)?;
659 println!("archived");
660 }
661 Commands::Invalidate { chunk_id, reason } => {
662 kb.invalidate(&chunk_id, &reason)?;
663 println!("invalidated");
664 }
665 Commands::Restore { chunk_id } => {
666 kb.restore(&chunk_id)?;
667 println!("restored");
668 }
669 Commands::MatureSpark { spark_id, to } => {
670 kb.mature_spark(&spark_id, &to)?;
671 println!("matured");
672 }
673 Commands::PromoteSpark { spark_id, to } => {
674 let id = kb.promote_spark(&spark_id, &to)?;
675 println!("{id}");
676 }
677 Commands::DropSpark { spark_id, reason } => {
678 kb.drop_spark(&spark_id, &reason)?;
679 println!("dropped");
680 }
681 Commands::Metrics { action } => match action {
682 MetricsAction::Snapshot => {
683 let kpis = kb.write_metric_snapshot()?;
684 println!("{}", serde_json::to_string_pretty(&kpis)?);
685 }
686 },
687 Commands::Vacuum => {
688 let (before, after) = kb.storage.vacuum()?;
689 let mb = |b: i64| b as f64 / 1_048_576.0;
690 println!(
691 "vacuumed: {:.2} MB → {:.2} MB (reclaimed {:.2} MB)",
692 mb(before),
693 mb(after),
694 mb(before - after)
695 );
696 }
697 Commands::Export {
698 out,
699 include_archived,
700 } => {
701 let rows = kb.storage.export_chunks(include_archived)?;
702 let mut buf = String::new();
703 for row in &rows {
704 buf.push_str(&serde_json::to_string(row)?);
705 buf.push('\n');
706 }
707 match out {
708 Some(path) => {
709 std::fs::write(&path, &buf)
710 .map_err(|e| anyhow::anyhow!("failed to write {}: {e}", path.display()))?;
711 eprintln!("exported {} chunk(s) → {}", rows.len(), path.display());
712 }
713 None => print!("{buf}"),
714 }
715 }
716 Commands::Import { file, source } => {
717 let text = std::fs::read_to_string(&file)
718 .map_err(|e| anyhow::anyhow!("failed to read {}: {e}", file.display()))?;
719 let mut known: std::collections::HashSet<String> = kb
723 .storage
724 .export_chunks(true)?
725 .into_iter()
726 .filter_map(|c| c.get("id").and_then(|v| v.as_str()).map(str::to_string))
727 .collect();
728 let mut imported = 0usize;
729 let mut skipped = 0usize;
730 let mut failed = 0usize;
731 for (lineno, line) in text.lines().enumerate() {
732 let line = line.trim();
733 if line.is_empty() {
734 continue;
735 }
736 let obj: serde_json::Value = match serde_json::from_str(line) {
737 Ok(v) => v,
738 Err(e) => {
739 eprintln!("line {}: invalid JSON ({e}), skipped", lineno + 1);
740 failed += 1;
741 continue;
742 }
743 };
744 let content = obj.get("content").and_then(|v| v.as_str()).unwrap_or("");
745 if content.trim().is_empty() {
746 eprintln!("line {}: missing/empty content, skipped", lineno + 1);
747 failed += 1;
748 continue;
749 }
750 let trigger = obj.get("trigger_desc").and_then(|v| v.as_str());
751 let anti = obj.get("anti_trigger_desc").and_then(|v| v.as_str());
752 let skill_name = obj.get("skill_name").and_then(|v| v.as_str());
753 let kind = if skill_name.is_some() { "skill" } else { "note" };
754 match kb.add(content, kind, trigger, anti, &source, skill_name) {
755 Ok(id) if id.is_empty() => skipped += 1,
757 Ok(id) if known.contains(&id) => skipped += 1,
760 Ok(id) => {
761 known.insert(id);
762 imported += 1;
763 }
764 Err(e) => {
765 eprintln!("line {}: {e}", lineno + 1);
766 failed += 1;
767 }
768 }
769 }
770 println!(
771 "{}",
772 serde_json::to_string_pretty(&json!({
773 "imported": imported,
774 "skipped": skipped,
775 "failed": failed,
776 }))?
777 );
778 }
779 Commands::RepairTraces { dry_run } => {
780 let r = kb.repair_traces(dry_run)?;
781 let tag = if dry_run {
782 "[dry-run] would repair"
783 } else {
784 "repaired"
785 };
786 println!(
787 "{tag}: deleted {} false daemon selection events, retired {} orphaned open logs, \
788 selected_count {} → {}",
789 r.daemon_events_deleted, r.open_logs_retired, r.selected_before, r.selected_after
790 );
791 }
792 Commands::RecallEval { labels, k, save } => {
793 let text = std::fs::read_to_string(&labels)
794 .map_err(|e| anyhow::anyhow!("read labels {}: {e}", labels.display()))?;
795 let mut n = 0usize;
796 let (mut sum_p1, mut sum_recall, mut sum_mrr, mut sum_ndcg) = (0.0, 0.0, 0.0, 0.0);
797 let mut misses: Vec<serde_json::Value> = Vec::new();
798 for (lineno, line) in text.lines().enumerate() {
799 let line = line.trim();
800 if line.is_empty() {
801 continue;
802 }
803 let row: serde_json::Value = serde_json::from_str(line)
804 .map_err(|e| anyhow::anyhow!("labels line {}: {e}", lineno + 1))?;
805 let query = row.get("query").and_then(|v| v.as_str()).unwrap_or("");
806 let relevant: std::collections::HashSet<String> = row
807 .get("relevant_ids")
808 .and_then(|v| v.as_array())
809 .map(|a| {
810 a.iter()
811 .filter_map(|v| v.as_str().map(str::to_string))
812 .collect()
813 })
814 .unwrap_or_default();
815 if query.is_empty() || relevant.is_empty() {
816 continue;
817 }
818 let result = kb.recall(RecallParams {
819 query,
820 budget: 100_000,
821 trace: false,
822 top: Some(k),
823 source: "cli",
824 ..Default::default()
825 })?;
826 let ranked: Vec<String> = result
827 .knowledge
828 .iter()
829 .filter_map(|c| c.get("id").and_then(|v| v.as_str()).map(str::to_string))
830 .collect();
831 let (p1, recall_k, mrr, ndcg) = recall_metrics(&ranked, &relevant, k);
832 sum_p1 += p1;
833 sum_recall += recall_k;
834 sum_mrr += mrr;
835 sum_ndcg += ndcg;
836 n += 1;
837 if recall_k == 0.0 {
840 misses.push(json!({
841 "query": query,
842 "relevant_ids": relevant.iter().cloned().collect::<Vec<_>>(),
843 "got_top_k": ranked,
844 }));
845 }
846 }
847 if n == 0 {
848 return Err(anyhow::anyhow!(
849 "no usable labeled queries (need lines with non-empty query + relevant_ids)"
850 ));
851 }
852 let nf = n as f64;
853 let out = json!({
854 "queries": n,
855 "k": k,
856 "p_at_1": (sum_p1 / nf * 1000.0).round() / 1000.0,
857 "recall_at_k": (sum_recall / nf * 1000.0).round() / 1000.0,
858 "mrr": (sum_mrr / nf * 1000.0).round() / 1000.0,
859 "ndcg_at_k": (sum_ndcg / nf * 1000.0).round() / 1000.0,
860 "params": kb.recall_weights(),
863 "misses": misses,
864 });
865 if save {
866 let mut summary = out.clone();
869 if let Some(o) = summary.as_object_mut() {
870 o.remove("misses");
871 o.insert("ts".to_string(), json!(crate::utils::utc_now_iso()));
872 }
873 let path = crate::paths::logs_dir().join("eval_runs.jsonl");
874 if let Some(parent) = path.parent() {
875 let _ = std::fs::create_dir_all(parent);
876 }
877 if let Ok(mut f) = std::fs::OpenOptions::new()
878 .create(true)
879 .append(true)
880 .open(&path)
881 {
882 use std::io::Write;
883 let _ = writeln!(f, "{}", serde_json::to_string(&summary)?);
884 eprintln!("eval run summary appended to {}", path.display());
885 }
886 }
887 println!("{}", serde_json::to_string_pretty(&out)?);
888 }
889 Commands::Web {
890 bind,
891 port,
892 no_token,
893 allow_remote,
894 } => {
895 let loopback = crate::web::is_loopback(&bind);
896 if !loopback && !allow_remote {
897 anyhow::bail!(
898 "refusing to bind non-loopback address {bind} without --allow-remote \
899 (this exposes the knowledge base to the network)"
900 );
901 }
902 if !loopback && no_token {
903 anyhow::bail!(
904 "--no-token cannot be combined with a non-loopback bind: a network-exposed \
905 server must keep the auth token to gate reads and writes"
906 );
907 }
908 crate::web::serve(kb, &bind, port, !no_token)?;
909 }
910 Commands::Mcp
911 | Commands::Install
912 | Commands::Uninstall { .. }
913 | Commands::Migrate
914 | Commands::Upgrade { .. }
915 | Commands::Daemon { .. }
916 | Commands::Backup { .. }
917 | Commands::Hook { .. } => unreachable!(),
918 }
919 Ok(())
920}
921
922pub(crate) fn recall_metrics(
927 ranked: &[String],
928 relevant: &std::collections::HashSet<String>,
929 k: usize,
930) -> (f64, f64, f64, f64) {
931 let topk = &ranked[..ranked.len().min(k)];
932 let p_at_1 = topk
933 .first()
934 .map(|id| relevant.contains(id) as u8 as f64)
935 .unwrap_or(0.0);
936 let hits = topk.iter().filter(|id| relevant.contains(*id)).count();
937 let recall_at_k = hits as f64 / relevant.len() as f64;
938 let mrr = ranked
940 .iter()
941 .position(|id| relevant.contains(id))
942 .map(|pos| 1.0 / (pos as f64 + 1.0))
943 .unwrap_or(0.0);
944 let dcg: f64 = topk
946 .iter()
947 .enumerate()
948 .filter(|(_, id)| relevant.contains(*id))
949 .map(|(i, _)| 1.0 / ((i as f64 + 2.0).log2()))
950 .sum();
951 let ideal_hits = relevant.len().min(k);
952 let idcg: f64 = (0..ideal_hits)
953 .map(|i| 1.0 / ((i as f64 + 2.0).log2()))
954 .sum();
955 let ndcg = if idcg > 0.0 { dcg / idcg } else { 0.0 };
956 (p_at_1, recall_at_k, mrr, ndcg)
957}
958
959#[cfg(test)]
960mod metric_tests {
961 use super::recall_metrics;
962 use std::collections::HashSet;
963
964 fn rel(ids: &[&str]) -> HashSet<String> {
965 ids.iter().map(|s| s.to_string()).collect()
966 }
967 fn ranked(ids: &[&str]) -> Vec<String> {
968 ids.iter().map(|s| s.to_string()).collect()
969 }
970
971 #[test]
972 fn perfect_ranking_scores_one() {
973 let (p1, r, mrr, ndcg) = recall_metrics(&ranked(&["a", "b", "x"]), &rel(&["a", "b"]), 5);
974 assert!((p1 - 1.0).abs() < 1e-9);
975 assert!((r - 1.0).abs() < 1e-9);
976 assert!((mrr - 1.0).abs() < 1e-9);
977 assert!((ndcg - 1.0).abs() < 1e-9);
978 }
979
980 #[test]
981 fn missed_first_lowers_p1_and_mrr() {
982 let (p1, r, mrr, _ndcg) = recall_metrics(&ranked(&["x", "a"]), &rel(&["a"]), 5);
984 assert_eq!(p1, 0.0);
985 assert!((mrr - 0.5).abs() < 1e-9);
986 assert!((r - 1.0).abs() < 1e-9);
987 }
988
989 #[test]
990 fn k_cutoff_limits_recall() {
991 let (_p1, r, _mrr, ndcg) = recall_metrics(&ranked(&["x", "a"]), &rel(&["a"]), 1);
993 assert_eq!(r, 0.0);
994 assert_eq!(ndcg, 0.0);
995 }
996
997 #[test]
998 fn no_hits_is_all_zero() {
999 let (p1, r, mrr, ndcg) = recall_metrics(&ranked(&["x", "y"]), &rel(&["a"]), 5);
1000 assert_eq!((p1, r, mrr, ndcg), (0.0, 0.0, 0.0, 0.0));
1001 }
1002}