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
3 changes: 2 additions & 1 deletion codex-rs/ext/skills/src/extension.rs
Original file line number Diff line number Diff line change
Expand Up @@ -395,11 +395,12 @@ where
let shadow_selected_entries =
collect_explicit_skill_mentions(&input.user_input, &shadow_catalog);
Some(self.shadow_selection.run(
&input.user_input,
&input,
&shadow_catalog,
&shadow_selected_entries,
host_snapshot.as_deref(),
Arc::clone(&thread_state.recent_skill_invocations),
Arc::clone(&thread_state.shadow_task_context),
))
} else {
None
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@
// This shadow-selection experiment is temporary and should be removed after evaluation.

mod task_context;

pub(crate) use task_context::ShadowTaskContext;

use std::collections::HashMap;
use std::collections::HashSet;
use std::collections::VecDeque;
Expand All @@ -10,6 +14,7 @@ use std::time::Duration;
use std::time::Instant;

use crate::HostSkillsSnapshot;
use codex_extension_api::TurnInputContext;
use codex_otel::MetricsClient;
use codex_protocol::user_input::UserInput;

Expand Down Expand Up @@ -64,14 +69,16 @@ impl ShadowSelectionExperiment {

pub(crate) fn run(
&self,
inputs: &[UserInput],
input: &TurnInputContext,
catalog: &SkillCatalog,
explicitly_selected: &[SkillCatalogEntry],
host_snapshot: Option<&HostSkillsSnapshot>,
recent_skill_invocations: Arc<RecentSkillInvocations>,
task_context: Arc<ShadowTaskContext>,
) -> ShadowSelectionTurnState {
let query = build_shadow_query(inputs);
let query = build_shadow_query(&input.user_input);
let query_script = query_script_tag(&query.text);
let task_snapshot = task_context.begin_turn(&input.turn_id, &query, &input.user_input);
let explicitly_selected_skill_resources = explicitly_selected
.iter()
.map(|entry| normalize_skill_resource(entry.main_prompt.as_str()))
Expand Down Expand Up @@ -132,9 +139,19 @@ impl ShadowSelectionExperiment {
lru_selector.clone(),
routing_selector.clone(),
);
let mut ranked_selections = Vec::with_capacity(self.selectors.len() + 5);
let task_selector = LruPlusLexicalCharacterRoutingSkillSelector::new(
LruSkillSelector::new(
task_snapshot
.recent_skills
.iter()
.filter_map(|resource| eligible_skill_ids_by_resource.get(resource).copied())
.collect(),
),
routing_selector.clone(),
);
let mut ranked_selections = Vec::with_capacity(self.selectors.len() + 6);

for selector in self
for (method, selector, query) in self
.selectors
.iter()
.map(std::convert::AsRef::as_ref)
Expand All @@ -145,14 +162,21 @@ impl ShadowSelectionExperiment {
&lru_plus_character_selector as &dyn CheapSkillSelector,
&lru_plus_lexical_character_selector as &dyn CheapSkillSelector,
])
.map(|selector| (selector.method(), selector, &query))
.chain([(
"task_context_fusion_v1",
&task_selector as &dyn CheapSkillSelector,
&task_snapshot.query,
)])
{
let query_script = query_script_tag(&query.text);
let start = Instant::now();
let selection =
selector.select(&query.text, &documents, /*limit*/ MAX_SHADOW_RESULTS);
let duration = start.elapsed();
let selected_ids = sanitize_selected_ids(&selection, &eligible_ids);
self.record_metrics(ShadowSelectionObservation {
method: selector.method(),
method,
selection: &selection,
query_truncated_before_selection: query.truncated,
query_script,
Expand All @@ -161,14 +185,14 @@ impl ShadowSelectionExperiment {
duration,
});
ranked_selections.push(RankedSelection {
method: selector.method(),
method,
skill_resources: selected_ids
.iter()
.map(|id| normalize_skill_resource(catalog.entries[*id].main_prompt.as_str()))
.collect(),
});
tracing::debug!(
method = selector.method(),
method,
catalog_entries = documents.len(),
selected_entries = selected_ids.len(),
query_terms = selection.query_term_count,
Expand All @@ -179,12 +203,29 @@ impl ShadowSelectionExperiment {
);
}

// Explicit intent is a relevance signal even if the subsequent prompt read fails.
// Keep it out of the implicit-only controls and freeze predictions before recording it.
for entry in explicitly_selected.iter().filter(|entry| {
entry.is_model_visible()
&& matches!(
&entry.authority.kind,
SkillSourceKind::Host | SkillSourceKind::Orchestrator
)
}) {
task_context.record(
&input.turn_id,
normalize_skill_resource(entry.main_prompt.as_str()),
);
}

ShadowSelectionTurnState {
ranked_selections,
turn_id: input.turn_id.clone(),
query_script,
eligible_skill_resources,
seen_skill_resources: Mutex::new(HashSet::new()),
recent_skill_invocations,
task_context,
}
}

Expand All @@ -204,6 +245,9 @@ impl ShadowSelectionExperiment {
state
.recent_skill_invocations
.record(skill_resource.clone());
state
.task_context
.record(&state.turn_id, skill_resource.clone());
let Some(metrics_client) = self.metrics_client.as_ref() else {
return;
};
Expand Down Expand Up @@ -274,10 +318,12 @@ impl ShadowSelectionExperiment {

pub(crate) struct ShadowSelectionTurnState {
ranked_selections: Vec<RankedSelection>,
turn_id: String,
query_script: &'static str,
eligible_skill_resources: HashSet<String>,
seen_skill_resources: Mutex<HashSet<String>>,
recent_skill_invocations: Arc<RecentSkillInvocations>,
task_context: Arc<ShadowTaskContext>,
}

#[derive(Default)]
Expand Down Expand Up @@ -435,6 +481,7 @@ fn is_cjk(character: char) -> bool {
)
}

#[derive(Clone, Debug, PartialEq, Eq)]
struct ShadowQuery {
text: String,
truncated: bool,
Expand Down Expand Up @@ -479,5 +526,5 @@ fn push_bounded(destination: &mut String, value: &str) -> bool {
}

#[cfg(test)]
#[path = "shadow_selection_experiment_tests.rs"]
#[path = "experiment_tests.rs"]
mod tests;
155 changes: 155 additions & 0 deletions codex-rs/ext/skills/src/shadow_selection_experiment/task_context.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,155 @@
use std::collections::VecDeque;
use std::sync::Mutex;
use std::sync::PoisonError;

use codex_protocol::user_input::UserInput;
use codex_utils_string::take_bytes_at_char_boundary;

use super::ShadowQuery;

const MAX_PRIOR_REQUESTS: usize = 2;
const MAX_REQUEST_BYTES: usize = 2 * 1024;
const MAX_QUERY_BYTES: usize = 4 * 1024;
const MAX_RECENT_SKILLS: usize = 50;

/// Shadow-only relevance history. New and reconstructed thread runtimes start cold.
#[derive(Default)]
pub(crate) struct ShadowTaskContext(Mutex<TaskContextState>);

#[derive(Default)]
struct TaskContextState {
prior_requests: VecDeque<ShadowQuery>,
recent_skills: VecDeque<String>,
pending: Option<PendingTurn>,
}

struct PendingTurn {
id: String,
request: Option<ShadowQuery>,
recent_skills: VecDeque<String>,
}

#[derive(Debug, PartialEq, Eq)]
pub(super) struct TaskContextSnapshot {
pub(super) query: ShadowQuery,
pub(super) recent_skills: Vec<String>,
}

impl ShadowTaskContext {
pub(super) fn begin_turn(
&self,
turn_id: &str,
current: &ShadowQuery,
inputs: &[UserInput],
) -> TaskContextSnapshot {
let mut state = self.0.lock().unwrap_or_else(PoisonError::into_inner);
if state.pending.as_ref().is_none_or(|turn| turn.id != turn_id) {
if let Some(previous) = state.pending.take() {
if let Some(request) = previous.request {
state
.prior_requests
.retain(|prior| prior.text != request.text);
state.prior_requests.push_front(request);
state.prior_requests.truncate(MAX_PRIOR_REQUESTS);
}
for resource in previous.recent_skills.into_iter().rev() {
remember_skill(&mut state.recent_skills, resource);
}
}
state.pending = Some(PendingTurn {
id: turn_id.to_string(),
request: None,
recent_skills: VecDeque::new(),
});
}

let retained = take_bytes_at_char_boundary(&current.text, MAX_REQUEST_BYTES);
let text = take_bytes_at_char_boundary(&current.text, MAX_QUERY_BYTES);
let mut query = ShadowQuery {
text: text.to_string(),
truncated: current.truncated || text.len() < current.text.len(),
};
for prior in &state.prior_requests {
if prior.text == retained {
continue;
}
if !query.text.is_empty() && query.text.len() < MAX_QUERY_BYTES {
query.text.push('\n');
}
let part = take_bytes_at_char_boundary(
&prior.text,
MAX_QUERY_BYTES.saturating_sub(query.text.len()),
);
query.text.push_str(part);
query.truncated |= prior.truncated || part.len() < prior.text.len();
if part.len() < prior.text.len() {
break;
}
}
let recent_skills = state.recent_skills.iter().cloned().collect();
if is_substantive(&current.text, inputs)
&& let Some(pending) = state.pending.as_mut()
{
pending.request = Some(ShadowQuery {
text: retained.to_string(),
truncated: current.truncated || retained.len() < current.text.len(),
});
}
TaskContextSnapshot {
query,
recent_skills,
}
}

/// Records relevance evidence for future turns, never the active turn's predictions.
pub(super) fn record(&self, turn_id: &str, resource: String) {
let mut state = self.0.lock().unwrap_or_else(PoisonError::into_inner);
if let Some(pending) = state.pending.as_mut().filter(|turn| turn.id == turn_id) {
remember_skill(&mut pending.recent_skills, resource);
}
}
}

fn remember_skill(skills: &mut VecDeque<String>, resource: String) {
skills.retain(|previous| previous != &resource);
skills.push_front(resource);
skills.truncate(MAX_RECENT_SKILLS);
}

fn is_substantive(text: &str, inputs: &[UserInput]) -> bool {
if inputs.iter().any(|input| {
matches!(input, UserInput::Skill { name, .. } | UserInput::Mention { name, .. } if !name.trim().is_empty())
}) {
return true;
}
let normalized = text
.split(|character: char| !character.is_alphanumeric())
.filter(|part| !part.is_empty())
.map(str::to_lowercase)
.collect::<Vec<_>>()
.join(" ");
!matches!(
normalized.as_str(),
"" | "yes"
| "yep"
| "yeah"
| "ok"
| "okay"
| "sure"
| "go"
| "go ahead"
| "continue"
| "please continue"
| "proceed"
| "do it"
| "do that"
| "try again"
| "retry"
| "thanks"
| "thank you"
)
}

#[cfg(test)]
#[path = "task_context_tests.rs"]
mod tests;
Loading
Loading