#![allow(missing_docs)]
use std::fmt;
use rvoip_sip_core::types::headers::{HeaderName, TypedHeader};
use rvoip_sip_core::types::Method;
use super::options::{HeaderNameDiagnostic, MethodDiagnostic};
use super::ViolationReason;
#[non_exhaustive]
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum HeaderRole {
StackManaged,
MethodShaped { setter: &'static str },
ApplicationControlled,
}
#[derive(Clone)]
pub struct MissingRequiredHeader {
pub method: Method,
pub name: HeaderName,
pub reason: &'static str,
}
impl fmt::Debug for MissingRequiredHeader {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("MissingRequiredHeader")
.field("method", &MethodDiagnostic(&self.method))
.field("name", &HeaderNameDiagnostic(&self.name))
.field("reason", &self.reason)
.finish()
}
}
pub fn classify(method: Method, name: &HeaderName) -> HeaderRole {
let canonical_name = name.canonical_wire_name();
if is_always_stack_managed(&canonical_name) {
return HeaderRole::StackManaged;
}
if matches!(&canonical_name, HeaderName::Route) {
return HeaderRole::StackManaged;
}
if let Some(setter) = method_shaped_setter(method, &canonical_name) {
return HeaderRole::MethodShaped { setter };
}
HeaderRole::ApplicationControlled
}
pub fn forbidden_for_carry_through(name: &HeaderName) -> bool {
let canonical_name = name.canonical_wire_name();
is_always_stack_managed(&canonical_name) || matches!(canonical_name, HeaderName::Route)
}
pub fn validate_outbound(
method: Method,
headers: &[TypedHeader],
) -> Result<(), Vec<MissingRequiredHeader>> {
let mut violations = Vec::new();
for h in headers {
let name = h.name();
if matches!(classify(method.clone(), &name), HeaderRole::StackManaged) {
violations.push(MissingRequiredHeader {
method: method.clone(),
name,
reason: "stack-managed header staged in application extras — \
this name is owned by the dialog/transaction layer",
});
}
}
if violations.is_empty() {
Ok(())
} else {
Err(violations)
}
}
fn is_always_stack_managed(name: &HeaderName) -> bool {
matches!(
name,
HeaderName::CallId
| HeaderName::CSeq
| HeaderName::Via
| HeaderName::MaxForwards
| HeaderName::ContentLength
| HeaderName::RecordRoute
)
}
fn method_shaped_setter(method: Method, name: &HeaderName) -> Option<&'static str> {
use HeaderName as H;
use Method as M;
if matches!(name, H::Authorization | H::ProxyAuthorization) {
return Some("with_credentials");
}
match (method, name) {
(M::Invite | M::Register | M::Subscribe, H::Contact) => Some("with_contact_uri"),
(M::Register | M::Subscribe, H::Expires) => Some("with_expires"),
(M::Refer, H::ReferTo) => Some("refer(.., refer_to)"),
(M::Notify | M::Subscribe, H::Event) => Some("notify(.., event_package)"),
(M::Notify, H::SubscriptionState) => Some("with_subscription_state"),
_ => None,
}
}
pub fn role_to_violation(role: &HeaderRole) -> Option<ViolationReason> {
match role {
HeaderRole::StackManaged => Some(ViolationReason::StackManaged),
HeaderRole::MethodShaped { setter } => Some(ViolationReason::UseDedicatedSetter(setter)),
HeaderRole::ApplicationControlled => None,
}
}
#[cfg(test)]
mod diagnostic_tests {
use super::*;
#[test]
fn missing_required_header_debug_redacts_extension_spellings() {
const METHOD_CANARY: &str = "CUSTOM\r\nX-Method-Canary: exposed";
const HEADER_CANARY: &str = "X-Header-Canary\r\nInjected";
let violation = MissingRequiredHeader {
method: Method::Extension(METHOD_CANARY.to_string()),
name: HeaderName::Other(HEADER_CANARY.to_string()),
reason: "fixed-policy-class",
};
let rendered = format!("{violation:?}");
assert!(rendered.starts_with("MissingRequiredHeader"));
assert!(rendered.contains(&format!("value_len: {}", METHOD_CANARY.len())));
assert!(rendered.contains(&format!("name_len: {}", HEADER_CANARY.len())));
assert!(!rendered.contains(METHOD_CANARY));
assert!(!rendered.contains(HEADER_CANARY));
assert_eq!(
violation.method,
Method::Extension(METHOD_CANARY.to_string())
);
assert_eq!(violation.name, HeaderName::Other(HEADER_CANARY.to_string()));
}
#[test]
fn structural_other_aliases_cannot_bypass_stack_ownership() {
for alias in [
"call-ID",
"I",
"cSeQ",
"V",
"MAX-forwards",
"L",
"ROUTE",
"record-ROUTE",
] {
let name = HeaderName::Other(alias.into());
assert_eq!(classify(Method::Invite, &name), HeaderRole::StackManaged);
assert!(forbidden_for_carry_through(&name));
}
}
#[test]
fn credential_other_aliases_require_the_dedicated_setter() {
for alias in ["AUTHORIZATION", "proxy-AUTHORIZATION"] {
assert_eq!(
classify(Method::Invite, &HeaderName::Other(alias.into())),
HeaderRole::MethodShaped {
setter: "with_credentials"
}
);
}
}
}