import pyarrow.parquet as pq import pytest from app.events import COLUMNS, ClientEvent, to_row from app.sink import HubParquetSink, LocalParquetSink, rows_to_parquet_bytes def rows(n): return [to_row(ClientEvent(event_type="search", model_family="Qwen", quantization="4-bit")) for _ in range(n)] def test_batch_becomes_one_shard(tmp_path): s = LocalParquetSink(tmp_path, flush_seconds=999) s.add(rows(7)) s.add(rows(5)) assert s.flush() shards = list((tmp_path / "data/events").rglob("*.parquet")) assert len(shards) == 1 t = pq.read_table(shards[0]) assert t.num_rows == 12 and t.column_names == list(COLUMNS) assert s.flush() and len(list((tmp_path / "data/events").rglob("*.parquet"))) == 1 # empty flush = no shard s.add(rows(1)) s.flush() assert len(list((tmp_path / "data/events").rglob("*.parquet"))) == 2 # appends, never rewrites assert len(s.read_existing()) == 13 class FlakyApi: def __init__(self, fail_times): self.fail_times = fail_times self.commits = [] self.token = None def create_commit(self, repo_id, repo_type, operations, commit_message): if self.fail_times > 0: self.fail_times -= 1 raise ConnectionError("hub down") assert repo_type == "dataset" self.commits.append([op.path_in_repo for op in operations]) def test_hub_failure_keeps_spool_and_retries(tmp_path): api = FlakyApi(fail_times=2) s = HubParquetSink("org/data", None, spool_dir=tmp_path, api=api, flush_seconds=999) s.add(rows(3)) assert not s.flush() and s.pending == 3 and not s.status()["healthy"] # a restart in between must not lose rows s2 = HubParquetSink("org/data", None, spool_dir=tmp_path, api=api, flush_seconds=999) assert s2.pending == 3 assert not s2.flush() assert s2.flush() and s2.pending == 0 and s2.status()["healthy"] assert len(api.commits) == 1 and len(api.commits[0]) == 1 assert api.commits[0][0].startswith("data/events/") and api.commits[0][0].endswith(".parquet") s3 = HubParquetSink("org/data", None, spool_dir=tmp_path, api=api, flush_seconds=999) assert s3.pending == 0 def test_local_refuses_overwrite(tmp_path): s = LocalParquetSink(tmp_path, flush_seconds=999) data = rows_to_parquet_bytes(rows(1)) s._write_shard("data/events/x.parquet", data) with pytest.raises(FileExistsError): s._write_shard("data/events/x.parquet", data)