File size: 2,864 Bytes
1607c63
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
"""
Phase 2 — ground-truth impact tracing.

Given an agent's predicted list of affected consumers, compare against
the cascade scenario's ground truth and emit a precision/recall-style
``ImpactTraceResult`` so the reward layer can score each consumer
decision independently (composable rubric).
"""

from dataclasses import dataclass, field
from typing import List

from .service_graph import CascadeScenario


@dataclass
class ImpactTraceResult:
    """Per-consumer outcome of one ``trace_impact`` action."""

    correct_hits: List[str] = field(default_factory=list)
    missed: List[str] = field(default_factory=list)
    false_flags: List[str] = field(default_factory=list)
    unknown_services: List[str] = field(default_factory=list)
    total_consumers: int = 0

    @property
    def precision(self) -> float:
        flagged = len(self.correct_hits) + len(self.false_flags)
        if flagged == 0:
            return 0.0
        return len(self.correct_hits) / flagged

    @property
    def recall(self) -> float:
        truly_affected = len(self.correct_hits) + len(self.missed)
        if truly_affected == 0:
            return 1.0  # nothing to find
        return len(self.correct_hits) / truly_affected

    @property
    def f1(self) -> float:
        p, r = self.precision, self.recall
        if p + r == 0:
            return 0.0
        return 2 * p * r / (p + r)


def _normalise(name: str) -> str:
    return name.strip().lower()


def trace_impact(
    scenario: CascadeScenario,
    predicted_affected: List[str],
) -> ImpactTraceResult:
    """Compare agent's predicted consumer list against ground truth.

    ``predicted_affected`` is matched case-insensitively. Names that match
    no consumer in the scenario are reported in ``unknown_services`` and
    treated as false flags for reward purposes.
    """
    truth_lookup = {_normalise(n): n for n in scenario.ground_truth_affected}
    known_consumers = {_normalise(c.name): c.name for c in scenario.consumers}

    seen: set = set()
    correct_hits: List[str] = []
    false_flags: List[str] = []
    unknown_services: List[str] = []

    for raw in predicted_affected:
        key = _normalise(raw)
        if key in seen:
            continue
        seen.add(key)

        if key in truth_lookup:
            correct_hits.append(truth_lookup[key])
        elif key in known_consumers:
            false_flags.append(known_consumers[key])
        else:
            unknown_services.append(raw)

    missed = [
        name
        for name in scenario.ground_truth_affected
        if _normalise(name) not in {_normalise(c) for c in correct_hits}
    ]

    return ImpactTraceResult(
        correct_hits=correct_hits,
        missed=missed,
        false_flags=false_flags,
        unknown_services=unknown_services,
        total_consumers=len(scenario.consumers),
    )