upload-guard 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.
- upload_guard/__init__.py +41 -0
- upload_guard/_version.py +1 -0
- upload_guard/adapters.py +214 -0
- upload_guard/cli.py +97 -0
- upload_guard/exceptions.py +33 -0
- upload_guard/extension.py +179 -0
- upload_guard/guard.py +266 -0
- upload_guard/integrations/__init__.py +9 -0
- upload_guard/integrations/_http.py +31 -0
- upload_guard/integrations/django.py +79 -0
- upload_guard/integrations/fastapi.py +73 -0
- upload_guard/integrations/flask.py +50 -0
- upload_guard/py.typed +0 -0
- upload_guard/security/__init__.py +12 -0
- upload_guard/security/archive.py +614 -0
- upload_guard/security/filename.py +169 -0
- upload_guard/security/polyglot.py +185 -0
- upload_guard/security/svg.py +411 -0
- upload_guard/sniff/__init__.py +4 -0
- upload_guard/sniff/containers.py +471 -0
- upload_guard/sniff/detector.py +56 -0
- upload_guard/sniff/signatures.py +293 -0
- upload_guard/sniff/text.py +173 -0
- upload_guard/types.py +182 -0
- upload_guard-0.1.0.dist-info/METADATA +419 -0
- upload_guard-0.1.0.dist-info/RECORD +30 -0
- upload_guard-0.1.0.dist-info/WHEEL +5 -0
- upload_guard-0.1.0.dist-info/entry_points.txt +2 -0
- upload_guard-0.1.0.dist-info/licenses/LICENSE +21 -0
- upload_guard-0.1.0.dist-info/top_level.txt +1 -0
upload_guard/__init__.py
ADDED
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
"""upload_guard: pure-Python upload validation.
|
|
2
|
+
|
|
3
|
+
Quick start::
|
|
4
|
+
|
|
5
|
+
from upload_guard import UploadGuard, UploadRejected
|
|
6
|
+
|
|
7
|
+
guard = UploadGuard(allowed=["image/*", "application/pdf"], max_size=10 * 1024 * 1024)
|
|
8
|
+
try:
|
|
9
|
+
result = guard.check(upload) # bytes, path, file object, FastAPI/Django/Flask upload
|
|
10
|
+
except UploadRejected as exc:
|
|
11
|
+
print(exc.codes) # e.g. ['extension.mismatch']
|
|
12
|
+
else:
|
|
13
|
+
print(result.mime, result.safe_filename, result.sanitized)
|
|
14
|
+
"""
|
|
15
|
+
from ._version import __version__
|
|
16
|
+
from .exceptions import UnsupportedSource, UploadGuardError, UploadRejected
|
|
17
|
+
from .extension import TypeRule, canonical_extension, extension_matches, normalize_extension, split_filename
|
|
18
|
+
from .guard import Policy, UploadGuard, check, detect_type, scan
|
|
19
|
+
from .security.archive import ArchiveLimits, inspect_archive, inspect_compressed, inspect_zip
|
|
20
|
+
from .security.filename import DANGEROUS_EXTENSIONS, check_filename, sanitize_filename
|
|
21
|
+
from .security.polyglot import check_trailing, scan_embedded
|
|
22
|
+
from .security.svg import SvgPolicy, SvgResult, sanitize_svg, scan_svg
|
|
23
|
+
from .sniff import HEADER_SIZE, detect, detect_stream
|
|
24
|
+
from .types import ArchiveReport, DetectedType, Finding, ScanResult, Severity
|
|
25
|
+
|
|
26
|
+
__all__ = [
|
|
27
|
+
"__version__",
|
|
28
|
+
# core
|
|
29
|
+
"UploadGuard", "Policy", "check", "scan", "detect_type",
|
|
30
|
+
"ScanResult", "Finding", "Severity", "DetectedType", "ArchiveReport",
|
|
31
|
+
"UploadGuardError", "UploadRejected", "UnsupportedSource",
|
|
32
|
+
# sniffing
|
|
33
|
+
"detect", "detect_stream", "HEADER_SIZE",
|
|
34
|
+
# extensions
|
|
35
|
+
"TypeRule", "canonical_extension", "normalize_extension", "extension_matches", "split_filename",
|
|
36
|
+
# security
|
|
37
|
+
"check_filename", "sanitize_filename", "DANGEROUS_EXTENSIONS",
|
|
38
|
+
"SvgPolicy", "SvgResult", "scan_svg", "sanitize_svg",
|
|
39
|
+
"ArchiveLimits", "inspect_archive", "inspect_zip", "inspect_compressed",
|
|
40
|
+
"scan_embedded", "check_trailing",
|
|
41
|
+
]
|
upload_guard/_version.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "0.1.0"
|
upload_guard/adapters.py
ADDED
|
@@ -0,0 +1,214 @@
|
|
|
1
|
+
"""Turn *anything upload-shaped* into a seekable binary stream plus metadata.
|
|
2
|
+
|
|
3
|
+
Supported inputs (duck-typed, no framework imports):
|
|
4
|
+
|
|
5
|
+
* ``bytes`` / ``bytearray`` / ``memoryview``
|
|
6
|
+
* ``str`` / ``os.PathLike`` paths
|
|
7
|
+
* any binary file object (``io.BytesIO``, open files, ``tempfile`` objects)
|
|
8
|
+
* Starlette / FastAPI ``UploadFile`` (``.file``, ``.filename``, ``.content_type``, ``.size``)
|
|
9
|
+
* Django ``UploadedFile`` (``.name``, ``.size``, ``.content_type``, ``.read``/``.seek``)
|
|
10
|
+
* Werkzeug / Flask ``FileStorage`` (``.stream``, ``.filename``, ``.mimetype``)
|
|
11
|
+
|
|
12
|
+
The stream position is captured on entry and restored on :meth:`FileSource.close`,
|
|
13
|
+
so frameworks can still save the file afterwards.
|
|
14
|
+
"""
|
|
15
|
+
from __future__ import annotations
|
|
16
|
+
|
|
17
|
+
import io
|
|
18
|
+
import os
|
|
19
|
+
import tempfile
|
|
20
|
+
from typing import Any, BinaryIO, Optional
|
|
21
|
+
|
|
22
|
+
from .exceptions import UnsupportedSource
|
|
23
|
+
|
|
24
|
+
_SPOOL_MAX = 8 * 1024 * 1024
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
class FileSource:
|
|
28
|
+
"""A seekable binary stream with optional ``filename`` / ``content_type`` / ``size``."""
|
|
29
|
+
|
|
30
|
+
def __init__(
|
|
31
|
+
self,
|
|
32
|
+
stream: BinaryIO,
|
|
33
|
+
*,
|
|
34
|
+
filename: Optional[str] = None,
|
|
35
|
+
content_type: Optional[str] = None,
|
|
36
|
+
size: Optional[int] = None,
|
|
37
|
+
owns_stream: bool = False,
|
|
38
|
+
oversized: bool = False,
|
|
39
|
+
) -> None:
|
|
40
|
+
self.stream = stream
|
|
41
|
+
self.filename = filename
|
|
42
|
+
self.content_type = content_type
|
|
43
|
+
self._size = size
|
|
44
|
+
self.owns_stream = owns_stream
|
|
45
|
+
self.oversized = oversized # non-seekable input exceeded the read cap
|
|
46
|
+
try:
|
|
47
|
+
self._start_pos = stream.tell()
|
|
48
|
+
except (OSError, ValueError, AttributeError):
|
|
49
|
+
self._start_pos = 0
|
|
50
|
+
|
|
51
|
+
# -- construction -------------------------------------------------------------
|
|
52
|
+
@classmethod
|
|
53
|
+
def from_any(
|
|
54
|
+
cls,
|
|
55
|
+
obj: Any,
|
|
56
|
+
*,
|
|
57
|
+
filename: Optional[str] = None,
|
|
58
|
+
content_type: Optional[str] = None,
|
|
59
|
+
max_read: Optional[int] = None,
|
|
60
|
+
) -> "FileSource":
|
|
61
|
+
if isinstance(obj, FileSource):
|
|
62
|
+
if filename is not None:
|
|
63
|
+
obj.filename = filename
|
|
64
|
+
if content_type is not None:
|
|
65
|
+
obj.content_type = content_type
|
|
66
|
+
return obj
|
|
67
|
+
|
|
68
|
+
if isinstance(obj, (bytes, bytearray, memoryview)):
|
|
69
|
+
data = bytes(obj)
|
|
70
|
+
return cls(io.BytesIO(data), filename=filename, content_type=content_type, size=len(data), owns_stream=True)
|
|
71
|
+
|
|
72
|
+
if isinstance(obj, (str, os.PathLike)):
|
|
73
|
+
path = os.fspath(obj)
|
|
74
|
+
fh = open(path, "rb") # noqa: SIM115 - closed by FileSource.close()
|
|
75
|
+
size = os.fstat(fh.fileno()).st_size
|
|
76
|
+
return cls(fh, filename=filename or os.path.basename(path), content_type=content_type, size=size, owns_stream=True)
|
|
77
|
+
|
|
78
|
+
# Starlette / FastAPI UploadFile
|
|
79
|
+
inner = getattr(obj, "file", None)
|
|
80
|
+
if inner is not None and hasattr(inner, "read") and hasattr(obj, "filename"):
|
|
81
|
+
return cls(
|
|
82
|
+
_ensure_seekable(inner, max_read)[0],
|
|
83
|
+
filename=filename or _str_or_none(getattr(obj, "filename", None)),
|
|
84
|
+
content_type=content_type or _str_or_none(getattr(obj, "content_type", None)),
|
|
85
|
+
size=_int_or_none(getattr(obj, "size", None)),
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
# Werkzeug FileStorage
|
|
89
|
+
inner = getattr(obj, "stream", None)
|
|
90
|
+
if inner is not None and hasattr(inner, "read") and hasattr(obj, "filename"):
|
|
91
|
+
stream, owns, oversized = _ensure_seekable(inner, max_read)
|
|
92
|
+
return cls(
|
|
93
|
+
stream,
|
|
94
|
+
filename=filename or _str_or_none(getattr(obj, "filename", None)),
|
|
95
|
+
content_type=content_type or _str_or_none(getattr(obj, "mimetype", None) or getattr(obj, "content_type", None)),
|
|
96
|
+
size=_int_or_none(getattr(obj, "content_length", None)) or None,
|
|
97
|
+
owns_stream=owns,
|
|
98
|
+
oversized=oversized,
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
# Django UploadedFile (proxies read/seek to .file) or any file-like object
|
|
102
|
+
if hasattr(obj, "read"):
|
|
103
|
+
django_name = getattr(obj, "name", None) if hasattr(obj, "chunks") else None
|
|
104
|
+
stream, owns, oversized = _ensure_seekable(obj, max_read)
|
|
105
|
+
name = filename
|
|
106
|
+
if name is None:
|
|
107
|
+
raw_name = django_name if django_name is not None else getattr(obj, "name", None)
|
|
108
|
+
if isinstance(raw_name, str):
|
|
109
|
+
name = os.path.basename(raw_name) if not hasattr(obj, "chunks") else raw_name
|
|
110
|
+
ctype = content_type or _str_or_none(getattr(obj, "content_type", None))
|
|
111
|
+
size = _int_or_none(getattr(obj, "size", None)) if hasattr(obj, "chunks") else None
|
|
112
|
+
return cls(stream, filename=name, content_type=ctype, size=size, owns_stream=owns, oversized=oversized)
|
|
113
|
+
|
|
114
|
+
raise UnsupportedSource("cannot read uploads from %r" % type(obj).__name__)
|
|
115
|
+
|
|
116
|
+
# -- reading helpers -------------------------------------------------------------
|
|
117
|
+
@property
|
|
118
|
+
def size(self) -> Optional[int]:
|
|
119
|
+
if self._size is None:
|
|
120
|
+
try:
|
|
121
|
+
pos = self.stream.tell()
|
|
122
|
+
self.stream.seek(0, os.SEEK_END)
|
|
123
|
+
self._size = self.stream.tell()
|
|
124
|
+
self.stream.seek(pos)
|
|
125
|
+
except (OSError, ValueError, AttributeError):
|
|
126
|
+
return None
|
|
127
|
+
return self._size
|
|
128
|
+
|
|
129
|
+
def read_at(self, offset: int, n: int) -> bytes:
|
|
130
|
+
pos = self.stream.tell()
|
|
131
|
+
try:
|
|
132
|
+
self.stream.seek(offset)
|
|
133
|
+
return self.stream.read(n) or b""
|
|
134
|
+
finally:
|
|
135
|
+
self.stream.seek(pos)
|
|
136
|
+
|
|
137
|
+
def read_head(self, n: int) -> bytes:
|
|
138
|
+
return self.read_at(0, n)
|
|
139
|
+
|
|
140
|
+
def read_all(self, cap: Optional[int] = None) -> Optional[bytes]:
|
|
141
|
+
"""Whole content, or None when it exceeds ``cap`` bytes."""
|
|
142
|
+
if cap is not None and self.size is not None and self.size > cap:
|
|
143
|
+
return None
|
|
144
|
+
pos = self.stream.tell()
|
|
145
|
+
try:
|
|
146
|
+
self.stream.seek(0)
|
|
147
|
+
if cap is None:
|
|
148
|
+
return self.stream.read() or b""
|
|
149
|
+
data = self.stream.read(cap + 1) or b""
|
|
150
|
+
return None if len(data) > cap else data
|
|
151
|
+
finally:
|
|
152
|
+
self.stream.seek(pos)
|
|
153
|
+
|
|
154
|
+
# -- lifecycle -------------------------------------------------------------------
|
|
155
|
+
def restore(self) -> None:
|
|
156
|
+
try:
|
|
157
|
+
self.stream.seek(self._start_pos)
|
|
158
|
+
except (OSError, ValueError, AttributeError):
|
|
159
|
+
pass
|
|
160
|
+
|
|
161
|
+
def close(self) -> None:
|
|
162
|
+
if self.owns_stream:
|
|
163
|
+
try:
|
|
164
|
+
self.stream.close()
|
|
165
|
+
except Exception: # noqa: BLE001
|
|
166
|
+
pass
|
|
167
|
+
else:
|
|
168
|
+
self.restore()
|
|
169
|
+
|
|
170
|
+
def __enter__(self) -> "FileSource":
|
|
171
|
+
return self
|
|
172
|
+
|
|
173
|
+
def __exit__(self, *exc: object) -> None:
|
|
174
|
+
self.close()
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
def _str_or_none(value: Any) -> Optional[str]:
|
|
178
|
+
return value if isinstance(value, str) and value else None
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
def _int_or_none(value: Any) -> Optional[int]:
|
|
182
|
+
return value if isinstance(value, int) and not isinstance(value, bool) and value >= 0 else None
|
|
183
|
+
|
|
184
|
+
|
|
185
|
+
def _is_seekable(stream: Any) -> bool:
|
|
186
|
+
try:
|
|
187
|
+
if hasattr(stream, "seekable") and not stream.seekable():
|
|
188
|
+
return False
|
|
189
|
+
pos = stream.tell()
|
|
190
|
+
stream.seek(pos)
|
|
191
|
+
return True
|
|
192
|
+
except (OSError, ValueError, AttributeError, io.UnsupportedOperation):
|
|
193
|
+
return False
|
|
194
|
+
|
|
195
|
+
|
|
196
|
+
def _ensure_seekable(stream: Any, max_read: Optional[int]):
|
|
197
|
+
"""Return ``(seekable_stream, owns, oversized)``; spools non-seekable input to a temp file."""
|
|
198
|
+
if _is_seekable(stream):
|
|
199
|
+
return stream, False, False
|
|
200
|
+
spool = tempfile.SpooledTemporaryFile(max_size=_SPOOL_MAX)
|
|
201
|
+
remaining = None if max_read is None else max_read + 1
|
|
202
|
+
oversized = False
|
|
203
|
+
while True:
|
|
204
|
+
chunk = stream.read(64 * 1024 if remaining is None else min(64 * 1024, remaining))
|
|
205
|
+
if not chunk:
|
|
206
|
+
break
|
|
207
|
+
spool.write(chunk)
|
|
208
|
+
if remaining is not None:
|
|
209
|
+
remaining -= len(chunk)
|
|
210
|
+
if remaining <= 0:
|
|
211
|
+
oversized = True
|
|
212
|
+
break
|
|
213
|
+
spool.seek(0)
|
|
214
|
+
return spool, True, oversized
|
upload_guard/cli.py
ADDED
|
@@ -0,0 +1,97 @@
|
|
|
1
|
+
"""``upload_guard`` command line interface."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import argparse
|
|
5
|
+
import json
|
|
6
|
+
import sys
|
|
7
|
+
from typing import List, Optional
|
|
8
|
+
|
|
9
|
+
from ._version import __version__
|
|
10
|
+
from .guard import UploadGuard, detect_type
|
|
11
|
+
from .types import Severity
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def _build_parser() -> argparse.ArgumentParser:
|
|
15
|
+
parser = argparse.ArgumentParser(prog="upload_guard", description="Validate uploaded files: magic-number sniffing, extension checks, SVG/zip-bomb/path-traversal protection.")
|
|
16
|
+
parser.add_argument("--version", action="version", version="upload_guard %s" % __version__)
|
|
17
|
+
sub = parser.add_subparsers(dest="command", required=True)
|
|
18
|
+
|
|
19
|
+
scan = sub.add_parser("scan", help="run the full policy against one or more files")
|
|
20
|
+
scan.add_argument("files", nargs="+")
|
|
21
|
+
scan.add_argument("--allow", action="append", default=None, metavar="RULE", help='allowed type rule, e.g. "image/*", "application/pdf", ".docx" (repeatable)')
|
|
22
|
+
scan.add_argument("--block", action="append", default=None, metavar="RULE", help="blocked type rule (repeatable; default blocks executables and scripts)")
|
|
23
|
+
scan.add_argument("--max-size", type=int, default=None, metavar="BYTES")
|
|
24
|
+
scan.add_argument("--threshold", choices=[s.name for s in Severity], default="MEDIUM", help="minimum severity that rejects (default MEDIUM)")
|
|
25
|
+
scan.add_argument("--no-sanitize", action="store_true", help="report SVG problems instead of sanitising")
|
|
26
|
+
scan.add_argument("--lenient-extension", action="store_true", help="do not require the extension to match the detected type")
|
|
27
|
+
scan.add_argument("--json", action="store_true", help="machine readable output")
|
|
28
|
+
scan.add_argument("--write-sanitized", metavar="PATH", help="write the sanitised SVG (single file only)")
|
|
29
|
+
|
|
30
|
+
det = sub.add_parser("detect", help="print the detected type only")
|
|
31
|
+
det.add_argument("files", nargs="+")
|
|
32
|
+
det.add_argument("--json", action="store_true")
|
|
33
|
+
return parser
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def _print_result(path: str, result) -> None:
|
|
37
|
+
status = "OK " if result.ok else "FAIL"
|
|
38
|
+
print("%s %s -> %s (%s)%s" % (status, path, result.detected.mime, result.detected.description, "" if result.size is None else " %d bytes" % result.size))
|
|
39
|
+
for f in result.findings:
|
|
40
|
+
tag = "fixed" if f.remediated else f.severity.name.lower()
|
|
41
|
+
print(" [%-8s] %-28s %s" % (tag, f.code, f.message))
|
|
42
|
+
if result.archive:
|
|
43
|
+
a = result.archive
|
|
44
|
+
print(" archive: %s, %d entries, %d -> %d bytes%s" % (a.format, a.entry_count, a.total_compressed, a.total_uncompressed, " (truncated)" if a.truncated else ""))
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def main(argv: Optional[List[str]] = None) -> int:
|
|
48
|
+
parser = _build_parser()
|
|
49
|
+
args = parser.parse_args(argv)
|
|
50
|
+
|
|
51
|
+
if args.command == "detect":
|
|
52
|
+
out = []
|
|
53
|
+
for path in args.files:
|
|
54
|
+
try:
|
|
55
|
+
t = detect_type(path)
|
|
56
|
+
except OSError as exc:
|
|
57
|
+
print("%s: %s" % (path, exc), file=sys.stderr)
|
|
58
|
+
return 2
|
|
59
|
+
out.append({"file": path, **t.to_dict()})
|
|
60
|
+
if not args.json:
|
|
61
|
+
print("%s: %s (%s)" % (path, t.mime, t.description))
|
|
62
|
+
if args.json:
|
|
63
|
+
print(json.dumps(out, indent=2))
|
|
64
|
+
return 0
|
|
65
|
+
|
|
66
|
+
guard = UploadGuard(
|
|
67
|
+
allowed=args.allow,
|
|
68
|
+
blocked=args.block if args.block is not None else ("category:executable", "category:script"),
|
|
69
|
+
max_size=args.max_size,
|
|
70
|
+
reject_threshold=Severity[args.threshold],
|
|
71
|
+
sanitize_svg=not args.no_sanitize,
|
|
72
|
+
strict_extension=not args.lenient_extension,
|
|
73
|
+
)
|
|
74
|
+
exit_code = 0
|
|
75
|
+
results = []
|
|
76
|
+
for path in args.files:
|
|
77
|
+
try:
|
|
78
|
+
result = guard.scan(path)
|
|
79
|
+
except OSError as exc:
|
|
80
|
+
print("%s: %s" % (path, exc), file=sys.stderr)
|
|
81
|
+
exit_code = 2
|
|
82
|
+
continue
|
|
83
|
+
results.append({"file": path, **result.to_dict()})
|
|
84
|
+
if not result.ok:
|
|
85
|
+
exit_code = max(exit_code, 1)
|
|
86
|
+
if not args.json:
|
|
87
|
+
_print_result(path, result)
|
|
88
|
+
if args.write_sanitized and result.sanitized is not None and len(args.files) == 1:
|
|
89
|
+
with open(args.write_sanitized, "wb") as fh:
|
|
90
|
+
fh.write(result.sanitized)
|
|
91
|
+
if args.json:
|
|
92
|
+
print(json.dumps(results, indent=2))
|
|
93
|
+
return exit_code
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
if __name__ == "__main__": # pragma: no cover
|
|
97
|
+
sys.exit(main())
|
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
"""Exception hierarchy."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
from typing import TYPE_CHECKING, List
|
|
5
|
+
|
|
6
|
+
if TYPE_CHECKING: # pragma: no cover
|
|
7
|
+
from .types import Finding, ScanResult
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class UploadGuardError(Exception):
|
|
11
|
+
"""Base class for all upload_guard errors."""
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class UnsupportedSource(UploadGuardError, TypeError):
|
|
15
|
+
"""The object passed in is not something upload_guard knows how to read."""
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class UploadRejected(UploadGuardError, ValueError):
|
|
19
|
+
"""Raised by :meth:`UploadGuard.check` when the upload fails the policy.
|
|
20
|
+
|
|
21
|
+
``result`` holds the full :class:`ScanResult`; ``findings`` is a shortcut
|
|
22
|
+
to the findings that caused the rejection.
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
def __init__(self, result: "ScanResult") -> None:
|
|
26
|
+
self.result = result
|
|
27
|
+
self.findings: List["Finding"] = result.errors
|
|
28
|
+
message = "; ".join(f"[{f.code}] {f.message}" for f in self.findings) or "upload rejected"
|
|
29
|
+
super().__init__(message)
|
|
30
|
+
|
|
31
|
+
@property
|
|
32
|
+
def codes(self) -> List[str]:
|
|
33
|
+
return [f.code for f in self.findings]
|
|
@@ -0,0 +1,179 @@
|
|
|
1
|
+
"""Extension handling: normalisation, aliases, declared-vs-detected comparison
|
|
2
|
+
and allow-list rules such as ``"image/*"``, ``"application/pdf"`` or ``".docx"``."""
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import os
|
|
6
|
+
import re
|
|
7
|
+
from typing import Dict, Iterable, List, Optional, Sequence, Tuple, Union
|
|
8
|
+
|
|
9
|
+
from .types import DetectedType
|
|
10
|
+
|
|
11
|
+
# Alternative spellings that should be treated as the same extension.
|
|
12
|
+
EXTENSION_ALIASES: Dict[str, str] = {
|
|
13
|
+
"jpeg": "jpg", "jpe": "jpg", "jfif": "jpg", "pjpeg": "jpg",
|
|
14
|
+
"tiff": "tif",
|
|
15
|
+
"htm": "html", "xhtml": "html",
|
|
16
|
+
"yml": "yaml",
|
|
17
|
+
"mpeg": "mpg", "mpe": "mpg",
|
|
18
|
+
"midi": "mid",
|
|
19
|
+
"markdown": "md",
|
|
20
|
+
"text": "txt",
|
|
21
|
+
"gzip": "gz",
|
|
22
|
+
"sqlite3": "sqlite",
|
|
23
|
+
"aiff": "aif",
|
|
24
|
+
"wave": "wav",
|
|
25
|
+
"heif": "heic",
|
|
26
|
+
"tar.gz": "tgz",
|
|
27
|
+
}
|
|
28
|
+
|
|
29
|
+
_EXT_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_\-+~]{0,15}$")
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def normalize_extension(ext: Optional[str]) -> str:
|
|
33
|
+
"""``".JPG "`` → ``"jpg"``. Returns ``""`` for empty/None."""
|
|
34
|
+
if not ext:
|
|
35
|
+
return ""
|
|
36
|
+
return ext.strip().lstrip(".").strip().lower()
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def canonical_extension(ext: Optional[str]) -> str:
|
|
40
|
+
norm = normalize_extension(ext)
|
|
41
|
+
return EXTENSION_ALIASES.get(norm, norm)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def split_filename(filename: Optional[str]) -> Tuple[str, Optional[str]]:
|
|
45
|
+
"""Return ``(stem, extension)`` using the *last* dot; extension is lower-case without the dot.
|
|
46
|
+
|
|
47
|
+
Trailing dots/spaces are ignored (``"a.txt. "`` → ``("a", "txt")``). A leading
|
|
48
|
+
dot alone (``".bashrc"``) is not an extension.
|
|
49
|
+
"""
|
|
50
|
+
if not filename:
|
|
51
|
+
return "", None
|
|
52
|
+
name = os.path.basename(filename.replace("\\", "/")).rstrip(". ")
|
|
53
|
+
if "." not in name:
|
|
54
|
+
return name, None
|
|
55
|
+
stem, _, ext = name.rpartition(".")
|
|
56
|
+
if not stem: # ".bashrc"
|
|
57
|
+
return name, None
|
|
58
|
+
ext = ext.lower()
|
|
59
|
+
if not _EXT_RE.match(ext):
|
|
60
|
+
return name, None
|
|
61
|
+
return stem, ext
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def all_extensions(filename: Optional[str]) -> List[str]:
|
|
65
|
+
"""Every dotted suffix: ``"a.tar.gz"`` → ``["tar", "gz"]``."""
|
|
66
|
+
if not filename:
|
|
67
|
+
return []
|
|
68
|
+
name = os.path.basename(filename.replace("\\", "/")).rstrip(". ").lstrip(".")
|
|
69
|
+
parts = name.split(".")
|
|
70
|
+
if len(parts) < 2:
|
|
71
|
+
return []
|
|
72
|
+
return [p.lower() for p in parts[1:] if _EXT_RE.match(p)]
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def extension_matches(ext: Optional[str], detected: DetectedType) -> bool:
|
|
76
|
+
"""True when ``ext`` is one of the extensions that ``detected`` may legitimately use."""
|
|
77
|
+
canon = canonical_extension(ext)
|
|
78
|
+
if not canon:
|
|
79
|
+
return False
|
|
80
|
+
return canon in {canonical_extension(e) for e in detected.extensions}
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
# ---------------------------------------------------------------------------
|
|
84
|
+
# Registry of every known type (for reverse look-ups and double-extension checks)
|
|
85
|
+
# ---------------------------------------------------------------------------
|
|
86
|
+
def _all_known_types() -> List[DetectedType]:
|
|
87
|
+
from .sniff import containers, signatures
|
|
88
|
+
|
|
89
|
+
seen: Dict[Tuple[str, Tuple[str, ...]], DetectedType] = {}
|
|
90
|
+
for module in (signatures, containers):
|
|
91
|
+
for value in vars(module).values():
|
|
92
|
+
if isinstance(value, DetectedType):
|
|
93
|
+
seen.setdefault((value.mime, value.extensions), value)
|
|
94
|
+
return list(seen.values())
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
_KNOWN_TYPES: Optional[List[DetectedType]] = None
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def known_types() -> List[DetectedType]:
|
|
101
|
+
global _KNOWN_TYPES
|
|
102
|
+
if _KNOWN_TYPES is None:
|
|
103
|
+
_KNOWN_TYPES = _all_known_types()
|
|
104
|
+
return _KNOWN_TYPES
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def known_extensions() -> frozenset:
|
|
108
|
+
return frozenset(canonical_extension(e) for t in known_types() for e in t.extensions)
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def types_for_extension(ext: Optional[str]) -> List[DetectedType]:
|
|
112
|
+
canon = canonical_extension(ext)
|
|
113
|
+
if not canon:
|
|
114
|
+
return []
|
|
115
|
+
return [t for t in known_types() if canon in {canonical_extension(e) for e in t.extensions}]
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def mime_for_extension(ext: Optional[str]) -> Optional[str]:
|
|
119
|
+
types = types_for_extension(ext)
|
|
120
|
+
return types[0].mime if types else None
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
# ---------------------------------------------------------------------------
|
|
124
|
+
# Allow-list rules
|
|
125
|
+
# ---------------------------------------------------------------------------
|
|
126
|
+
_MIME_RE = re.compile(r"^[a-z0-9!#$&^_.+-]+/(\*|[a-z0-9!#$&^_.+-]+)$", re.IGNORECASE)
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
class TypeRule:
|
|
130
|
+
"""One allow/deny rule: a MIME type, a MIME wildcard (``image/*``), an extension
|
|
131
|
+
(``.pdf`` / ``pdf``) or a category (``category:image``)."""
|
|
132
|
+
|
|
133
|
+
__slots__ = ("kind", "value", "raw")
|
|
134
|
+
|
|
135
|
+
def __init__(self, raw: str) -> None:
|
|
136
|
+
self.raw = raw
|
|
137
|
+
text = raw.strip().lower()
|
|
138
|
+
if not text:
|
|
139
|
+
raise ValueError("empty type rule")
|
|
140
|
+
if text.startswith("category:"):
|
|
141
|
+
self.kind, self.value = "category", text.split(":", 1)[1].strip()
|
|
142
|
+
elif "/" in text:
|
|
143
|
+
if not _MIME_RE.match(text):
|
|
144
|
+
raise ValueError("invalid MIME type rule: %r" % raw)
|
|
145
|
+
if text.endswith("/*"):
|
|
146
|
+
self.kind, self.value = "mime_prefix", text[:-1] # "image/"
|
|
147
|
+
else:
|
|
148
|
+
self.kind, self.value = "mime", text
|
|
149
|
+
else:
|
|
150
|
+
self.kind, self.value = "ext", canonical_extension(text)
|
|
151
|
+
if not self.value:
|
|
152
|
+
raise ValueError("invalid extension rule: %r" % raw)
|
|
153
|
+
|
|
154
|
+
def matches(self, detected: DetectedType) -> bool:
|
|
155
|
+
if self.kind == "mime":
|
|
156
|
+
return detected.mime.lower() == self.value
|
|
157
|
+
if self.kind == "mime_prefix":
|
|
158
|
+
return detected.mime.lower().startswith(self.value)
|
|
159
|
+
if self.kind == "category":
|
|
160
|
+
return detected.category == self.value
|
|
161
|
+
return any(canonical_extension(e) == self.value for e in detected.extensions)
|
|
162
|
+
|
|
163
|
+
def __repr__(self) -> str: # pragma: no cover - debugging aid
|
|
164
|
+
return "TypeRule(%r)" % self.raw
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
RuleInput = Union[str, TypeRule, Iterable[Union[str, TypeRule]], None]
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
def compile_rules(rules: RuleInput) -> List[TypeRule]:
|
|
171
|
+
if rules is None:
|
|
172
|
+
return []
|
|
173
|
+
if isinstance(rules, (str, TypeRule)):
|
|
174
|
+
rules = [rules]
|
|
175
|
+
return [r if isinstance(r, TypeRule) else TypeRule(r) for r in rules]
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
def any_rule_matches(rules: Sequence[TypeRule], detected: DetectedType) -> bool:
|
|
179
|
+
return any(r.matches(detected) for r in rules)
|