use async_trait::async_trait;
use bytes::Bytes;
use fern_protocol_postgresql::codec::backend;
use fern_proxy_interfaces::{SQLMessage, SQLMessageHandler};
use crate::strategies::MaskingStrategy;
pub use fern_proxy_interfaces::SQLHandlerConfig;
#[derive(Debug)]
pub struct DataMaskingHandler {
state: QueryState,
strategy: Box<dyn MaskingStrategy>,
columns_excluded: Vec<Bytes>,
columns_forced: Vec<Bytes>,
}
#[derive(Debug)]
enum QueryState {
Description,
Data(Vec<usize>),
}
#[async_trait]
impl SQLMessageHandler<backend::Message> for DataMaskingHandler {
fn new(config: &SQLHandlerConfig) -> Self {
let strategy: Box<dyn MaskingStrategy> =
if let Ok(strategy) = config.get::<String>("masking.strategy") {
match strategy.as_str() {
"caviar" => Box::new(strategies::CaviarMask::new(6)),
"caviar-preserve-shape" => Box::new(strategies::CaviarShapeMask::new()),
_ => Box::new(strategies::CaviarMask::new(6)),
}
} else {
Box::new(strategies::CaviarMask::new(6))
};
let mut columns_excluded = vec![];
if let Ok(columns) = config.get::<Vec<String>>("masking.exclude.columns") {
for column_name in columns.iter() {
columns_excluded.push(Bytes::from(column_name.clone()));
}
}
let mut columns_forced = vec![];
if let Ok(columns) = config.get::<Vec<String>>("masking.force.columns") {
for column_name in columns.iter() {
columns_forced.push(Bytes::from(column_name.clone()));
}
}
Self {
state: QueryState::Description,
strategy,
columns_excluded,
columns_forced,
}
}
async fn process(&mut self, msg: backend::Message) -> backend::Message {
match msg {
backend::Message::RowDescription(descriptions) => {
let mut no_mask = vec![];
if self.columns_excluded.len() == 1 && self.columns_excluded[0] == "*" {
for (idx, description) in descriptions.iter().enumerate() {
if !self.columns_forced.contains(&description.name) {
no_mask.push(idx);
}
}
} else {
for (idx, description) in descriptions.iter().enumerate() {
if self.columns_excluded.contains(&description.name)
&& !self.columns_forced.contains(&description.name)
{
no_mask.push(idx);
}
}
}
self.state = QueryState::Data(no_mask);
log::debug!("new masking exclusion state: {:?}", self.state);
backend::Message::RowDescription(descriptions)
}
backend::Message::CommandComplete(command) => {
self.state = QueryState::Description;
log::debug!("resetting masking state, awaiting next query");
backend::Message::CommandComplete(command)
}
backend::Message::DataRow(fields) => {
log::trace!("processing fields: {:?}", fields);
let mask = if let QueryState::Data(mask) = &self.state {
mask
} else {
panic!("unexpected state for `QueryState`");
};
let mut replaced_fields = vec![];
for (idx, field) in fields.iter().enumerate() {
if !mask.contains(&idx) {
log::debug!("applying masking to field #{}", idx);
let rewritten = self.strategy.mask(field);
replaced_fields.push(rewritten);
} else {
replaced_fields.push(field.clone());
}
}
backend::Message::DataRow(replaced_fields)
}
_ => msg,
}
}
}
#[derive(Debug)]
pub struct PassthroughHandler<M> {
_phantom: std::marker::PhantomData<M>,
}
#[async_trait]
impl<M> SQLMessageHandler<M> for PassthroughHandler<M>
where
M: SQLMessage,
{
fn new(_config: &SQLHandlerConfig) -> Self {
Self {
_phantom: std::marker::PhantomData,
}
}
}
mod strategies;
#[cfg(test)]
mod tests {
#[test]
fn it_works() {}
}