#[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::TaskId;
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::error;
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;
use crate::task::ExecutionResult;
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,
}
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>>>,
) -> 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),
),
});
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: TaskId,
task_name: &str,
tes_id: &str,
completed: oneshot::Receiver<Result<()>>,
) -> Result<NonEmpty<ExecutionResult>, 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_log = logs.last().context(
"invalid response from TES server: completed task is missing task logs",
)?;
Ok(
NonEmpty::collect(task_log.logs.iter().enumerate().map(|(idx, executor)| {
#[cfg(unix)]
let status = ExitStatus::from_raw(executor.exit_code << 8);
#[cfg(windows)]
let status = ExitStatus::from_raw(executor.exit_code as u32);
ExecutionResult {
image: task.executors.get(idx).map(|e| e.image.clone()),
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,
events: Option<broadcast::Sender<Event>>,
token: CancellationToken,
) -> Result<BoxFuture<'static, Result<NonEmpty<ExecutionResult>, 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(), events.clone(), completed_tx)
.await;
let mut tes_id = None;
let result = async {
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 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);
tes_id = Some(id);
let task_token = CancellationToken::new();
send_event!(
events,
Event::TaskCreated {
id: task_id,
name: task_name.clone(),
tes_id: tes_id.clone(),
token: task_token.clone()
}
);
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.as_deref().unwrap(), completed_rx) => {
res
}
}
}
.await;
monitor.remove_task(task_id).await;
if let (Some(tes_id), Err(TaskRunError::Canceled)) = (&tes_id, &result) {
match state.permits.acquire().await {
Ok(permit) => {
info!("canceling TES task `{tes_id}` (task `{task_name}`)");
if let Err(e) = state.client.cancel_task(&tes_id, state.policy()).await {
error!("failed to cancel task with TES server: {e:#}");
}
drop(permit);
}
Err(e) => {
error!("failed to acquire permit to cancel TES task: {e}");
}
}
}
if tes_id.is_some() {
match &result {
Ok(results) => send_event!(
events,
Event::TaskCompleted {
id: task_id,
exit_statuses: NonEmpty::collect(results.iter().map(|r| r.status))
.unwrap(),
}
),
Err(TaskRunError::Canceled) => {
send_event!(events, Event::TaskCanceled { id: task_id })
}
Err(TaskRunError::Preempted) => {
send_event!(events, Event::TaskPreempted { id: task_id })
}
Err(TaskRunError::Other(e)) => send_event!(
events,
Event::TaskFailed {
id: task_id,
message: format!("{e:#}")
}
),
}
}
result
}
.boxed())
}
}