feat: Implement LdapStream, begin event loop
This commit is contained in:
parent
d30df04064
commit
1ef02c58dd
8 changed files with 152 additions and 56 deletions
13
Cargo.lock
generated
13
Cargo.lock
generated
|
|
@ -69,6 +69,17 @@ version = "0.3.34"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "92d699e522242e69e3003b94ecc1f960f3a5e015aa7c5d7486e65ad01dd94f5e"
|
checksum = "92d699e522242e69e3003b94ecc1f960f3a5e015aa7c5d7486e65ad01dd94f5e"
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "futures-macro"
|
||||||
|
version = "0.3.34"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "9fb9654ba8355388abeb8dcb4fc62f511300867002afc858860463bdd9fe0c44"
|
||||||
|
dependencies = [
|
||||||
|
"proc-macro2",
|
||||||
|
"quote",
|
||||||
|
"syn 3.0.4",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "futures-sink"
|
name = "futures-sink"
|
||||||
version = "0.3.34"
|
version = "0.3.34"
|
||||||
|
|
@ -88,6 +99,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "0d50a92467f8ba5dd6e3ee5d4bd04d73ab2e4e1c44474a0674821dfce14b79bc"
|
checksum = "0d50a92467f8ba5dd6e3ee5d4bd04d73ab2e4e1c44474a0674821dfce14b79bc"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"futures-core",
|
"futures-core",
|
||||||
|
"futures-macro",
|
||||||
"futures-task",
|
"futures-task",
|
||||||
"pin-project-lite",
|
"pin-project-lite",
|
||||||
"slab",
|
"slab",
|
||||||
|
|
@ -160,6 +172,7 @@ version = "0.1.0"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"argh",
|
"argh",
|
||||||
"camino",
|
"camino",
|
||||||
|
"futures-util",
|
||||||
"ldap3_proto",
|
"ldap3_proto",
|
||||||
"serde",
|
"serde",
|
||||||
"tokio",
|
"tokio",
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@ edition = "2024"
|
||||||
[dependencies]
|
[dependencies]
|
||||||
argh = "0.1.19"
|
argh = "0.1.19"
|
||||||
camino = "1.2.5"
|
camino = "1.2.5"
|
||||||
|
futures-util = "0.3.34"
|
||||||
ldap3_proto = "0.8.1"
|
ldap3_proto = "0.8.1"
|
||||||
serde = { version = "1.0.229", features = ["derive"] }
|
serde = { version = "1.0.229", features = ["derive"] }
|
||||||
tokio = { version = "1.53.1", features = ["macros", "net", "rt", "time", "sync"] }
|
tokio = { version = "1.53.1", features = ["macros", "net", "rt", "time", "sync"] }
|
||||||
|
|
|
||||||
37
src/ldap/handler.rs
Normal file
37
src/ldap/handler.rs
Normal file
|
|
@ -0,0 +1,37 @@
|
||||||
|
use ldap3_proto::LdapMsg;
|
||||||
|
|
||||||
|
use crate::ldap::{LdapStream, LdapStreamError};
|
||||||
|
|
||||||
|
#[tracing::instrument(name = "ldap", skip(s), fields(session = %s.session))]
|
||||||
|
pub async fn ldap_handler(mut s: LdapStream) {
|
||||||
|
tracing::info! {
|
||||||
|
remote_addr = ?s.remote_addr,
|
||||||
|
"New client connection"
|
||||||
|
};
|
||||||
|
|
||||||
|
loop {
|
||||||
|
match s.next().await {
|
||||||
|
Ok(msg) => {
|
||||||
|
if let Err(e) = ldap_handler_inner(msg).await {
|
||||||
|
tracing::debug!(
|
||||||
|
reason = ?e,
|
||||||
|
"Failed to respond"
|
||||||
|
);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
tracing::debug!(
|
||||||
|
reason = ?e,
|
||||||
|
"Closing connection"
|
||||||
|
);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn ldap_handler_inner(msg: LdapMsg) -> Result<(), LdapStreamError> {
|
||||||
|
tracing::debug!(msg = ?msg, "Received LDAP message");
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
4
src/ldap/mod.rs
Normal file
4
src/ldap/mod.rs
Normal file
|
|
@ -0,0 +1,4 @@
|
||||||
|
mod handler;
|
||||||
|
pub use handler::ldap_handler;
|
||||||
|
mod stream;
|
||||||
|
pub use stream::{LdapStream, LdapStreamError};
|
||||||
69
src/ldap/stream.rs
Normal file
69
src/ldap/stream.rs
Normal file
|
|
@ -0,0 +1,69 @@
|
||||||
|
use futures_util::StreamExt;
|
||||||
|
use ldap3_proto::{LdapCodec, LdapMsg};
|
||||||
|
use tokio::time::{Duration, timeout};
|
||||||
|
use tokio_util::codec::Framed;
|
||||||
|
use uuid::Uuid;
|
||||||
|
|
||||||
|
use std::fmt;
|
||||||
|
|
||||||
|
use crate::stream::{AbstractSocketAddr, AbstractStreamKind};
|
||||||
|
|
||||||
|
#[derive(Debug)]
|
||||||
|
pub enum LdapStreamError {
|
||||||
|
ClientClosed,
|
||||||
|
Timeout(Duration),
|
||||||
|
LdapError(std::io::Error),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl fmt::Display for LdapStreamError {
|
||||||
|
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||||
|
match self {
|
||||||
|
Self::ClientClosed => write!(f, "Client closed the connection"),
|
||||||
|
Self::Timeout(d) => write!(f, "Client timeout after {}ms", d.as_millis()),
|
||||||
|
Self::LdapError(e) => write!(f, "Failed to parse LDAP message: {e:?}"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl std::error::Error for LdapStreamError {}
|
||||||
|
|
||||||
|
pub struct LdapStream {
|
||||||
|
pub remote_addr: AbstractSocketAddr,
|
||||||
|
pub inner: Framed<AbstractStreamKind, LdapCodec>,
|
||||||
|
/// An arbitrary ID to correlate logs with a specific TCP session
|
||||||
|
pub session: Uuid,
|
||||||
|
pub timeout: Duration,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl LdapStream {
|
||||||
|
pub fn new(
|
||||||
|
remote_addr: AbstractSocketAddr,
|
||||||
|
kind: AbstractStreamKind,
|
||||||
|
timeout: tokio::time::Duration,
|
||||||
|
) -> Self {
|
||||||
|
Self {
|
||||||
|
remote_addr,
|
||||||
|
inner: Framed::new(kind, LdapCodec::new(None, None)),
|
||||||
|
session: Uuid::new_v4(),
|
||||||
|
timeout,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn next(&mut self) -> Result<LdapMsg, LdapStreamError> {
|
||||||
|
// Check for timeout
|
||||||
|
let Ok(msg) = timeout(self.timeout, self.inner.next()).await else {
|
||||||
|
return Err(LdapStreamError::Timeout(self.timeout));
|
||||||
|
};
|
||||||
|
|
||||||
|
// Check for closed client connection
|
||||||
|
let Some(msg) = msg else {
|
||||||
|
return Err(LdapStreamError::ClientClosed);
|
||||||
|
};
|
||||||
|
|
||||||
|
// So far the message may or may not be a valid LDAP message,
|
||||||
|
// but using the codec some other conditions (eg. OOM) may trigger
|
||||||
|
// an error. We abort whether the message failed because it was malformed,
|
||||||
|
// because of transient IO error, or any other error condition.
|
||||||
|
msg.map_err(LdapStreamError::LdapError)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -1,12 +1,13 @@
|
||||||
use camino::Utf8PathBuf;
|
use camino::Utf8PathBuf;
|
||||||
use tokio::net::{TcpListener, UnixListener};
|
use tokio::net::{TcpListener, UnixListener};
|
||||||
|
use tokio::time::Duration;
|
||||||
|
|
||||||
use std::fmt;
|
use std::fmt;
|
||||||
use std::net::SocketAddr;
|
use std::net::SocketAddr;
|
||||||
use std::str::FromStr;
|
use std::str::FromStr;
|
||||||
|
|
||||||
use crate::error::GlobalError;
|
use crate::error::GlobalError;
|
||||||
use crate::stream::AbstractStream;
|
use crate::ldap::LdapStream;
|
||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub struct AcceptError {
|
pub struct AcceptError {
|
||||||
|
|
@ -98,10 +99,17 @@ impl Listener {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn accept(&self) -> Result<Option<AbstractStream>, AcceptError> {
|
pub async fn accept_ldap(&self) -> Result<Option<LdapStream>, AcceptError> {
|
||||||
|
// TODO: configurable timeout
|
||||||
|
let timeout = Duration::from_millis(500);
|
||||||
|
|
||||||
let res = match &self.kind {
|
let res = match &self.kind {
|
||||||
ListenerKind::Tcp(l) => l.accept().await.map(AbstractStream::from),
|
ListenerKind::Tcp(l) => l.accept().await.map(|(stream, remote_addr)| {
|
||||||
ListenerKind::Uds(l) => l.accept().await.map(AbstractStream::from),
|
LdapStream::new(remote_addr.into(), stream.into(), timeout)
|
||||||
|
}),
|
||||||
|
ListenerKind::Uds(l) => l.accept().await.map(|(stream, remote_addr)| {
|
||||||
|
LdapStream::new(remote_addr.into(), stream.into(), timeout)
|
||||||
|
}),
|
||||||
};
|
};
|
||||||
|
|
||||||
match res {
|
match res {
|
||||||
|
|
|
||||||
15
src/main.rs
15
src/main.rs
|
|
@ -2,21 +2,14 @@
|
||||||
|
|
||||||
mod cli;
|
mod cli;
|
||||||
mod error;
|
mod error;
|
||||||
|
mod ldap;
|
||||||
mod listener;
|
mod listener;
|
||||||
mod stream;
|
mod stream;
|
||||||
|
|
||||||
use cli::CliArgs;
|
use cli::CliArgs;
|
||||||
use error::GlobalError;
|
use error::GlobalError;
|
||||||
|
use ldap::ldap_handler;
|
||||||
use listener::ListenerPath;
|
use listener::ListenerPath;
|
||||||
use stream::AbstractStream;
|
|
||||||
|
|
||||||
#[tracing::instrument(name = "client", skip(s), fields(session = %s.session))]
|
|
||||||
async fn client(s: AbstractStream) {
|
|
||||||
tracing::info! {
|
|
||||||
remote_addr = ?s.remote_addr,
|
|
||||||
"New client connection"
|
|
||||||
};
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::main(flavor = "current_thread")]
|
#[tokio::main(flavor = "current_thread")]
|
||||||
async fn main() -> Result<(), GlobalError> {
|
async fn main() -> Result<(), GlobalError> {
|
||||||
|
|
@ -34,8 +27,8 @@ async fn main() -> Result<(), GlobalError> {
|
||||||
|
|
||||||
// If the connection is None, it's because the client aborted early
|
// If the connection is None, it's because the client aborted early
|
||||||
// so there's nothing to do about it.
|
// so there's nothing to do about it.
|
||||||
while let Some(stream) = listener.accept().await? {
|
while let Some(stream) = listener.accept_ldap().await? {
|
||||||
client(stream).await;
|
tokio::spawn(ldap_handler(stream));
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
|
|
|
||||||
|
|
@ -1,9 +1,8 @@
|
||||||
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
|
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
|
||||||
use tokio::net::unix::SocketAddr as UnixSocketAddr;
|
use tokio::net::unix::SocketAddr as UnixSocketAddr;
|
||||||
use tokio::net::{TcpStream, UnixStream};
|
use tokio::net::{TcpStream, UnixStream};
|
||||||
use uuid::Uuid;
|
|
||||||
|
|
||||||
use std::io::Result;
|
use std::io::Result as IOResult;
|
||||||
use std::marker::Unpin;
|
use std::marker::Unpin;
|
||||||
use std::net::SocketAddr as TcpSocketAddr;
|
use std::net::SocketAddr as TcpSocketAddr;
|
||||||
use std::pin::{Pin, pin};
|
use std::pin::{Pin, pin};
|
||||||
|
|
@ -47,70 +46,42 @@ impl From<UnixStream> for AbstractStreamKind {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug)]
|
impl Unpin for AbstractStreamKind {}
|
||||||
pub struct AbstractStream {
|
|
||||||
pub remote_addr: AbstractSocketAddr,
|
|
||||||
pub kind: AbstractStreamKind,
|
|
||||||
/// An arbitrary ID to correlate logs with a specific TCP session
|
|
||||||
pub session: Uuid,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl From<(UnixStream, UnixSocketAddr)> for AbstractStream {
|
impl AsyncRead for AbstractStreamKind {
|
||||||
fn from(res: (UnixStream, UnixSocketAddr)) -> Self {
|
|
||||||
Self {
|
|
||||||
remote_addr: res.1.into(),
|
|
||||||
kind: res.0.into(),
|
|
||||||
session: Uuid::new_v4(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl From<(TcpStream, TcpSocketAddr)> for AbstractStream {
|
|
||||||
fn from(res: (TcpStream, std::net::SocketAddr)) -> Self {
|
|
||||||
Self {
|
|
||||||
remote_addr: res.1.into(),
|
|
||||||
kind: res.0.into(),
|
|
||||||
session: Uuid::new_v4(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Unpin for AbstractStream {}
|
|
||||||
|
|
||||||
impl AsyncRead for AbstractStream {
|
|
||||||
fn poll_read(
|
fn poll_read(
|
||||||
mut self: Pin<&mut Self>,
|
mut self: Pin<&mut Self>,
|
||||||
cx: &mut Context<'_>,
|
cx: &mut Context<'_>,
|
||||||
buf: &mut ReadBuf<'_>,
|
buf: &mut ReadBuf<'_>,
|
||||||
) -> Poll<Result<()>> {
|
) -> Poll<IOResult<()>> {
|
||||||
match &mut self.kind {
|
match &mut *self {
|
||||||
AbstractStreamKind::Tcp(stream) => pin!(stream).poll_read(cx, buf),
|
AbstractStreamKind::Tcp(stream) => pin!(stream).poll_read(cx, buf),
|
||||||
AbstractStreamKind::Uds(stream) => pin!(stream).poll_read(cx, buf),
|
AbstractStreamKind::Uds(stream) => pin!(stream).poll_read(cx, buf),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl AsyncWrite for AbstractStream {
|
impl AsyncWrite for AbstractStreamKind {
|
||||||
fn poll_write(
|
fn poll_write(
|
||||||
mut self: Pin<&mut Self>,
|
mut self: Pin<&mut Self>,
|
||||||
cx: &mut Context<'_>,
|
cx: &mut Context<'_>,
|
||||||
buf: &[u8],
|
buf: &[u8],
|
||||||
) -> Poll<Result<usize>> {
|
) -> Poll<IOResult<usize>> {
|
||||||
match &mut self.kind {
|
match &mut *self {
|
||||||
AbstractStreamKind::Tcp(stream) => pin!(stream).poll_write(cx, buf),
|
AbstractStreamKind::Tcp(stream) => pin!(stream).poll_write(cx, buf),
|
||||||
AbstractStreamKind::Uds(stream) => pin!(stream).poll_write(cx, buf),
|
AbstractStreamKind::Uds(stream) => pin!(stream).poll_write(cx, buf),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<()>> {
|
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<IOResult<()>> {
|
||||||
match &mut self.kind {
|
match &mut *self {
|
||||||
AbstractStreamKind::Tcp(stream) => pin!(stream).poll_flush(cx),
|
AbstractStreamKind::Tcp(stream) => pin!(stream).poll_flush(cx),
|
||||||
AbstractStreamKind::Uds(stream) => pin!(stream).poll_flush(cx),
|
AbstractStreamKind::Uds(stream) => pin!(stream).poll_flush(cx),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<()>> {
|
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<IOResult<()>> {
|
||||||
match &mut self.kind {
|
match &mut *self {
|
||||||
AbstractStreamKind::Tcp(stream) => pin!(stream).poll_shutdown(cx),
|
AbstractStreamKind::Tcp(stream) => pin!(stream).poll_shutdown(cx),
|
||||||
AbstractStreamKind::Uds(stream) => pin!(stream).poll_shutdown(cx),
|
AbstractStreamKind::Uds(stream) => pin!(stream).poll_shutdown(cx),
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue