use std::{path::Path, sync::Arc};
use anyhow::{anyhow, bail};
use rho_providers::{
model::{
models_dev::{cached_model_metadata, ModelMetadata},
Message,
},
provider::ProviderId,
reasoning::ReasoningLevel,
};
use rho_sdk::model::context::estimate_text_tokens;
use rho_sdk::{
provider::ModelProvider, ApprovalRequest, CancellationToken, ProviderRequestUsageRecording,
SessionId,
};
use crate::{
agent::{effective_internal_agent_reasoning, PERMISSION_CLASSIFIER_AGENT_ID},
config::{
Config, InternalAgentModelConfig, InternalAgentTarget, ModelKind, RhoInternalAgentModel,
},
credential_store::build_provider_on,
};
use super::{
review_verdict, screen_allow_probability, screen_verdict, transcript::render_with_pending_call,
ClassifierVerdict, ScreenVerdict, TranscriptBudget, TranscriptOverBudget, CLASSIFIER_POLICY,
DEFAULT_SCREEN_ALLOW_PERCENT, REVIEW_QUESTION, SCREEN_ALLOW_PERCENT_RANGE, SCREEN_QUESTION,
};
use crate::decision::{self, EntryModel, TextModel};
use rho_sdk::decision::{
text::{questions_block, system_prompt, AnswerStyle},
Answer, DecisionModel, DecisionRequest, Question,
};
pub(crate) const DECISION_SCREEN_ID: &str = "permission-classifier-screen";
const CLASSIFIER_OUTPUT_RESERVE_TOKENS: u64 = 8_192;
pub(crate) struct ClassifyRequest<'a> {
pub history: &'a [Message],
pub pending: &'a ApprovalRequest,
pub cancellation: CancellationToken,
pub session_id: &'a SessionId,
pub workspace_path: &'a Path,
pub usage_recording: ProviderRequestUsageRecording,
}
impl ClassifyRequest<'_> {
pub(super) fn scope(&self) -> CallScope<'_> {
CallScope {
cancellation: &self.cancellation,
session_id: self.session_id,
workspace_path: self.workspace_path,
usage_recording: &self.usage_recording,
}
}
}
pub(super) struct CallScope<'a> {
pub cancellation: &'a CancellationToken,
pub session_id: &'a SessionId,
pub workspace_path: &'a Path,
pub usage_recording: &'a ProviderRequestUsageRecording,
}
pub(crate) async fn classify_capability_request(
config: &Config,
request: ClassifyRequest<'_>,
) -> ClassifierVerdict {
let result = match ClassifierModel::resolve(config).await {
Ok(model) => model.classify(request).await.result,
Err(error) => Err(error),
};
result.unwrap_or_else(classifier_unavailable)
}
pub(crate) struct ClassifierModel {
pub(super) provider: Arc<dyn ModelProvider>,
pub(super) reasoning: ReasoningLevel,
pub(super) budget: TranscriptBudget,
pub(super) screen: Screen,
}
pub(super) enum Screen {
Classifier,
Text {
provider: Arc<dyn ModelProvider>,
budget: Option<TranscriptBudget>,
},
Decision {
model: Box<dyn DecisionModel>,
allow_percent: u8,
},
}
pub(crate) fn check_screen_config(config: &Config) -> anyhow::Result<()> {
match decision::resolve(config, DECISION_SCREEN_ID)? {
Some(EntryModel::Decision(_)) => {
decision_allow_percent(config)?;
}
Some(EntryModel::Text(_)) | None => {}
}
Ok(())
}
fn decision_allow_percent(config: &Config) -> Result<u8, decision::ConfigError> {
config
.internal_agent_model(DECISION_SCREEN_ID)
.and_then(InternalAgentModelConfig::rho)
.map_or(Ok(DEFAULT_SCREEN_ALLOW_PERCENT), screen_allow_percent)
}
pub(crate) fn screen_allow_percent(
selection: &RhoInternalAgentModel,
) -> Result<u8, decision::ConfigError> {
let percent = selection
.allow_threshold_percent
.unwrap_or(DEFAULT_SCREEN_ALLOW_PERCENT);
if SCREEN_ALLOW_PERCENT_RANGE.contains(&percent) {
Ok(percent)
} else {
Err(decision::ConfigError::AllowThresholdOutOfRange {
entry: DECISION_SCREEN_ID,
percent,
min: *SCREEN_ALLOW_PERCENT_RANGE.start(),
max: *SCREEN_ALLOW_PERCENT_RANGE.end(),
})
}
}
impl ClassifierModel {
pub(crate) async fn resolve(config: &Config) -> anyhow::Result<Self> {
let screen = resolve_screen(config).await?;
let model = config
.internal_agent_model(PERMISSION_CLASSIFIER_AGENT_ID)
.ok_or_else(|| anyhow!("{PERMISSION_CLASSIFIER_AGENT_ID} model is not configured"))?;
let reasoning = effective_internal_agent_reasoning(PERMISSION_CLASSIFIER_AGENT_ID, model);
let InternalAgentTarget::Rho(selection) = &model.target else {
bail!("{PERMISSION_CLASSIFIER_AGENT_ID} cannot run on Claude Code runtime");
};
let provider = build_provider_on(
config,
&selection.provider,
&selection.model,
reasoning,
&selection.auth,
)
.await
.map_err(|_| {
anyhow!(
"failed to build {PERMISSION_CLASSIFIER_AGENT_ID} provider; check configured credentials"
)
})?;
let budget = transcript_budget(
cached_model_metadata(&selection.provider, &selection.model)
.and_then(|metadata| metadata.display_context_window()),
);
Ok(Self {
provider,
reasoning,
budget,
screen,
})
}
pub(crate) fn provider(&self) -> &dyn ModelProvider {
self.provider.as_ref()
}
pub(crate) fn reasoning(&self) -> ReasoningLevel {
self.reasoning
}
pub(crate) async fn classify(&self, request: ClassifyRequest<'_>) -> ClassifierTrace {
let pending_call_id = request.pending.tool_call_id().map(|id| id.as_str());
run_pipeline(
self.provider.as_ref(),
&self.screen,
self.reasoning,
self.budget,
pending_call_id,
&request,
)
.await
}
pub(crate) async fn classify_with_pending_call(
&self,
request: ClassifyRequest<'_>,
pending_call_id: &str,
) -> ClassifierTrace {
run_pipeline(
self.provider.as_ref(),
&self.screen,
self.reasoning,
self.budget,
Some(pending_call_id),
&request,
)
.await
}
}
async fn resolve_screen(config: &Config) -> anyhow::Result<Screen> {
let selection = match decision::resolve(config, DECISION_SCREEN_ID)? {
None => return Ok(Screen::Classifier),
Some(EntryModel::Decision(model)) => {
return Ok(Screen::Decision {
model,
allow_percent: decision_allow_percent(config)?,
});
}
Some(EntryModel::Text(selection)) => selection,
};
let provider = build_provider_on(
config,
&selection.provider,
&selection.model,
ReasoningLevel::Low,
&selection.auth,
)
.await
.map_err(|_| decision::ConfigError::TextModelUnavailable {
entry: DECISION_SCREEN_ID,
configured: rho_providers::provider::model_reference(&selection.provider, &selection.model),
})?;
let budget = text_screen_budget(
&selection.provider,
cached_model_metadata(&selection.provider, &selection.model),
);
Ok(Screen::Text { provider, budget })
}
pub(crate) fn screen_warning(selection: &RhoInternalAgentModel) -> Option<String> {
decision::kind_mismatch(selection).or_else(|| {
let unknown_window = decision::entry_kind(selection) == ModelKind::Text
&& text_screen_budget(
&selection.provider,
cached_model_metadata(&selection.provider, &selection.model),
)
.is_none();
unknown_window.then(|| {
format!(
"{} has no usable_context_window, so the screen escalates every request",
selection.model
)
})
})
}
pub(super) fn text_screen_budget(
provider: &str,
metadata: Option<ModelMetadata>,
) -> Option<TranscriptBudget> {
let ollama = rho_providers::provider::provider_descriptor(provider)
.is_some_and(|descriptor| descriptor.id == ProviderId::Ollama);
if ollama {
metadata
.and_then(|metadata| metadata.usable_context_window)
.map(|window| transcript_budget(Some(window)))
} else {
Some(transcript_budget(
metadata.and_then(|metadata| metadata.display_context_window()),
))
}
}
pub(crate) struct ClassifierTrace {
pub screen: ScreenOutcome,
pub screen_allow_probability: Option<f64>,
pub result: anyhow::Result<ClassifierVerdict>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) enum ScreenOutcome {
Skipped,
Allowed,
Escalated,
Failed(String),
}
#[cfg(test)]
pub(super) async fn classify_capability_request_with_provider(
provider: &dyn ModelProvider,
screen: &Screen,
reasoning: ReasoningLevel,
budget: TranscriptBudget,
request: ClassifyRequest<'_>,
) -> ClassifierVerdict {
let pending_call_id = request.pending.tool_call_id().map(|id| id.as_str());
run_pipeline(
provider,
screen,
reasoning,
budget,
pending_call_id,
&request,
)
.await
.result
.unwrap_or_else(classifier_unavailable)
}
pub(super) fn transcript_budget(context_window: Option<u64>) -> TranscriptBudget {
let Some(window) = context_window else {
return TranscriptBudget::Unbounded;
};
let questions_tokens = [SCREEN_STAGE, REVIEW_STAGE]
.iter()
.map(|stage| estimate_text_tokens(&questions_block(stage.questions, stage.style)))
.max()
.unwrap_or_default();
let overhead = estimate_text_tokens(&system_prompt(CLASSIFIER_POLICY))
.saturating_add(questions_tokens)
.saturating_add(CLASSIFIER_OUTPUT_RESERVE_TOKENS);
TranscriptBudget::Tokens(window.saturating_sub(overhead))
}
async fn run_pipeline(
provider: &dyn ModelProvider,
screen: &Screen,
reasoning: ReasoningLevel,
budget: TranscriptBudget,
pending_call_id: Option<&str>,
request: &ClassifyRequest<'_>,
) -> ClassifierTrace {
match screen_step(provider, screen, budget, pending_call_id, request).await {
ScreenStep::Done(trace) => trace,
ScreenStep::Review {
screen,
screen_allow_probability,
transcript,
} => ClassifierTrace {
screen,
screen_allow_probability,
result: review(provider, reasoning, &request.scope(), &transcript).await,
},
}
}
pub(crate) enum ScreenStep {
Done(ClassifierTrace),
Review {
screen: ScreenOutcome,
screen_allow_probability: Option<f64>,
transcript: String,
},
}
pub(super) async fn screen_step(
provider: &dyn ModelProvider,
screen: &Screen,
budget: TranscriptBudget,
pending_call_id: Option<&str>,
request: &ClassifyRequest<'_>,
) -> ScreenStep {
let transcript =
match render_with_pending_call(request.history, request.pending, pending_call_id, budget) {
Ok(transcript) => transcript,
Err(error) => {
return ScreenStep::Done(ClassifierTrace {
screen: ScreenOutcome::Skipped,
screen_allow_probability: None,
result: Err(error),
})
}
};
let (screen, screen_allow_probability) = match run_screen(
provider,
screen,
request,
pending_call_id,
&transcript,
)
.await
{
Ok((ScreenVerdict::Allow, allow_probability)) => {
return ScreenStep::Done(ClassifierTrace {
screen: ScreenOutcome::Allowed,
screen_allow_probability: allow_probability,
result: Ok(ClassifierVerdict::Allow),
})
}
Ok((ScreenVerdict::Escalate, allow_probability)) => {
(ScreenOutcome::Escalated, allow_probability)
}
Err(error) => {
tracing::warn!(error = %error, "permission classifier screen failed; running review");
(ScreenOutcome::Failed(format!("{error:#}")), None)
}
};
ScreenStep::Review {
screen,
screen_allow_probability,
transcript,
}
}
pub(super) async fn review(
provider: &dyn ModelProvider,
reasoning: ReasoningLevel,
scope: &CallScope<'_>,
transcript: &str,
) -> anyhow::Result<ClassifierVerdict> {
let model = text_model(
provider,
scope,
REVIEW_STAGE.usage_purpose,
REVIEW_STAGE.style,
reasoning,
);
let decision = DecisionRequest::new(CLASSIFIER_POLICY, transcript, REVIEW_STAGE.questions);
review_verdict(&ask(&model, scope.cancellation, decision).await?)
}
async fn run_screen(
provider: &dyn ModelProvider,
screen: &Screen,
request: &ClassifyRequest<'_>,
pending_call_id: Option<&str>,
transcript: &str,
) -> anyhow::Result<(ScreenVerdict, Option<f64>)> {
let text_screen;
let (model, budget, allow_percent): (&dyn DecisionModel, _, _) = match screen {
Screen::Classifier => {
text_screen = text_model(
provider,
&request.scope(),
SCREEN_STAGE.usage_purpose,
SCREEN_STAGE.style,
ReasoningLevel::Low,
);
(&text_screen, None, DEFAULT_SCREEN_ALLOW_PERCENT)
}
Screen::Text { provider, budget } => {
let Some(budget) = budget else {
let identity = provider.identity();
bail!(
"text screen {} has no known served context window, so it may truncate the transcript; set usable_context_window for it in ~/.rho/models.toml",
rho_providers::provider::model_reference(&identity.provider, &identity.model)
);
};
text_screen = text_model(
provider.as_ref(),
&request.scope(),
SCREEN_STAGE.usage_purpose,
SCREEN_STAGE.style,
ReasoningLevel::Low,
);
let own = match budget {
TranscriptBudget::Unbounded => None,
TranscriptBudget::Tokens(_) => Some(*budget),
};
(&text_screen, own, DEFAULT_SCREEN_ALLOW_PERCENT)
}
Screen::Decision {
model,
allow_percent,
} => (
model.as_ref(),
model.state_budget().map(TranscriptBudget::Tokens),
*allow_percent,
),
};
let fitted;
let state = match budget {
None => transcript,
Some(budget) => {
fitted = render_with_pending_call(
request.history,
request.pending,
pending_call_id,
budget,
)?;
&fitted
}
};
let decision = DecisionRequest::new(CLASSIFIER_POLICY, state, SCREEN_STAGE.questions);
let answers = ask(model, &request.cancellation, decision).await?;
Ok((
screen_verdict(&answers, allow_percent),
screen_allow_probability(&answers),
))
}
pub(super) struct Stage {
pub usage_purpose: &'static str,
pub questions: &'static [Question<'static>],
pub style: AnswerStyle,
}
const SCREEN_STAGE: Stage = Stage {
usage_purpose: "permission-classifier-screen",
questions: std::slice::from_ref(&SCREEN_QUESTION),
style: AnswerStyle::Direct,
};
pub(super) const REVIEW_STAGE: Stage = Stage {
usage_purpose: "permission-classifier-review",
questions: std::slice::from_ref(&REVIEW_QUESTION),
style: AnswerStyle::Reasoned,
};
pub(super) fn text_model<'a>(
provider: &'a dyn ModelProvider,
scope: &CallScope<'a>,
usage_purpose: &'static str,
style: AnswerStyle,
reasoning: ReasoningLevel,
) -> TextModel<'a> {
TextModel {
provider,
agent_id: PERMISSION_CLASSIFIER_AGENT_ID,
usage_purpose,
reasoning,
style,
session_id: scope.session_id,
workspace_path: scope.workspace_path,
usage_recording: scope.usage_recording.clone(),
}
}
pub(super) async fn ask(
model: &dyn DecisionModel,
cancellation: &CancellationToken,
decision: DecisionRequest<'_>,
) -> anyhow::Result<Vec<Answer>> {
let answers = model.decide(decision, cancellation).await?;
decision.check_answers(&answers)?;
Ok(answers)
}
fn classifier_unavailable(error: anyhow::Error) -> ClassifierVerdict {
if let Some(over_budget) = error.downcast_ref::<TranscriptOverBudget>() {
tracing::warn!(error = %over_budget, "permission classifier transcript over budget");
return ClassifierVerdict::Deny {
reason: over_budget.to_string(),
};
}
if let Some(screen) = error.downcast_ref::<decision::ConfigError>() {
tracing::warn!(error = %screen, "permission classifier screen misconfigured");
return ClassifierVerdict::Deny {
reason: screen.to_string(),
};
}
tracing::warn!(error = %error, "permission classifier unavailable");
ClassifierVerdict::Deny {
reason: "classifier unavailable".into(),
}
}