Commit 560e34fc authored by BO ZHANG's avatar BO ZHANG 🏀
Browse files

feat(dag): 新增数据模型注册表并扩展BrickDAG能力

parent 8a16250d
Loading
Loading
Loading
Loading
+11 −3
Original line number Diff line number Diff line
from .base import BaseDAG, DependencyBundle
from .brick import BrickDAG
from .brick import BrickDAG, CatalogBrickDAG, FileBrickDAG
from .dispatcher import Dispatcher
from .file import FileDAG

@@ -54,10 +54,16 @@ CSST_DAGS = {
        "csst-ifs-l2",
        dispatcher=Dispatcher.dispatch_obsid,
    ),
    "csst-msc-l2-mbi-mosaic": BrickDAG(
    "csst-msc-l2-mbi-mosaic": FileBrickDAG(
        "csst-msc-l2-mbi-mosaic",
    ),
    "csst-msc-l2-mbi-xcat": BrickDAG(
    "csst-msc-l2-mbi-fphot": FileBrickDAG(
        "csst-msc-l2-mbi-fphot",
    ),
    "csst-msc-l2-mbi-photoz": FileBrickDAG(
        "csst-msc-l2-mbi-photoz",
    ),
    "csst-msc-l2-mbi-xcat": CatalogBrickDAG(
        "csst-msc-l2-mbi-xcat",
    ),
}
@@ -65,8 +71,10 @@ CSST_DAGS = {
__all__ = [
    "BaseDAG",
    "BrickDAG",
    "CatalogBrickDAG",
    "CSST_DAGS",
    "DependencyBundle",
    "Dispatcher",
    "FileDAG",
    "FileBrickDAG",
]
+159 −74
Original line number Diff line number Diff line
from __future__ import annotations

from itertools import product
from typing import Any, Optional

from astropy.table import Table
from astropy.table import Column, Table

from ..config_loader import load_dag_config
from ..dfs import data as dfs_data
from ..dfs.model_registry import get_storage_level
from ..models import DagRun, DagRunGroup
from .base import BaseDAG, DependencyBundle


class BrickDAG(BaseDAG):
class BaseBrickDAG(BaseDAG):
    """
    v2 Brick DAG。
    v2 Brick DAG 基类

    Notes
    -----
    Brick DAG 的触发以 `healpix` 为分组键。
    当前逻辑直接复用 submission 页面里的 level1 查询条件,
    先筛出 level1 数据,再按唯一 `healpix` 生成一批 DagRun。
    `data_model` 视为本次 submission 的数据检索条件之一,
    需要由调用方在 `data` 中显式传入,且必须为 level1 data model。
    Brick DAG 通过 submission 页面 `data` 面板中的筛选条件,
    调用 DFS 的 `find_unique_healpix()` 检索唯一 healpix 列表。
    随后以 healpix 为基本分组键生成 DagRun;若 DAG 配置中声明了
    `brick_multipliers`,则进一步生成 `healpix × multiplier` 的排列组合。
    """

    source_level = ""
    UNIQUE_QUERY_KEYS = (
        "data_model",
        "dataset",
        "batch_id",
        "obs_group",
        "obs_id",
        "instrument",
        "detector",
        "obs_type",
        "filter",
        "healpix",
        "custom_id",
        "prc_status",
        "qc_status",
    )

    def __init__(self, dag_name: str):
        super().__init__(dag_name=dag_name, require_plan=False, require_data=True, require_upstream=True)
        self.dag_cfg = load_dag_config(dag_name)
        self.brick_multipliers_cfg = self.dag_cfg.get("brick_multipliers") or []

    def build_dependency_queries(
        self,
@@ -33,31 +54,15 @@ class BrickDAG(BaseDAG):
        plan_query, data_query, upstream_query = super().build_dependency_queries(data, proc)
        data_model = str(data_query.get("data_model") or "").strip()
        if not data_model or data_model.lower() == "raw":
            raise ValueError("BrickDAG 需要在 data 中显式提供非 raw 的 data_model。")
        return plan_query, data_query, upstream_query

    def resolve_data_dependency(self, data_query: dict[str, Any]) -> Table:
        return dfs_data.find(**data_query)
            raise ValueError(f"{self.__class__.__name__} 需要在 data 中显式提供非 raw 的 data_model。")

    @staticmethod
    def _pick_group_value(rows: Table, column: str, preferred: Any = None) -> Optional[str]:
        """优先使用用户显式输入,否则仅在组内唯一时回填该字段。"""
        if preferred not in (None, ""):
            return str(preferred)
        if column not in rows.colnames:
            return None

        values: list[str] = []
        seen: set[str] = set()
        for value in rows[column]:
            if value in (None, ""):
                continue
            text = str(value)
            if text in seen:
                continue
            seen.add(text)
            values.append(text)
        return values[0] if len(values) == 1 else None
        actual_level = get_storage_level(data_model)
        if actual_level != self.source_level:
            raise ValueError(
                f"{self.__class__.__name__} 需要 `{self.source_level}` data_model,"
                f"当前 `{data_model}` 注册在 `{actual_level}`。"
            )
        return plan_query, data_query, upstream_query

    @staticmethod
    def _normalize_healpix(value: Any) -> Optional[str]:
@@ -69,23 +74,79 @@ class BrickDAG(BaseDAG):
            return None
        return text

    def _build_unique_healpix_query(self, data_query: dict[str, Any]) -> dict[str, Any]:
        """仅保留 unique_healpix 接口需要的筛选条件。"""
        return {
            key: value
            for key, value in data_query.items()
            if key in self.UNIQUE_QUERY_KEYS and value not in (None, "")
        }

    def resolve_unique_healpix(self, data_query: dict[str, Any]) -> list[str]:
        """统一通过 v2 DFS wrapper 查询唯一 healpix。"""
        return dfs_data.find_unique_healpix(**data_query)

    def resolve_data_dependency(self, data_query: dict[str, Any]) -> Table:
        unique_healpix = self.resolve_unique_healpix(self._build_unique_healpix_query(data_query))
        if not unique_healpix:
            return Table([Column([], dtype=object, name="healpix")])
        return Table([{"healpix": value} for value in unique_healpix])

    @staticmethod
    def _iter_unique_healpix(rows: Table) -> list[str]:
        """保持原始顺序提取唯一的整数 healpix"""
        if "healpix" not in rows.colnames:
    def _normalize_multiplier_values(values: Any) -> list[str]:
        """规范化 multiplier 的 values 列表并去重"""
        if values in (None, ""):
            return []
        if not isinstance(values, (list, tuple, set)):
            values = [values]

        values: list[str] = []
        normalized: list[str] = []
        seen: set[str] = set()
        for value in rows["healpix"]:
            text = BrickDAG._normalize_healpix(value)
            if text is None:
        for item in values:
            if item in (None, ""):
                continue
            text = str(item)
            if text in seen:
                continue
            seen.add(text)
            values.append(text)
        return values
            normalized.append(text)
        return normalized

    def _resolve_run_multiplier_values(self, data: dict[str, Any], field: str, values: Any) -> list[str]:
        """优先使用 submission 显式值,否则回退到 DAG 配置中的 multiplier。"""
        explicit_value = data.get(field)
        if explicit_value not in (None, ""):
            return [str(explicit_value)]
        return self._normalize_multiplier_values(values)

    def _resolve_run_multiplier_combinations(self, data: dict[str, Any]) -> list[dict[str, str]]:
        """解析并展开本次 Brick 需要附加到 DagRun 的 multiplier 组合。"""
        if not isinstance(self.brick_multipliers_cfg, list):
            raise ValueError("`brick_multipliers` 必须是列表。")

        specs: list[tuple[str, list[str]]] = []
        for item in self.brick_multipliers_cfg:
            if not isinstance(item, dict):
                raise ValueError("`brick_multipliers` 的每一项都必须是字典。")

            field = str(item.get("field") or "").strip()
            if not field:
                raise ValueError("`brick_multipliers` 缺少 `field`。")

            values = self._resolve_run_multiplier_values(data, field, item.get("values"))
            if not values:
                return []
            specs.append((field, values))

        if not specs:
            return [{}]

        combinations: list[dict[str, str]] = []
        for combo_values in product(*[values for _, values in specs]):
            combinations.append({
                field: value for (field, _), value in zip(specs, combo_values)
            })
        return combinations

    def build_trigger_payload(
        self,
@@ -104,39 +165,63 @@ class BrickDAG(BaseDAG):
        if deps.data is None or len(deps.data) == 0:
            return dag_group

        unique_healpix = self._iter_unique_healpix(deps.data)
        for this_healpix in unique_healpix:
            rows = deps.data[
                [self._normalize_healpix(v) == this_healpix for v in deps.data["healpix"]]
        run_multiplier_combinations = self._resolve_run_multiplier_combinations(data)
        if not run_multiplier_combinations:
            return dag_group
        unique_healpix = [
            self._normalize_healpix(value) for value in deps.data["healpix"]
        ]
            run = DagRun(
                dag=self.dag_name,
                dag_run_group=dag_group["dag_run_group"],
                batch_id=dag_group.get("batch_id"),
                priority=dag_group.get("priority", "low"),
                data_model=self._pick_group_value(rows, "data_model", data.get("data_model")),
                dataset=self._pick_group_value(rows, "dataset", data.get("dataset")),
                instrument=self._pick_group_value(rows, "instrument", data.get("instrument")),
                obs_type=self._pick_group_value(rows, "obs_type", data.get("obs_type")),
                obs_group=self._pick_group_value(rows, "obs_group", data.get("obs_group")),
                obs_id=self._pick_group_value(rows, "obs_id", data.get("obs_id")),
                detector=self._pick_group_value(rows, "detector", data.get("detector")),
                filter=self._pick_group_value(rows, "filter", data.get("filter")),
                healpix=this_healpix,
                custom_id=data.get("custom_id"),
                prc_status=data.get("prc_status"),
                qc_status=data.get("qc_status"),
                n_file_expected=len(rows),
                n_file_found=len(rows),
                data_list=(
                    [str(v) for v in rows["data_uuid"] if v not in (None, "")]
                    if "data_uuid" in rows.colnames
                    else []

        base_run_kwargs = {
            "dag": self.dag_name,
            "dag_run_group": dag_group["dag_run_group"],
            "batch_id": dag_group.get("batch_id"),
            "priority": dag_group.get("priority", "low"),
            "data_model": (str(data.get("data_model")) if data.get("data_model") not in (None, "") else None),
            "dataset": (str(data.get("dataset")) if data.get("dataset") not in (None, "") else None),
            "source_batch_id": (
                str(data.get("source_batch_id"))
                if data.get("source_batch_id") not in (None, "")
                else None
            ),
                pmapname=proc.get("pmapname"),
                ref_cat=proc.get("ref_cat"),
                extra_kwargs=proc.get("extra_kwargs"),
            )
            "instrument": (str(data.get("instrument")) if data.get("instrument") not in (None, "") else None),
            "obs_type": (str(data.get("obs_type")) if data.get("obs_type") not in (None, "") else None),
            "obs_group": (str(data.get("obs_group")) if data.get("obs_group") not in (None, "") else None),
            "obs_id": (str(data.get("obs_id")) if data.get("obs_id") not in (None, "") else None),
            "detector": (str(data.get("detector")) if data.get("detector") not in (None, "") else None),
            "filter": (str(data.get("filter")) if data.get("filter") not in (None, "") else None),
            "custom_id": (str(data.get("custom_id")) if data.get("custom_id") not in (None, "") else None),
            "prc_status": data.get("prc_status"),
            "qc_status": data.get("qc_status"),
            "pmapname": proc.get("pmapname"),
            "ref_cat": proc.get("ref_cat"),
            "extra_kwargs": proc.get("extra_kwargs"),
        }

        for this_healpix in unique_healpix:
            if this_healpix is None:
                continue
            for multiplier_values in run_multiplier_combinations:
                run_kwargs = dict(base_run_kwargs)
                run_kwargs.update(multiplier_values)
                run_kwargs["healpix"] = this_healpix
                run = DagRun(**run_kwargs)
                dag_group.append_dag_run(run)

        return dag_group


class FileBrickDAG(BaseBrickDAG):
    """基于 `level1.find_unique_healpix()` 的文件型 Brick DAG。"""

    source_level = "level1"


class CatalogBrickDAG(BaseBrickDAG):
    """基于 `level2.find_unique_healpix()` 的星表型 Brick DAG。"""

    source_level = "level2"


class BrickDAG(FileBrickDAG):
    """兼容旧导入路径的别名,默认仍表示 level1 文件型 Brick。"""
+1 −0
Original line number Diff line number Diff line
@@ -127,6 +127,7 @@ class FileDAG(BaseDAG):
                priority=dag_group.get("priority", "low"),
                data_model=task.get("data_model"),
                dataset=task.get("dataset"),
                source_batch_id=data.get("source_batch_id"),
                instrument=task.get("instrument"),
                obs_type=task.get("obs_type"),
                obs_group=task.get("obs_group"),
+1 −0
Original line number Diff line number Diff line
@@ -46,4 +46,5 @@ tasks:
  - name: QC0
    image: csst-cpic-l1-qc0
    command: ["run", "{{ params.input }}"]
    pool_slots: 1
+1 −0
Original line number Diff line number Diff line
@@ -47,4 +47,5 @@ tasks:
  - name: CPIC
    image: csst-cpic-l1
    command: ["run", "{{ params.input }}"]
    pool_slots: 1
Loading