feat: Initial implementation (bind/search/whoami)

This commit is contained in:
selfhoster selfhoster 2026-08-20 21:05:09 +02:00
commit d42508b15e
14 changed files with 2325 additions and 0 deletions

11
src/cli.rs Normal file
View file

@ -0,0 +1,11 @@
use clap::Parser;
use std::path::PathBuf;
#[derive(Parser)]
#[command(version, about, long_about = None)]
pub struct Cli {
/// Sets a custom config file
#[arg(short, long, value_name = "FILE")]
pub config: PathBuf,
}

176
src/client.rs Normal file
View file

@ -0,0 +1,176 @@
use futures_util::sink::SinkExt;
use futures_util::stream::StreamExt;
use ldap3_proto::LdapCodec;
use ldap3_proto::control::LdapControl;
use ldap3_proto::proto::*;
use tokio::net::TcpStream;
use tokio::time::timeout;
use tokio_util::codec::{FramedRead, FramedWrite};
use std::time::Duration;
use crate::{CR, CW, LdapError};
pub struct BasicLdapClient {
r: FramedRead<CR, LdapCodec>,
w: FramedWrite<CW, LdapCodec>,
msg_counter: i32,
}
impl BasicLdapClient {
fn next_msgid(&mut self) -> i32 {
self.msg_counter += 1;
self.msg_counter
}
pub async fn build(addr: &str) -> Result<Self, LdapError> {
let tcpstream = match timeout(Duration::from_secs(1), TcpStream::connect(addr)).await {
Ok(Ok(t)) => {
trace!("connection established to {addr}");
t
}
Ok(Err(err)) => {
// trace!(?addr, ?err, "error");
error!("error to {addr}: {err}");
panic!();
}
Err(_) => {
warn!("timeout to {addr}");
panic!();
// continue;
}
};
let (r, w) = tokio::io::split(tcpstream);
let w = FramedWrite::new(w, LdapCodec::new(None, None));
let r = FramedRead::new(r, LdapCodec::new(None, None));
Ok(Self {
r,
w,
msg_counter: 0,
})
}
pub async fn bind(
&mut self,
lbr: LdapBindRequest,
ctrl: Vec<LdapControl>,
) -> Result<(LdapBindResponse, Vec<LdapControl>), LdapError> {
let ck_msgid = self.next_msgid();
let msg = LdapMsg {
msgid: ck_msgid,
op: LdapOp::BindRequest(lbr),
ctrl,
};
match self.w.send(msg).await {
Ok(_) => {}
Err(err) => {
error!("unable to transmit to ldap server: {err}");
return Err(LdapError::Transport);
}
};
match self.r.next().await {
Some(Ok(LdapMsg {
msgid,
op: LdapOp::BindResponse(bind_resp),
ctrl,
})) => {
if msgid == ck_msgid {
Ok((bind_resp, ctrl))
} else {
error!("invalid msgid, sequence error.");
Err(LdapError::InvalidProtocolState)
}
}
Some(Ok(msg)) => {
trace!("{:?}", msg);
Err(LdapError::InvalidProtocolState)
}
Some(Err(e)) => {
error!("unable to receive from ldap server: {e}");
Err(LdapError::Transport)
}
None => {
error!("connection closed");
Err(LdapError::Transport)
}
}
}
pub async fn search(
&mut self,
sr: LdapSearchRequest,
ctrl: Vec<LdapControl>,
) -> Result<
(
Vec<(LdapSearchResultEntry, Vec<LdapControl>)>,
LdapResult,
Vec<LdapControl>,
),
LdapError,
> {
let ck_msgid = self.next_msgid();
let msg = LdapMsg {
msgid: ck_msgid,
op: LdapOp::SearchRequest(sr),
ctrl,
};
match self.w.send(msg).await {
Ok(_) => {}
Err(err) => {
error!("unable to transmit to ldap server: {err}");
return Err(LdapError::Transport);
}
};
let mut entries = Vec::new();
loop {
match self.r.next().await {
// This terminates the iteration of entries.
Some(Ok(LdapMsg {
msgid,
op: LdapOp::SearchResultDone(search_res),
ctrl,
})) => {
if msgid == ck_msgid {
break Ok((entries, search_res, ctrl));
} else {
error!("invalid msgid, sequence error.");
break Err(LdapError::InvalidProtocolState);
}
}
Some(Ok(LdapMsg {
msgid,
op: LdapOp::SearchResultEntry(search_entry),
ctrl,
})) => {
if msgid == ck_msgid {
entries.push((search_entry, ctrl))
} else {
error!("invalid msgid, sequence error.");
break Err(LdapError::InvalidProtocolState);
}
}
Some(Ok(msg)) => {
trace!("{:?}", msg);
break Err(LdapError::InvalidProtocolState);
}
Some(Err(e)) => {
error!("unable to receive from ldap server: {e}");
break Err(LdapError::Transport);
}
None => {
error!("connection closed");
break Err(LdapError::Transport);
}
}
}
}
}

