codekingpro/portable-devtools
114k
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 