use std::num::NonZeroUsize;
use rho_sdk::{
model::{ModelEvent, ModelRequest, ModelResponse, ModelUsage},
provider::{
provider_event_channel, ModelProvider, ModelRequestOptions, ProviderRequestEvent,
ProviderStreamEvent,
},
ProviderError, ProviderRequestOutcome, ProviderRequestUsageContext, ProviderRequestUsageEvent,
ProviderRequestUsageRecording,
};
const EVENT_CAPACITY: usize = 16;
pub(crate) async fn send_recorded(
provider: &dyn ModelProvider,
request: ModelRequest<'_>,
context: ProviderRequestUsageContext,
recording: ProviderRequestUsageRecording,
) -> Result<(ModelResponse, ModelUsage), ProviderError> {
send_recorded_observing(provider, request, context, recording, 1, |_| {}).await
}
pub(crate) async fn send_recorded_observing(
provider: &dyn ModelProvider,
request: ModelRequest<'_>,
context: ProviderRequestUsageContext,
recording: ProviderRequestUsageRecording,
first_attempt_index: usize,
on_event: impl FnMut(&ProviderStreamEvent) + Send,
) -> Result<(ModelResponse, ModelUsage), ProviderError> {
let mut next_attempt_index = first_attempt_index;
send_recorded_with(
provider,
RecordedRequest {
request,
options: ModelRequestOptions::default(),
context,
recording,
},
&mut next_attempt_index,
&mut ModelUsage::default(),
on_event,
)
.await
}
pub(crate) struct RecordedRequest<'a> {
pub(crate) request: ModelRequest<'a>,
pub(crate) options: ModelRequestOptions,
pub(crate) context: ProviderRequestUsageContext,
pub(crate) recording: ProviderRequestUsageRecording,
}
pub(crate) async fn send_recorded_with(
provider: &dyn ModelProvider,
recorded: RecordedRequest<'_>,
next_attempt_index: &mut usize,
reported: &mut ModelUsage,
mut on_event: impl FnMut(&ProviderStreamEvent) + Send,
) -> Result<(ModelResponse, ModelUsage), ProviderError> {
let RecordedRequest {
request,
options,
context,
recording,
} = recorded;
let cancellation = request.cancellation.clone();
let (events, mut receiver) =
provider_event_channel(NonZeroUsize::new(EVENT_CAPACITY).expect("capacity is nonzero"));
let provider_call = provider.send_turn_stream_with_options(request, options, events);
let collect_usage = async {
let mut usage = ModelUsage::default();
let mut failed_attempts = Vec::new();
while let Some(event) = receiver.recv_stream_event().await {
on_event(&event);
match event {
ProviderStreamEvent::Model(ModelEvent::Usage(partial)) => {
usage = usage.saturating_add(&partial);
}
ProviderStreamEvent::Model(
ModelEvent::OutputDelta(_)
| ModelEvent::ReasoningDelta(_)
| ModelEvent::ReasoningSummaryDelta(_)
| ModelEvent::WebSearch(_)
| ModelEvent::ToolCallDelta { .. }
| ModelEvent::ProviderContext { .. }
| ModelEvent::GenerationOutputTokens(_)
| ModelEvent::HostedToolActivity { .. }
| ModelEvent::ServiceTierFallback { .. },
) => {}
ProviderStreamEvent::Request(ProviderRequestEvent::RequestAttemptFailed {
kind,
usage: attempt_usage,
}) => {
failed_attempts.push((kind, usage.saturating_add(&attempt_usage)));
usage = ModelUsage::default();
}
}
}
(usage, failed_attempts)
};
let (result, (usage, failed_attempts)) = tokio::join!(provider_call, collect_usage);
let outcome = match &result {
Ok(_) => ProviderRequestOutcome::Completed,
Err(_) if cancellation.is_cancelled() => ProviderRequestOutcome::Cancelled,
Err(error) => ProviderRequestOutcome::Failed(error.kind()),
};
*next_attempt_index = (*next_attempt_index).max(1);
for (kind, usage) in failed_attempts {
*reported = reported.saturating_add(&usage);
recording
.record(ProviderRequestUsageEvent::observed(
context.clone().with_attempt_index(*next_attempt_index),
usage,
ProviderRequestOutcome::Failed(kind),
))
.await;
*next_attempt_index += 1;
}
recording
.record(ProviderRequestUsageEvent::observed(
context.with_attempt_index(*next_attempt_index),
usage.clone(),
outcome,
))
.await;
*reported = reported.saturating_add(&usage);
*next_attempt_index += 1;
result.map(|response| (response, usage))
}