mod handlers;
mod state;
mod wire;
#[cfg(test)]
mod fidelity;
#[cfg(test)]
mod tests;
use std::collections::{HashMap, HashSet};
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::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::net::TcpListener;
use tokio::sync::Notify;
use tokio::task::{AbortHandle, JoinHandle};
use tracing::{debug, warn};
use crate::consumer::ConsumerRecord;
use crate::error::{ErrorCode, KrafkaError, Result};
use crate::protocol::{Decode, KafkaString, TaggedFields};
use crate::protocol::{RequestHeader, ResponseHeader};
pub use crate::protocol::ApiKey;
pub use state::{
BatchMetadata, BrokerNode, BrokerTransaction, ClassicGroupState, ClusterState, CommittedOffset,
ConsumerGroupHeartbeatSeen, ConsumerGroupMemberState, GroupMember, GroupState,
LeaveGroupMemberSeen, ListOffsetsLookup, PartitionState, ProducerEntry, ShareAckType,
ShareGroupState, ShareMemberState, SharePartitionState, ShareSession, ShareSessionClose,
StreamsGroupState, StreamsMemberState, TelemetryPush, TelemetrySubscription, TopicState,
TxnStatus,
};
const MEMORY_BUFFER: usize = 256 * 1024;
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,
ApplyThen(Box<Control>),
}
#[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,
pub connection: u64,
pub at: Duration,
}
#[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 {
started: tokio::time::Instant,
crashed: Mutex<HashSet<i32>>,
open: Mutex<HashMap<i32, Vec<AbortHandle>>>,
cluster: Mutex<ClusterState>,
hooks: Mutex<Hooks>,
log: Mutex<Vec<RecordedRequest>>,
sequence: AtomicU64,
connections: AtomicU64,
open_connections: AtomicU64,
changed: Notify,
}
impl Shared {
fn new(cluster: ClusterState) -> Self {
Self {
started: tokio::time::Instant::now(),
crashed: Mutex::new(HashSet::new()),
open: Mutex::new(HashMap::new()),
cluster: Mutex::new(cluster),
hooks: Mutex::new(Hooks::default()),
log: Mutex::new(Vec::new()),
sequence: AtomicU64::new(0),
connections: AtomicU64::new(0),
open_connections: AtomicU64::new(0),
changed: Notify::new(),
}
}
fn mutate<T>(&self, f: impl FnOnce(&mut ClusterState) -> T) -> T {
let out = f(&mut self.cluster.lock());
self.changed.notify_waiters();
out
}
fn record(&self, request: RecordedRequest) {
self.log.lock().push(request);
}
fn accept<S>(self: &Arc<Self>, stream: S, node_id: i32)
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
if self.crashed.lock().contains(&node_id) {
return;
}
let shared = Arc::clone(self);
let connection = shared.connections.fetch_add(1, Ordering::Relaxed);
shared.open_connections.fetch_add(1, Ordering::Relaxed);
let task = tokio::spawn(async move {
let _open = OpenConnection(Arc::clone(&shared));
if let Err(e) = serve(stream, node_id, connection, shared).await {
debug!(node_id, "fake broker connection ended: {e}");
}
});
let mut open = self.open.lock();
let tasks = open.entry(node_id).or_default();
tasks.retain(|t| !t.is_finished());
tasks.push(task.abort_handle());
}
fn api_call_index(&self, api_key: ApiKey) -> u64 {
self.log
.lock()
.iter()
.filter(|r| r.api_key == api_key)
.count() as u64
}
}
struct OpenConnection(Arc<Shared>);
impl Drop for OpenConnection {
fn drop(&mut self) {
self.0.open_connections.fetch_sub(1, Ordering::Relaxed);
}
}
pub struct FakeBroker {
shared: Arc<Shared>,
transport: Transport,
tasks: Vec<JoinHandle<()>>,
}
#[derive(Debug)]
enum Transport {
Tcp(Vec<SocketAddr>),
Memory(Vec<String>),
}
impl std::fmt::Debug for FakeBroker {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FakeBroker")
.field("transport", &self.transport)
.finish_non_exhaustive()
}
}
impl Drop for FakeBroker {
fn drop(&mut self) {
for task in &self.tasks {
task.abort();
}
for task in self.shared.open.lock().values().flatten() {
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::new(cluster));
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,
transport: Transport::Tcp(addrs),
tasks,
})
}
pub fn start_in_memory(broker_count: usize) -> Self {
let broker_count = broker_count.max(1);
let mut cluster = ClusterState::new(broker_count);
let mut addrs = Vec::with_capacity(broker_count);
for broker in &mut cluster.brokers {
broker.host = format!("broker-{}", broker.node_id);
broker.port = 9092;
addrs.push(format!("{}:{}", broker.host, broker.port));
}
Self {
shared: Arc::new(Shared::new(cluster)),
transport: Transport::Memory(addrs),
tasks: Vec::new(),
}
}
pub fn kafka(&self) -> crate::KafkaBuilder {
let builder = crate::Kafka::builder(self.bootstrap_servers());
match &self.transport {
Transport::Tcp(_) => builder,
Transport::Memory(addrs) => {
let shared = Arc::downgrade(&self.shared);
let addrs = addrs.clone();
builder.connector(Arc::new(move |address: &str| {
let refused = || io::Error::from(io::ErrorKind::ConnectionRefused);
let shared = shared.upgrade().ok_or_else(refused)?;
let node_id = addrs
.iter()
.position(|a| a == address)
.and_then(|i| i32::try_from(i).ok())
.ok_or_else(|| {
io::Error::new(
io::ErrorKind::NotFound,
format!("no in-memory broker at {address}"),
)
})?;
if shared.crashed.lock().contains(&node_id) {
return Err(refused());
}
let (client, server) = tokio::io::duplex(MEMORY_BUFFER);
shared.accept(server, node_id);
Ok(client)
}))
}
}
}
pub fn bootstrap_servers(&self) -> String {
match &self.transport {
Transport::Tcp(addrs) => addrs
.iter()
.map(SocketAddr::to_string)
.collect::<Vec<_>>()
.join(","),
Transport::Memory(addrs) => addrs.join(","),
}
}
pub fn broker_addr(&self, node_id: i32) -> Option<SocketAddr> {
match &self.transport {
Transport::Tcp(addrs) => addrs.get(usize::try_from(node_id).ok()?).copied(),
Transport::Memory(_) => None,
}
}
pub fn crash(&self, node_id: i32) {
self.shared.crashed.lock().insert(node_id);
if let Some(tasks) = self.shared.open.lock().remove(&node_id) {
for task in tasks {
task.abort();
}
}
}
pub fn restart(&self, node_id: i32) {
self.shared.crashed.lock().remove(&node_id);
}
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 open_connections(&self) -> u64 {
self.shared.open_connections.load(Ordering::Relaxed)
}
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 {
self.shared.mutate(f)
}
pub fn create_topic(&self, name: &str, partitions: i32) -> bool {
self.with_state(|s| s.create_topic(name, partitions))
}
pub fn delete_topic(&self, name: &str) -> bool {
self.with_state(|s| s.delete_topic(name))
}
pub fn topic_id(&self, name: &str) -> Option<[u8; 16]> {
self.shared
.cluster
.lock()
.topics
.get(name)
.map(|t| t.topic_id)
}
pub fn add_partitions(&self, topic: &str, partitions: i32) -> usize {
self.with_state(|s| s.add_partitions(topic, partitions))
}
pub fn set_leader(&self, topic: &str, partition: i32, node_id: i32) -> bool {
self.with_state(|cluster| 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 {
self.with_state(|cluster| 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_broker_rack(&self, node_id: i32, rack: Option<&str>) {
let mut cluster = self.shared.cluster.lock();
if let Some(broker) = cluster.brokers.iter_mut().find(|b| b.node_id == node_id) {
broker.rack = rack.map(str::to_string);
}
}
pub fn set_throttle(&self, api_key: ApiKey, throttle: Duration) {
let ms = i32::try_from(throttle.as_millis()).unwrap_or(i32::MAX);
let mut cluster = self.shared.cluster.lock();
if ms == 0 {
cluster.throttle_time_ms.remove(&api_key);
} else {
cluster.throttle_time_ms.insert(api_key, ms);
}
}
pub fn require_sasl_plain(&self, username: &str, password: &str) {
self.shared.mutate(|s| {
s.sasl_plain = Some((username.to_string(), password.to_string()));
});
}
pub fn set_idempotence(&self, enabled: bool) {
self.shared.cluster.lock().idempotence = enabled;
}
pub fn clear_producer_state(&self, topic: &str, partition: i32) -> bool {
self.with_state(|cluster| match cluster.partition_mut(topic, partition) {
Some(p) => {
p.producers.clear();
true
}
None => false,
})
}
pub fn hold_transaction_markers(&self, hold: bool) {
self.with_state(|cluster| {
cluster.hold_transaction_markers = hold;
if !hold {
let ids: Vec<String> = cluster.transactions.keys().cloned().collect();
for id in ids {
cluster.write_transaction_markers(&id);
}
}
});
}
pub fn abort_transaction(&self, transactional_id: &str) -> bool {
self.with_state(|cluster| {
let open = cluster
.transactions
.get(transactional_id)
.is_some_and(BrokerTransaction::is_open);
if open {
cluster.fence_transaction(transactional_id);
}
open
})
}
pub fn transaction_status(&self, transactional_id: &str) -> Option<TxnStatus> {
self.shared
.cluster
.lock()
.transactions
.get(transactional_id)
.map(|t| t.status)
}
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 set_telemetry(&self, subscription: Option<TelemetrySubscription>) {
self.shared.mutate(|state| state.telemetry = subscription);
}
pub fn telemetry_pushes(&self) -> Vec<TelemetryPush> {
self.shared.cluster.lock().telemetry_pushes.clone()
}
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 share_acknowledgements(
&self,
group_id: &str,
topic: &str,
partition: i32,
) -> std::collections::BTreeMap<i64, Vec<ShareAckType>> {
self.shared
.cluster
.lock()
.share_groups
.get(group_id)
.and_then(|g| g.partitions.get(&(topic.to_string(), partition)))
.map(|p| p.acknowledgements.clone())
.unwrap_or_default()
}
pub fn share_session_closes(&self) -> Vec<ShareSessionClose> {
self.shared.cluster.lock().share_session_closes.clone()
}
pub fn list_offsets_lookups(&self) -> Vec<ListOffsetsLookup> {
self.shared.cluster.lock().list_offsets_lookups.clone()
}
pub fn leave_group_members(&self) -> Vec<LeaveGroupMemberSeen> {
self.shared.cluster.lock().leave_group_members.clone()
}
pub fn consumer_group_heartbeats(&self) -> Vec<ConsumerGroupHeartbeatSeen> {
self.shared.cluster.lock().consumer_group_heartbeats.clone()
}
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(BrokerTransaction::is_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: std::sync::Arc::from(topic),
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,
key: record.key,
value: record.value,
headers: crate::consumer::headers_from_wire(record.headers),
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()
}
}
pub fn seed_rng(seed: u64) {
crate::util::seed_rng(seed);
}
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");
shared.accept(stream, node_id);
}
Err(e) => {
warn!(node_id, "fake broker accept failed: {e}");
return;
}
}
}
}
async fn serve<S>(mut stream: S, node_id: i32, connection: u64, shared: Arc<Shared>) -> Result<()>
where
S: AsyncRead + AsyncWrite + Unpin,
{
let mut authenticated = false;
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,
connection,
at: shared.started.elapsed(),
});
let sasl_plain = shared.cluster.lock().sasl_plain.clone();
if let Some(credentials) = sasl_plain {
let mut body = BytesMut::new();
match api_key {
ApiKey::ApiVersions => {}
ApiKey::SaslHandshake => {
handlers::sasl_handshake(&mut frame, &mut body)?;
write_response(&mut stream, &header, &body).await?;
continue;
}
ApiKey::SaslAuthenticate => {
authenticated =
handlers::sasl_authenticate(&mut frame, &credentials, &mut body)?;
write_response(&mut stream, &header, &body).await?;
continue;
}
_ if !authenticated => return Ok(()),
_ => {}
}
}
let mut control = shared.hooks.lock().take(&info).unwrap_or(Control::Pass);
control = run_delays(control).await;
let mut body = BytesMut::new();
match control {
Control::Disconnect => return Ok(()),
Control::Silence => return silence().await,
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)?;
}
Control::ApplyThen(after) => {
serve_default(&shared, &header, &frame, node_id, true, &mut body).await?;
match run_delays(*after).await {
Control::Pass => {}
Control::Disconnect => return Ok(()),
Control::Silence => return silence().await,
Control::Error(code) => {
body.clear();
handlers::dispatch_error(
api_key,
header.api_version,
&mut frame,
code,
&mut body,
)?;
}
other => {
return Err(KrafkaError::protocol_kind(
crate::error::ProtocolErrorKind::InvalidValue,
format!("fake broker: {other:?} cannot follow Control::ApplyThen"),
));
}
}
}
_ => serve_default(&shared, &header, &frame, node_id, false, &mut body).await?,
}
let throttle = shared.cluster.lock().throttle(api_key);
apply_throttle(api_key, header.api_version, throttle, &mut body);
write_response(&mut stream, &header, &body).await?;
}
}
async fn write_response<S>(stream: &mut S, header: &ParsedHeader, body: &[u8]) -> Result<()>
where
S: AsyncWrite + Unpin,
{
let mut out = BytesMut::with_capacity(body.len() + 8);
out.put_i32(0); write_response_header(
&mut out,
header.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)
}
async fn run_delays(mut control: Control) -> Control {
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;
}
other => return other,
}
}
}
async fn silence() -> Result<()> {
std::future::pending::<()>().await;
Ok(())
}
async fn serve_default(
shared: &Shared,
header: &ParsedHeader,
frame: &Bytes,
node_id: i32,
no_wait: bool,
body: &mut BytesMut,
) -> Result<()> {
let mut poll = handlers::LongPoll {
expired: no_wait,
..handlers::LongPoll::default()
};
let mut deadline: Option<tokio::time::Instant> = None;
loop {
let changed = shared.changed.notified();
tokio::pin!(changed);
changed.as_mut().enable();
body.clear();
let mut request = frame.clone();
let served = {
let mut cluster = shared.cluster.lock();
handlers::dispatch(
header.api_key,
header.api_version,
&mut request,
node_id,
header.client_id.as_deref(),
&mut cluster,
&mut poll,
body,
)?
};
match served {
handlers::Served::Done => break,
handlers::Served::Wait(max_wait) => {
let until = *deadline.get_or_insert_with(|| tokio::time::Instant::now() + max_wait);
tokio::select! {
() = &mut changed => {}
() = tokio::time::sleep_until(until) => poll.expired = true,
}
}
}
}
if !matches!(header.api_key, ApiKey::Fetch | ApiKey::ShareFetch) {
shared.changed.notify_waiters();
}
Ok(())
}
fn apply_throttle(api_key: ApiKey, api_version: i16, throttle_time_ms: i32, body: &mut BytesMut) {
let leads = api_key
.leading_throttle_time_min_version()
.is_some_and(|min| api_version >= min);
if throttle_time_ms > 0 && leads && body.len() >= 4 {
body[..4].copy_from_slice(&throttle_time_ms.to_be_bytes());
}
}
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));
}
}