use std::sync::Arc;
use aion_core::Event;
use aion_proto::{StreamedEvent, SubscriptionRequest, WireError, subscription_request};
use futures::stream::BoxStream;
use futures::{SinkExt, StreamExt, stream};
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::HeaderValue;
use tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode;
use tokio_tungstenite::tungstenite::{self, Message};
use tokio_tungstenite::{Connector, MaybeTlsStream, WebSocketStream};
use crate::client::{ClientConfig, TlsOptions};
use crate::error::ClientError;
use crate::transport::contract::SubscriptionAttempt;
pub const EVENT_STREAM_PATH: &str = "/events/stream";
pub(crate) type WsStream = WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>;
pub async fn open_subscription(
config: &ClientConfig,
request: SubscriptionRequest,
resume_from_sequence: Option<u64>,
) -> Result<SubscriptionAttempt, ClientError> {
let url = stream_url(config)?;
let frame = subscription_frame(request, resume_from_sequence)?;
let mut upgrade = url.as_str().into_client_request().map_err(|source| {
ClientError::invalid_argument(format!(
"stream endpoint {url} is not a valid websocket URL: {source}"
))
})?;
apply_headers(&mut upgrade, config)?;
let connector = tls_connector(config.tls.as_ref())?;
let (mut socket, _response) =
tokio_tungstenite::connect_async_tls_with_config(upgrade, None, false, connector)
.await
.map_err(map_connect_error)?;
socket
.send(Message::Text(frame.into()))
.await
.map_err(|source| {
ClientError::unavailable(format!(
"websocket subscription frame send failed: {source}"
))
})?;
Ok(SubscriptionAttempt::new(socket_events(socket)))
}
pub(crate) fn stream_url(config: &ClientConfig) -> Result<String, ClientError> {
let Some(endpoint) = config.stream_endpoint.as_deref() else {
return Err(ClientError::invalid_argument(format!(
"no stream endpoint is configured; event subscriptions require \
ClientBuilder::with_stream_endpoint pointing at the server's \
{EVENT_STREAM_PATH} WebSocket URL (the HTTP/WebSocket listener \
is a separate address from the gRPC endpoint)"
)));
};
let Some((scheme, rest)) = endpoint.split_once("://") else {
return Err(ClientError::invalid_argument(format!(
"stream endpoint {endpoint} is not an absolute URL; expected a \
ws://, wss://, http://, or https:// address"
)));
};
match scheme {
"ws" | "wss" => Ok(endpoint.to_owned()),
"http" => Ok(format!("ws://{rest}")),
"https" => Ok(format!("wss://{rest}")),
other => Err(ClientError::invalid_argument(format!(
"cannot derive a websocket stream URL from a {other}:// endpoint; \
expected ws://, wss://, http://, or https://"
))),
}
}
pub(crate) fn tls_connector(tls: Option<&TlsOptions>) -> Result<Option<Connector>, ClientError> {
let Some(tls) = tls else {
return Ok(None);
};
let mut roots = rustls::RootCertStore {
roots: webpki_roots::TLS_SERVER_ROOTS.to_vec(),
};
if let Some(pem) = &tls.ca_certificate_pem {
let mut added = 0_usize;
for certificate in rustls_pemfile::certs(&mut pem.as_slice()) {
let certificate = certificate.map_err(|source| {
ClientError::invalid_argument(format!(
"TLS ca_certificate_pem is not parseable PEM: {source}"
))
})?;
roots.add(certificate).map_err(|source| {
ClientError::invalid_argument(format!(
"TLS ca_certificate_pem holds a certificate the trust store rejects: {source}"
))
})?;
added += 1;
}
if added == 0 {
return Err(ClientError::invalid_argument(
"TLS ca_certificate_pem contains no CA certificate",
));
}
}
let tls_config = rustls::ClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth();
Ok(Some(Connector::Rustls(Arc::new(tls_config))))
}
pub(crate) fn subscription_frame(
request: SubscriptionRequest,
resume_from_sequence: Option<u64>,
) -> Result<String, ClientError> {
let (key, subscription) = match request.subscription {
Some(subscription_request::Subscription::PerWorkflow(mut per_workflow)) => {
if let Some(cursor) = resume_from_sequence {
if cursor == 0 {
return Err(ClientError::invalid_argument(
"resume_from_seq must be >= 1 (the first sequence number wanted)",
));
}
per_workflow.resume_from_seq = Some(cursor);
}
("per_workflow", encode_subscription(&per_workflow)?)
}
Some(subscription_request::Subscription::Filtered(filtered)) => {
reject_live_only_cursor("filtered", resume_from_sequence)?;
("filtered", encode_subscription(&filtered)?)
}
Some(subscription_request::Subscription::Firehose(firehose)) => {
reject_live_only_cursor("firehose", resume_from_sequence)?;
("firehose", encode_subscription(&firehose)?)
}
Some(subscription_request::Subscription::Cluster(cluster)) => {
reject_live_only_cursor("cluster", resume_from_sequence)?;
("cluster", encode_subscription(&cluster)?)
}
Some(subscription_request::Subscription::Transcript(transcript)) => {
reject_live_only_cursor("transcript", resume_from_sequence)?;
("transcript", encode_subscription(&transcript)?)
}
None => {
return Err(ClientError::invalid_argument(
"subscription request is missing its subscription variant",
));
}
};
serde_json::to_string(&serde_json::json!({ key: subscription })).map_err(|source| {
ClientError::invalid_argument(format!("failed to encode subscription request: {source}"))
})
}
fn encode_subscription<T: serde::Serialize>(value: &T) -> Result<serde_json::Value, ClientError> {
serde_json::to_value(value).map_err(|source| {
ClientError::invalid_argument(format!("failed to encode subscription request: {source}"))
})
}
fn reject_live_only_cursor(kind: &str, cursor: Option<u64>) -> Result<(), ClientError> {
if cursor.is_some() {
return Err(ClientError::invalid_argument(format!(
"{kind} event streams are live-only by design; resume cursors are \
valid for per-workflow subscriptions only"
)));
}
Ok(())
}
pub(crate) fn apply_headers(
upgrade: &mut tungstenite::handshake::client::Request,
config: &ClientConfig,
) -> Result<(), ClientError> {
let headers = upgrade.headers_mut();
if let Some(auth) = &config.auth {
let value = HeaderValue::from_str(&format!("Bearer {}", auth.token()))
.map_err(|_| ClientError::invalid_argument("auth token is not a valid header value"))?;
headers.insert("authorization", value);
}
if let Some(subject) = &config.subject {
let value = HeaderValue::from_str(subject).map_err(|_| {
ClientError::invalid_argument("subject is not a valid x-aion-subject header value")
})?;
headers.insert("x-aion-subject", value);
}
if !config.authorized_namespaces.is_empty() {
let value =
HeaderValue::from_str(&config.authorized_namespaces.join(",")).map_err(|_| {
ClientError::invalid_argument(
"authorized namespaces are not a valid x-aion-namespaces header value",
)
})?;
headers.insert("x-aion-namespaces", value);
}
Ok(())
}
pub(crate) fn map_connect_error(error: tungstenite::Error) -> ClientError {
match error {
tungstenite::Error::Http(response)
if response.status() == tungstenite::http::StatusCode::UNAUTHORIZED =>
{
ClientError::unauthenticated("websocket upgrade was rejected with HTTP 401")
}
other => ClientError::unavailable(format!("websocket connect failed: {other}")),
}
}
fn socket_events(socket: WsStream) -> BoxStream<'static, Result<Event, ClientError>> {
stream::unfold(Some(socket), |state| async move {
let mut socket = state?;
loop {
return match socket.next().await {
None | Some(Err(tungstenite::Error::ConnectionClosed)) => None,
Some(Ok(Message::Text(text))) => match decode_frame(text.as_bytes()) {
Ok(event) => Some((Ok(event), Some(socket))),
Err(error) => Some((Err(error), None)),
},
Some(Ok(Message::Binary(bytes))) => match decode_frame(&bytes) {
Ok(event) => Some((Ok(event), Some(socket))),
Err(error) => Some((Err(error), None)),
},
Some(Ok(Message::Close(frame))) => match frame {
Some(frame) if frame.code == CloseCode::Normal => None,
Some(frame) => Some((
Err(ClientError::unavailable(format!(
"websocket closed abnormally ({} {})",
frame.code, frame.reason
))),
None,
)),
None => Some((
Err(ClientError::unavailable(
"websocket closed without a close frame",
)),
None,
)),
},
Some(Ok(Message::Ping(_) | Message::Pong(_) | Message::Frame(_))) => continue,
Some(Err(source)) => Some((
Err(ClientError::unavailable(format!(
"websocket transport failed: {source}"
))),
None,
)),
};
}
})
.boxed()
}
fn decode_frame(bytes: &[u8]) -> Result<Event, ClientError> {
#[derive(serde::Deserialize)]
struct ErrorFrame {
error: WireError,
}
if let Ok(frame) = serde_json::from_slice::<ErrorFrame>(bytes) {
return Err(ClientError::from_wire_error(frame.error));
}
let streamed = serde_json::from_slice::<StreamedEvent>(bytes).map_err(|source| {
ClientError::server(format!(
"event stream frame is neither a StreamedEvent nor an error frame: {source}"
))
})?;
streamed
.decode_event()
.map_err(ClientError::from_wire_error)
}
#[cfg(test)]
#[path = "ws_socket_tests.rs"]
mod socket_tests;
#[cfg(test)]
#[path = "ws_unit_tests.rs"]
mod unit_tests;