zstdlib 0.0.1__tar.gz → 0.0.3__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (32) hide show
  1. {zstdlib-0.0.1/zstdlib.egg-info → zstdlib-0.0.3}/PKG-INFO +4 -1
  2. zstdlib-0.0.3/README.md +2 -0
  3. {zstdlib-0.0.1 → zstdlib-0.0.3}/pyproject.toml +8 -3
  4. zstdlib-0.0.3/tests/__init__.py +0 -0
  5. zstdlib-0.0.3/tests/log/__init__.py +0 -0
  6. zstdlib-0.0.3/tests/log/base.py +69 -0
  7. zstdlib-0.0.3/tests/log/test_cute.py +98 -0
  8. zstdlib-0.0.3/tests/log/test_trace.py +74 -0
  9. zstdlib-0.0.3/tests/test_ansi.py +56 -0
  10. zstdlib-0.0.3/tests/test_enum.py +137 -0
  11. zstdlib-0.0.3/tests/test_frozen.py +80 -0
  12. zstdlib-0.0.3/tests/test_singleton.py +126 -0
  13. zstdlib-0.0.3/zstdlib/__init__.py +6 -0
  14. zstdlib-0.0.3/zstdlib/ansi.py +172 -0
  15. zstdlib-0.0.3/zstdlib/enum.py +57 -0
  16. zstdlib-0.0.3/zstdlib/frozen.py +65 -0
  17. zstdlib-0.0.3/zstdlib/log/__init__.py +2 -0
  18. zstdlib-0.0.3/zstdlib/log/cute.py +92 -0
  19. zstdlib-0.0.3/zstdlib/log/trace.py +84 -0
  20. zstdlib-0.0.3/zstdlib/singleton.py +58 -0
  21. {zstdlib-0.0.1 → zstdlib-0.0.3/zstdlib.egg-info}/PKG-INFO +4 -1
  22. zstdlib-0.0.3/zstdlib.egg-info/SOURCES.txt +25 -0
  23. {zstdlib-0.0.1 → zstdlib-0.0.3}/zstdlib.egg-info/top_level.txt +1 -0
  24. zstdlib-0.0.1/zstdlib/__init__.py +0 -5
  25. zstdlib-0.0.1/zstdlib/enum.py +0 -32
  26. zstdlib-0.0.1/zstdlib/singleton.py +0 -28
  27. zstdlib-0.0.1/zstdlib/trace.py +0 -55
  28. zstdlib-0.0.1/zstdlib.egg-info/SOURCES.txt +0 -11
  29. {zstdlib-0.0.1 → zstdlib-0.0.3}/LICENSE +0 -0
  30. {zstdlib-0.0.1 → zstdlib-0.0.3}/setup.cfg +0 -0
  31. {zstdlib-0.0.1 → zstdlib-0.0.3}/zstdlib/py.typed +0 -0
  32. {zstdlib-0.0.1 → zstdlib-0.0.3}/zstdlib.egg-info/dependency_links.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: zstdlib
3
- Version: 0.0.1
3
+ Version: 0.0.3
4
4
  Summary: A set of useful python utilities
5
5
  License: GPLv3
6
6
  Project-URL: Homepage, https://github.com/zwimer/zstdlib
@@ -12,3 +12,6 @@ Classifier: License :: OSI Approved :: GNU General Public License v3 (GPLv3)
12
12
  Requires-Python: >=3.10
13
13
  Description-Content-Type: text/markdown
14
14
  License-File: LICENSE
15
+
16
+ # zstdlib
17
+ A set of useful python utilities
@@ -0,0 +1,2 @@
1
+ # zstdlib
2
+ A set of useful python utilities
@@ -49,11 +49,16 @@ version = {attr = "zstdlib.__version__"}
49
49
 
50
50
  # Tools
51
51
 
52
+ [tool.pylint.MASTER]
53
+ ignore-paths = '^tests/.*$'
52
54
  [tool.pylint."MESSAGES CONTROL"]
53
55
  disable = [
56
+ "unnecessary-lambda-assignment",
57
+ "method-cache-max-size-none",
54
58
  "missing-module-docstring",
55
- "invalid-name",
56
- "line-too-long"
59
+ "too-few-public-methods",
60
+ "line-too-long",
61
+ "invalid-name"
57
62
  ]
58
63
 
59
64
  [tool.black]
