Commit 48c192b7 authored by BO ZHANG's avatar BO ZHANG 🏀
Browse files

feat: 配置csst-hstdm-l1默认不返回数据列表,优化DFS模型层级选择逻辑

parent 7a28f4fe
Loading
Loading
Loading
Loading
+1 −0
Original line number Diff line number Diff line
@@ -53,6 +53,7 @@ CSST_DAGS = {
    "csst-hstdm-l1": FileDAG(
        "csst-hstdm-l1",
        dispatcher=Dispatcher.dispatch_obsgroup,
        default_return_data_list=False,
    ),
    "csst-ifs-l2": FileDAG(
        "csst-ifs-l2",
+10 −2
Original line number Diff line number Diff line
@@ -33,9 +33,11 @@ class FileDAG(BaseDAG):
        dag_name: str,
        *,
        dispatcher: Callable[[Table, Table], list[dict]] = Dispatcher.dispatch_file,
        default_return_data_list: bool = True,
    ):
        super().__init__(dag_name=dag_name, require_plan=True, require_data=True, require_upstream=True)
        self.dispatcher = dispatcher
        self.default_return_data_list = bool(default_return_data_list)
        self.dag_cfg = load_dag_config(dag_name)
        self.match_cfg = self.dag_cfg.get("match") or {}
        self.pattern_table = self.compile_match(self.match_cfg)
@@ -113,11 +115,13 @@ class FileDAG(BaseDAG):
        proc: dict[str, Any],
        dag_group: DagRunGroup,
        force_success: bool = False,
        return_data_list: bool = True,
        return_data_list: bool | None = None,
    ) -> DagRunGroup:
        if deps.plan is None or deps.data is None:
            return dag_group

        if return_data_list is None:
            return_data_list = self.default_return_data_list
        task_list = self.dispatcher(deps.plan, deps.data)
        for this_task in task_list:
            if not (force_success or this_task.get("success")):
@@ -178,7 +182,11 @@ class FileDAG(BaseDAG):

        task_list = self.dispatcher(deps.plan, deps.data)
        force_success = bool(data.get("force_success", False))
        return_data_list = bool(data.get("return_data_list", True))
        return_data_list = data.get("return_data_list")
        if return_data_list is None:
            return_data_list = self.default_return_data_list
        else:
            return_data_list = bool(return_data_list)
        for this_task in task_list:
            if not (force_success or this_task.get("success")):
                continue
+5 −0
Original line number Diff line number Diff line
@@ -74,6 +74,9 @@ def load_data_model_registry() -> Dict[str, str]:
    for model, levels in _load_data_model_level_candidates().items():
        if len(levels) == 1:
            registry[model] = levels[0]
            continue
        if levels == ("level1", "level2"):
            registry[model] = "level2"
    return registry


@@ -90,6 +93,8 @@ def get_storage_level(data_model: str | None) -> str:
    candidates = _load_data_model_level_candidates().get(model, ())
    if not candidates:
        raise ValueError(f"未注册的数据模型: `{model}`。请先确认 DFS level0/1/2 的 list_data_models 返回值。")
    if candidates == ("level1", "level2"):
        return "level2"
    if len(candidates) > 1:
        joined_levels = ", ".join(candidates)
        raise ValueError(f"数据模型 `{model}` 同时出现在多个 DFS level: {joined_levels}")
+1 −0
Original line number Diff line number Diff line
@@ -14,6 +14,7 @@ def test_v2_dag_registry_contains_level1_level2_and_brick():
    assert isinstance(CSST_DAGS["csst-msc-l2-mbi-fphot"], FileBrickDAG)
    assert isinstance(CSST_DAGS["csst-msc-l2-mbi-photoz"], CatalogBrickDAG)
    assert isinstance(CSST_DAGS["csst-msc-l2-mbi-xcat"], CatalogBrickDAG)
    assert CSST_DAGS["csst-hstdm-l1"].default_return_data_list is False


def test_v2_dispatch_obsgroup_detector():