Codeseys/composer-replication-framework
0
1"""s3_contract.py — THE single dataset layout + manifest (finding V8/D-7/D-8).2 3Supersedes BOTH prior contracts: design-F1's `runs/<id>/{sft_corpus,dpo_pairs,4rl_task_pool,divergence_pairs,wm_tuples,holdout,diloco_rendezvous}` and5design-F2's `{traces,tasks,replay,task_grades,corpus}/v1/run_id=<id>` — the two6were never reconciled and coexisted in the grounding doc. One layout, one7manifest, two explicit serializers with a unit-tested leak guard.8 9Deliberate exclusions from the run layout:10 * `diloco_rendezvous/` — training-comms state, not dataset; lives in its own11 prefix/bucket (finding D-19).12 * `wm_tuples/` — emitted only when the P4 world-model ablation is scheduled13 (finding D-14); not part of Stage 0.14 15Layout (root = any local path or fsspec URI):16 <root>/runs/<run_id>/17 tasks/manifest.jsonl policy-safe task rows (golden_diff -> sha256)18 tasks_full/manifest.jsonl construction-side full rows (RESTRICTED prefix)19 traj/*.jsonl CanonicalTrajectory records (audit trail)20 corpus_sft/rows.jsonl admitted SFT rows (to_policy_row output)21 corpus_dpo/rows.jsonl DPO-candidate rows22 holdout/tasks.jsonl held-out task ids+rows (never rolled out)23 quarantine/*.jsonl rejected trajectories w/ reasons (audit)24 manifest.json RunManifest25 DATASET_CARD.md human-readable card26"""27from __future__ import annotations28 29import dataclasses30import hashlib31import json32from dataclasses import dataclass, field33from typing import IO, Iterable34 35from composer_replication.datagen.schema import FeatureDeletionTask36 37SCHEMA_VERSION = "1"38 39 40def _is_local(root: str) -> bool:41 return "://" not in root or root.startswith("file://")42 43 44def _open(path: str, mode: str = "w") -> IO[str]:45 """Open a path for text IO; plain `open` locally, fsspec for s3:// etc.46 47 fsspec is lazy so the module (and all local-corpus runs) need no extra dep.48 """49 if _is_local(path):50 import os51 local = path.removeprefix("file://")52 os.makedirs(os.path.dirname(local), exist_ok=True)53 return open(local, mode, encoding="utf-8")54 try:55 import fsspec # noqa: PLC0415 — lazy heavy dep56 except ImportError as e:57 raise RuntimeError(58 "Non-local corpus roots require fsspec; install with "59 "`pip install -e .[serverless]`. Got: " + repr(e)60 ) from e61 return fsspec.open(path, mode, encoding="utf-8").open()62 63 64def _exists(path: str) -> bool:65 if _is_local(path):66 import os67 return os.path.exists(path.removeprefix("file://"))68 import fsspec # noqa: PLC041569 fs, _, paths = fsspec.get_fs_token_paths(path)70 return bool(fs.exists(paths[0]))71 72 73@dataclass(frozen=True)74class RunLayout:75 """Pure-path logic for one run's prefixes — testable without any IO."""76 77 root: str78 run_id: str79 80 def __post_init__(self) -> None:81 # Defense-in-depth (Wave-21 review P2): run_id is operator-supplied,82 # but a separator or `..` would silently escape the corpus root.83 if not self.run_id or "/" in self.run_id or "\\" in self.run_id \84 or ".." in self.run_id:85 raise ValueError(86 f"run_id {self.run_id!r} must be a single non-empty path "87 "segment (no separators, no '..')."88 )89 90 def _p(self, *parts: str) -> str:91 base = self.root.rstrip("/")92 return f"{base}/runs/{self.run_id}/" + "/".join(parts)93 94 @property95 def tasks_path(self) -> str:96 return self._p("tasks", "manifest.jsonl")97 98 @property99 def tasks_full_path(self) -> str:100 # RESTRICTED prefix: carries golden_diff/deleted_symbols. On S3 this101 # prefix gets a deny-by-default policy; locally it is still separated102 # so a naive `corpus_*` glob can never sweep it up.103 return self._p("tasks_full", "manifest.jsonl")104 105 @property106 def traj_path(self) -> str:107 return self._p("traj", "trajectories.jsonl")108 109 @property110 def sft_path(self) -> str:111 return self._p("corpus_sft", "rows.jsonl")112 113 @property114 def dpo_path(self) -> str:115 return self._p("corpus_dpo", "rows.jsonl")116 117 @property118 def holdout_path(self) -> str:119 return self._p("holdout", "tasks.jsonl")120 121 @property122 def quarantine_path(self) -> str:123 return self._p("quarantine", "rejected.jsonl")124 125 @property126 def manifest_path(self) -> str:127 return self._p("manifest.json")128 129 @property130 def card_path(self) -> str:131 return self._p("DATASET_CARD.md")132 133 134@dataclass135class RunManifest:136 """Run-level metadata: counts, cost, lineage, budget, acceptance status.137 138 `created_at` is CALLER-passed (never datetime.now() in here) so manifests139 are reproducible in tests. `parent_run_id` threads flywheel lineage so140 cross-generation dedup (finding D-12) can find prior signatures.141 """142 143 run_id: str144 created_at: str145 source: str = ""146 counts: dict = field(default_factory=dict)147 cost_usd: float = 0.0148 parent_run_id: str | None = None149 schema_version: str = SCHEMA_VERSION150 status: str = "building" # building | accepted | rejected | partial151 budget_usd: float | None = None152 153 def spend(self, usd: float) -> None:154 self.cost_usd += usd155 156 @property157 def over_budget(self) -> bool:158 return self.budget_usd is not None and self.cost_usd >= self.budget_usd159 160 def write(self, layout: RunLayout) -> None:161 with _open(layout.manifest_path) as f:162 json.dump(dataclasses.asdict(self), f, indent=2)163 164 @classmethod165 def read(cls, layout: RunLayout) -> RunManifest:166 with _open(layout.manifest_path, "r") as f:167 return cls(**json.load(f))168 169 170# ---------------------------------------------------------------------171# Writers — the leak guard lives here (finding D-8)172# ---------------------------------------------------------------------173 174 175def _task_row_policy_safe(task: FeatureDeletionTask) -> dict:176 """Task row with the construction-side secrets REPLACED, not just hidden.177 178 `asdict()` includes `golden_diff` despite `repr=False` — that is exactly179 the leak D-8 flagged. We keep provenance via a sha256 (verifiable, not180 recoverable) and drop `deleted_symbols` entirely (they name the answer).181 """182 row = dataclasses.asdict(task)183 gold = row.pop("golden_diff", "")184 row.pop("deleted_symbols", None)185 row["golden_diff_sha256"] = hashlib.sha256(gold.encode()).hexdigest() if gold else ""186 return row187 188 189def write_tasks(layout: RunLayout, tasks: Iterable[FeatureDeletionTask]) -> int:190 """Write the POLICY-SAFE task manifest (the default everything reads)."""191 n = 0192 with _open(layout.tasks_path) as f:193 for t in tasks:194 f.write(json.dumps(_task_row_policy_safe(t)) + "\n")195 n += 1196 return n197 198 199def write_tasks_full(layout: RunLayout, tasks: Iterable[FeatureDeletionTask]) -> int:200 """Write FULL task rows (incl. golden_diff) to the RESTRICTED prefix.201 202 Only the validator/monitor side reads this; never corpus consumers.203 """204 n = 0205 with _open(layout.tasks_full_path) as f:206 for t in tasks:207 f.write(json.dumps(dataclasses.asdict(t)) + "\n")208 n += 1209 return n210 211 212def _write_jsonl(path: str, rows: Iterable[dict]) -> int:213 n = 0214 with _open(path) as f:215 for r in rows:216 f.write(json.dumps(r) + "\n")217 n += 1218 return n219 220 221def write_sft_rows(layout: RunLayout, rows: Iterable[dict]) -> int:222 return _write_jsonl(layout.sft_path, rows)223 224 225def write_dpo_rows(layout: RunLayout, rows: Iterable[dict]) -> int:226 return _write_jsonl(layout.dpo_path, rows)227 228 229def write_quarantine(layout: RunLayout, rows: Iterable[dict]) -> int:230 return _write_jsonl(layout.quarantine_path, rows)231 232 233def write_holdout(layout: RunLayout, tasks: Iterable[FeatureDeletionTask]) -> int:234 return _write_jsonl(layout.holdout_path, (_task_row_policy_safe(t) for t in tasks))235 236 237def write_trajectories(layout: RunLayout, rows: Iterable[dict]) -> int:238 return _write_jsonl(layout.traj_path, rows)239 240 241def write_dataset_card(layout: RunLayout, manifest: RunManifest,242 *, license_tiers: dict[str, int] | None = None,243 dedup_stats: dict | None = None,244 decontamination_note: str = "") -> None:245 """A small human-readable dataset card (finding D-18)."""246 lines = [247 f"# Dataset card — run `{manifest.run_id}`",248 "",249 f"- **created:** {manifest.created_at}",250 f"- **source:** {manifest.source}",251 f"- **status:** {manifest.status}",252 f"- **schema_version:** {manifest.schema_version}",253 f"- **cost (USD):** {manifest.cost_usd:.2f}"254 + (f" / budget {manifest.budget_usd:.2f}" if manifest.budget_usd else ""),255 f"- **lineage:** parent_run_id={manifest.parent_run_id or 'none'}",256 "",257 "## Counts",258 "",259 ]260 for k, v in sorted(manifest.counts.items()):261 lines.append(f"- {k}: {v}")262 if license_tiers:263 lines += ["", "## License tiers seen", ""]264 lines += [f"- {k}: {v}" for k, v in sorted(license_tiers.items())]265 lines += ["", "## Decontamination", "",266 decontamination_note or267 "All source repos checked against the SWE-bench-family eval list "268 "(datagen.repo_gate.DECONTAMINATION_LIST) at ingest."]269 if dedup_stats:270 lines += ["", "## Dedup", ""]271 lines += [f"- {k}: {v}" for k, v in sorted(dedup_stats.items())]272 lines += ["", "Policy-safe rows only: `golden_diff` is sha256-hashed and "273 "`deleted_symbols` dropped in `tasks/`, `corpus_*/`, `holdout/` "274 "(full rows live in the restricted `tasks_full/`).", ""]275 with _open(layout.card_path) as f:276 f.write("\n".join(lines))277 278 279def manifest_exists(layout: RunLayout) -> bool:280 """Write-once guard for the driver (finding D-21 idempotency)."""281 return _exists(layout.manifest_path)282 283 284__all__ = [285 "SCHEMA_VERSION",286 "RunLayout",287 "RunManifest",288 "manifest_exists",289 "write_dataset_card",290 "write_dpo_rows",291 "write_holdout",292 "write_quarantine",293 "write_sft_rows",294 "write_tasks",295 "write_tasks_full",296 "write_trajectories",297]298 