"""Kafka Producer Fixture.""" from collections import defaultdict from unittest.mock import patch from confluent_kafka import KafkaError import pytest class MockMessage: """Mock Message. A mocked version of the Message class that confluent-kafka implements. The Message class is not user-instantiable. https://docs.confluent.io/platform/current/clients/confluent-kafka-python/html/index.html#pythonclient-message """ def __init__( self, topic, key, value, headers, timestamp=0, partition=-1, offset=0, error=None ): """Init.""" self._data = { 'topic': topic, 'key': key, 'value': value, 'partition': partition, 'offset': offset, 'timestamp': timestamp, 'headers': headers, 'error': error } def assert_values(self, **kwargs): """Match kwargs against the message. Checks the subset of fields passed as kwargs.""" for field_name, value in kwargs.items(): assert field_name in self._data, (field_name, value, self._data) assert self._data[field_name] == value, (field_name, value, self._data) return True def topic(self): """Get topic.""" return self._data['topic'] def key(self): """Get key.""" return self._data['key'] def value(self): """Get value.""" return self._data['value'] def partition(self): """Get partition.""" return self._data['partition'] def offset(self): """Get offset.""" return self._data['offset'] def timestamp(self): """Get timestamp.""" return self._data['timestamp'] def headers(self): """Get headers.""" return self._data['headers'] def error(self): """Get error.""" return self._data['error'] class MockHistory: """Mock History. Tracks the calls the calls made to MockProducer methods and produced Messages. """ def __init__(self): """Init.""" self.method_calls = defaultdict(int) self.messages = [] def add_message(self, message): """Add message to history.""" self.messages.append(message) class MockProducer: """Mock Producer. Mocks the confluent kafka producer class. Tracks a queue of messages which can be cleared via flush/poll calls. Can be configured to call the callback with an error state. https://docs.confluent.io/platform/current/clients/confluent-kafka-python/html/index.html#pythonclient-producer """ def __init__(self, conf): """Init.""" self.queue = [] self.history = MockHistory() self.return_error = False def reset(self): """Reset state.""" self.queue = [] self.history = MockHistory() self.return_error = False def set_return_error(self, return_error): """Set return_error flag.""" self.return_error = return_error def produce( self, topic, value=None, key=None, partition=-1, on_delivery=None, timestamp=0, headers=None ): """Produce.""" self.history.method_calls['produce'] += 1 err = None if self.return_error: err = KafkaError(17) mock_msg = MockMessage(topic, key, value, headers, timestamp, partition, err) self.queue.append((err, mock_msg, on_delivery)) self.history.add_message(mock_msg) def flush(self, interval=0): """Flush.""" self.history.method_calls['flush'] += 1 while len(self.queue): err, msg, callback = self.queue.pop() if callback is not None: callback(err, msg) def poll(self, interval=0): """Poll, triggering the callback for the most recent message only.""" self.history.method_calls['poll'] += 1 if not len(self.queue): return err, msg, callback = self.queue.pop() if callback is not None: callback(err, msg) def assert_messages(self, check_messages): """Assert a list of dicts match the message history.""" assert len(check_messages) == len(self.history.messages), \ (check_messages, self.history.messages) for idx, check_message in enumerate(check_messages): self.history.messages[idx].assert_values(**check_message) def assert_call_count(self, method_name, call_count): """Assert the call count of a producer method.""" assert self.history.method_calls[method_name] == call_count, \ (method_name, call_count, self.history.method_calls) def pytest_collection_modifyitems(session, config, items): """Inject mocked producer in all tests.""" for item in items: if 'no_kafka_utils_patch' in item.keywords: continue if 'kafka_mock' not in item.fixturenames: item.fixturenames.append('kafka_mock') @pytest.fixture def kafka_mock(request): """Fixture mocking kafka producer.""" global mock_producer mock_producer = MockProducer({}) def return_mock(_): return mock_producer with patch('kafka_utils.producer.event.Producer', new=return_mock): yield mock_producer mock_producer = None