RED-10200: Spike: Performant update logic for facts in working memory

This commit is contained in:
maverickstuder 2024-11-19 17:16:47 +01:00
parent 3f606ad567
commit 62591ebb1a

View File

@ -2,15 +2,20 @@ package com.iqser.red.service.redaction.v1.server.service.drools;
import static com.iqser.red.service.redaction.v1.server.service.drools.ComponentDroolsExecutionService.RULES_LOGGER_GLOBAL; import static com.iqser.red.service.redaction.v1.server.service.drools.ComponentDroolsExecutionService.RULES_LOGGER_GLOBAL;
import java.time.OffsetDateTime;
import java.util.ArrayList;
import java.util.Collections; import java.util.Collections;
import java.util.HashMap;
import java.util.HashSet; import java.util.HashSet;
import java.util.LinkedList; import java.util.LinkedList;
import java.util.List; import java.util.List;
import java.util.Map;
import java.util.Set; import java.util.Set;
import java.util.concurrent.CompletableFuture; import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ExecutionException; import java.util.concurrent.ExecutionException;
import java.util.concurrent.TimeUnit; import java.util.concurrent.TimeUnit;
import java.util.concurrent.TimeoutException; import java.util.concurrent.TimeoutException;
import java.util.stream.Collectors;
import org.kie.api.runtime.KieContainer; import org.kie.api.runtime.KieContainer;
import org.kie.api.runtime.KieSession; import org.kie.api.runtime.KieSession;
@ -30,12 +35,15 @@ import com.iqser.red.service.redaction.v1.server.logger.TrackingAgendaEventListe
import com.iqser.red.service.redaction.v1.server.model.NerEntities; import com.iqser.red.service.redaction.v1.server.model.NerEntities;
import com.iqser.red.service.redaction.v1.server.model.dictionary.Dictionary; import com.iqser.red.service.redaction.v1.server.model.dictionary.Dictionary;
import com.iqser.red.service.redaction.v1.server.model.document.nodes.Document; import com.iqser.red.service.redaction.v1.server.model.document.nodes.Document;
import com.iqser.red.service.redaction.v1.server.model.document.nodes.NodeType;
import com.iqser.red.service.redaction.v1.server.model.document.nodes.SemanticNode; import com.iqser.red.service.redaction.v1.server.model.document.nodes.SemanticNode;
import com.iqser.red.service.redaction.v1.server.model.document.nodes.SuperSection;
import com.iqser.red.service.redaction.v1.server.service.ManualChangesApplicationService; import com.iqser.red.service.redaction.v1.server.service.ManualChangesApplicationService;
import com.iqser.red.service.redaction.v1.server.service.document.EntityCreationService; import com.iqser.red.service.redaction.v1.server.service.document.EntityCreationService;
import com.iqser.red.service.redaction.v1.server.service.document.EntityEnrichmentService; import com.iqser.red.service.redaction.v1.server.service.document.EntityEnrichmentService;
import com.iqser.red.service.redaction.v1.server.service.websocket.WebSocketService; import com.iqser.red.service.redaction.v1.server.service.websocket.WebSocketService;
import com.iqser.red.service.redaction.v1.server.utils.exception.DroolsTimeoutException; import com.iqser.red.service.redaction.v1.server.utils.exception.DroolsTimeoutException;
import com.knecon.fforesight.tenantcommons.TenantContext;
import io.micrometer.core.annotation.Timed; import io.micrometer.core.annotation.Timed;
import io.micrometer.observation.ObservationRegistry; import io.micrometer.observation.ObservationRegistry;
@ -93,12 +101,32 @@ public class EntityDroolsExecutionService {
addNumberOfPagesAndSectionsToAnalyseToTrace(document.getNumberOfPages(), sectionsToAnalyze.size()); addNumberOfPagesAndSectionsToAnalyseToTrace(document.getNumberOfPages(), sectionsToAnalyze.size());
KieSession kieSession = kieContainer.newKieSession(); List<SuperSection> superSections = document.streamChildrenOfType(NodeType.SUPER_SECTION)
.map(SuperSection.class::cast)
.toList();
Set<SemanticNode> nodesInKieSession = sectionsToAnalyze.size() == document.streamAllSubNodes() Set<SemanticNode> nodesInKieSession = sectionsToAnalyze.size() == document.streamAllSubNodes()
.count() ? Collections.emptySet() : buildSet(sectionsToAnalyze, document); .count() ? Collections.emptySet() : buildSet(sectionsToAnalyze, document);
EntityCreationService entityCreationService = new EntityCreationService(entityEnrichmentService, kieSession, nodesInKieSession); List<CompletableFuture<SuperSectionResult>> futures = new ArrayList<>();
String tenantId = TenantContext.getTenantId();
superSections.parallelStream()
.forEach(superSection -> {
Set<SemanticNode> nodesInKieSessionInSuperSection = nodesInKieSession.stream()
.filter(node -> !node.getTreeId().isEmpty() && node.getTreeId()
.get(0)
.equals(superSection.getTreeId()
.get(0)))
.collect(Collectors.toSet());
CompletableFuture<SuperSectionResult> future = CompletableFuture.supplyAsync(() -> {
TenantContext.setTenantId(tenantId);
KieSession kieSession = kieContainer.newKieSession();
try {
EntityCreationService entityCreationService = new EntityCreationService(entityEnrichmentService, kieSession, nodesInKieSessionInSuperSection);
RulesLogger logger = new RulesLogger(webSocketService, context); RulesLogger logger = new RulesLogger(webSocketService, context);
if (settings.isDroolsDebug()) { if (settings.isDroolsDebug()) {
logger.enableAgendaTracking(); logger.enableAgendaTracking();
@ -118,12 +146,12 @@ public class EntityDroolsExecutionService {
kieSession.insert(document); kieSession.insert(document);
document.getEntities() superSection.getEntities()
.forEach(kieSession::insert); .forEach(kieSession::insert);
sectionsToAnalyze.forEach(kieSession::insert); nodesInKieSessionInSuperSection.forEach(kieSession::insert);
sectionsToAnalyze.stream() nodesInKieSessionInSuperSection.stream()
.flatMap(SemanticNode::streamAllSubNodes) .flatMap(SemanticNode::streamAllSubNodes)
.forEach(kieSession::insert); .forEach(kieSession::insert);
@ -145,31 +173,57 @@ public class EntityDroolsExecutionService {
kieSession.getAgenda().getAgendaGroup("LOCAL_DICTIONARY_ADDS").setFocus(); kieSession.getAgenda().getAgendaGroup("LOCAL_DICTIONARY_ADDS").setFocus();
CompletableFuture<Void> completableFuture = CompletableFuture.supplyAsync(() -> {
System.out.println("kieSession.getFactCount(): " + kieSession.getFactCount());
kieSession.fireAllRules(); kieSession.fireAllRules();
return null;
List<FileAttribute> resultingFileAttributes = getFileAttributes(kieSession);
kieSession.dispose();
return new SuperSectionResult(resultingFileAttributes);
} catch (Exception e) {
kieSession.dispose();
throw new RuntimeException(e);
}
});
futures.add(future);
}); });
List<FileAttribute> resultingFileAttributes = new ArrayList<>();
futures.parallelStream().forEach(future -> {
var start = System.currentTimeMillis();
try { try {
completableFuture.get(settings.getDroolsExecutionTimeoutSecs(document.getNumberOfPages()), TimeUnit.SECONDS); SuperSectionResult result = future.get(settings.getDroolsExecutionTimeoutSecs(document.getNumberOfPages()), TimeUnit.SECONDS);
addOrUpdate(resultingFileAttributes, result.getFileAttributes());
} catch (ExecutionException e) { } catch (ExecutionException e) {
logger.error(e, "Exception during rule execution"); Throwable cause = e.getCause();
kieSession.dispose(); if (cause instanceof TimeoutException) {
if (e.getCause() instanceof TimeoutException) { throw new DroolsTimeoutException(String.format("The file %s caused a timeout", context.getFileId()), cause, false, RuleFileType.ENTITY);
throw new DroolsTimeoutException(String.format("The file %s caused a timeout",context.getFileId()), e, false, RuleFileType.ENTITY);
} }
throw new RuntimeException(e); throw new RuntimeException(cause);
} catch (InterruptedException e) { } catch (InterruptedException e) {
logger.error(e, "Exception during rule execution");
kieSession.dispose();
throw new RuntimeException(e); throw new RuntimeException(e);
} catch (TimeoutException e) { } catch (TimeoutException e) {
throw new DroolsTimeoutException(String.format("The file %s caused a timeout", context.getFileId()), e, false, RuleFileType.ENTITY); throw new DroolsTimeoutException(String.format("The file %s caused a timeout", context.getFileId()), e, false, RuleFileType.ENTITY);
} }
System.out.printf("Total time in %s : %d ms\n", future, System.currentTimeMillis() - start);
});
List<FileAttribute> resultingFileAttributes = getFileAttributes(kieSession); addOrUpdate(fileAttributes, new ArrayList<>(resultingFileAttributes));
kieSession.dispose();
return resultingFileAttributes; return new ArrayList<>(resultingFileAttributes);
}
private static void addOrUpdate(List<FileAttribute> fileAttributes, List<FileAttribute> resultingFileAttributes) {
for (FileAttribute resultingFileAttribute : resultingFileAttributes) {
fileAttributes.removeIf(fa -> fa.getLabel().equals(resultingFileAttribute.getLabel()));
fileAttributes.add(resultingFileAttribute);
}
} }
@ -213,4 +267,23 @@ public class EntityDroolsExecutionService {
.anyMatch(global -> global.getName().equals(globalName)); .anyMatch(global -> global.getName().equals(globalName));
} }
private static class SuperSectionResult {
private final List<FileAttribute> fileAttributes;
public SuperSectionResult(List<FileAttribute> fileAttributes) {
this.fileAttributes = fileAttributes;
}
public List<FileAttribute> getFileAttributes() {
return fileAttributes;
}
}
} }