[Backport release-26.05] python3Packages.stanza: backport security fixes from 1.14.0 (#543283)

This commit is contained in:
Thomas Gerbet
2026-07-23 21:48:18 +00:00
committed by GitHub
4 changed files with 626 additions and 0 deletions

View 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

View File

@@ -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

View 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

View File

@@ -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 ];