Refactor Kafka activity to use local kafka_connector variable for improved readability and ensure proper closure of the Kafka consumer after polling.
101 lines
3.1 KiB
Python
101 lines
3.1 KiB
Python
from temporalio import workflow, activity
|
|
|
|
with workflow.unsafe.imports_passed_through():
|
|
from logging import Logger
|
|
from sientia_do.notifications.handlers import NotificationHandler
|
|
from sientia_do.temporal.activities.base import BaseActivity
|
|
from sientia_do.temporal.utils.logger import Logger
|
|
from typing import Any
|
|
from kafka import KafkaConsumer
|
|
from pandas import DataFrame
|
|
import json
|
|
|
|
|
|
class Kafka(BaseActivity):
|
|
def __init__(self, bootstrap_servers: str, polling_time: int,
|
|
group_id: str, logger: Logger, notification_handler: NotificationHandler):
|
|
self.polling_time = polling_time
|
|
|
|
self.kafka_connector = KafkaConsumer(
|
|
bootstrap_servers=bootstrap_servers,
|
|
auto_offset_reset="earliest",
|
|
enable_auto_commit=True,
|
|
group_id=group_id,
|
|
value_deserializer=lambda x: json.loads(x.decode("utf-8"))
|
|
)
|
|
|
|
BaseActivity.__init__(self, logger, notification_handler)
|
|
|
|
def close(self):
|
|
"""Closes the connector connection."""
|
|
self.info("Closing Kafka connector...")
|
|
self.kafka_connector.close()
|
|
|
|
def __del__(self):
|
|
self.close()
|
|
|
|
@activity.defn(name="load_from_kafka")
|
|
async def load_from_kafka(self, input_data: dict[str, Any]) -> dict[str, Any]:
|
|
"""
|
|
Loads data from a kafka topic. Polls the topic for a given time and returns the data.
|
|
|
|
Args:
|
|
input_data (dict[str, Any]): The data to load. Contains:
|
|
topic (str): The topic to load data from.
|
|
Returns:
|
|
dict[str, Any]: The data loaded from the topic.
|
|
"""
|
|
|
|
metadata = input_data['metadata']
|
|
|
|
self.debug(
|
|
f"Loading data from topic: {input_data['topic']}",
|
|
metadata=metadata
|
|
)
|
|
|
|
topic = input_data["topic"]
|
|
|
|
# Subscribe to the specified topic
|
|
kafka_connector = KafkaConsumer(
|
|
topic,
|
|
bootstrap_servers=self.bootstrap_servers,
|
|
auto_offset_reset="earliest",
|
|
enable_auto_commit=True,
|
|
group_id=self.group_id,
|
|
value_deserializer=lambda x: json.loads(x.decode("utf-8"))
|
|
)
|
|
|
|
# List to store message values
|
|
message_values = []
|
|
|
|
# Poll for messages
|
|
records = kafka_connector.poll(timeout_ms=self.polling_time)
|
|
|
|
self.debug(
|
|
f"Polled {len(records)} records from topic: {topic}",
|
|
metadata=metadata
|
|
)
|
|
|
|
# Process the polled records
|
|
for _topic_partition, msgs in records.items():
|
|
for msg in msgs:
|
|
message_values.append(msg.value)
|
|
|
|
# Return empty dict if no messages were received
|
|
if not message_values:
|
|
return {}
|
|
|
|
self.debug(
|
|
f"Loaded {len(message_values)} messages from topic: {topic}",
|
|
metadata=metadata
|
|
)
|
|
|
|
self.debug(
|
|
f"Loaded data: {message_values}",
|
|
metadata=metadata
|
|
)
|
|
|
|
kafka_connector.close()
|
|
|
|
return DataFrame(message_values).to_dict()
|