updater.py
20662 bytes
1import json
2import os
3import re
4import shutil
5import subprocess
6import sys
7import timeit
8from copy import deepcopy
9from typing import Literal, NotRequired, Optional, TypedDict
10
11import requests
12import yaml
13from semver import Version
14
15# Get TMP_DIR variable from environment
16TMP_DIR = os.path.join(os.environ.get("TMP_DIR", "/tmp"), "ohmyzsh")
17# Relative path to dependencies.yml file
18DEPS_YAML_FILE = ".github/dependencies.yml"
19# Dry run flag
20DRY_RUN = os.environ.get("DRY_RUN", "0") == "1"
21# GitHub Token is needed to avoid rate limiting
22GH_TOKEN = os.environ.get("GH_TOKEN")
23HEADERS = {
24 "Accept": "application/vnd.github+json",
25}
26if GH_TOKEN:
27 HEADERS["Authorization"] = f"Bearer {GH_TOKEN}"
28
29# utils for tag comparison
30BASEVERSION = re.compile(
31 r"""[vV]?
32 (?P<major>(0|[1-9])\d*)
33 (\.
34 (?P<minor>(0|[1-9])\d*)
35 (\.
36 (?P<patch>(0|[1-9])\d*)
37 )?
38 )?
39 """,
40 re.VERBOSE,
41)
42
43
44def coerce(version: str) -> Optional[Version]:
45 match = BASEVERSION.search(version)
46 if not match:
47 return None
48
49 # BASEVERSION looks for `MAJOR.minor.patch` in the string given
50 # it fills with None if any of them is missing (for example `2.1`)
51 ver = {
52 key: 0 if value is None else value for key, value in match.groupdict().items()
53 }
54 # Version takes `major`, `minor`, `patch` arguments
55 ver = Version(**ver) # pyright: ignore[reportArgumentType]
56 return ver
57
58
59class CodeTimer:
60 def __init__(self, name=None):
61 self.name = " '" + name + "'" if name else ""
62
63 def __enter__(self):
64 self.start = timeit.default_timer()
65
66 def __exit__(self, exc_type, exc_value, traceback):
67 self.took = (timeit.default_timer() - self.start) * 1000.0
68 print("Code block" + self.name + " took: " + str(self.took) + " ms")
69
70
71### YAML representation
72def str_presenter(dumper, data):
73 """
74 Configures yaml for dumping multiline strings
75 Ref: https://stackoverflow.com/a/33300001
76 """
77 if len(data.splitlines()) > 1: # check for multiline string
78 return dumper.represent_scalar("tag:yaml.org,2002:str", data, style="|")
79 return dumper.represent_scalar("tag:yaml.org,2002:str", data)
80
81
82yaml.add_representer(str, str_presenter)
83yaml.representer.SafeRepresenter.add_representer(str, str_presenter)
84
85
86# Types
87class DependencyDict(TypedDict):
88 repo: str
89 branch: str
90 version: str
91 precopy: NotRequired[str]
92 postcopy: NotRequired[str]
93
94
95class DependencyYAML(TypedDict):
96 dependencies: dict[str, DependencyDict]
97
98
99class UpdateStatusFalse(TypedDict):
100 has_updates: Literal[False]
101
102
103class UpdateStatusTrue(TypedDict):
104 has_updates: Literal[True]
105 version: str
106 compare_url: str
107 head_ref: str
108 head_url: str
109
110
111class CommandRunner:
112 class Exception(Exception):
113 def __init__(self, message, returncode, stage, stdout, stderr):
114 super().__init__(message)
115 self.returncode = returncode
116 self.stage = stage
117 self.stdout = stdout
118 self.stderr = stderr
119
120 @staticmethod
121 def run_or_fail(command: list[str], stage: str, *args, **kwargs):
122 if DRY_RUN and command[0] == "gh":
123 command.insert(0, "echo")
124
125 result = subprocess.run(command, *args, capture_output=True, **kwargs)
126
127 if result.returncode != 0:
128 raise CommandRunner.Exception(
129 f"{stage} command failed with exit code {result.returncode}",
130 returncode=result.returncode,
131 stage=stage,
132 stdout=result.stdout.decode("utf-8"),
133 stderr=result.stderr.decode("utf-8"),
134 )
135
136 return result
137
138
139class DependencyStore:
140 store: DependencyYAML = {"dependencies": {}}
141
142 @staticmethod
143 def set(data: DependencyYAML):
144 DependencyStore.store = data
145
146 @staticmethod
147 def update_dependency_version(path: str, version: str) -> DependencyYAML:
148 with CodeTimer(f"store deepcopy: {path}"):
149 store_copy = deepcopy(DependencyStore.store)
150
151 dependency = store_copy["dependencies"].get(path)
152 if dependency is None:
153 raise ValueError(f"Dependency {path} {version} not found")
154 dependency["version"] = version
155 store_copy["dependencies"][path] = dependency
156
157 return store_copy
158
159 @staticmethod
160 def write_store(file: str, data: DependencyYAML):
161 with open(file, "w") as yaml_file:
162 yaml.safe_dump(data, yaml_file, sort_keys=False)
163
164
165class Dependency:
166 def __init__(self, path: str, values: DependencyDict):
167 self.path = path
168 self.values = values
169
170 self.name: str = ""
171 self.desc: str = ""
172 self.kind: str = ""
173
174 match path.split("/"):
175 case ["plugins", name]:
176 self.name = name
177 self.kind = "plugin"
178 self.desc = f"{name} plugin"
179 case ["themes", name]:
180 self.name = name.replace(".zsh-theme", "")
181 self.kind = "theme"
182 self.desc = f"{self.name} theme"
183 case _:
184 self.name = self.desc = path
185
186 def __str__(self):
187 output: str = ""
188 for key in DependencyDict.__dict__["__annotations__"].keys():
189 if key not in self.values:
190 output += f"{key}: None\n"
191 continue
192
193 value = self.values[key]
194 if "\n" not in value:
195 output += f"{key}: {value}\n"
196 else:
197 output += f"{key}:\n "
198 output += value.replace("\n", "\n ", value.count("\n") - 1)
199 return output
200
201 def update_or_notify(self):
202 # Print dependency settings
203 print(f"Processing {self.desc}...", file=sys.stderr)
204 print(self, file=sys.stderr)
205
206 # Check for updates
207 repo = self.values["repo"]
208 remote_branch = self.values["branch"]
209 version = self.values["version"]
210 is_tag = version.startswith("tag:")
211
212 try:
213 with CodeTimer(f"update check: {repo}"):
214 if is_tag:
215 status = GitHub.check_newer_tag(repo, version.replace("tag:", ""))
216 else:
217 status = GitHub.check_updates(repo, remote_branch, version)
218
219 if status["has_updates"] is True:
220 short_sha = status["head_ref"][:8]
221 new_version = status["version"] if is_tag else short_sha
222 source_ref = new_version if is_tag else status["head_ref"]
223
224 try:
225 branch_name = f"update/{self.path}/{new_version}"
226
227 # Create new branch
228 branch = Git.checkout_or_create_branch(branch_name)
229
230 # Update dependency files
231 self.__apply_upstream_changes(source_ref)
232
233 if not Git.repo_is_clean():
234 # Update dependencies.yml file
235 self.__update_yaml(
236 f"tag:{new_version}" if is_tag else status["version"]
237 )
238
239 # Add all changes and commit
240 has_new_commit = Git.add_and_commit(self.name, new_version)
241
242 if has_new_commit:
243 # Push changes to remote
244 Git.push(branch)
245
246 # Create GitHub PR
247 GitHub.create_pr(
248 branch,
249 f"chore({self.name}): update to version {new_version}",
250 f"""## Description
251
252Update for **{self.desc}**: update to version [{new_version}]({status["head_url"]}).
253Check out the [list of changes]({status["compare_url"]}).
254""",
255 )
256
257 # Clean up repository
258 Git.clean_repo()
259 except (CommandRunner.Exception, shutil.Error) as e:
260 # Handle exception on automatic update
261 match type(e):
262 case CommandRunner.Exception:
263 # Print error message
264 print(
265 f"Error running {e.stage} command: {e.returncode}", # pyright: ignore[reportAttributeAccessIssue]
266 file=sys.stderr,
267 )
268 print(e.stderr, file=sys.stderr) # pyright: ignore[reportAttributeAccessIssue]
269 case shutil.Error:
270 print(f"Error copying files: {e}", file=sys.stderr)
271
272 try:
273 Git.clean_repo()
274 except CommandRunner.Exception as e:
275 print(
276 f"Error reverting repository to clean state: {e}",
277 file=sys.stderr,
278 )
279 sys.exit(1)
280
281 # Create a GitHub issue to notify maintainer
282 title = f"{self.path}: update to {new_version}"
283 body = f"""## Description
284
285There is a new version of `{self.name}` {self.kind} available.
286
287New version: [{new_version}]({status["head_url"]})
288Check out the [list of changes]({status["compare_url"]}).
289"""
290
291 print("Creating GitHub issue", file=sys.stderr)
292 print(f"{title}\n\n{body}", file=sys.stderr)
293 GitHub.create_issue(title, body)
294 except Exception as e:
295 print(e, file=sys.stderr)
296
297 def __update_yaml(self, new_version: str) -> None:
298 dep_yaml = DependencyStore.update_dependency_version(self.path, new_version)
299 DependencyStore.write_store(DEPS_YAML_FILE, dep_yaml)
300
301 def __apply_upstream_changes(self, ref: str) -> None:
302 # Patterns to ignore in copying files from upstream repo
303 GLOBAL_IGNORE = [".git", ".github", ".gitignore"]
304
305 path = os.path.abspath(self.path)
306 precopy = self.values.get("precopy")
307 postcopy = self.values.get("postcopy")
308
309 repo = self.values["repo"]
310 remote_url = f"https://github.com/{repo}.git"
311 repo_dir = os.path.join(TMP_DIR, repo)
312
313 # Clone repository
314 Git.clone(remote_url, ref, repo_dir, reclone=True)
315
316 # Run precopy on tmp repo
317 if precopy is not None:
318 print("Running precopy script:", end="\n ", file=sys.stderr)
319 print(
320 precopy.replace("\n", "\n ", precopy.count("\n") - 1), file=sys.stderr
321 )
322 CommandRunner.run_or_fail(
323 ["bash", "-c", precopy], cwd=repo_dir, stage="Precopy"
324 )
325
326 # Copy files from upstream repo
327 print(f"Copying files from {repo_dir} to {path}", file=sys.stderr)
328 shutil.copytree(
329 repo_dir,
330 path,
331 dirs_exist_ok=True,
332 ignore=shutil.ignore_patterns(*GLOBAL_IGNORE),
333 )
334
335 # Run postcopy on our repository
336 if postcopy is not None:
337 print("Running postcopy script:", end="\n ", file=sys.stderr)
338 print(
339 postcopy.replace("\n", "\n ", postcopy.count("\n") - 1),
340 file=sys.stderr,
341 )
342 CommandRunner.run_or_fail(
343 ["bash", "-c", postcopy], cwd=path, stage="Postcopy"
344 )
345
346
347class Git:
348 default_branch = "master"
349
350 @staticmethod
351 def clone(remote_url: str, ref: str, repo_dir: str, reclone=False):
352 # If repo needs to be fresh
353 if reclone and os.path.exists(repo_dir):
354 shutil.rmtree(repo_dir)
355
356 # Clone repo in tmp directory and checkout branch
357 if not os.path.exists(repo_dir):
358 print(
359 f"Cloning {remote_url} to {repo_dir} and checking out {ref}",
360 file=sys.stderr,
361 )
362 CommandRunner.run_or_fail(
363 ["git", "clone", "--depth=1", "--revision", ref, remote_url, repo_dir],
364 stage="Clone",
365 )
366
367 @staticmethod
368 def checkout_or_create_branch(branch_name: str):
369 # Get current branch name
370 result = CommandRunner.run_or_fail(
371 ["git", "rev-parse", "--abbrev-ref", "HEAD"], stage="GetDefaultBranch"
372 )
373 Git.default_branch = result.stdout.decode("utf-8").strip()
374
375 # Create new branch and return created branch name
376 try:
377 # try to checkout already existing branch
378 CommandRunner.run_or_fail(
379 ["git", "checkout", branch_name], stage="CreateBranch"
380 )
381 except CommandRunner.Exception:
382 # otherwise create new branch
383 CommandRunner.run_or_fail(
384 ["git", "checkout", "-b", branch_name], stage="CreateBranch"
385 )
386 return branch_name
387
388 @staticmethod
389 def repo_is_clean() -> bool:
390 """
391 Returns `True` if the repo is clean.
392 Returns `False` if the repo is dirty.
393 """
394 try:
395 result = CommandRunner.run_or_fail(
396 ["git", "status", "--porcelain", "--untracked-files=normal"],
397 stage="CheckRepoClean",
398 )
399 except CommandRunner.Exception:
400 return False
401
402 return result.stdout.strip() == b""
403
404 @staticmethod
405 def add_and_commit(scope: str, version: str) -> bool:
406 """
407 Returns `True` if there were changes and were indeed commited.
408 Returns `False` if the repo was clean and no changes were commited.
409 """
410 if Git.repo_is_clean():
411 return False
412
413 user_name = os.environ.get("GIT_APP_NAME")
414 user_email = os.environ.get("GIT_APP_EMAIL")
415
416 # Add all files to git staging
417 CommandRunner.run_or_fail(["git", "add", "-A", "-v"], stage="AddFiles")
418
419 # Reset environment and git config
420 clean_env = os.environ.copy()
421 clean_env["LANG"] = "C.UTF-8"
422 clean_env["GIT_CONFIG_GLOBAL"] = "/dev/null"
423 clean_env["GIT_CONFIG_NOSYSTEM"] = "1"
424
425 # Commit with settings above
426 CommandRunner.run_or_fail(
427 [
428 "git",
429 "-c",
430 f"user.name={user_name}",
431 "-c",
432 f"user.email={user_email}",
433 "commit",
434 "-m",
435 f"chore({scope}): update to {version}",
436 ],
437 stage="CreateCommit",
438 env=clean_env,
439 )
440 return True
441
442 @staticmethod
443 def push(branch: str):
444 CommandRunner.run_or_fail(
445 ["git", "push", "-u", "origin", branch], stage="PushBranch"
446 )
447
448 @staticmethod
449 def clean_repo():
450 CommandRunner.run_or_fail(
451 ["git", "reset", "--hard", "HEAD"], stage="ResetRepository"
452 )
453 CommandRunner.run_or_fail(
454 ["git", "checkout", Git.default_branch], stage="CheckoutDefaultBranch"
455 )
456
457
458class GitHub:
459 @staticmethod
460 def check_newer_tag(repo, current_tag) -> UpdateStatusFalse | UpdateStatusTrue:
461 # GET /repos/:owner/:repo/git/refs/tags
462 url = f"https://api.github.com/repos/{repo}/git/refs/tags"
463
464 # Send a GET request to the GitHub API
465 response = requests.get(url, headers=HEADERS)
466 current_version = coerce(current_tag)
467 if current_version is None:
468 raise ValueError(
469 f"Stored {current_version} from {repo} does not follow semver"
470 )
471
472 # If the request was successful
473 if response.status_code == 200:
474 # Parse the JSON response
475 data = response.json()
476
477 if len(data) == 0:
478 return {
479 "has_updates": False,
480 }
481
482 latest_ref = None
483 latest_version: Optional[Version] = None
484 for ref in data:
485 # we find the tag since GitHub returns it as plain git ref
486 tag_version = coerce(ref["ref"].replace("refs/tags/", ""))
487 if tag_version is None:
488 # we skip every tag that is not semver-complaint
489 continue
490 if latest_version is None or tag_version.compare(latest_version) > 0:
491 # if we have a "greater" semver version, set it as latest
492 latest_version = tag_version
493 latest_ref = ref
494
495 # raise if no valid semver tag is found
496 if latest_ref is None or latest_version is None:
497 raise ValueError(f"No tags following semver found in {repo}")
498
499 # we get the tag since GitHub returns it as plain git ref
500 latest_tag = latest_ref["ref"].replace("refs/tags/", "")
501
502 if latest_version.compare(current_version) <= 0:
503 return {
504 "has_updates": False,
505 }
506
507 return {
508 "has_updates": True,
509 "version": latest_tag,
510 "compare_url": f"https://github.com/{repo}/compare/{current_tag}...{latest_tag}",
511 "head_ref": latest_ref["object"]["sha"],
512 "head_url": f"https://github.com/{repo}/releases/tag/{latest_tag}",
513 }
514 else:
515 # If the request was not successful, raise an exception
516 raise Exception(
517 f"GitHub API request failed with status code {response.status_code}: {response.json()}"
518 )
519
520 @staticmethod
521 def check_updates(repo, branch, version) -> UpdateStatusFalse | UpdateStatusTrue:
522 url = f"https://api.github.com/repos/{repo}/compare/{version}...{branch}"
523
524 # Send a GET request to the GitHub API
525 response = requests.get(url, headers=HEADERS)
526
527 # If the request was successful
528 if response.status_code == 200:
529 # Parse the JSON response
530 data = response.json()
531
532 # If the base is behind the head, there is a newer version
533 has_updates = data["status"] != "identical"
534
535 if not has_updates:
536 return {
537 "has_updates": False,
538 }
539
540 return {
541 "has_updates": data["status"] != "identical",
542 "version": data["commits"][-1]["sha"],
543 "compare_url": data["permalink_url"],
544 "head_ref": data["commits"][-1]["sha"],
545 "head_url": data["commits"][-1]["html_url"],
546 }
547 else:
548 # If the request was not successful, raise an exception
549 raise Exception(
550 f"GitHub API request failed with status code {response.status_code}: {response.json()}"
551 )
552
553 @staticmethod
554 def create_issue(title: str, body: str) -> None:
555 cmd = ["gh", "issue", "create", "-t", title, "-b", body]
556 CommandRunner.run_or_fail(cmd, stage="CreateIssue")
557
558 @staticmethod
559 def create_pr(branch: str, title: str, body: str) -> None:
560 # first of all let's check if PR is already open
561 check_cmd = [
562 "gh",
563 "pr",
564 "list",
565 "--state",
566 "open",
567 "--head",
568 branch,
569 "--json",
570 "title",
571 ]
572 # returncode is 0 also if no PRs are found
573 output = json.loads(
574 CommandRunner.run_or_fail(check_cmd, stage="CheckPullRequestOpen")
575 .stdout.decode("utf-8")
576 .strip()
577 )
578 # we have PR in this case!
579 if len(output) > 0:
580 return
581 cmd = [
582 "gh",
583 "pr",
584 "create",
585 "-B",
586 Git.default_branch,
587 "-H",
588 branch,
589 "-t",
590 title,
591 "-b",
592 body,
593 ]
594 CommandRunner.run_or_fail(cmd, stage="CreatePullRequest")
595
596
597def main():
598 # Load the YAML file
599 with open(DEPS_YAML_FILE, "r") as yaml_file:
600 data: DependencyYAML = yaml.safe_load(yaml_file)
601
602 if "dependencies" not in data:
603 raise Exception("dependencies.yml not properly formatted")
604
605 # Cache YAML version
606 DependencyStore.set(data)
607
608 dependencies = data["dependencies"]
609 if len(sys.argv) > 1:
610 # argv is list of dependencies to run, default is all of them
611 dependency_list = sys.argv[1:]
612 else:
613 dependency_list = dependencies.keys()
614
615 for path in dependency_list:
616 dependency = Dependency(path, dependencies[path])
617 dependency.update_or_notify()
618
619
620if __name__ == "__main__":
621 main()