Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion codex-rs/Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

1 change: 0 additions & 1 deletion codex-rs/ext/skills/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,5 @@ url = { workspace = true }

[dev-dependencies]
codex-utils-absolute-path = { workspace = true }
opentelemetry_sdk = { workspace = true }
pretty_assertions = { workspace = true }
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
2 changes: 2 additions & 0 deletions codex-rs/ext/skills/src/dynamic_skill_selector.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
mod fielded_bm25;
mod weighted_lexical;
pub(crate) use fielded_bm25::FieldedBm25SkillSelector;
pub(crate) use weighted_lexical::WeightedLexicalSkillSelector;

/// Metadata searched by a cheap skill selector.
Expand Down
239 changes: 239 additions & 0 deletions codex-rs/ext/skills/src/dynamic_skill_selector/fielded_bm25.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,239 @@
use std::collections::HashMap;
use std::collections::HashSet;

use super::CheapSkillSelection;
use super::CheapSkillSelector;
use super::SkillSelectionDocument;

const MAX_QUERY_BYTES: usize = 4 * 1024;
const MAX_QUERY_TERMS: usize = 64;
const MAX_DOCUMENT_BYTES: usize = 4 * 1024;
const MAX_DOCUMENT_TERMS: usize = 256;
const MAX_CANDIDATES: usize = 1_000;
const MAX_RESULTS: usize = 50;
const FIELD_WEIGHTS: [f64; 3] = [8.0, 4.0, 1.0];
const K1: f64 = 1.2;
const B: f64 = 0.75;

const STOP_WORDS: &[&str] = &[
"a", "an", "and", "are", "as", "at", "be", "by", "do", "for", "from", "how", "i", "in", "is",
"it", "me", "my", "of", "on", "or", "please", "that", "the", "this", "to", "use", "we", "what",
"when", "where", "which", "with", "you", "your",
];

#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct FieldedBm25SkillSelector;

impl CheapSkillSelector for FieldedBm25SkillSelector {
fn method(&self) -> &'static str {
"fielded_bm25_v1"
}

fn select(
&self,
query: &str,
documents: &[SkillSelectionDocument<'_>],
limit: usize,
) -> CheapSkillSelection {
let (query, query_bytes_truncated) = bounded(query, MAX_QUERY_BYTES);
let (query_terms, query_terms_truncated) = query_terms(query);
let query_truncated = query_bytes_truncated || query_terms_truncated;
let candidate_set_truncated = documents.len() > MAX_CANDIDATES;
if query_terms.is_empty() || limit == 0 {
return CheapSkillSelection {
query_term_count: query_terms.len(),
query_truncated,
candidate_set_truncated,
..Default::default()
};
}

let prepared = documents
.iter()
.take(MAX_CANDIDATES)
.map(PreparedDocument::new)
.collect::<Vec<_>>();
let averages = average_field_lengths(&prepared);
let document_frequencies = document_frequencies(&prepared);
let document_count = prepared.len() as f64;
let mut scored = prepared
.iter()
.filter_map(|document| {
let score = score_document(
document,
&query_terms,
&document_frequencies,
document_count,
averages,
);
(score > 0.0).then_some((score, document.id, document.name))
})
.collect::<Vec<_>>();
scored.sort_by(|left, right| {
right
.0
.total_cmp(&left.0)
.then_with(|| left.2.cmp(right.2))
.then_with(|| left.1.cmp(&right.1))
});

CheapSkillSelection {
candidate_ids: scored
.into_iter()
.take(limit.min(MAX_RESULTS))
.map(|(_, id, _)| id)
.collect(),
query_term_count: query_terms.len(),
query_truncated,
candidate_set_truncated,
}
}
}

struct PreparedDocument<'a> {
id: usize,
name: &'a str,
fields: [Vec<String>; 3],
}

impl<'a> PreparedDocument<'a> {
fn new(document: &'a SkillSelectionDocument<'a>) -> Self {
Self {
id: document.id,
name: document.name,
fields: [
document_terms(document.name),
document_terms(document.short_description.unwrap_or_default()),
document_terms(document.description),
],
}
}
}

fn score_document(
document: &PreparedDocument<'_>,
query_terms: &[String],
document_frequencies: &HashMap<String, usize>,
document_count: f64,
average_field_lengths: [f64; 3],
) -> f64 {
query_terms.iter().fold(0.0, |score, query_term| {
let frequency = document_frequencies
.get(query_term)
.copied()
.unwrap_or_default() as f64;
if frequency == 0.0 {
return score;
}
let weighted_term_frequency =
document
.fields
.iter()
.enumerate()
.fold(0.0, |weighted, (field_index, terms)| {
let term_frequency =
terms.iter().filter(|term| *term == query_term).count() as f64;
if term_frequency == 0.0 {
return weighted;
}
let average_length = average_field_lengths[field_index];
let length_ratio = if average_length == 0.0 {
1.0
} else {
terms.len() as f64 / average_length
};
weighted
+ FIELD_WEIGHTS[field_index] * term_frequency / (1.0 - B + B * length_ratio)
});
if weighted_term_frequency == 0.0 {
return score;
}
let inverse_document_frequency =
(1.0 + (document_count - frequency + 0.5) / (frequency + 0.5)).ln();
score
+ inverse_document_frequency * weighted_term_frequency * (K1 + 1.0)
/ (weighted_term_frequency + K1)
})
}

