1use core::future::Future;
5
6use mkit_rpc::mkit::rpc::v1::ssh::{HelloResponse, SshFrame, ssh_frame};
7use mkit_rpc::mkit::rpc::v1::{ErrorCode, ProtocolVersion};
8
9use super::budget::Budget;
10use super::verbs;
11use crate::error::Redacted;
12use crate::pipeline::{AuthMode, HookSet, Pipeline};
13use crate::principal::Principal;
14use crate::rt::MaybeSend;
15use crate::store::{MultipartBlobStore, NamespaceStore};
16
17#[derive(Debug)]
19#[non_exhaustive]
20pub enum FrameIoError {
21 Eof,
23 Timeout,
25 Malformed,
27 Io(Redacted),
29}
30
31impl From<mkit_rpc::FrameError> for FrameIoError {
32 fn from(err: mkit_rpc::FrameError) -> Self {
36 match err {
37 mkit_rpc::FrameError::LengthTruncated => Self::Eof,
38 mkit_rpc::FrameError::Io(e) => Self::Io(Redacted::new(e.to_string())),
39 _ => Self::Malformed,
40 }
41 }
42}
43
44pub trait FrameSource: MaybeSend {
46 fn next_frame(&mut self) -> impl Future<Output = Result<SshFrame, FrameIoError>> + MaybeSend;
49}
50
51pub trait FrameSink: MaybeSend {
53 fn send(
55 &mut self,
56 frame: &SshFrame,
57 ) -> impl Future<Output = Result<(), FrameIoError>> + MaybeSend;
58}
59
60#[derive(Debug, Clone, PartialEq, Eq)]
62#[non_exhaustive]
63pub struct SessionConfig {
64 pub server_id: String,
67 pub stop_after_hello: bool,
71 pub repository: Option<String>,
76}
77
78impl SessionConfig {
79 #[must_use]
81 pub fn new(server_id: impl Into<String>) -> Self {
82 Self {
83 server_id: server_id.into(),
84 stop_after_hello: false,
85 repository: None,
86 }
87 }
88}
89
90#[derive(Debug, Clone, Copy, PartialEq, Eq)]
93pub enum SessionEnd {
94 Clean,
97 ProtocolError,
100 Timeout,
102 IoError,
104}
105
106#[derive(Debug, Clone, Copy, PartialEq, Eq)]
108pub(super) enum Stop {
109 Io,
111 Timeout,
113}
114
115impl From<FrameIoError> for Stop {
116 fn from(_: FrameIoError) -> Self {
117 Self::Io
118 }
119}
120
121pub(super) async fn emit_error<K: FrameSink>(
123 sink: &mut K,
124 code: ErrorCode,
125 message: &str,
126) -> Result<(), Stop> {
127 let frame = mkit_rpc::ssh_error_frame(code, message);
128 sink.send(&frame).await.map_err(Stop::from)
129}
130
131pub(super) async fn send_body<K: FrameSink>(
133 sink: &mut K,
134 body: ssh_frame::Body,
135) -> Result<(), Stop> {
136 let frame = SshFrame {
137 body: Some(body),
138 ..Default::default()
139 };
140 sink.send(&frame).await.map_err(Stop::from)
141}
142
143pub async fn handshake<S: FrameSource, K: FrameSink>(
152 src: &mut S,
153 sink: &mut K,
154 server_id: &str,
155) -> Result<(), SessionEnd> {
156 let frame = match src.next_frame().await {
157 Ok(frame) => frame,
158 Err(FrameIoError::Timeout) => return Err(SessionEnd::Timeout),
159 Err(_) => return Err(SessionEnd::ProtocolError),
160 };
161 let Some(ssh_frame::Body::Hello(hello)) = frame.body else {
162 let _ = emit_error(sink, ErrorCode::InvalidRequest, "first frame must be Hello").await;
163 return Err(SessionEnd::ProtocolError);
164 };
165 let proto = hello.proto.unwrap_or_default();
166 if proto != ProtocolVersion::ProtocolVersion1 {
167 let message = format!("unsupported proto_version {}", proto.to_i32());
168 let _ = emit_error(sink, ErrorCode::InvalidRequest, &message).await;
169 return Err(SessionEnd::ProtocolError);
170 }
171 let resp = ssh_frame::Body::HelloResponse(Box::new(HelloResponse {
172 proto: Some(ProtocolVersion::ProtocolVersion1.into()),
173 server_id: Some(server_id.to_owned()),
174 ..Default::default()
175 }));
176 send_body(sink, resp)
177 .await
178 .map_err(|_| SessionEnd::ProtocolError)
179}
180
181pub async fn serve_session<B, N, H, S, K>(
194 pipeline: &Pipeline<B, N, H>,
195 principal: Principal,
196 src: &mut S,
197 sink: &mut K,
198 cfg: &SessionConfig,
199) -> SessionEnd
200where
201 B: MultipartBlobStore,
202 N: NamespaceStore,
203 H: HookSet,
204 S: FrameSource,
205 K: FrameSink,
206{
207 if !matches!(pipeline.auth_mode(), AuthMode::TransportIdentity) {
208 tracing::error!("ssh session refused: the pipeline is not in TransportIdentity mode");
209 return SessionEnd::ProtocolError;
210 }
211 if let Err(end) = handshake(src, sink, &cfg.server_id).await {
212 return end;
213 }
214 if cfg.stop_after_hello {
215 return SessionEnd::Clean;
216 }
217 let mut verbs = verbs::Verbs::new(pipeline, principal, cfg.repository.clone());
218 let mut budget = Budget::default();
219 loop {
220 let frame = match src.next_frame().await {
221 Ok(frame) => frame,
222 Err(FrameIoError::Eof) => return SessionEnd::Clean,
223 Err(FrameIoError::Timeout) => return SessionEnd::Timeout,
224 Err(_) => {
225 let _ = emit_error(sink, ErrorCode::InvalidRequest, "frame parse error").await;
226 return SessionEnd::ProtocolError;
227 }
228 };
229 if let Err(message) = budget.charge(&frame) {
230 let _ = emit_error(sink, ErrorCode::InvalidRequest, message).await;
231 return SessionEnd::ProtocolError;
232 }
233 let body = match frame.body {
234 Some(ssh_frame::Body::Close(_)) => return SessionEnd::Clean,
235 body => body,
236 };
237 match verbs.dispatch(body, src, sink).await {
238 Ok(()) => {}
239 Err(Stop::Io) => return SessionEnd::IoError,
240 Err(Stop::Timeout) => return SessionEnd::Timeout,
241 }
242 }
243}