Parent directory

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()