diff --git a/pk-cli/src/main.rs b/pk-cli/src/main.rs index dfc3ca9..64458fd 100644 --- a/pk-cli/src/main.rs +++ b/pk-cli/src/main.rs @@ -145,7 +145,8 @@ enum Cmd { scopes: Vec, #[arg(long, default_value_t = 8)] limit: usize, - /// Maximum immutable snapshot candidates inspected across all scopes. + /// Maximum scored candidates kept across all scopes. Every snapshot entry + /// is scored first; the cap applies to the merged, ranked list. #[arg(long, default_value_t = 128)] max_candidates: usize, /// Maximum bytes emitted in hook format. @@ -659,6 +660,9 @@ struct ContextCandidate { struct ContextOutput { query: String, snapshot_generations: BTreeMap, + /// Snapshot entries scored across all readable scopes. + scored_count: usize, + /// Ranked candidates kept after de-duplication and the `--max-candidates` cap. candidate_count: usize, byte_count: usize, failures: Vec, @@ -705,10 +709,9 @@ async fn run_context( }; let max_candidates = max_candidates.clamp(1, 512); let max_bytes = max_bytes.clamp(256, 65_536); - let candidates_per_scope = max_candidates.div_ceil(scopes.len().max(1)); let mut failures = Vec::new(); let mut candidates = Vec::new(); - let mut inspected_candidates = 0usize; + let mut scored_count = 0usize; let mut generations = BTreeMap::new(); for scope in scopes { let Some(path) = knowledge_root_for_scope(scope, explicit_project_kb) else { @@ -732,16 +735,11 @@ async fn run_context( } }; generations.insert(scope.label().to_owned(), snapshot.generation); - let remaining = max_candidates - .saturating_sub(inspected_candidates) - .min(candidates_per_scope); - let bounded_entries = snapshot - .entries - .into_iter() - .take(remaining) - .collect::>(); - inspected_candidates += bounded_entries.len(); - candidates.extend(bounded_entries.into_iter().filter_map(|entry| { + // Score every entry before any budget applies: truncating first meant + // only the first entries in snapshot order were ever considered, and a + // failed scope's share of the budget was lost. + scored_count += snapshot.entries.len(); + candidates.extend(snapshot.entries.into_iter().filter_map(|entry| { let score = snapshot_score(query, &entry); (score > 0.0 || query.trim().is_empty()).then_some(ContextCandidate { scope, @@ -749,9 +747,6 @@ async fn run_context( score, }) })); - if inspected_candidates == max_candidates { - break; - } } // Select a canonical copy of duplicate IDs or duplicate content. Scope @@ -793,6 +788,8 @@ async fn run_context( .then_with(|| left.scope.priority().cmp(&right.scope.priority())) .then_with(|| left.entry.id.as_str().cmp(right.entry.id.as_str())) }); + selected.truncate(max_candidates); + let candidate_count = selected.len(); selected.truncate(limit.clamp(1, 32)); let results = selected @@ -811,7 +808,8 @@ async fn run_context( let mut output = ContextOutput { query: query.to_owned(), snapshot_generations: generations, - candidate_count: inspected_candidates, + scored_count, + candidate_count, byte_count: 0, failures, results, diff --git a/pk-cli/tests/context.rs b/pk-cli/tests/context.rs index 77bd49d..285b8ed 100644 --- a/pk-cli/tests/context.rs +++ b/pk-cli/tests/context.rs @@ -184,6 +184,9 @@ fn candidate_budget_is_shared_across_requested_scopes() { assert!(output.status.success()); let report: Value = serde_json::from_slice(&output.stdout).unwrap(); - assert_eq!(report["candidate_count"], 2); + // Every entry in both scopes is scored; only the matching one survives, + // and the cap bounds what is kept, not what is inspected. + assert_eq!(report["scored_count"], 5); + assert_eq!(report["candidate_count"], 1); assert_eq!(report["results"][0]["id"], "shared-target"); } diff --git a/pk-cli/tests/context_scoring.rs b/pk-cli/tests/context_scoring.rs new file mode 100644 index 0000000..cdec55e --- /dev/null +++ b/pk-cli/tests/context_scoring.rs @@ -0,0 +1,171 @@ +//! `pk context` must score every committed snapshot entry before any candidate +//! budget applies. Before this was fixed, only the first +//! `ceil(max_candidates / scopes)` entries of each scope, in snapshot order, +//! were ever scored, so a lesson that sorted late was never recalled. + +use serde_json::Value; +use std::{ + fs, + path::{Path, PathBuf}, + process::{Command, Output}, +}; + +fn write_entry(base: &Path, id: &str, body: &str) { + let wiki = base.join("wiki"); + fs::create_dir_all(&wiki).unwrap(); + fs::write( + wiki.join(format!("{id}.md")), + format!("---\ntype: Lesson\ntitle: {id}\ntags: [fixture]\n---\n\n{body}\n"), + ) + .unwrap(); +} + +struct Fixture { + _dir: tempfile::TempDir, + home: PathBuf, + project: PathBuf, +} + +fn fixture(project_has_root: bool) -> Fixture { + let dir = tempfile::tempdir().unwrap(); + let home = dir.path().join("home"); + let project = dir.path().join("project"); + fs::create_dir_all(&project).unwrap(); + if project_has_root { + fs::create_dir_all(project.join(".git")).unwrap(); + } + fs::create_dir_all(&home).unwrap(); + Fixture { + _dir: dir, + home, + project, + } +} + +fn pk(f: &Fixture, args: &[&str]) -> Output { + let output = Command::new(env!("CARGO_BIN_EXE_pk")) + .current_dir(&f.project) + .env("HOME", &f.home) + .env("RUST_LOG", "error") + .args(args) + .output() + .unwrap(); + assert!( + output.status.success(), + "pk {args:?} failed: {}", + String::from_utf8_lossy(&output.stderr) + ); + output +} + +fn json(output: &Output) -> Value { + serde_json::from_slice(&output.stdout).unwrap() +} + +#[test] +fn an_entry_sorting_last_in_a_large_scope_is_still_recalled_with_default_flags() { + let f = fixture(true); + let kb = f.project.join(".prometheus/knowledge"); + for index in 0..199 { + write_entry( + &kb, + &format!("entry-{index:03}"), + "unrelated filler lesson about build caching", + ); + } + // Sorts after every filler entry by id. + write_entry(&kb, "zzz-target", "latesortingtoken lesson that must be recalled"); + pk(&f, &["snapshot", "--scope", "project"]); + + let report = json(&pk( + &f, + &["context", "latesortingtoken", "--scope", "project", "--format", "json"], + )); + assert_eq!(report["scored_count"], 200, "{report:#}"); + assert_eq!(report["results"][0]["id"], "zzz-target", "{report:#}"); +} + +#[test] +fn a_failed_scope_does_not_reserve_budget_the_remaining_scopes_can_use() { + // No project root: the project scope fails, and the shared scope must be + // able to fill the whole cap on its own. + let f = fixture(false); + let shared = f.home.join(".prometheus/knowledge/shared"); + for index in 0..60 { + write_entry( + &shared, + &format!("shared-{index:02}"), + "budgettoken shared lesson", + ); + } + pk(&f, &["snapshot", "--scope", "shared"]); + + let report = json(&pk( + &f, + &[ + "context", + "budgettoken", + "--scope", + "project", + "--scope", + "shared", + "--max-candidates", + "60", + "--limit", + "32", + "--format", + "json", + ], + )); + assert_eq!(report["failures"].as_array().unwrap().len(), 1, "{report:#}"); + assert_eq!(report["candidate_count"], 60, "{report:#}"); + assert_eq!(report["results"].as_array().unwrap().len(), 32, "{report:#}"); +} + +#[test] +fn output_is_deterministic_and_ranked_by_score_then_scope_then_id() { + let f = fixture(true); + let project_kb = f.project.join(".prometheus/knowledge"); + let shared = f.home.join(".prometheus/knowledge/shared"); + // Equal scores across scopes: project must precede shared, then id order. + write_entry(&project_kb, "b-project", "ordertoken lesson"); + write_entry(&project_kb, "a-project", "ordertoken lesson"); + write_entry(&shared, "a-shared", "ordertoken different lesson"); + // Higher score (token appears twice) ranks first regardless of scope. + write_entry(&shared, "z-strong", "ordertoken ordertoken strong lesson"); + pk(&f, &["snapshot", "--scope", "project", "--scope", "shared"]); + + let args = [ + "context", + "ordertoken", + "--scope", + "project", + "--scope", + "shared", + "--format", + "json", + ]; + let first = pk(&f, &args).stdout; + let second = pk(&f, &args).stdout; + assert_eq!(first, second, "context output must be byte-identical across runs"); + + let report: Value = serde_json::from_slice(&first).unwrap(); + let results = report["results"].as_array().unwrap(); + let scores: Vec = results.iter().map(|r| r["score"].as_f64().unwrap()).collect(); + assert!( + scores.windows(2).all(|w| w[0] >= w[1]), + "results must be ordered by descending score: {report:#}" + ); + // Among entries with the top-but-one score, project entries come first in id order. + let tied: Vec<(&str, &str)> = results + .iter() + .filter(|r| r["score"].as_f64().unwrap() == scores[scores.len() - 1]) + .map(|r| (r["scope"].as_str().unwrap(), r["id"].as_str().unwrap())) + .collect(); + let mut expected = tied.clone(); + expected.sort_by(|left, right| { + let rank = |scope: &str| if scope == "project" { 0 } else { 1 }; + rank(left.0).cmp(&rank(right.0)).then(left.1.cmp(right.1)) + }); + assert_eq!(tied, expected, "{report:#}"); +}