use bytes::Bytes;
use futures_util::{
stream::{SplitSink, SplitStream},
SinkExt, StreamExt,
};
use serde::Serialize;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use tokio_tungstenite::{
connect_async,
tungstenite::{self, Message},
MaybeTlsStream, WebSocketStream,
};
use tracing::{debug, error, info, trace, warn};
use zeroize::Zeroizing;
use crate::ack::{self, IncomingBuffer, OutgoingBuffer, RetransmitAction};
use crate::binary_protocol::{
ClientMessage, ControlFlag, MessageType, PayloadType, MAX_PAYLOAD_SIZE,
};
use crate::errors::{Error, Result};
use crate::handshake::{
self, EncryptionChallengeRequest, HandshakeComplete, HandshakeHandler, HandshakeRequest,
};
use crate::metrics::{self, names};
use crate::session::{CloseReason, SessionCore};
const MESSAGE_SCHEMA_VERSION: &str = "1.0";
const RETRANSMIT_TICK: Duration = Duration::from_millis(100);
const MAX_RETRANSMIT_ATTEMPTS: u32 = 3000;
const MESSAGE_BUFFER_CAPACITY: usize = 10_000;
const WRITER_QUEUE_DEPTH: usize = 1024;
pub(crate) const COMMAND_QUEUE_DEPTH: usize = 256;
const CONNECT_ATTEMPTS: u32 = 3;
type WsStream = WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>;
type WsSink = SplitSink<WsStream, Message>;
type WsSource = SplitStream<WsStream>;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum EndpointPolicy {
#[default]
AwsOnly,
AllowAny,
}
impl EndpointPolicy {
pub(crate) fn validate(self, raw: &str) -> Result<()> {
let url = url::Url::parse(raw)
.map_err(|e| Error::Config(format!("stream URL is not a valid URL: {e}")))?;
if self == EndpointPolicy::AllowAny {
return match url.scheme() {
"ws" | "wss" => Ok(()),
other => Err(Error::Config(format!(
"stream URL scheme must be ws or wss, got {other}"
))),
};
}
if url.scheme() != "wss" {
return Err(Error::Config(format!(
"stream URL must use wss://, got {}://",
url.scheme()
)));
}
let host = url
.host_str()
.ok_or_else(|| Error::Config("stream URL has no host".into()))?;
let in_aws = host.ends_with(".amazonaws.com") || host.ends_with(".amazonaws.com.cn");
let is_ssm = host
.split('.')
.any(|label| label == "ssmmessages" || label.starts_with("ssmmessages-"));
if in_aws && is_ssm {
Ok(())
} else {
Err(Error::Config(format!(
"refusing to connect to {host}: not an AWS SSM messages endpoint. \
Set SessionConfig::endpoint_policy to EndpointPolicy::AllowAny to override."
)))
}
}
}
fn sanitize_url(raw: &str) -> String {
match url::Url::parse(raw) {
Ok(mut url) => {
let kept: Vec<(String, String)> = url
.query_pairs()
.filter(|(k, _)| !k.eq_ignore_ascii_case("tokenValue"))
.map(|(k, v)| (k.into_owned(), v.into_owned()))
.collect();
url.set_query(None);
if !kept.is_empty() {
url.query_pairs_mut().extend_pairs(kept);
}
url.into()
}
Err(_) => "<unparseable URL>".to_owned(),
}
}
#[derive(Serialize)]
#[serde(rename_all = "PascalCase")]
struct OpenDataChannelInput<'a> {
message_schema_version: &'a str,
request_id: String,
token_value: &'a str,
client_id: &'a str,
client_version: &'a str,
}
#[derive(Debug)]
pub(crate) enum Command {
Data(Bytes),
Control {
payload_type: PayloadType,
data: Bytes,
},
}
pub(crate) struct ConnectionParams {
pub session_id: String,
pub target: String,
pub stream_url: String,
pub token: Zeroizing<String>,
pub chunk_size: usize,
pub endpoint_policy: EndpointPolicy,
pub heartbeat_interval: Duration,
pub idle_timeout: Duration,
#[cfg(feature = "kms")]
pub kms: Option<aws_sdk_kms::Client>,
}
pub(crate) async fn connect(
core: Arc<SessionCore>,
params: ConnectionParams,
) -> Result<(mpsc::Sender<Command>, Vec<JoinHandle<()>>)> {
params.endpoint_policy.validate(¶ms.stream_url)?;
info!(
url = %sanitize_url(¶ms.stream_url),
"opening SSM data channel"
);
let ws = dial(¶ms.stream_url).await?;
let (mut sink, source) = ws.split();
let open = OpenDataChannelInput {
message_schema_version: MESSAGE_SCHEMA_VERSION,
request_id: uuid::Uuid::new_v4().to_string(),
token_value: ¶ms.token,
client_id: ¶ms.session_id,
client_version: handshake::CLIENT_PROTOCOL_VERSION,
};
let open_json = Zeroizing::new(serde_json::to_string(&open)?);
sink.send(Message::Text(open_json.as_str().into()))
.await
.map_err(|e| Error::transport(format!("failed to open the data channel: {e}")))?;
debug!("data channel open message sent");
let (writer_tx, writer_rx) = mpsc::channel::<Message>(WRITER_QUEUE_DEPTH);
let (command_tx, command_rx) = mpsc::channel::<Command>(COMMAND_QUEUE_DEPTH);
let shared = Arc::new(ChannelState {
core: Arc::clone(&core),
writer_tx: writer_tx.clone(),
outgoing: OutgoingBuffer::new(MESSAGE_BUFFER_CAPACITY, MAX_RETRANSMIT_ATTEMPTS),
sequence: tokio::sync::Mutex::new(0),
last_inbound: Mutexed::new(Instant::now()),
chunk_size: params.chunk_size.clamp(1, MAX_PAYLOAD_SIZE),
});
#[cfg(feature = "kms")]
let handshake_handler = match params.kms {
Some(client) => HandshakeHandler::with_kms(handshake::KmsContext {
client,
session_id: params.session_id.clone(),
target_id: params.target.clone(),
}),
None => HandshakeHandler::new(),
};
#[cfg(not(feature = "kms"))]
let handshake_handler = {
let _ = ¶ms.target;
HandshakeHandler::new()
};
let tasks = vec![
tokio::spawn(writer_task(sink, writer_rx, Arc::clone(&core))),
tokio::spawn(reader_task(source, Arc::clone(&shared), handshake_handler)),
tokio::spawn(command_task(command_rx, Arc::clone(&shared))),
tokio::spawn(retransmit_task(Arc::clone(&shared))),
tokio::spawn(heartbeat_task(
Arc::clone(&shared),
params.heartbeat_interval,
params.idle_timeout,
)),
];
Ok((command_tx, tasks))
}
async fn dial(url: &str) -> Result<WsStream> {
let mut backoff = Duration::from_millis(200);
for attempt in 1..=CONNECT_ATTEMPTS {
match connect_async(url).await {
Ok((ws, response)) => {
debug!(status = ?response.status(), attempt, "WebSocket connected");
return Ok(ws);
}
Err(e) => {
let retriable = matches!(e, tungstenite::Error::Io(_));
if !retriable || attempt == CONNECT_ATTEMPTS {
return Err(Error::transport(format!(
"could not open the SSM WebSocket: {e}"
)));
}
warn!(attempt, error = %e, "WebSocket dial failed; retrying");
tokio::time::sleep(backoff).await;
backoff *= 2;
}
}
}
unreachable!("the loop returns on the final attempt")
}
struct Mutexed<T>(std::sync::Mutex<T>);
impl<T> Mutexed<T> {
fn new(value: T) -> Self {
Self(std::sync::Mutex::new(value))
}
fn lock(&self) -> std::sync::MutexGuard<'_, T> {
self.0.lock().unwrap_or_else(|e| e.into_inner())
}
}
struct ChannelState {
core: Arc<SessionCore>,
writer_tx: mpsc::Sender<Message>,
outgoing: OutgoingBuffer,
sequence: tokio::sync::Mutex<i64>,
last_inbound: Mutexed<Instant>,
chunk_size: usize,
}
impl ChannelState {
fn mark_inbound(&self) {
*self.last_inbound.lock() = Instant::now();
}
fn idle_for(&self) -> Duration {
self.last_inbound.lock().elapsed()
}
async fn write(&self, message: Message) -> Result<()> {
self.writer_tx
.send(message)
.await
.map_err(|_| Error::transport("the data channel writer has stopped"))
}
fn try_write(&self, message: Message) -> bool {
self.writer_tx.try_send(message).is_ok()
}
}
async fn writer_task(mut sink: WsSink, mut rx: mpsc::Receiver<Message>, core: Arc<SessionCore>) {
loop {
tokio::select! {
biased;
() = core.closed() => break,
message = rx.recv() => {
let Some(message) = message else { break };
if let Err(e) = sink.send(message).await {
if core.is_closed() {
debug!(error = %e, "write failed during shutdown");
} else {
core.close(CloseReason::Transport(format!("WebSocket write failed: {e}")));
}
break;
}
}
}
}
let _ = sink.close().await;
debug!("writer task finished");
}
async fn heartbeat_task(state: Arc<ChannelState>, interval: Duration, idle_timeout: Duration) {
let mut ticker = tokio::time::interval_at(tokio::time::Instant::now() + interval, interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
tokio::select! {
biased;
() = state.core.closed() => break,
_ = ticker.tick() => {
let idle = state.idle_for();
if idle > idle_timeout {
error!(
idle_secs = idle.as_secs(),
timeout_secs = idle_timeout.as_secs(),
"no traffic from the SSM agent within the idle timeout; \
treating the connection as dead"
);
state.core.close(CloseReason::PeerUnresponsive { idle });
break;
}
trace!(idle_ms = idle.as_millis(), "sending keep-alive ping");
if !state.try_write(Message::Ping(Bytes::new())) {
debug!("writer queue full; skipping keep-alive ping");
}
}
}
}
debug!("heartbeat task finished");
}
async fn retransmit_task(state: Arc<ChannelState>) {
let mut ticker = tokio::time::interval(RETRANSMIT_TICK);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
tokio::select! {
biased;
() = state.core.closed() => break,
_ = ticker.tick() => match state.outgoing.poll_retransmit() {
RetransmitAction::Idle => {}
RetransmitAction::Resend { sequence, wire } => {
metrics::counter(names::RETRANSMISSIONS, 1);
tokio::select! {
biased;
() = state.core.closed() => break,
result = state.write(Message::Binary(wire)) => {
if result.is_err() {
debug!(sequence, "writer gone; stopping retransmission");
break;
}
}
}
}
RetransmitAction::GaveUp { sequence, attempts } => {
error!(
sequence,
attempts,
"the SSM agent never acknowledged a message; giving up on the session"
);
state.core.close(CloseReason::DeliveryFailed { sequence, attempts });
break;
}
},
}
}
debug!("retransmit task finished");
}
async fn command_task(mut rx: mpsc::Receiver<Command>, state: Arc<ChannelState>) {
loop {
let command = tokio::select! {
biased;
() = state.core.closed() => break,
command = rx.recv() => match command {
Some(c) => c,
None => break,
},
};
let permitted = tokio::select! {
biased;
() = state.core.closed() => false,
() = state.core.wait_sendable() => true,
};
if !permitted {
break;
}
let result = match command {
Command::Data(data) => send_stream_data(&state, data).await,
Command::Control { payload_type, data } => {
send_payload(&state, payload_type, data).await
}
};
if let Err(e) = result {
if state.core.is_closed() {
debug!(error = %e, "send failed during shutdown");
} else {
error!(error = %e, "failed to send on the data channel");
state.core.close(CloseReason::Transport(e.to_string()));
}
break;
}
}
debug!("command task finished");
}
async fn send_stream_data(state: &ChannelState, data: Bytes) -> Result<()> {
if data.is_empty() {
return Ok(());
}
let mut offset = 0;
while offset < data.len() {
let end = (offset + state.chunk_size).min(data.len());
send_payload(state, PayloadType::Output, data.slice(offset..end)).await?;
offset = end;
}
Ok(())
}
async fn send_payload(
state: &ChannelState,
payload_type: PayloadType,
payload: Bytes,
) -> Result<()> {
let payload = match (payload_type, state.core.crypto()) {
(PayloadType::Output, Some(crypto)) => crypto.encrypt(&payload)?,
_ => payload,
};
let mut sequence = state.sequence.lock().await;
let wire = ClientMessage::new(
MessageType::InputStreamData,
*sequence,
payload_type,
payload,
)
.serialize();
loop {
let acked = state.core.ack_notified();
if state.outgoing.track(wire.clone(), *sequence) {
*sequence += 1;
break;
}
debug!("outgoing buffer full; waiting for acknowledgements");
tokio::select! {
biased;
() = state.core.closed() => {
return Err(Error::SessionClosed(
"session closed while waiting for acknowledgements".into(),
))
}
() = acked => {}
}
}
metrics::counter(names::MESSAGES_SENT, 1);
metrics::counter(names::BYTES_SENT, wire.len() as u64);
state.write(Message::Binary(wire)).await
}
struct ReaderState {
handshake: HandshakeHandler,
expected_sequence: i64,
incoming: IncomingBuffer,
handshake_started: Option<Instant>,
}
async fn reader_task(mut source: WsSource, state: Arc<ChannelState>, handshake: HandshakeHandler) {
let mut reader = ReaderState {
handshake,
expected_sequence: 0,
incoming: IncomingBuffer::new(MESSAGE_BUFFER_CAPACITY),
handshake_started: None,
};
let reason = loop {
let frame = tokio::select! {
biased;
() = state.core.closed() => return,
frame = source.next() => frame,
};
state.mark_inbound();
match frame {
Some(Ok(Message::Binary(data))) => {
metrics::counter(names::MESSAGES_RECEIVED, 1);
metrics::counter(names::BYTES_RECEIVED, data.len() as u64);
let message = match ClientMessage::deserialize(data) {
Ok(m) => m,
Err(e) => {
break CloseReason::Protocol(e.to_string());
}
};
if let Err(e) = route(&state, &mut reader, message).await {
break CloseReason::Protocol(e.to_string());
}
}
Some(Ok(Message::Text(text))) => {
if let Some(reason) = handle_text(&state, text.as_str()) {
break reason;
}
}
Some(Ok(Message::Close(frame))) => {
info!(?frame, "gateway closed the WebSocket");
break CloseReason::AgentClosed {
exit_code: state.core.exit_code(),
detail: frame.and_then(|f| {
let reason = f.reason.trim().to_owned();
(!reason.is_empty()).then_some(reason)
}),
};
}
Some(Ok(Message::Ping(_) | Message::Pong(_) | Message::Frame(_))) => {}
Some(Err(e)) => break CloseReason::Transport(format!("WebSocket read failed: {e}")),
None => break CloseReason::Transport("WebSocket stream ended".into()),
}
};
state.core.close(reason);
debug!("reader task finished");
}
fn handle_text(state: &ChannelState, text: &str) -> Option<CloseReason> {
match text.trim() {
"start_publication" => {
debug!("gateway is ready to receive (text start_publication)");
state.core.set_sendable(true);
None
}
"pause_publication" => {
debug!("gateway asked us to pause sending (text pause_publication)");
state.core.set_sendable(false);
None
}
"channel_closed" => Some(CloseReason::AgentClosed {
exit_code: state.core.exit_code(),
detail: None,
}),
other => {
debug!(preview = %truncate(other, 200), "ignoring unrecognised text frame");
None
}
}
}
async fn route(
state: &ChannelState,
reader: &mut ReaderState,
message: ClientMessage,
) -> Result<()> {
match message.message_type {
MessageType::OutputStreamData => route_stream_data(state, reader, message).await,
MessageType::Acknowledge => {
let content = ack::parse_ack(&message)?;
metrics::counter(names::ACKS_RECEIVED, 1);
if state.outgoing.acknowledge(content.sequence_number) {
state.core.notify_ack();
let rtt = state.outgoing.rtt();
metrics::histogram(names::RTT_SECONDS, rtt.smoothed.as_secs_f64());
}
Ok(())
}
MessageType::StartPublication => {
debug!("gateway is ready to receive");
state.core.set_sendable(true);
Ok(())
}
MessageType::PausePublication => {
debug!("gateway asked us to pause sending");
state.core.set_sendable(false);
Ok(())
}
MessageType::ChannelClosed => {
let detail = channel_closed_detail(&message.payload);
info!(detail = ?detail, "agent closed the channel");
state.core.close(CloseReason::AgentClosed {
exit_code: state.core.exit_code(),
detail,
});
Ok(())
}
MessageType::InputStreamData => {
warn!("ignoring unexpected input_stream_data from the agent");
Ok(())
}
}
}
async fn route_stream_data(
state: &ChannelState,
reader: &mut ReaderState,
message: ClientMessage,
) -> Result<()> {
use std::cmp::Ordering as Ord;
match message.sequence_number.cmp(&reader.expected_sequence) {
Ord::Less => {
trace!(
sequence = message.sequence_number,
expected = reader.expected_sequence,
"dropping an already-processed message without acknowledging"
);
Ok(())
}
Ord::Greater => {
let sequence = message.sequence_number;
if reader.incoming.insert(message.clone()) {
acknowledge(state, &message, false);
trace!(
sequence,
expected = reader.expected_sequence,
buffered = reader.incoming.len(),
"buffered an out-of-order message"
);
} else {
warn!(
sequence,
"reorder buffer full; dropping without acknowledging"
);
}
Ok(())
}
Ord::Equal => {
acknowledge(state, &message, true);
deliver(state, reader, message).await?;
reader.expected_sequence += 1;
while let Some(buffered) = reader.incoming.take(reader.expected_sequence) {
deliver(state, reader, buffered).await?;
reader.expected_sequence += 1;
}
Ok(())
}
}
}
fn acknowledge(state: &ChannelState, message: &ClientMessage, sequential: bool) {
match ack::build_ack(message, sequential) {
Ok(reply) => {
if !state.try_write(Message::Binary(reply.serialize())) {
warn!(
sequence = message.sequence_number,
"writer queue full; dropped an acknowledgement"
);
}
}
Err(e) => error!(error = %e, "could not build an acknowledgement"),
}
}
async fn deliver(
state: &ChannelState,
reader: &mut ReaderState,
message: ClientMessage,
) -> Result<()> {
match message.payload_type {
PayloadType::Output | PayloadType::StdErr | PayloadType::Undefined => {
if !state.core.is_sendable() {
debug!("agent sent output without a handshake; enabling sending");
state.core.set_sendable(true);
}
if message.payload.is_empty() {
return Ok(());
}
let payload = decrypt_if_needed(state, &message)?;
state.core.emit_output(payload);
Ok(())
}
PayloadType::HandshakeRequest => {
let request: HandshakeRequest = serde_json::from_slice(&message.payload)
.map_err(|e| Error::protocol(format!("malformed HandshakeRequest: {e}")))?;
reader.handshake_started.get_or_insert_with(Instant::now);
let Some(response) = reader.handshake.on_request(request).await? else {
return Ok(()); };
let failed = reader.handshake.state() == handshake::HandshakeState::Failed;
let errors = response.errors.join("; ");
let payload = handshake::response_payload(&response)?;
send_payload(state, PayloadType::HandshakeResponse, payload).await?;
debug!("handshake response sent");
if failed {
return Err(Error::Unsupported(errors));
}
state
.core
.set_agent_version(reader.handshake.agent_version());
if let Some(crypto) = reader.handshake.crypto() {
info!("session encryption negotiated (AES-256-GCM)");
state.core.set_crypto(crypto);
}
Ok(())
}
PayloadType::HandshakeComplete => {
let complete: HandshakeComplete = serde_json::from_slice(&message.payload)
.map_err(|e| Error::protocol(format!("malformed HandshakeComplete: {e}")))?;
let banner = reader.handshake.on_complete(complete)?;
if let Some(started) = reader.handshake_started.take() {
metrics::histogram(names::HANDSHAKE_SECONDS, started.elapsed().as_secs_f64());
}
state.core.set_session_banner(banner);
state.core.set_sendable(true);
Ok(())
}
PayloadType::EncChallengeRequest => {
let request: EncryptionChallengeRequest = serde_json::from_slice(&message.payload)
.map_err(|e| Error::protocol(format!("malformed encryption challenge: {e}")))?;
let response = reader.handshake.answer_challenge(&request)?;
let payload = Bytes::from(serde_json::to_vec(&response)?);
send_payload(state, PayloadType::EncChallengeResponse, payload).await?;
debug!("answered the agent's encryption challenge");
Ok(())
}
PayloadType::ExitCode => {
let payload = decrypt_if_needed(state, &message)?;
let exit_code = std::str::from_utf8(&payload)
.ok()
.and_then(|s| s.trim().parse::<i32>().ok());
info!(?exit_code, "remote process exited");
state.core.set_exit_code(exit_code);
Ok(())
}
PayloadType::Flag => {
match ControlFlag::from_payload(&message.payload) {
Some(ControlFlag::TerminateSession) => {
info!("agent signalled session termination");
state.core.close(CloseReason::AgentClosed {
exit_code: state.core.exit_code(),
detail: None,
});
}
Some(flag) => debug!(?flag, "received a control flag"),
None => debug!("received an unrecognised control flag"),
}
Ok(())
}
PayloadType::Error => {
let text = String::from_utf8_lossy(&message.payload);
warn!(error = %truncate(&text, 500), "agent reported an error");
Ok(())
}
other => {
debug!(payload_type = ?other, "ignoring an unhandled payload type");
Ok(())
}
}
}
fn channel_closed_detail(payload: &[u8]) -> Option<String> {
#[derive(serde::Deserialize)]
struct ChannelClosed {
#[serde(rename = "Output")]
output: Option<String>,
}
let text = match serde_json::from_slice::<ChannelClosed>(payload) {
Ok(parsed) => parsed.output.unwrap_or_default(),
Err(_) => String::from_utf8_lossy(payload).into_owned(),
};
let trimmed = truncate(text.trim(), 500);
(!trimmed.is_empty()).then(|| trimmed.to_owned())
}
fn decrypt_if_needed(state: &ChannelState, message: &ClientMessage) -> Result<Bytes> {
match state.core.crypto() {
Some(crypto) if message.payload_type.is_encrypted_inbound() => {
crypto.decrypt(&message.payload)
}
_ => Ok(message.payload.clone()),
}
}
fn truncate(text: &str, max: usize) -> &str {
match text.char_indices().nth(max) {
Some((idx, _)) => &text[..idx],
None => text,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn aws_endpoints_are_accepted() {
for url in [
"wss://ssmmessages.us-east-1.amazonaws.com/v1/data-channel/abc",
"wss://ssmmessages.eu-central-1.amazonaws.com/v1/data-channel/abc?role=publish",
"wss://ssmmessages-fips.us-gov-west-1.amazonaws.com/v1/data-channel/abc",
"wss://ssmmessages.cn-north-1.amazonaws.com.cn/v1/data-channel/abc",
] {
EndpointPolicy::AwsOnly
.validate(url)
.unwrap_or_else(|e| panic!("{url} should be accepted: {e}"));
}
}
#[test]
fn vpc_endpoint_hostnames_are_accepted() {
EndpointPolicy::AwsOnly
.validate("wss://vpce-0abc1def.ssmmessages.eu-central-1.vpce.amazonaws.com/v1/x")
.expect("VPC endpoint hostnames must work");
}
#[test]
fn non_ssm_and_lookalike_hosts_are_rejected() {
for url in [
"wss://evil.com/v1/data-channel/abc",
"wss://ssmmessages.attacker.com/steal",
"wss://s3.us-east-1.amazonaws.com/bucket",
"wss://amazonaws.com.evil.com/fake",
"wss://evil-ssmmessages.us-east-1.amazonaws.com/v1/x",
] {
assert!(
EndpointPolicy::AwsOnly.validate(url).is_err(),
"{url} must be rejected"
);
}
}
#[test]
fn plaintext_websocket_is_rejected_under_the_default_policy() {
let err = EndpointPolicy::AwsOnly
.validate("ws://ssmmessages.us-east-1.amazonaws.com/v1/x")
.unwrap_err();
assert!(err.to_string().contains("wss://"), "{err}");
}
#[test]
fn allow_any_accepts_local_mocks_but_still_requires_websocket_scheme() {
assert!(EndpointPolicy::AllowAny
.validate("ws://127.0.0.1:9001/x")
.is_ok());
assert!(EndpointPolicy::AllowAny
.validate("wss://localhost/x")
.is_ok());
assert!(EndpointPolicy::AllowAny
.validate("http://127.0.0.1/x")
.is_err());
}
#[test]
fn sanitize_url_strips_the_token_and_keeps_everything_else() {
let raw = "wss://ssmmessages.us-east-1.amazonaws.com/v1/data-channel/abc\
?role=publish&tokenValue=SECRET&cell-number=7";
let clean = sanitize_url(raw);
assert!(!clean.contains("SECRET"), "{clean}");
assert!(!clean.contains("tokenValue"), "{clean}");
assert!(clean.contains("role=publish"), "{clean}");
assert!(clean.contains("cell-number=7"), "{clean}");
}
#[test]
fn sanitize_url_is_case_insensitive_about_the_token_parameter() {
let clean = sanitize_url("wss://h.amazonaws.com/?TokenValue=a&TOKENVALUE=b&tokenvalue=c");
for secret in ["a", "b", "c"] {
assert!(!clean.contains(&format!("={secret}")), "{clean}");
}
}
#[test]
fn sanitize_url_handles_garbage() {
assert_eq!(sanitize_url("not a url"), "<unparseable URL>");
}
#[test]
fn open_message_carries_the_token_not_the_url() {
let open = OpenDataChannelInput {
message_schema_version: MESSAGE_SCHEMA_VERSION,
request_id: "req-1".into(),
token_value: "SECRET-TOKEN",
client_id: "session-1",
client_version: handshake::CLIENT_PROTOCOL_VERSION,
};
let json = serde_json::to_string(&open).unwrap();
assert!(json.contains("\"TokenValue\":\"SECRET-TOKEN\""), "{json}");
assert!(json.contains("\"MessageSchemaVersion\":\"1.0\""), "{json}");
assert!(json.contains("\"ClientId\":\"session-1\""), "{json}");
}
#[test]
fn truncate_respects_char_boundaries() {
assert_eq!(truncate("hello", 3), "hel");
assert_eq!(truncate("hello", 99), "hello");
assert_eq!(truncate("äöü", 2), "äö");
}
}