File size: 3,996 Bytes
932bc69
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
132
# ============================================================================
# AUTO-GENERATED — do not edit directly.
# Source of truth: src/speculators/provenance.py
# Regenerate with: make style   (or: python scripts/sync_provenance.py)
# ============================================================================
"""Shared provenance helpers for reproducibility artifacts.

Used by training, evaluation, and vLLM-launch scripts to record
command lines, git state, and package versions.
"""

from __future__ import annotations

import importlib.metadata
import importlib.util
import os
import subprocess
import tempfile
from pathlib import Path

TRACKED_PACKAGES = (
    "speculators",
    "vllm",
    "transformers",
    "torch",
    "compressed-tensors",
)


def atomic_write(path: Path, content: str) -> None:
    """Write *content* to *path* atomically via tempfile + rename."""
    fd, tmp = tempfile.mkstemp(dir=path.parent, prefix=f".{path.name}_", suffix=".tmp")
    tmp_path = Path(tmp)
    try:
        with os.fdopen(fd, "w") as f:
            f.write(content)
        tmp_path.replace(path)
    finally:
        if tmp_path.exists():
            tmp_path.unlink()


def pkg_version(name: str) -> str:
    """Return the installed version of *name*, or ``'not installed'``."""
    try:
        return importlib.metadata.version(name)
    except importlib.metadata.PackageNotFoundError:
        return "not installed"


def find_repo_root(start: Path) -> Path | None:
    """Walk up from *start* to the nearest directory containing ``.git``."""
    try:
        d = start.resolve()
        if d.is_file():
            d = d.parent
        while d != d.parent:
            if (d / ".git").exists():
                return d
            d = d.parent
    except OSError:
        pass
    return None


def find_package_repo(package_name: str) -> Path | None:
    """Find the git repo root for an installed editable package."""
    try:
        spec = importlib.util.find_spec(package_name)
        if spec and spec.origin:
            return find_repo_root(Path(spec.origin))
    except (ModuleNotFoundError, ValueError):
        pass
    return None


def git_sha(repo_root: Path | None) -> str:
    """Return the HEAD SHA of *repo_root*, or ``'unknown'`` on failure."""
    if repo_root is None:
        return "unknown"
    try:
        result = subprocess.run(
            ["git", "rev-parse", "HEAD"],  # noqa: S607
            capture_output=True,
            text=True,
            cwd=repo_root,
            timeout=5,
            check=False,
        )
        return result.stdout.strip() or "unknown"
    except (OSError, subprocess.TimeoutExpired):
        return "unknown"


def git_diff(repo_root: Path | None, *, timeout: int = 30) -> str:
    """Return ``git diff HEAD`` for *repo_root*, or empty string on failure."""
    if repo_root is None:
        return ""
    try:
        result = subprocess.run(
            ["git", "diff", "HEAD"],  # noqa: S607
            capture_output=True,
            text=True,
            cwd=repo_root,
            timeout=timeout,
            check=False,
        )
        return result.stdout.strip() if result.returncode == 0 else ""
    except (OSError, subprocess.TimeoutExpired):
        return ""


def run_git(args: list[str], cwd: str | Path, *, timeout: int = 5) -> str:
    """Run a git command and return stdout, or empty string on failure."""
    try:
        result = subprocess.run(  # noqa: S603
            args,
            capture_output=True,
            text=True,
            cwd=str(cwd),
            timeout=timeout,
            check=False,
        )
        return result.stdout.strip() if result.returncode == 0 else ""
    except (OSError, subprocess.TimeoutExpired):
        return ""


def package_versions(packages: tuple[str, ...] = TRACKED_PACKAGES) -> list[str]:
    """Return ``['# pkg: version', ...]`` header lines for *packages*."""
    return [f"# {p}: {pkg_version(p)}" for p in packages]