use std::net::SocketAddr;
use std::sync::Arc;
use serde_json::json;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tonic::Request;
use crate::proto::udb::core::idp::services::v1 as idp_pb;
use idp_pb::identity_provider_service_server::IdentityProviderService;
use super::IdentityProviderServiceImpl;
#[cfg(test)]
use super::idp_saml_replay_rejected_status;
struct SamlHttpConfig {
addr: SocketAddr,
default_tenant: String,
default_provider: String,
acs_url: String,
}
impl SamlHttpConfig {
fn from_env() -> Option<Self> {
let raw_addr = std::env::var("UDB_SAML_HTTP_ADDR").ok()?;
let addr: SocketAddr = match raw_addr.trim().parse() {
Ok(addr) => addr,
Err(err) => {
tracing::warn!(value = %raw_addr, error = %err, "invalid UDB_SAML_HTTP_ADDR; SAML HTTP disabled");
return None;
}
};
let default_tenant = std::env::var("UDB_SAML_DEFAULT_TENANT")
.unwrap_or_default()
.trim()
.to_string();
let default_provider = std::env::var("UDB_SAML_DEFAULT_PROVIDER")
.unwrap_or_default()
.trim()
.to_string();
if default_tenant.is_empty() || default_provider.is_empty() {
tracing::warn!(
"UDB_SAML_HTTP_ADDR is set but UDB_SAML_DEFAULT_TENANT/UDB_SAML_DEFAULT_PROVIDER \
is empty; SAML HTTP refuses to start without a provider binding (no provider => \
no signing-cert trust anchor; fail closed)"
);
return None;
}
Some(Self {
addr,
default_tenant,
default_provider,
acs_url: std::env::var("UDB_SAML_ACS_URL")
.unwrap_or_default()
.trim()
.to_string(),
})
}
}
pub(crate) fn spawn_from_env_with_shutdown<F>(
service: Arc<IdentityProviderServiceImpl>,
shutdown: F,
) -> Option<tokio::task::JoinHandle<()>>
where
F: std::future::Future<Output = ()> + Send + 'static,
{
let cfg = SamlHttpConfig::from_env()?;
tracing::info!(addr = %cfg.addr, "SAML 2.0 Web SSO HTTP surface enabled");
Some(tokio::spawn(async move {
tokio::select! {
_ = serve(service, cfg) => {}
_ = shutdown => {
tracing::info!("SAML HTTP listener shutting down");
}
}
}))
}
async fn serve(service: Arc<IdentityProviderServiceImpl>, cfg: SamlHttpConfig) {
let listener = match tokio::net::TcpListener::bind(cfg.addr).await {
Ok(l) => l,
Err(err) => {
tracing::warn!(addr = %cfg.addr, error = %err, "SAML HTTP endpoint disabled");
return;
}
};
let cfg = Arc::new(cfg);
loop {
let Ok((mut socket, _peer)) = listener.accept().await else {
continue;
};
let service = service.clone();
let cfg = cfg.clone();
tokio::spawn(async move {
let response = match read_request(&mut socket).await {
Some(req) => dispatch(service.as_ref(), cfg.as_ref(), req).await,
None => http_response(400, "text/plain; charset=utf-8", "malformed HTTP request"),
};
let _ = socket.write_all(response.as_bytes()).await;
});
}
}
struct HttpRequest {
method: String,
path: String,
body: String,
}
async fn read_request(socket: &mut tokio::net::TcpStream) -> Option<HttpRequest> {
const MAX: usize = 1024 * 1024;
let mut buf: Vec<u8> = Vec::with_capacity(8192);
let mut chunk = [0u8; 8192];
let header_end = loop {
let n = tokio::time::timeout(std::time::Duration::from_secs(10), socket.read(&mut chunk))
.await
.ok()?
.ok()?;
if n == 0 {
break find_header_end(&buf);
}
buf.extend_from_slice(&chunk[..n]);
if let Some(end) = find_header_end(&buf) {
break Some(end);
}
if buf.len() > MAX {
return None;
}
}?;
let head = String::from_utf8_lossy(&buf[..header_end]).to_string();
let mut lines = head.lines();
let request_line = lines.next()?;
let mut parts = request_line.split_whitespace();
let method = parts.next()?.to_string();
let raw_target = parts.next()?.to_string();
let path = match raw_target.split_once('?') {
Some((p, _)) => p.to_string(),
None => raw_target,
};
let mut content_length = 0usize;
for line in lines {
if let Some((name, value)) = line.split_once(':') {
if name.trim().eq_ignore_ascii_case("content-length") {
content_length = value.trim().parse().unwrap_or(0);
}
}
}
let body_start = header_end + 4; let mut body_bytes: Vec<u8> = buf
.get(body_start..)
.map(|s| s.to_vec())
.unwrap_or_default();
while body_bytes.len() < content_length.min(MAX) {
let n = tokio::time::timeout(std::time::Duration::from_secs(10), socket.read(&mut chunk))
.await
.ok()?
.ok()?;
if n == 0 {
break;
}
body_bytes.extend_from_slice(&chunk[..n]);
}
let body = String::from_utf8_lossy(&body_bytes).to_string();
Some(HttpRequest { method, path, body })
}
fn find_header_end(buf: &[u8]) -> Option<usize> {
buf.windows(4).position(|w| w == b"\r\n\r\n")
}
async fn dispatch(
service: &IdentityProviderServiceImpl,
cfg: &SamlHttpConfig,
req: HttpRequest,
) -> String {
let Some(rest) = req.path.strip_prefix("/saml") else {
return http_response(404, "text/plain; charset=utf-8", "not a SAML endpoint");
};
let rest = rest.trim_start_matches('/');
let (tenant, provider, resource) = resolve_scope(rest, cfg);
match (req.method.as_str(), resource.as_str()) {
("GET", "metadata") => metadata_response(service, cfg, &tenant, &provider).await,
("POST", "acs") => acs_response(service, &tenant, &provider, &req.body).await,
("GET", "acs") => http_response(
405,
"text/plain; charset=utf-8",
"the ACS endpoint accepts only the SAML HTTP-POST binding (POST)",
),
_ => http_response(404, "text/plain; charset=utf-8", "unknown SAML endpoint"),
}
}
fn resolve_scope(rest: &str, cfg: &SamlHttpConfig) -> (String, String, String) {
let segs: Vec<&str> = rest.split('/').collect();
if segs.len() >= 4 && segs[0] == "t" && segs[2] == "p" {
let tenant = segs[1].to_string();
let provider = segs[3].to_string();
let resource = segs[4..].join("/");
return (tenant, provider, resource);
}
(
cfg.default_tenant.clone(),
cfg.default_provider.clone(),
rest.to_string(),
)
}
async fn acs_response(
service: &IdentityProviderServiceImpl,
tenant: &str,
provider: &str,
body: &str,
) -> String {
let Some(saml_response) = form_param(body, "SAMLResponse") else {
return http_response(
400,
"text/plain; charset=utf-8",
"missing SAMLResponse form field (HTTP-POST binding)",
);
};
let relay_state = form_param(body, "RelayState").unwrap_or_default();
let grpc = idp_pb::SamlAcsRequest {
provider_id: provider.to_string(),
tenant_id: tenant.to_string(),
saml_response,
relay_state,
context: None,
};
match service.saml_acs(Request::new(grpc)).await {
Ok(resp) => {
let r = resp.into_inner();
if !r.authenticated || !r.signature_verified {
return http_response(
401,
"application/json; charset=utf-8",
&json!({
"authenticated": false,
"signature_verified": r.signature_verified,
"detail": if r.detail.is_empty() {
"assertion rejected".to_string()
} else {
r.detail
},
})
.to_string(),
);
}
http_response(
200,
"application/json; charset=utf-8",
&json!({
"authenticated": true,
"signature_verified": true,
"subject": r.subject,
"user_id": r.user_id,
"email": r.email,
"email_verified": r.email_verified,
"groups": r.groups,
"roles": r.roles,
"assurance": r.assurance,
"attributes": serde_json::from_str::<serde_json::Value>(&r.attributes_json)
.unwrap_or(serde_json::Value::Null),
"detail": r.detail,
})
.to_string(),
)
}
Err(status) => status_response(status),
}
}
async fn metadata_response(
service: &IdentityProviderServiceImpl,
cfg: &SamlHttpConfig,
tenant: &str,
provider: &str,
) -> String {
let entity_id = match service
.get_provider(Request::new(idp_pb::GetProviderRequest {
provider_id: provider.to_string(),
tenant_id: tenant.to_string(),
}))
.await
{
Ok(resp) => resp
.into_inner()
.provider
.map(|p| p.entity_id)
.unwrap_or_default(),
Err(status) => return status_response(status),
};
let sp_entity_id = if entity_id.trim().is_empty() {
format!("urn:udb:sp:{tenant}")
} else {
entity_id
};
let acs_url = if cfg.acs_url.is_empty() {
"/saml/acs".to_string()
} else {
cfg.acs_url.clone()
};
let xml = format!(
"<?xml version=\"1.0\" encoding=\"UTF-8\"?>\
<md:EntityDescriptor xmlns:md=\"urn:oasis:names:tc:SAML:2.0:metadata\" entityID=\"{entity}\">\
<md:SPSSODescriptor AuthnRequestsSigned=\"false\" WantAssertionsSigned=\"true\" \
protocolSupportEnumeration=\"urn:oasis:names:tc:SAML:2.0:protocol\">\
<md:NameIDFormat>urn:oasis:names:tc:SAML:1.1:nameid-format:emailAddress</md:NameIDFormat>\
<md:AssertionConsumerService Binding=\"urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST\" \
Location=\"{acs}\" index=\"0\" isDefault=\"true\"/>\
</md:SPSSODescriptor></md:EntityDescriptor>",
entity = xml_escape(&sp_entity_id),
acs = xml_escape(&acs_url),
);
http_response(200, "application/samlmetadata+xml; charset=utf-8", &xml)
}
fn form_param(body: &str, key: &str) -> Option<String> {
for pair in body.split('&') {
if let Some((k, v)) = pair.split_once('=') {
if k == key {
let plus_decoded = v.replace('+', " ");
let decoded = urlencoding::decode(&plus_decoded)
.map(|c| c.into_owned())
.unwrap_or_else(|_| plus_decoded.clone());
let cleaned: String = decoded.chars().filter(|c| !c.is_whitespace()).collect();
if cleaned.is_empty() {
return None;
}
return Some(cleaned);
}
}
}
None
}
fn status_response(status: tonic::Status) -> String {
use tonic::Code;
let http = match status.code() {
Code::InvalidArgument | Code::FailedPrecondition => 400,
Code::Unauthenticated => 401,
Code::PermissionDenied => 403,
Code::NotFound => 404,
Code::AlreadyExists => 409,
Code::Unavailable => 503,
_ => 500,
};
http_response(
http,
"application/json; charset=utf-8",
&json!({ "authenticated": false, "detail": status.message() }).to_string(),
)
}
fn xml_escape(s: &str) -> String {
s.replace('&', "&")
.replace('<', "<")
.replace('>', ">")
.replace('"', """)
}
fn http_response(status: u16, content_type: &str, body: &str) -> String {
let reason = match status {
200 => "OK",
400 => "Bad Request",
401 => "Unauthorized",
403 => "Forbidden",
404 => "Not Found",
405 => "Method Not Allowed",
409 => "Conflict",
500 => "Internal Server Error",
503 => "Service Unavailable",
_ => "OK",
};
let mut head = format!(
"HTTP/1.1 {status} {reason}\r\ncontent-type: {content_type}\r\ncontent-length: {}\r\nconnection: close\r\n",
body.len()
);
if status == 401 {
head.push_str("www-authenticate: SAML\r\n");
}
if status == 405 {
head.push_str("allow: POST\r\n");
}
head.push_str("\r\n");
head.push_str(body);
head
}
#[cfg(test)]
mod saml_http_tests {
use super::*;
fn set_env(key: &str, value: Option<&str>) {
#[allow(unused_unsafe)]
unsafe {
match value {
Some(v) => std::env::set_var(key, v),
None => std::env::remove_var(key),
}
}
}
#[test]
fn round_trip_form_funnels_into_saml_acs_request() {
let raw_b64 = "PHNhbWxwOlJlc3BvbnNl+with/pad==";
let encoded = "PHNhbWxwOlJlc3BvbnNl%2Bwith%2Fpad%3D%3D";
let body = format!("SAMLResponse={encoded}&RelayState=app%2Fhome");
let saml_response = form_param(&body, "SAMLResponse").expect("SAMLResponse present");
let relay_state = form_param(&body, "RelayState").unwrap_or_default();
assert_eq!(saml_response, raw_b64);
assert_eq!(relay_state, "app/home");
let grpc = idp_pb::SamlAcsRequest {
provider_id: "okta".to_string(),
tenant_id: "acme".to_string(),
saml_response: saml_response.clone(),
relay_state,
context: None,
};
assert_eq!(grpc.saml_response, raw_b64);
assert_eq!(grpc.tenant_id, "acme");
assert_eq!(grpc.provider_id, "okta");
}
#[test]
fn tamper_response_is_transported_unvalidated_not_rejected_by_http() {
let tampered = "not-base64-!!!-tampered-bytes";
let body = format!("SAMLResponse={tampered}&RelayState=x");
assert_eq!(form_param(&body, "SAMLResponse").as_deref(), Some(tampered));
}
#[test]
fn empty_or_missing_saml_response_is_none() {
assert_eq!(form_param("RelayState=x", "SAMLResponse"), None);
assert_eq!(form_param("SAMLResponse=", "SAMLResponse"), None);
assert_eq!(form_param("SAMLResponse=%20%20", "SAMLResponse"), None);
}
#[test]
fn from_env_gating_disabled_failclosed_enabled() {
let prior_addr = std::env::var("UDB_SAML_HTTP_ADDR").ok();
let prior_tenant = std::env::var("UDB_SAML_DEFAULT_TENANT").ok();
let prior_provider = std::env::var("UDB_SAML_DEFAULT_PROVIDER").ok();
let prior_acs = std::env::var("UDB_SAML_ACS_URL").ok();
set_env("UDB_SAML_HTTP_ADDR", None);
set_env("UDB_SAML_DEFAULT_TENANT", Some("acme"));
set_env("UDB_SAML_DEFAULT_PROVIDER", Some("okta"));
assert!(
SamlHttpConfig::from_env().is_none(),
"unset UDB_SAML_HTTP_ADDR must keep the SAML HTTP surface OFF"
);
set_env("UDB_SAML_HTTP_ADDR", Some("127.0.0.1:0"));
set_env("UDB_SAML_DEFAULT_TENANT", None);
set_env("UDB_SAML_DEFAULT_PROVIDER", None);
assert!(
SamlHttpConfig::from_env().is_none(),
"missing provider binding must fail closed"
);
set_env("UDB_SAML_HTTP_ADDR", Some("127.0.0.1:0"));
set_env("UDB_SAML_DEFAULT_TENANT", Some("acme"));
set_env("UDB_SAML_DEFAULT_PROVIDER", Some("okta"));
set_env("UDB_SAML_ACS_URL", Some("https://sp.example.com/saml/acs"));
let cfg = SamlHttpConfig::from_env().expect("configured listener is enabled");
assert_eq!(cfg.default_tenant, "acme");
assert_eq!(cfg.default_provider, "okta");
assert_eq!(cfg.acs_url, "https://sp.example.com/saml/acs");
set_env("UDB_SAML_HTTP_ADDR", Some("not-an-addr"));
assert!(SamlHttpConfig::from_env().is_none());
set_env("UDB_SAML_HTTP_ADDR", prior_addr.as_deref());
set_env("UDB_SAML_DEFAULT_TENANT", prior_tenant.as_deref());
set_env("UDB_SAML_DEFAULT_PROVIDER", prior_provider.as_deref());
set_env("UDB_SAML_ACS_URL", prior_acs.as_deref());
}
#[test]
fn resolve_scope_explicit_and_default() {
let cfg = SamlHttpConfig {
addr: "127.0.0.1:0".parse().expect("addr"),
default_tenant: "acme".to_string(),
default_provider: "okta".to_string(),
acs_url: String::new(),
};
let (t, p, r) = resolve_scope("t/contoso/p/entra/acs", &cfg);
assert_eq!(
(t.as_str(), p.as_str(), r.as_str()),
("contoso", "entra", "acs")
);
let (t, p, r) = resolve_scope("metadata", &cfg);
assert_eq!(
(t.as_str(), p.as_str(), r.as_str()),
("acme", "okta", "metadata")
);
}
#[test]
fn status_response_never_authenticates() {
let body = status_response(super::idp_saml_replay_rejected_status());
assert!(body.starts_with("HTTP/1.1 403"));
assert!(
body.contains("\"authenticated\": false") || body.contains("\"authenticated\":false")
);
}
}