use std::collections::HashMap;
use std::collections::HashSet;
use std::sync::Arc;
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::mpsc;
use tokio::sync::oneshot;
use tokio::time::MissedTickBehavior;
use tracing::debug;
use tracing::info;
const MONITOR_CAPACITY: usize = 100;
pub const CRANKSHAFT_GROUP_TAG_NAME: &str = "crankshaft-task-group";
#[derive(Debug)]
struct AddTaskRequest {
id: u64,
name: String,
completed: oneshot::Sender<Result<()>>,
response: oneshot::Sender<AddTaskResponse>,
}
#[derive(Debug)]
struct AddTaskResponse {
tag: String,
}
#[derive(Debug)]
struct AssociateTaskIdRequest {
id: u64,
tes_id: String,
}
#[derive(Debug)]
struct RemoveTaskRequest {
tes_id: String,
}
#[derive(Debug)]
enum MonitorRequest {
AddTask(AddTaskRequest),
AssociateTaskId(AssociateTaskIdRequest),
RemoveTask(RemoveTaskRequest),
}
#[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(mpsc::Sender<MonitorRequest>);
impl TaskMonitor {
pub async fn new(name: String, backend_state: Arc<super::BackendState>) -> Self {
let (tx, rx) = mpsc::channel(MONITOR_CAPACITY);
tokio::spawn(Self::monitor(name, backend_state, rx));
Self(tx)
}
pub async fn add_task(
&self,
id: u64,
name: String,
completed: oneshot::Sender<Result<()>>,
) -> String {
let (tx, rx) = oneshot::channel();
self.0
.send(MonitorRequest::AddTask(AddTaskRequest {
id,
name,
completed,
response: tx,
}))
.await
.expect("failed to send request");
rx.await.map(|r| r.tag).expect("failed to receive response")
}
pub async fn associate_task_id(&self, id: u64, tes_id: String) {
self.0
.send(MonitorRequest::AssociateTaskId(AssociateTaskIdRequest {
id,
tes_id,
}))
.await
.expect("failed to send request");
}
pub async fn remove_task(&self, tes_id: String) {
self.0
.send(MonitorRequest::RemoveTask(RemoveTaskRequest { tes_id }))
.await
.expect("failed to send request");
}
fn handle_add_task(state: &mut TaskMonitorState, name: &str, req: AddTaskRequest) {
if state.tasks.is_empty() {
state.running.clear();
state.tag = format!(
"{name}-{timestamp}-{id}",
timestamp = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
id = req.id,
);
}
state.tasks.insert(
req.id,
Task {
name: req.name,
completed: req.completed,
},
);
req.response
.send(AddTaskResponse {
tag: state.tag.clone(),
})
.expect("failed to send add task response");
}
fn handle_associate_task_id(state: &mut TaskMonitorState, req: AssociateTaskIdRequest) {
state.ids.insert(req.tes_id, req.id);
}
fn handle_remove_task(state: &mut TaskMonitorState, req: RemoveTaskRequest) {
if let Some(id) = state.ids.get(&req.tes_id) {
state.tasks.remove(id);
state.running.remove(id);
}
}
async fn update_tasks(state: &mut TaskMonitorState, backend_state: &super::BackendState) {
if state.tasks.is_empty() {
return;
}
assert!(!state.tag.is_empty(), "should have a current tag");
let mut page_token = None;
loop {
debug!(
"querying for the state of TES tasks with tag `{tag}` and page token \
`{page_token:?}`",
tag = state.tag
);
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![state.tag.clone()]),
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,
}) => {
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) {
if let Some(Task { name, .. }) = state.tasks.get(id) {
if state.running.insert(*id) {
info!(
"TES task `{tes_id}` (task `{name}`) is now \
running",
tes_id = task.id
);
send_event!(
backend_state.events,
Event::TaskStarted { id: *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.is_none() {
break;
}
page_token = next_page_token;
}
Err(e) => {
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(
monitor_name: String,
backend_state: Arc<super::BackendState>,
mut rx: mpsc::Receiver<MonitorRequest>,
) {
info!(
"TES task monitor is starting with polling interval of {interval} seconds",
interval = backend_state.interval.as_secs()
);
let mut state = TaskMonitorState::default();
let mut timer = tokio::time::interval(backend_state.interval);
timer.set_missed_tick_behavior(MissedTickBehavior::Delay);
loop {
select! {
msg = rx.recv() => match msg {
Some(request) => match request {
MonitorRequest::AddTask(req) => Self::handle_add_task(&mut state, &monitor_name, req),
MonitorRequest::AssociateTaskId(req) => Self::handle_associate_task_id(&mut state, req),
MonitorRequest::RemoveTask(req) => Self::handle_remove_task(&mut state, req),
},
None => break,
},
_ = timer.tick() => Self::update_tasks(&mut state, backend_state.as_ref()).await,
}
}
info!("TES task monitor has shut down");
}
}