import threading import unittest class ThreadSafeCounter: """A lock-protected integer counter.""" def __init__(self, initial: int = 0) -> None: self._value = initial self._lock = threading.Lock() def increment(self, amount: int = 1) -> int: """Increment the counter and return its new value.""" with self._lock: self._value += amount return self._value def get(self) -> int: """Return the current value.""" with self._lock: return self._value def reset(self) -> int: """Reset to zero and return the previous value.""" with self._lock: previous = self._value self._value = 0 return previous class ThreadSafeCounterTests(unittest.TestCase): def test_increment_and_reset(self) -> None: counter = ThreadSafeCounter(2) self.assertEqual(counter.increment(3), 5) self.assertEqual(counter.reset(), 5) self.assertEqual(counter.get(), 0) def test_concurrent_increments(self) -> None: counter = ThreadSafeCounter() workers = 8 increments_per_worker = 10_000 def increment_many() -> None: for _ in range(increments_per_worker): counter.increment() threads = [ threading.Thread(target=increment_many) for _ in range(workers) ] for thread in threads: thread.start() for thread in threads: thread.join() self.assertEqual(counter.get(), workers * increments_per_worker) if __name__ == "__main__": unittest.main()