1use crate::{
7 Event, Offset, Partition, SegmentId, SegmentManager, StorageError,
8 StorageResult, Topic, memory::PooledBuffer
9};
10use serde::{Deserialize, Serialize};
11use std::{
12 collections::{HashMap, VecDeque},
13 path::PathBuf,
14 sync::Arc,
15};
16use tokio::sync::RwLock;
17
18#[derive(Debug, Clone, Serialize, Deserialize)]
22pub struct WalConfig {
23 pub data_dir: PathBuf,
25 pub segment_size: u64,
27 pub sync_interval_ms: u64,
29}
30
31impl Default for WalConfig {
32 fn default() -> Self {
33 Self {
34 data_dir: PathBuf::from("./data"),
35 segment_size: 1024 * 1024 * 1024, sync_interval_ms: 1000,
37 }
38 }
39}
40
41#[derive(Debug, Clone)]
45struct PartitionOffset {
46 latest_offset: Offset,
48 current_segment: SegmentId,
50 local_offset: u64,
52}
53
54#[derive(Debug)]
71pub struct WriteAheadLog {
72 segment_manager: Arc<SegmentManager>,
74 partition_offsets: RwLock<HashMap<(Topic, Partition), PartitionOffset>>,
76 current_write_segment: RwLock<Option<SegmentId>>,
78 config: WalConfig,
80 write_buffer: RwLock<VecDeque<Event>>,
82}
83
84impl WriteAheadLog {
85 pub async fn new(config: WalConfig) -> StorageResult<Self> {
98 let segment_manager = Arc::new(SegmentManager::new(&config.data_dir).await?);
99
100 let wal = Self {
101 segment_manager,
102 partition_offsets: RwLock::new(HashMap::new()),
103 current_write_segment: RwLock::new(None),
104 config,
105 write_buffer: RwLock::new(VecDeque::new()),
106 };
107
108 wal.recover_state().await?;
110
111 Ok(wal)
112 }
113
114 async fn recover_state(&self) -> StorageResult<()> {
121 tracing::info!("Recovering WAL state from existing segments");
122
123 tracing::info!("WAL state recovery completed");
127 Ok(())
128 }
129
130 pub async fn append(&mut self, event: &Event) -> StorageResult<Offset> {
143 let topic_partition = (event.topic.clone(), event.partition);
144
145 let segment_id = self.ensure_write_segment().await?;
147
148 let segment_arc = self.segment_manager
150 .get_segment(segment_id)
151 .await
152 .ok_or_else(|| StorageError::SegmentNotFound { segment_id: segment_id.0 })?;
153
154 let mut segment = segment_arc.write().await;
155 let local_offset = segment.append_event(event)?;
156
157 let mut partition_offsets = self.partition_offsets.write().await;
159 let global_offset = match partition_offsets.get_mut(&topic_partition) {
160 Some(partition_offset) => {
161 partition_offset.latest_offset = partition_offset.latest_offset.next();
162 partition_offset.local_offset = *local_offset;
163 partition_offset.latest_offset
164 }
165 None => {
166 let new_offset = Offset::new(0);
167 partition_offsets.insert(topic_partition, PartitionOffset {
168 latest_offset: new_offset,
169 current_segment: segment_id,
170 local_offset: *local_offset,
171 });
172 new_offset
173 }
174 };
175
176 tracing::debug!(
177 "Event appended: topic={}, partition={}, offset={}, segment={}",
178 event.topic.as_str(),
179 event.partition.0,
180 global_offset.0,
181 segment_id
182 );
183
184 Ok(global_offset)
185 }
186
187 pub async fn read_events(
204 &self,
205 topic: &Topic,
206 partition: Partition,
207 start_offset: Offset,
208 max_events: usize,
209 ) -> StorageResult<Vec<Event>> {
210 let topic_partition = (topic.clone(), partition);
211
212 let partition_offsets = self.partition_offsets.read().await;
214 let partition_offset = partition_offsets.get(&topic_partition);
215
216 if partition_offset.is_none() || start_offset > partition_offset.unwrap().latest_offset {
217 return Ok(Vec::new());
218 }
219
220 let current_segment_id = partition_offset.unwrap().current_segment;
222 let segment_arc = self.segment_manager
223 .get_segment(current_segment_id)
224 .await
225 .ok_or_else(|| StorageError::SegmentNotFound { segment_id: current_segment_id.0 })?;
226
227 let segment = segment_arc.read().await;
228
229 let segment_stats = segment.stats();
232 let all_events = segment.read_events_range(Offset::new(0), segment_stats.entry_count as usize)?;
233
234 let mut partition_events: Vec<Event> = all_events.into_iter()
236 .filter(|event| event.topic == *topic && event.partition == partition)
237 .collect();
238
239 for (idx, event) in partition_events.iter_mut().enumerate() {
241 event.set_offset(Offset::new(idx as u64));
242 }
243
244 let start_idx = start_offset.0 as usize;
246 let end_idx = (start_idx + max_events).min(partition_events.len());
247
248 let filtered_events = if start_idx < partition_events.len() {
249 partition_events[start_idx..end_idx].to_vec()
250 } else {
251 Vec::new()
252 };
253
254 tracing::debug!(
255 "Read {} events: topic={}, partition={}, start_offset={}, found_total={}",
256 filtered_events.len(),
257 topic.as_str(),
258 partition.0,
259 start_offset.0,
260 partition_events.len()
261 );
262
263 Ok(filtered_events)
264 }
265
266 pub async fn get_latest_offset(
275 &self,
276 topic: &Topic,
277 partition: Partition,
278 ) -> StorageResult<Option<Offset>> {
279 let topic_partition = (topic.clone(), partition);
280 let partition_offsets = self.partition_offsets.read().await;
281
282 Ok(partition_offsets
283 .get(&topic_partition)
284 .map(|offset_info| offset_info.latest_offset))
285 }
286
287 async fn ensure_write_segment(&self) -> StorageResult<SegmentId> {
299 let mut current_segment = self.current_write_segment.write().await;
300
301 if let Some(segment_id) = *current_segment {
302 if let Some(segment_arc) = self.segment_manager.get_segment(segment_id).await {
304 let segment = segment_arc.read().await;
305 let stats = segment.stats();
306
307 if stats.used_size < stats.total_size * 9 / 10 { return Ok(segment_id);
310 }
311 }
312 }
313
314 let new_segment_id = self.segment_manager.create_segment().await?;
316 *current_segment = Some(new_segment_id);
317
318 tracing::info!("Created new write segment: {}", new_segment_id);
319
320 Ok(new_segment_id)
321 }
322
323 pub async fn sync(&self) -> StorageResult<()> {
333 if let Some(segment_id) = *self.current_write_segment.read().await {
335 if let Some(segment_arc) = self.segment_manager.get_segment(segment_id).await {
336 let mut segment = segment_arc.write().await;
337 segment.flush()?;
338 }
339 }
340
341 tracing::debug!("WAL sync completed");
342 Ok(())
343 }
344
345 pub async fn get_stats(&self) -> WalStats {
350 let partition_offsets = self.partition_offsets.read().await;
351 let segment_stats = self.segment_manager.get_all_stats().await;
352
353 WalStats {
354 total_partitions: partition_offsets.len(),
355 total_segments: segment_stats.len(),
356 total_events: partition_offsets.values().map(|p| p.latest_offset.0 + 1).sum(),
357 total_size_bytes: segment_stats.iter().map(|s| s.used_size).sum(),
358 }
359 }
360
361 pub async fn append_batch(
373 &mut self,
374 events: &[Event],
375 buffer: &mut PooledBuffer
376 ) -> StorageResult<Vec<Offset>> {
377 if events.is_empty() {
378 return Ok(Vec::new());
379 }
380
381 let mut results = Vec::with_capacity(events.len());
382 buffer.clear();
383
384 let mut partition_groups: HashMap<(Topic, Partition), Vec<&Event>> = HashMap::new();
386 for event in events {
387 let key = (event.topic.clone(), event.partition);
388 partition_groups.entry(key).or_default().push(event);
389 }
390
391 for ((topic, partition), partition_events) in partition_groups {
393 let partition_offsets = self.append_partition_batch(
394 &topic,
395 partition,
396 &partition_events,
397 buffer
398 ).await?;
399 results.extend(partition_offsets);
400 }
401
402 tracing::debug!(
403 "Batch append completed: {} events across {} partitions",
404 events.len(),
405 results.len()
406 );
407
408 Ok(results)
409 }
410
411 async fn append_partition_batch(
413 &mut self,
414 topic: &Topic,
415 partition: Partition,
416 events: &[&Event],
417 buffer: &mut PooledBuffer,
418 ) -> StorageResult<Vec<Offset>> {
419 let topic_partition = (topic.clone(), partition);
420 let mut results = Vec::with_capacity(events.len());
421
422 let segment_id = self.ensure_write_segment().await?;
424
425 let segment_arc = self.segment_manager
427 .get_segment(segment_id)
428 .await
429 .ok_or_else(|| StorageError::SegmentNotFound { segment_id: segment_id.0 })?;
430
431 let mut segment = segment_arc.write().await;
432
433 buffer.clear();
435 let mut event_sizes = Vec::with_capacity(events.len());
436
437 for event in events {
438 let event_data = bincode::serialize(event)?;
439 event_sizes.push(event_data.len());
440 buffer.write(&event_data)?;
441 }
442
443 let start_offset = segment.append_batch_data(buffer.data(), &event_sizes)?;
445
446 for i in 0..events.len() {
448 let offset = Offset::new(start_offset.0 + i as u64);
449 results.push(offset);
450 }
451
452 let mut partition_offsets = self.partition_offsets.write().await;
454 if let Some(last_offset) = results.last() {
455 let partition_offset = PartitionOffset {
456 latest_offset: *last_offset,
457 current_segment: segment_id,
458 local_offset: segment.stats().used_size,
459 };
460 partition_offsets.insert(topic_partition, partition_offset);
461 }
462
463 tracing::debug!(
464 "Partition batch written: topic={}, partition={}, events={}, offsets={:?}",
465 topic.as_str(),
466 partition.0,
467 events.len(),
468 results.first().zip(results.last())
469 );
470
471 Ok(results)
472 }
473
474 pub async fn read_events_optimized(
486 &self,
487 topic: &Topic,
488 partition: Partition,
489 start_offset: Offset,
490 max_events: usize,
491 buffer: &mut PooledBuffer,
492 ) -> StorageResult<Vec<Event>> {
493 buffer.clear();
494
495 let segments = self.segment_manager.list_segments().await;
497 let mut all_events = Vec::new();
498
499 let mut read_tasks: Vec<tokio::task::JoinHandle<StorageResult<Vec<Event>>>> = Vec::new();
501
502 for segment_id in segments {
503 if all_events.len() >= max_events {
504 break;
505 }
506
507 let segment_arc = match self.segment_manager.get_segment(segment_id).await {
508 Some(segment) => segment,
509 None => continue,
510 };
511
512 let segment = segment_arc.read().await;
513 let segment_events = segment.read_events_range(
514 start_offset,
515 max_events - all_events.len()
516 )?;
517
518 let filtered_events: Vec<Event> = segment_events
520 .into_iter()
521 .filter(|event| event.topic == *topic && event.partition == partition)
522 .take(max_events - all_events.len())
523 .collect();
524
525 all_events.extend(filtered_events);
526 }
527
528 all_events.sort_by_key(|event| event.offset);
530
531 let start_idx = start_offset.0 as usize;
533 if start_idx < all_events.len() {
534 let end_idx = (start_idx + max_events).min(all_events.len());
535 Ok(all_events[start_idx..end_idx].to_vec())
536 } else {
537 Ok(Vec::new())
538 }
539 }
540}
541
542#[derive(Debug, Clone, Serialize, Deserialize)]
544pub struct WalStats {
545 pub total_partitions: usize,
547 pub total_segments: usize,
549 pub total_events: u64,
551 pub total_size_bytes: u64,
553}