File size: 4,962 Bytes
34f3bc9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
"""Disk-persistent scan-cache infrastructure shared by the v2.1 and v3.0
adapters.

``lerobot_v21`` re-exports the names defined here for any out-of-tree caller.
"""
from __future__ import annotations

import json
import logging
from pathlib import Path
from typing import Optional

logger = logging.getLogger(__name__)

_SCAN_CACHE_FILE = ".labvla_scan_cache.json"


def _shard_fingerprint(data_root: Optional[Path]) -> dict:
    """Summarize the data/ shard parquet files by count + total size + max mtime.

    A cheap content fingerprint so the scan cache invalidates on in-place shard
    repair/replacement that keeps the episode list unchanged (keying only on
    episodes.jsonl mtime missed those). O(num_shards) ``stat()`` calls, no
    parquet open. Returns zeros when ``data_root`` is absent (degrades to the
    previous episode-list-only behavior).
    """
    if data_root is None:
        return {"shard_count": 0, "shard_total_size": 0, "shard_max_mtime": 0.0}
    count = 0
    total_size = 0
    max_mtime = 0.0
    if data_root.is_dir():
        for p in data_root.rglob("*.parquet"):
            try:
                st = p.stat()
            except OSError:
                continue
            count += 1
            total_size += int(st.st_size)
            if st.st_mtime > max_mtime:
                max_mtime = st.st_mtime
    return {
        "shard_count": count,
        "shard_total_size": total_size,
        "shard_max_mtime": max_mtime,
    }


def _scan_cache_key(
    meta_root: Path,
    schema_id: str,
    chunks_size: int,
    data_root: Optional[Path] = None,
) -> dict:
    """Fingerprint used to invalidate the cached scan.

    Invalidated when any of these change from the stored key:
      - schema_id: a different schema filters different episodes
      - chunks_size: repartitioning changes chunk_ok semantics
      - episodes_jsonl_mtime: the canonical "episode list changed" signal
      - shard_*: count/total-size/max-mtime of data/**/*.parquet, so the cache
        also invalidates when shard CONTENTS change even though the episode list
        is identical (v3 has no episodes.jsonl and relies entirely on this).
    """
    ep_jsonl = meta_root / "episodes.jsonl"
    from src.utils import env_flags as _env_flags
    return {
        "schema_id": str(schema_id),
        "chunks_size": int(chunks_size),
        # These env gates change which episodes/shards the integrity filter
        # accepts; flipping them must invalidate the cache. Read via the
        # env_flags registry (not raw os.environ) so an unset flag and an
        # explicit "=0" resolve to the same default rather than colliding with
        # opposite scan behavior.
        "allow_truncate": _env_flags.get("LABVLA_ALLOW_TRUNCATE") == "1",
        "v21_validate_per_file": str(_env_flags.get("LABVLA_V21_VALIDATE_PER_FILE")),
        "episodes_jsonl_mtime": (
            ep_jsonl.stat().st_mtime if ep_jsonl.exists() else 0.0
        ),
        **_shard_fingerprint(data_root),
    }


def _load_scan_cache(meta_root: Path, expected_key: dict) -> Optional[tuple[dict, set]]:
    """Return (chunk_ok_map, existing_episodes) if cache is valid, else None.

    Cache file is `meta/.labvla_scan_cache.json`; hidden to avoid cluttering.
    Bypass by setting env `LABVLA_SCAN_CACHE=0`.
    """
    from src.utils import env_flags as _env_flags
    if _env_flags.get("LABVLA_SCAN_CACHE") == "0":
        return None
    cache_path = meta_root / _SCAN_CACHE_FILE
    if not cache_path.exists():
        return None
    try:
        cached = json.loads(cache_path.read_text())
    except Exception:
        return None
    # MA2: expected_key may contain list[tuple] (declared_dims); the cached key
    # round-tripped through JSON is list[list]. Compare canonical JSON so tuples
    # and lists match instead of always reporting a (false) cache miss.
    if json.dumps(cached.get("key"), sort_keys=True) != json.dumps(expected_key, sort_keys=True):
        return None
    # json keys are strings → restore int dict
    chunk_ok = {int(k): bool(v) for k, v in cached.get("chunk_ok", {}).items()}
    existing = set(int(x) for x in cached.get("existing_episodes", []))
    return chunk_ok, existing


def _save_scan_cache(
    meta_root: Path,
    key: dict,
    chunk_ok: dict[int, bool],
    existing_episodes: set[int],
) -> None:
    """Atomically write scan result. Failure is non-fatal (just log)."""
    cache_path = meta_root / _SCAN_CACHE_FILE
    tmp_path = cache_path.with_suffix(cache_path.suffix + ".tmp")
    payload = {
        "key": key,
        "chunk_ok": {str(k): bool(v) for k, v in chunk_ok.items()},
        "existing_episodes": sorted(int(x) for x in existing_episodes),
    }
    try:
        tmp_path.write_text(json.dumps(payload))
        import os
        os.replace(tmp_path, cache_path)
    except Exception as e:
        logger.warning("[scan-cache] failed to write scan cache at %s: %s", cache_path, e)