fn average_field_lengths(documents: &[PreparedDocument<'_>]) -> [f64; 3] {
if documents.is_empty() {
return [0.0; 3];
}
let totals = documents.iter().fold([0usize; 3], |mut totals, document| {
for (index, field) in document.fields.iter().enumerate() {
totals[index] = totals[index].saturating_add(field.len());
}
totals
});
totals.map(|total| total as f64 / documents.len() as f64)
}

fn document_frequencies(documents: &[PreparedDocument<'_>]) -> HashMap<String, usize> {
let mut frequencies = HashMap::new();
for document in documents {
let terms = document
.fields
.iter()
.flatten()
.map(String::as_str)
.collect::<HashSet<_>>();
for term in terms {
*frequencies.entry(term.to_string()).or_default() += 1;
}
}
frequencies
}

fn query_terms(query: &str) -> (Vec<String>, bool) {
let mut seen = HashSet::new();
let mut terms = Vec::new();
for term in normalized_terms(query)
.into_iter()
.filter(|term| term.chars().count() >= 2 && !STOP_WORDS.contains(&term.as_str()))
{
if !seen.insert(term.clone()) {
continue;
}
if terms.len() == MAX_QUERY_TERMS {
return (terms, true);
}
terms.push(term);
}
(terms, false)
}

fn document_terms(value: &str) -> Vec<String> {
let (value, _) = bounded(value, MAX_DOCUMENT_BYTES);
normalized_terms(value)
.into_iter()
.take(MAX_DOCUMENT_TERMS)
.collect()
}

fn normalized_terms(value: &str) -> Vec<String> {
let mut normalized = String::with_capacity(value.len());
for character in value.chars() {
if character.is_alphanumeric() {
normalized.extend(character.to_lowercase());
} else {
normalized.push(' ');
}
}
normalized.split_whitespace().map(str::to_string).collect()
}

fn bounded(value: &str, max_bytes: usize) -> (&str, bool) {
if value.len() <= max_bytes {
return (value, false);
}
let mut end = max_bytes;
while !value.is_char_boundary(end) {
end = end.saturating_sub(1);
}
(&value[..end], true)
}

#[cfg(test)]
#[path = "fielded_bm25_tests.rs"]
mod tests;
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
use super::*;
use pretty_assertions::assert_eq;

#[test]
fn bm25_prioritizes_rare_terms() {
let documents = [
document(/*id*/ 1, "review-helper", "Review code and prose."),
document(
/*id*/ 2,
"terraform-review",
"Review Terraform infrastructure.",
),
document(/*id*/ 3, "document-review", "Review Word documents."),
];

let selection =
FieldedBm25SkillSelector.select("review terraform", &documents, /*limit*/ 20);

assert_eq!(vec![2, 3, 1], selection.candidate_ids);
}

#[test]
fn bm25_weights_names_above_descriptions() {
let documents = [
document(/*id*/ 1, "slides", "Create presentations."),
document(/*id*/ 2, "presentations", "Create and edit slides."),
];

let selection = FieldedBm25SkillSelector.select("slides", &documents, /*limit*/ 20);

assert_eq!(vec![1, 2], selection.candidate_ids);
}

#[test]
fn bm25_drops_candidates_without_matching_terms() {
let documents = [document(
/*id*/ 1,
"spreadsheets",
"Analyze tabular data.",
)];

let selection =
FieldedBm25SkillSelector.select("render a video", &documents, /*limit*/ 20);

assert!(selection.candidate_ids.is_empty());
}

#[test]
fn bm25_reports_bounded_inputs() {
let long_query = "match ".repeat(MAX_QUERY_BYTES);
let names = (0..=MAX_CANDIDATES)
.map(|index| format!("candidate-{index}"))
.collect::<Vec<_>>();
let documents = names
.iter()
.enumerate()
.map(|(id, name)| SkillSelectionDocument {
id,
name,
short_description: None,
description: "match",
})
.collect::<Vec<_>>();

let selection = FieldedBm25SkillSelector.select(&long_query, &documents, /*limit*/ 20);

assert!(selection.query_truncated);
assert!(selection.candidate_set_truncated);
assert_eq!(20, selection.candidate_ids.len());
}

fn document<'a>(id: usize, name: &'a str, description: &'a str) -> SkillSelectionDocument<'a> {
SkillSelectionDocument {
id,
name,
short_description: None,
description,
}
}
6 changes: 5 additions & 1 deletion codex-rs/ext/skills/src/shadow_selection_experiment.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ use crate::catalog::SkillCatalog;
use crate::catalog::SkillSourceKind;
use crate::dynamic_skill_selector::CheapSkillSelection;
use crate::dynamic_skill_selector::CheapSkillSelector;
use crate::dynamic_skill_selector::FieldedBm25SkillSelector;
use crate::dynamic_skill_selector::SkillSelectionDocument;
use crate::dynamic_skill_selector::WeightedLexicalSkillSelector;

Expand All @@ -35,7 +36,10 @@ pub(crate) struct ShadowSelectionExperiment {
impl ShadowSelectionExperiment {
pub(crate) fn new(metrics_client: Option<MetricsClient>) -> Self {
Self {
selectors: vec![Box::new(WeightedLexicalSkillSelector)],
selectors: vec![
Box::new(WeightedLexicalSkillSelector),
Box::new(FieldedBm25SkillSelector),
],
metrics_client,
}
}
Expand Down
Loading
Loading