diff --git a/LICENSE b/LICENSE.txt similarity index 100% rename from LICENSE rename to LICENSE.txt diff --git a/etw/etw.py b/etw/etw.py index 8fb428b..ca8b5fd 100644 --- a/etw/etw.py +++ b/etw/etw.py @@ -168,7 +168,7 @@ class EventProvider: status = et.StartTraceW(ct.byref(self.session_handle), self.session_name, self.session_properties.get()) if status != tdh.ERROR_SUCCESS: - raise ct.WinError() + raise ct.WinError(status) if self.kernel_trace is False: for provider in self.providers: @@ -185,7 +185,7 @@ class EventProvider: 0, provider.params) if status != tdh.ERROR_SUCCESS: - raise ct.WinError() + raise ct.WinError(status) def stop(self): """ @@ -208,14 +208,14 @@ class EventProvider: 0, None) if status != tdh.ERROR_SUCCESS: - raise ct.WinError() + raise ct.WinError(status) status = et.ControlTraceW(self.session_handle, self.session_name, self.session_properties.get(), et.EVENT_TRACE_CONTROL_STOP) if status != tdh.ERROR_SUCCESS: - raise ct.WinError() + raise ct.WinError(status) et.CloseTrace(self.session_handle) @@ -255,7 +255,7 @@ class EventConsumer: self.index = 0 self.task_name_filters = task_name_filters self.event_id_filters = event_id_filters - self.callback_data_flag = callback_data_flag + self.callback_data_flag = callback_data_flag if not callback_data_flag else self.check_callback_flag(callback_data_flag) # NOQA if not trace_logfile: # Construct the EVENT_TRACE_LOGFILE structure @@ -306,6 +306,21 @@ class EventConsumer: # before pulling the rug out from underneath it. self.process_thread.join() + @staticmethod + def check_callback_flag(flag): + """ + Checks callback flags. + + :return: Returns flags on success, on failure raises exception + """ + flags = [RETURN_RAW_DATA_ONLY, + RETURN_RAW_DATA_ON_ERROR, + RETURN_ONLY_RAW_DATA_ON_ERROR, + RETURN_RAW_UNFORMATTED_DATA] + if flag not in flags: + raise Exception('Callback flag value {:d} passed into EventConsumer is invalid'.format(flag)) + return flag + @staticmethod def _run(trace_handle, end_capture): """ @@ -342,7 +357,7 @@ class EventConsumer: status = tdh.TdhGetEventInformation(record, 0, None, info, ct.byref(buffer_size)) if tdh.ERROR_SUCCESS != status: - raise ct.WinError() + raise ct.WinError(status) return info @@ -371,11 +386,11 @@ class EventConsumer: status = tdh.TdhGetPropertySize(record, 0, None, 1, ct.byref(data_descriptor), ct.byref(property_size)) if tdh.ERROR_SUCCESS != status: - raise ct.WinError() + raise ct.WinError(status) status = tdh.TdhGetProperty(record, 0, None, 1, ct.byref(data_descriptor), property_size, ct.byref(count)) if tdh.ERROR_SUCCESS != status: - raise ct.WinError() + raise ct.WinError(status) return count if flags & tdh.PropertyParamFixedCount: @@ -410,7 +425,7 @@ class EventConsumer: status = tdh.TdhGetPropertySize(record, 0, None, 1, ct.byref(data_descriptor), ct.byref(property_size)) if tdh.ERROR_SUCCESS != status: - raise ct.WinError() + raise ct.WinError(status) status = tdh.TdhGetProperty(record, 0, @@ -420,7 +435,7 @@ class EventConsumer: property_size, ct.cast(ct.byref(length), ct.POINTER(ct.c_byte))) if tdh.ERROR_SUCCESS != status: - raise ct.WinError() + raise ct.WinError(status) return length.value in_type = event_property.epi_u1.nonStructType.InType @@ -613,7 +628,6 @@ class EventConsumer: parsed_data = {} field_parse_error = False - raw_msg = True if self.callback_data_flag == RETURN_RAW_UNFORMATTED_DATA: event_id = 0 @@ -662,8 +676,11 @@ class EventConsumer: task_name = task_name.strip().upper() - # Add a description for the event - description = rel_ptr_to_str(info, info.contents.EventMessageOffset) + # Add a description for the event, if present + if info.contents.EventMessageOffset: + description = rel_ptr_to_str(info, info.contents.EventMessageOffset) + else: + description = '' # Windows 7 does not support predicate filters. Instead, we use a whitelist to filter things on the # consumer. @@ -698,7 +715,6 @@ class EventConsumer: # Add the description field in parsed_data['Description'] = description parsed_data['Task Name'] = task_name - raw_msg = False except Exception as e: logger.warning('Unable to parse event: {}'.format(e)) @@ -709,20 +725,14 @@ class EventConsumer: out['UserData'] = b''.join([ct.cast(record.contents.UserData + i, wt.PBYTE).contents for i in range(record.contents.UserDataLength)]) - if (self.callback_data_flag == RETURN_RAW_DATA_ON_ERROR and field_parse_error is True) or \ - self.callback_data_flag == 0: + if (self.callback_data_flag == RETURN_ONLY_RAW_DATA_ON_ERROR and field_parse_error is False) or \ + self.callback_data_flag == RETURN_RAW_DATA_ON_ERROR or self.callback_data_flag == 0: out.update(parsed_data) # Call the user's specified callback function - if (self.callback_data_flag == RETURN_RAW_DATA_ONLY and raw_msg is True) or \ - ((self.callback_data_flag == RETURN_RAW_DATA_ON_ERROR or - self.callback_data_flag == RETURN_ONLY_RAW_DATA_ON_ERROR) and - (field_parse_error is True or raw_msg is True)) or \ - (self.callback_data_flag == RETURN_RAW_UNFORMATTED_DATA and raw_msg is True) or\ - (self.callback_data_flag == 0 and raw_msg is False): - - if self.event_callback: + if self.event_callback: self.event_callback((event_id, out)) + except Exception as e: logger.error('Exception during callback: {}'.format(e)) logger.error(traceback.format_exc()) @@ -832,7 +842,8 @@ class ETW: try: self.provider.start() except WindowsError as wex: - if ct.GetLastError() == tdh.ERROR_ALREADY_EXISTS and not self.ignore_exists_error: + if (wex.winerror == tdh.ERROR_ALREADY_EXISTS and not self.ignore_exists_error) or \ + wex.winerror != tdh.ERROR_ALREADY_EXISTS: raise wex # Start the consumer @@ -1030,7 +1041,7 @@ def get_keywords_bitmask(guid, keywords): ct.byref(providers_size)) if tdh.ERROR_SUCCESS != status and tdh.ERROR_NOT_FOUND != status: - raise ct.WinError() + raise ct.WinError(status) if provider_info: field_info_array = ct.cast(provider_info.contents.FieldInfoArray, ct.POINTER(tdh.PROVIDER_FIELD_INFO)) diff --git a/setup.cfg b/setup.cfg new file mode 100644 index 0000000..e91ee50 --- /dev/null +++ b/setup.cfg @@ -0,0 +1,5 @@ +[metadata] +description-file = README.md + +[flake8] +max-line-length = 120 diff --git a/setup.py b/setup.py index 07135b1..74d59c8 100644 --- a/setup.py +++ b/setup.py @@ -36,10 +36,12 @@ if sys.version_info < (3, 4): setup(name='wintrace', - version='0.1.0', + version='0.1.1', description='ETW Tracing', author='Anthony Berglund', + author_email='anthony.berglund@fireeye.com', url='https://github.com/fireeye/pywintrace', + download_url='https://github.com/fireeye/pywintrace/archive/v0.1.1.tar.gz', platforms=['Windows'], license='Apache', packages=['etw'], diff --git a/tests/test_etw.py b/tests/test_etw.py index 5a9ff88..b500f26 100644 --- a/tests/test_etw.py +++ b/tests/test_etw.py @@ -21,7 +21,7 @@ import ctypes.wintypes as wt import subprocess as sp from etw import ETW, ProviderInfo -from etw.etw import TraceProperties, ProviderParameters +from etw.etw import TraceProperties, ProviderParameters, EventConsumer from etw.GUID import GUID from etw import evntrace as et from etw import evntprov as ep @@ -166,8 +166,8 @@ class TestETW(unittest.TestCase): capture.start() - # start powershell - args = ['powershell'] + # start ping + args = ['ping.exe'] p = sp.Popen(args, stdout=sp.DEVNULL, stderr=sp.DEVNULL) time.sleep(5) p.kill() @@ -250,8 +250,8 @@ class TestETW(unittest.TestCase): event_callback=lambda event_tufo: self.event_tufo_list.append(event_tufo)) capture.start() - # start powershell - args = ['powershell'] + # start ping.exe + args = ['ping.exe'] p = sp.Popen(args, stdout=sp.DEVNULL, stderr=sp.DEVNULL) time.sleep(2) p.kill() @@ -312,6 +312,30 @@ class TestETW(unittest.TestCase): return + def test_callback_flag_good(self): + """ + Test to check good flag value + + :return: None + """ + self.assertNotEqual(EventConsumer('test', None, None, None, common.RETURN_RAW_DATA_ONLY), None) + self.assertNotEqual(EventConsumer('test', None, None, None, common.RETURN_RAW_DATA_ON_ERROR), None) + self.assertNotEqual(EventConsumer('test', None, None, None, common.RETURN_ONLY_RAW_DATA_ON_ERROR), None) + self.assertNotEqual(EventConsumer('test', None, None, None, common.RETURN_RAW_UNFORMATTED_DATA), None) + + def test_callback_flag_bad(self): + """ + Test to check bad flag value + + :return: None + """ + consumer = None + try: + consumer = EventConsumer('test', None, None, None, 1234) + except: + pass + self.assertEqual(consumer, None) + if __name__ == '__main__': unittest.main()