22
src/config.rs Normal file
View file

@ -0,0 +1,22 @@
use serde::Deserialize;
use std::path::Path;
#[derive(Clone, Debug, Deserialize)]
pub struct Config {
pub mapping: Vec<Mapping>,
}
impl Config {
pub async fn from_path(path: &Path) -> anyhow::Result<Self> {
let content = tokio::fs::read(path).await?;
Ok(toml::from_slice(&content)?)
}
}
#[derive(Clone, Debug, Deserialize)]
pub struct Mapping {
pub from: String,
pub to: String,
pub backend: String,
}

100
src/dn.rs Normal file
View file

@ -0,0 +1,100 @@
use indexmap::IndexMap;
use ldap3::dn_escape;
use crate::LdapError;
/// A Dn is a key-value mapping which can contain the same key several times.
///
/// This implementation is not fully RFC compliant and will only parse simple DNs for
/// basic attribute manipulation.
///
/// Keys are lowercased, but scrambled keys in a broken order will be reordered. For example,
/// `dc=example,ou=people,dc=com` will become `dc=example,dc=com,ou=people`.
pub struct Dn {
pub keys: IndexMap<String, Vec<String>>,
}
impl Dn {
/// Parse a DN string. Here parsing escaped characters is not critical, if a
/// client sends us funny characters, the domain name simply won't match
/// and their request won't go anywhere.
///
/// However, if we find a really funny request such as `dc=foo=bar`, then
/// we return an error to the client.
pub fn from_dn_str(input: &str) -> Result<Self, LdapError> {
let mut keys: IndexMap<String, Vec<String>> = IndexMap::new();
for query in input.split(',') {
let mut query_parts = query.split('=');
// Here we have key=val pairs
if query_parts.clone().count() != 2 {
// Bad request
log::debug!("Invalid query DN: {input}");
return Err(LdapError::InvalidQuery);
}
let key = query_parts.next().unwrap();
let val = query_parts.next().unwrap();
if let Some(previous) = keys.get_mut(key) {
previous.push(val.to_string());
} else {
keys.insert(key.to_string(), vec![val.to_string()]);
}
}
Ok(Self { keys })
}
pub fn to_dn_string(&self) -> String {
let mut s = String::new();
let mut first = true;
for (key, values) in &self.keys {
for value in values {
if first {
first = false;
} else {
s.push(',');
}
s.push_str(key);
s.push('=');
s.push_str(value);
}
}
s
}
/// Gets the hostname defined in the `dc` fields of the DN.
///
/// For example, `dc=example,dc=com` becomes `Some(example.com)`.
///
/// The returned domain is not normalized and may require casing treatment
/// to compare meaningfully.
pub fn get_hostname(&self) -> Option<String> {
let domain_components = self.keys.get("dc")?;
// We don't populate the dc key if there was no value at all, so
// we have at least one component.
let mut domain_components = domain_components.iter();
let mut s = String::from(domain_components.next().unwrap());
for domain_component in domain_components {
s.push('.');
s.push_str(domain_component);
}
Some(s)
}
/// Overrides the DN hostname (`dc` fields) with the provided host.
///
/// When no `dc` fields are present, they are only added when `force` is true.
pub fn set_hostname(&mut self, host: &str, force: bool) {
let domain_components: Vec<String> =
host.split('.').map(|x| dn_escape(x).to_string()).collect();
if !force && !self.keys.contains_key("dc") {
return;
}
self.keys.insert("dc".to_string(), domain_components);
}
}

