diff --git a/avbroot.py b/avbroot.py index b691041..01ebbfc 100644 --- a/avbroot.py +++ b/avbroot.py @@ -211,7 +211,7 @@ def patch_subcommand(args): # Set default temp directory to the output directory because this is the # only way to control where external libraries put their temp files - tempfile.tempdir = os.path.dirname(os.path.abspath(output)) + util.set_default_temp_dir(os.path.dirname(os.path.abspath(output))) # Decrypt keys to temp directory with tempfile.TemporaryDirectory(dir=util.tmpfs_path()) as key_dir: @@ -257,7 +257,7 @@ def patch_subcommand(args): def extract_subcommand(args): # Set default temp directory to the output directory because this is the # only way to control where external libraries put their temp files - tempfile.tempdir = os.path.abspath(args.directory) + util.set_default_temp_dir(os.path.abspath(args.directory)) with zipfile.ZipFile(args.input, 'r') as z: info = z.getinfo(PATH_PAYLOAD) diff --git a/avbroot/util.py b/avbroot/util.py index d8ade89..9cfe523 100644 --- a/avbroot/util.py +++ b/avbroot/util.py @@ -57,6 +57,19 @@ def tmpfs_path(): return None +def set_default_temp_dir(path): + ''' + Set the default temporary directory for use if the user hasn't overridden + it with standard environment variables. + ''' + + if any(os.getenv(k) for k in ('TMPDIR', 'TEMP', 'TMP')): + # Rely on tempfile's default handling + return + + tempfile.tempdir = path + + @contextlib.contextmanager def open_output_file(path): '''