use std::num::NonZeroU32;
use std::path::Path;
use std::sync::Arc;
use futures_util::stream::{SplitSink, SplitStream};
use futures_util::{SinkExt, StreamExt};
use governor::clock::DefaultClock;
use governor::state::{InMemoryState, NotKeyed};
use governor::{Quota, RateLimiter};
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::net::TcpListener;
use tokio::sync::{mpsc, Semaphore};
use tokio_rustls::TlsAcceptor;
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::{accept_async, WebSocketStream};
use crate::core::{BrokerCore, PushItem};
use crate::error::{BrokerError, BrokerResult};
use crate::proto::{BrokerFrame, ClientFrame};
type DeliverLimiter = RateLimiter<NotKeyed, InMemoryState, DefaultClock>;
#[derive(Debug, Clone, Copy)]
pub struct BrokerLimits {
pub max_connections: usize,
pub messages_per_second: NonZeroU32,
pub message_burst: NonZeroU32,
}
impl Default for BrokerLimits {
fn default() -> Self {
Self {
max_connections: 1000,
messages_per_second: NonZeroU32::new(500).expect("500 is non-zero"),
message_burst: NonZeroU32::new(1000).expect("1000 is non-zero"),
}
}
}
pub struct BrokerServer {
core: Arc<BrokerCore>,
token: String,
limits: BrokerLimits,
connection_slots: Arc<Semaphore>,
tls: Option<TlsAcceptor>,
}
impl BrokerServer {
pub fn new(core: Arc<BrokerCore>, token: impl Into<String>) -> Self {
Self::with_limits(core, token, BrokerLimits::default())
}
pub fn with_limits(
core: Arc<BrokerCore>,
token: impl Into<String>,
limits: BrokerLimits,
) -> Self {
Self {
core,
token: token.into(),
connection_slots: Arc::new(Semaphore::new(limits.max_connections)),
limits,
tls: None,
}
}
pub fn with_tls(mut self, cert_file: &Path, key_file: &Path) -> BrokerResult<Self> {
let server_config = bamboo_subagent::transport::build_server_config(cert_file, key_file)
.map_err(BrokerError::Tls)?;
self.tls = Some(TlsAcceptor::from(Arc::new(server_config)));
Ok(self)
}
pub fn is_tls(&self) -> bool {
self.tls.is_some()
}
pub async fn serve(self: Arc<Self>, listener: TcpListener) -> BrokerResult<()> {
loop {
let (stream, peer) = listener
.accept()
.await
.map_err(|e| BrokerError::Transport(format!("accept: {e}")))?;
let server = Arc::clone(&self);
match Arc::clone(&server.connection_slots).try_acquire_owned() {
Ok(permit) => {
tokio::spawn(async move {
let _permit = permit; let result = match &server.tls {
Some(acceptor) => match acceptor.accept(stream).await {
Ok(tls_stream) => server.handle_conn(tls_stream).await,
Err(e) => Err(BrokerError::Tls(format!(
"tls accept handshake ({peer}): {e}"
))),
},
None => server.handle_conn(stream).await,
};
if let Err(e) = result {
tracing::debug!("broker connection ended: {e}");
}
});
}
Err(_) => {
tracing::warn!(
max_connections = server.limits.max_connections,
%peer,
"broker: max_connections reached — rejecting connection"
);
drop(stream);
}
}
}
}
async fn handle_conn<S>(&self, stream: S) -> BrokerResult<()>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let ws = accept_async(stream)
.await
.map_err(|e| BrokerError::Transport(format!("ws accept: {e}")))?;
let (mut sink, mut source) = ws.split();
let deliver_limiter: DeliverLimiter = RateLimiter::direct(
Quota::per_second(self.limits.messages_per_second)
.allow_burst(self.limits.message_burst),
);
let (session_id, role) = match read_client_frame(&mut source).await? {
Some(ClientFrame::Hello { agent, token }) => {
if token != self.token {
let _ = send(
&mut sink,
BrokerFrame::Error {
reason: "invalid token".into(),
id: None,
},
)
.await;
return Err(BrokerError::Auth("invalid token".into()));
}
(agent.session_id, agent.role)
}
Some(_) => {
let _ = send(
&mut sink,
BrokerFrame::Error {
reason: "expected hello".into(),
id: None,
},
)
.await;
return Err(BrokerError::Protocol("expected hello first".into()));
}
None => return Ok(()), };
send(&mut sink, BrokerFrame::Welcome).await?;
let mut sub_rx: Option<mpsc::UnboundedReceiver<PushItem>> = None;
let outcome = loop {
tokio::select! {
biased;
frame = read_client_frame(&mut source) => {
match frame {
Ok(Some(ClientFrame::Deliver { to, message })) => {
deliver_limiter.until_ready().await;
let msg_id = message.id.clone();
match self.core.deliver(&to, &message).await {
Ok(id) => {
if send(&mut sink, BrokerFrame::Delivered { id }).await.is_err() {
break Ok(());
}
}
Err(e) => {
let _ = send(
&mut sink,
BrokerFrame::Error {
reason: e.to_string(),
id: Some(msg_id),
},
)
.await;
}
}
}
Ok(Some(ClientFrame::Subscribe)) => match self.core.subscribe(&session_id, role.as_deref()).await {
Ok(rx) => sub_rx = Some(rx),
Err(e) => {
let _ = send(&mut sink, BrokerFrame::Error { reason: e.to_string(), id: None }).await;
}
},
Ok(Some(ClientFrame::Ack { id })) => {
if let Err(e) = self.core.ack(&session_id, &id).await {
let _ = send(&mut sink, BrokerFrame::Error { reason: e.to_string(), id: None }).await;
}
}
Ok(Some(ClientFrame::Hello { .. })) => {}
Ok(Some(ClientFrame::Cancel { to, correlation_id })) => {
self.core.cancel(&to, &correlation_id).await;
}
Ok(Some(ClientFrame::ListConnected { role })) => {
let ids = self.core.connected_by_role(&role).await;
if send(&mut sink, BrokerFrame::Connected { ids }).await.is_err() {
break Ok(());
}
}
Ok(None) => break Ok(()), Err(e) => break Err(e),
}
}
pushed = next_pushed(&mut sub_rx) => {
match pushed {
Some(PushItem::Message(m)) => {
if send(&mut sink, BrokerFrame::Message { message: m }).await.is_err() {
break Ok(());
}
}
Some(PushItem::Cancel(correlation_id)) => {
if send(&mut sink, BrokerFrame::Cancel { correlation_id }).await.is_err() {
break Ok(());
}
}
None => sub_rx = None, }
}
}
};
self.core.unsubscribe(&session_id).await;
outcome
}
}
async fn next_pushed(rx: &mut Option<mpsc::UnboundedReceiver<PushItem>>) -> Option<PushItem> {
match rx {
Some(r) => r.recv().await,
None => std::future::pending().await,
}
}
async fn read_client_frame<S>(
source: &mut SplitStream<WebSocketStream<S>>,
) -> BrokerResult<Option<ClientFrame>>
where
S: AsyncRead + AsyncWrite + Unpin,
{
loop {
match source.next().await {
Some(Ok(Message::Text(t))) => {
return ClientFrame::from_text(&t)
.map(Some)
.map_err(|e| BrokerError::Protocol(format!("bad client frame: {e}")));
}
Some(Ok(Message::Close(_))) | None => return Ok(None),
Some(Ok(_)) => continue,
Some(Err(e)) => return Err(BrokerError::Transport(format!("ws: {e}"))),
}
}
}
async fn send<S>(
sink: &mut SplitSink<WebSocketStream<S>, Message>,
frame: BrokerFrame,
) -> BrokerResult<()>
where
S: AsyncRead + AsyncWrite + Unpin,
{
sink.send(Message::text(frame.to_text()))
.await
.map_err(|e| BrokerError::Transport(format!("ws send: {e}")))
}