Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
tool_support.py202 linesDownload Raw Back to util
1# util/tool_support.py
2# Copyright (C) 2005-2024 the SQLAlchemy authors and contributors
3# <see AUTHORS file>
4#
5# This module is part of SQLAlchemy and is released under
6# the MIT License: https://www.opensource.org/licenses/mit-license.php
7# mypy: allow-untyped-defs, allow-untyped-calls
8"""support routines for the helpers in tools/.
9
10These aren't imported by the enclosing util package as the are not
11needed for normal library use.
12
13"""
14from __future__ import annotations
15
16from argparse import ArgumentParser
17from argparse import Namespace
18import contextlib
19import difflib
20import os
21from pathlib import Path
22import shlex
23import shutil
24import subprocess
25import sys
26from typing import Any
27from typing import Dict
28from typing import Iterator
29from typing import Optional
30from typing import Union
31
32from . import compat
33
34
35class code_writer_cmd:
36    parser: ArgumentParser
37    args: Namespace
38    suppress_output: bool
39    diffs_detected: bool
40    source_root: Path
41    pyproject_toml_path: Path
42
43    def __init__(self, tool_script: str):
44        self.source_root = Path(tool_script).parent.parent
45        self.pyproject_toml_path = self.source_root / Path("pyproject.toml")
46        assert self.pyproject_toml_path.exists()
47
48        self.parser = ArgumentParser()
49        self.parser.add_argument(
50            "--stdout",
51            action="store_true",
52            help="Write to stdout instead of saving to file",
53        )
54        self.parser.add_argument(
55            "-c",
56            "--check",
57            help="Don't write the files back, just return the "
58            "status. Return code 0 means nothing would change. "
59            "Return code 1 means some files would be reformatted",
60            action="store_true",
61        )
62
63    def run_zimports(self, tempfile: str) -> None:
64        self._run_console_script(
65            str(tempfile),
66            {
67                "entrypoint": "zimports",
68                "options": f"--toml-config {self.pyproject_toml_path}",
69            },
70        )
71
72    def run_black(self, tempfile: str) -> None:
73        self._run_console_script(
74            str(tempfile),
75            {
76                "entrypoint": "black",
77                "options": f"--config {self.pyproject_toml_path}",
78            },
79        )
80
81    def _run_console_script(self, path: str, options: Dict[str, Any]) -> None:
82        """Run a Python console application from within the process.
83
84        Used for black, zimports
85
86        """
87
88        is_posix = os.name == "posix"
89
90        entrypoint_name = options["entrypoint"]
91
92        for entry in compat.importlib_metadata_get("console_scripts"):
93            if entry.name == entrypoint_name:
94                impl = entry
95                break
96        else:
97            raise Exception(
98                f"Could not find entrypoint console_scripts.{entrypoint_name}"
99            )
100        cmdline_options_str = options.get("options", "")
101        cmdline_options_list = shlex.split(
102            cmdline_options_str, posix=is_posix
103        ) + [path]
104
105        kw: Dict[str, Any] = {}
106        if self.suppress_output:
107            kw["stdout"] = kw["stderr"] = subprocess.DEVNULL
108
109        subprocess.run(
110            [
111                sys.executable,
112                "-c",
113                "import %s; %s.%s()" % (impl.module, impl.module, impl.attr),
114            ]
115            + cmdline_options_list,
116            cwd=str(self.source_root),
117            **kw,
118        )
119
120    def write_status(self, *text: str) -> None:
121        if not self.suppress_output:
122            sys.stderr.write(" ".join(text))
123
124    def write_output_file_from_text(
125        self, text: str, destination_path: Union[str, Path]
126    ) -> None:
127        if self.args.check:
128            self._run_diff(destination_path, source=text)
129        elif self.args.stdout:
130            print(text)
131        else:
132            self.write_status(f"Writing {destination_path}...")
133            Path(destination_path).write_text(
134                text, encoding="utf-8", newline="\n"
135            )
136            self.write_status("done\n")
137
138    def write_output_file_from_tempfile(
139        self, tempfile: str, destination_path: str
140    ) -> None:
141        if self.args.check:
142            self._run_diff(destination_path, source_file=tempfile)
143            os.unlink(tempfile)
144        elif self.args.stdout:
145            with open(tempfile) as tf:
146                print(tf.read())
147            os.unlink(tempfile)
148        else:
149            self.write_status(f"Writing {destination_path}...")
150            shutil.move(tempfile, destination_path)
151            self.write_status("done\n")
152
153    def _run_diff(
154        self,
155        destination_path: Union[str, Path],
156        *,
157        source: Optional[str] = None,
158        source_file: Optional[str] = None,
159    ) -> None:
160        if source_file:
161            with open(source_file, encoding="utf-8") as tf:
162                source_lines = list(tf)
163        elif source is not None:
164            source_lines = source.splitlines(keepends=True)
165        else:
166            assert False, "source or source_file is required"
167
168        with open(destination_path, encoding="utf-8") as dp:
169            d = difflib.unified_diff(
170                list(dp),
171                source_lines,
172                fromfile=Path(destination_path).as_posix(),
173                tofile="<proposed changes>",
174                n=3,
175                lineterm="\n",
176            )
177            d_as_list = list(d)
178            if d_as_list:
179                self.diffs_detected = True
180                print("".join(d_as_list))
181
182    @contextlib.contextmanager
183    def add_arguments(self) -> Iterator[ArgumentParser]:
184        yield self.parser
185
186    @contextlib.contextmanager
187    def run_program(self) -> Iterator[None]:
188        self.args = self.parser.parse_args()
189        if self.args.check:
190            self.diffs_detected = False
191            self.suppress_output = True
192        elif self.args.stdout:
193            self.suppress_output = True
194        else:
195            self.suppress_output = False
196        yield
197
198        if self.args.check and self.diffs_detected:
199            sys.exit(1)
200        else:
201            sys.exit(0)
202 
codekingpro/portable-devtools · Team Ai