use std::sync::Arc;
use lapin::options::{
BasicConsumeOptions, BasicQosOptions, ExchangeDeclareOptions, QueueBindOptions,
QueueDeclareOptions,
};
use lapin::types::{AMQPValue, FieldTable, ShortString};
use lapin::{Channel, Connection, ConnectionProperties};
use ruststream::{Broker, DescribeServer, ServerSpec, Subscribe};
use tokio::sync::OnceCell;
use crate::convert;
use crate::delay::{Delay, DelayContext, DelayTarget};
use crate::error::AmqpError;
use crate::publisher::LapinPublisher;
use crate::queue::{QueueType, RabbitQueue};
use crate::requester::LapinRequester;
use crate::subscriber::LapinSubscriber;
#[derive(Debug)]
pub(crate) struct ConnState {
connection: Connection,
publish_channel: Channel,
}
impl ConnState {
pub(crate) fn connection(&self) -> &Connection {
&self.connection
}
pub(crate) fn publish_channel(&self) -> &Channel {
&self.publish_channel
}
}
pub(crate) type SharedConn = Arc<OnceCell<ConnState>>;
#[derive(Debug, Clone)]
pub struct LapinBroker {
conn: SharedConn,
uri: String,
connection_name: Option<String>,
prefetch: Option<u16>,
declare: bool,
default_queue_type: Option<QueueType>,
}
impl LapinBroker {
#[must_use]
pub fn new(uri: impl Into<String>) -> Self {
Self {
conn: Arc::new(OnceCell::new()),
uri: uri.into(),
connection_name: None,
prefetch: None,
declare: false,
default_queue_type: None,
}
}
pub async fn connect(uri: impl Into<String>) -> Result<Self, AmqpError> {
let broker = Self::new(uri);
Broker::connect(&broker).await?;
Ok(broker)
}
#[must_use]
pub fn connection_name(mut self, name: impl Into<String>) -> Self {
self.connection_name = Some(name.into());
self
}
#[must_use]
pub fn prefetch(mut self, prefetch: u16) -> Self {
self.prefetch = Some(prefetch);
self
}
#[must_use]
pub fn declare_topology(mut self, declare: bool) -> Self {
self.declare = declare;
self
}
#[must_use]
pub fn default_queue_type(mut self, queue_type: QueueType) -> Self {
self.default_queue_type = Some(queue_type);
self
}
fn connected(&self) -> Result<&ConnState, AmqpError> {
self.conn.get().ok_or(AmqpError::NotConnected)
}
pub async fn subscribe(&self, def: RabbitQueue) -> Result<LapinSubscriber, AmqpError> {
let state = self.connected()?;
let channel = state
.connection
.create_channel()
.await
.map_err(AmqpError::subscribe)?;
if self.declare {
declare_topology(&channel, &def, self.default_queue_type).await?;
}
if let Some(prefetch) = def.prefetch_or(self.prefetch) {
channel
.basic_qos(prefetch, BasicQosOptions::default())
.await
.map_err(AmqpError::subscribe)?;
}
let queue = def.name().to_owned();
let delay = def
.delay_config()
.map(|delay| DelayContext::new(channel.clone(), delay.target_for(&queue)));
let consumer = channel
.basic_consume(
convert::short(&queue, "queue name")?,
ShortString::default(),
BasicConsumeOptions::default(),
FieldTable::default(),
)
.await
.map_err(AmqpError::subscribe)?;
Ok(LapinSubscriber::new(channel, consumer, queue, delay))
}
#[must_use]
pub fn publisher(&self) -> LapinPublisher {
LapinPublisher::new(Arc::clone(&self.conn))
}
#[must_use]
pub fn requester(&self) -> LapinRequester {
LapinRequester::new(Arc::clone(&self.conn))
}
}
impl Broker for LapinBroker {
type Error = AmqpError;
async fn connect(&self) -> Result<(), Self::Error> {
self.conn
.get_or_try_init(|| async {
let mut properties = ConnectionProperties::default();
if let Some(name) = &self.connection_name {
properties = properties.with_connection_name(name.as_str().into());
}
let connection = Connection::connect(&self.uri, properties)
.await
.map_err(AmqpError::connect)?;
let publish_channel = connection
.create_channel()
.await
.map_err(AmqpError::connect)?;
Ok(ConnState {
connection,
publish_channel,
})
})
.await?;
Ok(())
}
async fn shutdown(&self) -> Result<(), Self::Error> {
if let Some(state) = self.conn.get()
&& state.connection.status().connected()
{
state
.connection
.close(200, ShortString::from("OK"))
.await
.map_err(AmqpError::connect)?;
}
Ok(())
}
}
#[allow(clippy::use_self)]
impl Subscribe for LapinBroker {
type Subscriber = LapinSubscriber;
async fn subscribe(&self, name: &str) -> Result<Self::Subscriber, Self::Error> {
LapinBroker::subscribe(self, RabbitQueue::new(name)).await
}
}
impl DescribeServer for LapinBroker {
fn describe_server(&self) -> ServerSpec {
ServerSpec::new(host_of(&self.uri), "amqp")
}
}
fn host_of(uri: &str) -> String {
let after_scheme = uri.split_once("://").map_or(uri, |(_, rest)| rest);
let after_auth = after_scheme
.rsplit_once('@')
.map_or(after_scheme, |(_, rest)| rest);
let host = after_auth.split(['/', '?']).next().unwrap_or(after_auth);
host.to_owned()
}
async fn declare_topology(
channel: &Channel,
def: &RabbitQueue,
broker_default: Option<QueueType>,
) -> Result<(), AmqpError> {
for (exchange, _) in def.bindings() {
if exchange.name().is_empty() || exchange.name().starts_with("amq.") {
continue;
}
channel
.exchange_declare(
convert::short(exchange.name(), "exchange name")?,
exchange.kind().clone(),
ExchangeDeclareOptions {
durable: exchange.is_durable(),
auto_delete: exchange.is_auto_delete(),
..ExchangeDeclareOptions::default()
},
FieldTable::default(),
)
.await
.map_err(AmqpError::declare)?;
}
let queue_type = def.queue_type_or(broker_default);
if queue_type == Some(QueueType::Quorum) && !def.is_durable() {
return Err(AmqpError::InvalidOptions(format!(
"queue {:?} is a quorum queue and must stay durable; drop `.durable(false)` or pick \
`QueueType::Classic`",
def.name(),
)));
}
let mut arguments = def.declare_arguments().clone();
if let Some(queue_type) = queue_type {
arguments.insert(
ShortString::from("x-queue-type"),
AMQPValue::LongString(queue_type.as_str().into()),
);
}
channel
.queue_declare(
convert::short(def.name(), "queue name")?,
QueueDeclareOptions {
durable: def.is_durable(),
exclusive: def.is_exclusive(),
auto_delete: def.is_auto_delete(),
..QueueDeclareOptions::default()
},
arguments,
)
.await
.map_err(AmqpError::declare)?;
for (exchange, routing_key) in def.bindings() {
channel
.queue_bind(
convert::short(def.name(), "queue name")?,
convert::short(exchange.name(), "exchange name")?,
convert::short(routing_key, "routing key")?,
QueueBindOptions::default(),
FieldTable::default(),
)
.await
.map_err(AmqpError::declare)?;
}
if let Some(delay) = def.delay_config() {
declare_delay_backend(channel, delay, def.name()).await?;
}
Ok(())
}
async fn declare_delay_backend(
channel: &Channel,
delay: &Delay,
origin: &str,
) -> Result<(), AmqpError> {
match delay.target_for(origin) {
DelayTarget::WaitingQueue { waiting_queue } => {
declare_delay_queue(channel, &waiting_queue, origin).await
}
#[cfg(feature = "plugin-dme")]
DelayTarget::DelayedExchange {
exchange,
routing_key,
} => declare_delayed_exchange(channel, &exchange, origin, &routing_key).await,
}
}
async fn declare_delay_queue(
channel: &Channel,
waiting_queue: &str,
origin: &str,
) -> Result<(), AmqpError> {
let mut arguments = FieldTable::default();
arguments.insert(
ShortString::from("x-dead-letter-exchange"),
AMQPValue::LongString(String::new().into()),
);
arguments.insert(
ShortString::from("x-dead-letter-routing-key"),
AMQPValue::LongString(origin.into()),
);
channel
.queue_declare(
convert::short(waiting_queue, "waiting queue name")?,
QueueDeclareOptions {
durable: true,
..QueueDeclareOptions::default()
},
arguments,
)
.await
.map_err(AmqpError::declare)?;
Ok(())
}
#[cfg(feature = "plugin-dme")]
async fn declare_delayed_exchange(
channel: &Channel,
exchange: &str,
origin: &str,
routing_key: &str,
) -> Result<(), AmqpError> {
let mut arguments = FieldTable::default();
arguments.insert(
ShortString::from("x-delayed-type"),
AMQPValue::LongString("direct".into()),
);
channel
.exchange_declare(
convert::short(exchange, "delayed exchange name")?,
lapin::ExchangeKind::Custom("x-delayed-message".to_owned()),
ExchangeDeclareOptions {
durable: true,
..ExchangeDeclareOptions::default()
},
arguments,
)
.await
.map_err(AmqpError::declare)?;
channel
.queue_bind(
convert::short(origin, "queue name")?,
convert::short(exchange, "delayed exchange name")?,
convert::short(routing_key, "routing key")?,
QueueBindOptions::default(),
FieldTable::default(),
)
.await
.map_err(AmqpError::declare)?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::host_of;
#[test]
fn host_extraction_handles_auth_vhost_and_bare_forms() {
assert_eq!(host_of("amqp://localhost:5672"), "localhost:5672");
assert_eq!(host_of("amqp://user:pass@rabbit:5672/prod"), "rabbit:5672");
assert_eq!(host_of("amqps://rabbit/vhost"), "rabbit");
assert_eq!(host_of("rabbit:5672"), "rabbit:5672");
}
}