use crate::{Error, SqlReadBytes};
use futures_util::io::AsyncReadExt;
const FED_AUTH_INFO_ID_STSURL: u8 = 0x01;
const FED_AUTH_INFO_ID_SPN: u8 = 0x02;
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct TokenFedAuthInfo {
pub sts_url: Option<String>,
pub spn: Option<String>,
}
impl TokenFedAuthInfo {
pub(crate) async fn decode<R>(src: &mut R) -> crate::Result<Self>
where
R: SqlReadBytes + Unpin,
{
let token_length = src.read_u32_le().await? as usize;
if token_length > super::MAX_TOKEN_BODY {
return Err(Error::Protocol(
format!("FEDAUTHINFO token length {token_length} exceeds the maximum").into(),
));
}
let mut body = vec![0u8; token_length];
src.read_exact(&mut body).await?;
Self::parse(&body)
}
fn parse(body: &[u8]) -> crate::Result<Self> {
let read_u32 = |buf: &[u8], at: usize| -> crate::Result<u32> {
buf.get(at..at + 4)
.map(|b| u32::from_le_bytes([b[0], b[1], b[2], b[3]]))
.ok_or_else(|| {
Error::Protocol("FEDAUTHINFO token truncated while reading a DWORD".into())
})
};
let count = read_u32(body, 0)? as usize;
let mut info = TokenFedAuthInfo::default();
for i in 0..count {
let opt = 4 + i * 9;
let id = *body.get(opt).ok_or_else(|| {
Error::Protocol("FEDAUTHINFO token truncated while reading an option id".into())
})?;
let data_len = read_u32(body, opt + 1)? as usize;
let data_offset = read_u32(body, opt + 5)? as usize;
let data = body
.get(data_offset..data_offset + data_len)
.ok_or_else(|| {
Error::Protocol("FEDAUTHINFO token data offset out of bounds".into())
})?;
if data_len & 1 != 0 {
return Err(Error::Protocol(
"FEDAUTHINFO token data is not valid UTF-16".into(),
));
}
let mut utf16 = Vec::with_capacity(data_len / 2);
let mut idx = 0;
while idx < data_len {
utf16.push(u16::from_le_bytes([data[idx], data[idx + 1]]));
idx += 2;
}
let value = String::from_utf16(&utf16).map_err(|_| {
Error::Protocol("FEDAUTHINFO token data is not valid UTF-16".into())
})?;
match id {
FED_AUTH_INFO_ID_STSURL => info.sts_url = Some(value),
FED_AUTH_INFO_ID_SPN => info.spn = Some(value),
_ => (),
}
}
Ok(info)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn utf16le(s: &str) -> Vec<u8> {
s.encode_utf16().flat_map(|u| u.to_le_bytes()).collect()
}
#[test]
fn parses_stsurl_and_spn() {
let sts = utf16le("https://login.microsoftonline.com/");
let spn = utf16le("https://database.windows.net/");
let count: u32 = 2;
let header_len = 4 + 2 * 9; let sts_offset = header_len;
let spn_offset = header_len + sts.len();
let mut body = Vec::new();
body.extend_from_slice(&count.to_le_bytes());
body.push(FED_AUTH_INFO_ID_STSURL);
body.extend_from_slice(&(sts.len() as u32).to_le_bytes());
body.extend_from_slice(&(sts_offset as u32).to_le_bytes());
body.push(FED_AUTH_INFO_ID_SPN);
body.extend_from_slice(&(spn.len() as u32).to_le_bytes());
body.extend_from_slice(&(spn_offset as u32).to_le_bytes());
body.extend_from_slice(&sts);
body.extend_from_slice(&spn);
let info = TokenFedAuthInfo::parse(&body).unwrap();
assert_eq!(
info.sts_url.as_deref(),
Some("https://login.microsoftonline.com/")
);
assert_eq!(info.spn.as_deref(), Some("https://database.windows.net/"));
}
#[test]
fn ignores_unknown_info_id() {
let count: u32 = 1;
let mut body = Vec::new();
body.extend_from_slice(&count.to_le_bytes());
body.push(0x7F); body.extend_from_slice(&0u32.to_le_bytes()); body.extend_from_slice(&13u32.to_le_bytes());
let info = TokenFedAuthInfo::parse(&body).unwrap();
assert_eq!(info, TokenFedAuthInfo::default());
}
#[test]
fn rejects_out_of_bounds_offset() {
let count: u32 = 1;
let mut body = Vec::new();
body.extend_from_slice(&count.to_le_bytes());
body.push(FED_AUTH_INFO_ID_STSURL);
body.extend_from_slice(&8u32.to_le_bytes()); body.extend_from_slice(&1000u32.to_le_bytes());
assert!(TokenFedAuthInfo::parse(&body).is_err());
}
}