Browse Source

feat(indicator): IndicatorCorrelationService 关联CRUD+对序强制+Spearman校准

iwt 2 days ago
parent
commit
ee93761377

+ 267 - 0
cfc-backend/src/main/java/com/etotem/cfc/service/IndicatorCorrelationService.java

@@ -0,0 +1,267 @@
+package com.etotem.cfc.service;
+
+import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper;
+import com.etotem.cfc.entity.IndicatorCorrelation;
+import com.etotem.cfc.entity.IndicatorDefinition;
+import com.etotem.cfc.entity.IndicatorValue;
+import com.etotem.cfc.mapper.IndicatorCorrelationMapper;
+import com.etotem.cfc.mapper.IndicatorDefinitionMapper;
+import com.etotem.cfc.mapper.IndicatorValueMapper;
+import com.etotem.cfc.util.SpearmanUtil;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+import org.springframework.stereotype.Service;
+
+import javax.annotation.Resource;
+import java.math.BigDecimal;
+import java.math.RoundingMode;
+import java.util.ArrayList;
+import java.util.Date;
+import java.util.LinkedHashMap;
+import java.util.List;
+import java.util.Map;
+import java.util.stream.Collectors;
+
+/**
+ * 指标关联:人工定义(候选对+方向+人工强度)+ 真实数据校准(Spearman)。
+ * 对序规则:写入前按 definition.code 字典序比较,小的存 a、大的存 b(唯一键对称)。
+ * effective_strength:有效样本 >= 30 取 computed,否则取 manual;manual_locked=1 时校准不覆盖。
+ */
+@Service
+public class IndicatorCorrelationService {
+
+    private static final Logger log = LoggerFactory.getLogger(IndicatorCorrelationService.class);
+
+    @Resource
+    private IndicatorCorrelationMapper correlationMapper;
+
+    @Resource
+    private IndicatorDefinitionMapper definitionMapper;
+
+    @Resource
+    private IndicatorValueMapper valueMapper;
+
+    /** 有效样本阈值:>=30 才用 computed 覆盖 manual(设计 4.6) */
+    private static final int EFFECTIVE_SAMPLE_THRESHOLD = 30;
+
+    /** 人工定义/编辑关联对(强制对序) */
+    public void saveCorrelation(Long defA, Long defB, String direction, BigDecimal manualStrength,
+                                String description, Boolean manualLocked, String sourceNote, Long createdBy) {
+        if (defA == null || defB == null || defA.equals(defB)) {
+            throw new IllegalArgumentException("关联对不能为空且不能相同");
+        }
+        IndicatorDefinition da = definitionMapper.selectById(defA);
+        IndicatorDefinition db = definitionMapper.selectById(defB);
+        if (da == null || db == null) {
+            throw new IllegalArgumentException("指标定义不存在: " + defA + "/" + defB);
+        }
+        Long a = da.getCode().compareTo(db.getCode()) <= 0 ? defA : defB;
+        Long b = da.getCode().compareTo(db.getCode()) <= 0 ? defB : defA;
+
+        // 已存在则该对更新
+        IndicatorCorrelation existing = correlationMapper.selectOne(new LambdaQueryWrapper<IndicatorCorrelation>()
+                .eq(IndicatorCorrelation::getDefinitionIdA, a)
+                .eq(IndicatorCorrelation::getDefinitionIdB, b)
+                .last("LIMIT 1"));
+        IndicatorCorrelation c = existing != null ? existing : new IndicatorCorrelation();
+        c.setDefinitionIdA(a);
+        c.setDefinitionIdB(b);
+        if (c.getRelationType() == null) c.setRelationType("correlation");
+        c.setDirection(direction);
+        c.setManualStrength(manualStrength);
+        c.setEffectiveStrength(manualStrength); // 未校准时 = manual
+        if (manualLocked != null) c.setManualLocked(manualLocked ? 1 : 0);
+        if (c.getManualLocked() == null) c.setManualLocked(0);
+        c.setDescription(description);
+        c.setSourceNote(sourceNote);
+        c.setCreatedBy(createdBy);
+        if (c.getCohortScope() == null) c.setCohortScope("subject");
+        if (c.getStatus() == null || !c.getStatus().equals("inactive")) c.setStatus("active");
+        if (existing != null) {
+            c.setUpdatedAt(new Date());
+            correlationMapper.updateById(c);
+        } else {
+            c.setCreatedAt(new Date());
+            c.setUpdatedAt(new Date());
+            correlationMapper.insert(c);
+        }
+        log.info("关联对已保存: {} ({} ↔ {})", a, b, direction);
+    }
+
+    /** 关联查询(设计 6.4):返回该指标的全部关联 + 当前成员的配对证据 */
+    public Map<String, Object> listRelations(Long subjectId, Long definitionId) {
+        Map<String, Object> data = new LinkedHashMap<>();
+        List<Map<String, Object>> relations = new ArrayList<>();
+        if (definitionId == null) {
+            data.put("relations", relations);
+            return data;
+        }
+        List<IndicatorCorrelation> pairs = correlationMapper.selectList(new LambdaQueryWrapper<IndicatorCorrelation>()
+                .and(w -> w.eq(IndicatorCorrelation::getDefinitionIdA, definitionId)
+                        .or().eq(IndicatorCorrelation::getDefinitionIdB, definitionId))
+                .eq(IndicatorCorrelation::getStatus, "active"));
+
+        for (IndicatorCorrelation pair : pairs) {
+            Long otherId = pair.getDefinitionIdA().equals(definitionId)
+                    ? pair.getDefinitionIdB() : pair.getDefinitionIdA();
+            IndicatorDefinition other = definitionMapper.selectById(otherId);
+            Map<String, Object> rel = new LinkedHashMap<>();
+            rel.put("pairId", pair.getId());
+            Map<String, Object> indA = new LinkedHashMap<>();
+            indA.put("definitionId", pair.getDefinitionIdA());
+            IndicatorDefinition defA = definitionMapper.selectById(pair.getDefinitionIdA());
+            indA.put("name", defA != null ? defA.getName() : null);
+            rel.put("indicatorA", indA);
+            Map<String, Object> indB = new LinkedHashMap<>();
+            indB.put("definitionId", pair.getDefinitionIdB());
+            indB.put("name", other != null ? other.getName() : null);
+            rel.put("indicatorB", indB);
+            rel.put("direction", pair.getDirection());
+            rel.put("effectiveStrength", pair.getEffectiveStrength());
+            rel.put("strengthSource", pair.getComputedStrength() != null
+                    && pair.getSampleSize() != null && pair.getSampleSize() >= EFFECTIVE_SAMPLE_THRESHOLD
+                    ? "computed" : "manual");
+            rel.put("sampleSize", pair.getSampleSize());
+            rel.put("confidence", pair.getConfidence());
+            rel.put("manualLocked", pair.getManualLocked() != null && pair.getManualLocked() == 1);
+            rel.put("description", pair.getDescription());
+
+            // 该成员自己的配对数据支撑(如有)
+            Map<String, Object> evidence = subjectEvidence(subjectId, pair.getDefinitionIdA(), pair.getDefinitionIdB());
+            rel.put("subjectEvidence", evidence);
+            relations.add(rel);
+        }
+        data.put("relations", relations);
+        return data;
+    }
+
+    /** 成员级配对证据:点数 + 个人 rho + 是否足够(>=3) */
+    private Map<String, Object> subjectEvidence(Long subjectId, Long defA, Long defB) {
+        Map<String, Object> evidence = new LinkedHashMap<>();
+        evidence.put("points", 0);
+        evidence.put("personalRho", 0.0);
+        evidence.put("hasEnough", false);
+        if (subjectId == null) {
+            return evidence;
+        }
+        List<IndicatorValue> rowsA = valueMapper.selectList(new LambdaQueryWrapper<IndicatorValue>()
+                .eq(IndicatorValue::getSubjectId, subjectId)
+                .eq(IndicatorValue::getDefinitionId, defA)
+                .isNotNull(IndicatorValue::getReportDate)
+                .orderByAsc(IndicatorValue::getReportDate));
+        List<IndicatorValue> rowsB = valueMapper.selectList(new LambdaQueryWrapper<IndicatorValue>()
+                .eq(IndicatorValue::getSubjectId, subjectId)
+                .eq(IndicatorValue::getDefinitionId, defB)
+                .isNotNull(IndicatorValue::getReportDate)
+                .orderByAsc(IndicatorValue::getReportDate));
+        if (rowsA.size() < 3 || rowsB.size() < 3) {
+            return evidence;
+        }
+        // 按 report_date 对齐:累计最近 12 个共同日期点
+        List<java.math.BigDecimal> xs = new ArrayList<>();
+        List<java.math.BigDecimal> ys = new ArrayList<>();
+        int i = 0;
+        int j = 0;
+        while (i < rowsA.size() && j < rowsB.size() && xs.size() < 12) {
+            int cmp = rowsA.get(i).getReportDate().compareTo(rowsB.get(j).getReportDate());
+            if (cmp == 0) {
+                if (rowsA.get(i).getNumericValue() != null && rowsB.get(j).getNumericValue() != null) {
+                    xs.add(rowsA.get(i).getNumericValue());
+                    ys.add(rowsB.get(j).getNumericValue());
+                }
+                i++;
+                j++;
+            } else if (cmp < 0) {
+                i++;
+            } else {
+                j++;
+            }
+        }
+        if (xs.size() >= 3) {
+            evidence.put("points", xs.size());
+            evidence.put("personalRho", SpearmanUtil.spearmanRho(xs, ys));
+            evidence.put("hasEnough", true);
+        }
+        return evidence;
+    }
+
+    /**
+     * 全量校准(管理端手动触发 / @Scheduled 每日):
+     * 对 status=active 且 manual_locked=0 的关联对,按 subject 分组配对序列 → Spearman → Fisher z 聚合。
+     */
+    public int calibrateAll() {
+        List<IndicatorCorrelation> pairs = correlationMapper.selectList(new LambdaQueryWrapper<IndicatorCorrelation>()
+                .eq(IndicatorCorrelation::getStatus, "active"));
+        int updated = 0;
+        for (IndicatorCorrelation pair : pairs) {
+            if (pair.getManualLocked() != null && pair.getManualLocked() == 1) {
+                continue;
+            }
+            if (calibratePair(pair)) {
+                updated++;
+            }
+        }
+        log.info("指标关联校准完成,共更新 {} 对", updated);
+        return updated;
+    }
+
+    private boolean calibratePair(IndicatorCorrelation pair) {
+        // 全部观测(含所有 subject),按 (subject_id, report_date) 分组配对
+        List<IndicatorValue> rowsA = valueMapper.selectList(new LambdaQueryWrapper<IndicatorValue>()
+                .eq(IndicatorValue::getDefinitionId, pair.getDefinitionIdA())
+                .isNotNull(IndicatorValue::getReportDate));
+        List<IndicatorValue> rowsB = valueMapper.selectList(new LambdaQueryWrapper<IndicatorValue>()
+                .eq(IndicatorValue::getDefinitionId, pair.getDefinitionIdB())
+                .isNotNull(IndicatorValue::getReportDate));
+        // key = subjectId + '-' + reportDate
+        Map<String, IndicatorValue> mapA = rowsA.stream().filter(v -> v.getSubjectId() != null)
+                .collect(Collectors.toMap(v -> v.getSubjectId() + "-" + v.getReportDate(), v -> v, (x, y) -> x));
+        Map<String, IndicatorValue> mapB = rowsB.stream().filter(v -> v.getSubjectId() != null)
+                .collect(Collectors.toMap(v -> v.getSubjectId() + "-" + v.getReportDate(), v -> v, (x, y) -> x));
+
+        // 按 subject 分组配对点
+        Map<Long, List<java.math.BigDecimal[]>> bySubject = new java.util.HashMap<>();
+        int sampleSize = 0;
+        for (Map.Entry<String, IndicatorValue> e : mapA.entrySet()) {
+            IndicatorValue vb = mapB.get(e.getKey());
+            if (vb == null) continue;
+            IndicatorValue va = e.getValue();
+            if (va.getNumericValue() == null || vb.getNumericValue() == null) continue;
+            bySubject.computeIfAbsent(va.getSubjectId(), k -> new ArrayList<>())
+                    .add(new java.math.BigDecimal[]{va.getNumericValue(), vb.getNumericValue()});
+            sampleSize++;
+        }
+
+        if (sampleSize < 3) {
+            return false;
+        }
+        List<Double> rhos = new ArrayList<>();
+        for (List<java.math.BigDecimal[]> points : bySubject.values()) {
+            if (points.size() < 3) {
+                continue;
+            }
+            List<java.math.BigDecimal> xs = points.stream().map(p -> p[0]).collect(Collectors.toList());
+            List<java.math.BigDecimal> ys = points.stream().map(p -> p[1]).collect(Collectors.toList());
+            rhos.add(SpearmanUtil.spearmanRho(xs, ys));
+        }
+        if (rhos.isEmpty()) {
+            return false;
+        }
+        double avgRho = SpearmanUtil.fisherZAverage(rhos.stream().mapToDouble(Double::doubleValue).toArray());
+        String confidence = sampleSize >= 30 ? "high" : (sampleSize >= 10 ? "medium" : "low");
+
+        if (sampleSize >= EFFECTIVE_SAMPLE_THRESHOLD) {
+            pair.setEffectiveStrength(BigDecimal.valueOf(avgRho).setScale(3, RoundingMode.HALF_UP));
+        }
+        pair.setComputedStrength(BigDecimal.valueOf(avgRho).setScale(3, RoundingMode.HALF_UP));
+        pair.setMethod("spearman");
+        pair.setSampleSize(sampleSize);
+        pair.setConfidence(confidence);
+        pair.setComputedAt(new Date());
+        pair.setUpdatedAt(new Date());
+        correlationMapper.updateById(pair);
+        log.info("校准关联对 {}: rho={}, n={}, confidence={}",
+                pair.getId(), avgRho, sampleSize, confidence);
+        return true;
+    }
+}

