use arc_swap::ArcSwap;
use dashmap::DashMap;
use std::net::{SocketAddr, ToSocketAddrs};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicI32, Ordering};
use std::time::Duration;
use tokio::sync::Mutex;
use tracing::{debug, warn};
use crate::connection::{Builder, Connection};
use crate::error::{KafkaError, Result};
use crate::sasl::SaslCredentials;
use crate::transport::SecurityProtocol;
use kafka_client_protocol::{ApiVersionsRequest, MetadataResponseBroker};
struct BrokerEntry {
addr: SocketAddr,
conn: ArcSwap<Mutex<Connection>>,
healthy: AtomicBool,
}
impl BrokerEntry {
fn new(addr: SocketAddr, conn: Connection) -> Self {
Self {
addr,
conn: ArcSwap::new(Arc::new(Mutex::new(conn))),
healthy: AtomicBool::new(true),
}
}
fn is_healthy(&self) -> bool {
self.healthy.load(Ordering::Relaxed)
}
fn mark_unhealthy(&self) {
self.healthy.store(false, Ordering::Relaxed);
}
fn mark_healthy(&self) {
self.healthy.store(true, Ordering::Relaxed);
}
fn load_conn(&self) -> Arc<Mutex<Connection>> {
self.conn.load_full()
}
fn swap_conn(&self, new_conn: Connection) {
self.conn.store(Arc::new(Mutex::new(new_conn)));
self.mark_healthy();
}
}
pub struct BrokerManager {
bootstrap_servers: Vec<SocketAddr>,
security_protocol: SecurityProtocol,
client_id: String,
client_name: String,
client_version: String,
sasl: Option<SaslCredentials>,
brokers: DashMap<i32, BrokerEntry>,
addr_to_node: DashMap<SocketAddr, i32>,
next_unknown_node_id: AtomicI32,
}
impl BrokerManager {
pub fn new(
bootstrap_servers: Vec<SocketAddr>,
security_protocol: SecurityProtocol,
client_id: String,
client_name: String,
client_version: String,
sasl: Option<SaslCredentials>,
) -> Self {
Self {
bootstrap_servers,
security_protocol,
client_id,
client_name,
client_version,
sasl,
brokers: DashMap::new(),
addr_to_node: DashMap::new(),
next_unknown_node_id: AtomicI32::new(i32::MIN),
}
}
async fn connect_to_broker(&self, addr: SocketAddr) -> Result<Connection> {
let mut builder = Builder::new(
addr,
self.security_protocol.clone(),
self.client_name.clone(),
self.client_version.clone(),
)
.with_client_id(self.client_id.clone());
if let Some(ref sasl) = self.sasl {
builder = builder.with_sasl(sasl.mechanism, sasl.clone());
}
builder.build().await
}
pub async fn bootstrap(&self) -> Result<SocketAddr> {
let addrs: Vec<SocketAddr> = self.bootstrap_servers.clone();
for addr in addrs {
match self.connect_to_broker(addr).await {
Ok(conn) => {
let node_id = self.next_unknown_node_id.fetch_sub(1, Ordering::SeqCst);
self.register_broker(node_id, addr, conn).await;
debug!("Connected to bootstrap broker {}", addr);
return Ok(addr);
}
Err(e) => {
warn!("Failed to connect to bootstrap broker {}: {}", addr, e);
continue;
}
}
}
Err(KafkaError::NoBootstrapBrokerAvailable)
}
async fn register_broker(&self, node_id: i32, addr: SocketAddr, conn: Connection) {
if let Some(old_node_id) = self.addr_to_node.get(&addr).map(|e| *e) {
if old_node_id != node_id {
self.brokers.remove(&old_node_id);
}
}
self.addr_to_node.insert(addr, node_id);
self.brokers
.insert(node_id, BrokerEntry::new(addr, conn));
}
pub async fn get_connection(&self, addr: SocketAddr) -> Result<Arc<Mutex<Connection>>> {
if let Some(node_id) = self.addr_to_node.get(&addr).map(|e| *e) {
if let Some(entry) = self.brokers.get(&node_id) {
if entry.is_healthy() {
return Ok(entry.load_conn());
}
drop(entry);
match self.try_swap_connection(node_id, addr).await {
Some(conn) => return Ok(conn),
None => {
}
}
}
}
let conn = self.connect_to_broker(addr).await?;
let node_id = self.next_unknown_node_id.fetch_sub(1, Ordering::SeqCst);
self.register_broker(node_id, addr, conn).await;
self.brokers
.get(&node_id)
.map(|e| e.load_conn())
.ok_or_else(|| {
KafkaError::InvalidConfiguration(
"Failed to register new broker connection".to_string(),
)
})
}
async fn try_swap_connection(
&self,
node_id: i32,
addr: SocketAddr,
) -> Option<Arc<Mutex<Connection>>> {
match self.connect_to_broker(addr).await {
Ok(new_conn) => {
if let Some(entry) = self.brokers.get(&node_id) {
entry.swap_conn(new_conn);
debug!("Reconnected broker {} at {}", node_id, addr);
return Some(entry.load_conn());
}
None
}
Err(e) => {
warn!("Reconnect to broker {} at {} failed: {}", node_id, addr, e);
None
}
}
}
pub fn get_any_healthy_broker(&self) -> Option<(SocketAddr, Arc<Mutex<Connection>>)> {
self.brokers
.iter()
.find(|e| e.is_healthy())
.map(|e| (e.addr, e.load_conn()))
}
pub fn all_broker_addresses(&self) -> Vec<SocketAddr> {
self.brokers.iter().map(|e| e.addr).collect()
}
pub async fn refresh_from_metadata(&self, brokers: Vec<MetadataResponseBroker>) -> Result<()> {
for broker in brokers {
let addr = resolve_broker_address(&broker.host, broker.port)?;
let node_id = broker.node_id;
if let Some(entry) = self.brokers.get(&node_id) {
if entry.addr == addr && entry.is_healthy() {
continue;
}
}
match self.connect_to_broker(addr).await {
Ok(conn) => {
self.register_broker(node_id, addr, conn).await;
debug!("Registered/updated broker {} at {}", node_id, addr);
}
Err(e) => {
warn!("Could not connect to broker {} at {}: {}", node_id, addr, e);
}
}
}
Ok(())
}
pub fn mark_unhealthy(&self, addr: SocketAddr) {
if let Some(node_id) = self.addr_to_node.get(&addr).map(|e| *e) {
if let Some(entry) = self.brokers.get(&node_id) {
entry.mark_unhealthy();
warn!("Marked broker {} at {} as unhealthy", node_id, addr);
}
}
}
pub async fn force_close_connection(&self, addr: SocketAddr) {
let node_id = match self.addr_to_node.get(&addr).map(|e| *e) {
Some(id) => id,
None => return,
};
match self.connect_to_broker(addr).await {
Ok(new_conn) => {
if let Some(entry) = self.brokers.get(&node_id) {
let old = entry.conn.swap(Arc::new(Mutex::new(new_conn)));
entry.mark_healthy();
tokio::spawn(async move {
if let Ok(conn) = Arc::try_unwrap(old) {
let conn = conn.into_inner();
let _ = conn.close().await;
}
});
warn!("Force-closed (and replaced) connection to broker {} at {}", node_id, addr);
}
}
Err(e) => {
warn!(
"Force-close for broker {} at {} failed to reconnect: {}",
node_id, addr, e
);
self.brokers.remove(&node_id);
self.addr_to_node.remove(&addr);
}
}
}
pub async fn close(&self) -> Result<()> {
let node_ids: Vec<i32> = self.brokers.iter().map(|e| *e.key()).collect();
for node_id in node_ids {
if let Some((_, entry)) = self.brokers.remove(&node_id) {
let old = entry.conn.into_inner();
if let Ok(conn) = Arc::try_unwrap(old) {
let conn = conn.into_inner();
if let Err(e) = conn.close().await {
warn!(
"Error closing connection to broker {} at {}: {}",
node_id, entry.addr, e
);
}
}
}
}
self.addr_to_node.clear();
Ok(())
}
#[allow(dead_code)]
pub fn spawn_health_check(this: &Arc<Self>, check_interval: Duration) {
let this = this.clone();
tokio::spawn(async move {
let mut interval = tokio::time::interval(check_interval);
interval.tick().await;
loop {
interval.tick().await;
debug!("Starting broker health check");
let entries: Vec<(i32, SocketAddr)> = this
.brokers
.iter()
.map(|e| (*e.key(), e.addr))
.collect();
for (node_id, addr) in &entries {
let healthy = this
.brokers
.get(node_id)
.map(|e| e.is_healthy())
.unwrap_or(false);
if healthy {
let conn = this.brokers.get(node_id).map(|e| e.load_conn());
if let Some(ref conn) = conn {
let mut guard = conn.lock().await;
let request = ApiVersionsRequest {
client_software_name: String::new(),
client_software_version: String::new(),
};
match guard
.send_request::<_, crate::protocol::ApiVersionsResponse>(&request)
.await
{
Ok(_) => continue,
Err(e) => {
warn!(
"Health check failed for broker {} at {}: {}",
node_id, addr, e
);
}
}
}
drop(conn);
this.mark_unhealthy(*addr);
}
this.force_close_connection(*addr).await;
}
}
});
}
}
fn resolve_broker_address(host: &str, port: i32) -> Result<SocketAddr> {
let addr_str = format!("{}:{}", host, port);
addr_str
.to_socket_addrs()
.map_err(|_| {
KafkaError::InvalidConfiguration(format!("Invalid broker address: {}", addr_str))
})?
.next()
.ok_or_else(|| {
KafkaError::InvalidConfiguration(format!("No address resolved for: {}", addr_str))
})
}