use core::future::Future;
use mkit_rpc::mkit::rpc::v1::ssh::{HelloResponse, SshFrame, ssh_frame};
use mkit_rpc::mkit::rpc::v1::{ErrorCode, ProtocolVersion};
use super::budget::Budget;
use super::verbs;
use crate::error::Redacted;
use crate::pipeline::{AuthMode, HookSet, Pipeline};
use crate::principal::Principal;
use crate::rt::MaybeSend;
use crate::store::{MultipartBlobStore, NamespaceStore};
#[derive(Debug)]
#[non_exhaustive]
pub enum FrameIoError {
Eof,
Timeout,
Malformed,
Io(Redacted),
}
impl From<mkit_rpc::FrameError> for FrameIoError {
fn from(err: mkit_rpc::FrameError) -> Self {
match err {
mkit_rpc::FrameError::LengthTruncated => Self::Eof,
mkit_rpc::FrameError::Io(e) => Self::Io(Redacted::new(e.to_string())),
_ => Self::Malformed,
}
}
}
pub trait FrameSource: MaybeSend {
fn next_frame(&mut self) -> impl Future<Output = Result<SshFrame, FrameIoError>> + MaybeSend;
}
pub trait FrameSink: MaybeSend {
fn send(
&mut self,
frame: &SshFrame,
) -> impl Future<Output = Result<(), FrameIoError>> + MaybeSend;
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct SessionConfig {
pub server_id: String,
pub stop_after_hello: bool,
pub repository: Option<String>,
}
impl SessionConfig {
#[must_use]
pub fn new(server_id: impl Into<String>) -> Self {
Self {
server_id: server_id.into(),
stop_after_hello: false,
repository: None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SessionEnd {
Clean,
ProtocolError,
Timeout,
IoError,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum Stop {
Io,
Timeout,
}
impl From<FrameIoError> for Stop {
fn from(_: FrameIoError) -> Self {
Self::Io
}
}
pub(super) async fn emit_error<K: FrameSink>(
sink: &mut K,
code: ErrorCode,
message: &str,
) -> Result<(), Stop> {
let frame = mkit_rpc::ssh_error_frame(code, message);
sink.send(&frame).await.map_err(Stop::from)
}
pub(super) async fn send_body<K: FrameSink>(
sink: &mut K,
body: ssh_frame::Body,
) -> Result<(), Stop> {
let frame = SshFrame {
body: Some(body),
..Default::default()
};
sink.send(&frame).await.map_err(Stop::from)
}
pub async fn handshake<S: FrameSource, K: FrameSink>(
src: &mut S,
sink: &mut K,
server_id: &str,
) -> Result<(), SessionEnd> {
let frame = match src.next_frame().await {
Ok(frame) => frame,
Err(FrameIoError::Timeout) => return Err(SessionEnd::Timeout),
Err(_) => return Err(SessionEnd::ProtocolError),
};
let Some(ssh_frame::Body::Hello(hello)) = frame.body else {
let _ = emit_error(sink, ErrorCode::InvalidRequest, "first frame must be Hello").await;
return Err(SessionEnd::ProtocolError);
};
let proto = hello.proto.unwrap_or_default();
if proto != ProtocolVersion::ProtocolVersion1 {
let message = format!("unsupported proto_version {}", proto.to_i32());
let _ = emit_error(sink, ErrorCode::InvalidRequest, &message).await;
return Err(SessionEnd::ProtocolError);
}
let resp = ssh_frame::Body::HelloResponse(Box::new(HelloResponse {
proto: Some(ProtocolVersion::ProtocolVersion1.into()),
server_id: Some(server_id.to_owned()),
..Default::default()
}));
send_body(sink, resp)
.await
.map_err(|_| SessionEnd::ProtocolError)
}
pub async fn serve_session<B, N, H, S, K>(
pipeline: &Pipeline<B, N, H>,
principal: Principal,
src: &mut S,
sink: &mut K,
cfg: &SessionConfig,
) -> SessionEnd
where
B: MultipartBlobStore,
N: NamespaceStore,
H: HookSet,
S: FrameSource,
K: FrameSink,
{
if !matches!(pipeline.auth_mode(), AuthMode::TransportIdentity) {
tracing::error!("ssh session refused: the pipeline is not in TransportIdentity mode");
return SessionEnd::ProtocolError;
}
if let Err(end) = handshake(src, sink, &cfg.server_id).await {
return end;
}
if cfg.stop_after_hello {
return SessionEnd::Clean;
}
let mut verbs = verbs::Verbs::new(pipeline, principal, cfg.repository.clone());
let mut budget = Budget::default();
loop {
let frame = match src.next_frame().await {
Ok(frame) => frame,
Err(FrameIoError::Eof) => return SessionEnd::Clean,
Err(FrameIoError::Timeout) => return SessionEnd::Timeout,
Err(_) => {
let _ = emit_error(sink, ErrorCode::InvalidRequest, "frame parse error").await;
return SessionEnd::ProtocolError;
}
};
if let Err(message) = budget.charge(&frame) {
let _ = emit_error(sink, ErrorCode::InvalidRequest, message).await;
return SessionEnd::ProtocolError;
}
let body = match frame.body {
Some(ssh_frame::Body::Close(_)) => return SessionEnd::Clean,
body => body,
};
match verbs.dispatch(body, src, sink).await {
Ok(()) => {}
Err(Stop::Io) => return SessionEnd::IoError,
Err(Stop::Timeout) => return SessionEnd::Timeout,
}
}
}