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 tracing::{debug, warn};
use crate::connection::{Builder, ConnectionHandle};
use crate::error::{KafkaError, Result};
use crate::sasl::SaslCredentials;
use crate::transport::SecurityProtocol;
use kafka_client_protocol::{ApiVersionsRequest, MetadataResponseBroker};
use krb5_gss::KerberosCredentials;
struct BrokerEntry {
addr: SocketAddr,
conn: ArcSwap<ConnectionHandle>,
healthy: AtomicBool,
}
impl BrokerEntry {
fn new(addr: SocketAddr, conn: ConnectionHandle) -> Self {
Self {
addr,
conn: ArcSwap::new(Arc::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<ConnectionHandle> {
self.conn.load_full()
}
fn swap_conn(&self, new_conn: ConnectionHandle) {
self.conn.store(Arc::new(new_conn));
self.mark_healthy();
}
}
pub(crate) struct BrokerManager {
bootstrap_servers: Vec<SocketAddr>,
security_protocol: SecurityProtocol,
client_id: String,
client_name: String,
client_version: String,
sasl: Option<SaslCredentials>,
kerberos: Option<KerberosCredentials>,
kdc_host: Option<String>,
kdc_port: u16,
broker_hostname: Option<String>,
brokers: DashMap<i32, BrokerEntry>,
addr_to_node: DashMap<SocketAddr, i32>,
next_unknown_node_id: AtomicI32,
}
impl BrokerManager {
pub(crate) fn new(
bootstrap_servers: Vec<SocketAddr>,
security_protocol: SecurityProtocol,
client_id: String,
client_name: String,
client_version: String,
sasl: Option<SaslCredentials>,
kerberos: Option<KerberosCredentials>,
) -> Self {
Self {
bootstrap_servers,
security_protocol,
client_id,
client_name,
client_version,
sasl,
kerberos,
kdc_host: None,
kdc_port: 88,
broker_hostname: None,
brokers: DashMap::new(),
addr_to_node: DashMap::new(),
next_unknown_node_id: AtomicI32::new(i32::MIN),
}
}
pub(crate) fn with_kdc(mut self, host: Option<String>, port: u16) -> Self {
self.kdc_host = host;
self.kdc_port = port;
self
}
pub(crate) fn with_broker_hostname(mut self, host: Option<String>) -> Self {
self.broker_hostname = host;
self
}
async fn connect_to_broker(
&self,
addr: SocketAddr,
broker_hostname: Option<&str>,
) -> Result<ConnectionHandle> {
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())
.with_kdc(self.kdc_host.clone().unwrap_or_default(), self.kdc_port);
if let Some(host) = broker_hostname {
builder = builder.with_broker_hostname(host);
}
if let Some(ref sasl) = self.sasl {
builder = builder.with_sasl(sasl.mechanism(), sasl.clone());
}
if let Some(ref krb) = self.kerberos {
builder = builder.with_kerberos(krb.clone());
}
builder.build().await
}
pub(crate) async fn bootstrap(&self) -> Result<SocketAddr> {
let addrs: Vec<SocketAddr> = self.bootstrap_servers.clone();
let mut errors: Vec<crate::error::BrokerConnError> = Vec::new();
for addr in addrs {
match self
.connect_to_broker(addr, self.broker_hostname.as_deref())
.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);
errors.push(crate::error::BrokerConnError {
addr: addr.to_string(),
error: e,
});
continue;
}
}
}
Err(KafkaError::NoBootstrapBrokerAvailable(
crate::error::BrokerErrors(errors),
))
}
async fn register_broker(&self, node_id: i32, addr: SocketAddr, conn: ConnectionHandle) {
if let Some(old_node_id) = self.addr_to_node.get(&addr).map(|e| *e)
&& 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(crate) async fn get_connection(&self, addr: SocketAddr) -> Result<ConnectionHandle> {
if let Some(node_id) = self.addr_to_node.get(&addr).map(|e| *e)
&& let Some(entry) = self.brokers.get(&node_id)
{
if entry.is_healthy() {
return Ok(entry.load_conn().as_ref().clone());
}
drop(entry);
match self.try_swap_connection(node_id, addr).await {
Some(conn) => return Ok(conn),
None => {
}
}
}
let conn = self.connect_to_broker(addr, None).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().as_ref().clone())
.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<ConnectionHandle> {
match self.connect_to_broker(addr, None).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().as_ref().clone());
}
None
}
Err(e) => {
warn!("Reconnect to broker {} at {} failed: {}", node_id, addr, e);
None
}
}
}
pub(crate) fn get_any_healthy_broker(&self) -> Option<(SocketAddr, ConnectionHandle)> {
self.brokers
.iter()
.find(|e| e.is_healthy())
.map(|e| (e.addr, e.load_conn().as_ref().clone()))
}
pub(crate) fn all_broker_addresses(&self) -> Vec<SocketAddr> {
self.brokers.iter().map(|e| e.addr).collect()
}
pub(crate) 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)
&& entry.addr == addr
&& entry.is_healthy()
{
continue;
}
match self.connect_to_broker(addr, Some(&broker.host)).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(crate) fn mark_unhealthy(&self, addr: SocketAddr) {
if let Some(node_id) = self.addr_to_node.get(&addr).map(|e| *e)
&& let Some(entry) = self.brokers.get(&node_id)
{
entry.mark_unhealthy();
warn!("Marked broker {} at {} as unhealthy", node_id, addr);
}
}
pub(crate) 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, None).await {
Ok(new_conn) => {
if let Some(entry) = self.brokers.get(&node_id) {
entry.swap_conn(new_conn);
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(crate) async fn close(&self) -> Result<()> {
let node_ids: Vec<i32> = self.brokers.iter().map(|e| *e.key()).collect();
for node_id in node_ids {
self.brokers.remove(&node_id);
self.addr_to_node.retain(|_, v| *v != node_id);
}
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(conn) = conn {
let handle = conn.as_ref().clone();
let request = ApiVersionsRequest {
client_software_name: String::new(),
client_software_version: String::new(),
};
match handle
.send_request::<_, crate::protocol::ApiVersionsResponse>(&request)
.await
{
Ok(_) => continue,
Err(e) => {
warn!(
"Health check failed for broker {} at {}: {}",
node_id, addr, e
);
}
}
}
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))
})
}