use std::collections::HashMap;
use std::collections::HashSet;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::SystemTime;
use anyhow::Context;
use anyhow::Result;
use anyhow::anyhow;
use crankshaft_events::Event;
use crankshaft_events::send_event;
use tes::v1::types::requests::ListTasksParams;
use tes::v1::types::requests::MAX_PAGE_SIZE;
use tes::v1::types::requests::View;
use tes::v1::types::responses::ListTasks;
use tes::v1::types::task::State as TesState;
use tokio::select;
use tokio::sync::oneshot;
use tokio::time::MissedTickBehavior;
use tracing::debug;
use tracing::info;
pub const CRANKSHAFT_GROUP_TAG_NAME: &str = "crankshaft-task-group";
#[derive(Debug)]
struct Task {
name: String,
completed: oneshot::Sender<Result<()>>,
}
#[derive(Debug, Default)]
struct TaskMonitorState {
tag: String,
tasks: HashMap<u64, Task>,
ids: HashMap<String, u64>,
running: HashSet<u64>,
}
#[derive(Debug, Clone)]
pub struct TaskMonitor {
name: Arc<String>,
state: Arc<Mutex<TaskMonitorState>>,
_drop: Arc<oneshot::Sender<()>>,
}
impl TaskMonitor {
pub async fn new(name: String, backend_state: Arc<super::BackendState>) -> Self {
let state: Arc<Mutex<TaskMonitorState>> = Default::default();
let (tx, rx) = oneshot::channel();
tokio::spawn(Self::monitor(state.clone(), backend_state, rx));
Self {
name: name.into(),
state,
_drop: tx.into(),
}
}
pub async fn add_task(
&self,
id: u64,
name: String,
completed: oneshot::Sender<Result<()>>,
) -> String {
let mut state = self.state.lock().expect("failed to lock TES monitor state");
if state.tasks.is_empty() {
state.running.clear();
state.tag = format!(
"{name}-{timestamp}-{id}",
name = self.name,
timestamp = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
);
}
state.tasks.insert(id, Task { name, completed });
state.tag.clone()
}
pub async fn associate_task_id(&self, id: u64, tes_id: String) {
let mut state = self.state.lock().expect("failed to lock TES monitor state");
state.ids.insert(tes_id, id);
}
pub async fn remove_task(&self, tes_id: &str) {
let mut state = self.state.lock().expect("failed to lock TES monitor state");
if let Some(id) = state.ids.get(tes_id).copied() {
state.tasks.remove(&id);
state.running.remove(&id);
}
}
async fn update_tasks(
state: &Arc<Mutex<TaskMonitorState>>,
backend_state: &super::BackendState,
) {
let mut page_token = None;
loop {
let tag = {
let state = state.lock().expect("failed to TES lock monitor state");
if state.tasks.is_empty() {
return;
}
assert!(!state.tag.is_empty(), "should have a current tag");
debug!(
"querying for the state of TES tasks with tag `{tag}` and page token \
`{page_token:?}`",
tag = state.tag
);
state.tag.clone()
};
let list = async {
let permit = backend_state
.permits
.acquire()
.await
.context("failed to acquire network request permit")?;
let result = backend_state
.client
.list_tasks(
Some(&ListTasksParams {
tag_keys: Some(vec![CRANKSHAFT_GROUP_TAG_NAME.to_string()]),
tag_values: Some(vec![tag]),
page_size: Some(MAX_PAGE_SIZE - 1),
page_token,
view: Some(View::Minimal),
..Default::default()
}),
backend_state.policy(),
)
.await
.context("failed to get task information from TES server");
drop(permit);
result
};
match list.await {
Ok(ListTasks {
tasks: tes_tasks,
next_page_token,
}) => {
let mut state = state.lock().expect("failed to TES lock monitor state");
for task in tes_tasks
.into_iter()
.map(|t| t.into_minimal().expect("task should be minimal"))
{
match task.state.unwrap_or_default() {
TesState::Running | TesState::Paused => {
if let Some(id) = state.ids.get(&task.id).copied()
&& state.running.insert(id)
{
if let Some(Task { name, .. }) = state.tasks.get(&id) {
info!(
"TES task `{tes_id}` (task `{name}`) is now running",
tes_id = task.id
);
}
send_event!(backend_state.events, Event::TaskStarted { id });
}
}
TesState::Complete
| TesState::ExecutorError
| TesState::SystemError
| TesState::Canceled
| TesState::Preempted => {
if let Some(id) = state.ids.remove(&task.id) {
state.running.remove(&id);
if let Some(task) = state.tasks.remove(&id) {
let _ = task.completed.send(Ok(()));
}
}
}
_ => {}
}
}
if next_page_token
.as_ref()
.map(|t| t.is_empty())
.unwrap_or(true)
{
break;
}
page_token = next_page_token;
}
Err(e) => {
let mut state = state.lock().expect("failed to TES lock monitor state");
state.running.clear();
for (_, task) in state.tasks.drain() {
let _ = task
.completed
.send(Err(anyhow!("failed to monitor TES tasks: {e:#}")));
}
break;
}
}
}
}
async fn monitor(
state: Arc<Mutex<TaskMonitorState>>,
backend_state: Arc<super::BackendState>,
mut drop: oneshot::Receiver<()>,
) {
info!(
"TES task monitor is starting with polling interval of {interval} seconds",
interval = backend_state.interval.as_secs()
);
let mut timer = tokio::time::interval(backend_state.interval);
timer.set_missed_tick_behavior(MissedTickBehavior::Delay);
loop {
select! {
_ = &mut drop => break,
_ = timer.tick() => Self::update_tasks(&state, backend_state.as_ref()).await,
}
}
info!("TES task monitor has shut down");
}
}