#!/usr/bin/env python3 """Export three-turn SWE-chat candidates without retaining later-turn text. The source snapshot is gated and deliberately stays below ``.artifacts/``. This importer does not download it: callers supply an already accepted, revision-pinned Parquet snapshot. It reads Parquet in record batches, joins the small sessions table only for repository/user grouping, and writes an unlabeled JSONL that contains the first prompt plus hashes (never text) for the two teacher-context prompts. """ from __future__ import annotations import argparse import json import sys from collections import Counter, defaultdict from dataclasses import dataclass from datetime import datetime, timezone from pathlib import Path from typing import Any, Iterable, Iterator, Sequence from purpose_data import DataError, canonical_json, file_sha256, prompt_hash, write_json SCRIPT_DIR = Path(__file__).resolve().parent DEFAULT_RAW_DIR = SCRIPT_DIR / ".artifacts" / "swe-chat" / "raw" DEFAULT_OUTPUT = SCRIPT_DIR / ".artifacts" / "swe-chat" / "candidates.jsonl" DEFAULT_MANIFEST = SCRIPT_DIR / ".artifacts" / "swe-chat" / "export-manifest.json" SCHEMA_VERSION = 1 REPOSITORY_ID = "SALT-NLP/SWE-chat" LICENSE = "ODC-By-1.0" @dataclass(frozen=True) class Turn: session_id: str turn_id: str conversation_turn_number: int turn_number: int prompt: str @dataclass(frozen=True) class Candidate: session_id: str repo_id: str | None user_id: str | None turns: tuple[Turn, Turn, Turn] def json(self, revision: str) -> dict[str, Any]: first, second, third = self.turns return { "schemaVersion": SCHEMA_VERSION, "repoID": self.repo_id, "userID": self.user_id, "sessionID": self.session_id, "sourceTurnIDs": [turn.turn_id for turn in self.turns], "sourceRevision": revision, "promptHash": prompt_hash(first.prompt), "contextPromptHashes": [prompt_hash(second.prompt), prompt_hash(third.prompt)], "prompt": first.prompt, # This ignored pre-labeling file is the only artifact allowed to carry # later text. Canonical labeled JSONL contains only the seven data fields. "teacherContext": [second.prompt, third.prompt], } def _as_text(value: Any) -> str | None: return value if isinstance(value, str) and value.strip() else None def _as_int(value: Any) -> int | None: if isinstance(value, bool): return None if isinstance(value, int): return value if isinstance(value, float) and value.is_integer(): return int(value) return None def _paths(raw_dir: Path, stem: str) -> list[Path]: candidates = sorted(raw_dir.rglob(f"*{stem}*.parquet")) if raw_dir.is_dir() else [] if not candidates: raise DataError( f"{raw_dir}: no {stem} Parquet files found; place the accepted pinned " f"SWE-chat snapshot under this directory or pass --{stem}" ) return candidates def parquet_rows(paths: Sequence[Path], columns: Sequence[str]) -> Iterator[dict[str, Any]]: try: import pyarrow.parquet as pq except ImportError as error: raise DataError( "Parquet import requires pyarrow; install it in the purpose-classifier " "environment (the raw gated snapshot is not read otherwise)" ) from error for path in paths: try: parquet = pq.ParquetFile(path) except Exception as error: raise DataError(f"{path}: cannot open Parquet: {error}") from error available = set(parquet.schema_arrow.names) missing = sorted(set(columns) - available) if missing: raise DataError( f"{path}: missing required columns {missing}; available columns are " f"{sorted(available)}" ) for batch in parquet.iter_batches(columns=list(columns), batch_size=16_384): values = batch.to_pydict() for index in range(batch.num_rows): yield {column: values[column][index] for column in columns} def session_metadata(rows: Iterable[dict[str, Any]]) -> dict[str, tuple[str | None, str | None]]: result: dict[str, tuple[str | None, str | None]] = {} for row in rows: session_id = _as_text(row.get("session_id")) if session_id is None: continue metadata = (_as_text(row.get("repo_id")), _as_text(row.get("user_id"))) previous = result.get(session_id) if previous is not None and previous != metadata: raise DataError(f"sessions config assigns conflicting repository/user to {session_id!r}") result[session_id] = metadata return result def select_candidates( conversation_rows: Iterable[dict[str, Any]], *, sessions: dict[str, tuple[str | None, str | None]], max_per_repo: int, max_per_user: int, ) -> tuple[list[Candidate], dict[str, int]]: """Apply documented row filters and deterministic session-level selection.""" funnel: Counter[str] = Counter() by_session: dict[str, list[Turn]] = defaultdict(list) for row in conversation_rows: funnel["conversationRows"] += 1 if row.get("turn_type") != "user_prompt": funnel["rejectedTurnType"] += 1 continue if row.get("role") != "user": funnel["rejectedRole"] += 1 continue if row.get("is_conversational") is not True: funnel["rejectedNonConversational"] += 1 continue if row.get("is_continuation") is True: funnel["rejectedContinuation"] += 1 continue session_id = _as_text(row.get("session_id")) turn_id = _as_text(row.get("turn_id")) prompt = row.get("content") conversation_turn_number = _as_int(row.get("conversation_turn_number")) turn_number = _as_int(row.get("turn_number")) if ( session_id is None or turn_id is None or not isinstance(prompt, str) or not prompt.strip() or "\x00" in prompt or conversation_turn_number is None or turn_number is None ): funnel["rejectedMalformedOrEmpty"] += 1 continue by_session[session_id].append( Turn(session_id, turn_id, conversation_turn_number, turn_number, prompt) ) funnel["eligibleRows"] += 1 preliminary: list[Candidate] = [] for session_id, turns in sorted(by_session.items()): ordered = sorted(turns, key=lambda turn: (turn.conversation_turn_number, turn.turn_number, turn.turn_id)) if len(ordered) < 3: funnel["sessionsFewerThanThreeEligiblePrompts"] += 1 continue ordinals = [(turn.conversation_turn_number, turn.turn_number) for turn in ordered[:3]] if len(set(ordinals)) != len(ordinals): funnel["sessionsAmbiguousTurnOrder"] += 1 continue repo_id, user_id = sessions.get(session_id, (None, None)) preliminary.append(Candidate(session_id, repo_id, user_id, tuple(ordered[:3]))) funnel["sessionsWithThreeEligiblePrompts"] = len(preliminary) first_by_hash: dict[str, str] = {} deduped: list[Candidate] = [] for candidate in preliminary: digest = prompt_hash(candidate.turns[0].prompt) if digest in first_by_hash: funnel["rejectedDuplicateFirstPrompt"] += 1 continue first_by_hash[digest] = candidate.session_id deduped.append(candidate) accepted: list[Candidate] = [] repo_counts: Counter[str] = Counter() user_counts: Counter[str] = Counter() for candidate in deduped: repo_key = candidate.repo_id or "" user_key = candidate.user_id or "" if repo_counts[repo_key] >= max_per_repo: funnel["rejectedRepoCap"] += 1 continue if user_counts[user_key] >= max_per_user: funnel["rejectedUserCap"] += 1 continue accepted.append(candidate) repo_counts[repo_key] += 1 user_counts[user_key] += 1 funnel["exportedCandidates"] = len(accepted) return accepted, dict(sorted(funnel.items())) def export( *, conversations: Sequence[Path], sessions_path: Sequence[Path], revision: str, output: Path, manifest_path: Path, max_per_repo: int, max_per_user: int, ) -> dict[str, Any]: if not revision or revision == "main": raise DataError("--revision must be an accepted immutable SWE-chat commit, never main") if max_per_repo <= 0 or max_per_user <= 0: raise DataError("source concentration caps must be positive") sessions = session_metadata(parquet_rows(sessions_path, ("session_id", "repo_id", "user_id"))) candidates, funnel = select_candidates( parquet_rows( conversations, ("session_id", "turn_id", "conversation_turn_number", "turn_number", "turn_type", "role", "is_conversational", "is_continuation", "content"), ), sessions=sessions, max_per_repo=max_per_repo, max_per_user=max_per_user, ) output.parent.mkdir(parents=True, exist_ok=True) output.write_text("".join(f"{canonical_json(candidate.json(revision))}\n" for candidate in candidates), encoding="utf-8") manifest = { "schemaVersion": SCHEMA_VERSION, "generatedAt": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"), "source": { "repoID": REPOSITORY_ID, "revision": revision, "license": LICENSE, "attribution": "SALT-NLP/SWE-chat; SWE-chat paper arXiv:2604.20779", "rawFiles": [{"path": str(path), "sha256": file_sha256(path)} for path in sorted([*conversations, *sessions_path])], }, "selection": { "rowFilter": "turn_type=user_prompt, role=user, is_conversational=true, not is_continuation", "perSession": "first three eligible non-empty prompts ordered by conversation_turn_number then turn_number", "studentText": "first prompt only", "teacherContext": "second and third prompts retained only in ignored candidate JSONL until labeling", "dedupe": "exact normalized first prompt", "maxPerRepo": max_per_repo, "maxPerUser": max_per_user, }, "funnel": funnel, "output": {"path": str(output), "records": len(candidates), "sha256": file_sha256(output)}, "removalLineage": "sourceTurnIDs and prompt hashes are retained in ignored audit sidecars; rebuild the next dataset version after a source tombstone.", } write_json(manifest_path, manifest) return manifest def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--raw-dir", type=Path, default=DEFAULT_RAW_DIR) parser.add_argument("--conversations", action="append", type=Path) parser.add_argument("--sessions", action="append", type=Path) parser.add_argument("--revision", required=True, help="accepted immutable Hugging Face revision") parser.add_argument("--output", type=Path, default=DEFAULT_OUTPUT) parser.add_argument("--manifest", type=Path, default=DEFAULT_MANIFEST) parser.add_argument("--max-per-repo", type=int, default=100) parser.add_argument("--max-per-user", type=int, default=50) return parser def main(argv: Sequence[str] | None = None) -> int: args = build_parser().parse_args(argv) try: raw_dir = args.raw_dir.expanduser().resolve() conversations = [path.expanduser().resolve() for path in args.conversations] if args.conversations else _paths(raw_dir, "conversations") sessions = [path.expanduser().resolve() for path in args.sessions] if args.sessions else _paths(raw_dir, "sessions") manifest = export( conversations=conversations, sessions_path=sessions, revision=args.revision, output=args.output.expanduser().resolve(), manifest_path=args.manifest.expanduser().resolve(), max_per_repo=args.max_per_repo, max_per_user=args.max_per_user, ) except (DataError, OSError, ValueError) as error: print(f"error: {error}", file=sys.stderr) return 1 print(json.dumps(manifest["funnel"], sort_keys=True)) print(f"Candidates: {manifest['output']['path']}") print(f"Manifest: {args.manifest}") return 0 if __name__ == "__main__": raise SystemExit(main())