use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use alkcall::channels::adapter::ChannelsAdapter;
use alkcall::channels::operations::{ChannelCore, OpenEstablisher, OpenHandler};
use alkcall::channels::policy::{ChannelLifecyclePolicy, NoCap};
use alkcall::core::auth::{AuthContext, Identity};
use alkcall::core::types::{Connection, ProtocolHandler};
use alkcall::registry::registration::OperationRegistry;
use alkcall::registry::spec::OperationSpec;
use axum::extract::ws::WebSocketUpgrade;
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use parking_lot::Mutex;
use super::byte_adapter::{
split_ws_to_bytes_idle_with_write, WsPumps, INBOUND_WS_FRAME_CAP, INBOUND_WS_MESSAGE_CAP,
};
#[derive(Clone, Default)]
pub struct WsSessions {
counter: Arc<AtomicU64>,
sessions: Arc<Mutex<HashMap<u64, Arc<WsPumps>>>>,
connections: Arc<Mutex<HashMap<u64, Arc<alkcall::protocol::connection::CallConnection>>>>,
}
pub const DEFAULT_WS_MAX_SESSIONS: usize = 64;
#[derive(Clone)]
pub struct SessionState {
registry: Arc<OperationRegistry>,
sessions: Arc<WsSessions>,
session_slots: Arc<tokio::sync::Semaphore>,
idle_timeout: Option<std::time::Duration>,
openable_alpns: Option<Arc<[OpenableAlpn]>>,
op_register_acl: alkcall::registry::spec::AccessControl,
}
impl SessionState {
pub(crate) fn from_registry(registry: &Arc<OperationRegistry>) -> Self {
Self {
registry: Arc::clone(registry),
sessions: Arc::new(WsSessions::new()),
session_slots: Arc::new(tokio::sync::Semaphore::new(DEFAULT_WS_MAX_SESSIONS)),
idle_timeout: Some(crate::websocket::DEFAULT_WS_IDLE_TIMEOUT),
openable_alpns: None,
op_register_acl: alkcall::registry::spec::AccessControl::default(),
}
}
pub(crate) fn new(
registry: Arc<OperationRegistry>,
sessions: Arc<WsSessions>,
session_slots: Arc<tokio::sync::Semaphore>,
idle_timeout: Option<std::time::Duration>,
openable_alpns: Option<Arc<[OpenableAlpn]>>,
op_register_acl: alkcall::registry::spec::AccessControl,
) -> Self {
Self {
registry,
sessions,
session_slots,
idle_timeout,
openable_alpns,
op_register_acl,
}
}
pub(crate) fn registry(&self) -> &Arc<OperationRegistry> {
&self.registry
}
pub(crate) fn sessions(&self) -> &Arc<WsSessions> {
&self.sessions
}
pub(crate) fn session_slots(&self) -> &Arc<tokio::sync::Semaphore> {
&self.session_slots
}
pub(crate) fn idle_timeout(&self) -> Option<std::time::Duration> {
self.idle_timeout
}
pub(crate) fn openable_alpns(&self) -> Option<Arc<[OpenableAlpn]>> {
self.openable_alpns.clone()
}
pub(crate) fn op_register_acl(&self) -> alkcall::registry::spec::AccessControl {
self.op_register_acl.clone()
}
}
impl axum::extract::FromRef<Arc<OperationRegistry>> for SessionState {
fn from_ref(registry: &Arc<OperationRegistry>) -> Self {
SessionState::from_registry(registry)
}
}
impl WsSessions {
pub fn new() -> Self {
Self::default()
}
pub fn abort(&self) {
for (_, pumps) in self.sessions.lock().drain() {
pumps.abort();
}
}
pub fn len(&self) -> usize {
self.sessions.lock().len()
}
pub fn is_empty(&self) -> bool {
self.sessions.lock().is_empty()
}
pub fn live_connections(&self) -> Vec<Arc<alkcall::protocol::connection::CallConnection>> {
self.connections.lock().values().cloned().collect()
}
pub fn live_connection_count(&self) -> usize {
self.connections.lock().len()
}
fn insert(&self, pumps: Arc<WsPumps>) -> u64 {
let id = self.counter.fetch_add(1, Ordering::Relaxed);
self.sessions.lock().insert(id, pumps);
id
}
fn remove(&self, id: u64) {
self.sessions.lock().remove(&id);
}
fn insert_connection(&self, conn: Arc<alkcall::protocol::connection::CallConnection>) -> u64 {
let id = self.counter.fetch_add(1, Ordering::Relaxed);
self.connections.lock().insert(id, conn);
id
}
fn remove_connection(&self, id: u64) {
self.connections.lock().remove(&id);
}
}
#[derive(Clone)]
pub struct OpenableAlpn {
pub spec: OperationSpec,
pub open_handler: OpenHandler,
pub establisher: Option<OpenEstablisher>,
pub establisher_timeout: Option<std::time::Duration>,
}
impl OpenableAlpn {
pub fn new(spec: OperationSpec, open_handler: OpenHandler) -> Self {
Self {
spec,
open_handler,
establisher: None,
establisher_timeout: None,
}
}
pub fn with_establisher(
mut self,
establisher: OpenEstablisher,
timeout: Option<std::time::Duration>,
) -> Self {
self.establisher = Some(establisher);
self.establisher_timeout = timeout;
self
}
}
#[allow(clippy::too_many_arguments)]
pub async fn run_channels_session(
socket: axum::extract::ws::WebSocket,
registry: Arc<OperationRegistry>,
identity: Identity,
policy: Arc<dyn ChannelLifecyclePolicy>,
sessions: Option<WsSessions>,
idle_timeout: Option<std::time::Duration>,
write_timeout: Option<std::time::Duration>,
openable_alpns: Option<Arc<[OpenableAlpn]>>,
op_register_acl: alkcall::registry::spec::AccessControl,
) {
let (byte_stream, pumps) =
split_ws_to_bytes_idle_with_write(socket, idle_timeout, write_timeout);
let pumps = Arc::new(pumps);
let _guard = sessions.as_ref().map(|s| {
let id = s.insert(Arc::clone(&pumps));
SessionGuard {
sessions: s.clone(),
id,
}
});
let conn = Connection::from_bidi(byte_stream, b"alk/channels".to_vec(), None);
let _ = conn.set_identity(identity.clone());
let adapter = ChannelsAdapter::new(
install_channel_zero(
registry,
sessions,
Arc::clone(&policy),
openable_alpns,
op_register_acl,
),
policy,
);
let auth = AuthContext {
identity: Some(identity),
alpn: b"alk/channels".to_vec(),
remote_addr: None,
tls_client_fingerprint: None,
};
if let Err(e) = ProtocolHandler::handle(&adapter, conn, &auth).await {
tracing::warn!(error = %e, "channels session ended");
}
}
struct SessionGuard {
sessions: WsSessions,
id: u64,
}
impl Drop for SessionGuard {
fn drop(&mut self) {
self.sessions.remove(self.id);
}
}
struct ConnectionGuard {
sessions: WsSessions,
id: u64,
}
impl Drop for ConnectionGuard {
fn drop(&mut self) {
self.sessions.remove_connection(self.id);
}
}
fn install_channel_zero(
registry: Arc<OperationRegistry>,
sessions: Option<WsSessions>,
policy: Arc<dyn ChannelLifecyclePolicy>,
openable_alpns: Option<Arc<[OpenableAlpn]>>,
op_register_acl: alkcall::registry::spec::AccessControl,
) -> alkcall::channels::adapter::InstallChannelZero {
Arc::new(move |manager, channel0_conn, auth| {
let registry = Arc::clone(®istry);
let sessions = sessions.clone();
let policy = Arc::clone(&policy);
let openable_alpns = openable_alpns.clone();
let op_register_acl = op_register_acl.clone();
tokio::spawn(async move {
if let Some(identity) = auth.identity.clone() {
let _ = channel0_conn.set_identity(identity);
}
let channel0_bidi = match channel0_conn.accept_bi().await {
Ok(s) => s,
Err(_) => return,
};
let (writer, reader) =
alkcall::protocol::connection::split_single_stream(channel0_bidi);
let call_connection = Arc::new(
alkcall::protocol::connection::CallConnection::new_single_stream(
channel0_conn,
Arc::clone(&writer),
),
);
let fork = Arc::new(registry.fork());
let register_result: Result<(), String> = (|| {
alkcall::channels::operations::ChannelOperations::new(
manager.clone(),
Arc::clone(&policy),
)
.register_on(&fork)?;
if let Some(openables) = openable_alpns.as_ref() {
let core = ChannelCore::new(manager, Arc::clone(&policy));
for openable in openables.iter() {
core.register_openable_with_establisher(
openable.spec.clone(),
openable.establisher.clone(),
Arc::clone(&openable.open_handler),
&fork,
auth.clone(),
openable.establisher_timeout,
)?;
}
}
alkcall::registry::discovery::install_bootstrap_discovery(&fork)?;
fork.register(alkcall::registry::registration::HandlerRegistration::new(
alkcall::registry::op_register::op_register_spec(op_register_acl),
alkcall::registry::registration::HandlerKind::Once(
alkcall::registry::op_register::op_register_handler(
Arc::clone(&call_connection),
Arc::clone(&fork),
),
),
alkcall::registry::registration::OperationProvenance::Local,
None,
None,
alkcall::core::types::Capabilities::new(),
))?;
Ok(())
})();
if let Err(e) = register_result {
tracing::error!(error = %e, "channel-0 session registry setup failed");
return;
}
let _conn_guard = sessions.as_ref().map(|sessions| {
let id = sessions.insert_connection(Arc::clone(&call_connection));
ConnectionGuard {
sessions: sessions.clone(),
id,
}
});
let dispatcher = alkcall::protocol::dispatch::Dispatcher::new(
fork,
std::sync::Arc::new(NoopProvider),
);
dispatcher
.run_loop_single_stream(call_connection, reader, writer)
.await;
})
})
}
#[cfg(all(test, feature = "server", feature = "wss"))]
pub(crate) fn adapter_install_channel_zero(
registry: Arc<OperationRegistry>,
) -> alkcall::channels::adapter::InstallChannelZero {
install_channel_zero(
registry,
None,
Arc::new(NoCap),
None,
alkcall::registry::spec::AccessControl::default(),
)
}
struct NoopProvider;
impl alkcall::core::auth::IdentityProvider for NoopProvider {
fn resolve_from_fingerprint(&self, _: &str) -> Option<Identity> {
None
}
fn resolve_from_token(&self, _: &alkcall::core::auth::AuthToken) -> Option<Identity> {
None
}
}
#[derive(Clone)]
pub struct ChannelsPolicy(pub Arc<dyn ChannelLifecyclePolicy>);
#[derive(Clone, Copy, Debug)]
pub struct WsTimeouts {
pub idle: Option<std::time::Duration>,
pub write: Option<std::time::Duration>,
}
#[derive(Clone)]
pub struct OpenableAlpns(pub Arc<[OpenableAlpn]>);
#[derive(Clone, Debug)]
pub struct OpRegisterAcl(pub alkcall::registry::spec::AccessControl);
#[derive(Clone)]
pub struct SessionSlots(pub Arc<tokio::sync::Semaphore>);
#[allow(clippy::too_many_arguments)]
pub async fn ws_upgrade_handler(
sessions: Option<axum::Extension<WsSessions>>,
axum::extract::State(state): axum::extract::State<SessionState>,
axum::Extension(identity): axum::Extension<Identity>,
policy: Option<axum::Extension<ChannelsPolicy>>,
timeouts: Option<axum::Extension<WsTimeouts>>,
openables: Option<axum::Extension<OpenableAlpns>>,
op_register_acl: Option<axum::Extension<OpRegisterAcl>>,
session_slots: Option<axum::Extension<SessionSlots>>,
ws_upgrade: WebSocketUpgrade,
) -> Response {
let session_slots = session_slots
.map(|axum::Extension(s)| s.0)
.unwrap_or_else(|| Arc::clone(state.session_slots()));
let Ok(permit) = session_slots.try_acquire_owned() else {
return (
StatusCode::SERVICE_UNAVAILABLE,
"503 Service Unavailable: WS session cap reached",
)
.into_response();
};
let sessions = Some(
sessions
.map(|axum::Extension(s)| s)
.unwrap_or_else(|| WsSessions::clone(state.sessions())),
);
let policy = policy
.map(|axum::Extension(p)| p.0)
.unwrap_or_else(|| Arc::new(NoCap));
let registry = Arc::clone(state.registry());
let idle_timeout = match timeouts {
Some(axum::Extension(t)) => t.idle,
None => state.idle_timeout(),
};
let write_timeout = timeouts.and_then(|t| t.write);
let openable_alpns = openables
.map(|axum::Extension(o)| o.0)
.or_else(|| state.openable_alpns());
let op_register_acl = op_register_acl
.map(|axum::Extension(a)| a.0)
.unwrap_or_else(|| state.op_register_acl());
ws_upgrade
.max_frame_size(INBOUND_WS_FRAME_CAP)
.max_message_size(INBOUND_WS_MESSAGE_CAP)
.on_upgrade(move |socket| async move {
let _permit = permit;
run_channels_session(
socket,
registry,
identity,
policy,
sessions,
idle_timeout,
Some(write_timeout.unwrap_or(crate::websocket::DEFAULT_WS_WRITE_TIMEOUT)),
openable_alpns,
op_register_acl,
)
.await
})
}
pub async fn ws_bearer_auth(
axum::extract::State(provider): axum::extract::State<
Arc<dyn alkcall::core::auth::IdentityProvider>,
>,
mut req: axum::http::Request<axum::body::Body>,
next: axum::middleware::Next,
) -> Response {
let identity = crate::server::auth::extract_bearer_identity(&req, provider.as_ref());
match identity {
Some(identity) => {
req.extensions_mut().insert(identity);
next.run(req).await
}
None => (StatusCode::UNAUTHORIZED, "401 Unauthorized").into_response(),
}
}
#[cfg(any(test, feature = "test-support"))]
pub mod test_support {
use alkcall::protocol::wire::EventEnvelope;
use futures::StreamExt;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode;
pub fn frame_channel0_chunk(envelope: &EventEnvelope) -> Vec<u8> {
let body = serde_json::to_vec(envelope).unwrap();
let mut out = Vec::with_capacity(8 + 4 + body.len());
out.extend_from_slice(&0u32.to_be_bytes());
out.extend_from_slice(&((body.len() + 4) as u32).to_be_bytes());
out.extend_from_slice(&(body.len() as u32).to_be_bytes());
out.extend_from_slice(&body);
out
}
#[derive(Default)]
pub struct ChunkAssembler {
buf: Vec<u8>,
}
impl ChunkAssembler {
pub fn new() -> Self {
Self::default()
}
pub fn push(&mut self, bytes: &[u8]) {
self.buf.extend_from_slice(bytes);
}
pub fn next_chunk(&mut self) -> Option<(u32, Vec<u8>)> {
if self.buf.len() < 8 {
return None;
}
let channel_id =
u32::from_be_bytes([self.buf[0], self.buf[1], self.buf[2], self.buf[3]]);
let len =
u32::from_be_bytes([self.buf[4], self.buf[5], self.buf[6], self.buf[7]]) as usize;
if self.buf.len() < 8 + len {
return None;
}
let payload = self.buf.drain(..8 + len).skip(8).collect();
Some((channel_id, payload))
}
}
#[derive(Default)]
pub struct FrameAssembler {
buf: Vec<u8>,
}
impl FrameAssembler {
pub fn new() -> Self {
Self::default()
}
pub fn push(&mut self, bytes: &[u8]) {
self.buf.extend_from_slice(bytes);
}
pub fn next_frame(&mut self) -> Option<EventEnvelope> {
if self.buf.len() < 4 {
return None;
}
let len =
u32::from_be_bytes([self.buf[0], self.buf[1], self.buf[2], self.buf[3]]) as usize;
if self.buf.len() < 4 + len {
return None;
}
let frame: Vec<u8> = self.buf.drain(..4 + len).collect();
serde_json::from_slice(&frame[4..]).ok()
}
}
pub struct WsClient {
sink: futures::stream::SplitSink<
tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
tokio_tungstenite::tungstenite::Message,
>,
stream: futures::stream::SplitStream<
tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
>,
}
impl WsClient {
pub async fn connect_authorized(url: &str, token: &str) -> Result<Self, String> {
let mut request = url
.into_client_request()
.map_err(|e| format!("bad url: {e}"))?;
request.headers_mut().insert(
http::header::AUTHORIZATION,
http::HeaderValue::from_str(&format!("Bearer {token}"))
.map_err(|e| format!("bad token: {e}"))?,
);
let (stream, _resp): (
tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
_,
) = tokio_tungstenite::connect_async(request)
.await
.map_err(|e| format!("connect failed: {e}"))?;
Ok(Self::from_stream(stream))
}
pub async fn connect_status(url: &str, token: Option<&str>) -> Option<u16> {
let mut request = url.into_client_request().ok()?;
if let Some(t) = token {
request.headers_mut().insert(
http::header::AUTHORIZATION,
http::HeaderValue::from_str(&format!("Bearer {t}")).ok()?,
);
}
type WsStream = tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>;
type WsConnectResult = Result<
(
WsStream,
tokio_tungstenite::tungstenite::http::Response<Option<Vec<u8>>>,
),
tokio_tungstenite::tungstenite::Error,
>;
let result: WsConnectResult = tokio_tungstenite::connect_async(request).await;
match result {
Ok((stream, resp)) => {
drop(stream);
Some(resp.status().as_u16())
}
Err(tokio_tungstenite::tungstenite::Error::Http(resp)) => {
Some(resp.status().as_u16())
}
Err(_) => None,
}
}
fn from_stream(
stream: tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
) -> Self {
let (sink, stream) = stream.split();
Self { sink, stream }
}
pub async fn send_binary(&mut self, bytes: Vec<u8>) {
self.send_binary_piece(&bytes).await;
}
pub async fn send_binary_piece(&mut self, bytes: &[u8]) {
use futures::SinkExt;
self.sink
.send(tokio_tungstenite::tungstenite::Message::Binary(
bytes.to_vec().into(),
))
.await
.unwrap();
}
pub async fn send_text(&mut self, text: &str) {
use futures::SinkExt;
self.sink
.send(tokio_tungstenite::tungstenite::Message::Text(
text.to_string().into(),
))
.await
.unwrap();
}
pub async fn next_binary(&mut self, timeout: std::time::Duration) -> Option<Vec<u8>> {
use futures::StreamExt;
loop {
match tokio::time::timeout(timeout, self.stream.next()).await {
Err(_) => return None,
Ok(None) => return None,
Ok(Some(Err(_))) => return None,
Ok(Some(Ok(m))) => match m {
tokio_tungstenite::tungstenite::Message::Binary(b) => {
return Some(b.to_vec())
}
tokio_tungstenite::tungstenite::Message::Close(_) => return None,
_ => continue,
},
}
}
}
pub async fn next_close(&mut self, timeout: std::time::Duration) -> Option<Option<u16>> {
use futures::StreamExt;
match tokio::time::timeout(timeout, self.stream.next()).await {
Err(_) => None,
Ok(None) => Some(None),
Ok(Some(Ok(tokio_tungstenite::tungstenite::Message::Close(cf)))) => {
Some(cf.map(|f| match f.code {
CloseCode::Error => 1011,
other => other.into(),
}))
}
Ok(Some(Ok(_))) => Box::pin(self.next_close(timeout)).await,
Ok(Some(Err(_))) => Some(None),
}
}
pub async fn close(&mut self) {
use futures::SinkExt;
let _ = self.sink.close().await;
}
}
}