use std::future::Future;
#[must_use = "dropping an intermediary abandons both PostgreSQL sessions"]
#[derive(Debug)]
pub struct Intermediary<Downstream, Upstream> {
downstream: Downstream,
upstream: Upstream,
}
impl<Downstream, Upstream> Intermediary<Downstream, Upstream> {
pub const fn new(downstream: Downstream, upstream: Upstream) -> Self {
Self {
downstream,
upstream,
}
}
pub const fn downstream(&self) -> &Downstream {
&self.downstream
}
pub const fn upstream(&self) -> &Upstream {
&self.upstream
}
pub const fn sides_mut(&mut self) -> (&mut Downstream, &mut Upstream) {
(&mut self.downstream, &mut self.upstream)
}
pub fn into_parts(self) -> (Downstream, Upstream) {
(self.downstream, self.upstream)
}
pub fn transition_downstream<Next, Output, Error>(
self,
transition: impl FnOnce(Downstream) -> Result<(Next, Output), (Downstream, Error)>,
) -> Result<(Intermediary<Next, Upstream>, Output), (Self, Error)> {
let Self {
downstream,
upstream,
} = self;
match transition(downstream) {
Ok((downstream, output)) => Ok((
Intermediary {
downstream,
upstream,
},
output,
)),
Err((downstream, error)) => Err((
Intermediary {
downstream,
upstream,
},
error,
)),
}
}
pub fn transition_upstream<Next, Output, Error>(
self,
transition: impl FnOnce(Upstream) -> Result<(Next, Output), (Upstream, Error)>,
) -> Result<(Intermediary<Downstream, Next>, Output), (Self, Error)> {
let Self {
downstream,
upstream,
} = self;
match transition(upstream) {
Ok((upstream, output)) => Ok((
Intermediary {
downstream,
upstream,
},
output,
)),
Err((upstream, error)) => Err((
Intermediary {
downstream,
upstream,
},
error,
)),
}
}
pub fn inspect<Message, Output, Error>(
&mut self,
message: Message,
inspect: impl FnOnce(&mut Downstream, &mut Upstream, Message) -> Result<Output, Error>,
) -> Result<Output, Error> {
inspect(&mut self.downstream, &mut self.upstream, message)
}
pub async fn inspect_async<Message, Output, Error, Work>(
&mut self,
message: Message,
inspect: impl FnOnce(&mut Downstream, &mut Upstream, Message) -> Work,
) -> Result<Output, Error>
where
Work: Future<Output = Result<Output, Error>>,
{
inspect(&mut self.downstream, &mut self.upstream, message).await
}
}
#[cfg(test)]
mod tests {
use bytes::Bytes;
use super::Intermediary;
use crate::{
Conn, Pristine,
auth::{Auth, AuthOffer, SaslInitial, TlsServerEndPoint},
codec::Authentication,
grammar::{backend, frontend},
server_auth::{ServerAuth, ServerPassword},
};
#[derive(Debug)]
struct ClientFacingTls;
#[derive(Debug)]
struct UpstreamTls;
impl TlsServerEndPoint for ClientFacingTls {
fn tls_server_end_point(&self) -> &[u8] {
b"client-facing-certificate"
}
}
impl TlsServerEndPoint for UpstreamTls {
fn tls_server_end_point(&self) -> &[u8] {
b"upstream-certificate"
}
}
#[test]
fn each_side_transitions_without_coupling_the_other() {
#[derive(Debug)]
struct Clean;
let downstream: backend::TypedSession<(), backend::Ready, Clean> =
backend::TypedSession::with_transport(());
let upstream: frontend::TypedSession<(), frontend::Ready, Clean> =
frontend::TypedSession::with_transport(());
let intermediary = Intermediary::new(downstream, upstream);
let (intermediary, downstream_query) = intermediary
.transition_downstream(|session| {
session.query(Bytes::from_static(b"select 1"), |(), query| {
Ok::<_, &'static str>(query)
})
})
.expect("downstream inspection succeeds");
assert_eq!(downstream_query, Bytes::from_static(b"select 1"));
let (intermediary, upstream_query) = intermediary
.transition_upstream(|session| {
session.query(Bytes::from_static(b"select 2"), |(), query| {
Ok::<_, &'static str>(query)
})
})
.expect("upstream inspection succeeds");
assert_eq!(upstream_query, Bytes::from_static(b"select 2"));
let (_downstream, _upstream): (
backend::TypedSession<(), backend::Simple, backend::Dirty>,
frontend::TypedSession<(), frontend::Simple, frontend::Dirty>,
) = intermediary.into_parts();
}
#[test]
fn tls_and_authentication_mechanisms_remain_asymmetric() {
let downstream: Conn<ClientFacingTls, ServerAuth, Pristine> =
Conn::new(ClientFacingTls).transition();
let upstream: Conn<UpstreamTls, Auth, Pristine> = Conn::new(UpstreamTls).transition();
let (downstream, cleartext_request) = downstream.request_cleartext().unwrap();
let AuthOffer::Sasl {
conn: upstream,
mechanisms,
} = upstream
.offer(Authentication::Sasl {
mechanisms: vec![Bytes::from_static(b"SCRAM-SHA-256-PLUS")],
})
.unwrap()
else {
panic!("upstream did not independently select SASL")
};
assert_eq!(cleartext_request.tag, b'R');
assert_eq!(mechanisms, [Bytes::from_static(b"SCRAM-SHA-256-PLUS")]);
let intermediary: Intermediary<
Conn<ClientFacingTls, ServerPassword, Pristine>,
Conn<UpstreamTls, SaslInitial, Pristine>,
> = Intermediary::new(downstream, upstream);
assert_eq!(
intermediary.downstream().tls_server_end_point(),
b"client-facing-certificate"
);
assert_eq!(
intermediary.upstream().tls_server_end_point(),
b"upstream-certificate"
);
let (downstream, upstream) = intermediary.into_parts();
let _downstream_transport = downstream.into_transport();
let _upstream_transport = upstream.into_transport();
}
#[test]
fn rejected_transition_reconstructs_both_original_sides() {
let intermediary = Intermediary::new(vec![1_u8], vec![2_u8]);
let (intermediary, error) = intermediary
.transition_downstream(|downstream| Err::<(Vec<u8>, ()), _>((downstream, "reject")))
.unwrap_err();
assert_eq!(error, "reject");
assert_eq!(intermediary.into_parts(), (vec![1], vec![2]));
}
#[tokio::test]
async fn asynchronous_policy_can_modify_or_replace_a_typed_message() {
let mut intermediary = Intermediary::new(Vec::<Bytes>::new(), Vec::<Bytes>::new());
let rewritten = intermediary
.inspect_async(
Bytes::from_static(b"select secret"),
|downstream, upstream, query| {
downstream.push(query);
upstream.push(Bytes::from_static(b"select public"));
std::future::ready(Ok::<_, std::convert::Infallible>(Bytes::from_static(
b"select public",
)))
},
)
.await
.unwrap();
assert_eq!(rewritten, Bytes::from_static(b"select public"));
assert_eq!(intermediary.downstream().len(), 1);
assert_eq!(intermediary.upstream().len(), 1);
}
}