From fadcccb89dccbf60190e7b796306462dcc0e6920 Mon Sep 17 00:00:00 2001 From: Viet-Anh Nguyen Date: Sun, 30 Aug 2026 10:40:04 +0700 Subject: [PATCH] fix: close label files after reading and writing --- anylabeling/views/labeling/label_file.py | 3 +- tests/test_label_file.py | 50 ++++++++++++++++++++++++ 2 files changed, 52 insertions(+), 1 deletion(-) create mode 100644 tests/test_label_file.py diff --git a/anylabeling/views/labeling/label_file.py b/anylabeling/views/labeling/label_file.py index ec82c4d..28992a0 100644 --- a/anylabeling/views/labeling/label_file.py +++ b/anylabeling/views/labeling/label_file.py @@ -17,7 +17,8 @@ def io_open(name, mode): assert mode in ["r", "w"] encoding = "utf-8" - yield open(name, mode, encoding=encoding) + with open(name, mode, encoding=encoding) as file: + yield file class LabelFileError(Exception): diff --git a/tests/test_label_file.py b/tests/test_label_file.py new file mode 100644 index 0000000..b2ebd1d --- /dev/null +++ b/tests/test_label_file.py @@ -0,0 +1,50 @@ +"""Tests for reliable label-file reads and overwrites.""" + +import io +import os +import tempfile +import unittest + +from PIL import Image + +from anylabeling.views.labeling.label_file import LabelFile, io_open + + +class TestLabelFileIO(unittest.TestCase): + def test_io_open_closes_file_after_context(self): + with tempfile.TemporaryDirectory() as directory: + filename = os.path.join(directory, "labels.json") + + with io_open(filename, "w") as file: + file.write("{}") + opened_file = file + + self.assertTrue(opened_file.closed) + + def test_existing_label_file_can_be_overwritten_and_removed(self): + image_buffer = io.BytesIO() + Image.new("RGB", (2, 2)).save(image_buffer, format="PNG") + + with tempfile.TemporaryDirectory() as directory: + filename = os.path.join(directory, "labels.json") + label_file = LabelFile() + for revision in (1, 2): + label_file.save( + filename=filename, + shapes=[], + image_path="image.png", + image_height=2, + image_width=2, + image_data=image_buffer.getvalue(), + other_data={"revision": revision}, + ) + + loaded = LabelFile(filename) + self.assertEqual(loaded.other_data["revision"], 2) + + os.remove(filename) + self.assertFalse(os.path.exists(filename)) + + +if __name__ == "__main__": + unittest.main()