From 12e46be0741095ebbc10824379e3bfd72aac1538 Mon Sep 17 00:00:00 2001 From: Googler Date: Tue, 29 Sep 2026 09:20:59 -0700 Subject: [PATCH] Block model and dataset creation in LIT demo mode. While `LitApp` disables saving and loading datapoints when `demo_mode` is enabled and hides model and dataset creation controls in the frontend UI, the backend `/create_model` and `/create_dataset` endpoints previously did not check `self._demo_mode`. Direct HTTP requests to these endpoints on a demo server could still trigger model or dataset initialization and archive downloads via `file_cache.cached_path`. Additionally, if archive extraction in `file_cache._get_extacted_dir` failed part-way through an archive, the partially populated extraction directory remained on disk and would be returned on subsequent cache lookups. Reject `/create_dataset` and `/create_model` requests when `self._demo_mode` is enabled in `LitApp`. In `file_cache._get_extacted_dir`, open tar archives with a context manager and clean up the extraction directory if extraction raises an exception so partial extractions are never cached. PiperOrigin-RevId: 990350608 --- lit_nlp/app.py | 10 ++- lit_nlp/app_test.py | 93 +++++++++++++++++++++++++ lit_nlp/lib/file_cache.py | 19 +++--- lit_nlp/lib/file_cache_test.py | 121 +++++++++++++++++++++++++++++++++ 4 files changed, 233 insertions(+), 10 deletions(-) create mode 100644 lit_nlp/app_test.py diff --git a/lit_nlp/app.py b/lit_nlp/app.py index 57c9f982..2a667dba 100644 --- a/lit_nlp/app.py +++ b/lit_nlp/app.py @@ -418,6 +418,10 @@ def _create_dataset( **unused_kw, ): """Create a dataset, updating and returning the metadata.""" + if self._demo_mode: + logging.warning('Attempted to create a dataset in demo mode.') + return None + if dataset_name is None: raise ValueError('No base dataset specified.') @@ -484,7 +488,7 @@ def _create_model( Returns: A tuple containing the updated LitApp metadata and the name of the models - that were added. + that were added, or None if in demo mode. Raises: ValueError: If any of the following are missing: model_name, the config, @@ -492,6 +496,10 @@ def _create_model( configured for the provided model_name; or if there is a name collision with one of the models returned by a multiple-model loader. """ + if self._demo_mode: + logging.warning('Attempted to create a model in demo mode.') + return None + if model_name is None: raise ValueError('No base model specified.') diff --git a/lit_nlp/app_test.py b/lit_nlp/app_test.py new file mode 100644 index 00000000..4ef3e6a8 --- /dev/null +++ b/lit_nlp/app_test.py @@ -0,0 +1,93 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from unittest import mock + +from absl.testing import absltest +from lit_nlp import app as lit_app +from lit_nlp.api import dataset as lit_dataset +from lit_nlp.api import types as lit_types +from lit_nlp.lib import testing_utils + + +class _TestDataset(lit_dataset.Dataset): + + def __init__(self, val: float = 1.0): + self._examples = [{'val': val}] + + def spec(self) -> lit_types.Spec: + return {'val': lit_types.Scalar()} + + def init_spec(self) -> lit_types.Spec: + return {'val': lit_types.Scalar(required=False)} + + +class AppTest(absltest.TestCase): + + def _make_app(self, demo_mode: bool) -> lit_app.LitApp: + test_model = testing_utils.IdentityRegressionModelForTesting() + return lit_app.LitApp( + models={'test_model': test_model}, + datasets={'test_ds': _TestDataset()}, + generators={}, + interpreters={}, + metrics={}, + client_root=self.create_tempdir().full_path, + demo_mode=demo_mode, + ) + + def test_create_dataset_blocked_in_demo_mode(self): + app = self._make_app(demo_mode=True) + result = app._create_dataset( + data={'config': {'new_name': 'new_ds', 'val': 2.0}}, + dataset_name='test_ds', + ) + self.assertIsNone(result) + self.assertNotIn('new_ds', app._datasets) + + def test_create_dataset_succeeds_when_not_demo_mode(self): + app = self._make_app(demo_mode=False) + result = app._create_dataset( + data={'config': {'new_name': 'new_ds', 'val': 2.0}}, + dataset_name='test_ds', + ) + self.assertIsNotNone(result) + _, new_name = result + self.assertEqual(new_name, 'new_ds') + self.assertIn('new_ds', app._datasets) + + def test_create_model_blocked_in_demo_mode(self): + app = self._make_app(demo_mode=True) + result = app._create_model( + data={'config': {'new_name': 'new_model'}}, + model_name='test_model', + ) + self.assertIsNone(result) + self.assertNotIn('new_model', app._models) + + def test_create_model_succeeds_when_not_demo_mode(self): + app = self._make_app(demo_mode=False) + result = app._create_model( + data={'config': {'new_name': 'new_model'}}, + model_name='test_model', + ) + self.assertIsNotNone(result) + _, new_names = result + self.assertEqual(new_names, ['new_model']) + self.assertIn('new_model', app._models) + + +if __name__ == '__main__': + absltest.main() diff --git a/lit_nlp/lib/file_cache.py b/lit_nlp/lib/file_cache.py index 177edad8..5205a41d 100644 --- a/lit_nlp/lib/file_cache.py +++ b/lit_nlp/lib/file_cache.py @@ -133,15 +133,16 @@ def _get_extacted_dir(output_path: str) -> str: with filelock.FileLock(lock_path): shutil.rmtree(output_extracted_path, ignore_errors=True) os.makedirs(output_extracted_path) - - if is_zip: - with zipfile.ZipFile(output_path, 'r') as zip_file: - _safe_zip_file_extractall(zip_file, output_extracted_path) - zip_file.close() - else: - tar_file = tarfile.open(output_path) - tar_file.extractall(output_extracted_path, filter='data') - tar_file.close() + try: + if is_zip: + with zipfile.ZipFile(output_path, 'r') as zip_file: + _safe_zip_file_extractall(zip_file, output_extracted_path) + else: + with tarfile.open(output_path) as tar_file: + tar_file.extractall(output_extracted_path, filter='data') + except Exception: + shutil.rmtree(output_extracted_path, ignore_errors=True) + raise return output_extracted_path diff --git a/lit_nlp/lib/file_cache_test.py b/lit_nlp/lib/file_cache_test.py index 67092ff9..36aaaac3 100644 --- a/lit_nlp/lib/file_cache_test.py +++ b/lit_nlp/lib/file_cache_test.py @@ -13,6 +13,11 @@ # limitations under the License. # ============================================================================== +import io +import os +import tarfile +import zipfile + from absl.testing import absltest from absl.testing import parameterized from lit_nlp.lib import file_cache @@ -94,6 +99,122 @@ def test_is_remote(self, url: str, expected: bool): is_remote = file_cache.is_remote(url) self.assertEqual(is_remote, expected) + def test_cached_path_extracts_valid_tar(self): + temp_dir = self.create_tempdir().full_path + tar_path = os.path.join(temp_dir, 'model.tar.gz') + payload = b'{"model": "ok"}' + with tarfile.open(tar_path, 'w:gz') as tar: + info = tarfile.TarInfo(name='subdir/config.json') + info.size = len(payload) + tar.addfile(info, io.BytesIO(payload)) + + extracted_dir = file_cache.cached_path( + tar_path, extract_compressed_file=True + ) + extracted_file = os.path.join(extracted_dir, 'subdir', 'config.json') + self.assertTrue(os.path.isfile(extracted_file)) + with open(extracted_file, 'rb') as f: + self.assertEqual(f.read(), payload) + + def test_cached_path_extracts_valid_zip(self): + temp_dir = self.create_tempdir().full_path + zip_path = os.path.join(temp_dir, 'model.zip') + payload = b'{"model": "ok"}' + with zipfile.ZipFile(zip_path, 'w') as zf: + zf.writestr('subdir/config.json', payload) + + extracted_dir = file_cache.cached_path( + zip_path, extract_compressed_file=True + ) + extracted_file = os.path.join(extracted_dir, 'subdir', 'config.json') + self.assertTrue(os.path.isfile(extracted_file)) + with open(extracted_file, 'rb') as f: + self.assertEqual(f.read(), payload) + + @parameterized.named_parameters( + ('parent_traversal', '../escaped.txt'), + ('nested_parent_traversal', 'subdir/../../escaped.txt'), + ('absolute_parent_traversal', '/../escaped.txt'), + ) + def test_cached_path_rejects_tar_traversal(self, malicious_name: str): + temp_dir = self.create_tempdir().full_path + archive_dir = os.path.join(temp_dir, 'cache') + os.makedirs(archive_dir) + tar_path = os.path.join(archive_dir, 'malicious.tar.gz') + with tarfile.open(tar_path, 'w:gz') as tar: + valid_info = tarfile.TarInfo(name='valid.txt') + valid_info.size = 2 + tar.addfile(valid_info, io.BytesIO(b'ok')) + bad_info = tarfile.TarInfo(name=malicious_name) + bad_info.size = 4 + tar.addfile(bad_info, io.BytesIO(b'evil')) + + with self.assertRaises(tarfile.FilterError): + file_cache.cached_path(tar_path, extract_compressed_file=True) + + self.assertFalse(os.path.exists(os.path.join(temp_dir, 'escaped.txt'))) + self.assertFalse(os.path.exists(os.path.join(archive_dir, 'escaped.txt'))) + expected_extracted = os.path.join(archive_dir, 'malicious-tar-gz-extracted') + self.assertFalse(os.path.exists(expected_extracted)) + + def test_cached_path_sanitizes_tar_absolute_path(self): + temp_dir = self.create_tempdir().full_path + archive_dir = os.path.join(temp_dir, 'cache') + os.makedirs(archive_dir) + outside_target = os.path.join(temp_dir, 'outside_abs.txt') + tar_path = os.path.join(archive_dir, 'abs_path.tar.gz') + with tarfile.open(tar_path, 'w:gz') as tar: + abs_info = tarfile.TarInfo(name=outside_target) + abs_info.size = 4 + tar.addfile(abs_info, io.BytesIO(b'safe')) + + extracted_dir = file_cache.cached_path( + tar_path, extract_compressed_file=True + ) + self.assertFalse(os.path.exists(outside_target)) + self.assertTrue( + os.path.isfile(os.path.join(extracted_dir, outside_target.lstrip('/'))) + ) + + def test_cached_path_rejects_tar_external_symlink(self): + temp_dir = self.create_tempdir().full_path + archive_dir = os.path.join(temp_dir, 'cache') + os.makedirs(archive_dir) + tar_path = os.path.join(archive_dir, 'symlink.tar.gz') + with tarfile.open(tar_path, 'w:gz') as tar: + sym_info = tarfile.TarInfo(name='link_out') + sym_info.type = tarfile.SYMTYPE + sym_info.linkname = '../../outside_target' + tar.addfile(sym_info) + + with self.assertRaises(tarfile.FilterError): + file_cache.cached_path(tar_path, extract_compressed_file=True) + + expected_extracted = os.path.join(archive_dir, 'symlink-tar-gz-extracted') + self.assertFalse(os.path.exists(expected_extracted)) + + @parameterized.named_parameters( + ('parent_traversal', '../escaped.txt'), + ('nested_parent_traversal', 'subdir/../../escaped.txt'), + ('absolute_path', '/tmp/escaped_abs.txt'), + ) + def test_cached_path_rejects_zip_traversal(self, malicious_name: str): + temp_dir = self.create_tempdir().full_path + archive_dir = os.path.join(temp_dir, 'cache') + os.makedirs(archive_dir) + zip_path = os.path.join(archive_dir, 'malicious.zip') + with zipfile.ZipFile(zip_path, 'w') as zf: + zf.writestr('valid.txt', b'ok') + zf.writestr(malicious_name, b'evil') + + with self.assertRaises(ValueError): + file_cache.cached_path(zip_path, extract_compressed_file=True) + + self.assertFalse(os.path.exists(os.path.join(temp_dir, 'escaped.txt'))) + self.assertFalse(os.path.exists(os.path.join(archive_dir, 'escaped.txt'))) + expected_extracted = os.path.join(archive_dir, 'malicious-zip-extracted') + self.assertFalse(os.path.exists(expected_extracted)) + # TODO(b/285157349, b/254110131): Add UT/ITs for file_cache.cached_path(). # Conditions should include: #