Commit 719d1757 authored by BO ZHANG's avatar BO ZHANG 🏀
Browse files

feat(dag): 为brick DAG添加healpix支持并统一submission配置

parent d03908af
Loading
Loading
Loading
Loading
+86 −2
Original line number Diff line number Diff line
@@ -88,10 +88,11 @@ DagRunGroup.trigger(
- `obs_id`
- `detector`
- `filter`
- `healpix`
- `custom_id`
- `prc_status`
- `qc_status`
- `batch_id`
- `source_batch_id`
- `priority`
- `force_success`
- `return_data_list`
@@ -100,7 +101,10 @@ DagRunGroup.trigger(

- `data_model="raw"` 时,底层会走 level0 数据查询。
-`raw` 时,会走 level1 数据查询。
- `batch_id` 也可以只从 `trigger(..., batch_id=...)` 顶层传入,框架会自动补到 `data` 中。
- `batch_id` 顶层参数表示当前新生成的 `DagRunGroup` / 输出产品所属批次。
- `data.source_batch_id` 表示检索上游 level1 数据时使用的批次条件。
-`data.source_batch_id` 为空时,框架会默认回退到顶层 `batch_id`
- 为兼容旧调用方式,若调用方仍传 `data.batch_id`,当前实现会将其视为 `source_batch_id` 的别名。

#### `proc` 支持的参数

@@ -108,6 +112,86 @@ DagRunGroup.trigger(
- `ref_cat`
- `extra_kwargs`

#### DAG YAML 中的 `match` 配置

v2 中,DAG YAML 里的 `match` 只用于描述“这个 DAG 适用于哪些输入数据特征”,也就是匹配键本身。

推荐把这类字段放在 `match` 中:

- `instrument`
- `obs_type`
- `detector_group`

例如:

```yaml
match:
  instrument: "MSC"
  obs_type: ["WIDE", "DEEP", "CALF"]
  detector_group: "MBI"
```

约定说明:

- `match` 不应再包含 `dag_type`
- `dag_type` 不是输入匹配键,而是某个 DAG 的固有属性。
- 在当前 v2 实现里,DAG 类型由注册时使用的 Python 类决定,例如 `FileDAG``BrickDAG`
- 因此新增 DAG 时,不需要在 YAML 的 `match` 中再写 `dag_type`

#### DAG YAML 中的 `submission` 配置

v2 中,Submission 页面在选中某个 DAG 后,`data` / `proc` 两组表单默认值不再写死在前端,而是从该 DAG 自己的 YAML 中读取 `submission` 配置。

推荐写法如下:

```yaml
submission:
  data:
    data_model:
      value: "raw"
      fixed: true
    dataset:
      value: "test-msc-c9-25sqdeg-v3"
    source_batch_id:
      value: "default"
    instrument:
      value: "MSC"
    obs_type:
      value: "WIDE"
    obs_group:
      value: "W5"

  proc:
    pmapname:
      value: "csst_000155.pmap"
    ref_cat:
      value: "trilegal_093"
    extra_kwargs:
      value: "{}"
```

字段语义:

- `value`: 该字段在 Submission 页面中的默认值。
- `fixed`: 可选,默认是 `false`。当为 `true` 时,前端会将该字段置为不可编辑,后端生成模板时也会强制覆盖用户输入。

约定说明:

- `submission.data.<field>` 用于配置 `data` 组字段,例如 `data_model``dataset``source_batch_id``instrument``obs_type``healpix``custom_id`
- `submission.proc.<field>` 用于配置 `proc` 组字段,例如 `pmapname``ref_cat``extra_kwargs`
- `data_model` 虽然属于 `data` 检索条件,但通常和 DAG 定义强绑定,因此推荐在 YAML 中显式给出。
- `source_batch_id` 建议只在非 `raw` 的 level1 / L2 检索场景中使用;raw 数据场景通常不需要该字段。
- 大部分 L1 DAG 推荐固定为 `data_model.value: "raw"``fixed: true`
- 某些 L2 / brick DAG 需要固定到特定 level1 模型,例如:
  - `csst-msc-l2-mbi-mosaic` 推荐固定为 `csst-msc-l1-mbi`
  - `csst-msc-l2-mbi-xcat` 推荐固定为 `csst-msc-l1-mbi-catmix`

实现细节:

- 前端在调用 `/api/utils/dagrun-group/group-defaults` 时会读取该 DAG 的 `submission` 配置,并自动填充表单。
- 后端在调用 `/api/utils/dagrun-group/template` 生成模板时,也会再次套用 `submission` 中的默认值与固定值,避免前端绕过。
- 当前 `config_loader` 同时兼容新旧两种写法,但新增 DAG 时应统一使用上面的字段对象配置格式。

#### 返回值

`trigger()` 返回的是一个 `DagRunGroup` 对象,而不是普通字典。
+55 −0
Original line number Diff line number Diff line
@@ -6,6 +6,55 @@ DAG_CONFIG_DIR_V2 = os.path.join(
    "dag_config",
)


def _normalize_field_group(field_group: dict | None) -> tuple[dict, dict]:
    """将字段中心的 submission 配置展开为 defaults/fixed 两张表。"""
    defaults: dict = {}
    fixed: dict = {}
    if not isinstance(field_group, dict):
        return defaults, fixed

    for field_name, field_cfg in field_group.items():
        if isinstance(field_cfg, dict):
            if "value" in field_cfg and field_cfg.get("value") is not None:
                defaults[field_name] = field_cfg.get("value")
            if bool(field_cfg.get("fixed")) and "value" in field_cfg:
                fixed[field_name] = field_cfg.get("value")
            continue
        if field_cfg is not None:
            defaults[field_name] = field_cfg
    return defaults, fixed


def _normalize_submission_section(submission_cfg: dict | None) -> dict:
    """将 YAML 中的 submission 段规范化为固定结构。"""
    normalized = {
        "data_defaults": {},
        "data_fixed": {},
        "proc_defaults": {},
        "proc_fixed": {},
    }
    if not isinstance(submission_cfg, dict):
        return normalized

    # 新结构:submission.data.<field>.value/fixed 与 submission.proc.<field>.value/fixed
    data_defaults, data_fixed = _normalize_field_group(submission_cfg.get("data"))
    proc_defaults, proc_fixed = _normalize_field_group(submission_cfg.get("proc"))
    if data_defaults or data_fixed or proc_defaults or proc_fixed:
        normalized["data_defaults"] = data_defaults
        normalized["data_fixed"] = data_fixed
        normalized["proc_defaults"] = proc_defaults
        normalized["proc_fixed"] = proc_fixed
        return normalized

    # 兼容旧结构:data_defaults/data_fixed/proc_defaults/proc_fixed
    for key in normalized.keys():
        value = submission_cfg.get(key)
        if isinstance(value, dict):
            normalized[key] = dict(value)
    return normalized


def load_dag_config(dag_name: str) -> dict:
    """
    从 YAML 文件中加载增强版的 DAG 配置。
@@ -36,3 +85,9 @@ def load_dag_config(dag_name: str) -> dict:
        
    assert dag_cfg.get("name") == dag_name, f"配置名称 '{dag_cfg.get('name')}' 与文件名 '{dag_name}' 不匹配"
    return dag_cfg


def load_submission_config(dag_name: str) -> dict:
    """读取 DAG YAML 中的 submission 表单默认值/固定值配置。"""
    dag_cfg = load_dag_config(dag_name)
    return _normalize_submission_section(dag_cfg.get("submission"))
+0 −4
Original line number Diff line number Diff line
@@ -56,13 +56,9 @@ CSST_DAGS = {
    ),
    "csst-msc-l2-mbi-mosaic": BrickDAG(
        "csst-msc-l2-mbi-mosaic",
        required_data_models=["csst-msc-l1-mbi"],
        required_dags=[],
    ),
    "csst-msc-l2-mbi-xcat": BrickDAG(
        "csst-msc-l2-mbi-xcat",
        required_data_models=["csst-msc-l2-mbi-cat"],
        required_dags=[],
    ),
}

+4 −1
Original line number Diff line number Diff line
@@ -70,6 +70,8 @@ class BaseDAG(ABC):
                return None
            return str(v)

        source_batch_id = _as_str(data.get("source_batch_id")) or _as_str(data.get("batch_id"))

        plan_query = {
            "dataset": _as_str(data.get("dataset")),
            "obs_type": _as_str(data.get("obs_type")),
@@ -88,11 +90,12 @@ class BaseDAG(ABC):
            "instrument": _as_str(data.get("instrument")),
            "detector": _as_str(data.get("detector")),
            "filter": _as_str(data.get("filter")),
            "healpix": _as_str(data.get("healpix")),
            "data_model": _as_str(data.get("data_model")) or "raw",
            "custom_id": _as_str(data.get("custom_id")),
            "prc_status": data.get("prc_status"),
            "qc_status": data.get("qc_status"),
            "batch_id": _as_str(data.get("batch_id")),
            "batch_id": source_batch_id,
            "pmapname": _as_str(proc.get("pmapname")),
            "page": 1,
            "limit": 0,
+80 −63
Original line number Diff line number Diff line
@@ -2,7 +2,7 @@ from __future__ import annotations

from typing import Any, Optional

from astropy.table import Table, unique, vstack
from astropy.table import Table

from ..dfs import data as dfs_data
from ..models import DagRun, DagRunGroup
@@ -15,21 +15,15 @@ class BrickDAG(BaseDAG):

    Notes
    -----
    Brick DAG 的触发以砖块/自定义分片 ID 为单位。
    当前先支持 `required_data_models` 的根 DAG 模式;
    `required_dags` 依赖仍保留占位,后续对接任务状态查询。
    Brick DAG 的触发以 `healpix` 为分组键。
    当前逻辑直接复用 submission 页面里的 level1 查询条件,
    先筛出 level1 数据,再按唯一 `healpix` 生成一批 DagRun。
    `data_model` 视为本次 submission 的数据检索条件之一,
    需要由调用方在 `data` 中显式传入,且必须为 level1 data model。
    """

    def __init__(
        self,
        dag_name: str,
        *,
        required_data_models: Optional[list[str]] = None,
        required_dags: Optional[list[str]] = None,
    ):
    def __init__(self, dag_name: str):
        super().__init__(dag_name=dag_name, require_plan=False, require_data=True, require_upstream=True)
        self.required_data_models = list(required_data_models or [])
        self.required_dags = list(required_dags or [])

    def build_dependency_queries(
        self,
@@ -37,43 +31,61 @@ class BrickDAG(BaseDAG):
        proc: dict[str, Any],
    ) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any]]:
        plan_query, data_query, upstream_query = super().build_dependency_queries(data, proc)
        data_query["required_data_models"] = list(self.required_data_models)
        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:
        required_data_models = list(data_query.pop("required_data_models", []))
        if not required_data_models:
            return Table()

        if self.required_dags:
            raise NotImplementedError("required_dags 尚未接入任务状态查询。")

        tables = []
        for data_model in required_data_models:
            query = dict(data_query)
            query["data_model"] = data_model
            tables.append(dfs_data.find(**query))

        if not tables:
            return Table()

        if len(tables) == 1:
            return tables[0]

        # 目前以 custom_id 作为 brick 键做交集;没有 custom_id 时返回空表。
        if any("custom_id" not in t.colnames for t in tables):
            return Table()

        key_sets = [set(str(v) for v in t["custom_id"] if v not in ("", None)) for t in tables]
        common_keys = set.intersection(*key_sets) if key_sets else set()
        if not common_keys:
            return tables[0][:0]

        filtered = []
        for t in tables:
            mask = [str(v) in common_keys for v in t["custom_id"]]
            filtered.append(t[mask])
        return unique(vstack(filtered), keys=["custom_id", "data_model", "data_uuid"])
        return dfs_data.find(**data_query)

    @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

    @staticmethod
    def _normalize_healpix(value: Any) -> Optional[str]:
        """只接受非负整数形式的 healpix,过滤占位符和脏值。"""
        if value in (None, ""):
            return None
        text = str(value).strip()
        if not text or not text.isdigit():
            return None
        return text

    @staticmethod
    def _iter_unique_healpix(rows: Table) -> list[str]:
        """保持原始顺序提取唯一的整数 healpix。"""
        if "healpix" not in rows.colnames:
            return []

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

    def build_trigger_payload(
        self,
@@ -85,37 +97,42 @@ class BrickDAG(BaseDAG):
    ) -> DagRunGroup:
        dag_group = DagRunGroup(
            dag=self.dag_name,
            batch_id=data.get("batch_id"),
            batch_id=data.get("group_batch_id", data.get("batch_id")),
            priority=data.get("priority", "low"),
            docker_images=docker_images,
        )
        if deps.data is None or len(deps.data) == 0:
            return dag_group

        key_col = "custom_id" if "custom_id" in deps.data.colnames else "data_uuid"
        unique_keys = [str(v) for v in unique(deps.data[key_col]).tolist() if v not in ("", None)]
        for this_key in unique_keys:
            rows = deps.data[[str(v) == this_key for v in deps.data[key_col]]]
            first = rows[0]
        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 = 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=data.get("data_model"),
                dataset=data.get("dataset") or first["dataset"],
                instrument=data.get("instrument") or first["instrument"],
                obs_type=data.get("obs_type") or first["obs_type"],
                obs_group=data.get("obs_group") or first["obs_group"],
                obs_id=data.get("obs_id") or first["obs_id"],
                detector=data.get("detector"),
                filter=data.get("filter"),
                custom_id=this_key,
                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 "data_uuid" in rows.colnames else [],
                data_list=(
                    [str(v) for v in rows["data_uuid"] if v not in (None, "")]
                    if "data_uuid" in rows.colnames
                    else []
                ),
                pmapname=proc.get("pmapname"),
                ref_cat=proc.get("ref_cat"),
                extra_kwargs=proc.get("extra_kwargs"),
Loading