mquery-toolkit 0.1.0__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.
- mquery_toolkit/THIRD_PARTY_NOTICES.txt +34 -0
- mquery_toolkit/__init__.py +35 -0
- mquery_toolkit/_bridge.cjs +45531 -0
- mquery_toolkit/cli.py +118 -0
- mquery_toolkit/core.py +600 -0
- mquery_toolkit/fabric.py +213 -0
- mquery_toolkit/pqtest.py +93 -0
- mquery_toolkit/py.typed +0 -0
- mquery_toolkit-0.1.0.dist-info/METADATA +189 -0
- mquery_toolkit-0.1.0.dist-info/RECORD +15 -0
- mquery_toolkit-0.1.0.dist-info/WHEEL +5 -0
- mquery_toolkit-0.1.0.dist-info/entry_points.txt +2 -0
- mquery_toolkit-0.1.0.dist-info/licenses/LICENSE +21 -0
- mquery_toolkit-0.1.0.dist-info/licenses/NOTICE +9 -0
- mquery_toolkit-0.1.0.dist-info/top_level.txt +1 -0
mquery_toolkit/core.py
ADDED
|
@@ -0,0 +1,600 @@
|
|
|
1
|
+
"""Offline, deliberately narrow Power Query M source operations."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import contextlib
|
|
6
|
+
import difflib
|
|
7
|
+
import importlib
|
|
8
|
+
import json
|
|
9
|
+
import os
|
|
10
|
+
import re
|
|
11
|
+
import shutil
|
|
12
|
+
import stat
|
|
13
|
+
import subprocess
|
|
14
|
+
import tempfile
|
|
15
|
+
import threading
|
|
16
|
+
import time
|
|
17
|
+
from collections import Counter
|
|
18
|
+
from collections.abc import Callable
|
|
19
|
+
from dataclasses import dataclass
|
|
20
|
+
from functools import lru_cache
|
|
21
|
+
from pathlib import Path
|
|
22
|
+
from typing import Any
|
|
23
|
+
|
|
24
|
+
MAX_BYTES = 10 * 1024 * 1024
|
|
25
|
+
NODE_TIMEOUT_SECONDS = 30
|
|
26
|
+
_IDENTIFIER = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
|
|
27
|
+
_FILE_SUFFIXES = (".pq", ".m", ".pqm")
|
|
28
|
+
_RESERVED = {
|
|
29
|
+
"and",
|
|
30
|
+
"as",
|
|
31
|
+
"each",
|
|
32
|
+
"else",
|
|
33
|
+
"error",
|
|
34
|
+
"false",
|
|
35
|
+
"if",
|
|
36
|
+
"in",
|
|
37
|
+
"is",
|
|
38
|
+
"let",
|
|
39
|
+
"meta",
|
|
40
|
+
"not",
|
|
41
|
+
"null",
|
|
42
|
+
"or",
|
|
43
|
+
"otherwise",
|
|
44
|
+
"section",
|
|
45
|
+
"shared",
|
|
46
|
+
"then",
|
|
47
|
+
"true",
|
|
48
|
+
"try",
|
|
49
|
+
"type",
|
|
50
|
+
}
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class MQueryError(Exception):
|
|
54
|
+
code = "MQUERY_ERROR"
|
|
55
|
+
|
|
56
|
+
def __init__(self, message: str) -> None:
|
|
57
|
+
super().__init__(message)
|
|
58
|
+
self.message = message
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
class NodeError(MQueryError):
|
|
62
|
+
code = "NODE_ERROR"
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
class ParseError(MQueryError):
|
|
66
|
+
code = "M_PARSE_ERROR"
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
class RenameRefusal(MQueryError):
|
|
70
|
+
code = "M_RENAME_REFUSED"
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
class SafeWriteError(MQueryError):
|
|
74
|
+
code = "M_SAFE_WRITE_REFUSED"
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
class AdapterError(MQueryError):
|
|
78
|
+
code = "M_ADAPTER_ERROR"
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
class _ProcessOutputLimit(Exception):
|
|
82
|
+
pass
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
@dataclass(frozen=True)
|
|
86
|
+
class Diagnostic:
|
|
87
|
+
file: str = "<string>"
|
|
88
|
+
line: int = 1
|
|
89
|
+
column: int = 1
|
|
90
|
+
code: str = "M000"
|
|
91
|
+
severity: str = "error"
|
|
92
|
+
message: str = ""
|
|
93
|
+
|
|
94
|
+
def as_dict(self) -> dict[str, Any]:
|
|
95
|
+
return self.__dict__.copy()
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
@dataclass(frozen=True)
|
|
99
|
+
class FileSnapshot:
|
|
100
|
+
data: bytes
|
|
101
|
+
mode: int
|
|
102
|
+
device: int
|
|
103
|
+
inode: int
|
|
104
|
+
size: int
|
|
105
|
+
mtime_ns: int
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def _node_binary() -> str:
|
|
109
|
+
configured = os.environ.get("MQUERY_NODE")
|
|
110
|
+
if configured:
|
|
111
|
+
return configured
|
|
112
|
+
if os.name == "nt":
|
|
113
|
+
found = shutil.which("node", path=os.environ.get("PATH", ""))
|
|
114
|
+
else:
|
|
115
|
+
found = shutil.which("node")
|
|
116
|
+
if found is None:
|
|
117
|
+
raise NodeError("Node.js 22 or newer is required")
|
|
118
|
+
if Path(found).resolve().parent == Path.cwd().resolve():
|
|
119
|
+
raise NodeError("refusing to run node from the current directory")
|
|
120
|
+
return found
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
def _newline(source: str) -> str:
|
|
124
|
+
return "\r\n" if "\r\n" in source else "\n"
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
def _run_process_bounded(
|
|
128
|
+
command: list[str], input_data: bytes | None, timeout: int
|
|
129
|
+
) -> subprocess.CompletedProcess[bytes]:
|
|
130
|
+
deadline = time.monotonic() + timeout
|
|
131
|
+
with subprocess.Popen(
|
|
132
|
+
command,
|
|
133
|
+
stdin=subprocess.PIPE if input_data is not None else subprocess.DEVNULL,
|
|
134
|
+
stdout=subprocess.PIPE,
|
|
135
|
+
stderr=subprocess.PIPE,
|
|
136
|
+
) as process:
|
|
137
|
+
buffers = [bytearray(), bytearray()]
|
|
138
|
+
exceeded = threading.Event()
|
|
139
|
+
|
|
140
|
+
def read(stream: Any, buffer: bytearray) -> None:
|
|
141
|
+
while chunk := stream.read(65536):
|
|
142
|
+
if len(buffer) + len(chunk) > MAX_BYTES:
|
|
143
|
+
exceeded.set()
|
|
144
|
+
process.kill()
|
|
145
|
+
return
|
|
146
|
+
buffer.extend(chunk)
|
|
147
|
+
|
|
148
|
+
threads = [
|
|
149
|
+
threading.Thread(target=read, args=(stream, buffer), daemon=True)
|
|
150
|
+
for stream, buffer in zip(
|
|
151
|
+
(process.stdout, process.stderr), buffers, strict=True
|
|
152
|
+
)
|
|
153
|
+
]
|
|
154
|
+
if process.stdin is not None:
|
|
155
|
+
stdin = process.stdin
|
|
156
|
+
input_payload = input_data or b""
|
|
157
|
+
|
|
158
|
+
def write() -> None:
|
|
159
|
+
try:
|
|
160
|
+
stdin.write(input_payload)
|
|
161
|
+
except BrokenPipeError:
|
|
162
|
+
pass
|
|
163
|
+
finally:
|
|
164
|
+
with contextlib.suppress(OSError):
|
|
165
|
+
stdin.close()
|
|
166
|
+
|
|
167
|
+
threads.append(threading.Thread(target=write, daemon=True))
|
|
168
|
+
for thread in threads:
|
|
169
|
+
thread.start()
|
|
170
|
+
try:
|
|
171
|
+
process.wait(timeout=timeout)
|
|
172
|
+
except subprocess.TimeoutExpired:
|
|
173
|
+
process.kill()
|
|
174
|
+
process.wait()
|
|
175
|
+
raise
|
|
176
|
+
finally:
|
|
177
|
+
for thread in threads:
|
|
178
|
+
thread.join(max(0.0, deadline - time.monotonic()))
|
|
179
|
+
if any(thread.is_alive() for thread in threads):
|
|
180
|
+
if process.poll() is None:
|
|
181
|
+
process.kill()
|
|
182
|
+
# An abandoned thread may still hold the read lock on this
|
|
183
|
+
# pipe (e.g. a grandchild keeping it open); detach it so the
|
|
184
|
+
# context manager's close() below does not block on it.
|
|
185
|
+
process.stdout = process.stderr = process.stdin = None
|
|
186
|
+
raise subprocess.TimeoutExpired(command, timeout)
|
|
187
|
+
if exceeded.is_set():
|
|
188
|
+
raise _ProcessOutputLimit
|
|
189
|
+
return subprocess.CompletedProcess(
|
|
190
|
+
command, process.returncode, *map(bytes, buffers)
|
|
191
|
+
)
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
@lru_cache(maxsize=4)
|
|
195
|
+
def _require_node(binary: str) -> None:
|
|
196
|
+
try:
|
|
197
|
+
result = _run_process_bounded([binary, "--version"], None, 5)
|
|
198
|
+
except (OSError, subprocess.SubprocessError, _ProcessOutputLimit) as error:
|
|
199
|
+
raise NodeError("Node.js 22 or newer is required") from error
|
|
200
|
+
if result.returncode or not re.fullmatch(
|
|
201
|
+
rb"v(2[2-9]|[3-9]\d|\d{3,})\.\d+\.\d+\S*\s*", result.stdout
|
|
202
|
+
):
|
|
203
|
+
raise NodeError("Node.js 22 or newer is required")
|
|
204
|
+
|
|
205
|
+
|
|
206
|
+
def _bridge(source: str, kind: str, **options: str) -> dict[str, Any]:
|
|
207
|
+
stripped = source[1:] if source.startswith("\ufeff") else source
|
|
208
|
+
try:
|
|
209
|
+
payload = json.dumps(
|
|
210
|
+
{"source": stripped, "kind": kind, "newline": _newline(source), **options}
|
|
211
|
+
).encode("utf-8", "strict")
|
|
212
|
+
source_size = len(source.encode("utf-8", "strict"))
|
|
213
|
+
except UnicodeEncodeError as error:
|
|
214
|
+
raise MQueryError("source must be valid UTF-8") from error
|
|
215
|
+
if source_size > MAX_BYTES or len(payload) > MAX_BYTES:
|
|
216
|
+
raise MQueryError("input exceeds 10 MiB")
|
|
217
|
+
bridge = Path(__file__).with_name("_bridge.cjs")
|
|
218
|
+
node = _node_binary()
|
|
219
|
+
_require_node(node)
|
|
220
|
+
try:
|
|
221
|
+
result = _run_process_bounded(
|
|
222
|
+
[node, str(bridge)], payload, NODE_TIMEOUT_SECONDS
|
|
223
|
+
)
|
|
224
|
+
except FileNotFoundError as error:
|
|
225
|
+
raise NodeError("Node.js 22 or newer is required") from error
|
|
226
|
+
except subprocess.TimeoutExpired as error:
|
|
227
|
+
raise NodeError("Node subprocess timed out after 30 seconds") from error
|
|
228
|
+
except _ProcessOutputLimit as error:
|
|
229
|
+
raise NodeError("Node output exceeds 10 MiB") from error
|
|
230
|
+
if result.returncode:
|
|
231
|
+
raise NodeError(f"Node bridge failed with exit {result.returncode}")
|
|
232
|
+
try:
|
|
233
|
+
response: dict[str, Any] = json.loads(result.stdout.decode("utf-8", "strict"))
|
|
234
|
+
except (UnicodeDecodeError, json.JSONDecodeError) as error:
|
|
235
|
+
raise NodeError("Node bridge returned invalid JSON") from error
|
|
236
|
+
if response.get("error") == "PARSE_ERROR":
|
|
237
|
+
raise ParseError(
|
|
238
|
+
f"parse error at {response.get('line', 1)}:{response.get('column', 1)}: "
|
|
239
|
+
f"{response.get('message', 'invalid Power Query source')}"
|
|
240
|
+
)
|
|
241
|
+
if response.get("error"):
|
|
242
|
+
raise NodeError(str(response["error"]))
|
|
243
|
+
return response
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
def parse(source: str) -> dict[str, Any]:
|
|
247
|
+
"""Return a stable, JSON-safe view from Microsoft's pinned parser."""
|
|
248
|
+
return _bridge(source, "parse")
|
|
249
|
+
|
|
250
|
+
|
|
251
|
+
def _preserve_layout(updated: str, original: str) -> str:
|
|
252
|
+
if _newline(original) == "\r\n":
|
|
253
|
+
updated = updated.replace("\r\n", "\n").replace("\n", "\r\n")
|
|
254
|
+
else:
|
|
255
|
+
updated = updated.replace("\r\n", "\n")
|
|
256
|
+
if original.startswith("\ufeff"):
|
|
257
|
+
if not updated.startswith("\ufeff"):
|
|
258
|
+
updated = "\ufeff" + updated
|
|
259
|
+
elif updated.startswith("\ufeff"):
|
|
260
|
+
updated = updated[1:]
|
|
261
|
+
if not original.endswith(("\n", "\r")):
|
|
262
|
+
return updated.rstrip("\r\n")
|
|
263
|
+
if not updated.endswith(("\n", "\r")):
|
|
264
|
+
updated += _newline(original)
|
|
265
|
+
return updated
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
def format_source(source: str) -> str:
|
|
269
|
+
return _preserve_layout(str(_bridge(source, "format")["formatted"]), source)
|
|
270
|
+
|
|
271
|
+
|
|
272
|
+
def dependencies(source: str) -> list[str]:
|
|
273
|
+
return _dependencies_from(parse(source))
|
|
274
|
+
|
|
275
|
+
|
|
276
|
+
def _dependencies_from(parsed: dict[str, Any]) -> list[str]:
|
|
277
|
+
tokens = parsed["tokens"]
|
|
278
|
+
bound = {
|
|
279
|
+
str(binding["name"])
|
|
280
|
+
for binding in (parsed.get("analysis") or {}).get("bindings", [])
|
|
281
|
+
}
|
|
282
|
+
return sorted(
|
|
283
|
+
{
|
|
284
|
+
str(item["text"])
|
|
285
|
+
for index, item in enumerate(tokens[:-1])
|
|
286
|
+
if item["kind"] == "Identifier"
|
|
287
|
+
and tokens[index + 1]["kind"] == "LeftParenthesis"
|
|
288
|
+
and str(item["text"]) not in bound
|
|
289
|
+
}
|
|
290
|
+
)
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
def _rename_plan(source: str, old: str) -> dict[str, Any]:
|
|
294
|
+
if '#"' in source or "[" in source or "=>" in source or not source.isascii():
|
|
295
|
+
raise RenameRefusal(
|
|
296
|
+
"quoted, record, lambda, or non-ASCII rename is unsupported"
|
|
297
|
+
)
|
|
298
|
+
try:
|
|
299
|
+
return _bridge(source, "rename", old=old)
|
|
300
|
+
except NodeError as error:
|
|
301
|
+
raise RenameRefusal(
|
|
302
|
+
"rename supports one unquoted top-level let scope"
|
|
303
|
+
) from error
|
|
304
|
+
|
|
305
|
+
|
|
306
|
+
def rename(source: str, old: str, new: str) -> str:
|
|
307
|
+
"""Rename one unquoted top-level binding and its unambiguous references."""
|
|
308
|
+
has_bom = source.startswith("\ufeff")
|
|
309
|
+
if has_bom:
|
|
310
|
+
source = source[1:]
|
|
311
|
+
if not _IDENTIFIER.fullmatch(old) or not _IDENTIFIER.fullmatch(new) or old == new:
|
|
312
|
+
raise RenameRefusal("rename requires distinct unquoted identifiers")
|
|
313
|
+
if new.lower() in _RESERVED:
|
|
314
|
+
raise RenameRefusal("rename target is a reserved M keyword")
|
|
315
|
+
plan = _rename_plan(source, old)
|
|
316
|
+
declarations = list(plan["bindings"])
|
|
317
|
+
if new in declarations:
|
|
318
|
+
raise RenameRefusal("rename target collides with an existing let binding")
|
|
319
|
+
parsed = parse(source)
|
|
320
|
+
if any(
|
|
321
|
+
token["kind"] == "Identifier" and str(token["text"]) == new
|
|
322
|
+
for token in parsed["tokens"]
|
|
323
|
+
):
|
|
324
|
+
raise RenameRefusal("rename target already appears in the source")
|
|
325
|
+
edits = [(int(start), int(end)) for start, end in plan["spans"]]
|
|
326
|
+
if not edits:
|
|
327
|
+
raise RenameRefusal("target must name exactly one top-level let binding")
|
|
328
|
+
edits = sorted(edits)
|
|
329
|
+
for index in range(len(edits) - 1):
|
|
330
|
+
if edits[index][1] > edits[index + 1][0]:
|
|
331
|
+
raise RenameRefusal("rename spans overlap")
|
|
332
|
+
for start, end in reversed(edits):
|
|
333
|
+
source = source[:start] + new + source[end:]
|
|
334
|
+
if has_bom:
|
|
335
|
+
source = "\ufeff" + source
|
|
336
|
+
parse(source)
|
|
337
|
+
return source
|
|
338
|
+
|
|
339
|
+
|
|
340
|
+
def replace_source(source: str, replacement: str) -> str:
|
|
341
|
+
"""Replace complete source only - never an unsafe partial-text match."""
|
|
342
|
+
try:
|
|
343
|
+
size = len(replacement.encode("utf-8", "strict"))
|
|
344
|
+
except UnicodeEncodeError as error:
|
|
345
|
+
raise MQueryError("source must be valid UTF-8") from error
|
|
346
|
+
if size > MAX_BYTES:
|
|
347
|
+
raise MQueryError("replacement exceeds 10 MiB")
|
|
348
|
+
parse(replacement)
|
|
349
|
+
return replacement
|
|
350
|
+
|
|
351
|
+
|
|
352
|
+
def check(source: str, file: str = "<string>") -> list[Diagnostic]:
|
|
353
|
+
try:
|
|
354
|
+
parsed = parse(source)
|
|
355
|
+
except ParseError as error:
|
|
356
|
+
match = re.search(r"(\d+):(\d+)", error.message)
|
|
357
|
+
line, column = (int(item) for item in match.groups()) if match else (1, 1)
|
|
358
|
+
return [Diagnostic(file, line, column, error.code, "error", error.message)]
|
|
359
|
+
diagnostics: list[Diagnostic] = []
|
|
360
|
+
analysis = parsed.get("analysis") or {}
|
|
361
|
+
bindings = analysis.get("bindings", [])
|
|
362
|
+
names = [str(binding["name"]) for binding in bindings]
|
|
363
|
+
counts = Counter(names)
|
|
364
|
+
for binding in bindings:
|
|
365
|
+
if counts[str(binding["name"])] > 1:
|
|
366
|
+
diagnostics.append(
|
|
367
|
+
Diagnostic(
|
|
368
|
+
file,
|
|
369
|
+
int(binding["line"]),
|
|
370
|
+
int(binding["column"]),
|
|
371
|
+
"M001",
|
|
372
|
+
"error",
|
|
373
|
+
f"duplicate let binding: {binding['name']}",
|
|
374
|
+
)
|
|
375
|
+
)
|
|
376
|
+
|
|
377
|
+
by_name = {str(binding["name"]): binding for binding in bindings}
|
|
378
|
+
result_references = list(analysis.get("resultReferences", []))
|
|
379
|
+
reachable = {
|
|
380
|
+
str(reference["name"])
|
|
381
|
+
for reference in result_references
|
|
382
|
+
if str(reference["name"]) in by_name
|
|
383
|
+
}
|
|
384
|
+
pending = list(reachable)
|
|
385
|
+
while pending:
|
|
386
|
+
binding = by_name[pending.pop()]
|
|
387
|
+
for reference in binding.get("references", []):
|
|
388
|
+
name = str(reference["name"])
|
|
389
|
+
if name in by_name and name not in reachable:
|
|
390
|
+
reachable.add(name)
|
|
391
|
+
pending.append(name)
|
|
392
|
+
for binding in bindings:
|
|
393
|
+
if str(binding["name"]) not in reachable:
|
|
394
|
+
diagnostics.append(
|
|
395
|
+
Diagnostic(
|
|
396
|
+
file,
|
|
397
|
+
int(binding["line"]),
|
|
398
|
+
int(binding["column"]),
|
|
399
|
+
"M004",
|
|
400
|
+
"warning",
|
|
401
|
+
f"unreachable let binding: {binding['name']}",
|
|
402
|
+
)
|
|
403
|
+
)
|
|
404
|
+
|
|
405
|
+
references = result_references + [
|
|
406
|
+
reference for binding in bindings for reference in binding.get("references", [])
|
|
407
|
+
]
|
|
408
|
+
for reference in references:
|
|
409
|
+
name = str(reference["name"])
|
|
410
|
+
if name not in by_name and "." not in name:
|
|
411
|
+
diagnostics.append(
|
|
412
|
+
Diagnostic(
|
|
413
|
+
file,
|
|
414
|
+
int(reference["line"]),
|
|
415
|
+
int(reference["column"]),
|
|
416
|
+
"M005",
|
|
417
|
+
"warning",
|
|
418
|
+
f"unresolved unqualified reference: {name}",
|
|
419
|
+
)
|
|
420
|
+
)
|
|
421
|
+
|
|
422
|
+
for web_match in re.finditer(r"Web\.Contents\s*\(\s*", source):
|
|
423
|
+
if source[web_match.end() : web_match.end() + 1] != '"':
|
|
424
|
+
line = source.count("\n", 0, web_match.start()) + 1
|
|
425
|
+
column = web_match.start() - source.rfind("\n", 0, web_match.start())
|
|
426
|
+
diagnostics.append(
|
|
427
|
+
Diagnostic(
|
|
428
|
+
file, line, column, "M002", "warning", "dynamic Web.Contents URL"
|
|
429
|
+
)
|
|
430
|
+
)
|
|
431
|
+
for credential_match in re.finditer(
|
|
432
|
+
r"(?i)(password|token|secret)\s*=\s*\"", source
|
|
433
|
+
):
|
|
434
|
+
line = source.count("\n", 0, credential_match.start()) + 1
|
|
435
|
+
column = credential_match.start() - source.rfind(
|
|
436
|
+
"\n", 0, credential_match.start()
|
|
437
|
+
)
|
|
438
|
+
diagnostics.append(
|
|
439
|
+
Diagnostic(file, line, column, "M003", "warning", "credential-like literal")
|
|
440
|
+
)
|
|
441
|
+
for dependency in _dependencies_from(parsed):
|
|
442
|
+
if dependency.endswith(".Contents"):
|
|
443
|
+
diagnostics.append(
|
|
444
|
+
Diagnostic(file, 1, 1, "M006", "info", f"source function: {dependency}")
|
|
445
|
+
)
|
|
446
|
+
return diagnostics
|
|
447
|
+
|
|
448
|
+
|
|
449
|
+
def _snapshot(path: Path) -> FileSnapshot:
|
|
450
|
+
flags = os.O_RDONLY | getattr(os, "O_NOFOLLOW", 0)
|
|
451
|
+
try:
|
|
452
|
+
info = os.lstat(path)
|
|
453
|
+
if stat.S_ISLNK(info.st_mode) or (
|
|
454
|
+
getattr(info, "st_file_attributes", 0)
|
|
455
|
+
& getattr(stat, "FILE_ATTRIBUTE_REPARSE_POINT", 0)
|
|
456
|
+
):
|
|
457
|
+
raise SafeWriteError(
|
|
458
|
+
"writes require a regular, non-symlink, single-link file"
|
|
459
|
+
)
|
|
460
|
+
descriptor = os.open(path, flags)
|
|
461
|
+
except OSError as error:
|
|
462
|
+
raise SafeWriteError(
|
|
463
|
+
"writes require a regular, non-symlink, single-link file"
|
|
464
|
+
) from error
|
|
465
|
+
try:
|
|
466
|
+
info = os.fstat(descriptor)
|
|
467
|
+
if not stat.S_ISREG(info.st_mode) or info.st_nlink != 1:
|
|
468
|
+
raise SafeWriteError(
|
|
469
|
+
"writes require a regular, non-symlink, single-link file"
|
|
470
|
+
)
|
|
471
|
+
if info.st_size > MAX_BYTES:
|
|
472
|
+
raise SafeWriteError("input exceeds 10 MiB")
|
|
473
|
+
data = bytearray()
|
|
474
|
+
while True:
|
|
475
|
+
chunk = os.read(descriptor, min(65536, MAX_BYTES + 1 - len(data)))
|
|
476
|
+
if not chunk:
|
|
477
|
+
break
|
|
478
|
+
data.extend(chunk)
|
|
479
|
+
if len(data) > MAX_BYTES:
|
|
480
|
+
raise SafeWriteError("input exceeds 10 MiB")
|
|
481
|
+
finally:
|
|
482
|
+
os.close(descriptor)
|
|
483
|
+
return FileSnapshot(
|
|
484
|
+
bytes(data),
|
|
485
|
+
stat.S_IMODE(info.st_mode),
|
|
486
|
+
info.st_dev,
|
|
487
|
+
info.st_ino,
|
|
488
|
+
info.st_size,
|
|
489
|
+
info.st_mtime_ns,
|
|
490
|
+
)
|
|
491
|
+
|
|
492
|
+
|
|
493
|
+
def _diff(path: Path, original: bytes, updated: str) -> str:
|
|
494
|
+
try:
|
|
495
|
+
original_text = original.decode("utf-8", "strict")
|
|
496
|
+
except UnicodeDecodeError as error:
|
|
497
|
+
raise SafeWriteError("source must be valid UTF-8") from error
|
|
498
|
+
return "".join(
|
|
499
|
+
difflib.unified_diff(
|
|
500
|
+
original_text.splitlines(True),
|
|
501
|
+
updated.splitlines(True),
|
|
502
|
+
fromfile=str(path),
|
|
503
|
+
tofile=str(path),
|
|
504
|
+
)
|
|
505
|
+
)
|
|
506
|
+
|
|
507
|
+
|
|
508
|
+
def _lock_file(descriptor: int) -> None:
|
|
509
|
+
if os.name == "nt":
|
|
510
|
+
if os.fstat(descriptor).st_size == 0:
|
|
511
|
+
os.write(descriptor, b"\0")
|
|
512
|
+
os.lseek(descriptor, 0, os.SEEK_SET)
|
|
513
|
+
locker = importlib.import_module("msvcrt")
|
|
514
|
+
locker.locking(descriptor, locker.LK_LOCK, 1)
|
|
515
|
+
else:
|
|
516
|
+
locker = importlib.import_module("fcntl")
|
|
517
|
+
locker.flock(descriptor, locker.LOCK_EX)
|
|
518
|
+
|
|
519
|
+
|
|
520
|
+
def _unlock_file(descriptor: int) -> None:
|
|
521
|
+
if os.name == "nt":
|
|
522
|
+
os.lseek(descriptor, 0, os.SEEK_SET)
|
|
523
|
+
locker = importlib.import_module("msvcrt")
|
|
524
|
+
locker.locking(descriptor, locker.LK_UNLCK, 1)
|
|
525
|
+
else:
|
|
526
|
+
locker = importlib.import_module("fcntl")
|
|
527
|
+
locker.flock(descriptor, locker.LOCK_UN)
|
|
528
|
+
|
|
529
|
+
|
|
530
|
+
def update_file(
|
|
531
|
+
path: Path, transform: Callable[[str], str], write: bool = False
|
|
532
|
+
) -> str:
|
|
533
|
+
if path.suffix not in _FILE_SUFFIXES and not path.name.endswith(".query.pq"):
|
|
534
|
+
raise SafeWriteError("unsupported source file extension")
|
|
535
|
+
if not write:
|
|
536
|
+
snapshot = _snapshot(path)
|
|
537
|
+
try:
|
|
538
|
+
original = snapshot.data.decode("utf-8", "strict")
|
|
539
|
+
except UnicodeDecodeError as error:
|
|
540
|
+
raise SafeWriteError("source must be valid UTF-8") from error
|
|
541
|
+
updated = _preserve_layout(transform(original), original)
|
|
542
|
+
return _diff(path, snapshot.data, updated)
|
|
543
|
+
lock = path.with_name(f".{path.name}.lock")
|
|
544
|
+
lock_flags = os.O_CREAT | os.O_RDWR | getattr(os, "O_NOFOLLOW", 0)
|
|
545
|
+
temporary: Path | None = None
|
|
546
|
+
try:
|
|
547
|
+
lock_fd = os.open(lock, lock_flags, 0o600)
|
|
548
|
+
except OSError as error:
|
|
549
|
+
raise SafeWriteError("unable to acquire safe source lock") from error
|
|
550
|
+
acquired = False
|
|
551
|
+
try:
|
|
552
|
+
lock_info = os.fstat(lock_fd)
|
|
553
|
+
if not stat.S_ISREG(lock_info.st_mode) or lock_info.st_nlink != 1:
|
|
554
|
+
raise SafeWriteError("source lock must be a regular single-link file")
|
|
555
|
+
try:
|
|
556
|
+
_lock_file(lock_fd)
|
|
557
|
+
except OSError as error:
|
|
558
|
+
raise SafeWriteError("unable to acquire safe source lock") from error
|
|
559
|
+
acquired = True
|
|
560
|
+
snapshot = _snapshot(path)
|
|
561
|
+
try:
|
|
562
|
+
original = snapshot.data.decode("utf-8", "strict")
|
|
563
|
+
except UnicodeDecodeError as error:
|
|
564
|
+
raise SafeWriteError("source must be valid UTF-8") from error
|
|
565
|
+
updated = _preserve_layout(transform(original), original)
|
|
566
|
+
diff = _diff(path, snapshot.data, updated)
|
|
567
|
+
if updated == original:
|
|
568
|
+
return diff
|
|
569
|
+
if _snapshot(path) != snapshot:
|
|
570
|
+
raise SafeWriteError("source changed during operation")
|
|
571
|
+
try:
|
|
572
|
+
descriptor, name = tempfile.mkstemp(
|
|
573
|
+
dir=path.parent, prefix=f".{path.name}.", suffix=".tmp"
|
|
574
|
+
)
|
|
575
|
+
except OSError as error:
|
|
576
|
+
raise SafeWriteError("unable to create temporary file") from error
|
|
577
|
+
temporary = Path(name)
|
|
578
|
+
with os.fdopen(descriptor, "wb") as handle:
|
|
579
|
+
handle.write(updated.encode())
|
|
580
|
+
handle.flush()
|
|
581
|
+
os.fsync(handle.fileno())
|
|
582
|
+
os.chmod(temporary, snapshot.mode)
|
|
583
|
+
if _snapshot(path) != snapshot:
|
|
584
|
+
raise SafeWriteError("source changed before atomic replacement")
|
|
585
|
+
os.replace(temporary, path)
|
|
586
|
+
return diff
|
|
587
|
+
finally:
|
|
588
|
+
if temporary is not None:
|
|
589
|
+
with contextlib.suppress(OSError):
|
|
590
|
+
temporary.unlink()
|
|
591
|
+
try:
|
|
592
|
+
if acquired:
|
|
593
|
+
with contextlib.suppress(OSError):
|
|
594
|
+
_unlock_file(lock_fd)
|
|
595
|
+
finally:
|
|
596
|
+
os.close(lock_fd)
|
|
597
|
+
# Lock removal is best-effort; the snapshot re-checks above remain
|
|
598
|
+
# the correctness guard against a concurrent writer.
|
|
599
|
+
with contextlib.suppress(OSError):
|
|
600
|
+
os.unlink(lock)
|