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