use std::process::Command;
use std::sync::Arc;
use bytes::BytesMut;
use futures_util::{SinkExt, StreamExt};
use tokio_util::codec::{Decoder, Encoder, Framed, FramedParts};
use crate::broker::protocol_v2::{session_frame, SessionExit, SessionFrame, SessionStart};
use crate::broker::session_codec::{encode_session_frame, try_decode_session_frame};
use crate::broker::session_pump::{run_child_session, FrameSink};
use crate::broker::session_server::spawn_contained_session_with_environment;
use crate::containment::ContainedProcessGroup;
impl FrameSink for tokio::sync::mpsc::Sender<SessionFrame> {
fn send(&self, frame: SessionFrame) -> Result<(), SessionFrame> {
self.blocking_send(frame).map_err(|e| e.0)
}
}
#[derive(Default)]
pub struct SessionFrameCodec {
seq: u64,
}
impl Decoder for SessionFrameCodec {
type Item = SessionFrame;
type Error = std::io::Error;
fn decode(&mut self, buf: &mut BytesMut) -> Result<Option<SessionFrame>, std::io::Error> {
match try_decode_session_frame(buf) {
Ok(Some(decoded)) => {
let _ = buf.split_to(decoded.consumed);
Ok(Some(decoded.frame))
}
Ok(None) => Ok(None),
Err(err) => Err(std::io::Error::new(std::io::ErrorKind::InvalidData, err)),
}
}
}
impl Encoder<SessionFrame> for SessionFrameCodec {
type Error = std::io::Error;
fn encode(&mut self, frame: SessionFrame, buf: &mut BytesMut) -> Result<(), std::io::Error> {
let wire = encode_session_frame(&frame, self.seq)
.map_err(|err| std::io::Error::new(std::io::ErrorKind::InvalidData, err))?;
self.seq = self.seq.wrapping_add(1);
buf.extend_from_slice(&wire);
Ok(())
}
}
pub fn session_framed<T>(io: T) -> Framed<T, SessionFrameCodec>
where
T: tokio::io::AsyncRead + tokio::io::AsyncWrite,
{
Framed::new(io, SessionFrameCodec::default())
}
const OUTBOUND_FRAME_CAPACITY: usize = 64;
fn command_from_start(
start: &SessionStart,
) -> std::io::Result<(Command, crate::EnvironmentPolicy)> {
let mut command = Command::new(&start.program);
command.args(&start.args);
if !start.cwd.is_empty() {
command.current_dir(&start.cwd);
}
for entry in &start.env {
command.env(&entry.key, &entry.value);
}
let policy =
crate::EnvironmentPolicy::from_wire(start.environment_policy, start.clear_inherited_env)
.map_err(|message| std::io::Error::new(std::io::ErrorKind::InvalidInput, message))?;
Ok((command, policy))
}
pub async fn session_takeover_from_buffered<T>(
io: T,
prebuffered: BytesMut,
group: Arc<ContainedProcessGroup>,
) -> std::io::Result<SessionExit>
where
T: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let mut parts = FramedParts::new(io, SessionFrameCodec::default());
parts.read_buf = prebuffered;
serve_session(Framed::from_parts(parts), group).await
}
pub async fn serve_session<T>(
mut framed: Framed<T, SessionFrameCodec>,
group: Arc<ContainedProcessGroup>,
) -> std::io::Result<SessionExit>
where
T: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let first = framed
.next()
.await
.ok_or_else(|| std::io::Error::other("session closed before SessionStart"))?
.map_err(|e| std::io::Error::other(format!("session recv failed before start: {e}")))?;
let start = match first.kind {
Some(session_frame::Kind::Start(start)) => start,
other => {
return Err(std::io::Error::other(format!(
"session must open with SessionStart, got {other:?}"
)))
}
};
let (command, policy) = command_from_start(&start)?;
run_compile_session_with_environment(framed, command, group, policy).await
}
pub async fn run_compile_session<T>(
framed: Framed<T, SessionFrameCodec>,
command: Command,
group: Arc<ContainedProcessGroup>,
) -> std::io::Result<SessionExit>
where
T: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
run_compile_session_with_environment(framed, command, group, crate::EnvironmentPolicy::Inherit)
.await
}
async fn run_compile_session_with_environment<T>(
mut framed: Framed<T, SessionFrameCodec>,
mut command: Command,
group: Arc<ContainedProcessGroup>,
environment_policy: crate::EnvironmentPolicy,
) -> std::io::Result<SessionExit>
where
T: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let child = spawn_contained_session_with_environment(&group, &mut command, environment_policy)?;
let (out_tx, mut out_rx) = tokio::sync::mpsc::channel::<SessionFrame>(OUTBOUND_FRAME_CAPACITY);
let (stdin_tx, stdin_rx) = std::sync::mpsc::channel::<SessionFrame>();
let pump = tokio::task::spawn_blocking(move || run_child_session(child, out_tx, stdin_rx));
let mut stdin_tx = Some(stdin_tx);
loop {
tokio::select! {
outbound = out_rx.recv() => {
match outbound {
Some(frame) => {
let terminal = matches!(frame.kind, Some(session_frame::Kind::Exit(_)));
framed
.send(frame)
.await
.map_err(|e| std::io::Error::other(format!("session send failed: {e}")))?;
if terminal {
break;
}
}
None => break,
}
}
inbound = framed.next() => {
match inbound {
Some(Ok(frame)) => {
let is_eof =
matches!(frame.kind, Some(session_frame::Kind::StdinEof(_)));
if let Some(tx) = stdin_tx.as_ref() {
if tx.send(frame).is_err() {
stdin_tx = None;
}
}
if is_eof {
stdin_tx = None;
}
}
Some(Err(e)) => {
return Err(std::io::Error::other(format!("session recv failed: {e}")))
}
None => {
stdin_tx = None;
}
}
}
}
}
drop(stdin_tx);
match pump.await {
Ok(result) => result,
Err(join_err) => Err(std::io::Error::other(format!(
"compile-session pump panicked: {join_err}"
))),
}
}