from __future__ import annotations import json import os from pathlib import Path import numpy as np import pandas as pd import pytest pytestmark = pytest.mark.integration def test_tensor_bridge_smoke_artifacts_contract() -> None: out = os.environ.get("SCTM_TENSOR_BRIDGE_OUTPUT") if not out: return root = Path(out) tensor_dir = root / "tensor" stage0_dir = root / "stage0" manifest = json.loads((root / "tensor_builder_manifest.json").read_text(encoding="utf-8")) split_counts = pd.read_csv(root / "tensor_split_counts.csv") meta = json.loads((tensor_dir / "tensor_metadata.json").read_text(encoding="utf-8")) vocab = json.loads((tensor_dir / "cat_value_vocab.json").read_text(encoding="utf-8")) assert "变动情况" not in "".join(meta["cat_cols"]) assert "失访原因" not in "".join(meta["cat_cols"]) assert not any("变动情况=" in key or "失访原因=" in key for key in vocab) for field in [ "治疗方式", "服药方式", "服药时间", "是否免费服药", "是否转诊", "转诊类型", "随访类型", "本次访视方式", "本次访视对象", "是否应急处置", "健康体检", "专科医生意见-治疗", "基础管理分级", ]: assert field in meta["cat_cols"] assert "公安滋事总数" in meta["num_cols"] assert "公安肇事总数" in meta["num_cols"] assert "has_static" in meta["static_cols"] assert manifest["n_windowed_patients"] == manifest["n_patients_used"] assert {"n_singleton_patients", "n_transition_eligible_windows", "n_windowed_patients"}.issubset(split_counts.columns) assert int(split_counts["n_windowed_patients"].sum()) == manifest["n_patients_used"] assert int(split_counts["n_transition_eligible_windows"].sum()) > 0 assert (stage0_dir / "ordinal_direction_table.json").exists() assert (stage0_dir / "service_state_transitions_train.json").exists() with np.load(Path(meta["split_paths"]["val"]), allow_pickle=True) as z: assert z["cat_value_ids"].shape[1:] == (meta["max_seq_len"], len(meta["cat_cols"])) assert z["numeric_values"].shape[1:] == (meta["max_seq_len"], len(meta["num_cols"])) assert z["ordinal_cbe"].shape[1:] == (meta["max_seq_len"], len(meta["ord_cols"]), meta["cbe_dim"]) assert z["static_value_ids"].shape[1] == len(meta["static_cols"]) assert z["valid_mask"].sum() > 0