#[cfg(test)]
mod client_test;
pub mod binding;
pub mod permission;
mod proto;
pub mod relay;
pub mod transaction;
use bytes::BytesMut;
use log::{debug, trace};
use std::collections::{HashMap, VecDeque};
use std::net::SocketAddr;
use std::time::{Duration, Instant};
use stun::attributes::*;
use stun::integrity::*;
use stun::message::*;
use stun::textattrs::*;
use stun::xoraddr::*;
use binding::*;
use transaction::*;
use crate::client::relay::{Relay, RelayState};
use crate::proto::chandata::*;
use crate::proto::channum::ChannelNumber;
use crate::proto::data::*;
use crate::proto::lifetime::Lifetime;
use crate::proto::peeraddr::*;
use crate::proto::relayaddr::RelayedAddress;
use crate::proto::reqtrans::RequestedTransport;
use crate::proto::{PROTO_TCP, PROTO_UDP};
use shared::error::{Error, Result};
use shared::util::lookup_host;
use shared::{TransportContext, TransportMessage, TransportProtocol};
use stun::error_code::ErrorCodeAttribute;
use stun::fingerprint::FINGERPRINT;
const DEFAULT_RTO_IN_MS: u64 = 200;
const MAX_DATA_BUFFER_SIZE: usize = u16::MAX as usize; const MAX_READ_QUEUE_SIZE: usize = 1024;
pub type RelayedAddr = SocketAddr;
pub type ReflexiveAddr = SocketAddr;
pub type PeerAddr = SocketAddr;
#[derive(Debug)]
pub enum Event {
TransactionTimeout(TransactionId),
BindingResponse(TransactionId, ReflexiveAddr),
BindingError(TransactionId, Error),
AllocateResponse(TransactionId, RelayedAddr),
AllocateError(TransactionId, Error),
CreatePermissionResponse(TransactionId, PeerAddr),
CreatePermissionError(TransactionId, Error),
DataIndicationOrChannelData(Option<ChannelNumber>, PeerAddr, BytesMut),
}
enum AllocateState {
Attempting,
Requesting(TextAttribute),
}
pub struct ClientConfig {
pub stun_serv_addr: String, pub turn_serv_addr: String, pub local_addr: SocketAddr,
pub transport_protocol: TransportProtocol,
pub username: String,
pub password: String,
pub realm: String,
pub software: String,
pub rto_in_ms: u64,
}
pub struct Client {
stun_serv_addr: Option<SocketAddr>,
turn_serv_addr: Option<SocketAddr>,
local_addr: SocketAddr,
transport_protocol: TransportProtocol,
username: Username,
password: String,
realm: Realm,
integrity: MessageIntegrity,
software: Software,
tr_map: TransactionMap,
binding_mgr: BindingManager,
rto_in_ms: u64,
relays: HashMap<RelayedAddr, RelayState>,
transmits: VecDeque<TransportMessage<BytesMut>>,
events: VecDeque<Event>,
}
impl Client {
pub fn new(config: ClientConfig) -> Result<Self> {
let stun_serv_addr = if config.stun_serv_addr.is_empty() {
None
} else {
Some(lookup_host(
config.local_addr.is_ipv4(),
config.stun_serv_addr.as_str(),
)?)
};
let turn_serv_addr = if config.turn_serv_addr.is_empty() {
None
} else {
Some(lookup_host(
config.local_addr.is_ipv4(),
config.turn_serv_addr.as_str(),
)?)
};
Ok(Client {
stun_serv_addr,
turn_serv_addr,
local_addr: config.local_addr,
transport_protocol: config.transport_protocol,
username: Username::new(ATTR_USERNAME, config.username),
password: config.password,
realm: Realm::new(ATTR_REALM, config.realm),
software: Software::new(ATTR_SOFTWARE, config.software),
tr_map: TransactionMap::new(),
binding_mgr: BindingManager::new(),
rto_in_ms: if config.rto_in_ms != 0 {
config.rto_in_ms
} else {
DEFAULT_RTO_IN_MS
},
integrity: MessageIntegrity::new_short_term_integrity(String::new()),
relays: HashMap::new(),
transmits: VecDeque::new(),
events: VecDeque::new(),
})
}
fn handle_inbound(&mut self, data: &[u8], from: SocketAddr) -> Result<()> {
if is_stun_message(data) {
self.handle_stun_message(data)
} else if ChannelData::is_channel_data(data) {
self.handle_channel_data(data)
} else if self.stun_serv_addr.is_some() && &from == self.stun_serv_addr.as_ref().unwrap() {
Err(Error::ErrNonStunmessage)
} else {
trace!("non-STUN/TURN packet, unhandled");
Ok(())
}
}
fn handle_stun_message(&mut self, data: &[u8]) -> Result<()> {
let mut msg = Message::new();
msg.raw = data.to_vec();
msg.decode()?;
if msg.typ.class == CLASS_REQUEST {
return Err(Error::Other(format!(
"{:?} : {}",
Error::ErrUnexpectedStunrequestMessage,
msg
)));
}
if msg.typ.class == CLASS_INDICATION {
if msg.typ.method == METHOD_DATA {
let mut peer_addr = PeerAddress::default();
peer_addr.get_from(&msg)?;
let from = SocketAddr::new(peer_addr.ip, peer_addr.port);
let mut data = Data::default();
data.get_from(&msg)?;
debug!("data indication received from {}", from);
self.events.push_back(Event::DataIndicationOrChannelData(
None,
from,
BytesMut::from(&data.0[..]),
))
}
return Ok(());
}
if self.tr_map.find(&msg.transaction_id).is_none() {
debug!("no transaction for {}", msg);
return Ok(());
}
if let Some(tr) = self.tr_map.delete(&msg.transaction_id) {
match msg.typ.method {
METHOD_BINDING => {
if msg.typ.class == CLASS_ERROR_RESPONSE {
let mut code = ErrorCodeAttribute::default();
let err = if code.get_from(&msg).is_err() {
Error::Other(format!("{}", msg.typ))
} else {
Error::Other(format!("{} (error {})", msg.typ, code))
};
self.events
.push_back(Event::BindingError(tr.transaction_id, err));
} else {
let mut refl_addr = XorMappedAddress::default();
match refl_addr.get_from(&msg) {
Ok(_) => {
self.events.push_back(Event::BindingResponse(
tr.transaction_id,
ReflexiveAddr::new(refl_addr.ip, refl_addr.port),
));
}
Err(err) => {
self.events
.push_back(Event::BindingError(tr.transaction_id, err));
}
}
}
}
METHOD_ALLOCATE => {
self.handle_allocate_response(msg, tr.transaction_type)?;
}
METHOD_CREATE_PERMISSION => {
if let TransactionType::CreatePermissionRequest(relayed_addr, peer_addr) =
tr.transaction_type
{
let mut relay = Relay {
relayed_addr,
client: self,
};
relay.handle_create_permission_response(msg, peer_addr)?;
}
}
METHOD_REFRESH => {
if let TransactionType::RefreshRequest(relayed_addr) = tr.transaction_type {
let mut relay = Relay {
relayed_addr,
client: self,
};
relay.handle_refresh_allocation_response(msg)?;
}
}
METHOD_CHANNEL_BIND => {
if let TransactionType::ChannelBindRequest(relayed_addr, bind_addr) =
tr.transaction_type
{
let mut relay = Relay {
relayed_addr,
client: self,
};
relay.handle_channel_bind_response(msg, bind_addr)?;
}
}
_ => {}
}
}
Ok(())
}
fn handle_channel_data(&mut self, data: &[u8]) -> Result<()> {
let mut ch_data = ChannelData {
raw: data.to_vec(),
..Default::default()
};
ch_data.decode()?;
let addr = self
.find_addr_by_channel_number(ch_data.number.0)
.ok_or(Error::ErrChannelBindNotFound)?;
trace!(
"channel data received from {} (ch={})",
addr, ch_data.number.0
);
self.events.push_back(Event::DataIndicationOrChannelData(
Some(ch_data.number),
addr,
BytesMut::from(&ch_data.data[..]),
));
Ok(())
}
pub fn relay(&mut self, relayed_addr: SocketAddr) -> Result<Relay<'_>> {
if !self.relays.contains_key(&relayed_addr) {
Err(Error::ErrStreamNotExisted)
} else {
Ok(Relay {
relayed_addr,
client: self,
})
}
}
pub fn send_binding_request_to(&mut self, to: SocketAddr) -> Result<TransactionId> {
let msg = {
let attrs: Vec<Box<dyn Setter>> = if !self.software.text.is_empty() {
vec![
Box::new(TransactionId::new()),
Box::new(BINDING_REQUEST),
Box::new(self.software.clone()),
]
} else {
vec![Box::new(TransactionId::new()), Box::new(BINDING_REQUEST)]
};
let mut msg = Message::new();
msg.build(&attrs)?;
msg
};
debug!("client.SendBindingRequestTo call PerformTransaction 1");
Ok(self.perform_transaction(&msg, to, TransactionType::BindingRequest))
}
pub fn send_binding_request(&mut self) -> Result<TransactionId> {
if let Some(stun_serv_addr) = &self.stun_serv_addr {
self.send_binding_request_to(*stun_serv_addr)
} else {
Err(Error::ErrStunserverAddressNotSet)
}
}
fn find_addr_by_channel_number(&self, ch_num: u16) -> Option<SocketAddr> {
self.binding_mgr.find_by_number(ch_num).map(|b| b.addr)
}
fn stun_server_addr(&self) -> Option<SocketAddr> {
self.stun_serv_addr
}
pub fn update_credentials(&mut self, username: String, password: String) {
self.username = Username::new(ATTR_USERNAME, username);
self.password = password;
self.integrity = MessageIntegrity::new_long_term_integrity(
self.username.text.clone(),
self.realm.text.clone(),
self.password.clone(),
);
for relay in self.relays.values_mut() {
relay.integrity = self.integrity.clone();
}
}
pub fn refresh_allocations(&mut self) -> Result<()> {
let relays: Vec<(RelayedAddr, Duration)> = self
.relays
.iter()
.map(|(addr, relay)| (*addr, relay.lifetime))
.collect();
for (relayed_addr, lifetime) in relays {
self.relay(relayed_addr)?.refresh_allocation(lifetime)?;
}
Ok(())
}
pub fn allocate(&mut self) -> Result<TransactionId> {
let mut msg = Message::new();
msg.build(&[
Box::new(TransactionId::new()),
Box::new(MessageType::new(METHOD_ALLOCATE, CLASS_REQUEST)),
Box::new(RequestedTransport {
protocol: if self.transport_protocol == TransportProtocol::UDP {
PROTO_UDP
} else {
PROTO_TCP
},
}),
Box::new(FINGERPRINT),
])?;
debug!("client.Allocate call PerformTransaction 1");
let mut tid = self.perform_transaction(
&msg,
self.turn_server_addr()?,
TransactionType::AllocateAttempt,
);
tid.0[TRANSACTION_ID_SIZE - 1] = tid.0[TRANSACTION_ID_SIZE - 1].wrapping_add(1);
Ok(tid)
}
fn handle_allocate_response(
&mut self,
response: Message,
allocate_state: TransactionType,
) -> Result<()> {
match allocate_state {
TransactionType::AllocateAttempt => {
let nonce = match Nonce::get_from_as(&response, ATTR_NONCE) {
Ok(nonce) => nonce,
Err(err) => {
self.events
.push_back(Event::AllocateError(response.transaction_id, err));
return Ok(());
}
};
self.realm = match Realm::get_from_as(&response, ATTR_REALM) {
Ok(realm) => realm,
Err(err) => {
self.events
.push_back(Event::AllocateError(response.transaction_id, err));
return Ok(());
}
};
self.integrity = MessageIntegrity::new_long_term_integrity(
self.username.text.clone(),
self.realm.text.clone(),
self.password.clone(),
);
let mut msg = Message::new();
let mut tid = response.transaction_id;
tid.0[TRANSACTION_ID_SIZE - 1] = tid.0[TRANSACTION_ID_SIZE - 1].wrapping_add(1);
msg.build(&[
Box::new(tid),
Box::new(MessageType::new(METHOD_ALLOCATE, CLASS_REQUEST)),
Box::new(RequestedTransport {
protocol: if self.transport_protocol == TransportProtocol::UDP {
PROTO_UDP
} else {
PROTO_TCP
},
}),
Box::new(self.username.clone()),
Box::new(self.realm.clone()),
Box::new(nonce.clone()),
Box::new(self.integrity.clone()),
Box::new(FINGERPRINT),
])?;
debug!("client.Allocate call PerformTransaction 2");
self.perform_transaction(
&msg,
self.turn_server_addr()?,
TransactionType::AllocateRequest(nonce),
);
}
TransactionType::AllocateRequest(nonce) => {
if response.typ.class == CLASS_ERROR_RESPONSE {
let mut code = ErrorCodeAttribute::default();
let err = if code.get_from(&response).is_err() {
Error::Other(format!("{}", response.typ))
} else {
Error::Other(format!("{} (error {})", response.typ, code))
};
self.events
.push_back(Event::AllocateError(response.transaction_id, err));
return Ok(());
}
let mut relayed = RelayedAddress::default();
relayed.get_from(&response)?;
let relayed_addr = RelayedAddr::new(relayed.ip, relayed.port);
let mut lifetime = Lifetime::default();
lifetime.get_from(&response)?;
self.relays.insert(
relayed_addr,
RelayState::new(relayed_addr, self.integrity.clone(), nonce, lifetime.0),
);
self.events.push_back(Event::AllocateResponse(
response.transaction_id,
relayed_addr,
));
}
_ => {}
}
Ok(())
}
fn turn_server_addr(&self) -> Result<SocketAddr> {
self.turn_serv_addr.ok_or(Error::ErrNilTurnSocket)
}
fn username(&self) -> Username {
self.username.clone()
}
fn realm(&self) -> Realm {
self.realm.clone()
}
fn write_to(&mut self, data: &[u8], remote: SocketAddr) {
self.transmits.push_back(TransportMessage {
now: Instant::now(),
transport: TransportContext {
local_addr: self.local_addr,
peer_addr: remote,
transport_protocol: self.transport_protocol,
ecn: None,
},
message: BytesMut::from(data),
});
}
fn perform_transaction(
&mut self,
msg: &Message,
to: SocketAddr,
transaction_type: TransactionType,
) -> TransactionId {
let tr = Transaction::new(TransactionConfig {
transaction_id: msg.transaction_id,
transaction_type,
raw: BytesMut::from(&msg.raw[..]),
local_addr: self.local_addr,
peer_addr: to,
transport_protocol: self.transport_protocol,
interval: self.rto_in_ms,
});
trace!(
"start {} transaction {:?} to {}",
msg.typ, msg.transaction_id, tr.peer_addr
);
self.tr_map.insert(msg.transaction_id, tr);
self.write_to(&msg.raw, to);
msg.transaction_id
}
}