use futures_util::{future::BoxFuture, FutureExt};
use rho_providers::model::{ModelMetadata, ReasoningRequestSource};
use rho_sdk::ReasoningLevel;
use tokio::task::{AbortHandle, JoinError, JoinHandle};
use super::{
changelog_command::ChangelogFetchResult, github_pr::GithubPrLookup,
limits_command::LimitsFetchResult, limits_command::LimitsSectionId, App, DefaultTerminal,
InteractiveRuntime,
};
use crate::doctor::{DoctorProbeId, DoctorProbeOutcome};
#[cfg(test)]
#[path = "background_tasks_tests.rs"]
mod tests;
#[derive(Clone, Debug, PartialEq, Eq)]
pub(super) enum TaskId {
ModelMetadata,
UpdateNotice,
CustomModels,
CursorModels,
SyntaxWarmup,
InteractiveLogin,
GithubPr,
UsageLimits(LimitsSectionId),
DoctorProbe(DoctorProbeId),
InfoRuntimes,
InfoTree,
SpendLoad,
DiffPatch,
Changelog,
WebSearchTest,
}
pub(super) enum TaskOutput {
Ui(UiOutput),
Session(SessionOutput),
}
pub(super) enum UiOutput {
GithubPr(Result<GithubPrLookup, JoinError>),
UsageLimits(LimitsSectionId, LimitsFetchResult),
DoctorProbe(DoctorProbeOutcome),
InfoRuntimes(Result<Vec<String>, JoinError>),
InfoTree(Result<anyhow::Result<crate::session::tree::SessionTreeFacts>, JoinError>),
SpendLoad(super::spend_overlay::LoadResult),
DiffPatch(usize, anyhow::Result<Vec<rho_tools::tool_card::DiffRow>>),
Changelog(Result<ChangelogFetchResult, JoinError>),
WebSearchTest(Result<Result<usize, rho_tools::tool::ToolError>, JoinError>),
}
pub(super) enum SessionOutput {
ModelMetadata {
reasoning_at_start: (ReasoningLevel, ReasoningRequestSource),
metadata: Result<Option<ModelMetadata>, JoinError>,
},
UpdateNotice(Result<Option<String>, JoinError>),
CustomModels,
CursorModels(Result<crate::cursor_runtime::models::RefreshResult, JoinError>),
SyntaxWarmup,
InteractiveLogin(super::login::FinishedInteractiveLogin),
}
impl From<UiOutput> for TaskOutput {
fn from(output: UiOutput) -> Self {
Self::Ui(output)
}
}
impl From<SessionOutput> for TaskOutput {
fn from(output: SessionOutput) -> Self {
Self::Session(output)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) enum OnCancel {
Await,
Detach,
RunToCompletion,
}
struct RunningTask {
id: TaskId,
abort: AbortHandle,
on_cancel: OnCancel,
output: BoxFuture<'static, TaskOutput>,
}
#[derive(Default)]
pub(super) struct BackgroundTasks {
running: Vec<RunningTask>,
ready: Vec<(TaskId, TaskOutput)>,
}
impl BackgroundTasks {
pub(super) fn from_startup(services: &mut super::ApplicationServices) -> Self {
let mut tasks = Self::default();
if let Some(handle) = services.pending_update_notice.take() {
tasks.track(TaskId::UpdateNotice, handle, OnCancel::Await, |result| {
SessionOutput::UpdateNotice(result).into()
});
}
if let Some(handle) = services.pending_custom_models.take() {
tasks.track(
TaskId::CustomModels,
handle,
OnCancel::RunToCompletion,
|_| SessionOutput::CustomModels.into(),
);
}
if let Some(handle) = services.pending_syntax_warmup.take() {
tasks.track(TaskId::SyntaxWarmup, handle, OnCancel::Detach, |_| {
SessionOutput::SyntaxWarmup.into()
});
}
tasks
}
pub(super) fn spawn<T: Send + 'static>(
&mut self,
id: TaskId,
future: impl std::future::Future<Output = T> + Send + 'static,
wrap: impl FnOnce(Result<T, JoinError>) -> TaskOutput + Send + 'static,
) {
self.track(id, tokio::spawn(future), OnCancel::Await, wrap);
}
pub(super) fn spawn_blocking<T: Send + 'static>(
&mut self,
id: TaskId,
work: impl FnOnce() -> T + Send + 'static,
wrap: impl FnOnce(Result<T, JoinError>) -> TaskOutput + Send + 'static,
) {
self.track(
id,
tokio::task::spawn_blocking(work),
OnCancel::Detach,
wrap,
);
}
pub(super) fn track<T: Send + 'static>(
&mut self,
id: TaskId,
handle: JoinHandle<T>,
on_cancel: OnCancel,
wrap: impl FnOnce(Result<T, JoinError>) -> TaskOutput + Send + 'static,
) {
self.running.push(RunningTask {
id,
abort: handle.abort_handle(),
on_cancel,
output: async move { wrap(handle.await) }.boxed(),
});
}
pub(super) fn ids(&self) -> impl Iterator<Item = &TaskId> {
self.running
.iter()
.map(|task| &task.id)
.chain(self.ready.iter().map(|(id, _)| id))
}
pub(super) fn contains(&self, matches: impl Fn(&TaskId) -> bool) -> bool {
self.ids().any(matches)
}
pub(super) fn has_pending(&self) -> bool {
!self.running.is_empty() || !self.ready.is_empty()
}
pub(super) fn has_finished(&self) -> bool {
!self.ready.is_empty() || self.running.iter().any(|task| task.abort.is_finished())
}
pub(super) fn take_next_finished(
&mut self,
accepts: impl Fn(&TaskOutput) -> bool,
) -> Option<TaskOutput> {
let mut index = 0;
while index < self.running.len() {
if let Some(output) = (&mut self.running[index].output).now_or_never() {
let task = self.running.remove(index);
self.ready.push((task.id, output));
} else {
index += 1;
}
}
let index = self.ready.iter().position(|(_, output)| accepts(output))?;
Some(self.ready.remove(index).1)
}
pub(super) fn abort(&mut self, matches: impl Fn(&TaskId) -> bool) {
for task in self.remove(matches) {
if task.on_cancel != OnCancel::RunToCompletion {
task.abort.abort();
}
}
}
pub(super) async fn cancel(&mut self, matches: impl Fn(&TaskId) -> bool) {
for task in self.remove(matches) {
match task.on_cancel {
OnCancel::Await => {
task.abort.abort();
let _ = task.output.await;
}
OnCancel::Detach => task.abort.abort(),
OnCancel::RunToCompletion => {}
}
}
}
pub(super) async fn cancel_all(&mut self) {
self.cancel(|_| true).await;
}
fn remove(&mut self, matches: impl Fn(&TaskId) -> bool) -> Vec<RunningTask> {
self.ready.retain(|(id, _)| !matches(id));
let (removed, kept) = std::mem::take(&mut self.running)
.into_iter()
.partition(|task| matches(&task.id));
self.running = kept;
removed
}
}
impl SessionOutput {
fn waits_for_idle_session(&self) -> bool {
matches!(self, Self::ModelMetadata { .. })
}
}
impl App {
pub(super) fn apply_finished_ui_tasks(&mut self) -> bool {
let mut changed = false;
while let Some(output) = self
.tasks
.take_next_finished(|output| matches!(output, TaskOutput::Ui(_)))
{
let TaskOutput::Ui(output) = output else {
unreachable!("only Ui outputs are accepted");
};
changed |= self.apply_ui_output(output);
}
changed
}
pub(super) async fn apply_finished_tasks(
&mut self,
terminal: &mut DefaultTerminal,
agent: &mut InteractiveRuntime,
) -> anyhow::Result<bool> {
let mut changed = false;
while let Some(output) = self.take_next_finished_task(agent) {
changed |= match output {
TaskOutput::Ui(output) => self.apply_ui_output(output),
TaskOutput::Session(output) => {
self.apply_session_output(output, terminal, agent).await?
}
};
}
Ok(changed)
}
pub(super) fn take_next_finished_task(
&mut self,
agent: &InteractiveRuntime,
) -> Option<TaskOutput> {
let session_busy = agent.is_session_busy();
self.tasks.take_next_finished(|output| match output {
TaskOutput::Session(output) => !(session_busy && output.waits_for_idle_session()),
TaskOutput::Ui(_) => true,
})
}
fn apply_ui_output(&mut self, output: UiOutput) -> bool {
match output {
UiOutput::GithubPr(result) => self.apply_github_pr(result),
UiOutput::UsageLimits(id, result) => self.apply_limits_fetch(id, result),
UiOutput::DoctorProbe(outcome) => self.apply_doctor_probe(&outcome),
UiOutput::InfoRuntimes(result) => self.apply_info_runtimes_result(result),
UiOutput::InfoTree(result) => self.apply_info_tree_result(result),
UiOutput::SpendLoad(result) => self.apply_spend_load(result),
UiOutput::DiffPatch(index, result) => self.apply_diff_patch(index, result),
UiOutput::Changelog(result) => self.apply_changelog_fetch(result),
UiOutput::WebSearchTest(result) => self.apply_web_search_test(result),
}
}
async fn apply_session_output(
&mut self,
output: SessionOutput,
terminal: &mut DefaultTerminal,
agent: &mut InteractiveRuntime,
) -> anyhow::Result<bool> {
match output {
SessionOutput::ModelMetadata {
reasoning_at_start,
metadata,
} => {
self.apply_model_metadata(agent, reasoning_at_start, metadata)
.await;
Ok(false)
}
SessionOutput::UpdateNotice(result) => {
if let Ok(Some(notice)) = result {
self.info.services.update_notice = Some(notice);
}
Ok(false)
}
SessionOutput::CustomModels => Ok(false),
SessionOutput::CursorModels(result) => {
self.apply_cursor_model_refresh(result);
Ok(false)
}
SessionOutput::SyntaxWarmup => {
self.history.invalidate_from(0);
Ok(true)
}
SessionOutput::InteractiveLogin(finished) => {
self.apply_interactive_login(finished, terminal, agent)
.await?;
Ok(false)
}
}
}
}