use std::path::Path;
use futures_util::future::join_all;
use rho_providers::model::Message;
use rho_sdk::{
decision::{
text::{questions_block, system_prompt, AnswerStyle},
DecisionRequest,
},
model::context::estimate_text_tokens,
ApprovalRequest, CancellationToken, ProviderRequestUsageRecording, SessionId,
};
use super::{
batch_review_questions,
classify::{
ask, review, screen_step, text_model, CallScope, ClassifierModel, ClassifyRequest,
ScreenStep, REVIEW_STAGE,
},
review_verdict,
transcript::{render_with_pending_call, render_with_pending_calls, LabeledPending},
ClassifierVerdict, TranscriptBudget, TranscriptOverBudget, BATCH_POLICY, CLASSIFIER_POLICY,
};
const BATCH_REVIEW_PURPOSE: &str = "permission-classifier-batch-review";
pub(crate) struct BatchMember<'a> {
pub pending: &'a ApprovalRequest,
pub call_id: Option<&'a str>,
}
pub(crate) struct BatchRequest<'a> {
pub history: &'a [Message],
pub members: &'a [BatchMember<'a>],
pub cancellation: CancellationToken,
pub session_id: &'a SessionId,
pub workspace_path: &'a Path,
pub usage_recording: ProviderRequestUsageRecording,
}
impl BatchRequest<'_> {
fn member(&self, index: usize) -> ClassifyRequest<'_> {
ClassifyRequest {
history: self.history,
pending: self.members[index].pending,
cancellation: self.cancellation.clone(),
session_id: self.session_id,
workspace_path: self.workspace_path,
usage_recording: self.usage_recording.clone(),
}
}
fn scope(&self) -> CallScope<'_> {
CallScope {
cancellation: &self.cancellation,
session_id: self.session_id,
workspace_path: self.workspace_path,
usage_recording: &self.usage_recording,
}
}
}
fn label(index: usize) -> String {
format!("request_{}", index + 1)
}
impl ClassifierModel {
pub(crate) async fn screen_batch(&self, batch: &BatchRequest<'_>) -> Vec<ScreenStep> {
let requests: Vec<_> = (0..batch.members.len())
.map(|index| batch.member(index))
.collect();
join_all(requests.iter().zip(batch.members).map(|(request, member)| {
screen_step(
self.provider.as_ref(),
&self.screen,
self.budget,
member.call_id,
request,
)
}))
.await
}
pub(crate) async fn review_batch(
&self,
batch: &BatchRequest<'_>,
members: &[usize],
) -> Vec<anyhow::Result<ClassifierVerdict>> {
match members {
[] => Vec::new(),
[_] => self.review_each(batch, members).await,
_ => match self.review_together(batch, members).await {
Ok(verdicts) => verdicts.into_iter().map(Ok).collect(),
Err(error) => members.iter().map(|_| Err(copy_error(&error))).collect(),
},
}
}
pub(crate) async fn review_each(
&self,
batch: &BatchRequest<'_>,
members: &[usize],
) -> Vec<anyhow::Result<ClassifierVerdict>> {
let scope = batch.scope();
let mut results = Vec::with_capacity(members.len());
for &index in members {
let member = &batch.members[index];
let transcript = render_with_pending_call(
batch.history,
member.pending,
member.call_id,
self.budget,
);
results.push(match transcript {
Ok(transcript) => {
review(self.provider.as_ref(), self.reasoning, &scope, &transcript).await
}
Err(error) => Err(error),
});
}
results
}
async fn review_together(
&self,
batch: &BatchRequest<'_>,
members: &[usize],
) -> anyhow::Result<Vec<ClassifierVerdict>> {
let labels: Vec<String> = members.iter().map(|&index| label(index)).collect();
let questions = batch_review_questions(&labels);
let pendings: Vec<LabeledPending<'_>> = members
.iter()
.zip(&labels)
.map(|(&index, label)| LabeledPending {
label,
pending: batch.members[index].pending,
call_id: batch.members[index].call_id,
})
.collect();
let budget = batch_budget(self.budget, &questions);
let transcript = render_with_pending_calls(batch.history, &pendings, budget)?;
let scope = batch.scope();
let model = text_model(
self.provider.as_ref(),
&scope,
BATCH_REVIEW_PURPOSE,
REVIEW_STAGE.style,
self.reasoning,
);
let decision = DecisionRequest::new(BATCH_POLICY, &transcript, &questions);
let answers = ask(&model, scope.cancellation, decision).await?;
answers
.iter()
.map(|answer| review_verdict(std::slice::from_ref(answer)))
.collect()
}
}
fn batch_budget(
budget: TranscriptBudget,
questions: &[rho_sdk::decision::Question<'_>],
) -> TranscriptBudget {
let tokens = |text: &str| estimate_text_tokens(text);
let batch = tokens(&system_prompt(BATCH_POLICY))
+ tokens(&questions_block(questions, AnswerStyle::Reasoned));
let single = tokens(&system_prompt(CLASSIFIER_POLICY))
+ tokens(&questions_block(
REVIEW_STAGE.questions,
AnswerStyle::Reasoned,
));
budget.less(batch.saturating_sub(single))
}
fn copy_error(error: &anyhow::Error) -> anyhow::Error {
match error.downcast_ref::<TranscriptOverBudget>() {
Some(over_budget) => (*over_budget).into(),
None => anyhow::anyhow!("{error:#}"),
}
}