pub mod layer_builder;
use crate::{artifact_pusher, preparer, progress::ProgressTracker, router};
use anyhow::{anyhow, bail, Error, Result};
use maelstrom_base::proto::Hello;
use maelstrom_client_base::{
spec::{self, ContainerSpec},
AcceptInvalidRemoteContainerTlsCerts, CacheDir, ClusterCommunicationStrategy,
GitHubClusterCommunicationStrategy, IntrospectResponse, JobStatus, ProjectDir,
TcpClusterCommunicationStrategy, MANIFEST_DIR, STUB_MANIFEST_DIR, SYMLINK_MANIFEST_DIR,
};
use maelstrom_container::ContainerImageDepotDir;
use maelstrom_github::GitHubClient;
use maelstrom_util::{
async_fs,
broker_connection::{
BrokerConnectionFactory as _, BrokerReadConnection as _, BrokerWriteConnection as _,
GitHubQueueBrokerConnectionFactory, TcpBrokerConnectionFactory,
},
config::common::{CacheSize, InlineLimit, Slots},
root::RootBuf,
sync::EventSender,
};
use maelstrom_worker::local_worker;
use slog::{debug, o, Logger};
use std::mem;
use tokio::{
sync::{mpsc, oneshot, RwLock},
task::JoinSet,
};
pub struct Client(RwLock<ClientStateWrapper>);
enum ClientStateWrapper {
NotRunning(Option<EventSender>),
Running(ClientState),
}
impl ClientStateWrapper {
fn get_running(&self) -> Result<&ClientState> {
match self {
Self::NotRunning(Some(_)) => Err(anyhow!("client not yet started")),
Self::NotRunning(None) => Err(anyhow!("client started and stopped")),
Self::Running(state) => Ok(state),
}
}
fn take_to_start(&mut self) -> Result<EventSender> {
match self {
Self::NotRunning(Some(_)) => {}
Self::NotRunning(None) => {
bail!("client already started and stopped");
}
Self::Running(_) => {
bail!("client already started");
}
}
let Self::NotRunning(Some(done)) = mem::replace(self, Self::NotRunning(None)) else {
unreachable!();
};
Ok(done)
}
fn take_to_stop(&mut self) -> Result<ClientState> {
match self {
Self::NotRunning(Some(_)) => {
bail!("client not yet started");
}
Self::NotRunning(None) => {
bail!("client already stopped");
}
Self::Running(_) => {}
}
let Self::Running(state) = mem::replace(self, Self::NotRunning(None)) else {
unreachable!();
};
Ok(state)
}
}
struct ClientState {
router_sender: router::Sender,
artifact_upload_tracker: ProgressTracker,
image_download_tracker: ProgressTracker,
log: Logger,
preparer_sender: preparer::task::Sender,
tasks: local_worker::Tasks,
}
const MANIFEST_INLINE_LIMIT: u64 = 200 * 1024;
const MAX_PENDING_LAYER_BUILDS: usize = 10;
impl Client {
pub fn new(done: EventSender) -> Self {
Self(RwLock::new(ClientStateWrapper::NotRunning(Some(done))))
}
#[allow(clippy::too_many_arguments)]
async fn try_to_start(
accept_invalid_remote_container_tls_certs: AcceptInvalidRemoteContainerTlsCerts,
cache_dir: RootBuf<CacheDir>,
cache_size: CacheSize,
cluster_communication_strategy: Option<ClusterCommunicationStrategy>,
container_image_depot_cache_dir: RootBuf<ContainerImageDepotDir>,
done: EventSender,
inline_limit: InlineLimit,
log: Logger,
project_dir: RootBuf<ProjectDir>,
slots: Slots,
) -> Result<ClientState> {
let fs = async_fs::Fs::new();
debug!(log, "client starting";
"cluster_communication_strategy" => ?cluster_communication_strategy,
"project_dir" => ?project_dir,
"cache_dir" => ?cache_dir,
"container_image_depot_cache_dir" => ?container_image_depot_cache_dir,
"cache_size" => ?cache_size,
"inline_limit" => ?inline_limit,
"slots" => ?slots,
);
let extra = 1 +
MAX_PENDING_LAYER_BUILDS * 3 +
artifact_pusher::MAX_CLIENT_UPLOADS * 2; local_worker::check_open_file_limit(&log, slots, extra as u64)?;
if fs.exists((**cache_dir).join(MANIFEST_DIR)).await {
fs.remove_dir_all((**cache_dir).join(MANIFEST_DIR)).await?;
}
const LOCAL_WORKER_DIR: &str = "local-worker";
for d in [STUB_MANIFEST_DIR, SYMLINK_MANIFEST_DIR, LOCAL_WORKER_DIR] {
fs.create_dir_all((**cache_dir).join(d)).await?;
}
let artifact_upload_tracker = ProgressTracker::default();
let image_download_tracker = ProgressTracker::default();
let mut join_set = JoinSet::new();
let (artifact_pusher_sender, artifact_pusher_receiver) = artifact_pusher::channel();
let (broker_sender, broker_receiver) = mpsc::unbounded_channel();
let (local_worker_sender, local_worker_receiver) = local_worker::channel();
let (preparer_sender, preparer_receiver) = preparer::task::channel();
let (router_sender, router_receiver) = router::channel();
let standalone;
if let Some(cluster_communication_strategy) = cluster_communication_strategy {
standalone = false;
match cluster_communication_strategy {
ClusterCommunicationStrategy::Tcp(TcpClusterCommunicationStrategy { broker }) => {
let broker_connection_factory = TcpBrokerConnectionFactory::new(broker, &log);
let (broker_socket_read_half, broker_socket_write_half) =
broker_connection_factory.connect(&Hello::Client).await?;
join_set.spawn(broker_socket_read_half.read_messages(
router_sender.clone(),
log.new(o!("task" => "broker socket reader")),
router::Message::Broker,
));
join_set.spawn(broker_socket_write_half.write_messages(
broker_receiver,
log.new(o!("task" => "broker socket writer")),
));
}
ClusterCommunicationStrategy::GitHub(GitHubClusterCommunicationStrategy {
ref token,
ref url,
}) => {
let broker_connection_factory =
GitHubQueueBrokerConnectionFactory::new(&log, token.clone(), url.clone())?;
let (broker_socket_read_half, broker_socket_write_half) =
broker_connection_factory.connect(&Hello::Client).await?;
join_set.spawn(broker_socket_read_half.read_messages(
router_sender.clone(),
log.new(o!("task" => "broker socket reader")),
router::Message::Broker,
));
join_set.spawn(broker_socket_write_half.write_messages(
broker_receiver,
log.new(o!("task" => "broker socket writer")),
));
}
}
match cluster_communication_strategy {
ClusterCommunicationStrategy::Tcp(TcpClusterCommunicationStrategy { broker }) => {
artifact_pusher::tcp::start_task(
&mut join_set,
artifact_pusher_receiver,
broker,
artifact_upload_tracker.clone(),
log.new(o!("task" => "artifact pusher")),
);
}
ClusterCommunicationStrategy::GitHub(GitHubClusterCommunicationStrategy {
token,
url,
}) => {
let github_client = GitHubClient::new(token.as_str(), url.clone())?;
artifact_pusher::github::start_task(
github_client,
&mut join_set,
artifact_pusher_receiver,
artifact_upload_tracker.clone(),
);
}
}
} else {
standalone = true;
drop(artifact_pusher_receiver);
drop(broker_receiver);
}
router::start_task(
artifact_pusher_sender,
broker_sender,
&mut join_set,
local_worker_sender.clone(),
router_receiver,
standalone,
);
preparer::task::start(
accept_invalid_remote_container_tls_certs,
cache_dir.clone(),
container_image_depot_cache_dir,
image_download_tracker.clone(),
&mut join_set,
log.clone(),
MANIFEST_INLINE_LIMIT,
MAX_PENDING_LAYER_BUILDS.try_into().unwrap(),
project_dir,
preparer_receiver,
router_sender.clone(),
preparer_sender.clone(),
)?;
let router_sender_clone_1 = router_sender.clone();
let router_sender_clone_2 = router_sender.clone();
let tasks = local_worker::start_task(
move |digest| {
router_sender_clone_1.send(router::Message::LocalWorkerStartArtifactFetch(digest))
},
move |msg| router_sender_clone_2.send(router::Message::LocalWorker(msg)),
cache_dir.join(LOCAL_WORKER_DIR),
cache_size,
done,
inline_limit,
&log,
local_worker_receiver,
local_worker_sender,
slots,
join_set,
)?;
Ok(ClientState {
router_sender,
artifact_upload_tracker,
image_download_tracker,
log,
preparer_sender,
tasks,
})
}
#[allow(clippy::too_many_arguments)]
pub async fn start(
&self,
accept_invalid_remote_container_tls_certs: AcceptInvalidRemoteContainerTlsCerts,
cache_dir: RootBuf<CacheDir>,
cache_size: CacheSize,
cluster_communication_strategy: Option<ClusterCommunicationStrategy>,
container_image_depot_cache_dir: RootBuf<ContainerImageDepotDir>,
inline_limit: InlineLimit,
log: Logger,
project_dir: RootBuf<ProjectDir>,
slots: Slots,
) -> Result<()> {
let mut guard = self.0.write().await;
let done = guard.take_to_start()?;
*guard = ClientStateWrapper::Running(
Self::try_to_start(
accept_invalid_remote_container_tls_certs,
cache_dir,
cache_size,
cluster_communication_strategy,
container_image_depot_cache_dir,
done,
inline_limit,
log.clone(),
project_dir,
slots,
)
.await?,
);
debug!(log, "client started successfully");
Ok(())
}
pub async fn stop(&self) -> Result<()> {
let mut guard = self.0.write().await;
let state = guard.take_to_stop()?;
let _ = state
.tasks
.shut_down(anyhow!("client process stopped by client"))
.await;
Ok(())
}
pub async fn run_job(
&self,
spec: spec::JobSpec,
) -> Result<futures::channel::mpsc::UnboundedReceiver<JobStatus>> {
let guard = self.0.read().await;
let state = guard.get_running()?;
debug!(state.log, "run_job"; "spec" => ?spec);
let (sender, receiver) = oneshot::channel();
state
.preparer_sender
.send(preparer::Message::PrepareJob(sender, spec))?;
let spec = receiver.await?.map_err(Error::msg)?;
let (sender, receiver) = futures::channel::mpsc::unbounded();
state
.router_sender
.send(router::Message::RunJob(spec, sender))?;
Ok(receiver)
}
pub async fn add_container(&self, name: String, container: ContainerSpec) -> Result<()> {
let guard = self.0.read().await;
let state = guard.get_running()?;
debug!(state.log, "add_container"; "name" => ?name, "container" => ?container);
let (sender, receiver) = oneshot::channel();
state
.preparer_sender
.send(preparer::Message::AddContainer(sender, name, container))?;
if let Some(existing) = receiver.await? {
debug!(state.log, "add_container replacing existing"; "existing" => ?existing);
}
Ok(())
}
pub async fn introspect(&self) -> Result<IntrospectResponse> {
let guard = self.0.read().await;
let state = guard.get_running()?;
let artifact_uploads = state.artifact_upload_tracker.get_remote_progresses();
let image_downloads = state.image_download_tracker.get_remote_progresses();
Ok(IntrospectResponse {
artifact_uploads,
image_downloads,
})
}
}