#[cfg(unix)]
use std::os::unix::process::ExitStatusExt;
#[cfg(windows)]
use std::os::windows::process::ExitStatusExt;
use std::process::ExitStatus;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::Duration;
use anyhow::Context;
use anyhow::Result;
use anyhow::anyhow;
use async_trait::async_trait;
use crankshaft_config::backend::tes::Config;
use crankshaft_events::Event;
use crankshaft_events::next_task_id;
use crankshaft_events::send_event;
use futures::FutureExt as _;
use futures::future::BoxFuture;
use nonempty::NonEmpty;
use tes::v1::Client;
use tes::v1::client::strategy::ExponentialFactorBackoff;
use tes::v1::types::requests::GetTaskParams;
use tes::v1::types::requests::View;
use tes::v1::types::task::State as TesState;
use tokio::select;
use tokio::sync::Semaphore;
use tokio::sync::broadcast;
use tokio::sync::oneshot;
use tokio_util::sync::CancellationToken;
use tracing::info;
use super::TaskRunError;
use crate::Task;
use crate::service::name::GeneratorIterator;
use crate::service::name::UniqueAlphanumeric;
use crate::service::runner::backend::tes::monitor::TaskMonitor;
mod monitor;
const DEFAULT_INTERVAL: Duration = Duration::from_secs(1);
const MAX_RETRY_DELAY: Duration = Duration::from_secs(30);
const DEFAULT_MAX_CONCURRENT_REQUESTS: usize = 10;
#[derive(Debug)]
struct BackendState {
client: Client,
interval: Duration,
retries: usize,
policy: ExponentialFactorBackoff,
permits: Semaphore,
events: Option<broadcast::Sender<Event>>,
}
impl BackendState {
fn policy(&self) -> impl Iterator<Item = Duration> + use<'_> {
self.policy.clone().take(self.retries)
}
}
#[derive(Debug)]
pub struct Backend {
names: Arc<Mutex<GeneratorIterator<UniqueAlphanumeric>>>,
state: Arc<BackendState>,
monitor: TaskMonitor,
}
impl Backend {
pub async fn initialize(
config: Config,
names: Arc<Mutex<GeneratorIterator<UniqueAlphanumeric>>>,
events: Option<broadcast::Sender<Event>>,
) -> Self {
let (url, http, interval) = config.into_parts();
let mut builder = Client::builder().url(url);
if let Some(auth) = &http.auth {
builder = builder.insert_header("Authorization", auth.header_value());
}
let state = Arc::new(BackendState {
client: builder.try_build().expect("client to build"),
interval: interval
.map(Duration::from_secs)
.unwrap_or(DEFAULT_INTERVAL),
retries: http.retries.unwrap_or_default() as usize,
policy: ExponentialFactorBackoff::from_millis(1000, 2.0).max_delay(MAX_RETRY_DELAY),
permits: Semaphore::new(
http.max_concurrency
.unwrap_or(DEFAULT_MAX_CONCURRENT_REQUESTS),
),
events,
});
let monitor_name = names.lock().unwrap().next().unwrap();
let monitor = TaskMonitor::new(monitor_name, state.clone()).await;
Self {
names,
state,
monitor,
}
}
async fn wait_task(
state: &BackendState,
monitor: &TaskMonitor,
task_id: u64,
task_name: &str,
tes_id: &str,
completed: oneshot::Receiver<Result<()>>,
) -> Result<NonEmpty<ExitStatus>, TaskRunError> {
info!(
"TES task `{tes_id}` (task `{task_name}`) has been created; waiting for task to start"
);
monitor.associate_task_id(task_id, tes_id.to_string()).await;
completed
.await
.context("failed to wait for task completion")??;
let permit = state
.permits
.acquire()
.await
.context("failed to acquire network request permit")?;
let task = state
.client
.get_task(
tes_id,
Some(&GetTaskParams { view: View::Full }),
state.policy(),
)
.await
.context("failed to get task information from TES server")?
.into_task()
.context("returned task is not a full view")?;
drop(permit);
let task_state = task.state.unwrap_or_default();
match task_state {
TesState::Unknown
| TesState::Queued
| TesState::Initializing
| TesState::Running
| TesState::Paused
| TesState::Canceling => Err(TaskRunError::Other(anyhow!(
"TES task is not in a completed state"
))),
TesState::Complete | TesState::ExecutorError => {
if task_state == TesState::Complete {
info!("TES task `{tes_id}` (task `{task_name}`) has completed");
} else {
info!("TES task `{tes_id}` (task `{task_name}`) has failed");
}
let logs = task.logs.unwrap_or_default();
let task = logs.last().context(
"invalid response from TES server: completed task is missing task logs",
)?;
Ok(NonEmpty::collect(task.logs.iter().map(|executor| {
#[cfg(unix)]
let status = ExitStatus::from_raw(executor.exit_code << 8);
#[cfg(windows)]
let status = ExitStatus::from_raw(executor.exit_code as u32);
status
}))
.context(
"invalid response from TES server: completed task is missing executor logs",
)?)
}
TesState::SystemError => {
info!("TES task `{tes_id}` (task `{task_name}`) has failed with a system error");
let messages = task
.logs
.unwrap_or_default()
.last()
.and_then(|l| l.system_logs.as_ref().map(|l| l.join("\n")))
.unwrap_or_default();
Err(TaskRunError::Other(anyhow!(
"task failed due to system error:\n\n{messages}"
)))
}
TesState::Canceled => {
info!("TES task `{tes_id}` (task `{task_name}`) has been canceled");
Err(TaskRunError::Canceled)
}
TesState::Preempted => {
info!("TES task `{tes_id}` (task `{task_name}`) has been preempted");
Err(TaskRunError::Preempted)
}
}
}
}
#[async_trait]
impl crate::Backend for Backend {
fn default_name(&self) -> &'static str {
"tes"
}
fn run(
&self,
task: Task,
token: CancellationToken,
) -> Result<BoxFuture<'static, Result<NonEmpty<ExitStatus>, TaskRunError>>> {
let task_id = next_task_id();
let names = self.names.clone();
let monitor = self.monitor.clone();
let state = self.state.clone();
Ok(async move {
let task_name = task.name.clone().unwrap_or_else(|| {
names.lock().unwrap().next().unwrap()
});
let mut task = tes::v1::types::requests::Task::try_from(task)?;
let (completed_tx, completed_rx) = oneshot::channel();
let tag = monitor.add_task(task_id, task_name.clone(), completed_tx).await;
task.tags
.get_or_insert_default()
.insert(monitor::CRANKSHAFT_GROUP_TAG_NAME.to_string(), tag);
let permit = state
.permits
.acquire()
.await
.context("failed to acquire network request permit")?;
let tes_id = select! {
biased;
_ = token.cancelled() => {
return Err(TaskRunError::Canceled);
}
res = state.client.create_task(&task, state.policy()) => {
res.context("failed to create task with TES server")?.id
}
};
drop(permit);
let task_token = CancellationToken::new();
send_event!(
state.events,
Event::TaskCreated {
id: task_id,
name: task_name.clone(),
tes_id: Some(tes_id.clone()),
token: task_token.clone()
}
);
let result = select! {
biased;
_ = task_token.cancelled() =>{
Err(TaskRunError::Canceled)
}
_ = token.cancelled() => {
Err(TaskRunError::Canceled)
}
res = Self::wait_task(&state, &monitor, task_id, &task_name, &tes_id, completed_rx) => {
res
}
};
if let Err(TaskRunError::Canceled) = &result {
let permit = state
.permits
.acquire()
.await
.context("failed to acquire permit")?;
info!("canceling TES task `{tes_id}` (task `{task_name}`)");
state
.client
.cancel_task(&tes_id, state.policy())
.await
.context("failed to cancel task with TES server")?;
drop(permit);
}
monitor.remove_task(&tes_id).await;
match &result {
Ok(statuses) => send_event!(
state.events,
Event::TaskCompleted {
id: task_id,
exit_statuses: statuses.clone(),
}
),
Err(TaskRunError::Canceled) => {
send_event!(state.events, Event::TaskCanceled { id: task_id })
}
Err(TaskRunError::Preempted) => {
send_event!(state.events, Event::TaskPreempted { id: task_id })
}
Err(TaskRunError::Other(e)) => send_event!(
state.events,
Event::TaskFailed {
id: task_id,
message: format!("{e:#}")
}
),
}
result
}
.boxed())
}
}