use std::borrow::Cow;
use std::ops::Range;
use bytes::Bytes;
use crate::error::{AddressEditError, HeaderError, WarningEditError};
use crate::headers::address::{AddressValueSpan, value_spans};
use crate::headers::grammar::is_token_char;
use crate::headers::warning::{WarningValueSpan, value_spans as warning_value_spans};
use crate::name::HeaderName;
use crate::uri::Uri;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum Method {
Invite,
Ack,
Bye,
Cancel,
Register,
Options,
Info,
Prack,
Update,
Subscribe,
Notify,
Refer,
Message,
Publish,
Other(Bytes),
}
impl Method {
#[must_use]
pub fn parse(raw: &Bytes) -> Self {
match raw.as_ref() {
b"INVITE" => Self::Invite,
b"ACK" => Self::Ack,
b"BYE" => Self::Bye,
b"CANCEL" => Self::Cancel,
b"REGISTER" => Self::Register,
b"OPTIONS" => Self::Options,
b"INFO" => Self::Info,
b"PRACK" => Self::Prack,
b"UPDATE" => Self::Update,
b"SUBSCRIBE" => Self::Subscribe,
b"NOTIFY" => Self::Notify,
b"REFER" => Self::Refer,
b"MESSAGE" => Self::Message,
b"PUBLISH" => Self::Publish,
_ => Self::Other(raw.clone()),
}
}
#[must_use]
pub fn as_bytes(&self) -> &[u8] {
match self {
Self::Invite => b"INVITE",
Self::Ack => b"ACK",
Self::Bye => b"BYE",
Self::Cancel => b"CANCEL",
Self::Register => b"REGISTER",
Self::Options => b"OPTIONS",
Self::Info => b"INFO",
Self::Prack => b"PRACK",
Self::Update => b"UPDATE",
Self::Subscribe => b"SUBSCRIBE",
Self::Notify => b"NOTIFY",
Self::Refer => b"REFER",
Self::Message => b"MESSAGE",
Self::Publish => b"PUBLISH",
Self::Other(raw) => raw,
}
}
}
impl std::fmt::Display for Method {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", String::from_utf8_lossy(self.as_bytes()))
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Version {
Sip20,
Other(Bytes),
}
impl Version {
#[must_use]
pub(crate) fn parse(raw: &Bytes) -> Self {
if raw.eq_ignore_ascii_case(b"SIP/2.0") {
Self::Sip20
} else {
Self::Other(raw.clone())
}
}
#[must_use]
pub fn as_bytes(&self) -> &[u8] {
match self {
Self::Sip20 => b"SIP/2.0",
Self::Other(raw) => raw,
}
}
#[must_use]
pub fn is_supported(&self) -> bool {
matches!(self, Self::Sip20)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct StatusCode(u16);
impl StatusCode {
#[must_use]
pub fn new(code: u16) -> Option<Self> {
(100..=699).contains(&code).then_some(Self(code))
}
#[must_use]
pub fn code(self) -> u16 {
self.0
}
#[must_use]
pub fn class(self) -> u16 {
self.0 / 100
}
#[must_use]
pub fn is_provisional(self) -> bool {
self.class() == 1
}
#[must_use]
pub fn is_final(self) -> bool {
!self.is_provisional()
}
#[must_use]
pub fn is_success(self) -> bool {
self.class() == 2
}
}
impl std::fmt::Display for StatusCode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.0)
}
}
#[derive(Debug, Clone)]
pub struct Header {
name: HeaderName,
repr: HeaderRepr,
}
#[derive(Debug, Clone)]
enum HeaderRepr {
Wire { line: Bytes, value_offset: usize },
Built { value: Bytes },
}
#[derive(Debug)]
struct AddressLayout {
spans: Vec<AddressValueSpan>,
source_map: Vec<Range<usize>>,
raw_len: usize,
}
#[derive(Debug)]
struct WarningLayout {
spans: Vec<WarningValueSpan>,
source_map: Vec<Range<usize>>,
raw_len: usize,
}
impl Header {
#[must_use]
pub(crate) fn new_unchecked(name: HeaderName, value: impl Into<Bytes>) -> Self {
Self {
name,
repr: HeaderRepr::Built {
value: value.into(),
},
}
}
pub(crate) fn from_wire(name: HeaderName, line: Bytes, value_offset: usize) -> Self {
Self {
name,
repr: HeaderRepr::Wire { line, value_offset },
}
}
#[must_use]
pub fn name(&self) -> &HeaderName {
&self.name
}
#[must_use]
pub fn raw_value(&self) -> &[u8] {
match &self.repr {
HeaderRepr::Wire { line, value_offset } => line.get(*value_offset..).unwrap_or(&[]),
HeaderRepr::Built { value } => value,
}
}
#[must_use]
pub fn value(&self) -> Cow<'_, [u8]> {
let raw = self.raw_value();
if raw.iter().any(|&b| b == b'\r' || b == b'\n') {
let mut out = Vec::with_capacity(raw.len());
let mut i = 0;
while let Some(&b) = raw.get(i) {
if b == b'\r' && raw.get(i + 1) == Some(&b'\n') {
let mut j = i + 2;
while matches!(raw.get(j), Some(b' ' | b'\t')) {
j += 1;
}
out.push(b' ');
i = j;
} else {
out.push(b);
i += 1;
}
}
Cow::Owned(trim(&out).to_vec())
} else {
Cow::Borrowed(trim(raw))
}
}
pub fn address_value_count(&self) -> Result<usize, AddressEditError> {
self.address_layout().map(|layout| layout.spans.len())
}
pub fn replace_address_uri(
&mut self,
value_index: usize,
uri: &Uri,
) -> Result<(), AddressEditError> {
let encoded = validate_replacement_uri(uri)?;
let layout = self.address_layout()?;
let span = layout
.spans
.get(value_index)
.ok_or(AddressEditError::IndexOutOfRange { index: value_index })?;
let raw_span =
project_range(&layout, &span.uri).ok_or_else(|| malformed_address(self.name()))?;
let rewritten = self
.with_value_span_replaced(&raw_span, &encoded)
.ok_or_else(|| malformed_address(self.name()))?;
let candidate = rewritten.address_layout()?;
if candidate.spans.len() != layout.spans.len() {
return Err(malformed_address(self.name()));
}
let candidate_span = candidate
.spans
.get(value_index)
.ok_or_else(|| malformed_address(self.name()))?;
let candidate_raw_span = project_range(&candidate, &candidate_span.uri)
.ok_or_else(|| malformed_address(self.name()))?;
if rewritten.raw_value().get(candidate_raw_span) != Some(encoded.as_ref()) {
return Err(malformed_address(self.name()));
}
*self = rewritten;
Ok(())
}
pub fn replace_address_presentation(
&mut self,
value_index: usize,
display_name: Option<&str>,
uri: &Uri,
) -> Result<(), AddressEditError> {
let encoded = encode_address_presentation(display_name, uri)?;
let layout = self.address_layout()?;
let span = layout
.spans
.get(value_index)
.ok_or(AddressEditError::IndexOutOfRange { index: value_index })?;
let raw_span = project_range(&layout, &span.presentation)
.ok_or_else(|| malformed_address(self.name()))?;
let rewritten = self
.with_value_span_replaced(&raw_span, &encoded)
.ok_or_else(|| malformed_address(self.name()))?;
let candidate = rewritten.address_layout()?;
if candidate.spans.len() != layout.spans.len() {
return Err(malformed_address(self.name()));
}
let candidate_span = candidate
.spans
.get(value_index)
.ok_or_else(|| malformed_address(self.name()))?;
let candidate_raw_span = project_range(&candidate, &candidate_span.presentation)
.ok_or_else(|| malformed_address(self.name()))?;
if rewritten.raw_value().get(candidate_raw_span) != Some(encoded.as_ref()) {
return Err(malformed_address(self.name()));
}
*self = rewritten;
Ok(())
}
pub fn warning_value_count(&self) -> Result<usize, WarningEditError> {
self.warning_layout().map(|layout| layout.spans.len())
}
pub fn replace_warning_agent_with_pseudonym(
&mut self,
value_index: usize,
pseudonym: &[u8],
) -> Result<(), WarningEditError> {
validate_warning_pseudonym(pseudonym)?;
let layout = self.warning_layout()?;
let span = layout
.spans
.get(value_index)
.ok_or(WarningEditError::IndexOutOfRange { index: value_index })?;
let raw_span = project_source_range(&layout.source_map, layout.raw_len, &span.agent)
.ok_or_else(malformed_warning)?;
let rewritten = self
.with_value_span_replaced(&raw_span, pseudonym)
.ok_or_else(malformed_warning)?;
let candidate = rewritten.warning_layout()?;
if candidate.spans.len() != layout.spans.len() {
return Err(malformed_warning());
}
let candidate_span = candidate
.spans
.get(value_index)
.ok_or_else(malformed_warning)?;
let candidate_raw_span = project_source_range(
&candidate.source_map,
candidate.raw_len,
&candidate_span.agent,
)
.ok_or_else(malformed_warning)?;
if rewritten.raw_value().get(candidate_raw_span) != Some(pseudonym) {
return Err(malformed_warning());
}
*self = rewritten;
Ok(())
}
pub fn without_address_value(
&self,
value_index: usize,
) -> Result<Option<Self>, AddressEditError> {
let layout = self.address_layout()?;
let selected = layout
.spans
.get(value_index)
.ok_or(AddressEditError::IndexOutOfRange { index: value_index })?;
if layout.spans.len() == 1 {
return Ok(None);
}
let unfolded = if let Some(next) = layout.spans.get(value_index.saturating_add(1)) {
selected.item.start..next.item.start
} else {
let previous = value_index
.checked_sub(1)
.and_then(|index| layout.spans.get(index))
.ok_or_else(|| malformed_address(self.name()))?;
previous.part.end..selected.item.end
};
let raw_span =
project_range(&layout, &unfolded).ok_or_else(|| malformed_address(self.name()))?;
self.with_value_span_replaced(&raw_span, &[])
.map(Some)
.ok_or_else(|| malformed_address(self.name()))
}
fn address_layout(&self) -> Result<AddressLayout, AddressEditError> {
let (header, list) = address_grammar(self.name())?;
let raw = self.raw_value();
let (unfolded, source_map) = unfold_with_source_map(raw);
let spans = value_spans(&unfolded, header, list).map_err(AddressEditError::Malformed)?;
Ok(AddressLayout {
spans,
source_map,
raw_len: raw.len(),
})
}
fn warning_layout(&self) -> Result<WarningLayout, WarningEditError> {
if self.name() != &HeaderName::Warning {
return Err(malformed_warning());
}
let raw = self.raw_value();
let (unfolded, source_map) = unfold_with_source_map(raw);
let spans = warning_value_spans(&unfolded).map_err(WarningEditError::Malformed)?;
Ok(WarningLayout {
spans,
source_map,
raw_len: raw.len(),
})
}
fn with_value_span_replaced(&self, span: &Range<usize>, replacement: &[u8]) -> Option<Self> {
let repr = match &self.repr {
HeaderRepr::Wire { line, value_offset } => {
let start = value_offset.checked_add(span.start)?;
let end = value_offset.checked_add(span.end)?;
HeaderRepr::Wire {
line: replace_byte_span(line, &(start..end), replacement)?,
value_offset: *value_offset,
}
}
HeaderRepr::Built { value } => HeaderRepr::Built {
value: replace_byte_span(value, span, replacement)?,
},
};
Some(Self {
name: self.name.clone(),
repr,
})
}
pub fn write_to(&self, out: &mut Vec<u8>) {
match &self.repr {
HeaderRepr::Wire { line, .. } => out.extend_from_slice(line),
HeaderRepr::Built { value } => {
out.extend_from_slice(self.name.canonical());
out.extend_from_slice(b": ");
out.extend_from_slice(value);
}
}
}
}
fn trim(mut b: &[u8]) -> &[u8] {
while let Some((first, rest)) = b.split_first() {
if matches!(first, b' ' | b'\t') {
b = rest;
} else {
break;
}
}
while let Some((last, rest)) = b.split_last() {
if matches!(last, b' ' | b'\t') {
b = rest;
} else {
break;
}
}
b
}
#[derive(Debug, Clone, Default)]
pub struct Headers {
entries: Vec<Header>,
}
impl Headers {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn len(&self) -> usize {
self.entries.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn push(&mut self, header: Header) {
self.entries.push(header);
}
pub fn push_front(&mut self, header: Header) {
self.entries.insert(0, header);
}
pub fn iter(&self) -> impl Iterator<Item = &Header> {
self.entries.iter()
}
#[must_use]
pub fn get(&self, name: &HeaderName) -> Option<&Header> {
self.entries.iter().find(|h| h.name() == name)
}
pub fn get_all<'a>(&'a self, name: &'a HeaderName) -> impl Iterator<Item = &'a Header> {
self.entries.iter().filter(move |h| h.name() == name)
}
#[must_use]
pub fn count(&self, name: &HeaderName) -> usize {
self.entries.iter().filter(|h| h.name() == name).count()
}
pub fn remove_all(&mut self, name: &HeaderName) -> usize {
let before = self.entries.len();
self.entries.retain(|h| h.name() != name);
before - self.entries.len()
}
pub fn remove_first(&mut self, name: &HeaderName) -> Option<Header> {
let index = self.entries.iter().position(|h| h.name() == name)?;
Some(self.entries.remove(index))
}
pub fn insert(&mut self, index: usize, header: Header) {
let index = index.min(self.entries.len());
self.entries.insert(index, header);
}
pub fn retain(&mut self, f: impl FnMut(&Header) -> bool) {
self.entries.retain(f);
}
pub fn replace_address_uri(
&mut self,
name: &HeaderName,
value_index: usize,
uri: &Uri,
) -> Result<(), AddressEditError> {
address_grammar(name)?;
validate_replacement_uri(uri)?;
let rows = self.address_rows(name)?;
let (entry_index, row_index) = locate_address_value(&rows, value_index)?;
let header = self
.entries
.get_mut(entry_index)
.ok_or(AddressEditError::IndexOutOfRange { index: value_index })?;
header.replace_address_uri(row_index, uri)
}
pub fn replace_address_presentation(
&mut self,
name: &HeaderName,
value_index: usize,
display_name: Option<&str>,
uri: &Uri,
) -> Result<(), AddressEditError> {
address_grammar(name)?;
encode_address_presentation(display_name, uri)?;
let rows = self.address_rows(name)?;
let (entry_index, row_index) = locate_address_value(&rows, value_index)?;
let header = self
.entries
.get_mut(entry_index)
.ok_or(AddressEditError::IndexOutOfRange { index: value_index })?;
header.replace_address_presentation(row_index, display_name, uri)
}
pub fn replace_warning_agent_with_pseudonym(
&mut self,
value_index: usize,
pseudonym: &[u8],
) -> Result<(), WarningEditError> {
validate_warning_pseudonym(pseudonym)?;
let rows = self.warning_rows()?;
let (entry_index, row_index) = locate_flattened_value(&rows, value_index)
.map_err(|()| WarningEditError::IndexOutOfRange { index: value_index })?;
let header = self
.entries
.get_mut(entry_index)
.ok_or(WarningEditError::IndexOutOfRange { index: value_index })?;
header.replace_warning_agent_with_pseudonym(row_index, pseudonym)
}
pub fn remove_address_value(
&mut self,
name: &HeaderName,
value_index: usize,
) -> Result<(), AddressEditError> {
address_grammar(name)?;
let rows = self.address_rows(name)?;
let (entry_index, row_index) = locate_address_value(&rows, value_index)?;
let replacement = self
.entries
.get(entry_index)
.ok_or(AddressEditError::IndexOutOfRange { index: value_index })?
.without_address_value(row_index)?;
if let Some(header) = replacement {
let slot = self
.entries
.get_mut(entry_index)
.ok_or(AddressEditError::IndexOutOfRange { index: value_index })?;
*slot = header;
} else {
self.entries.remove(entry_index);
}
Ok(())
}
fn address_rows(&self, name: &HeaderName) -> Result<Vec<(usize, usize)>, AddressEditError> {
self.entries
.iter()
.enumerate()
.filter(|(_, header)| header.name() == name)
.map(|(index, header)| header.address_value_count().map(|count| (index, count)))
.collect()
}
fn warning_rows(&self) -> Result<Vec<(usize, usize)>, WarningEditError> {
self.entries
.iter()
.enumerate()
.filter(|(_, header)| header.name() == &HeaderName::Warning)
.map(|(index, header)| header.warning_value_count().map(|count| (index, count)))
.collect()
}
#[must_use]
pub fn value(&self, name: &HeaderName) -> Option<Cow<'_, [u8]>> {
self.get(name).map(Header::value)
}
pub fn write_to(&self, out: &mut Vec<u8>) {
for h in &self.entries {
h.write_to(out);
out.extend_from_slice(b"\r\n");
}
}
}
fn address_grammar(name: &HeaderName) -> Result<(&'static str, bool), AddressEditError> {
match name {
HeaderName::From => Ok(("From", false)),
HeaderName::To => Ok(("To", false)),
HeaderName::Contact => Ok(("Contact", true)),
HeaderName::Route => Ok(("Route", true)),
HeaderName::RecordRoute => Ok(("Record-Route", true)),
HeaderName::Path => Ok(("Path", true)),
HeaderName::ServiceRoute => Ok(("Service-Route", true)),
HeaderName::PAssertedIdentity => Ok(("P-Asserted-Identity", true)),
HeaderName::PPreferredIdentity => Ok(("P-Preferred-Identity", true)),
_ => Err(AddressEditError::UnsupportedHeader),
}
}
fn malformed_address(name: &HeaderName) -> AddressEditError {
let header = address_grammar(name).map_or("address", |(header, _)| header);
AddressEditError::Malformed(HeaderError::Syntax { header })
}
fn malformed_warning() -> WarningEditError {
WarningEditError::Malformed(HeaderError::Syntax { header: "Warning" })
}
fn validate_warning_pseudonym(pseudonym: &[u8]) -> Result<(), WarningEditError> {
if pseudonym.is_empty() || !pseudonym.iter().copied().all(is_token_char) {
return Err(WarningEditError::InvalidPseudonym);
}
Ok(())
}
fn validate_replacement_uri(uri: &Uri) -> Result<Bytes, AddressEditError> {
let encoded = uri.to_bytes();
Uri::parse(encoded.clone()).map_err(AddressEditError::InvalidUri)?;
Ok(encoded)
}
fn encode_address_presentation(
display_name: Option<&str>,
uri: &Uri,
) -> Result<Bytes, AddressEditError> {
let uri = validate_replacement_uri(uri)?;
let mut encoded = Vec::new();
if let Some(display_name) = display_name {
if display_name
.as_bytes()
.iter()
.any(|byte| *byte < 0x20 || *byte == 0x7f)
{
return Err(AddressEditError::InvalidDisplayName);
}
encoded.push(b'"');
for byte in display_name.as_bytes() {
if matches!(byte, b'"' | b'\\') {
encoded.push(b'\\');
}
encoded.push(*byte);
}
encoded.extend_from_slice(b"\" ");
}
encoded.push(b'<');
encoded.extend_from_slice(&uri);
encoded.push(b'>');
Ok(Bytes::from(encoded))
}
fn locate_address_value(
rows: &[(usize, usize)],
value_index: usize,
) -> Result<(usize, usize), AddressEditError> {
locate_flattened_value(rows, value_index)
.map_err(|()| AddressEditError::IndexOutOfRange { index: value_index })
}
fn locate_flattened_value(
rows: &[(usize, usize)],
value_index: usize,
) -> Result<(usize, usize), ()> {
let mut first = 0usize;
for &(entry_index, count) in rows {
let end = first.checked_add(count).ok_or(())?;
if value_index < end {
return Ok((entry_index, value_index - first));
}
first = end;
}
Err(())
}
fn unfold_with_source_map(raw: &[u8]) -> (Vec<u8>, Vec<Range<usize>>) {
let mut unfolded = Vec::with_capacity(raw.len());
let mut source_map = Vec::with_capacity(raw.len());
let mut i = 0usize;
while let Some(&byte) = raw.get(i) {
if byte == b'\r'
&& raw.get(i + 1) == Some(&b'\n')
&& matches!(raw.get(i + 2), Some(b' ' | b'\t'))
{
let mut end = i + 2;
while matches!(raw.get(end), Some(b' ' | b'\t')) {
end += 1;
}
unfolded.push(b' ');
source_map.push(i..end);
i = end;
} else {
unfolded.push(byte);
source_map.push(i..i + 1);
i += 1;
}
}
(unfolded, source_map)
}
fn project_range(layout: &AddressLayout, span: &Range<usize>) -> Option<Range<usize>> {
project_source_range(&layout.source_map, layout.raw_len, span)
}
fn project_source_range(
source_map: &[Range<usize>],
raw_len: usize,
span: &Range<usize>,
) -> Option<Range<usize>> {
if span.start > span.end || span.end > source_map.len() {
return None;
}
let start = source_boundary(source_map, raw_len, span.start)?;
let end = source_boundary(source_map, raw_len, span.end)?;
(start <= end).then_some(start..end)
}
fn source_boundary(source_map: &[Range<usize>], raw_len: usize, position: usize) -> Option<usize> {
if position == source_map.len() {
Some(raw_len)
} else {
source_map.get(position).map(|source| source.start)
}
}
#[derive(Debug, Clone)]
pub struct Request {
pub method: Method,
pub uri: Uri,
pub version: Version,
pub headers: Headers,
body: Bytes,
raw_start_line: Option<Bytes>,
raw_uri_span: Option<Range<usize>>,
}
#[derive(Debug, Clone)]
pub struct Response {
pub version: Version,
pub status: StatusCode,
pub reason: Bytes,
pub headers: Headers,
body: Bytes,
raw_start_line: Option<Bytes>,
}
#[derive(Debug, Clone)]
pub enum Message {
Request(Request),
Response(Response),
}
impl Request {
pub(crate) fn from_wire(
method: Method,
uri: Uri,
version: Version,
raw_start_line: Bytes,
raw_uri_span: Range<usize>,
headers: Headers,
body: Bytes,
) -> Self {
Self {
method,
uri,
version,
headers,
body,
raw_start_line: Some(raw_start_line),
raw_uri_span: Some(raw_uri_span),
}
}
pub fn set_uri(&mut self, uri: Uri) -> Result<(), crate::error::UriError> {
let encoded = uri.to_bytes();
Uri::parse(encoded.clone())?;
let rewritten = match (&self.raw_start_line, &self.raw_uri_span) {
(Some(raw), Some(span)) => Some(
replace_byte_span(raw, span, &encoded)
.ok_or(crate::error::UriError::RetainedSpan)?,
),
(None, None) => None,
_ => return Err(crate::error::UriError::RetainedSpan),
};
if let Some(raw) = rewritten {
let start = self
.raw_uri_span
.as_ref()
.map(|span| span.start)
.ok_or(crate::error::UriError::RetainedSpan)?;
let end = start
.checked_add(encoded.len())
.ok_or(crate::error::UriError::RetainedSpan)?;
self.raw_start_line = Some(raw);
self.raw_uri_span = Some(start..end);
} else {
self.raw_start_line = None;
self.raw_uri_span = None;
}
self.uri = uri;
Ok(())
}
#[must_use]
pub fn new(method: Method, uri: Uri) -> Self {
Self {
method,
uri,
version: Version::Sip20,
headers: Headers::new(),
body: Bytes::new(),
raw_start_line: None,
raw_uri_span: None,
}
}
#[must_use]
pub fn body(&self) -> &Bytes {
&self.body
}
pub fn set_body(&mut self, body: Bytes) {
self.body = body;
}
pub fn write_to(&self, out: &mut Vec<u8>) {
if let Some(raw) = &self.raw_start_line {
out.extend_from_slice(raw);
} else {
out.extend_from_slice(self.method.as_bytes());
out.push(b' ');
self.uri.write_to(out);
out.push(b' ');
out.extend_from_slice(self.version.as_bytes());
}
out.extend_from_slice(b"\r\n");
self.headers.write_to(out);
out.extend_from_slice(b"\r\n");
out.extend_from_slice(&self.body);
}
}
fn replace_byte_span(source: &Bytes, span: &Range<usize>, replacement: &[u8]) -> Option<Bytes> {
let prefix = source.get(..span.start)?;
let suffix = source.get(span.end..)?;
let capacity = prefix
.len()
.checked_add(replacement.len())?
.checked_add(suffix.len())?;
let mut out = Vec::with_capacity(capacity);
out.extend_from_slice(prefix);
out.extend_from_slice(replacement);
out.extend_from_slice(suffix);
Some(Bytes::from(out))
}
impl Response {
pub(crate) fn from_wire(
version: Version,
status: StatusCode,
reason: Bytes,
raw_start_line: Bytes,
headers: Headers,
body: Bytes,
) -> Self {
Self {
version,
status,
reason,
headers,
body,
raw_start_line: Some(raw_start_line),
}
}
#[must_use]
pub fn new(status: StatusCode, reason: impl Into<Bytes>) -> Self {
Self {
version: Version::Sip20,
status,
reason: reason.into(),
headers: Headers::new(),
body: Bytes::new(),
raw_start_line: None,
}
}
#[must_use]
pub fn body(&self) -> &Bytes {
&self.body
}
pub fn set_body(&mut self, body: Bytes) {
self.body = body;
}
pub fn write_to(&self, out: &mut Vec<u8>) {
if let Some(raw) = &self.raw_start_line {
out.extend_from_slice(raw);
} else {
out.extend_from_slice(self.version.as_bytes());
out.push(b' ');
out.extend_from_slice(self.status.to_string().as_bytes());
out.push(b' ');
out.extend_from_slice(&self.reason);
}
out.extend_from_slice(b"\r\n");
self.headers.write_to(out);
out.extend_from_slice(b"\r\n");
out.extend_from_slice(&self.body);
}
}
impl Message {
#[must_use]
pub fn headers(&self) -> &Headers {
match self {
Self::Request(r) => &r.headers,
Self::Response(r) => &r.headers,
}
}
pub fn headers_mut(&mut self) -> &mut Headers {
match self {
Self::Request(r) => &mut r.headers,
Self::Response(r) => &mut r.headers,
}
}
#[must_use]
pub fn body(&self) -> &Bytes {
match self {
Self::Request(r) => r.body(),
Self::Response(r) => r.body(),
}
}
#[must_use]
pub fn as_request(&self) -> Option<&Request> {
match self {
Self::Request(r) => Some(r),
Self::Response(_) => None,
}
}
#[must_use]
pub fn as_response(&self) -> Option<&Response> {
match self {
Self::Response(r) => Some(r),
Self::Request(_) => None,
}
}
pub fn write_to(&self, out: &mut Vec<u8>) {
match self {
Self::Request(r) => r.write_to(out),
Self::Response(r) => r.write_to(out),
}
}
#[must_use]
pub fn to_bytes(&self) -> Bytes {
let mut out = Vec::new();
self.write_to(&mut out);
Bytes::from(out)
}
}
pub trait TypedHeader: Sized {
const NAME: HeaderName;
const VALIDATE_LIST: bool = false;
fn decode(value: &[u8]) -> Result<Self, HeaderError>;
fn decode_list(value: &[u8]) -> Result<Vec<Self>, HeaderError> {
Self::decode(value).map(|one| vec![one])
}
fn validate_list(_values: &[&Self]) -> Result<(), HeaderError> {
Ok(())
}
}
struct TypedAll<'a, H: TypedHeader> {
entries: std::slice::Iter<'a, Header>,
row: std::vec::IntoIter<H>,
validated: Option<std::vec::IntoIter<Result<H, HeaderError>>>,
}
impl<'a, H: TypedHeader> TypedAll<'a, H> {
fn new(headers: &'a Headers) -> Self {
let validated = H::VALIDATE_LIST.then(|| {
let mut decoded: Vec<Result<H, HeaderError>> = headers
.entries
.iter()
.filter(|header| header.name() == &H::NAME)
.flat_map(|header| match H::decode_list(&header.value()) {
Ok(values) => values.into_iter().map(Ok).collect::<Vec<_>>(),
Err(error) => vec![Err(error)],
})
.collect();
let decode_error = decoded.iter().find_map(|result| match result {
Ok(_) => None,
Err(error) => Some(error.clone()),
});
if let Some(error) = decode_error {
decoded = vec![Err(error)];
} else if !decoded.is_empty() {
let values: Vec<&H> = decoded
.iter()
.filter_map(|result| result.as_ref().ok())
.collect();
if let Err(error) = H::validate_list(&values) {
decoded = vec![Err(error)];
}
}
decoded.into_iter()
});
Self {
entries: headers.entries.iter(),
row: Vec::new().into_iter(),
validated,
}
}
}
impl<H: TypedHeader> Iterator for TypedAll<'_, H> {
type Item = Result<H, HeaderError>;
fn next(&mut self) -> Option<Self::Item> {
if let Some(validated) = &mut self.validated {
return validated.next();
}
loop {
if let Some(value) = self.row.next() {
return Some(Ok(value));
}
let header = self.entries.find(|header| header.name() == &H::NAME)?;
match H::decode_list(&header.value()) {
Ok(values) => self.row = values.into_iter(),
Err(error) => return Some(Err(error)),
}
}
}
}
impl Headers {
#[must_use]
pub fn typed<H: TypedHeader>(&self) -> Option<Result<H, HeaderError>> {
self.get(&H::NAME).map(|h| H::decode(&h.value()))
}
pub fn typed_all<'a, H: TypedHeader + 'a>(
&'a self,
) -> impl Iterator<Item = Result<H, HeaderError>> + 'a {
TypedAll::new(self)
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
clippy::indexing_slicing
)]
mod tests {
use super::*;
#[test]
fn unfolding_collapses_continuations_to_a_single_space() {
let line = Bytes::from_static(b"Subject: one\r\n two\r\n\tthree");
let h = Header::from_wire(HeaderName::Subject, line, 9);
assert_eq!(h.value().as_ref(), b"one two three");
assert_eq!(h.raw_value(), b"one\r\n two\r\n\tthree");
}
#[test]
fn unfolded_value_borrows_when_there_is_no_folding() {
let line = Bytes::from_static(b"Subject: plain");
let h = Header::from_wire(HeaderName::Subject, line, 9);
assert!(matches!(h.value(), Cow::Borrowed(_)));
}
#[test]
fn status_code_range_is_enforced() {
assert!(StatusCode::new(99).is_none());
assert!(StatusCode::new(700).is_none());
assert_eq!(StatusCode::new(200).map(StatusCode::code), Some(200));
assert!(StatusCode::new(180).unwrap().is_provisional());
assert!(StatusCode::new(200).unwrap().is_success());
assert!(StatusCode::new(486).unwrap().is_final());
}
#[test]
fn methods_compare_case_sensitively() {
assert_ne!(
Method::parse(&Bytes::from_static(b"Invite")),
Method::Invite
);
assert_eq!(
Method::parse(&Bytes::from_static(b"INVITE")),
Method::Invite
);
}
#[test]
fn remove_first_takes_only_the_topmost_via() {
let mut headers = Headers::new();
for value in [&b"first"[..], b"second", b"third"] {
headers.push(Header::new_unchecked(
HeaderName::Via,
Bytes::copy_from_slice(value),
));
}
headers.insert(
1,
Header::new_unchecked(HeaderName::Route, Bytes::from_static(b"r")),
);
let taken = headers.remove_first(&HeaderName::Via).expect("a Via");
assert_eq!(taken.value().as_ref(), b"first");
assert_eq!(
headers
.get_all(&HeaderName::Via)
.map(|h| h.value().to_vec())
.collect::<Vec<_>>(),
vec![b"second".to_vec(), b"third".to_vec()],
"the remaining Vias keep their order"
);
assert_eq!(
headers.iter().map(|h| h.name().clone()).collect::<Vec<_>>(),
vec![HeaderName::Route, HeaderName::Via, HeaderName::Via],
"and every other header stays where it was"
);
}
#[test]
fn remove_first_on_a_name_that_is_absent_yields_nothing_and_changes_nothing() {
let mut headers = Headers::new();
headers.push(Header::new_unchecked(
HeaderName::Via,
Bytes::from_static(b"v"),
));
assert!(headers.remove_first(&HeaderName::Route).is_none());
assert_eq!(headers.len(), 1);
}
#[test]
fn inserting_past_the_end_appends_rather_than_panicking() {
let mut headers = Headers::new();
headers.push(Header::new_unchecked(
HeaderName::Via,
Bytes::from_static(b"v"),
));
headers.insert(
9999,
Header::new_unchecked(HeaderName::Route, Bytes::from_static(b"r")),
);
assert_eq!(headers.len(), 2);
assert_eq!(
headers.iter().last().map(|h| h.name().clone()),
Some(HeaderName::Route)
);
}
#[test]
fn insert_places_a_header_at_an_absolute_position() {
let mut headers = Headers::new();
for name in [HeaderName::Via, HeaderName::To, HeaderName::From] {
headers.push(Header::new_unchecked(name, Bytes::from_static(b"x")));
}
headers.insert(
1,
Header::new_unchecked(HeaderName::RecordRoute, Bytes::from_static(b"rr")),
);
assert_eq!(
headers.iter().map(|h| h.name().clone()).collect::<Vec<_>>(),
vec![
HeaderName::Via,
HeaderName::RecordRoute,
HeaderName::To,
HeaderName::From
]
);
headers.insert(
0,
Header::new_unchecked(HeaderName::Via, Bytes::from_static(b"newest")),
);
assert_eq!(headers.value(&HeaderName::Via).unwrap().as_ref(), b"newest");
}
#[test]
fn retain_filters_in_place_and_keeps_order() {
let mut headers = Headers::new();
for (name, value) in [
(HeaderName::Via, &b"keep"[..]),
(HeaderName::Route, b"drop"),
(HeaderName::Via, b"drop"),
(HeaderName::To, b"keep"),
] {
headers.push(Header::new_unchecked(name, Bytes::copy_from_slice(value)));
}
headers.retain(|header| header.value().as_ref() == b"keep");
assert_eq!(
headers.iter().map(|h| h.name().clone()).collect::<Vec<_>>(),
vec![HeaderName::Via, HeaderName::To]
);
}
#[test]
fn header_order_is_preserved_including_duplicates() {
let mut headers = Headers::new();
headers.push(Header::new_unchecked(
HeaderName::Via,
Bytes::from_static(b"first"),
));
headers.push(Header::new_unchecked(
HeaderName::Route,
Bytes::from_static(b"r"),
));
headers.push(Header::new_unchecked(
HeaderName::Via,
Bytes::from_static(b"second"),
));
let vias: Vec<_> = headers
.get_all(&HeaderName::Via)
.map(|h| h.value().to_vec())
.collect();
assert_eq!(vias, vec![b"first".to_vec(), b"second".to_vec()]);
assert_eq!(headers.count(&HeaderName::Via), 2);
headers.push_front(Header::new_unchecked(
HeaderName::Via,
Bytes::from_static(b"newest"),
));
assert_eq!(headers.value(&HeaderName::Via).unwrap().as_ref(), b"newest");
}
#[test]
fn typed_all_yields_each_element_of_a_comma_separated_row() {
use crate::headers::Contact;
let mut headers = Headers::new();
headers.push(Header::new_unchecked(
HeaderName::Contact,
Bytes::from_static(b"<sip:a@b.com>, <sip:c@d.com>"),
));
headers.push(Header::new_unchecked(
HeaderName::Contact,
Bytes::from_static(b"<sip:e@f.org>"),
));
let contacts: Vec<Contact> = headers
.typed_all::<Contact>()
.collect::<Result<_, _>>()
.unwrap();
let uris: Vec<_> = contacts.iter().map(|c| c.uri.to_bytes()).collect();
assert_eq!(
uris,
vec![
Bytes::from_static(b"sip:a@b.com"),
Bytes::from_static(b"sip:c@d.com"),
Bytes::from_static(b"sip:e@f.org"),
]
);
}
#[test]
fn typed_all_keeps_unconstrained_headers_lazy() {
use std::sync::atomic::{AtomicUsize, Ordering};
static DECODES: AtomicUsize = AtomicUsize::new(0);
struct CountingSubject;
impl TypedHeader for CountingSubject {
const NAME: HeaderName = HeaderName::Subject;
fn decode(_value: &[u8]) -> Result<Self, HeaderError> {
DECODES.fetch_add(1, Ordering::SeqCst);
Ok(Self)
}
}
DECODES.store(0, Ordering::SeqCst);
let mut headers = Headers::new();
headers.push(Header::new_unchecked(
HeaderName::Subject,
Bytes::from_static(b"first"),
));
headers.push(Header::new_unchecked(
HeaderName::Subject,
Bytes::from_static(b"second"),
));
let mut values = headers.typed_all::<CountingSubject>();
assert_eq!(DECODES.load(Ordering::SeqCst), 0);
assert!(values.next().is_some_and(|value| value.is_ok()));
assert_eq!(DECODES.load(Ordering::SeqCst), 1);
drop(values);
assert_eq!(DECODES.load(Ordering::SeqCst), 1);
}
#[test]
fn built_headers_serialize_canonically() {
let mut headers = Headers::new();
headers.push(Header::new_unchecked(
HeaderName::MaxForwards,
Bytes::from_static(b"70"),
));
let mut out = Vec::new();
headers.write_to(&mut out);
assert_eq!(out, b"Max-Forwards: 70\r\n");
}
#[test]
fn wire_headers_serialize_verbatim() {
let line = Bytes::from_static(b"MaX-fOrWaRdS : 0068");
let h = Header::from_wire(HeaderName::MaxForwards, line.clone(), 17);
let mut out = Vec::new();
h.write_to(&mut out);
assert_eq!(out, line);
assert_eq!(h.value().as_ref(), b"0068");
}
}