use bytes::Bytes;
use crate::error::BuildError;
use crate::headers::grammar::is_token_char;
use crate::message::{Header, Headers, Method, Request, Response, StatusCode};
use crate::name::HeaderName;
use crate::uri::Uri;
pub(crate) fn check_value(value: &[u8], field: &'static str) -> Result<(), BuildError> {
if let Some(pos) = value
.iter()
.position(|&b| matches!(b, b'\r' | b'\n' | b'\0'))
{
return Err(BuildError::IllegalCharacter {
field,
offset: pos,
byte: value.get(pos).copied().unwrap_or(0),
});
}
Ok(())
}
pub(crate) fn check_token(value: &[u8], field: &'static str) -> Result<(), BuildError> {
check_value(value, field)?;
if value.is_empty() || !value.iter().all(|&b| is_token_char(b)) {
return Err(BuildError::NotAToken { field });
}
Ok(())
}
impl Header {
pub fn build(name: HeaderName, value: impl Into<Bytes>) -> Result<Self, BuildError> {
let value = value.into();
check_value(&value, "header value")?;
if let HeaderName::Other(raw) = &name {
check_token(raw, "header name")?;
}
Ok(Self::new_unchecked(name, value))
}
}
#[derive(Debug)]
pub struct RequestBuilder {
method: Method,
uri: Uri,
headers: Headers,
body: Bytes,
}
impl RequestBuilder {
#[must_use]
pub fn new(method: Method, uri: Uri) -> Self {
Self {
method,
uri,
headers: Headers::new(),
body: Bytes::new(),
}
}
pub fn header(mut self, name: HeaderName, value: impl Into<Bytes>) -> Result<Self, BuildError> {
self.headers.push(Header::build(name, value)?);
Ok(self)
}
pub fn set_header(
mut self,
name: &HeaderName,
value: impl Into<Bytes>,
) -> Result<Self, BuildError> {
let header = Header::build(name.clone(), value)?;
self.headers.remove_all(name);
self.headers.push(header);
Ok(self)
}
#[must_use]
pub fn max_forwards(mut self, hops: u8) -> Self {
self.headers.push(Header::new_unchecked(
HeaderName::MaxForwards,
Bytes::from(hops.to_string()),
));
self
}
pub fn cseq(mut self, sequence: u32, method: &Method) -> Result<Self, BuildError> {
check_token(method.as_bytes(), "CSeq method")?;
let mut value = sequence.to_string().into_bytes();
value.push(b' ');
value.extend_from_slice(method.as_bytes());
self.headers
.push(Header::new_unchecked(HeaderName::CSeq, Bytes::from(value)));
Ok(self)
}
#[must_use]
pub fn body(mut self, body: impl Into<Bytes>) -> Self {
let body = body.into();
self.headers.remove_all(&HeaderName::ContentLength);
self.headers.push(Header::new_unchecked(
HeaderName::ContentLength,
Bytes::from(body.len().to_string()),
));
self.body = body;
self
}
#[must_use]
pub fn build(mut self) -> Request {
if self.headers.get(&HeaderName::ContentLength).is_none() {
self.headers.push(Header::new_unchecked(
HeaderName::ContentLength,
Bytes::from_static(b"0"),
));
}
let mut request = Request::new(self.method, self.uri);
request.headers = self.headers;
request.set_body(self.body);
request
}
}
#[derive(Debug)]
pub struct ResponseBuilder {
status: StatusCode,
reason: Bytes,
headers: Headers,
body: Bytes,
}
impl ResponseBuilder {
pub fn new(status: StatusCode, reason: impl Into<Bytes>) -> Result<Self, BuildError> {
let reason = reason.into();
check_value(&reason, "reason phrase")?;
Ok(Self {
status,
reason,
headers: Headers::new(),
body: Bytes::new(),
})
}
pub fn to_request(
request: &Request,
status: StatusCode,
reason: impl Into<Bytes>,
) -> Result<Self, BuildError> {
let mut builder = Self::new(status, reason)?;
for (name, label) in [
(HeaderName::Via, "Via"),
(HeaderName::From, "From"),
(HeaderName::To, "To"),
(HeaderName::CallId, "Call-ID"),
(HeaderName::CSeq, "CSeq"),
] {
if request.headers.get(&name).is_none() {
return Err(BuildError::MissingRequiredResponseHeader { header: label });
}
for header in request.headers.get_all(&name) {
builder.headers.push(header.clone());
}
}
if let Some(history) = crate::headers::history::for_response(request, status) {
builder
.headers
.push(Header::new_unchecked(HeaderName::HistoryInfo, history));
}
Ok(builder)
}
pub fn header(mut self, name: HeaderName, value: impl Into<Bytes>) -> Result<Self, BuildError> {
self.headers.push(Header::build(name, value)?);
Ok(self)
}
pub fn set_header(
mut self,
name: &HeaderName,
value: impl Into<Bytes>,
) -> Result<Self, BuildError> {
let header = Header::build(name.clone(), value)?;
self.headers.remove_all(name);
self.headers.push(header);
Ok(self)
}
#[must_use]
pub fn body(mut self, body: impl Into<Bytes>) -> Self {
let body = body.into();
self.headers.remove_all(&HeaderName::ContentLength);
self.headers.push(Header::new_unchecked(
HeaderName::ContentLength,
Bytes::from(body.len().to_string()),
));
self.body = body;
self
}
#[must_use]
pub fn build(mut self) -> Response {
if self.headers.get(&HeaderName::ContentLength).is_none() {
self.headers.push(Header::new_unchecked(
HeaderName::ContentLength,
Bytes::from_static(b"0"),
));
}
let mut response = Response::new(self.status, self.reason);
response.headers = self.headers;
response.set_body(self.body);
response
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
clippy::indexing_slicing
)]
mod tests {
use super::*;
use crate::headers::CSeq;
use crate::uri::{Host, HostName};
fn uri() -> Uri {
Uri::sip(Host::Name(
HostName::new(Bytes::from_static(b"example.com")).expect("a valid host"),
))
}
fn answerable_invite() -> RequestBuilder {
RequestBuilder::new(Method::Invite, uri())
.header(HeaderName::Via, "SIP/2.0/UDP host;branch=z9hG4bK1")
.expect("valid Via")
.header(HeaderName::From, "<sip:a@example.com>;tag=1")
.expect("valid From")
.header(HeaderName::To, "<sip:b@example.com>")
.expect("valid To")
.header(HeaderName::CallId, "call@example.com")
.expect("valid Call-ID")
.cseq(1, &Method::Invite)
.expect("valid CSeq")
}
#[test]
fn crlf_injection_rejected_in_every_user_supplied_field() {
let payloads: &[&[u8]] = &[
b"value\r\nInjected: yes",
b"value\rInjected: yes",
b"value\nInjected: yes",
b"value\0truncated",
b"\r\n",
b"\n\n",
];
for payload in payloads {
let p = Bytes::copy_from_slice(payload);
assert!(
Header::build(HeaderName::Subject, p.clone()).is_err(),
"header value accepted {payload:?}"
);
assert!(
Header::build(HeaderName::Other(p.clone()), Bytes::from_static(b"x")).is_err(),
"header name accepted {payload:?}"
);
assert!(
RequestBuilder::new(Method::Options, uri())
.header(HeaderName::CallId, p.clone())
.is_err(),
"request builder accepted {payload:?}"
);
assert!(
ResponseBuilder::new(StatusCode::new(200).unwrap(), p.clone()).is_err(),
"reason phrase accepted {payload:?}"
);
assert!(
ResponseBuilder::new(StatusCode::new(200).unwrap(), "OK")
.unwrap()
.header(HeaderName::Server, p.clone())
.is_err(),
"response builder accepted {payload:?}"
);
assert!(
RequestBuilder::new(Method::Options, uri())
.cseq(1, &Method::Other(p.clone()))
.is_err(),
"CSeq method accepted {payload:?}"
);
assert!(
HostName::new(p.clone()).is_err(),
"host name accepted {payload:?}"
);
}
}
#[test]
fn host_names_are_validated_not_just_screened() {
assert!(HostName::new("example.com").is_ok());
assert!(HostName::new("host-5.sub.example.com").is_ok());
for bad in [
"",
"exa mple.com",
"host@evil.com",
"host;lr",
"<host>",
"host/path",
] {
assert!(HostName::new(bad).is_err(), "{bad:?} should be rejected");
}
}
#[test]
fn a_built_request_frames_correctly() {
let request = RequestBuilder::new(Method::Options, uri())
.header(HeaderName::CallId, "abc@example.com")
.unwrap()
.cseq(1, &Method::Options)
.unwrap()
.max_forwards(70)
.build();
let mut out = Vec::new();
request.write_to(&mut out);
let text = String::from_utf8_lossy(&out);
assert!(text.starts_with("OPTIONS sip:example.com SIP/2.0\r\n"));
assert!(text.contains("CSeq: 1 OPTIONS\r\n"));
assert!(text.contains("Content-Length: 0\r\n"));
assert!(text.ends_with("\r\n\r\n"));
}
#[test]
fn setting_a_body_sets_the_matching_content_length() {
let request = RequestBuilder::new(Method::Options, uri())
.body(Bytes::from_static(b"hello"))
.build();
assert_eq!(request.body().len(), 5);
assert_eq!(
request
.headers
.value(&HeaderName::ContentLength)
.as_deref()
.map(<[u8]>::to_vec),
Some(b"5".to_vec())
);
let request = RequestBuilder::new(Method::Options, uri())
.body(Bytes::from_static(b"hello"))
.body(Bytes::from_static(b"hi"))
.build();
assert_eq!(request.headers.count(&HeaderName::ContentLength), 1);
assert_eq!(
request
.headers
.value(&HeaderName::ContentLength)
.as_deref()
.map(<[u8]>::to_vec),
Some(b"2".to_vec())
);
}
#[test]
fn a_response_copies_the_via_stack_in_order() {
let request = RequestBuilder::new(Method::Invite, uri())
.header(HeaderName::Via, "SIP/2.0/UDP first;branch=z9hG4bK1")
.unwrap()
.header(HeaderName::Via, "SIP/2.0/UDP second;branch=z9hG4bK2")
.unwrap()
.header(HeaderName::From, "<sip:a@b>;tag=1")
.unwrap()
.header(HeaderName::To, "<sip:c@d>")
.unwrap()
.header(HeaderName::CallId, "x@y")
.unwrap()
.cseq(7, &Method::Invite)
.unwrap()
.build();
let response =
ResponseBuilder::to_request(&request, StatusCode::new(180).unwrap(), "Ringing")
.unwrap()
.build();
let vias: Vec<_> = response
.headers
.get_all(&HeaderName::Via)
.map(|h| h.value().to_vec())
.collect();
assert_eq!(
vias,
vec![
b"SIP/2.0/UDP first;branch=z9hG4bK1".to_vec(),
b"SIP/2.0/UDP second;branch=z9hG4bK2".to_vec(),
],
"the Via stack must be copied in order; a response walks it back"
);
assert_eq!(
response.headers.typed::<CSeq>().and_then(Result::ok),
Some(CSeq {
sequence: 7,
method: Method::Invite
})
);
}
#[test]
fn history_is_returned_in_responses_other_than_100() {
let request = answerable_invite()
.header(HeaderName::Supported, "histinfo")
.unwrap()
.build();
let trying = ResponseBuilder::to_request(&request, StatusCode::new(100).unwrap(), "Trying")
.unwrap()
.build();
assert!(trying.headers.get(&HeaderName::HistoryInfo).is_none());
let ringing =
ResponseBuilder::to_request(&request, StatusCode::new(180).unwrap(), "Ringing")
.unwrap()
.build();
assert_eq!(
ringing.headers.value(&HeaderName::HistoryInfo).as_deref(),
Some(&b"<sip:example.com>;index=1"[..])
);
}
#[test]
fn repeated_history_rows_are_one_ordered_cache() {
let request = answerable_invite()
.header(HeaderName::HistoryInfo, "<sip:first@example.com>;index=1")
.unwrap()
.header(
HeaderName::HistoryInfo,
"<sip:second@example.com>;index=1.1;mp=1",
)
.unwrap()
.build();
let response =
ResponseBuilder::to_request(&request, StatusCode::new(180).unwrap(), "Ringing")
.unwrap()
.build();
assert_eq!(
response.headers.value(&HeaderName::HistoryInfo).as_deref(),
Some(&b"<sip:first@example.com>;index=1, <sip:second@example.com>;index=1.1;mp=1"[..])
);
}
#[test]
fn history_privacy_anonymizes_a_response_cache() {
let request = answerable_invite()
.header(
HeaderName::HistoryInfo,
"<sip:alice@example.com?Reason=SIP%3Bcause%3D302>;index=1",
)
.unwrap()
.header(HeaderName::Privacy, "history")
.unwrap()
.build();
let response =
ResponseBuilder::to_request(&request, StatusCode::new(486).unwrap(), "Busy Here")
.unwrap()
.build();
assert_eq!(
response.headers.value(&HeaderName::HistoryInfo).as_deref(),
Some(&b"<sip:anonymous@anonymous.invalid>;index=1"[..])
);
}
#[test]
fn a_response_can_be_built_for_a_request_with_an_unparseable_header() {
use crate::{Limits, parse_datagram};
let text = "OPTIONS sip:a@b.com SIP/2.0\r\n\
Via: SIP/2.0/UDP h;branch=z9hG4bKx\r\n\
To: \"unterminated <sip:a@b.com>\r\n\
From: <sip:c@d>;tag=1\r\n\
Call-ID: x@y\r\n\
CSeq: 1 OPTIONS\r\n\
Content-Length: 0\r\n\r\n";
let msg = parse_datagram(Bytes::from(text), &Limits::datagram()).expect("frames");
let request = msg.as_request().expect("a request");
assert!(
request
.headers
.typed::<crate::headers::To>()
.is_some_and(|r| r.is_err())
);
let response =
ResponseBuilder::to_request(request, StatusCode::new(400).unwrap(), "Bad Request")
.expect("must still build")
.build();
let mut out = Vec::new();
response.write_to(&mut out);
assert!(String::from_utf8_lossy(&out).starts_with("SIP/2.0 400 Bad Request\r\n"));
}
}