use std::io;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc::{self, RecvError, SyncSender, TryRecvError, TrySendError};
use std::sync::Arc;
use std::thread::{self, JoinHandle};
use std::time::Duration;
use crate::serve::kv_persist::block_store::{DiskBlockStore, WriteJob};
pub struct AsyncWriterHandle {
tx: Option<SyncSender<WriteJob>>,
join_handle: Option<JoinHandle<()>>,
pending: Arc<AtomicUsize>,
}
impl AsyncWriterHandle {
pub fn spawn(store: Arc<DiskBlockStore>, channel_capacity: usize) -> Self {
let (tx, rx) = mpsc::sync_channel::<WriteJob>(channel_capacity);
let pending = Arc::new(AtomicUsize::new(0));
let pending_for_worker = Arc::clone(&pending);
let join_handle = thread::Builder::new()
.name("hf2q-kv-writer".to_string())
.spawn(move || {
worker_loop(store, rx, pending_for_worker);
})
.expect("spawn hf2q-kv-writer thread");
Self {
tx: Some(tx),
join_handle: Some(join_handle),
pending,
}
}
#[allow(clippy::result_large_err)]
pub fn enqueue(&self, job: WriteJob) -> Result<(), TrySendError<WriteJob>> {
match self.tx.as_ref() {
Some(tx) => match tx.try_send(job) {
Ok(()) => {
self.pending.fetch_add(1, Ordering::Relaxed);
Ok(())
}
Err(e) => Err(e),
},
None => Err(TrySendError::Disconnected(job)),
}
}
#[allow(clippy::result_large_err)]
pub fn enqueue_blocking(&self, job: WriteJob) -> Result<(), mpsc::SendError<WriteJob>> {
match self.tx.as_ref() {
Some(tx) => match tx.send(job) {
Ok(()) => {
self.pending.fetch_add(1, Ordering::Relaxed);
Ok(())
}
Err(e) => Err(e),
},
None => Err(mpsc::SendError(job)),
}
}
pub fn pending_jobs(&self) -> usize {
self.pending.load(Ordering::Relaxed)
}
pub fn shutdown(mut self) -> io::Result<()> {
self.tx.take();
if let Some(jh) = self.join_handle.take() {
jh.join().map_err(|panic_payload| {
io::Error::other(format!(
"kv_persist writer worker panicked: {:?}",
panic_payload
.downcast_ref::<&str>()
.copied()
.or_else(|| { panic_payload.downcast_ref::<String>().map(|s| s.as_str()) })
.unwrap_or("<non-string panic payload>")
))
})?;
}
Ok(())
}
}
impl Drop for AsyncWriterHandle {
fn drop(&mut self) {
self.tx.take();
if let Some(jh) = self.join_handle.take() {
let _ = jh.join();
}
}
}
fn worker_loop(
store: Arc<DiskBlockStore>,
rx: mpsc::Receiver<WriteJob>,
pending: Arc<AtomicUsize>,
) {
loop {
let job = match rx.recv() {
Ok(j) => j,
Err(RecvError) => {
break;
}
};
process_job(&store, job, &pending);
loop {
match rx.try_recv() {
Ok(j) => process_job(&store, j, &pending),
Err(TryRecvError::Empty) => break,
Err(TryRecvError::Disconnected) => return,
}
}
}
}
fn process_job(store: &DiskBlockStore, job: WriteJob, pending: &Arc<AtomicUsize>) {
let WriteJob {
header,
body,
completion_tx,
} = job;
let result = store.write_block_sync(&header, &body).map(|_| ());
if let Err(ref e) = result {
tracing::warn!(
target: "hf2q::kv_persist::writer",
error = %e,
block_hash = %header.block_hash,
"kv_persist writer: write_block_sync failed; continuing"
);
}
if let Some(tx) = completion_tx {
let _ = tx.try_send(result);
}
if let Err(e) = store.evict_lru_until_under_budget(|_| false) {
tracing::warn!(
target: "hf2q::kv_persist::writer",
error = %e,
"kv_persist writer: post-write budget eviction failed; continuing"
);
}
pending.fetch_sub(1, Ordering::Relaxed);
}
pub fn completion_channel() -> (SyncSender<io::Result<()>>, mpsc::Receiver<io::Result<()>>) {
mpsc::sync_channel::<io::Result<()>>(1)
}
pub const DEFAULT_CHANNEL_CAPACITY: usize = 8;
pub const DEFAULT_COMPLETION_TIMEOUT: Duration = Duration::from_secs(5);
#[cfg(test)]
mod tests {
use super::*;
use crate::serve::kv_persist::block_store::DiskBlockStore;
use crate::serve::kv_persist::format::{
compute_model_fingerprint, BlockHash, EnvelopeHeader, ModelFingerprint, ParentBlockHash,
BLOCK_TOKENS, CURRENT_FORMAT_VERSION,
};
use sha2::{Digest, Sha256};
use std::path::PathBuf;
use std::process;
use std::sync::atomic::{AtomicU32, Ordering};
use std::time::SystemTime;
fn temp_dir(label: &str) -> PathBuf {
static COUNTER: AtomicU32 = AtomicU32::new(0);
let n = COUNTER.fetch_add(1, Ordering::SeqCst);
let pid = process::id();
let nanos = SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
let dir = std::env::temp_dir().join(format!("hf2q-kv-writer-{label}-{pid}-{nanos}-{n}"));
std::fs::create_dir_all(&dir).expect("temp_dir mkdir");
dir
}
fn fixture_fp(seed: &str) -> ModelFingerprint {
compute_model_fingerprint(
seed,
"Q4_0",
"hf2q-test-1.0.0",
"deadbeefcafebabe1122334455667788",
"<|im_start|>...<|im_end|>",
)
}
fn make_block(
fp: ModelFingerprint,
parent: ParentBlockHash,
seed: u32,
) -> (Vec<u8>, EnvelopeHeader) {
let body: Vec<u8> = (0..512u32)
.flat_map(|i| (i.wrapping_add(seed)).to_le_bytes())
.collect();
let mut h = Sha256::new();
h.update(&body);
let bh: [u8; 32] = h.finalize().into();
let header = EnvelopeHeader {
format_version: CURRENT_FORMAT_VERSION.0,
model_fingerprint: fp,
block_hash: BlockHash(bh),
parent_block_hash: parent,
payload_kind: "kv-dense-bf16".into(),
codec_version: 1,
n_tokens: BLOCK_TOKENS,
};
(body, header)
}
#[test]
fn spawn_then_shutdown_clean() {
let dir = temp_dir("spawnshut");
let store = Arc::new(DiskBlockStore::new(dir.clone(), 0).expect("new"));
let handle = AsyncWriterHandle::spawn(Arc::clone(&store), 8);
handle.shutdown().expect("clean shutdown");
assert_eq!(store.index().block_count(), 0, "no work, no blocks");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn enqueue_then_shutdown_drains_pending_jobs() {
let dir = temp_dir("drain");
let store = Arc::new(DiskBlockStore::new(dir.clone(), 0).expect("new"));
let handle = AsyncWriterHandle::spawn(Arc::clone(&store), 4);
let fp = fixture_fp("drain");
let mut hashes: Vec<BlockHash> = Vec::new();
for s in 0u32..10 {
let (body, header) = make_block(fp, ParentBlockHash(None), s);
hashes.push(header.block_hash);
let job = WriteJob {
header,
body,
completion_tx: None,
};
handle.enqueue_blocking(job).expect("enqueue");
}
handle.shutdown().expect("shutdown drains");
assert_eq!(store.index().block_count(), 10, "all 10 drained");
for h in &hashes {
assert!(store.index().lookup(h).is_some(), "hash indexed");
let p = store.block_path(&fp, h);
assert!(p.exists(), "file at {} present", p.display());
}
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn enqueue_completion_ack_fires() {
let dir = temp_dir("ack");
let store = Arc::new(DiskBlockStore::new(dir.clone(), 0).expect("new"));
let handle = AsyncWriterHandle::spawn(Arc::clone(&store), 4);
let fp = fixture_fp("ack");
let (body, header) = make_block(fp, ParentBlockHash(None), 1);
let body_clone = body.clone();
let block_hash = header.block_hash;
let (ack_tx, ack_rx) = completion_channel();
let job = WriteJob {
header,
body,
completion_tx: Some(ack_tx),
};
handle.enqueue_blocking(job).expect("enqueue");
let result = ack_rx
.recv_timeout(DEFAULT_COMPLETION_TIMEOUT)
.expect("completion received within timeout");
result.expect("write succeeded");
let body_back = store.read_block(&block_hash).expect("read");
assert_eq!(body_back, body_clone, "body bytes round-trip via writer");
handle.shutdown().expect("shutdown");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn worker_does_not_panic_on_io_error() {
let dir = temp_dir("ioerr");
let store = Arc::new(DiskBlockStore::new(dir.clone(), 0).expect("new"));
store.set_max_block_bytes_override(1); let handle = AsyncWriterHandle::spawn(Arc::clone(&store), 4);
let fp = fixture_fp("ioerr");
let (body_a, header_a) = make_block(fp, ParentBlockHash(None), 1);
let (ack_a_tx, ack_a_rx) = completion_channel();
handle
.enqueue_blocking(WriteJob {
header: header_a,
body: body_a,
completion_tx: Some(ack_a_tx),
})
.expect("enqueue a");
let r_a = ack_a_rx
.recv_timeout(DEFAULT_COMPLETION_TIMEOUT)
.expect("ack a");
assert!(r_a.is_err(), "first job reported error");
store.set_max_block_bytes_override(0); let (body_b, header_b) = make_block(fp, ParentBlockHash(None), 2);
let block_hash_b = header_b.block_hash;
let (ack_b_tx, ack_b_rx) = completion_channel();
handle
.enqueue_blocking(WriteJob {
header: header_b,
body: body_b,
completion_tx: Some(ack_b_tx),
})
.expect("enqueue b");
let r_b = ack_b_rx
.recv_timeout(DEFAULT_COMPLETION_TIMEOUT)
.expect("ack b");
r_b.expect("second job succeeded — worker survived previous error");
assert!(store.index().lookup(&block_hash_b).is_some());
handle.shutdown().expect("shutdown");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn enqueue_returns_full_error_when_channel_capacity_reached() {
let (tx, _rx) = mpsc::sync_channel::<WriteJob>(1);
let dir = temp_dir("full");
let store = DiskBlockStore::new(dir.clone(), 0).expect("new");
let fp = fixture_fp("full");
let (body_a, header_a) = make_block(fp, ParentBlockHash(None), 1);
let (body_b, header_b) = make_block(fp, ParentBlockHash(None), 2);
tx.try_send(WriteJob {
header: header_a,
body: body_a,
completion_tx: None,
})
.expect("first try_send fits");
let err = tx
.try_send(WriteJob {
header: header_b,
body: body_b,
completion_tx: None,
})
.expect_err("must fail");
match err {
TrySendError::Full(_) => {}
other => panic!("expected Full, got {other:?}"),
}
drop(tx);
drop(store);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn enqueue_after_shutdown_returns_disconnected() {
let dir = temp_dir("disc");
let store = Arc::new(DiskBlockStore::new(dir.clone(), 0).expect("new"));
let handle = AsyncWriterHandle::spawn(Arc::clone(&store), 1);
handle.shutdown().expect("shutdown");
let _ = std::fs::remove_dir_all(&dir);
}
}