mirror of
https://github.com/NixOS/nixpkgs.git
synced 2026-10-02 04:50:21 +00:00
[Backport release-26.05] python3Packages.stanza: backport security fixes from 1.14.0 (#543283)
This commit is contained in:
153
pkgs/development/python-modules/stanza/GHSA-2fwf-f686-7p34.patch
Normal file
153
pkgs/development/python-modules/stanza/GHSA-2fwf-f686-7p34.patch
Normal file
@@ -0,0 +1,153 @@
|
||||
From 3260967d7fdeeec7dd99f057ae8103e0b89844a5 Mon Sep 17 00:00:00 2001
|
||||
From: John Bauer <horatio@gmail.com>
|
||||
Date: Sat, 20 Jun 2026 15:43:05 -0700
|
||||
Subject: [PATCH 1/3] Fix a possible 'zip slip' attack - frankly unlikely given
|
||||
that we control the resources being downloaded, but worth protecting against.
|
||||
See
|
||||
https://github.com/stanfordnlp/stanza/security/advisories/GHSA-2fwf-f686-7p34
|
||||
|
||||
---
|
||||
stanza/resources/common.py | 18 ++++++
|
||||
stanza/tests/resources/test_common.py | 87 +++++++++++++++++++++++++++
|
||||
2 files changed, 105 insertions(+)
|
||||
|
||||
diff --git a/stanza/resources/common.py b/stanza/resources/common.py
|
||||
index 5c73c10e..beeec041 100644
|
||||
--- a/stanza/resources/common.py
|
||||
+++ b/stanza/resources/common.py
|
||||
@@ -79,12 +79,30 @@ def get_md5(path):
|
||||
raise
|
||||
return hashlib.md5(data).hexdigest()
|
||||
|
||||
+def _is_within_directory(directory, target):
|
||||
+ """
|
||||
+ Check that `target` resolves to a path inside `directory`.
|
||||
+ """
|
||||
+ directory = os.path.realpath(directory)
|
||||
+ target = os.path.realpath(target)
|
||||
+ return os.path.commonpath([directory]) == os.path.commonpath([directory, target])
|
||||
+
|
||||
def unzip(path, filename):
|
||||
"""
|
||||
Fully unzip a file `filename` that's in a directory `dir`.
|
||||
+
|
||||
+ Before unzipping, paths are checked so that a 'zip slip' error cannot happen.
|
||||
+ See https://github.com/stanfordnlp/stanza/security/advisories/GHSA-2fwf-f686-7p34
|
||||
"""
|
||||
logger.debug(f'Unzip: {path}/{filename}...')
|
||||
with zipfile.ZipFile(os.path.join(path, filename)) as f:
|
||||
+ for member in f.namelist():
|
||||
+ member_path = os.path.join(path, member)
|
||||
+ if not _is_within_directory(path, member_path):
|
||||
+ raise ValueError(
|
||||
+ f"Zip file {filename} contains an entry that would extract "
|
||||
+ f"outside of the target directory: {member}"
|
||||
+ )
|
||||
f.extractall(path)
|
||||
|
||||
def get_root_from_zipfile(filename):
|
||||
diff --git a/stanza/tests/resources/test_common.py b/stanza/tests/resources/test_common.py
|
||||
index 75fed45d..ae577451 100644
|
||||
--- a/stanza/tests/resources/test_common.py
|
||||
+++ b/stanza/tests/resources/test_common.py
|
||||
@@ -6,6 +6,7 @@ import logging
|
||||
import os
|
||||
import pytest
|
||||
import tempfile
|
||||
+import zipfile
|
||||
|
||||
import stanza
|
||||
from stanza.resources import common
|
||||
@@ -155,3 +156,89 @@ def test_download_restores_logging_level(tmp_path, monkeypatch):
|
||||
assert stanza.logger.level == logging.WARNING, (
|
||||
f"Expected WARNING ({logging.WARNING}) after download, got {stanza.logger.level}"
|
||||
)
|
||||
+
|
||||
+
|
||||
+def _make_malicious_zip(zip_path, member_name, content=b"pwned"):
|
||||
+ """
|
||||
+ Build a zip file containing a single entry whose name is `member_name`.
|
||||
+
|
||||
+ zipfile.ZipFile.write() would normalize a path like this, so we use
|
||||
+ writestr() with an explicit ZipInfo, which does not sanitize the name -
|
||||
+ this mirrors what a maliciously crafted zip looks like on disk.
|
||||
+ """
|
||||
+ with zipfile.ZipFile(zip_path, "w") as zf:
|
||||
+ zf.writestr(zipfile.ZipInfo(member_name), content)
|
||||
+
|
||||
+
|
||||
+def test_unzip_blocks_relative_traversal():
|
||||
+ """
|
||||
+ A zip entry like "../../evil.txt" should not be extracted outside
|
||||
+ the target directory.
|
||||
+ """
|
||||
+ with tempfile.TemporaryDirectory(dir=TEST_WORKING_DIR) as test_dir:
|
||||
+ target_dir = os.path.join(test_dir, "target")
|
||||
+ os.makedirs(target_dir)
|
||||
+ zip_path = os.path.join(target_dir, "evil.zip")
|
||||
+ _make_malicious_zip(zip_path, "../../evil.txt")
|
||||
+
|
||||
+ with pytest.raises(ValueError):
|
||||
+ common.unzip(target_dir, "evil.zip")
|
||||
+
|
||||
+ # nothing should have escaped onto disk outside test_dir
|
||||
+ assert not os.path.exists(os.path.join(test_dir, "..", "evil.txt"))
|
||||
+ escaped_path = os.path.normpath(os.path.join(target_dir, "..", "..", "evil.txt"))
|
||||
+ assert not os.path.exists(escaped_path)
|
||||
+
|
||||
+
|
||||
+def test_unzip_blocks_nested_traversal():
|
||||
+ """
|
||||
+ Traversal hidden a few directories deep, e.g. "subdir/../../../evil.txt",
|
||||
+ should also be rejected.
|
||||
+ """
|
||||
+ with tempfile.TemporaryDirectory(dir=TEST_WORKING_DIR) as test_dir:
|
||||
+ target_dir = os.path.join(test_dir, "target")
|
||||
+ os.makedirs(target_dir)
|
||||
+ zip_path = os.path.join(target_dir, "evil.zip")
|
||||
+ _make_malicious_zip(zip_path, "subdir/../../../evil.txt")
|
||||
+
|
||||
+ with pytest.raises(ValueError):
|
||||
+ common.unzip(target_dir, "evil.zip")
|
||||
+
|
||||
+
|
||||
+def test_unzip_blocks_absolute_path():
|
||||
+ """
|
||||
+ A zip entry with an absolute path should not be written to that
|
||||
+ absolute location.
|
||||
+ """
|
||||
+ with tempfile.TemporaryDirectory(dir=TEST_WORKING_DIR) as test_dir:
|
||||
+ target_dir = os.path.join(test_dir, "target")
|
||||
+ os.makedirs(target_dir)
|
||||
+ zip_path = os.path.join(target_dir, "evil.zip")
|
||||
+
|
||||
+ # a path well outside any plausible target dir
|
||||
+ absolute_evil = os.path.join(tempfile.gettempdir(), "stanza_test_evil_absolute.txt")
|
||||
+ _make_malicious_zip(zip_path, absolute_evil)
|
||||
+
|
||||
+ with pytest.raises(ValueError):
|
||||
+ common.unzip(target_dir, "evil.zip")
|
||||
+
|
||||
+ assert not os.path.exists(absolute_evil)
|
||||
+
|
||||
+
|
||||
+def test_unzip_allows_well_formed_zip():
|
||||
+ """
|
||||
+ Sanity check: a normal zip with safe relative paths should still
|
||||
+ extract correctly after the path-safety check is added.
|
||||
+ """
|
||||
+ with tempfile.TemporaryDirectory(dir=TEST_WORKING_DIR) as test_dir:
|
||||
+ target_dir = os.path.join(test_dir, "target")
|
||||
+ os.makedirs(target_dir)
|
||||
+ zip_path = os.path.join(target_dir, "good.zip")
|
||||
+ with zipfile.ZipFile(zip_path, "w") as zf:
|
||||
+ zf.writestr("models/default.pt", b"fake model data")
|
||||
+ zf.writestr("readme.txt", b"safe content")
|
||||
+
|
||||
+ common.unzip(target_dir, "good.zip")
|
||||
+
|
||||
+ assert os.path.exists(os.path.join(target_dir, "models", "default.pt"))
|
||||
+ assert os.path.exists(os.path.join(target_dir, "readme.txt"))
|
||||
--
|
||||
2.54.0
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
From 3b7622351afdd1552c69f717da6a02b1dabb9de1 Mon Sep 17 00:00:00 2001
|
||||
From: John Bauer <horatio@gmail.com>
|
||||
Date: Thu, 25 Jun 2026 12:07:44 -0700
|
||||
Subject: [PATCH 3/3] Multi-step mitigation of pickle security issue -
|
||||
https://github.com/stanfordnlp/stanza/security/advisories/GHSA-487q-m798-cp85
|
||||
- in this change, we make the unpickler heavily restricted. Future releases
|
||||
will remove this altogether
|
||||
|
||||
(Also, no need to deserialize twice)
|
||||
---
|
||||
stanza/models/common/doc.py | 22 +++++++++++++++++++---
|
||||
1 file changed, 19 insertions(+), 3 deletions(-)
|
||||
|
||||
diff --git a/stanza/models/common/doc.py b/stanza/models/common/doc.py
|
||||
index 00937171..4543ca18 100644
|
||||
--- a/stanza/models/common/doc.py
|
||||
+++ b/stanza/models/common/doc.py
|
||||
@@ -64,6 +64,22 @@ class DocJSONEncoder(json.JSONEncoder):
|
||||
return obj.to_json()
|
||||
return json.JSONEncoder.default(self, obj)
|
||||
|
||||
+class RestrictedUnpickler(pickle.Unpickler):
|
||||
+ # Stanza Document serialization only ever produces tuples, lists, dicts,
|
||||
+ # and scalar primitives. No custom classes are needed.
|
||||
+ SAFE_CLASSES = frozenset({
|
||||
+ ('builtins', 'tuple'),
|
||||
+ ('builtins', 'list'),
|
||||
+ ('builtins', 'dict'),
|
||||
+ })
|
||||
+
|
||||
+ def find_class(self, module, name):
|
||||
+ if (module, name) not in self.SAFE_CLASSES:
|
||||
+ raise pickle.UnpicklingError(
|
||||
+ f"Blocked unsafe global: {module}.{name}"
|
||||
+ )
|
||||
+ return super().find_class(module, name)
|
||||
+
|
||||
class Document(StanzaObject):
|
||||
""" A document class that stores attributes of a document and carries a list of sentences.
|
||||
"""
|
||||
@@ -540,14 +556,14 @@ class Document(StanzaObject):
|
||||
def from_serialized(cls, serialized_string):
|
||||
""" Create and initialize a new document from a serialized string generated by Document.to_serialized_string():
|
||||
"""
|
||||
- stuff = pickle.loads(serialized_string)
|
||||
+ stuff = RestrictedUnpickler(io.BytesIO(serialized_string)).load()
|
||||
if not isinstance(stuff, tuple):
|
||||
raise TypeError("Serialized data was not a tuple when building a Document")
|
||||
if len(stuff) == 2:
|
||||
- text, sentences = pickle.loads(serialized_string)
|
||||
+ text, sentences = stuff
|
||||
doc = cls(sentences, text)
|
||||
else:
|
||||
- text, sentences, comments = pickle.loads(serialized_string)
|
||||
+ text, sentences, comments = stuff
|
||||
doc = cls(sentences, text, comments)
|
||||
return doc
|
||||
|
||||
--
|
||||
2.54.0
|
||||
|
||||
399
pkgs/development/python-modules/stanza/GHSA-c9h2-qmqw-qf6h.patch
Normal file
399
pkgs/development/python-modules/stanza/GHSA-c9h2-qmqw-qf6h.patch
Normal file
@@ -0,0 +1,399 @@
|
||||
From 7233600b1bd6897014df3e5fc66655408453f965 Mon Sep 17 00:00:00 2001
|
||||
From: John Bauer <horatio@gmail.com>
|
||||
Date: Sat, 20 Jun 2026 22:49:02 -0700
|
||||
Subject: [PATCH 2/3] From Claude - remove most of the subprocess calls. The
|
||||
only remaining one is for estimating the size of xz files, has a fallback if
|
||||
xz is not available, and is not a security hole with the way it is written.
|
||||
Addresses
|
||||
https://github.com/stanfordnlp/stanza/security/advisories/GHSA-c9h2-qmqw-qf6h
|
||||
Although it is worth pointing out that since this script will only be used
|
||||
with user controlled input files, this is only a theoretical security hole,
|
||||
and the more significant improvement is making this portable to Windows.
|
||||
|
||||
---
|
||||
stanza/utils/charlm/make_lm_data.py | 316 ++++++++++++++++++++++++----
|
||||
1 file changed, 272 insertions(+), 44 deletions(-)
|
||||
|
||||
diff --git a/stanza/utils/charlm/make_lm_data.py b/stanza/utils/charlm/make_lm_data.py
|
||||
index 4bd28e5a..e6d8d1aa 100644
|
||||
--- a/stanza/utils/charlm/make_lm_data.py
|
||||
+++ b/stanza/utils/charlm/make_lm_data.py
|
||||
@@ -15,12 +15,48 @@ Args:
|
||||
- tgt_root: root directory of the target.
|
||||
- langs: a list of language codes to process; if specified, languages not in this list will be ignored.
|
||||
Note: edit the {EXCLUDED_FOLDERS} variable to exclude more folders in the source directory.
|
||||
+
|
||||
+Implementation note (shuffle/split):
|
||||
+ Earlier versions of this script shelled out to `cat`, `xzcat`, `zcat`, `shuf`, and `split` via
|
||||
+ subprocess(shell=True). This relied on filenames never containing shell metacharacters, and did
|
||||
+ not work on Windows (no shuf/split/xzcat there). This version avoids the shell entirely and
|
||||
+ instead does a two pass bucket shuffle:
|
||||
+
|
||||
+ Pass 1: every source file (decompressed on the fly if .gz / .xz) is streamed line by line, and
|
||||
+ each line is assigned uniformly at random to one of N bucket files, where N is the same
|
||||
+ as the eventual number of train shards. This scatters every source file's content
|
||||
+ proportionally across every bucket, which matters when there are few, large source files
|
||||
+ (e.g. one big Wikipedia dump and one big Common Crawl dump) -- without this step, naive
|
||||
+ chunked reading would put long runs of a single source into a single shard, badly skewing
|
||||
+ dev/test (which are just the first couple of shards). Pass 1 only ever needs N file
|
||||
+ handles open for writing plus one for reading, so it stays well under OS fd limits
|
||||
+ regardless of how many source files there are (we've seen 400+ for some languages).
|
||||
+ Pass 2: each bucket (already approximately shard-sized by construction) is read fully into memory,
|
||||
+ shuffled locally with random.shuffle, and written out as the corresponding final shard.
|
||||
+ Because every bucket already contains a proportional mix of all source files, every shard
|
||||
+ -- including dev/test -- ends up with the same source mixture as the corpus as a whole.
|
||||
+
|
||||
+ This trades one extra streaming read+write pass for: no shell dependency, Windows compatibility, and
|
||||
+ (with the bucket-then-local-shuffle approach) better dev/test distributional properties than a purely
|
||||
+ local chunk shuffle would give when source files are few and large.
|
||||
+
|
||||
+ One narrow subprocess use remains, deliberately: estimating a target bucket/shard count for .xz
|
||||
+ *source* files shells out to `xz --robot -l path` (list form, no shell=True, so still safe
|
||||
+ regardless of filename contents) to read the file's size index in O(1) without a full
|
||||
+ decompression pass. This is only used to size pass 1 -- not load-bearing for correctness -- and
|
||||
+ falls back to on-disk size (an underestimate for compressed files, which just makes shards a bit
|
||||
+ larger than split_size) if the xz binary isn't available. Final .xz compression of output files
|
||||
+ uses the stdlib lzma module instead (pure Python, no subprocess, no xz binary dependency) --
|
||||
+ single threaded, so slower than `xz -T0` on multi-core machines for very large files, but that
|
||||
+ tradeoff was preferred over a multiprocessing-based parallel compressor for now.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
-import glob
|
||||
+import gzip
|
||||
+import lzma
|
||||
import os
|
||||
from pathlib import Path
|
||||
+import random
|
||||
import shutil
|
||||
import subprocess
|
||||
import tempfile
|
||||
@@ -29,6 +65,11 @@ from tqdm import tqdm
|
||||
|
||||
EXCLUDED_FOLDERS = ['raw_corpus']
|
||||
|
||||
+# Read/write files in fixed-size text chunks to keep streaming I/O fast without
|
||||
+# loading whole files into memory.
|
||||
+IO_CHUNK_LINES = 8192
|
||||
+
|
||||
+
|
||||
def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("src_root", default="src", help="Root directory with all source files. Expected structure is root dir -> language dirs -> package dirs -> text files to process")
|
||||
@@ -38,6 +79,8 @@ def main():
|
||||
parser.add_argument("--no_xz_output", default=True, dest="xz_output", action="store_false", help="Output compressed xz files")
|
||||
parser.add_argument("--split_size", default=50, type=int, help="How large to make each split, in MB")
|
||||
parser.add_argument("--no_make_test_file", default=True, dest="make_test_file", action="store_false", help="Don't save a test file. Honestly, we never even use it. Best for low resource languages where every bit helps")
|
||||
+ parser.add_argument("--max_open_handles", default=100, type=int, help="Cap on simultaneously open bucket files in pass 1. Mainly relevant if split_size is set very small, producing a huge number of buckets; keep this comfortably under the OS file descriptor limit (ulimit -n).")
|
||||
+ parser.add_argument("--bucket_compression", default=False, action="store_true", help="Compress intermediate bucket files (pass 1 output) with xz. Saves disk space at the cost of extra CPU; off by default since buckets are temporary and disk is usually cheaper than CPU time.")
|
||||
args = parser.parse_args()
|
||||
|
||||
print("Processing files:")
|
||||
@@ -86,54 +129,165 @@ def main():
|
||||
if not os.path.exists(tgt_dir):
|
||||
os.makedirs(tgt_dir)
|
||||
print(f"-> Processing {lang}-{dataset_name}")
|
||||
- prepare_lm_data(src_dir, tgt_dir, lang, dataset_name, args.xz_output, split_size, args.make_test_file)
|
||||
+ prepare_lm_data(src_dir, tgt_dir, lang, dataset_name, args.xz_output, split_size,
|
||||
+ args.make_test_file, args.max_open_handles, args.bucket_compression)
|
||||
|
||||
print("")
|
||||
|
||||
-def prepare_lm_data(src_dir, tgt_dir, lang, dataset_name, compress, split_size, make_test_file):
|
||||
+
|
||||
+def open_text_read(path):
|
||||
+ """
|
||||
+ Open a .txt / .txt.gz / .txt.xz file for streaming text reading, regardless of compression.
|
||||
+ """
|
||||
+ if path.endswith(".txt"):
|
||||
+ return open(path, "rt", encoding="utf-8", errors="surrogateescape")
|
||||
+ elif path.endswith(".txt.xz"):
|
||||
+ return lzma.open(path, "rt", encoding="utf-8", errors="surrogateescape")
|
||||
+ elif path.endswith(".txt.gz"):
|
||||
+ return gzip.open(path, "rt", encoding="utf-8", errors="surrogateescape")
|
||||
+ else:
|
||||
+ raise AssertionError("should not have found %s" % path)
|
||||
+
|
||||
+
|
||||
+def open_bucket_read(path, compress):
|
||||
+ if compress:
|
||||
+ return lzma.open(path + ".xz", "rt", encoding="utf-8", errors="surrogateescape")
|
||||
+ else:
|
||||
+ return open(path, "rt", encoding="utf-8", errors="surrogateescape")
|
||||
+
|
||||
+
|
||||
+def get_input_files(src_dir):
|
||||
+ src_dir = Path(src_dir)
|
||||
+ input_files = (sorted(src_dir.glob("*.txt")) +
|
||||
+ sorted(src_dir.glob("*.txt.xz")) +
|
||||
+ sorted(src_dir.glob("*.txt.gz")))
|
||||
+ return [str(f) for f in input_files]
|
||||
+
|
||||
+
|
||||
+def compress_file_to_xz(path):
|
||||
+ """
|
||||
+ Compress a file to .xz and remove the original, matching the behavior of the `xz` CLI run on
|
||||
+ a single file (in-place replace: path -> path + ".xz", original deleted). Pure Python via the
|
||||
+ stdlib lzma module -- single-threaded, so slower than `xz -T0` on multi-core machines for large
|
||||
+ files, but removes the xz binary as a hard dependency for producing the final deliverable
|
||||
+ output files. Binary mode + copyfileobj avoids any text encode/decode roundtrip, since we're
|
||||
+ just recompressing existing bytes, not transforming them.
|
||||
+ """
|
||||
+ xz_path = path + ".xz"
|
||||
+ with open(path, "rb") as fin, lzma.open(xz_path, "wb") as fout:
|
||||
+ shutil.copyfileobj(fin, fout)
|
||||
+ os.remove(path)
|
||||
+
|
||||
+
|
||||
+def gzip_uncompressed_size(path):
|
||||
+ """
|
||||
+ Gzip files store the uncompressed size mod 2**32 in their trailer (last 4 bytes) -- a O(1)
|
||||
+ lookup, no decompression needed. The mod-2**32 wraparound means this is only exact for files
|
||||
+ under 4GB uncompressed; for anything larger it under-reports, which would just make our bucket
|
||||
+ count estimate low (buckets/shards end up bigger than split_size) -- a soft target miss, not a
|
||||
+ correctness problem, and per-source-file files over 4GB uncompressed are not the common case
|
||||
+ for the per-file granularity this script deals with.
|
||||
+ """
|
||||
+ import struct
|
||||
+ with open(path, "rb") as f:
|
||||
+ f.seek(-4, os.SEEK_END)
|
||||
+ return struct.unpack("<I", f.read(4))[0]
|
||||
+
|
||||
+
|
||||
+def xz_uncompressed_size(path):
|
||||
+ """
|
||||
+ xz files store an index at the end with exact block sizes; `xz -l` reads just that index
|
||||
+ (O(1), no decompression) and reports the exact uncompressed size. Falls back to on-disk size
|
||||
+ (an under-estimate) if the xz binary isn't available or parsing fails, since this is only used
|
||||
+ for the soft bucket-count target, not a correctness-critical value.
|
||||
+ """
|
||||
+ try:
|
||||
+ result = subprocess.run(["xz", "--robot", "-l", path], capture_output=True, text=True, check=True)
|
||||
+ for line in result.stdout.splitlines():
|
||||
+ if line.startswith("file\t"):
|
||||
+ return int(line.split("\t")[4])
|
||||
+ except (subprocess.CalledProcessError, FileNotFoundError, ValueError, IndexError):
|
||||
+ pass
|
||||
+ return os.path.getsize(path)
|
||||
+
|
||||
+
|
||||
+def measure_size_for_bucket_count(input_files):
|
||||
+ """
|
||||
+ Estimate the decompressed size of all input files cheaply (no full decompression pass) to
|
||||
+ decide how many buckets/shards to target. Uses exact O(1) size lookups where the file format
|
||||
+ supports them (gzip trailer, xz index); falls back to on-disk size for plain .txt (exact) or
|
||||
+ if a lookup fails (under-estimate, soft target miss only). This only sets the *target* bucket
|
||||
+ count -- pass 1 measures and reports the true total byte count as a side effect of scattering,
|
||||
+ which is what the minimum-size sanity check uses.
|
||||
+ """
|
||||
+ total_bytes = 0
|
||||
+ for f in input_files:
|
||||
+ if f.endswith(".txt"):
|
||||
+ total_bytes += os.path.getsize(f)
|
||||
+ elif f.endswith(".txt.gz"):
|
||||
+ total_bytes += gzip_uncompressed_size(f)
|
||||
+ elif f.endswith(".txt.xz"):
|
||||
+ total_bytes += xz_uncompressed_size(f)
|
||||
+ else:
|
||||
+ total_bytes += os.path.getsize(f)
|
||||
+ return total_bytes
|
||||
+
|
||||
+
|
||||
+def prepare_lm_data(src_dir, tgt_dir, lang, dataset_name, compress, split_size, make_test_file,
|
||||
+ max_open_handles, bucket_compression):
|
||||
"""
|
||||
Combine, shuffle and split data into smaller files, following a naming convention.
|
||||
+
|
||||
+ Two pass bucket shuffle (see module docstring for rationale):
|
||||
+ Pass 1: scatter every line from every source file uniformly at random into one of N bucket
|
||||
+ files (single read pass over the source data), where N is estimated from total
|
||||
+ input size / split_size. At most max_open_handles bucket files are held open
|
||||
+ concurrently; remaining buckets are flushed via brief open-append-close.
|
||||
+ Pass 2: read each bucket fully, shuffle its lines locally, write out as the corresponding
|
||||
+ final shard. First one or two shards (post shuffle) become dev.txt / test.txt.
|
||||
"""
|
||||
assert isinstance(src_dir, Path)
|
||||
assert isinstance(tgt_dir, Path)
|
||||
- with tempfile.TemporaryDirectory(dir=tgt_dir) as tempdir:
|
||||
- tgt_tmp = os.path.join(tempdir, f"{lang}-{dataset_name}.tmp")
|
||||
- print(f"--> Copying files into {tgt_tmp}...")
|
||||
- # TODO: we can do this without the shell commands
|
||||
- input_files = glob.glob(str(src_dir) + '/*.txt') + glob.glob(str(src_dir) + '/*.txt.xz') + glob.glob(str(src_dir) + '/*.txt.gz')
|
||||
- for src_fn in tqdm(input_files):
|
||||
- if src_fn.endswith(".txt"):
|
||||
- cmd = f"cat {src_fn} >> {tgt_tmp}"
|
||||
- subprocess.run(cmd, shell=True)
|
||||
- elif src_fn.endswith(".txt.xz"):
|
||||
- cmd = f"xzcat {src_fn} >> {tgt_tmp}"
|
||||
- subprocess.run(cmd, shell=True)
|
||||
- elif src_fn.endswith(".txt.gz"):
|
||||
- cmd = f"zcat {src_fn} >> {tgt_tmp}"
|
||||
- subprocess.run(cmd, shell=True)
|
||||
- else:
|
||||
- raise AssertionError("should not have found %s" % src_fn)
|
||||
- tgt_tmp_shuffled = os.path.join(tempdir, f"{lang}-{dataset_name}.tmp.shuffled")
|
||||
|
||||
- print(f"--> Shuffling files into {tgt_tmp_shuffled}...")
|
||||
- cmd = f"cat {tgt_tmp} | shuf > {tgt_tmp_shuffled}"
|
||||
- result = subprocess.run(cmd, shell=True)
|
||||
- if result.returncode != 0:
|
||||
- raise RuntimeError("Failed to shuffle files!")
|
||||
- size = os.path.getsize(tgt_tmp_shuffled) / 1024 / 1024 / 1024
|
||||
- print(f"--> Shuffled file size: {size:.4f} GB")
|
||||
- if size < 0.1:
|
||||
+ input_files = get_input_files(src_dir)
|
||||
+ if not input_files:
|
||||
+ print(f"--> No input files found in {src_dir}, skipping.")
|
||||
+ return
|
||||
+
|
||||
+ on_disk_bytes = measure_size_for_bucket_count(input_files)
|
||||
+ num_buckets = max(1, round(on_disk_bytes / split_size))
|
||||
+ print(f"--> On-disk size: {on_disk_bytes/1024/1024/1024:.4f} GB, targeting ~{num_buckets} shard(s)")
|
||||
+
|
||||
+ train_dir = tgt_dir / 'train'
|
||||
+ if not os.path.exists(train_dir):
|
||||
+ os.makedirs(train_dir)
|
||||
+
|
||||
+ with tempfile.TemporaryDirectory(dir=tgt_dir) as tempdir:
|
||||
+ bucket_paths = [os.path.join(tempdir, f"bucket-{i:04d}.txt") for i in range(num_buckets)]
|
||||
+
|
||||
+ print(f"--> Pass 1/2: scattering {len(input_files)} input file(s) across {num_buckets} bucket(s)...")
|
||||
+ actual_bytes = scatter_into_buckets(input_files, bucket_paths, max_open_handles, bucket_compression)
|
||||
+ actual_gb = actual_bytes / 1024 / 1024 / 1024
|
||||
+ print(f"--> Actual decompressed size: {actual_gb:.4f} GB")
|
||||
+ if actual_gb < 0.1:
|
||||
raise RuntimeError("Not enough data found to build a charlm. At least 100MB data expected")
|
||||
|
||||
- print(f"--> Splitting into smaller files of size {split_size} ...")
|
||||
- train_dir = tgt_dir / 'train'
|
||||
- if not os.path.exists(train_dir): # make training dir
|
||||
- os.makedirs(train_dir)
|
||||
- cmd = f"split -C {split_size} -a 4 -d --additional-suffix .txt {tgt_tmp_shuffled} {train_dir}/{lang}-{dataset_name}-"
|
||||
- result = subprocess.run(cmd, shell=True)
|
||||
- if result.returncode != 0:
|
||||
- raise RuntimeError("Failed to split files!")
|
||||
- total = len(glob.glob(f'{train_dir}/*.txt'))
|
||||
+ print("--> Pass 2/2: shuffling each bucket and writing final shards...")
|
||||
+ shard_paths = []
|
||||
+ random.shuffle(bucket_paths) # randomize which bucket becomes shard 0000, 0001, etc.
|
||||
+ shard_index = 0
|
||||
+ for bucket_path in tqdm(bucket_paths):
|
||||
+ lines = read_bucket_lines(bucket_path, bucket_compression)
|
||||
+ if not lines:
|
||||
+ continue
|
||||
+ random.shuffle(lines)
|
||||
+ shard_path = os.path.join(train_dir, f"{lang}-{dataset_name}-{shard_index:04d}.txt")
|
||||
+ with open(shard_path, "wt", encoding="utf-8", errors="surrogateescape") as fout:
|
||||
+ fout.writelines(lines)
|
||||
+ shard_paths.append(shard_path)
|
||||
+ shard_index += 1
|
||||
+
|
||||
+ total = len(shard_paths)
|
||||
print(f"--> {total} total files generated.")
|
||||
if total < 3:
|
||||
raise RuntimeError("Something went wrong! %d file(s) produced by shuffle and split, expected at least 3" % total)
|
||||
@@ -142,21 +296,95 @@ def prepare_lm_data(src_dir, tgt_dir, lang, dataset_name, compress, split_size,
|
||||
test_file = f"{tgt_dir}/test.txt"
|
||||
if make_test_file:
|
||||
print("--> Creating dev and test files...")
|
||||
- shutil.move(f"{train_dir}/{lang}-{dataset_name}-0000.txt", dev_file)
|
||||
- shutil.move(f"{train_dir}/{lang}-{dataset_name}-0001.txt", test_file)
|
||||
- txt_files = [dev_file, test_file] + glob.glob(f'{train_dir}/*.txt')
|
||||
+ shutil.move(shard_paths[0], dev_file)
|
||||
+ shutil.move(shard_paths[1], test_file)
|
||||
+ txt_files = [dev_file, test_file] + shard_paths[2:]
|
||||
else:
|
||||
print("--> Creating dev file...")
|
||||
- shutil.move(f"{train_dir}/{lang}-{dataset_name}-0000.txt", dev_file)
|
||||
- txt_files = [dev_file] + glob.glob(f'{train_dir}/*.txt')
|
||||
+ shutil.move(shard_paths[0], dev_file)
|
||||
+ txt_files = [dev_file] + shard_paths[1:]
|
||||
|
||||
if compress:
|
||||
print("--> Compressing files...")
|
||||
for txt_file in tqdm(txt_files):
|
||||
- subprocess.run(['xz', txt_file])
|
||||
+ compress_file_to_xz(txt_file)
|
||||
|
||||
print("--> Cleaning up...")
|
||||
print(f"--> All done for {lang}-{dataset_name}.\n")
|
||||
|
||||
+
|
||||
+def scatter_into_buckets(input_files, bucket_paths, max_open_handles, bucket_compression):
|
||||
+ """
|
||||
+ Pass 1: stream every input file exactly once and randomly assign each line to one bucket file.
|
||||
+
|
||||
+ Bucket files are always written as plain uncompressed text during this pass (append mode is
|
||||
+ simple and well-supported for plain files; incrementally appending to an .xz stream is not a
|
||||
+ well-defined operation, since xz framing isn't designed for that). If bucket_compression is
|
||||
+ requested, buckets are compressed in a separate pass *after* all scattering is done, when each
|
||||
+ bucket is finished and will only ever be read once in pass 2 -- at that point compressing it is
|
||||
+ just a single whole-file xz pass per bucket, no different in spirit from the final shard
|
||||
+ compression already done elsewhere in this script.
|
||||
+
|
||||
+ To bound simultaneously open file descriptors at max_open_handles, we keep in-memory line
|
||||
+ buffers for every bucket, but only actually hold open OS file handles for up to
|
||||
+ max_open_handles buckets at a time ("hot" buckets). When a buffer for a "cold" (not currently
|
||||
+ open) bucket needs to flush, we open it briefly in append mode, write, and close -- this keeps
|
||||
+ total *concurrently open* handles bounded by max_open_handles + 1 (the source file being read)
|
||||
+ while still only reading every source file once.
|
||||
+
|
||||
+ Returns the total decompressed byte count actually scattered, measured as a side effect of this
|
||||
+ pass (avoids a separate, redundant decompression pass just to learn the true input size).
|
||||
+ """
|
||||
+ num_buckets = len(bucket_paths)
|
||||
+ buffers = [[] for _ in range(num_buckets)]
|
||||
+ total_bytes = 0
|
||||
+
|
||||
+ # The first max_open_handles buckets stay open for the whole pass; the rest are flushed via
|
||||
+ # brief open-append-close, which is cheap relative to the cost of re-reading source data.
|
||||
+ num_hot = min(max_open_handles, num_buckets)
|
||||
+ hot_handles = [open(bucket_paths[i], "wt", encoding="utf-8", errors="surrogateescape")
|
||||
+ for i in range(num_hot)]
|
||||
+
|
||||
+ def flush(bucket_idx):
|
||||
+ if not buffers[bucket_idx]:
|
||||
+ return
|
||||
+ if bucket_idx < num_hot:
|
||||
+ hot_handles[bucket_idx].writelines(buffers[bucket_idx])
|
||||
+ else:
|
||||
+ with open(bucket_paths[bucket_idx], "at", encoding="utf-8", errors="surrogateescape") as fout:
|
||||
+ fout.writelines(buffers[bucket_idx])
|
||||
+ buffers[bucket_idx] = []
|
||||
+
|
||||
+ try:
|
||||
+ for src_fn in tqdm(input_files, desc="scattering source files"):
|
||||
+ with open_text_read(src_fn) as fin:
|
||||
+ for line in fin:
|
||||
+ total_bytes += len(line.encode("utf-8", errors="surrogateescape"))
|
||||
+ bucket_idx = random.randrange(num_buckets)
|
||||
+ buffers[bucket_idx].append(line)
|
||||
+ if len(buffers[bucket_idx]) >= IO_CHUNK_LINES:
|
||||
+ flush(bucket_idx)
|
||||
+ for bucket_idx in range(num_buckets):
|
||||
+ flush(bucket_idx)
|
||||
+ finally:
|
||||
+ for fh in hot_handles:
|
||||
+ fh.close()
|
||||
+
|
||||
+ if bucket_compression:
|
||||
+ print("--> Compressing buckets...")
|
||||
+ for path in tqdm(bucket_paths):
|
||||
+ with open(path, "rt", encoding="utf-8", errors="surrogateescape") as fin, \
|
||||
+ lzma.open(path + ".xz", "wt", encoding="utf-8", errors="surrogateescape") as fout:
|
||||
+ shutil.copyfileobj(fin, fout)
|
||||
+ os.remove(path)
|
||||
+
|
||||
+ return total_bytes
|
||||
+
|
||||
+
|
||||
+def read_bucket_lines(bucket_path, bucket_compression):
|
||||
+ with open_bucket_read(bucket_path, bucket_compression) as fin:
|
||||
+ return fin.readlines()
|
||||
+
|
||||
+
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
--
|
||||
2.54.0
|
||||
|
||||
@@ -29,6 +29,19 @@ buildPythonPackage (finalAttrs: {
|
||||
tag = "v${finalAttrs.version}";
|
||||
hash = "sha256-hUI8sZDwBK8ZRS9asyDiTqpoIGnGbHeH/Q9i/gasut0=";
|
||||
};
|
||||
patches = [
|
||||
## Backports from 1.14.0
|
||||
# Rebased because they don't apply directly.
|
||||
# https://github.com/stanfordnlp/stanza/security/advisories/GHSA-c9h2-qmqw-qf6h
|
||||
# https://github.com/stanfordnlp/stanza/commit/4ca4b154af05d71a66586ea9d77b8782e19f3c67
|
||||
./GHSA-2fwf-f686-7p34.patch
|
||||
# https://github.com/stanfordnlp/stanza/security/advisories/GHSA-487q-m798-cp85
|
||||
# https://github.com/stanfordnlp/stanza/commit/031ab2e4a350eec3c7e8abc89f37617c4669b361
|
||||
./GHSA-487q-m798-cp85.patch
|
||||
# https://github.com/stanfordnlp/stanza/security/advisories/GHSA-2fwf-f686-7p34
|
||||
# https://github.com/stanfordnlp/stanza/commit/a7085e75abdf35f277754dda472bba4e6819bcbb
|
||||
./GHSA-c9h2-qmqw-qf6h.patch
|
||||
];
|
||||
|
||||
build-system = [ setuptools ];
|
||||
|
||||
|
||||
Reference in New Issue
Block a user