+ 57 - 0
cfc-backend/src/test/java/com/etotem/cfc/service/IndicatorCorrelationServiceTest.java

@@ -0,0 +1,57 @@
+package com.etotem.cfc.service;
+
+import com.etotem.cfc.entity.IndicatorCorrelation;
+import com.etotem.cfc.entity.IndicatorDefinition;
+import com.etotem.cfc.mapper.IndicatorCorrelationMapper;
+import com.etotem.cfc.mapper.IndicatorDefinitionMapper;
+import org.junit.jupiter.api.BeforeEach;
+import org.junit.jupiter.api.Test;
+import org.springframework.test.util.ReflectionTestUtils;
+
+import java.math.BigDecimal;
+
+import static org.junit.jupiter.api.Assertions.*;
+import static org.mockito.ArgumentMatchers.any;
+import static org.mockito.Mockito.*;
+
+class IndicatorCorrelationServiceTest {
+
+    private IndicatorCorrelationService service;
+    private IndicatorCorrelationMapper correlationMapper;
+    private IndicatorDefinitionMapper definitionMapper;
+
+    @BeforeEach
+    void setUp() {
+        service = new IndicatorCorrelationService();
+        correlationMapper = mock(IndicatorCorrelationMapper.class);
+        definitionMapper = mock(IndicatorDefinitionMapper.class);
+        ReflectionTestUtils.setField(service, "correlationMapper", correlationMapper);
+        ReflectionTestUtils.setField(service, "definitionMapper", definitionMapper);
+    }
+
+    @Test
+    void saveOrdersByCodeLexicographic() {
+        // 定义 A code="bacteria.bb" 应存 a;定义 B code="bacteria.aa" 应存 a
+        IndicatorDefinition defAA = new IndicatorDefinition();
+        defAA.setId(1L);
+        defAA.setCode("bacteria.aa");
+        IndicatorDefinition defBB = new IndicatorDefinition();
+        defBB.setId(2L);
+        defBB.setCode("bacteria.bb");
+        when(definitionMapper.selectById(1L)).thenReturn(defAA);
+        when(definitionMapper.selectById(2L)).thenReturn(defBB);
+
+        service.saveCorrelation(1L, 2L, "positive", new BigDecimal("0.72"),
+                "描述", false, "来源", 9L);
+
+        org.mockito.ArgumentCaptor<IndicatorCorrelation> captor =
+                org.mockito.ArgumentCaptor.forClass(IndicatorCorrelation.class);
+        verify(correlationMapper).insert(captor.capture());
+        IndicatorCorrelation saved = captor.getValue();
+        // 1L 的 code("bacteria.aa") 字典序更小 → 应存 definition_id_a
+        assertEquals(1L, saved.getDefinitionIdA());
+        assertEquals(2L, saved.getDefinitionIdB());
+        // effective_strength 初始 = manual(未校准)
+        assertEquals(0, new BigDecimal("0.72").compareTo(saved.getEffectiveStrength()));
+    }
+}