232
src/main.rs Normal file
View file

@ -0,0 +1,232 @@
#[macro_use]
extern crate log;
use clap::Parser;
use futures_util::StreamExt;
use ldap3_proto::LdapCodec;
use ldap3_proto::proto::*;
use tokio::io::{AsyncRead, AsyncWrite, ReadHalf, WriteHalf};
use tokio::net::{TcpListener, TcpStream};
use tokio::time::timeout;
use tokio_util::codec::{FramedRead, FramedWrite};
use std::sync::Arc;
use std::time::Duration;
mod cli;
use cli::Cli;
mod client;
use crate::client::BasicLdapClient;
mod config;
use config::Config;
mod dn;
mod op;
use crate::dn::Dn;
const LDAP_CLIENT_IO_TIMEOUT: Duration = Duration::from_secs(1);
type CR = ReadHalf<TcpStream>;
type CW = WriteHalf<TcpStream>;
#[derive(Debug, Clone)]
pub enum LdapError {
TlsError,
ConnectError,
Transport,
InvalidProtocolState,
InvalidQuery,
}
// We allow the large enum to exist as we always do a mem swap from unbound to authenticated, so
// the memory layout penalty doesn't apply.
#[allow(clippy::large_enum_variant)]
enum ClientState {
Unbound,
Authenticated {
#[allow(dead_code)]
request_dn: String,
backend_dn: String,
client: BasicLdapClient,
},
}
pub async fn client_process<W: AsyncWrite + Unpin, R: AsyncRead + Unpin>(
mut r: FramedRead<R, LdapCodec>,
mut w: FramedWrite<W, LdapCodec>,
config: Arc<Config>,
) {
info!("Received new connection");
// We always start unbound.
let mut state = ClientState::Unbound;
// Start to wait for incoming packets
while let Ok(Some(Ok(protomsg))) = timeout(LDAP_CLIENT_IO_TIMEOUT, r.next()).await {
let next_state = match (&mut state, protomsg) {
// Doesn't matter what state we are in, any bind will trigger this process.
(
_,
LdapMsg {
msgid,
op: LdapOp::BindRequest(lbr),
ctrl,
},
) => match op::bind::bind(&mut w, lbr, config.clone(), msgid, ctrl).await {
Ok(ns) => ns,
Err(_) => break,
},
// Unbinds are always actioned.
(
_,
LdapMsg {
msgid: _,
op: LdapOp::UnbindRequest,
ctrl: _,
},
) => {
break;
}
// Unbound handler
(
ClientState::Unbound,
LdapMsg {
msgid,
op: LdapOp::SearchRequest(sr),
ctrl,
},
) => {
// We have to trigger a bind first in case we have a mapping.
let lbr = LdapBindRequest {
dn: "".to_string(),
cred: LdapBindCred::Simple("".to_string()),
};
let mut next_state =
match op::bind::bind(&mut w, lbr, config.clone(), 0, Vec::default()).await {
Ok(ns) => ns,
Err(_) => break,
};
match &mut next_state {
Some(ClientState::Unbound) | None => {
error!("Invalid state, bind did not return an authenticated state!");
break;
}
Some(ClientState::Authenticated {
client,
request_dn: _,
backend_dn: _,
}) => {
let search_req = op::search::SearchRequest {
sr,
msgid,
ctrl,
client,
};
match op::search::search(&mut w, search_req).await {
Ok(()) => {}
Err(_) => break,
}
}
}
next_state
}
// Authenticated message handler.
// - Search
(
ClientState::Authenticated {
client,
backend_dn: _,
request_dn: _,
},
LdapMsg {
msgid,
op: LdapOp::SearchRequest(sr),
ctrl,
},
) => {
let search_req = op::search::SearchRequest {
sr,
msgid,
ctrl,
client,
};
match op::search::search(&mut w, search_req).await {
Ok(()) => None,
Err(_) => break,
}
}
// Extended Requests - Generally whoami.
(
ClientState::Authenticated {
request_dn: _,
backend_dn,
client: _,
},
LdapMsg {
msgid,
op: LdapOp::ExtendedRequest(ler),
ctrl: _,
},
) => match op::ext::extop(&mut w, ler, msgid, backend_dn).await {
Ok(ns) => ns,
Err(_) => break,
},
_ => {
log::debug!("unimplemented");
None
}
};
if let Some(next_state) = next_state {
// Update the client state, dropping any former state.
state = next_state;
}
}
}
#[tokio::main(flavor = "current_thread")]
async fn main() {
if let Err(_) = std::env::var("RUST_LOG") {
unsafe { std::env::set_var("RUST_LOG", "info"); }
}
pretty_env_logger::formatted_timed_builder()
.parse_default_env()
.init();
let cli = Cli::parse();
let config = match Config::from_path(&cli.config).await {
Ok(config) => config,
Err(e) => {
error!(
"Loading configuration file {} failed!",
cli.config.display()
);
error!("{}", e);
std::process::exit(1);
}
};
let state = Arc::new(config);
// Let the listening port ready.
let listener = TcpListener::bind("127.0.0.1:3389").await.unwrap();
info!("Listening on {:?}", listener);
loop {
match listener.accept().await {
Ok((tcpstream, client_socket_addr)) => {
log::debug!("New connection from {client_socket_addr}");
let (r, w) = tokio::io::split(tcpstream);
let r = FramedRead::new(r, LdapCodec::new(None, None));
let w = FramedWrite::new(w, LdapCodec::new(None, None));
tokio::spawn(client_process(r, w, state.clone()));
}
Err(_e) => continue,
}
}
}

