import logging from datetime import datetime, timezone from criticalpath import Node from airflow.api import AirflowAPI from internal.queries import rds_query AVG_TIME_BETWEEN_TASKS = 0.9 # sec, calculated empirically AVG_SKIPPED_TASK_DURATION = 0.27 # sec, calculated empirically MIN_TASK_EXECUTION_TIME = 1 # sec, used for new tasks or as a lower bound for uncertain predictions ML_ENRICHMENT_ECS_VS_BATCH_THRESHOLD = 10000 # has to be in sync with fansifter-airflow value def _get_collection_airflow_data(schema, collection_id): return rds_query( f""" SELECT c.id AS collection_id, ca.dag_id, ca.dag_run_id, ca.fan_count FROM {schema}.collection_airflow ca RIGHT JOIN {schema}.collection c ON ca.collection_id = c.id WHERE c.id = %(collection_id)s ORDER BY ca.execution_date DESC """, dict(collection_id=collection_id), ) def _get_dag_graph_edges(dag_id): dag_tasks = AirflowAPI.get_dag_tasks(dag_id)["tasks"] edges = [] for task in dag_tasks: for downstream_task_id in task["downstream_task_ids"]: edges.append([task["task_id"], downstream_task_id]) return edges def _get_airflow_runtime_models(task_ids): return rds_query( f""" SELECT task_id, source_attribute_ids, model, model_params, coef_, intercept_ FROM commons.airflow_runtime_models WHERE task_id IN %(task_ids)s; """, dict(task_ids=tuple(task_ids)), ) def _estimate_runtime(task_instance, task_model, fan_count, collection_source_atrribute_ids): task_id = task_model["task_id"] if ( task_id == "EnrichMachineLearningClusters.Training-ECS" and fan_count >= ML_ENRICHMENT_ECS_VS_BATCH_THRESHOLD ) or ( task_id == "EnrichMachineLearningClusters.Training-Batch" and fan_count < ML_ENRICHMENT_ECS_VS_BATCH_THRESHOLD ): return AVG_SKIPPED_TASK_DURATION + AVG_TIME_BETWEEN_TASKS X = [fan_count] + [ 1 if task_source_attribute_id in collection_source_atrribute_ids else 0 for task_source_attribute_id in task_model["source_attribute_ids"] ] # constructing model features: fan_count, attribute_i, attribute_j, attribute_k, ... # Linear Regression: y = a_i * x_i + a_j * x_j + ... + b runtime = sum([a_i * x_i for a_i, x_i in zip(task_model["coef_"], X)]) + task_model["intercept_"] if task_instance["state"] == "running": task_start_date = datetime.fromisoformat(task_instance["start_date"]) seconds_already_passed_from_start = (datetime.now(timezone.utc) - task_start_date).seconds runtime -= seconds_already_passed_from_start # deducting time which has already passed return max(runtime, 0) + 0.5 * MIN_TASK_EXECUTION_TIME else: # for each task, 1/2 of AVG_TIME_BETWEEN_TASKS is accounted when task is started, # the other 1/2 - when it is finished return max(runtime, MIN_TASK_EXECUTION_TIME) + AVG_TIME_BETWEEN_TASKS def get_enrichment_progress(schema, collection_id, is_test=False): collection_airflow_data = _get_collection_airflow_data(schema, collection_id) if len(collection_airflow_data) == 0: raise ValueError(f"Collection {collection_id} does not exist.") _, dag_id, dag_run_id, fan_count = collection_airflow_data[0].values() if dag_run_id is None: raise ValueError(f"There's no enrichment runtime data available for collection {collection_id}.") try: dag_run = AirflowAPI.get_dag_run(dag_id, dag_run_id) except Exception as e: logging.error(e) raise ValueError(f"There's no enrichment runtime data available in Airflow for collection {collection_id}.") task_instances = AirflowAPI.get_task_instances(dag_id, dag_run_id)["task_instances"] if len(task_instances) == 0: raise ValueError(f"There's no enrichment tasks runtime data available in Airflow for collection {collection_id}.") # TODO rewrite w/o pandas # if some tasks were retried, we keep only the latest ones (which have the highest 'try_number') # task_instances_df = task_instances_df.sort_values("try_number", ascending=False).drop_duplicates(["task_id"]) tasks_total = len(task_instances) tasks_completed = len([task for task in task_instances if task["state"] in ["success", "skipped"]]) if dag_run["state"] in ["success", "failed"] and not is_test: return dict( status=dag_run["state"], tasksCompleted=tasks_completed, tasksTotal=tasks_total, estimatedTimeRemaining=0 ) dag_edges = _get_dag_graph_edges(dag_id) runtime_models = _get_airflow_runtime_models([ti["task_id"] for ti in task_instances]) task_id_to_model_dict = {model["task_id"]: model for model in runtime_models} graph = Node("DAG") vertices = {} for task_instance in task_instances: if task_instance["state"] in ["success", "skipped"] and not is_test: duration = 0 # task is already completed elif task_instance["task_id"] == "CalculateAnalyticsCache": # no need to estimate this task - the collection status during this task is already "analyzing" duration = 0 elif task_instance["task_id"] in task_id_to_model_dict: duration = _estimate_runtime( task_instance, task_id_to_model_dict[task_instance["task_id"]], fan_count, dag_run["conf"].get("source_attribute_ids", []), ) else: duration = MIN_TASK_EXECUTION_TIME + AVG_TIME_BETWEEN_TASKS vertices[task_instance["task_id"]] = graph.add(Node(task_instance["task_id"], duration=duration)) for vertex_from, vertex_to in dag_edges: graph.link(vertices[vertex_from], vertices[vertex_to]) graph.update_all() estimated_time_remaining = graph.duration return dict( status=dag_run["state"], tasksCompleted=tasks_completed, tasksTotal=tasks_total, estimatedTimeRemaining=estimated_time_remaining, ) def save_airflow_runtime_models(endpoint, req, event): import pandas as pd df = pd.DataFrame(event["df"]) try: for i, row in df.iterrows(): rds_query( [ f""" DELETE FROM commons.airflow_runtime_models WHERE task_id = %(task_id)s; """, f""" INSERT INTO commons.airflow_runtime_models(task_id, source_attribute_ids, model, model_params, coef_, intercept_) VALUES(%(task_id)s, %(source_attribute_ids)s, %(model)s, %(model_params)s, %(coef_)s, %(intercept_)s); """, ], vars=dict( task_id=row["task_id"], source_attribute_ids=row["source_attribute_ids"], model=row["model"], model_params=row["model_params"], coef_=row["coef_"], intercept_=row["intercept_"], ), fetch=False, ) except Exception as e: logging.error(e) raise e def get_enrichment_progress_endpoint(endpoint, req, event): workspace_schema, alliance_schema, collection_id = ( req.get("workspace_schema"), req.get("alliance_schema"), req["collectionId"], ) return get_enrichment_progress(workspace_schema or alliance_schema, collection_id) progress_bar_endpoints = {"airflow/progress/save-runtime-models": {"query": save_airflow_runtime_models}}