use bytes::Bytes;
use kafka_client::{Client, ConsumerConfig, ProducerConfig, ProducerRecord, admin::NewTopic};
use std::collections::HashSet;
use std::time::{Duration, Instant};
fn get_bootstrap_addrs() -> Vec<String> {
let bootstrap =
std::env::var("KAFKA_BOOTSTRAP").unwrap_or_else(|_| "127.0.0.1:9092".to_string());
bootstrap.split(',').map(|s| s.trim().to_string()).collect()
}
fn get_topic_name() -> String {
std::env::var("KAFKA_TOPIC").unwrap_or_else(|_| "stress-test-topic".to_string())
}
fn get_message_count() -> usize {
std::env::var("MESSAGE_COUNT")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(10000)
}
fn get_batch_size() -> usize {
std::env::var("BATCH_SIZE")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(500)
}
#[tokio::main]
async fn main() {
let _ = tracing_subscriber::fmt()
.with_env_filter(
tracing_subscriber::EnvFilter::try_from_default_env().unwrap_or_else(|_| "warn".into()),
)
.try_init();
let addrs = get_bootstrap_addrs();
let topic = get_topic_name();
let msg_count = get_message_count();
let batch_size = get_batch_size();
println!("============================================");
println!(" Kafka 生产-消费压力测试");
println!("============================================");
println!(" Bootstrap: {:?}", addrs);
println!(" Topic: {}", topic);
println!(" Partitions: 3");
println!(" Message 数: {}", msg_count);
println!(" Batch 大小: {}", batch_size);
println!("============================================\n");
print!("[1/5] 连接 Kafka...");
std::io::Write::flush(&mut std::io::stdout()).ok();
let client = match Client::builder(addrs)
.with_client_id("stress-test")
.with_metadata_ttl(Duration::from_secs(5))
.build()
.await
{
Ok(c) => c,
Err(e) => {
eprintln!(" 失败: {}", e);
std::process::exit(1);
}
};
println!(" ✓");
let cluster_info = client.admin().describe_cluster().await.unwrap();
let rf = (3).min(cluster_info.brokers.len()).max(1) as i16;
print!("[2/5] 创建 Topic '{}' (3 partitions, rf={})...", topic, rf);
std::io::Write::flush(&mut std::io::stdout()).ok();
let result = client
.admin()
.create_topic(&NewTopic::new(&topic, 3, rf))
.await
.unwrap();
use kafka_client::KafkaErrorCode;
if !result.error_code.is_ok() && result.error_code != KafkaErrorCode::TOPIC_ALREADY_EXISTS {
eprintln!(" 失败: {}", result.error_code);
std::process::exit(1);
}
println!(" ✓");
tokio::time::sleep(Duration::from_secs(2)).await;
client.refresh_metadata().await.expect("刷新 metadata 失败");
println!("\n[3/5] 开始生产 {} 条消息...", msg_count);
let producer_config = ProducerConfig::new()
.with_timeout(15000)
.with_batch_size(65536)
.with_linger(10);
let producer = client.producer(producer_config).await;
let warmup: Vec<ProducerRecord> = (0..100)
.map(|i| {
ProducerRecord::new(&topic, Bytes::from(format!("warmup-{}", i)))
.with_key(Bytes::from(format!("key-{}", i)))
})
.collect();
let _ = producer.send_batch(warmup).await;
producer.flush().await.unwrap();
let produce_start = Instant::now();
let mut total_produced = 0usize;
let mut batch_num = 0usize;
while total_produced < msg_count {
let remaining = msg_count - total_produced;
let this_batch = batch_size.min(remaining);
let batch: Vec<ProducerRecord> = (0..this_batch)
.map(|j| {
let seq = total_produced + j;
ProducerRecord::new(&topic, Bytes::from(format!("msg-{:06}", seq)))
.with_key(Bytes::from(format!("key-{:06}", seq)))
})
.collect();
match producer.send_batch(batch).await {
Ok(n) => total_produced += n,
Err(e) => eprintln!(" Batch {} 发送失败: {}", batch_num, e),
}
batch_num += 1;
if batch_num.is_multiple_of(20) {
print!(
"\r 已发送: {}/{} (batch #{})",
total_produced, msg_count, batch_num
);
std::io::Write::flush(&mut std::io::stdout()).ok();
}
}
producer.flush().await.expect("Flush 失败");
let produce_elapsed = produce_start.elapsed();
let produce_throughput = if produce_elapsed.as_secs_f64() > 0.0 {
(total_produced as f64) / produce_elapsed.as_secs_f64()
} else {
0.0
};
print!("\r 已发送: {}/{} ✓\n", total_produced, msg_count);
println!(" 生产耗时: {:.2?}", produce_elapsed);
println!(" 生产吞吐: {:.0} msgs/sec", produce_throughput);
println!("\n[4/5] 开始消费所有消息...");
let consumer_config = ConsumerConfig::new()
.with_group_id("stress-test-group")
.with_auto_commit_interval(Duration::from_secs(2))
.with_earliest()
.with_max_bytes(50 * 1024 * 1024)
.with_partition_max_bytes(10 * 1024 * 1024)
.with_max_wait(Duration::from_secs(3))
.with_session_timeout(Duration::from_secs(15));
let mut consumer = client.consumer(consumer_config);
consumer.subscribe(vec![topic.clone()]).await.unwrap();
let mut stream = consumer.into_stream();
let consume_start = Instant::now();
let mut consumed_ids = HashSet::new();
let mut total_consumed = 0usize;
let mut empty_polls = 0u32;
const MAX_EMPTY_POLLS: u32 = 20;
let deadline = Instant::now() + Duration::from_secs(120);
let mut last_progress_time = Instant::now();
while empty_polls < MAX_EMPTY_POLLS && Instant::now() < deadline {
match tokio::time::timeout(Duration::from_secs(3), stream.recv()).await {
Ok(Some(r)) => {
empty_polls = 0;
total_consumed += 1;
let value_str = String::from_utf8_lossy(&r.value);
if let Some(seq_str) = value_str.strip_prefix("msg-")
&& let Ok(seq) = seq_str.parse::<usize>()
{
consumed_ids.insert(seq);
}
let now = Instant::now();
if now - last_progress_time >= Duration::from_secs(1) {
let elapsed_secs = now - consume_start;
let throughput = total_consumed as f64 / elapsed_secs.as_secs_f64();
print!(
"\r 已消费: {} 条 (去重: {}) | {:.0} msgs/sec",
total_consumed,
consumed_ids.len(),
throughput,
);
std::io::Write::flush(&mut std::io::stdout()).ok();
last_progress_time = now;
}
}
Ok(None) => break,
Err(_) => {
empty_polls += 1; }
}
}
let consume_elapsed = consume_start.elapsed();
let consume_throughput = if consume_elapsed.as_secs_f64() > 0.0 {
(total_consumed as f64) / consume_elapsed.as_secs_f64()
} else {
0.0
};
print!(
"\r 已消费: {} 条消息 (去重后: {}) ✓\n",
total_consumed,
consumed_ids.len()
);
println!(" 消费耗时: {:.2?}", consume_elapsed);
println!(" 消费吞吐: {:.0} msgs/sec", consume_throughput);
println!("\n[5/5] 验证准确性...");
let expected_count = msg_count;
let mut errors = Vec::new();
if consumed_ids.len() != expected_count {
errors.push(format!(
"消息数量不匹配: 预期 {},实际 {} (差值: {})",
expected_count,
consumed_ids.len(),
if consumed_ids.len() > expected_count {
consumed_ids.len() - expected_count
} else {
expected_count - consumed_ids.len()
}
));
}
let mut missing = Vec::new();
for seq in 0..expected_count {
if !consumed_ids.contains(&seq) {
missing.push(seq);
if missing.len() >= 10 {
break;
}
}
}
if !missing.is_empty() {
errors.push(format!("缺失消息 (前10): {:?}", missing));
}
let mut unexpected: Vec<usize> = consumed_ids
.iter()
.filter(|&&seq| seq >= expected_count)
.copied()
.collect();
unexpected.sort();
if !unexpected.is_empty() {
errors.push(format!(
"非预期消息 (前10): {:?}",
&unexpected[..unexpected.len().min(10)]
));
}
println!("\n============================================");
println!(" 压力测试报告");
println!("============================================");
println!(" 消息总数: {}", msg_count);
println!(" 生产耗时: {:.2?}", produce_elapsed);
println!(
" 生产吞吐: {:.0} msgs/sec ({:.1} MB/s)",
produce_throughput,
produce_throughput * 64.0 / 1024.0 / 1024.0
);
println!(" 消费耗时: {:.2?}", consume_elapsed);
println!(
" 消费吞吐: {:.0} msgs/sec ({:.1} MB/s)",
consume_throughput,
consume_throughput * 64.0 / 1024.0 / 1024.0
);
println!(" 总耗时: {:.2?}", produce_elapsed + consume_elapsed);
println!(" 有效消息: {} (去重)", consumed_ids.len());
println!(" 总接收: {} (含 warmup)", total_consumed);
println!(" Warmup: 100 条");
println!("--------------------------------------------");
if errors.is_empty() {
println!(" ✅ 准确性: 完全正确,零丢失、零重复!");
} else {
println!(" ❌ 准确性: 发现 {} 个问题:", errors.len());
for e in &errors {
println!(" - {}", e);
}
}
println!("============================================");
print!("\n 关闭连接...");
std::io::Write::flush(&mut std::io::stdout()).ok();
drop(stream);
if let Err(e) = client.close().await {
eprintln!(" Client 关闭失败: {}", e);
}
println!(" ✓");
}