Refactor activity methods and update requirements.txt to enhance functionality and remove deprecated filters. Added detailed docstrings for clarity and improved error handling in data processing workflows.
169 lines
5.5 KiB
Python
169 lines
5.5 KiB
Python
import traceback
|
|
from temporalio import workflow, activity
|
|
|
|
from laborious.activities.base import BaseActivity
|
|
with workflow.unsafe.imports_passed_through():
|
|
from psycopg2.pool import ThreadedConnectionPool
|
|
from pandas import read_sql_query, DataFrame
|
|
from logging import Logger
|
|
from sientia_do.notifications.handlers import NotificationHandler
|
|
from sientia_do.notifications.models import NotificationLevel
|
|
from typing import Any
|
|
|
|
|
|
class Postgres(BaseActivity):
|
|
def __init__(self, host: str, port: int,
|
|
user: str, password: str, dbname: str,
|
|
min_connections: int, max_connections: int,
|
|
logger: Logger, notification_handler: NotificationHandler):
|
|
self.host = host
|
|
self.port = port
|
|
self.user = user
|
|
self.password = password
|
|
self.dbname = dbname
|
|
|
|
self.pool = ThreadedConnectionPool(
|
|
minconn=min_connections,
|
|
maxconn=max_connections,
|
|
host=self.host,
|
|
port=self.port,
|
|
user=self.user,
|
|
password=self.password,
|
|
dbname=self.dbname)
|
|
|
|
super().__init__(logger, notification_handler)
|
|
|
|
def close(self):
|
|
self.pool.closeall()
|
|
|
|
def __del__(self):
|
|
self.close()
|
|
|
|
@activity.defn(name="load_custom_query")
|
|
async def load_custom_query(self, query: str) -> dict[str, dict]:
|
|
"""
|
|
Loads data from a custom query.
|
|
|
|
Args:
|
|
query (str): The query to load data from.
|
|
|
|
Returns:
|
|
dict[str, dict]: The data from the query.
|
|
"""
|
|
self.logger.info(f"Fetching data from query: {query}")
|
|
|
|
conn = self.pool.getconn()
|
|
try:
|
|
data = read_sql_query(query, conn)
|
|
|
|
except Exception as e:
|
|
trace = traceback.format_exc()
|
|
self.notification_handler.build_and_send_notification(
|
|
notification_id="ERROR_LOADING_CUSTOM_QUERY",
|
|
message=f"Error fetching data from query: {e}",
|
|
block="load_custom_query",
|
|
level=NotificationLevel.ERROR,
|
|
attachment_content=trace
|
|
)
|
|
|
|
self.logger.error(trace)
|
|
|
|
return {}
|
|
finally:
|
|
self.pool.putconn(conn)
|
|
|
|
self.logger.info(f"Fetched {len(data)} rows")
|
|
self.logger.debug(f"Data: {data.to_string()}")
|
|
|
|
return data.to_dict()
|
|
|
|
@activity.defn(name="repeat_last_prediction")
|
|
async def repeat_last_prediction(self, query_items: dict[str, str]):
|
|
"""
|
|
Repeats the last prediction for a given model.
|
|
|
|
Args:
|
|
query_items (dict[str, str]): The query items. Contains:
|
|
schema (str): The schema of the table.
|
|
table_name (str): The name of the table.
|
|
model (str): The model to repeat the prediction for.
|
|
|
|
Returns:
|
|
None
|
|
"""
|
|
schema = query_items["schema"]
|
|
table_name = query_items["table_name"]
|
|
model = query_items["model"]
|
|
|
|
repeat_query = f"""
|
|
INSERT INTO \"{schema}\".{table_name} (model_id, prediction, timestamp, response_time, prediction_status, prediction_confidence, created_at)
|
|
SELECT model_id, prediction, timestamp, response_time, prediction_status, prediction_confidence, NOW()
|
|
FROM \"{schema}\".{table_name}
|
|
WHERE model_id = {model}
|
|
ORDER BY timestamp DESC
|
|
LIMIT 1;
|
|
"""
|
|
self.logger.info(f"Repeating last prediction for model {model}")
|
|
self.logger.debug(f"Query: {repeat_query}")
|
|
|
|
conn = self.pool.getconn()
|
|
|
|
try:
|
|
cursor = conn.cursor()
|
|
cursor.execute(repeat_query)
|
|
conn.commit()
|
|
cursor.close()
|
|
|
|
except Exception as e:
|
|
trace = traceback.format_exc()
|
|
self.notification_handler.build_and_send_notification(
|
|
notification_id="ERROR_REPEATING_LAST_PREDICTION",
|
|
message=f"Error repeating last prediction: {e}",
|
|
block="repeat_last_prediction",
|
|
level=NotificationLevel.ERROR,
|
|
attachment_content=trace
|
|
)
|
|
|
|
self.logger.error(trace)
|
|
|
|
finally:
|
|
self.pool.putconn(conn)
|
|
|
|
@activity.defn(name="export_data_to_postgres")
|
|
async def export_data_to_postgres(self, input_data: dict[str, Any]):
|
|
"""
|
|
Exports data to a postgres table.
|
|
|
|
Args:
|
|
input_data (dict[str, Any]): The data to export. Contains:
|
|
schema (str): The schema of the table.
|
|
table_name (str): The name of the table.
|
|
data (DataFrame): The data to export.
|
|
"""
|
|
|
|
schema = input_data["schema"]
|
|
table_name = input_data["table_name"]
|
|
data = DataFrame(input_data["data"])
|
|
|
|
conn = self.pool.getconn()
|
|
|
|
try:
|
|
data.to_sql(table_name, conn, schema=schema,
|
|
if_exists="append", index=False)
|
|
conn.commit()
|
|
|
|
except Exception as e:
|
|
trace = traceback.format_exc()
|
|
self.notification_handler.build_and_send_notification(
|
|
notification_id="ERROR_EXPORTING_DATA_TO_POSTGRES",
|
|
message=f"Error exporting data to postgres: {e}",
|
|
block="export_data_to_postgres",
|
|
level=NotificationLevel.ERROR,
|
|
attachment_content=trace
|
|
)
|
|
|
|
self.logger.error(trace)
|
|
|
|
finally:
|
|
self.pool.putconn(conn)
|