120
src/op/bind.rs Normal file
View file

@ -0,0 +1,120 @@
use futures_util::SinkExt;
use ldap3_proto::LdapCodec;
use ldap3_proto::control::*;
use ldap3_proto::proto::*;
use tokio::io::AsyncWrite;
use tokio_util::codec::FramedWrite;
use std::sync::Arc;
use crate::{BasicLdapClient, ClientState, Config, Dn, LdapError};
pub fn bind_operror(msgid: i32, msg: &str) -> LdapMsg {
LdapMsg {
msgid,
op: LdapOp::BindResponse(LdapBindResponse {
res: LdapResult {
code: LdapResultCode::OperationsError,
matcheddn: "".to_string(),
message: msg.to_string(),
referral: vec![],
},
saslcreds: None,
}),
ctrl: vec![],
}
}
pub async fn bind<W: AsyncWrite + Unpin>(
w: &mut FramedWrite<W, LdapCodec>,
mut lbr: LdapBindRequest,
config: Arc<Config>,
msgid: i32,
ctrl: Vec<LdapControl>,
) -> Result<Option<ClientState>, LdapError> {
trace!("{:?}", lbr);
let request_dn = lbr.dn.clone();
debug!("Received bind request on DN: {}", lbr.dn);
let mut dn = Dn::from_dn_str(&lbr.dn)?;
let Some(requested_domain) = dn.get_hostname() else {
debug!("No domain name CN found in DN: {}", lbr.dn);
return Err(LdapError::InvalidQuery);
};
// Lowercase the domain systematically to allow matches
let requested_domain = requested_domain.to_lowercase();
let Some(mapping) = config
.mapping
.iter()
.find(|x| x.from.to_lowercase() == requested_domain)
else {
debug!("No mapping found for domain {requested_domain}");
return Err(LdapError::InvalidQuery);
};
debug!(
"Redirecting {} to {} with domain {}",
requested_domain, mapping.backend, mapping.to
);
dn.set_hostname(&mapping.to, false);
lbr.dn = dn.to_dn_string();
let dn = lbr.dn.clone();
// We need the client to connect *and* bind to proceed here!
let mut client = match BasicLdapClient::build(&mapping.backend).await {
Ok(c) => c,
Err(e) => {
error!("A client build error has occurred: {e:?}");
let resp_msg = bind_operror(msgid, "unable to bind");
w.send(resp_msg).await.map_err(|err| {
error!("Unable to send response: {err}");
LdapError::Transport
})?;
// Always bail.
return Ok(None);
}
};
let valid = match client.bind(lbr, ctrl).await {
Ok((bind_resp, ctrl)) => {
// Almost there, lets check the bind result.
let valid = bind_resp.res.code == LdapResultCode::Success;
let resp_msg = LdapMsg {
msgid,
op: LdapOp::BindResponse(bind_resp),
ctrl,
};
w.send(resp_msg).await.map_err(|err| {
error!("Unable to send response: {err}");
LdapError::Transport
})?;
valid
}
Err(e) => {
error!("A client bind error has occurred: {e:?}");
let resp_msg = bind_operror(msgid, "unable to bind");
w.send(resp_msg).await.map_err(|err| {
error!("Unable to send response: {err}");
LdapError::Transport
})?;
// Always bail.
return Ok(None);
}
};
if valid {
info!("Successful bind for {}", dn);
Ok(Some(ClientState::Authenticated {
request_dn,
backend_dn: dn,
client,
}))
} else {
Ok(None)
}
}

