use std::io;
use serde::Deserialize;
use crate::core::{ReadHalf, WriteHalf};
pub mod tcp;
pub mod unix;
pub use self::tcp::TcpDialConfig;
pub use self::unix::UnixDialConfig;
pub type DialedHalves = (ReadHalf, WriteHalf);
#[async_trait::async_trait]
pub trait Dialer: Send + Sync {
async fn dial(&self) -> io::Result<DialedHalves>;
}
#[derive(Debug, thiserror::Error)]
pub enum DialConfigError {
#[error("transport: no dial transport configured (expected exactly one of [Dial.Unix] or [Dial.Tcp])")]
NoTransport,
#[error("transport: exactly one dial transport must be configured, got {0}")]
MultipleTransports(usize),
}
impl From<DialConfigError> for io::Error {
fn from(e: DialConfigError) -> io::Error {
io::Error::new(io::ErrorKind::InvalidInput, e)
}
}
#[derive(Clone, Debug, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct DialConfig {
#[serde(rename = "Unix", default)]
pub unix: Option<UnixDialConfig>,
#[serde(rename = "Tcp", default)]
pub tcp: Option<TcpDialConfig>,
}
impl DialConfig {
pub fn validate(&self) -> Result<(), DialConfigError> {
let n = (self.unix.is_some() as usize) + (self.tcp.is_some() as usize);
match n {
0 => Err(DialConfigError::NoTransport),
1 => Ok(()),
_ => Err(DialConfigError::MultipleTransports(n)),
}
}
pub fn resolve(&self) -> Result<&dyn Dialer, DialConfigError> {
self.validate()?;
if let Some(ref u) = self.unix {
return Ok(u);
}
if let Some(ref t) = self.tcp {
return Ok(t);
}
Err(DialConfigError::NoTransport)
}
pub async fn dial(&self) -> io::Result<DialedHalves> {
let dialer = self.resolve()?;
dialer.dial().await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn validate_zero_subtables_rejected() {
let cfg = DialConfig::default();
let err = cfg.validate().unwrap_err();
assert!(matches!(err, DialConfigError::NoTransport));
}
#[test]
fn validate_multiple_subtables_rejected() {
let cfg = DialConfig {
unix: Some(UnixDialConfig { address: "/tmp/x.sock".into() }),
tcp: Some(TcpDialConfig { address: "localhost:0".into(), network: None }),
};
let err = cfg.validate().unwrap_err();
assert!(matches!(err, DialConfigError::MultipleTransports(2)));
}
#[test]
fn validate_single_unix_accepted() {
let cfg = DialConfig {
unix: Some(UnixDialConfig { address: "/tmp/x.sock".into() }),
tcp: None,
};
cfg.validate().unwrap();
}
#[test]
fn validate_single_tcp_accepted() {
let cfg = DialConfig {
unix: None,
tcp: Some(TcpDialConfig { address: "localhost:0".into(), network: None }),
};
cfg.validate().unwrap();
}
#[test]
fn resolve_returns_correct_dialer_type() {
let cfg = DialConfig {
unix: None,
tcp: Some(TcpDialConfig { address: "localhost:0".into(), network: None }),
};
let dialer = cfg.resolve().unwrap();
assert!(cfg.tcp.is_some() && cfg.unix.is_none());
let _ = dialer as *const _;
}
#[test]
fn toml_parses_dial_tcp_subtable() {
let src = r#"
[Tcp]
Address = "localhost:64331"
"#;
let cfg: DialConfig = toml::from_str(src).unwrap();
assert!(cfg.tcp.is_some());
assert!(cfg.unix.is_none());
cfg.validate().unwrap();
let tcp = cfg.tcp.as_ref().unwrap();
assert_eq!(tcp.address, "localhost:64331");
}
#[test]
fn toml_parses_dial_unix_subtable() {
let src = r#"
[Unix]
Address = "/tmp/kp.sock"
"#;
let cfg: DialConfig = toml::from_str(src).unwrap();
assert!(cfg.unix.is_some());
assert!(cfg.tcp.is_none());
cfg.validate().unwrap();
assert_eq!(cfg.unix.as_ref().unwrap().address, "/tmp/kp.sock");
}
#[test]
fn toml_rejects_flat_old_fields() {
let src = r#"
Network = "tcp"
Address = "localhost:64331"
"#;
let err = toml::from_str::<DialConfig>(src).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("unknown field") && msg.contains("Network"),
"expected unknown-field rejection, got: {msg}",
);
}
}