From 6d2b596bb41729e17cafad2f40a3c2933a195878 Mon Sep 17 00:00:00 2001 From: sebres Date: Tue, 8 Jul 2014 10:56:19 +0200 Subject: [PATCH] observer inherits from JailThread + test cases for bad run; --- fail2ban/server/jailthread.py | 4 ++-- fail2ban/server/observer.py | 4 ++-- fail2ban/tests/observertestcase.py | 30 +++++++++++++++++++++++++----- 3 files changed, 29 insertions(+), 9 deletions(-) diff --git a/fail2ban/server/jailthread.py b/fail2ban/server/jailthread.py index 2627cebf..fa135b71 100644 --- a/fail2ban/server/jailthread.py +++ b/fail2ban/server/jailthread.py @@ -47,8 +47,8 @@ class JailThread(Thread): The time the thread sleeps for in the loop. """ - def __init__(self): - super(JailThread, self).__init__() + def __init__(self, name=None): + super(JailThread, self).__init__(name=name) ## Control the state of the thread. self.active = False ## Control the idle state of the thread. diff --git a/fail2ban/server/observer.py b/fail2ban/server/observer.py index 2db8144d..7f73d3d4 100644 --- a/fail2ban/server/observer.py +++ b/fail2ban/server/observer.py @@ -26,6 +26,7 @@ __copyright__ = "Copyright (c) 2014 Serg G. Brester" __license__ = "GPL" import threading +from .jailthread import JailThread import os, logging, time, datetime, math, json, random import sys from ..helpers import getLogger @@ -34,7 +35,7 @@ from .mytime import MyTime # Gets the instance of the logger. logSys = getLogger(__name__) -class ObserverThread(threading.Thread): +class ObserverThread(JailThread): """Handles observing a database, managing bad ips and ban increment. Parameters @@ -236,7 +237,6 @@ class ObserverThread(threading.Thread): def start(self): with self._queue_lock: if not self.active: - self.active = True super(ObserverThread, self).start() def stop(self): diff --git a/fail2ban/tests/observertestcase.py b/fail2ban/tests/observertestcase.py index e8d813cf..bc484f07 100644 --- a/fail2ban/tests/observertestcase.py +++ b/fail2ban/tests/observertestcase.py @@ -551,16 +551,15 @@ class BanTimeIncrDB(unittest.TestCase): # stop observer obs.stop() -class ObserverTest(unittest.TestCase): +class ObserverTest(LogCaptureTestCase): def setUp(self): """Call before every test case.""" - #super(ObserverTest, self).setUp() - pass + super(ObserverTest, self).setUp() def tearDown(self): - #super(ObserverTest, self).tearDown() - pass + """Call after every test case.""" + super(ObserverTest, self).tearDown() def testObserverBanTimeIncr(self): obs = ObserverThread() @@ -592,3 +591,24 @@ class ObserverTest(unittest.TestCase): self.assertTrue(obs.is_alive()) obs.stop() obs = None + + class _BadObserver(ObserverThread): + def run(self): + raise RuntimeError('run bad thread exception') + + def testObserverBadRun(self): + obs = ObserverTest._BadObserver() + # don't wait for empty by stop + obs.wait_empty = lambda v:() + # save previous hook, prevent write stderr and check hereafter __excepthook__ was executed + prev_exchook = sys.__excepthook__ + x = [] + sys.__excepthook__ = lambda *args: x.append(args) + obs.start() + obs.stop() + obs = None + self.assertTrue(self._is_logged("Unhandled exception")) + sys.__excepthook__ = prev_exchook + self.assertEqual(len(x), 1) + self.assertEqual(x[0][0], RuntimeError) + self.assertEqual(str(x[0][1]), 'run bad thread exception')