flexlock 0.8.2__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
flexlock/diff.py ADDED
@@ -0,0 +1,253 @@
1
+ """Diff utilities for FlexLock."""
2
+
3
+ import os
4
+ from typing import Any, Dict, Set
5
+ from omegaconf import OmegaConf, DictConfig, ListConfig
6
+ from loguru import logger
7
+
8
+
9
+ class RunDiff:
10
+ def __init__(
11
+ self,
12
+ current: dict,
13
+ target: dict,
14
+ ignore_keys: list = None,
15
+ current_save_dir: str = None,
16
+ target_save_dir: str = None,
17
+ match_include: list = None,
18
+ match_exclude: list = None,
19
+ ):
20
+ """
21
+ Initialize RunDiff for comparing two run snapshots.
22
+
23
+ Args:
24
+ current: Current/proposed run snapshot
25
+ target: Target/existing run snapshot
26
+ ignore_keys: Additional keys to ignore during comparison
27
+ current_save_dir: Save directory of current run (for normalization)
28
+ target_save_dir: Save directory of target run (for normalization)
29
+ match_include: Override include patterns for git comparison
30
+ match_exclude: Override exclude patterns for git comparison
31
+ """
32
+ self.current = current
33
+ self.target = target
34
+
35
+ # Keys that legitimately appear inside the user config and are
36
+ # FlexLock-managed (or explicitly requested) — ignored at *any* depth.
37
+ self.always_ignore = set(ignore_keys or []) | {"save_dir", "_snapshot_"}
38
+
39
+ # Keys FlexLock injects at the snapshot top level (run metadata). These
40
+ # are only ignored at the top of the config subtree — a user
41
+ # hyperparameter named e.g. ``time`` or ``system`` nested deeper must
42
+ # NOT be silently ignored (issue 7), which would cause false cache hits.
43
+ self.toplevel_ignore = {
44
+ "timestamp",
45
+ "note",
46
+ "system",
47
+ "job_id",
48
+ "work_dir",
49
+ "cwd",
50
+ "date",
51
+ "time",
52
+ "datetime",
53
+ }
54
+
55
+ # Back-compat: some callers/tests read ``ignore_keys``.
56
+ self.ignore_keys = self.always_ignore | self.toplevel_ignore
57
+
58
+ # For value normalization (handling interpolation)
59
+ self.c_dir = str(current_save_dir) if current_save_dir else None
60
+ self.t_dir = str(target_save_dir) if target_save_dir else None
61
+
62
+ # Override patterns for git comparison (takes priority over snapshot-level patterns)
63
+ self.match_include = match_include
64
+ self.match_exclude = match_exclude
65
+
66
+ self.diffs = {}
67
+
68
+ def _normalize_val(self, val: Any, root_dir: str) -> Any:
69
+ """
70
+ If the value is a string and contains the save_dir path,
71
+ replace it with a placeholder <SAVE_DIR> to allow comparison.
72
+
73
+ Args:
74
+ val: Value to normalize
75
+ root_dir: Root directory path to replace
76
+
77
+ Returns:
78
+ Normalized value
79
+ """
80
+ if not (root_dir and isinstance(val, str)):
81
+ return val
82
+ # Prefix-only: replace an exact match or a genuine path prefix, never an
83
+ # arbitrary substring (issue 8). ``outputs/other`` must not be rewritten
84
+ # just because save_dir is ``outputs``.
85
+ if val == root_dir:
86
+ return "<SAVE_DIR>"
87
+ if val.startswith(root_dir + os.sep):
88
+ return "<SAVE_DIR>" + val[len(root_dir):]
89
+ return val
90
+
91
+ def compare_git(self):
92
+ """Compare git tree hashes, with optional include/exclude filtering."""
93
+ diff = []
94
+ c_repos = self.current.get("repos", {})
95
+ t_repos = self.target.get("repos", {})
96
+
97
+ # Iterate the union so a repo present on only one side is flagged
98
+ # symmetrically (issue 9) instead of being silently ignored.
99
+ for name in set(c_repos) | set(t_repos):
100
+ c_info = c_repos.get(name)
101
+ t_info = t_repos.get(name)
102
+ if not c_info:
103
+ diff.append(f"Repo {name} only in target")
104
+ continue
105
+ if not t_info:
106
+ diff.append(f"Repo {name} missing")
107
+ continue
108
+
109
+ # Compare Tree Hashes (Content Identity)
110
+ if c_info.get("tree") != t_info.get("tree"):
111
+ # Trees differ — check if RELEVANT files changed
112
+ # Priority: RunDiff-level override > snapshot-level patterns
113
+ include = (
114
+ self.match_include or c_info.get("include") or t_info.get("include")
115
+ )
116
+ exclude = (
117
+ self.match_exclude or c_info.get("exclude") or t_info.get("exclude")
118
+ )
119
+
120
+ if include or exclude:
121
+ repo_path = c_info.get("path") or t_info.get("path")
122
+ if repo_path and self._trees_match_filtered(
123
+ repo_path, c_info["tree"], t_info["tree"], include, exclude
124
+ ):
125
+ continue # Relevant files unchanged — match
126
+
127
+ diff.append(f"Repo {name}: Content changed")
128
+
129
+ if diff:
130
+ self.diffs["git"] = diff
131
+ return len(diff) == 0
132
+
133
+ def _trees_match_filtered(
134
+ self, repo_path, tree1, tree2, include=None, exclude=None
135
+ ):
136
+ """
137
+ Check if two trees match when filtered by include/exclude patterns.
138
+
139
+ Uses git diff-tree with pathspec filtering to compare only relevant files.
140
+ Returns True if no relevant files differ, False otherwise.
141
+ """
142
+ try:
143
+ from git.repo import Repo as GitRepo
144
+
145
+ repo = GitRepo(repo_path, search_parent_directories=True)
146
+
147
+ # Build git pathspec: include patterns + :(exclude) patterns
148
+ pathspec = list(include or [])
149
+ if exclude:
150
+ pathspec.extend(f":(exclude){pat}" for pat in exclude)
151
+
152
+ args = ["-r", "--name-only", "--no-commit-id", tree1, tree2]
153
+ if pathspec:
154
+ args.append("--")
155
+ args.extend(pathspec)
156
+
157
+ output = repo.git.diff_tree(*args).strip()
158
+ return len(output) == 0 # No relevant files changed
159
+ except Exception as e:
160
+ logger.debug(f"Filtered git comparison failed: {e}")
161
+ return False # Conservative: treat as different
162
+
163
+ def compare_config(self):
164
+ """
165
+ Compare configurations with recursive diff, ignoring specified keys
166
+ and normalizing path values.
167
+ """
168
+ c_cfg = self.current.get("config", {})
169
+ t_cfg = self.target.get("config", {})
170
+
171
+ def _recursive_diff(d1, d2, path=""):
172
+ """Recursively compare two config structures."""
173
+ diff = []
174
+
175
+ # Handle DictConfigs vs Primitives
176
+ if isinstance(d1, (dict, DictConfig)) and isinstance(
177
+ d2, (dict, DictConfig)
178
+ ):
179
+ # Injected run-metadata keys are ignored only at the top of the
180
+ # config subtree; user keys are compared at every depth.
181
+ ignore_here = self.always_ignore
182
+ if path == "":
183
+ ignore_here = ignore_here | self.toplevel_ignore
184
+
185
+ all_keys = set(d1.keys()) | set(d2.keys())
186
+ for k in all_keys:
187
+ if k in ignore_here:
188
+ continue
189
+
190
+ new_path = f"{path}.{k}" if path else k
191
+
192
+ if k not in d1:
193
+ diff.append(f"Missing in current: {new_path}")
194
+ elif k not in d2:
195
+ diff.append(f"Extra in current: {new_path}")
196
+ else:
197
+ diff.extend(_recursive_diff(d1[k], d2[k], new_path))
198
+
199
+ elif isinstance(d1, (list, tuple, ListConfig)) and isinstance(
200
+ d2, (list, tuple, ListConfig)
201
+ ):
202
+ if len(d1) != len(d2):
203
+ diff.append(f"List length mismatch {path}: {len(d1)} vs {len(d2)}")
204
+ else:
205
+ for i, (v1, v2) in enumerate(zip(d1, d2)):
206
+ diff.extend(_recursive_diff(v1, v2, f"{path}[{i}]"))
207
+
208
+ else:
209
+ # Value Comparison with Normalization
210
+ logger.debug(
211
+ f"Comparing values at {path}: {d1} vs {d2} after normalization with dirs {self.c_dir}, {self.t_dir}"
212
+ )
213
+
214
+ v1_norm = self._normalize_val(d1, self.c_dir)
215
+ v2_norm = self._normalize_val(d2, self.t_dir)
216
+
217
+ if v1_norm != v2_norm:
218
+ diff.append(f"Value mismatch {path}: {v1_norm} != {v2_norm}")
219
+
220
+ return diff
221
+
222
+ diff = _recursive_diff(c_cfg, t_cfg)
223
+
224
+ if diff:
225
+ self.diffs["config"] = diff
226
+
227
+ return len(diff) == 0
228
+
229
+ def compare_data(self):
230
+ """Compare data hashes."""
231
+ c_data = self.current.get("data", {})
232
+ t_data = self.target.get("data", {})
233
+
234
+ if c_data == t_data:
235
+ return True
236
+
237
+ # Name which keys/hashes differ instead of a bare "Data differs" (issue 9).
238
+ detail = []
239
+ for key in sorted(set(c_data) | set(t_data)):
240
+ cv = c_data.get(key)
241
+ tv = t_data.get(key)
242
+ if key not in c_data:
243
+ detail.append(f"{key}: only in target")
244
+ elif key not in t_data:
245
+ detail.append(f"{key}: only in current")
246
+ elif cv != tv:
247
+ detail.append(f"{key}: {cv} != {tv}")
248
+ self.diffs["data"] = detail or ["Data differs"]
249
+ return False
250
+
251
+ def is_match(self):
252
+ """Check if the current run matches the target run."""
253
+ return self.compare_git() and self.compare_config() and self.compare_data()
flexlock/diff_cli.py ADDED
@@ -0,0 +1,155 @@
1
+ """CLI for comparing FlexLock snapshots from various sources.
2
+
3
+ Exit codes: ``0`` match, ``1`` differ, ``2`` error (so the command is usable
4
+ in scripts and CI — the old behaviour always exited 0).
5
+ """
6
+
7
+ import argparse
8
+ import json
9
+ import sys
10
+ import yaml
11
+ from pathlib import Path
12
+ from loguru import logger
13
+ from flexlock.diff import RunDiff
14
+ from flexlock.taskdb import get_task_snapshot
15
+
16
+
17
+ def load_snapshot_from_dir(dir_path: Path) -> dict:
18
+ """Load snapshot from a directory (finds run.lock file)."""
19
+ lock_file = dir_path / "run.lock"
20
+ if not lock_file.exists():
21
+ raise FileNotFoundError(f"No run.lock found in {dir_path}")
22
+
23
+ with open(lock_file) as f:
24
+ return yaml.safe_load(f)
25
+
26
+
27
+ def load_snapshot_from_db(db_path: Path, task_id: str) -> dict:
28
+ """Load snapshot from database by task_id."""
29
+ snapshot = get_task_snapshot(db_path, task_id)
30
+ if snapshot is None:
31
+ raise ValueError(f"No snapshot found for task_id '{task_id}' in {db_path}")
32
+ return snapshot
33
+
34
+
35
+ def run_comparison(snap1: dict, snap2: dict) -> "tuple[bool, dict]":
36
+ """Compare two snapshots, returning ``(is_match, diffs)``.
37
+
38
+ The single source of truth for both the text/JSON printers here and
39
+ ``flexlock why``. ``diffs`` is a JSON-serializable dict of lists keyed by
40
+ ``git``/``config``/``data`` (only the categories that differ appear).
41
+ """
42
+ diff = RunDiff(snap1, snap2)
43
+ # is_match() runs all three comparisons, populating diff.diffs as a side
44
+ # effect. Call it first so diffs is complete regardless of short-circuiting.
45
+ is_match = diff.is_match()
46
+ return is_match, dict(diff.diffs)
47
+
48
+
49
+ def compare_snapshots(snap1: dict, snap2: dict, show_details: bool = False) -> bool:
50
+ """Compare two snapshots and print a human-readable report."""
51
+ is_match, diffs = run_comparison(snap1, snap2)
52
+
53
+ print("\n=== Snapshot Comparison ===\n")
54
+ for label, key in (("Git", "git"), ("Config", "config"), ("Data", "data")):
55
+ section_match = key not in diffs
56
+ print(f"{label:6s}: {'✓ Match' if section_match else '✗ Differ'}")
57
+ if not section_match and show_details:
58
+ for d in diffs.get(key, []):
59
+ print(f" - {d}")
60
+
61
+ print(f"\nOverall: {'✓ Snapshots Match' if is_match else '✗ Snapshots Differ'}\n")
62
+ return is_match
63
+
64
+
65
+ def main():
66
+ """CLI entry point for flexlock diff command."""
67
+ parser = argparse.ArgumentParser(
68
+ description="Compare FlexLock snapshots from various sources"
69
+ )
70
+
71
+ # Shared options for every subcommand.
72
+ common = argparse.ArgumentParser(add_help=False)
73
+ common.add_argument(
74
+ "--details", action="store_true", help="Show detailed differences (text)"
75
+ )
76
+ common.add_argument(
77
+ "--format", choices=["text", "json"], default="text",
78
+ help="Output format (default: text)",
79
+ )
80
+
81
+ subparsers = parser.add_subparsers(
82
+ dest="mode", required=True, help="Comparison mode"
83
+ )
84
+
85
+ # Mode 1: Compare two directories (traditional)
86
+ dir_parser = subparsers.add_parser(
87
+ "dirs", parents=[common], help="Compare two directory-based snapshots"
88
+ )
89
+ dir_parser.add_argument("dir1", type=Path, help="First directory")
90
+ dir_parser.add_argument("dir2", type=Path, help="Second directory")
91
+
92
+ # Mode 2: Compare two tasks in DB
93
+ db_parser = subparsers.add_parser(
94
+ "db", parents=[common], help="Compare two DB-based snapshots"
95
+ )
96
+ db_parser.add_argument("db_path", type=Path, help="Path to tasks database")
97
+ db_parser.add_argument("task_id1", help="First task ID (hash)")
98
+ db_parser.add_argument("task_id2", help="Second task ID (hash)")
99
+
100
+ # Mode 3: Compare directory to DB task
101
+ mixed_parser = subparsers.add_parser(
102
+ "mixed", parents=[common], help="Compare directory snapshot to DB snapshot"
103
+ )
104
+ mixed_parser.add_argument("dir_path", type=Path, help="Directory path")
105
+ mixed_parser.add_argument("db_path", type=Path, help="Database path")
106
+ mixed_parser.add_argument("task_id", help="Task ID in database")
107
+
108
+ args = parser.parse_args()
109
+
110
+ try:
111
+ if args.mode == "dirs":
112
+ if not args.dir1.exists():
113
+ logger.error(f"Directory not found: {args.dir1}")
114
+ sys.exit(2)
115
+ if not args.dir2.exists():
116
+ logger.error(f"Directory not found: {args.dir2}")
117
+ sys.exit(2)
118
+ snap1 = load_snapshot_from_dir(args.dir1)
119
+ snap2 = load_snapshot_from_dir(args.dir2)
120
+
121
+ elif args.mode == "db":
122
+ if not args.db_path.exists():
123
+ logger.error(f"Database not found: {args.db_path}")
124
+ sys.exit(2)
125
+ snap1 = load_snapshot_from_db(args.db_path, args.task_id1)
126
+ snap2 = load_snapshot_from_db(args.db_path, args.task_id2)
127
+
128
+ elif args.mode == "mixed":
129
+ if not args.dir_path.exists():
130
+ logger.error(f"Directory not found: {args.dir_path}")
131
+ sys.exit(2)
132
+ if not args.db_path.exists():
133
+ logger.error(f"Database not found: {args.db_path}")
134
+ sys.exit(2)
135
+ snap1 = load_snapshot_from_dir(args.dir_path)
136
+ snap2 = load_snapshot_from_db(args.db_path, args.task_id)
137
+
138
+ if args.format == "json":
139
+ is_match, diffs = run_comparison(snap1, snap2)
140
+ print(json.dumps({"match": is_match, "diffs": diffs}, indent=2, default=str))
141
+ else:
142
+ is_match = compare_snapshots(snap1, snap2, args.details)
143
+
144
+ # 0 match / 1 differ.
145
+ sys.exit(0 if is_match else 1)
146
+
147
+ except SystemExit:
148
+ raise
149
+ except Exception as e:
150
+ logger.error(f"Comparison failed: {e}")
151
+ sys.exit(2)
152
+
153
+
154
+ if __name__ == "__main__":
155
+ main()
flexlock/exceptions.py ADDED
@@ -0,0 +1,50 @@
1
+ """Custom exceptions for FlexLock."""
2
+
3
+
4
+ class FlexLockError(Exception):
5
+ """Base exception for all FlexLock errors."""
6
+
7
+ pass
8
+
9
+
10
+ class FlexLockConfigError(FlexLockError):
11
+ """Raised when there is an error in configuration."""
12
+
13
+ pass
14
+
15
+
16
+ class FlexLockExecutionError(FlexLockError):
17
+ """Raised when there is an error during execution."""
18
+
19
+ pass
20
+
21
+
22
+ class FlexLockSnapshotError(FlexLockError):
23
+ """Raised when there is an error creating or loading snapshots."""
24
+
25
+ pass
26
+
27
+
28
+ class FlexLockValidationError(FlexLockError):
29
+ """Raised when validation of inputs fails."""
30
+
31
+ pass
32
+
33
+
34
+ class FlexLockCacheError(FlexLockError):
35
+ """Raised when there is an error with smart run caching."""
36
+
37
+ pass
38
+
39
+
40
+ class FlexLockBackendError(FlexLockError):
41
+ """Raised when there is an error with HPC backends (Slurm/PBS)."""
42
+
43
+ pass
44
+
45
+
46
+ class UnresolvedInterpolationError(FlexLockConfigError):
47
+ """Raised when a sub-node interpolation references a key that is not present
48
+ in the sub-tree and cannot be resolved against the root config either."""
49
+
50
+ pass
flexlock/export.py ADDED
@@ -0,0 +1,134 @@
1
+ """Export utilities for FlexLock - extract snapshots from database to files."""
2
+
3
+ import argparse
4
+ import sys
5
+ import json
6
+ import yaml
7
+ import tempfile
8
+ import os
9
+ from pathlib import Path
10
+ from loguru import logger
11
+ from flexlock.taskdb import get_task_snapshot, list_task_snapshots
12
+
13
+
14
+ def export_task(db_path: Path, task_id: str, output_dir: Path) -> None:
15
+ """
16
+ Export a single task snapshot from database to a standalone directory.
17
+
18
+ Args:
19
+ db_path: Path to SQLite database
20
+ task_id: Hash of the task to export
21
+ output_dir: Directory to write the exported snapshot
22
+
23
+ Raises:
24
+ ValueError: If task not found in database
25
+ """
26
+ snapshot_data = get_task_snapshot(db_path, task_id)
27
+ if not snapshot_data:
28
+ raise ValueError(f"Task {task_id} not found in {db_path}")
29
+
30
+ output_dir.mkdir(parents=True, exist_ok=True)
31
+
32
+ # Atomic write of run.lock file
33
+ with tempfile.NamedTemporaryFile("w", dir=output_dir, delete=False) as tf:
34
+ yaml.dump(snapshot_data, tf, sort_keys=False)
35
+ tmp_name = tf.name
36
+ os.replace(tmp_name, output_dir / "run.lock")
37
+
38
+ logger.info(f"Exported task {task_id} to {output_dir}")
39
+
40
+
41
+ def export_all_tasks(db_path: Path, output_base_dir: Path, status: str = None) -> None:
42
+ """
43
+ Export all tasks (or filtered by status) from database to separate directories.
44
+
45
+ Args:
46
+ db_path: Path to SQLite database
47
+ output_base_dir: Base directory for exports (each task gets a subdirectory)
48
+ status: Optional filter by status (pending, running, done, failed)
49
+ """
50
+ tasks = list_task_snapshots(db_path, status)
51
+
52
+ if not tasks:
53
+ logger.warning(
54
+ f"No tasks found in {db_path}"
55
+ + (f" with status={status}" if status else "")
56
+ )
57
+ return
58
+
59
+ output_base_dir.mkdir(parents=True, exist_ok=True)
60
+
61
+ for task_id, snapshot_data, task_status in tasks:
62
+ if snapshot_data:
63
+ task_output_dir = output_base_dir / f"task_{task_id[:8]}"
64
+ task_output_dir.mkdir(parents=True, exist_ok=True)
65
+
66
+ # Atomic write
67
+ with tempfile.NamedTemporaryFile(
68
+ "w", dir=task_output_dir, delete=False
69
+ ) as tf:
70
+ yaml.dump(snapshot_data, tf, sort_keys=False)
71
+ tmp_name = tf.name
72
+ os.replace(tmp_name, task_output_dir / "run.lock")
73
+
74
+ logger.info(
75
+ f"Exported task {task_id[:8]} (status={task_status}) to {task_output_dir}"
76
+ )
77
+
78
+ logger.info(f"Exported {len(tasks)} tasks to {output_base_dir}")
79
+
80
+
81
+ def main():
82
+ """CLI entry point for flexlock export command."""
83
+ parser = argparse.ArgumentParser(
84
+ description="Export FlexLock task snapshots from database to files"
85
+ )
86
+
87
+ parser.add_argument(
88
+ "--db",
89
+ type=Path,
90
+ required=True,
91
+ help="Path to tasks database (e.g., outputs/sweep/tasks.db)",
92
+ )
93
+
94
+ parser.add_argument(
95
+ "--task", help="Task ID to export (hash). If not specified, exports all tasks."
96
+ )
97
+
98
+ parser.add_argument(
99
+ "--out",
100
+ type=Path,
101
+ required=True,
102
+ help="Output directory for exported snapshot(s)",
103
+ )
104
+
105
+ parser.add_argument(
106
+ "--status",
107
+ choices=["pending", "running", "done", "failed"],
108
+ help="Filter tasks by status (only used when exporting all tasks)",
109
+ )
110
+
111
+ args = parser.parse_args()
112
+
113
+ # Validate database exists
114
+ if not args.db.exists():
115
+ logger.error(f"Database not found: {args.db}")
116
+ sys.exit(1)
117
+
118
+ try:
119
+ if args.task:
120
+ # Export single task
121
+ export_task(args.db, args.task, args.out)
122
+ else:
123
+ # Export all tasks
124
+ export_all_tasks(args.db, args.out, args.status)
125
+
126
+ logger.success("Export completed successfully")
127
+
128
+ except Exception as e:
129
+ logger.error(f"Export failed: {e}")
130
+ sys.exit(1)
131
+
132
+
133
+ if __name__ == "__main__":
134
+ main()