use std::time::Duration;
use crate::message::{Header, Message, Method, Request, Response};
use crate::name::HeaderName;
use crate::transaction::timing::{Timer, Timers};
use crate::transaction::{Output, Reason, Reliability, TuEvent};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ClientState {
Calling,
Trying,
Proceeding,
Completed,
Accepted,
Terminated,
}
impl ClientState {
#[must_use]
pub fn is_terminated(self) -> bool {
matches!(self, Self::Terminated)
}
}
#[derive(Debug)]
pub struct ClientTransaction {
request: Request,
is_invite: bool,
state: ClientState,
reliability: Reliability,
timers: Timers,
interval: Duration,
ack: Option<Request>,
}
impl ClientTransaction {
#[must_use]
pub fn new(request: Request, reliability: Reliability, timers: Timers) -> (Self, Vec<Output>) {
let is_invite = request.method == Method::Invite;
let mut tx = Self {
request,
is_invite,
state: if is_invite {
ClientState::Calling
} else {
ClientState::Trying
},
reliability,
timers,
interval: timers.t1,
ack: None,
};
let mut out = vec![Output::send(Message::Request(tx.request.clone()))];
if !reliability.is_reliable() {
out.push(Output::SetTimer {
timer: if is_invite { Timer::A } else { Timer::E },
after: tx.interval,
});
}
out.push(Output::SetTimer {
timer: if is_invite { Timer::B } else { Timer::F },
after: timers.timeout(),
});
tx.state = if is_invite {
ClientState::Calling
} else {
ClientState::Trying
};
(tx, out)
}
#[must_use]
pub fn state(&self) -> ClientState {
self.state
}
#[must_use]
pub fn request(&self) -> &Request {
&self.request
}
pub fn on_response(&mut self, response: Response) -> Vec<Output> {
if self.is_invite {
self.invite_response(response)
} else {
self.non_invite_response(response)
}
}
fn invite_response(&mut self, response: Response) -> Vec<Output> {
let status = response.status;
match self.state {
ClientState::Calling | ClientState::Proceeding => {
let from_calling = self.state == ClientState::Calling;
let mut out = Vec::new();
if status.is_provisional() {
if from_calling {
out.push(Output::ClearTimer(Timer::A));
out.push(Output::ClearTimer(Timer::B));
}
self.state = ClientState::Proceeding;
out.push(Output::to_tu(TuEvent::Response(Box::new(response))));
return out;
}
if from_calling {
out.push(Output::ClearTimer(Timer::A));
}
out.push(Output::ClearTimer(Timer::B));
if status.is_success() {
self.state = ClientState::Accepted;
out.push(Output::to_tu(TuEvent::Response(Box::new(response))));
out.push(Output::SetTimer {
timer: Timer::M,
after: self.timers.timeout(),
});
} else {
let ack = make_ack(&self.request, &response);
out.push(Output::send(Message::Request(ack.clone())));
self.ack = Some(ack);
self.state = ClientState::Completed;
out.push(Output::to_tu(TuEvent::Response(Box::new(response))));
out.push(Output::SetTimer {
timer: Timer::D,
after: self.timers.timer_d(self.reliability),
});
}
out
}
ClientState::Completed => {
self.ack
.as_ref()
.map(|ack| Output::send(Message::Request(ack.clone())))
.into_iter()
.collect()
}
ClientState::Accepted => {
if status.is_success() {
vec![Output::to_tu(TuEvent::Response(Box::new(response)))]
} else {
Vec::new()
}
}
ClientState::Trying | ClientState::Terminated => Vec::new(),
}
}
fn non_invite_response(&mut self, response: Response) -> Vec<Output> {
match self.state {
ClientState::Trying | ClientState::Proceeding => {
if response.status.is_provisional() {
self.state = ClientState::Proceeding;
return vec![Output::to_tu(TuEvent::Response(Box::new(response)))];
}
self.state = ClientState::Completed;
vec![
Output::ClearTimer(Timer::E),
Output::ClearTimer(Timer::F),
Output::to_tu(TuEvent::Response(Box::new(response))),
Output::SetTimer {
timer: Timer::K,
after: self.timers.absorb(self.reliability),
},
]
}
_ => Vec::new(),
}
}
pub fn on_timer(&mut self, timer: Timer) -> Vec<Output> {
match (self.state, timer) {
(ClientState::Calling, Timer::A) => {
self.interval = self.timers.double(self.interval);
vec![
Output::send(Message::Request(self.request.clone())),
Output::SetTimer {
timer: Timer::A,
after: self.interval,
},
]
}
(ClientState::Trying, Timer::E) => {
self.interval = self.timers.double_capped(self.interval);
vec![
Output::send(Message::Request(self.request.clone())),
Output::SetTimer {
timer: Timer::E,
after: self.interval,
},
]
}
(ClientState::Proceeding, Timer::E) if !self.is_invite => {
self.interval = self.timers.t2;
vec![
Output::send(Message::Request(self.request.clone())),
Output::SetTimer {
timer: Timer::E,
after: self.interval,
},
]
}
(ClientState::Calling, Timer::B)
| (ClientState::Trying | ClientState::Proceeding, Timer::F) => {
self.state = ClientState::Terminated;
vec![
Output::to_tu(TuEvent::Timeout),
Output::Terminated(Reason::Timeout),
]
}
(ClientState::Completed, Timer::D | Timer::K) | (ClientState::Accepted, Timer::M) => {
self.state = ClientState::Terminated;
vec![Output::Terminated(Reason::Completed)]
}
_ => Vec::new(),
}
}
pub fn on_transport_error(&mut self) -> Vec<Output> {
if self.state.is_terminated() {
return Vec::new();
}
self.state = ClientState::Terminated;
vec![
Output::to_tu(TuEvent::TransportError),
Output::Terminated(Reason::TransportError),
]
}
}
fn make_ack(request: &Request, response: &Response) -> Request {
let mut ack = Request::new(Method::Ack, request.uri.clone());
if let Some(via) = request.headers.get(&HeaderName::Via) {
ack.headers.push(via.clone());
}
if let Some(from) = request.headers.get(&HeaderName::From) {
ack.headers.push(from.clone());
}
if let Some(to) = response.headers.get(&HeaderName::To) {
ack.headers.push(to.clone());
}
if let Some(call_id) = request.headers.get(&HeaderName::CallId) {
ack.headers.push(call_id.clone());
}
for route in request.headers.get_all(&HeaderName::Route) {
ack.headers.push(route.clone());
}
ack.headers.push(Header::new_unchecked(
HeaderName::MaxForwards,
bytes::Bytes::from_static(b"70"),
));
let sequence = request
.headers
.value(&HeaderName::CSeq)
.and_then(|v| {
let digits: Vec<u8> = v.iter().copied().take_while(u8::is_ascii_digit).collect();
String::from_utf8(digits).ok()
})
.unwrap_or_default();
let mut cseq = sequence.into_bytes();
cseq.extend_from_slice(b" ACK");
ack.headers.push(Header::new_unchecked(
HeaderName::CSeq,
bytes::Bytes::from(cseq),
));
ack.headers.push(Header::new_unchecked(
HeaderName::ContentLength,
bytes::Bytes::from_static(b"0"),
));
ack
}