@@ -71,6 +76,6 @@ ignore=["E731"]
71
76
  skips = ["B101", "B104", "B201"]
72
77
 
73
78
  [tool.vulture]
74
- ignore_names = ["cli", "_help", "_show_version", "_channel", "strict_slashes"]
79
+ ignore_names = []
75
80
  min_confidence = 70
76
81
  paths = ["zstdlib"]
File without changes
File without changes
@@ -0,0 +1,69 @@
1
+ from logging.handlers import QueueHandler
2
+ from contextlib import contextmanager
3
+ from collections.abc import Callable
4
+ from threading import Lock
5
+ from typing import Any
6
+ import logging
7
+ import queue
8
+
9
+
10
+ @contextmanager
11
+ def _cml(logger: logging.Logger, exit_func: Callable[[], Any]):
12
+ try:
13
+ yield logger
14
+ finally:
15
+ exit_func()
16
+
17
+
18
+ class LeftBase:
19
+ """
20
+ A base class for tests that hijack loggers
21
+ """
22
+
23
+ _hijacked: set[logging.Logger] = set() # Set of all loggers that have ever been hijacked
24
+ _qmap: dict[logging.Logger, queue.Queue[logging.LogRecord]] = {}
25
+ _lb_old: dict[logging.Logger, list[logging.Handler]] = {}
26
+ _lb_lock = Lock()
27
+
28
+ messages: dict[logging.Logger, list[str]] = {}
29
+
30
+ @classmethod
31
+ def hijack(cls, name: str, reuse: bool = False, fmt: logging.Formatter | None = None):
32
+ """
33
+ :return: A logger configured to write to an internal queue
34
+ """
35
+ log = logging.getLogger(name)
36
+ q: queue.Queue[logging.LogRecord] = queue.Queue()
37
+ with cls._lb_lock:
38
+ if not reuse and log in cls._hijacked:
39
+ raise RuntimeError("Logger already hijacked")
40
+ cls._hijacked.add(log)
41
+ cls._lb_old[log] = log.handlers
42
+ cls._qmap[log] = q
43
+ log.handlers = [QueueHandler(q)]
44
+ if fmt is not None:
45
+ log.handlers[0].setFormatter(fmt)
46
+ log.setLevel(1) # Not 0 to ensure that the parent logger level is not used
47
+ return _cml(log, lambda: cls._restore(log))
48
+
49
+ @classmethod
50
+ def _restore(cls, logger: logging.Logger) -> None:
51
+ """
52
+ Restore the logger and read all messages from the queue
53
+ :return: The messages stored by the logger queue
54
+ """
55
+ if len(logger.handlers) != 1 or not isinstance(qh := logger.handlers[0], QueueHandler):
56
+ raise RuntimeError("Logger not hijacked")
57
+ # Restore old logger and extract q
58
+ with cls._lb_lock:
59
+ if logger not in cls._hijacked:
60
+ raise RuntimeError("Logger not hijacked")
61
+ logger.handlers = cls._lb_old.pop(logger)
62
+ q = cls._qmap.pop(logger)
63
+ # Read queue
64
+ qh.flush()
65
+ messages = []
66
+ while not q.empty():
67
+ messages.append(q.get(False).msg)
68
+ with cls._lb_lock:
69
+ cls.messages[logger] = messages
@@ -0,0 +1,98 @@
1
+ # pylint: disable=missing-module-docstring,missing-class-docstring,missing-function-docstring,unused-variable
2
+ import unittest
3
+
4
+ from zstdlib.log import CuteFormatter
5
+ from zstdlib.ansi import Color
6
+
7
+ from .base import LeftBase
8
+
9
+
10
+ class TestCuteFormatter(LeftBase, unittest.TestCase):
11
+
12
+ def test_hardcoded_colors(self) -> None:
13
+ color = Color.strikethrough_green
14
+ name = "TestCuteFormatter.test_hardcoded_colors"
15
+ with self.hijack(name, fmt=CuteFormatter(colors={name: color})) as log:
16
+ log.debug("test")
17
+ msg = self.messages[log][0].rsplit("|", 1)[-1]
18
+ self.assertEqual(f" {color.code}test{Color.RESET}", msg)
19
+
20
+ def test_name_width(self) -> None:
21
+ width = 123
22
+ color = Color.red
23
+ name = "TestCuteFormatter.name_width"
24
+ with self.hijack(name, fmt=CuteFormatter(colors={name: color}, name_width=width)) as log:
25
+ log.debug("test")
26
+ lvl, _, nam, msg = self.messages[log][0].split("|")
27
+ self.assertEqual("DEBUG", lvl.strip())
28
+ self.assertEqual(f" {color(name.ljust(width))} ", nam)
29
+ self.assertEqual(f" {color('test')}", msg)
30
+
31
+ def test_no_color(self) -> None:
32
+ name = "TestCuteFormatter.test_no_color"
33
+ with self.hijack(name, fmt=CuteFormatter(colored=False)) as log:
34
+ log.debug("test")
35
+ lvl, _, nam, msg = self.messages[log][0].split("|")
36
+ self.assertEqual("DEBUG", lvl.strip())
37
+ self.assertEqual(name, nam.strip())
38
+ self.assertEqual(" test", msg)
39
+
40
+ def test_edit_colors(self) -> None:
41
+ name = "TestCuteFormatter.edit_colors"
42
+ cf = CuteFormatter(colored=False)
43
+ color = Color.strikethrough_red
44
+ with self.hijack(name, fmt=cf) as log:
45
+ log.critical("test")
46
+ msg = self.messages[log][0].rsplit("|", 1)[-1]
47
+ self.assertEqual(" test", msg)
48
+ with self.hijack(name, reuse=True, fmt=cf) as log:
49
+ cf.colored = True
50
+ cf.update({name: color})
51
+ log.critical("test")
52
+ msg = self.messages[log][0].rsplit("|", 1)[-1]
53
+ self.assertEqual(f" {color.code}test{Color.RESET}", msg)
54
+
55
+ def test_level_color(self) -> None:
56
+ name = "TestCuteFormatter.test_level_color"
57
+ msg = "test"
58
+ levels = (5, 10, 20, 30, 40, 50, 60)
59
+ with self.hijack(name, fmt=CuteFormatter()) as log:
60
+ for i in levels:
61
+ log.log(i, msg)
62
+ self.assertEqual(len(levels), len(self.messages[log]))
63
+ l_name = lambda idx: self.messages[log][idx].split("|")[0].strip()
64
+ self.assertEqual(Color.dim("Level 5".ljust(8)), l_name(0))
65
+ self.assertEqual("DEBUG", l_name(1))
66
+ self.assertEqual(Color.blue("INFO".ljust(8)), l_name(2))
67
+ self.assertEqual(Color.yellow("WARNING".ljust(8)), l_name(3))
68
+ self.assertEqual(Color.red("ERROR".ljust(8)), l_name(4))
69
+ self.assertEqual(Color.bright_red_bg_yellow("CRITICAL".ljust(8)), l_name(5))
70
+ self.assertEqual(Color.bright_red_bg_yellow("Level 60".ljust(8)), l_name(6))
71
+
72
+ def test_exception(self) -> None:
73
+ name = "TestCuteFormatter.test_exception"
74
+ with self.hijack(name, fmt=CuteFormatter()) as log:
75
+ try:
76
+ raise ValueError(name)
77
+ except ValueError:
78
+ log.error("test", exc_info=True)
79
+ spt = self.messages[log][0].split("\n")
80
+ self.assertGreater(len(spt), 2)
81
+ self.assertEqual("Traceback (most recent call last):", spt[1].strip())
82
+ self.assertEqual("raise ValueError(name)", spt[-2].strip())
83
+ self.assertEqual(f"ValueError: {name}", spt[-1].strip())
84
+
85
+ def test_multi_color(self) -> None:
86
+ base = "TestCuteFormatter.test_multi_color."
87
+ cf = CuteFormatter()
88
+ messages: set[str] = set()
89
+ for i in range(100):
90
+ with self.hijack(f"{base}{i}", fmt=cf) as log:
91
+ log.info(base)
92
+ messages.add(self.messages[log][0].rsplit("|", 1)[-1].strip())
93
+ cols = ("red", "green", "yellow", "blue", "magenta", "cyan", "default")
94
+ self.assertSetEqual({getattr(Color, i)(base) for i in cols}, messages)
95
+
96
+
97
+ if __name__ == "__main__":
98
+ unittest.main()
@@ -0,0 +1,74 @@
1
+ # pylint: disable=missing-module-docstring,missing-class-docstring,missing-function-docstring
2
+ import unittest
3
+ import logging
4
+
5
+ from zstdlib.log import trace
6
+
7
+ from .base import LeftBase
8
+
9
+
10
+ # mypy: disable_error_code="attr-defined"
11
+ class TestTrace(LeftBase, unittest.TestCase):
12
+
13
+ def test_trace(self) -> None:
14
+ """
15
+ Avoid splitting up into multiple functions since trace.install() affects global state
16
+ Keeping this as one function ensures that the tests are run in order
17
+ """
18
+ # Test bad installs
19
+ with self.assertRaises(ValueError):
20
+ trace.install(value=-1)
21
+ with self.assertRaises(ValueError):
22
+ trace.install(value=logging.DEBUG + 1)
23
+ logging.TRACE = None
24
+ with self.assertRaises(AttributeError):
25
+ trace.install()
26
+ del logging.TRACE
27
+ logging.trace = None
28
+ with self.assertRaises(AttributeError):
29
+ trace.install()
30
+ del logging.trace
31
+ logging.getLoggerClass().trace = None
32
+ with self.assertRaises(AttributeError):
33
+ trace.install()
34
+ del logging.getLoggerClass().trace
35
+ trace._State.start = True # pylint: disable=protected-access
36
+ with self.assertRaises(RuntimeError):
37
+ trace.install()
38
+ trace._State.start = False # pylint: disable=protected-access
39
+ # Good install
40
+ trace.install(value=5)
41
+ # Attribute check
42
+ self.assertTrue(hasattr(logging, "TRACE"))
43
+ self.assertTrue(hasattr(logging, "trace"))
44
+ self.assertTrue(hasattr(logging.getLogger(), "trace"))
45
+ self.assertEqual(logging.TRACE, 5)
46
+ self.assertEqual(logging.getLevelName(logging.TRACE), "TRACE") # type: ignore[call-overload]
47
+ # Root check
48
+ with self.hijack("") as log:
49
+ old = log.getEffectiveLevel()
50
+ # pylint: disable=not-callable
51
+ logging.trace("test1") # type: ignore[misc]
52
+ log.setLevel(logging.DEBUG)
53
+ # pylint: disable=not-callable
54
+ logging.trace("test2") # type: ignore[misc]
55
+ log.setLevel(old)
56
+ self.assertEqual(self.messages[log], ["test1"])
57
+ # Logger check
58
+ with self.hijack("TestTrace.t1") as log:
59
+ old = log.getEffectiveLevel()
60
+ # pylint: disable=not-callable
61
+ log.trace("test3")
62
+ log.setLevel(logging.DEBUG)
63
+ # pylint: disable=not-callable
64
+ log.trace("test4")
65
+ log.setLevel(old)
66
+ self.assertEqual(self.messages[log], ["test3"])
67
+ # Test re-installation
68
+ with self.assertRaises(RuntimeError):
69
+ trace.install(value=5)
70
+ trace.install(value=5, force=True)
71
+
72
+
73
+ if __name__ == "__main__":
74
+ unittest.main()
@@ -0,0 +1,56 @@
1
+ # pylint: disable=missing-module-docstring,missing-class-docstring,missing-function-docstring
2
+ import unittest
3
+
4
+ from zstdlib.ansi import PureColor, RawColor, Color
5
+
6
+
7
+ class TestAnsiColor(unittest.TestCase):
8
+
9
+ def test_raw_color(self) -> None:
10
+ msg = "Hello, World!"
11
+ for i, k in PureColor.__members__.items():
12
+ self.assertEqual(f"\033[{k.value}m{msg}\033[0m", getattr(Color, i)(msg))
13
+
14
+ def test_bright_background_color(self) -> None:
15
+ msg = "Hello, World!"
16
+ for i, k in PureColor.__members__.items():
17
+ self.assertEqual(f"\033[{k.value+70}m{msg}\033[0m", getattr(Color, f"bg_bright_{i}")(msg))
18
+
19
+ def test_init(self) -> None:
20
+ msg = "Hello, World!"
21
+ cname = "black"
22
+ c = RawColor(getattr(PureColor, cname))
23
+ self.assertEqual(f"\033[{c.value}m{msg}\033[0m", getattr(Color, cname)(msg)) # Sanity check
24
+ self.assertEqual(f"\033[{c.value}m{msg}\033[0m", Color(c)(msg))
25
+ self.assertEqual(f"\033[{c.value}m{msg}\033[0m", Color(c)(msg))
26
+ self.assertEqual(f"\033[{c.value}m{msg}\033[0m", Color(Color(c))(msg))
27
+ self.assertEqual(f"\033[{c.value}m{msg}\033[0m", Color(code=Color(c).code)(msg))
28
+ # Errors
29
+ with self.assertRaises(ValueError):
30
+ _ = Color("blue_blue")
31
+ with self.assertRaises(ValueError):
32
+ _ = Color("bg_blue_bg_bright_blue")
33
+ with self.assertRaises(ValueError):
34
+ _ = Color(fmt="1", code="1")
35
+ with self.assertRaises(ValueError):
36
+ _ = Color(code="1")
37
+ with self.assertRaises(ValueError):
38
+ _ = Color("bright_bg_red")
39
+ # Error converted
40
+ with self.assertRaises(AttributeError):
41
+ _ = Color.bright_bg_red
42
+
43
+ def test_modifiers(self) -> None:
44
+ msg = "Hello, World!"
45
+ self.assertEqual(f"\033[1;34m{msg}\033[0m", Color.bold_blue(msg)) # Small modifier
46
+ self.assertEqual(f"\033[9;32m{msg}\033[0m", Color.strikethrough_green(msg)) # Large modifier
47
+ self.assertEqual(f"\033[2;9;31m{msg}\033[0m", Color.strikethrough_dim_red(msg))
48
+ self.assertEqual(f"\033[3;4;5;39m{msg}\033[0m", Color.underline_italic_blinking_default(msg))
49
+ # Test no colors / multiple colors
50
+ self.assertEqual(f"\033[1m{msg}\033[0m", Color.bold(msg))
51
+ self.assertEqual(f"\033[31;42m{msg}\033[0m", Color.red_bg_green(msg))
52
+ self.assertEqual(f"\033[1;3;32;101m{msg}\033[0m", Color.bg_bright_red_bold_italic_green(msg))
53
+
54
+
55
+ if __name__ == "__main__":
56
+ unittest.main()
@@ -0,0 +1,137 @@
1
+ # pylint: disable=missing-module-docstring,missing-class-docstring,missing-function-docstring,unused-variable
2
+ import unittest
3
+
4
+ from zstdlib.enum import EnumType, Enum
5
+
6
+
7
+ class TestEnumType(unittest.TestCase):
8
+
9
+ def test_valid(self) -> None:
10
+ class ET1(metaclass=EnumType):
11
+ arg1: int = 0
12
+ arg2: int = 1
13
+
14
+ def test_empty(self) -> None:
15
+ class ET2(metaclass=EnumType, empty_ok=True):
16
+ pass
17
+
18
+ with self.assertRaises(ValueError):
19
+
20
+ class ET3(metaclass=EnumType):
21
+ pass
22
+
23
+ def test_dupe(self) -> None:
24
+ with self.assertRaises(ValueError):
25
+
26
+ class ET4(metaclass=EnumType):
27
+ arg1: int = 0
28
+ arg2: int = 0
29
+
30
+ def test_annotations(self) -> None:
31
+ with self.assertRaises(ValueError):
32
+
33
+ class ET5(metaclass=EnumType):
34
+ arg1 = 1
35
+
36
+ with self.assertRaises(ValueError):
37
+
38
+ class ET6(metaclass=EnumType):
39
+ arg1: int
40
+
41
+ with self.assertRaises(TypeError):
42
+
43
+ class ET7(metaclass=EnumType):
44
+ arg1: str = 0 # type: ignore[assignment]
45
+
46
+ def test_instantiation(self) -> None:
47
+ with self.assertRaises(AttributeError):
48
+
49
+ class ET8(metaclass=EnumType):
50
+ arg1: int = 1
51
+
52
+ def __init__(self):
53
+ pass
54
+
55
+ with self.assertRaises(AttributeError):
56
+
57
+ class ET9(metaclass=EnumType):
58
+ arg1: int = 1
59
+
60
+ def __new__(cls):
61
+ pass
62
+
63
+ class ET10(metaclass=EnumType):
64
+ arg1: int = 0
65
+
66
+ with self.assertRaises(NotImplementedError):
67
+ ET10()
68
+ with self.assertRaises(NotImplementedError):
69
+ ET10.__init__({})
70
+
71
+
72
+ class TestEnum(unittest.TestCase):
73
+ def test_valid(self) -> None:
74
+ class E1(Enum):
75
+ arg1: int = 0
76
+ arg2: int = 1
77
+
78
+ def test_empty(self) -> None:
79
+ class E2(Enum, empty_ok=True):
80
+ pass
81
+
82
+ with self.assertRaises(ValueError):
83
+
84
+ class E3(Enum):
85
+ pass
86
+
87
+ def test_dupe(self) -> None:
88
+ with self.assertRaises(ValueError):
89
+
90
+ class E4(Enum):
91
+ arg1: int = 0
92
+ arg2: int = 0
93
+
94
+ def test_annotations(self) -> None:
95
+ with self.assertRaises(ValueError):
96
+
97
+ class E5(Enum):
98
+ arg1 = 1
99
+
100
+ with self.assertRaises(ValueError):
101
+
102
+ class E6(Enum):
103
+ arg1: int
104
+
105
+ with self.assertRaises(TypeError):
106
+
107
+ class E7(Enum):
108
+ arg1: str = 0 # type: ignore[assignment]
109
+
110
+ def test_instantiation(self) -> None:
111
+ with self.assertRaises(AttributeError):
112
+
113
+ class E8(Enum):
114
+ arg1: int = 1
115
+
116
+ def __init__(self):
117
+ pass
118
+
119
+ with self.assertRaises(AttributeError):
120
+
121
+ class E9(Enum):
122
+ arg1: int = 1
123
+
124
+ def __new__(cls):
125
+ pass
126
+
127
+ class E10(Enum):
128
+ arg1: int = 0
129
+
130
+ with self.assertRaises(NotImplementedError):
131
+ E10()
132
+ with self.assertRaises(NotImplementedError):
133
+ E10.__init__({})
134
+
135
+
136
+ if __name__ == "__main__":
137
+ unittest.main()
@@ -0,0 +1,80 @@
1
+ # pylint: disable=missing-module-docstring,missing-class-docstring,missing-function-docstring,unused-variable,attribute-defined-outside-init
2
+ import unittest
3
+
4
+ from zstdlib.frozen import Freezable, frozen
5
+
6
+
7
+ class TestFreezable(unittest.TestCase):
8
+
9
+ def test_unfrozen(self):
10
+ class F1(Freezable):
11
+ pass
12
+
13
+ f1 = F1()
14
+ f1.a = 1
15
+ self.assertEqual(f1.a, 1)
16
+ del f1.a
17
+ self.assertIs(getattr(f1, "a", None), None)
18
+
19
+ def test_frozen(self):
20
+ class F1(Freezable):
21
+ pass
22
+
23
+ f1 = F1()
24
+ f1.a = 1
25
+ f1.freeze()
26
+ self.assertEqual(f1.a, 1)
27
+ with self.assertRaises(AttributeError):
28
+ f1.a = 2
29
+ with self.assertRaises(AttributeError):
30
+ del f1.a
31
+
32
+ def test_thaw(self):
33
+ class F1(Freezable):
34
+ pass
35
+
36
+ f1 = F1()
37
+ f1.a = 1
38
+ f1.freeze()
39
+ self.assertEqual(f1.a, 1)
40
+ f1.thaw()
41
+ f1.a = 2
42
+ self.assertEqual(f1.a, 2)
43
+ del f1.a
44
+ self.assertIs(getattr(f1, "a", None), None)
45
+
46
+ def test_permanent_freeze(self):
47
+ class F1(Freezable):
48
+ pass
49
+
50
+ f1 = F1()
51
+ f1.a = 1
52
+ f1.freeze(permanent=True)
53
+ self.assertEqual(f1.a, 1)
54
+ with self.assertRaises(RuntimeError):
55
+ f1.thaw()
56
+ with self.assertRaises(AttributeError):
57
+ f1.a = 1
58
+
59
+
60
+ class TestFrozen(unittest.TestCase):
61
+
62
+ def test_frozen(self):
63
+ @frozen
64
+ class F1:
65
+ def __init__(self):
66
+ self.a = 1
67
+ self.b = 1
68
+ del self.b
69
+
70
+ f1 = F1()
71
+ self.assertEqual(f1.a, 1)
72
+ self.assertIs(getattr(f1, "b", None), None)
73
+ with self.assertRaises(AttributeError):
74
+ f1.a = 1
75
+ with self.assertRaises(AttributeError):
76
+ del f1.a
77
+
78
+
79
+ if __name__ == "__main__":
80
+ unittest.main()
@@ -0,0 +1,126 @@
1
+ # pylint: disable=missing-module-docstring,missing-class-docstring,missing-function-docstring,unused-variable
2
+ from threading import Thread, Lock
3
+ from time import sleep
4
+ import unittest
5
+
6
+ from zstdlib.singleton import SingletonType, Singleton
7
+
8
+
9
+ class TestSingletonType(unittest.TestCase):
10
+
11
+ def test_valid(self) -> None:
12
+ class ST1(metaclass=SingletonType):
13
+ pass
14
+
15
+ self.assertIs(ST1(), ST1())
16
+
17
+ class ST2(metaclass=SingletonType):
18
+ pass
19
+
20
+ self.assertIsNot(ST1(), ST2())
21
+
22
+ def test_subclass(self) -> None:
23
+ class ST3(metaclass=SingletonType):
24
+ pass
25
+
26
+ with self.assertRaises(NotImplementedError):
27
+
28
+ class ST4(ST3):
29
+ pass
30
+
31
+ def test_multi_thread(self):
32
+ """
33
+ Ensure that SingletonType is thread safe and that constructing an object doesn't delay other threads
34
+ Technically this is more of a heuristic, but it failing is extremely unlikely
35
+ """
36
+
37
+ class ST5(metaclass=SingletonType):
38
+ def __init__(self):
39
+ sleep(0.2)
40
+
41
+ class ST6(metaclass=SingletonType):
42
+ def __init__(self):
43
+ sleep(0.8)
44
+
45
+ lock = Lock()
46
+ events = []
47
+ results = []
48
+
49
+ def t1() -> None:
50
+ """
51
+ Construct an ST6 immediately
52
+ """
53
+ with lock:
54
+ pass
55
+ events.append("START: ST6()")
56
+ results.append(ST6())
57
+ events.append("END: ST6()")
58
+
59
+ def t2() -> None:
60
+ """
61
+ Construct an ST6 after the first ST6 has started construction but before it has finished
62
+ """
63
+ with lock:
64
+ pass
65
+ sleep(0.2)
66
+ events.append("START: ST6()")
67
+ results.append(ST6())
68
+ events.append("END: ST6()")
69
+
70
+ def t3() -> None:
71
+ """
72
+ Construct ST5's after both ST6s have started construction, finishing before either end
73
+ """
74
+ with lock:
75
+ pass
76
+ sleep(0.4)
77
+ # Loop Enough times that ST6 wil be complete if it ST5 actually constructed each time
78
+ for i in range(10):
79
+ events.append("START: ST5()")
80
+ results.append(ST5())
81
+ events.append("END: ST5()")
82
+
83
+ threads = (Thread(target=t1), Thread(target=t2), Thread(target=t3))
84
+ with lock:
85
+ for i in threads:
86
+ i.start()
87
+ # Give threads a moment to construct then let them go
88
+ sleep(0.2)
89
+ for i in threads:
90
+ i.join()
91
+ # Check results
92
+ wanted = ["START: ST6()"] * 2 + ["START: ST5()", "END: ST5()"] * 10 + ["END: ST6()"] * 2
93
+ self.assertEqual(events, wanted)
94
+ # Check constructed objects
95
+ self.assertEqual(len(results), 2 + 10)
96
+ for i in range(9):
97
+ self.assertIs(results[0], results[i + 1])
98
+ self.assertIsNot(results[0], results[-1])
99
+ self.assertIs(results[-1], results[-2])
100
+
101
+
102
+ class TestSingleton(unittest.TestCase):
103
+
104
+ def test_valid(self):
105
+ class S1(Singleton):
106
+ pass
107
+
108
+ self.assertIs(S1(), S1())
109
+
110
+ class S2(Singleton):
111
+ pass
112
+
113
+ self.assertIsNot(S1(), S2())
114
+
115
+ def test_subclass(self):
116
+ class S3(Singleton):
117
+ pass
118
+
119
+ with self.assertRaises(NotImplementedError):
120
+
121
+ class S4(S3):
122
+ pass
123
+
124
+
125
+ if __name__ == "__main__":
126
+ unittest.main()