mod handlers;
mod state;
mod wire;
#[cfg(test)]
mod tests;
use std::collections::HashMap;
use std::io;
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use bytes::{BufMut, Bytes, BytesMut};
use parking_lot::Mutex;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::task::JoinHandle;
use tracing::{debug, warn};
use crate::consumer::ConsumerRecord;
use crate::error::{ErrorCode, KrafkaError, Result};
use crate::protocol::ApiKey;
use crate::protocol::{Decode, KafkaString, TaggedFields};
use crate::protocol::{RequestHeader, ResponseHeader};
pub use state::{
BrokerNode, ClusterState, CommittedOffset, GroupMember, GroupState, PartitionState, TopicState,
};
const MAX_FRAME_LEN: usize = 64 * 1024 * 1024;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum Control {
Pass,
Error(ErrorCode),
Delay(Duration),
DelayThen(Duration, Box<Control>),
Disconnect,
Silence,
CorruptRecords,
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct RecordedRequest {
pub api_key: ApiKey,
pub api_version: i16,
pub correlation_id: i32,
pub client_id: Option<String>,
pub node_id: i32,
pub sequence: u64,
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct RequestInfo {
pub api_key: ApiKey,
pub api_version: i16,
pub correlation_id: i32,
pub client_id: Option<String>,
pub node_id: i32,
pub api_call_index: u64,
}
type HookFn = Arc<dyn Fn(&RequestInfo) -> Control + Send + Sync>;
struct Hook {
apply: HookFn,
remaining: Option<u32>,
}
#[derive(Default)]
struct Hooks {
by_api: HashMap<ApiKey, Vec<Hook>>,
}
impl Hooks {
fn take(&mut self, info: &RequestInfo) -> Option<Control> {
let hooks = self.by_api.get_mut(&info.api_key)?;
let hook = hooks.first_mut()?;
let control = (hook.apply)(info);
if let Some(remaining) = hook.remaining.as_mut() {
*remaining = remaining.saturating_sub(1);
if *remaining == 0 {
hooks.remove(0);
}
}
Some(control)
}
}
struct Shared {
cluster: Mutex<ClusterState>,
hooks: Mutex<Hooks>,
log: Mutex<Vec<RecordedRequest>>,
sequence: AtomicU64,
}
impl Shared {
fn record(&self, request: RecordedRequest) {
self.log.lock().push(request);
}
fn api_call_index(&self, api_key: ApiKey) -> u64 {
self.log
.lock()
.iter()
.filter(|r| r.api_key == api_key)
.count() as u64
}
}
pub struct FakeBroker {
shared: Arc<Shared>,
addrs: Vec<SocketAddr>,
tasks: Vec<JoinHandle<()>>,
}
impl std::fmt::Debug for FakeBroker {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FakeBroker")
.field("addrs", &self.addrs)
.finish_non_exhaustive()
}
}
impl Drop for FakeBroker {
fn drop(&mut self) {
for task in &self.tasks {
task.abort();
}
}
}
impl FakeBroker {
pub async fn start() -> Result<Self> {
Self::start_cluster(1).await
}
pub async fn start_cluster(broker_count: usize) -> Result<Self> {
let broker_count = broker_count.max(1);
let mut listeners = Vec::with_capacity(broker_count);
let mut addrs = Vec::with_capacity(broker_count);
for _ in 0..broker_count {
let listener = TcpListener::bind("127.0.0.1:0").await.map_err(io_error)?;
addrs.push(listener.local_addr().map_err(io_error)?);
listeners.push(listener);
}
let mut cluster = ClusterState::new(broker_count);
for (broker, addr) in cluster.brokers.iter_mut().zip(&addrs) {
broker.host = addr.ip().to_string();
broker.port = i32::from(addr.port());
}
let shared = Arc::new(Shared {
cluster: Mutex::new(cluster),
hooks: Mutex::new(Hooks::default()),
log: Mutex::new(Vec::new()),
sequence: AtomicU64::new(0),
});
let tasks = listeners
.into_iter()
.enumerate()
.map(|(index, listener)| {
let shared = Arc::clone(&shared);
let node_id = index as i32;
tokio::spawn(async move { accept_loop(listener, node_id, shared).await })
})
.collect();
Ok(Self {
shared,
addrs,
tasks,
})
}
pub fn bootstrap_servers(&self) -> String {
self.addrs
.iter()
.map(SocketAddr::to_string)
.collect::<Vec<_>>()
.join(",")
}
pub fn broker_addr(&self, node_id: i32) -> Option<SocketAddr> {
self.addrs.get(usize::try_from(node_id).ok()?).copied()
}
pub fn on<F>(&self, api_key: ApiKey, hook: F)
where
F: Fn(&RequestInfo) -> Control + Send + Sync + 'static,
{
self.register(api_key, hook, None);
}
pub fn on_once<F>(&self, api_key: ApiKey, hook: F)
where
F: Fn(&RequestInfo) -> Control + Send + Sync + 'static,
{
self.register(api_key, hook, Some(1));
}
pub fn on_times<F>(&self, api_key: ApiKey, times: u32, hook: F)
where
F: Fn(&RequestInfo) -> Control + Send + Sync + 'static,
{
self.register(api_key, hook, Some(times.max(1)));
}
fn register<F>(&self, api_key: ApiKey, hook: F, remaining: Option<u32>)
where
F: Fn(&RequestInfo) -> Control + Send + Sync + 'static,
{
self.shared
.hooks
.lock()
.by_api
.entry(api_key)
.or_default()
.push(Hook {
apply: Arc::new(hook),
remaining,
});
}
pub fn clear_hooks(&self) {
self.shared.hooks.lock().by_api.clear();
}
pub fn requests(&self) -> Vec<RecordedRequest> {
self.shared.log.lock().clone()
}
pub fn request_count(&self, api_key: ApiKey) -> usize {
self.shared
.log
.lock()
.iter()
.filter(|r| r.api_key == api_key)
.count()
}
pub fn request_nodes(&self, api_key: ApiKey) -> Vec<i32> {
self.shared
.log
.lock()
.iter()
.filter(|r| r.api_key == api_key)
.map(|r| r.node_id)
.collect()
}
pub fn clear_requests(&self) {
self.shared.log.lock().clear();
}
pub async fn wait_for_requests(
&self,
api_key: ApiKey,
count: usize,
timeout: Duration,
) -> bool {
let deadline = tokio::time::Instant::now() + timeout;
loop {
if self.request_count(api_key) >= count {
return true;
}
if tokio::time::Instant::now() >= deadline {
return false;
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
}
pub async fn wait_for_request_on_node(
&self,
api_key: ApiKey,
node_id: i32,
timeout: Duration,
) -> bool {
let deadline = tokio::time::Instant::now() + timeout;
loop {
if self.request_nodes(api_key).contains(&node_id) {
return true;
}
if tokio::time::Instant::now() >= deadline {
return false;
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
}
pub fn with_state<T>(&self, f: impl FnOnce(&mut ClusterState) -> T) -> T {
f(&mut self.shared.cluster.lock())
}
pub fn create_topic(&self, name: &str, partitions: i32) -> bool {
self.shared.cluster.lock().create_topic(name, partitions)
}
pub fn add_partitions(&self, topic: &str, partitions: i32) -> usize {
self.shared.cluster.lock().add_partitions(topic, partitions)
}
pub fn set_leader(&self, topic: &str, partition: i32, node_id: i32) -> bool {
let mut cluster = self.shared.cluster.lock();
match cluster.partition_mut(topic, partition) {
Some(p) => {
p.leader = node_id;
p.leader_epoch += 1;
if !p.replicas.contains(&node_id) {
p.replicas.push(node_id);
}
if !p.isr.contains(&node_id) {
p.isr.push(node_id);
}
true
}
None => false,
}
}
pub fn bump_leader_epoch(&self, topic: &str, partition: i32) -> bool {
let mut cluster = self.shared.cluster.lock();
match cluster.partition_mut(topic, partition) {
Some(p) => {
p.leader_epoch += 1;
true
}
None => false,
}
}
pub fn set_group_coordinator(&self, group_id: &str, node_id: i32) {
self.shared
.cluster
.lock()
.group_coordinators
.insert(group_id.to_string(), node_id);
}
pub fn set_txn_coordinator(&self, transactional_id: &str, node_id: i32) {
self.shared
.cluster
.lock()
.txn_coordinators
.insert(transactional_id.to_string(), node_id);
}
pub fn set_controller(&self, node_id: i32) {
self.shared.cluster.lock().controller_id = node_id;
}
pub fn set_broker_online(&self, node_id: i32, online: bool) {
let mut cluster = self.shared.cluster.lock();
if let Some(broker) = cluster.brokers.iter_mut().find(|b| b.node_id == node_id) {
broker.online = online;
}
}
pub fn set_api_versions(&self, api_key: ApiKey, min_version: i16, max_version: i16) {
self.shared
.cluster
.lock()
.api_version_overrides
.insert(api_key, (min_version, max_version));
}
pub fn finalized_feature(&self, feature: &str) -> Option<i16> {
self.shared
.cluster
.lock()
.finalized_features
.get(feature)
.copied()
}
pub fn committed_offset(&self, group_id: &str, topic: &str, partition: i32) -> Option<i64> {
self.shared
.cluster
.lock()
.groups
.get(group_id)
.and_then(|g| g.offsets.get(&(topic.to_string(), partition)))
.map(|c| c.offset)
}
pub fn next_offset(&self, topic: &str, partition: i32) -> Option<i64> {
self.shared
.cluster
.lock()
.partition(topic, partition)
.map(|p| p.next_offset)
}
pub fn set_transaction_version(&self, level: i16) {
let mut cluster = self.shared.cluster.lock();
cluster
.finalized_features
.insert("transaction.version".to_string(), level);
cluster.finalized_features_epoch += 1;
}
pub fn transactional_producer(&self, transactional_id: &str) -> Option<(i64, i16)> {
self.shared
.cluster
.lock()
.transactions
.get(transactional_id)
.map(|t| (t.producer_id, t.producer_epoch))
}
pub fn transaction_is_open(&self, transactional_id: &str) -> bool {
self.shared
.cluster
.lock()
.transactions
.get(transactional_id)
.is_some_and(|t| t.open)
}
pub fn last_stable_offset(&self, topic: &str, partition: i32) -> Option<i64> {
self.shared
.cluster
.lock()
.partition(topic, partition)
.map(|p| p.last_stable_offset())
}
pub fn committed_records(&self, topic: &str) -> Result<Vec<ConsumerRecord>> {
self.read_records(topic, true)
}
pub fn all_records(&self, topic: &str) -> Result<Vec<ConsumerRecord>> {
self.read_records(topic, false)
}
fn read_records(&self, topic: &str, committed_only: bool) -> Result<Vec<ConsumerRecord>> {
use crate::protocol::RecordBatch;
let cluster = self.shared.cluster.lock();
let Some(topic_state) = cluster.topics.get(topic) else {
return Ok(Vec::new());
};
let mut out = Vec::new();
for (index, partition) in topic_state.partitions.iter().enumerate() {
let partition_id = i32::try_from(index).unwrap_or(i32::MAX);
let limit = if committed_only {
partition.last_stable_offset()
} else {
i64::MAX
};
for stored in &partition.log {
let mut buf = stored.clone();
let batch = RecordBatch::decode(&mut buf)?;
let base = batch.base_offset;
let last = base.saturating_add(i64::from(batch.last_offset_delta));
if last >= limit {
continue;
}
if batch.attributes.is_control_batch {
continue;
}
if committed_only
&& batch.attributes.is_transactional
&& partition.aborted_transactions.iter().any(
|(producer_id, first_offset, marker_offset)| {
*producer_id == batch.producer_id
&& base >= *first_offset
&& base < *marker_offset
},
)
{
continue;
}
for record in batch.records {
out.push(ConsumerRecord {
topic: topic.to_string(),
partition: partition_id,
offset: base.saturating_add(i64::from(record.offset_delta)),
timestamp: batch.base_timestamp.saturating_add(record.timestamp_delta),
timestamp_type: batch.attributes.timestamp_type as i8,
key: record.key,
value: record.value,
headers: record
.headers
.into_iter()
.map(|h| (h.key, h.value))
.collect(),
leader_epoch: Some(batch.partition_leader_epoch),
delivery_count: None,
});
}
}
}
Ok(out)
}
pub fn aborted_transactions(&self, topic: &str, partition: i32) -> Vec<(i64, i64)> {
self.shared
.cluster
.lock()
.partition(topic, partition)
.map(|p| p.aborted_transactions_from(0))
.unwrap_or_default()
}
}
fn io_error(e: io::Error) -> KrafkaError {
KrafkaError::network(e)
}
async fn accept_loop(listener: TcpListener, node_id: i32, shared: Arc<Shared>) {
loop {
match listener.accept().await {
Ok((stream, peer)) => {
debug!(node_id, %peer, "fake broker accepted a connection");
let shared = Arc::clone(&shared);
tokio::spawn(async move {
if let Err(e) = serve(stream, node_id, shared).await {
debug!(node_id, "fake broker connection ended: {e}");
}
});
}
Err(e) => {
warn!(node_id, "fake broker accept failed: {e}");
return;
}
}
}
}
async fn serve(mut stream: TcpStream, node_id: i32, shared: Arc<Shared>) -> Result<()> {
loop {
let mut len_buf = [0u8; 4];
match stream.read_exact(&mut len_buf).await {
Ok(_) => {}
Err(e) if e.kind() == io::ErrorKind::UnexpectedEof => return Ok(()),
Err(e) => return Err(io_error(e)),
}
let len = i32::from_be_bytes(len_buf);
if len < 0 || len as usize > MAX_FRAME_LEN {
return Err(KrafkaError::protocol_kind(
crate::error::ProtocolErrorKind::Malformed,
format!("fake broker: implausible request frame length {len}"),
));
}
const CHUNK: usize = 8 * 1024;
let len = len as usize;
let mut frame = Vec::with_capacity(len.min(CHUNK));
let mut chunk = [0u8; CHUNK];
while frame.len() < len {
let want = (len - frame.len()).min(CHUNK);
let read = stream.read(&mut chunk[..want]).await.map_err(io_error)?;
if read == 0 {
return Err(io_error(io::Error::new(
io::ErrorKind::UnexpectedEof,
format!(
"fake broker: peer closed after {} of {len} frame bytes",
frame.len()
),
)));
}
frame.extend_from_slice(&chunk[..read]);
}
let mut frame = Bytes::from(frame);
let header = read_request_header(&mut frame)?;
let api_key = header.api_key;
let sequence = shared.sequence.fetch_add(1, Ordering::Relaxed);
let info = RequestInfo {
api_key,
api_version: header.api_version,
correlation_id: header.correlation_id,
client_id: header.client_id.clone(),
node_id,
api_call_index: shared.api_call_index(api_key),
};
shared.record(RecordedRequest {
api_key,
api_version: header.api_version,
correlation_id: header.correlation_id,
client_id: header.client_id.clone(),
node_id,
sequence,
});
let mut control = shared.hooks.lock().take(&info).unwrap_or(Control::Pass);
loop {
match control {
Control::Delay(d) => {
tokio::time::sleep(d).await;
control = Control::Pass;
}
Control::DelayThen(d, inner) => {
tokio::time::sleep(d).await;
control = *inner;
}
_ => break,
}
}
match control {
Control::Disconnect => return Ok(()),
Control::Silence => {
std::future::pending::<()>().await;
return Ok(());
}
_ => {}
}
let mut body = BytesMut::new();
let outcome = match control {
Control::Error(code) => {
handlers::dispatch_error(api_key, header.api_version, &mut frame, code, &mut body)
}
Control::CorruptRecords => {
let mut cluster = shared.cluster.lock();
handlers::dispatch_corrupt(api_key, &mut frame, node_id, &mut cluster, &mut body)
}
_ => {
let mut cluster = shared.cluster.lock();
handlers::dispatch(
api_key,
header.api_version,
&mut frame,
node_id,
&mut cluster,
&mut body,
)
}
};
outcome?;
let mut out = BytesMut::with_capacity(body.len() + 8);
out.put_i32(0); write_response_header(&mut out, api_key, header.api_version, header.correlation_id);
out.put_slice(&body);
let frame_len = i32::try_from(out.len() - 4).map_err(|_| {
KrafkaError::protocol_kind(
crate::error::ProtocolErrorKind::Malformed,
"fake broker: response frame exceeds i32::MAX",
)
})?;
out[0..4].copy_from_slice(&frame_len.to_be_bytes());
stream.write_all(&out).await.map_err(io_error)?;
stream.flush().await.map_err(io_error)?;
}
}
struct ParsedHeader {
api_key: ApiKey,
api_version: i16,
correlation_id: i32,
client_id: Option<String>,
}
fn read_request_header(buf: &mut Bytes) -> Result<ParsedHeader> {
let api_key = ApiKey::from_i16(i16::decode(buf)?);
let api_version = i16::decode(buf)?;
let correlation_id = i32::decode(buf)?;
let client_id = KafkaString::decode(buf)?.0;
if RequestHeader::header_version(api_key, api_version) == 2 {
let _ = TaggedFields::decode(buf)?;
}
Ok(ParsedHeader {
api_key,
api_version,
correlation_id,
client_id,
})
}
fn write_response_header(
out: &mut BytesMut,
api_key: ApiKey,
api_version: i16,
correlation_id: i32,
) {
out.put_i32(correlation_id);
if ResponseHeader::header_version(api_key, api_version) == 1 {
out.put_u8(0);
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod unit_tests {
use super::*;
use bytes::Buf;
#[test]
fn once_hooks_fire_exactly_once_then_fall_through() {
let mut hooks = Hooks::default();
hooks
.by_api
.entry(ApiKey::Metadata)
.or_default()
.push(Hook {
apply: Arc::new(|_| Control::Error(ErrorCode::NotController)),
remaining: Some(1),
});
let info = RequestInfo {
api_key: ApiKey::Metadata,
api_version: 8,
correlation_id: 1,
client_id: None,
node_id: 0,
api_call_index: 0,
};
assert!(matches!(
hooks.take(&info),
Some(Control::Error(ErrorCode::NotController))
));
assert!(hooks.take(&info).is_none(), "the hook must not fire twice");
}
#[test]
fn queued_hooks_are_consumed_in_registration_order() {
let mut hooks = Hooks::default();
let entry = hooks.by_api.entry(ApiKey::Produce).or_default();
entry.push(Hook {
apply: Arc::new(|_| Control::Error(ErrorCode::NotLeaderForPartition)),
remaining: Some(2),
});
entry.push(Hook {
apply: Arc::new(|_| Control::Disconnect),
remaining: Some(1),
});
let info = RequestInfo {
api_key: ApiKey::Produce,
api_version: 8,
correlation_id: 1,
client_id: None,
node_id: 0,
api_call_index: 0,
};
assert!(matches!(hooks.take(&info), Some(Control::Error(_))));
assert!(matches!(hooks.take(&info), Some(Control::Error(_))));
assert!(matches!(hooks.take(&info), Some(Control::Disconnect)));
assert!(hooks.take(&info).is_none());
}
#[test]
fn request_headers_round_trip_against_the_client_encoder() {
for (api_key, version) in [(ApiKey::Metadata, 8i16), (ApiKey::Metadata, 12i16)] {
let header = RequestHeader::new(api_key, version, 77).with_client_id("krafka-test");
let mut buf = BytesMut::new();
header.encode(&mut buf).unwrap();
let mut buf = buf.freeze();
let parsed = read_request_header(&mut buf).unwrap();
assert_eq!(parsed.api_key, api_key);
assert_eq!(parsed.api_version, version);
assert_eq!(parsed.correlation_id, 77);
assert_eq!(parsed.client_id.as_deref(), Some("krafka-test"));
assert_eq!(buf.remaining(), 0, "header reader left bytes behind");
}
}
#[tokio::test]
async fn a_started_cluster_advertises_one_address_per_broker() {
let broker = FakeBroker::start_cluster(3).await.unwrap();
assert_eq!(broker.bootstrap_servers().split(',').count(), 3);
assert!(broker.broker_addr(2).is_some());
assert!(broker.broker_addr(3).is_none());
}
#[tokio::test]
async fn moving_a_leader_bumps_the_epoch() {
let broker = FakeBroker::start_cluster(2).await.unwrap();
assert!(broker.create_topic("orders", 1));
let before = broker.with_state(|s| {
let p = s.partition("orders", 0).expect("partition exists");
(p.leader, p.leader_epoch)
});
assert_eq!(before, (0, 0));
assert!(broker.set_leader("orders", 0, 1));
let after = broker.with_state(|s| {
let p = s.partition("orders", 0).expect("partition exists");
(p.leader, p.leader_epoch)
});
assert_eq!(after, (1, 1));
assert!(!broker.set_leader("missing", 0, 1));
}
}