package ru.pdguard;
import static org.junit.jupiter.api.Assertions.assertTrue;
import java.util.ArrayList;
import java.util.Comparator;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import org.junit.jupiter.api.Test;
import ru.pdguard.config.SystemPolicy;
import ru.pdguard.core.PayloadStore;
import ru.pdguard.core.Pipeline;
import ru.pdguard.detect.NameCascade;
import ru.pdguard.detect.RuleRegistry;
import ru.pdguard.detect.Span;
import ru.pdguard.mask.Masker;
/**
* Оценка двухмодельной архитектуры: WikiNEuRal для имён, ruBERT для адресов.
*
*
Набор {@code benchmark-two-model.txt} проверяет, что имена клиентов и адреса маскируются, а
* известные личности — нет. Для каждого типа считается посимвольная точность, полнота и F1.
*/
class TwoModelBenchmarkTest {
private final Pipeline pipeline =
new Pipeline(
new RuleRegistry(),
new Masker(),
new PayloadStore(30),
new NameCascade(
new NameCascade.EngineConfig(
"wikineural", Optional.of("models/wikineural-ner"),
"rubert", Optional.of("models/rubert-ner"),
"off", Optional.empty()),
16,
4));
private static final class Score {
private int truePositive;
private int falsePositive;
private int falseNegative;
private int gold() {
return truePositive + falseNegative;
}
private double precision() {
int found = truePositive + falsePositive;
return found == 0 ? 1.0 : (double) truePositive / found;
}
private double recall() {
return gold() == 0 ? 1.0 : (double) truePositive / gold();
}
private double f1() {
double p = precision();
double r = recall();
return p + r == 0 ? 0.0 : 2 * p * r / (p + r);
}
}
@Test
void twoModelEfficiency() {
List samples = BenchmarkFixtures.load("/benchmark-two-model.txt");
Map byType = new LinkedHashMap<>();
Map> missed = new LinkedHashMap<>();
for (BenchmarkFixtures.Sample sample : samples) {
List found = pipeline.findPersonalData(sample.text(), SystemPolicy.DEFAULT);
String[] goldChars = paint(sample.text().length(), sample.gold());
String[] foundChars = paint(sample.text().length(), found);
for (int i = 0; i < sample.text().length(); i++) {
account(byType, goldChars[i], foundChars[i]);
}
for (Span gold : sample.gold()) {
boolean hit =
found.stream().anyMatch(f -> f.type().equals(gold.type()) && f.overlaps(gold));
if (!hit) {
missed
.computeIfAbsent(gold.type(), t -> new ArrayList<>())
.add(sample.text().substring(gold.start(), gold.end()));
}
}
}
report(byType);
reportMissed(missed);
// Каждый тип должен быть найден с F1 не ниже 0.8.
for (Map.Entry e : byType.entrySet()) {
assertTrue(
e.getValue().f1() >= 0.8,
String.format("F1 по типу %s упал до %.3f", e.getKey(), e.getValue().f1()));
}
}
private static String[] paint(int length, List spans) {
String[] painted = new String[length];
for (Span span : spans) {
for (int i = span.start(); i < Math.min(span.end(), length); i++) {
painted[i] = span.type();
}
}
return painted;
}
private static void account(Map byType, String gold, String found) {
if (gold != null) {
Score score = byType.computeIfAbsent(gold, t -> new Score());
if (gold.equals(found)) {
score.truePositive++;
} else {
score.falseNegative++;
}
}
if (found != null && !found.equals(gold)) {
byType.computeIfAbsent(found, t -> new Score()).falsePositive++;
}
}
private void report(Map byType) {
StringBuilder out = new StringBuilder(2048);
out.append("\n=== Эффективность двухмодельной архитектуры ===\n\n");
out.append(
String.format("%-20s %8s %8s %8s %8s%n", "тип", "знаков", "точность", "полнота", "F1"));
byType.entrySet().stream()
.sorted(
Comparator.comparingInt((Map.Entry e) -> e.getValue().gold()).reversed())
.forEach(
e ->
out.append(
String.format(
"%-20s %8d %8.3f %8.3f %8.3f%n",
e.getKey(),
e.getValue().gold(),
e.getValue().precision(),
e.getValue().recall(),
e.getValue().f1())));
System.out.println(out);
}
private void reportMissed(Map> missed) {
if (missed.isEmpty()) {
return;
}
StringBuilder out = new StringBuilder();
out.append("\n=== Не распознанные значения по типам ===\n");
missed.forEach(
(type, values) ->
out.append(type).append(": ").append(String.join(" | ", values)).append('\n'));
System.out.println(out);
}
}