mod message;
pub(crate) mod providers;
pub use providers::{default_providers, provider_names};
use crate::error::ProviderError;
use crate::provider::Provider;
use crate::types::{IpVersion, Protocol};
use message::{StunMessage, StunMethod};
use std::future::Future;
use std::net::{IpAddr, SocketAddr};
use std::pin::Pin;
use tokio::net::UdpSocket;
#[derive(Debug, Clone)]
pub struct StunProvider {
name: String,
server: String,
port: u16,
}
impl StunProvider {
pub fn new(name: impl Into<String>, server: impl Into<String>, port: u16) -> Self {
Self {
name: name.into(),
server: server.into(),
port,
}
}
async fn binding_request(&self, version: IpVersion) -> Result<IpAddr, ProviderError> {
let server_addr = format!("{}:{}", self.server, self.port);
let addrs: Vec<SocketAddr> = tokio::net::lookup_host(&server_addr)
.await
.map_err(|e| ProviderError::new(&self.name, e))?
.collect();
let addr = addrs
.iter()
.find(|a| match version {
IpVersion::V4 => a.is_ipv4(),
IpVersion::V6 => a.is_ipv6(),
IpVersion::Any => true,
})
.ok_or_else(|| {
ProviderError::message(&self.name, "no suitable address for IP version")
})?;
let local_addr = if addr.is_ipv4() {
SocketAddr::from(([0, 0, 0, 0], 0))
} else {
SocketAddr::from(([0u16; 8], 0))
};
let socket = UdpSocket::bind(local_addr)
.await
.map_err(|e| ProviderError::new(&self.name, e))?;
socket
.connect(addr)
.await
.map_err(|e| ProviderError::new(&self.name, e))?;
let request = StunMessage::new(StunMethod::Request);
let request_bytes = request.encode();
socket
.send(&request_bytes)
.await
.map_err(|e| ProviderError::new(&self.name, e))?;
let mut buf = [0u8; 576]; let len = socket
.recv(&mut buf)
.await
.map_err(|e| ProviderError::new(&self.name, e))?;
let response =
StunMessage::decode(&buf[..len]).map_err(|e| ProviderError::message(&self.name, e))?;
if response.transaction_id() != request.transaction_id() {
return Err(ProviderError::message(
&self.name,
"transaction ID mismatch",
));
}
response
.get_mapped_address()
.ok_or_else(|| ProviderError::message(&self.name, "no mapped address in response"))
}
}
impl Provider for StunProvider {
fn name(&self) -> &str {
&self.name
}
fn protocol(&self) -> Protocol {
Protocol::Stun
}
fn supports_v4(&self) -> bool {
true
}
fn supports_v6(&self) -> bool {
true
}
fn get_ip(
&self,
version: IpVersion,
) -> Pin<Box<dyn Future<Output = Result<IpAddr, ProviderError>> + Send + '_>> {
Box::pin(self.binding_request(version))
}
}