51
src/op/ext.rs Normal file
View file

@ -0,0 +1,51 @@
use futures_util::SinkExt;
use ldap3_proto::LdapCodec;
use ldap3_proto::proto::*;
use tokio::io::AsyncWrite;
use tokio_util::codec::FramedWrite;
use crate::{ClientState, LdapError};
pub async fn extop<W: AsyncWrite + Unpin>(
w: &mut FramedWrite<W, LdapCodec>,
ler: LdapExtendedRequest,
msgid: i32,
display_dn: &str,
) -> Result<Option<ClientState>, LdapError> {
let op = match ler.name.as_str() {
"1.3.6.1.4.1.4203.1.11.3" => LdapOp::ExtendedResponse(LdapExtendedResponse {
res: LdapResult {
code: LdapResultCode::Success,
matcheddn: "".to_string(),
message: "".to_string(),
referral: vec![],
},
name: None,
value: Some(Vec::from(display_dn)),
}),
_ => LdapOp::ExtendedResponse(LdapExtendedResponse {
res: LdapResult {
code: LdapResultCode::OperationsError,
matcheddn: "".to_string(),
message: "".to_string(),
referral: vec![],
},
name: None,
value: None,
}),
};
w.send(LdapMsg {
msgid,
op,
ctrl: vec![],
})
.await
.map_err(|err| {
error!("Unable to send response: {err}");
LdapError::Transport
})?;
Ok(None)
}

3
src/op/mod.rs Normal file
View file

@ -0,0 +1,3 @@
pub mod bind;
pub mod ext;
pub mod search;

70
src/op/search.rs Normal file
View file

@ -0,0 +1,70 @@
use futures_util::SinkExt;
use ldap3_proto::LdapCodec;
use ldap3_proto::control::*;
use ldap3_proto::proto::*;
use tokio::io::AsyncWrite;
use tokio_util::codec::FramedWrite;
use crate::op::bind::bind_operror;
use crate::{BasicLdapClient, LdapError};
pub struct SearchRequest<'a> {
pub sr: LdapSearchRequest,
pub msgid: i32,
pub ctrl: Vec<LdapControl>,
pub client: &'a mut BasicLdapClient,
}
pub async fn search<W: AsyncWrite + Unpin>(
w: &mut FramedWrite<W, LdapCodec>,
search_request: SearchRequest<'_>,
) -> Result<(), LdapError> {
let SearchRequest {
sr,
msgid,
ctrl,
client,
} = search_request;
let (entries, result, ctrl) = match client.search(sr, ctrl).await {
Ok(data) => data,
Err(e) => {
error!("A client search error has occurred: {e:?}");
let resp_msg = bind_operror(msgid, "unable to search");
w.send(resp_msg).await.map_err(|err| {
error!("Unable to send response: {err}");
LdapError::Transport
})?;
// Error sent, return with no state change.
return Ok(());
}
};
for (entry, ctrl) in entries {
w.send(LdapMsg {
msgid,
op: LdapOp::SearchResultEntry(entry),
ctrl,
})
.await
.map_err(|err| {
error!("Unable to send response: {err}");
LdapError::Transport
})?;
}
w.send(LdapMsg {
msgid,
op: LdapOp::SearchResultDone(result),
ctrl,
})
.await
.map_err(|err| {
error!("Unable to send response: {err}");
LdapError::Transport
})?;
// No state change
Ok(())
}