Skip to content
This repository was archived by the owner on Aug 13, 2026. It is now read-only.
Prev Previous commit
Next Next commit
Refactored _close
  • Loading branch information
gkevinzheng committed Mar 13, 2025
commit 1b037b6ebe7ebc8af6fa732cf840458961599656
45 changes: 22 additions & 23 deletions google/cloud/logging_v2/handlers/transports/background_thread.py
Original file line number Diff line number Diff line change
Expand Up @@ -171,7 +171,7 @@ def start(self):
)
self._thread.daemon = True
self._thread.start()
atexit.register(self._close)
atexit.register(self._handle_exit)

def stop(self, *, grace_period=None):
"""Signals the background thread to stop.
Expand Down Expand Up @@ -211,33 +211,20 @@ def stop(self, *, grace_period=None):

return success

def _close(self):
"""Callback that attempts to send pending logs before termination."""
def _close(self, close_msg):
"""Callback that attempts to send pending logs before termination if the main thread is alive."""
if not self.is_alive:
return

# Print different messages to the user depending on whether or not the
# program is shutting down. This is because this function now handles both
# the atexit handler and the regular close.
if not self._queue.empty():
if threading.main_thread().is_alive():
print(
"Background thread shutting down, attempting to send %d queued log "
"entries to Cloud Logging..." % (self._queue.qsize(),),
file=sys.stderr,
)
else:
print(
_CLOSE_THREAD_SHUTDOWN_ERROR_MSG,
file=sys.stderr,
)
print(close_msg, file=sys.stderr)

if (
threading.main_thread().is_alive() and
self.stop(grace_period=self._grace_period)
and threading.main_thread().is_alive()
):
print("Sent all pending logs.", file=sys.stderr)
else:
elif not self._queue.empty():
print(
"Failed to send %d pending logs." % (self._queue.qsize(),),
file=sys.stderr,
Expand Down Expand Up @@ -277,10 +264,22 @@ def flush(self):
def close(self):
"""Signals the worker thread to stop, then closes the transport thread.

This call should be followed up by disowning the transport object.
This call will attempt to send pending logs before termination, and
should be followed up by disowning the transport object.
"""
atexit.unregister(self._handle_exit)
self._close(
"Background thread shutting down, attempting to send %d queued log "
"entries to Cloud Logging..." % (self._queue.qsize(),)
)

def _handle_exit(self):
"""Handle system exit.

Since we cannot send pending logs during system shutdown due to thread errors,
log an error message to stderr to notify the user.
"""
atexit.unregister(self._close)
self._close()
self._close(_CLOSE_THREAD_SHUTDOWN_ERROR_MSG)


class BackgroundThreadTransport(Transport):
Expand Down Expand Up @@ -342,4 +341,4 @@ def flush(self):

def close(self):
"""Closes the worker thread."""
self.worker.stop(grace_period=self.grace_period)
self.worker.close()
52 changes: 40 additions & 12 deletions tests/unit/handlers/transports/test_background_thread.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,7 @@ def test_start(self):
self.assertTrue(worker._thread.daemon)
self.assertEqual(worker._thread._target, worker._thread_main)
self.assertEqual(worker._thread._name, background_thread._WORKER_THREAD_NAME)
self.assertIn(worker._close, atexit_mock.registered_funcs)
self.assertIn(worker._handle_exit, atexit_mock.registered_funcs)

# Calling start again should not start a new thread.
current_thread = worker._thread
Expand Down Expand Up @@ -291,21 +291,25 @@ def test__close(self):
worker = self._make_one(_Logger(self.NAME))

self._start_with_thread_patch(worker)
worker._close()
worker._close("")

self.assertFalse(worker.is_alive)

# Calling twice should not be an error
worker._close()
worker._close("")

def test__close_non_empty_queue(self):
worker = self._make_one(_Logger(self.NAME))
msg = "My Message"

self._start_with_thread_patch(worker)
record = mock.Mock()
record.created = time.time()
worker.enqueue(record, "")
worker._close()

with mock.patch("sys.stderr", new_callable=StringIO) as stderr_mock:
worker._close(msg)
self.assertIn(msg, stderr_mock.getvalue())

self.assertFalse(worker.is_alive)

Expand All @@ -317,11 +321,11 @@ def test__close_did_not_join(self):
record = mock.Mock()
record.created = time.time()
worker.enqueue(record, "")
worker._close()
worker._close("")

self.assertFalse(worker.is_alive)

def test__close_main_thread_not_alive(self):
def test__handle_exit(self):
from google.cloud.logging_v2.handlers.transports.background_thread import (
_CLOSE_THREAD_SHUTDOWN_ERROR_MSG,
)
Expand All @@ -333,21 +337,45 @@ def test__close_main_thread_not_alive(self):
with self._init_atexit_mock():
self._start_with_thread_patch(worker)
self._enqueue_record(worker, "test")
worker._close()
worker._handle_exit()

self.assertRegex(
stderr_mock.getvalue(),
re.compile("^%s$" % _CLOSE_THREAD_SHUTDOWN_ERROR_MSG, re.MULTILINE),
)

self.assertRegex(
stderr_mock.getvalue(),
re.compile(r"^Failed to send %d pending logs\.$" % worker._queue.qsize(), re.MULTILINE),
)

def test__handle_exit_no_items(self):
worker = self._make_one(_Logger(self.NAME))

with mock.patch("sys.stderr", new_callable=StringIO) as stderr_mock:
with self._init_main_thread_is_alive_mock(False):
with self._init_atexit_mock():
self._start_with_thread_patch(worker)
worker._handle_exit()

self.assertEqual(stderr_mock.getvalue(), "")

def test_close_unregister_atexit(self):
worker = self._make_one(_Logger(self.NAME))

with self._init_atexit_mock() as atexit_mock:
self._start_with_thread_patch(worker)
self.assertIn(worker._close, atexit_mock.registered_funcs)
worker.close()
self.assertNotIn(worker._close, atexit_mock.registered_funcs)
with mock.patch("sys.stderr", new_callable=StringIO) as stderr_mock:
with self._init_atexit_mock() as atexit_mock:
self._start_with_thread_patch(worker)
self.assertIn(worker._handle_exit, atexit_mock.registered_funcs)
worker.close()
self.assertNotIn(worker._handle_exit, atexit_mock.registered_funcs)

self.assertNotRegex(
stderr_mock.getvalue(),
re.compile(r"^Failed to send %d pending logs\.$" % worker._queue.qsize(), re.MULTILINE),
)

self.assertFalse(worker.is_alive)

@staticmethod
def _enqueue_record(worker, message, levelno=logging.INFO, **kw):
Expand Down