diff --git a/.drone.yml b/.drone.yml index 3be19cd955a..29a5386056b 100644 --- a/.drone.yml +++ b/.drone.yml @@ -15,6 +15,7 @@ steps: MINIO_ENDPOINT: minio MINIO_PORT: 9000 ENGINE_URL: http://elastic:9200 + OPENBAS_RABBITMQ_HOSTNAME: rabbitmq commands: - mvn spotless:check - sleep 60 @@ -53,6 +54,7 @@ steps: OPENAEV_ADMIN_PASSWORD: admin OPENAEV_ADMIN_TOKEN: 0d17ce9a-f3a8-4c6d-9721-c98dc3dc023f SPRING_PROFILES_ACTIVE: ci + OPENBAS_RABBITMQ_HOSTNAME: rabbitmq-e2e commands: - apt update && apt install -y gnupg - curl -sS https://dl.yarnpkg.com/debian/pubkey.gpg | apt-key add - @@ -177,6 +179,13 @@ services: POSTGRES_USER: openaev POSTGRES_PASSWORD: openaev POSTGRES_DB: openaev + commands: + - docker-entrypoint.sh -c 'max_connections=500' + - name: rabbitmq + image: rabbitmq:4.1-management + environment: + RABBITMQ_DEFAULT_USER: guest + RABBITMQ_DEFAULT_PASS: guest - name: minio-e2e image: minio/minio:RELEASE.2025-06-13T11-33-47Z environment: @@ -189,6 +198,13 @@ services: POSTGRES_USER: openaev POSTGRES_PASSWORD: openaev POSTGRES_DB: openaev + commands: + - docker-entrypoint.sh -c 'max_connections=500' + - name: rabbitmq-e2e + image: rabbitmq:4.1-management + environment: + RABBITMQ_DEFAULT_USER: guest + RABBITMQ_DEFAULT_PASS: guest - name: elastic image: docker.elastic.co/elasticsearch/elasticsearch:8.18.3 environment: diff --git a/openaev-api/src/main/java/io/openaev/config/CachingConfig.java b/openaev-api/src/main/java/io/openaev/config/CachingConfig.java index 00efe70ee7a..c11aaea29a4 100644 --- a/openaev-api/src/main/java/io/openaev/config/CachingConfig.java +++ b/openaev-api/src/main/java/io/openaev/config/CachingConfig.java @@ -2,23 +2,41 @@ import com.github.benmanes.caffeine.cache.Caffeine; import java.time.Duration; +import lombok.extern.slf4j.Slf4j; import org.springframework.cache.CacheManager; +import org.springframework.cache.annotation.CacheEvict; import org.springframework.cache.annotation.EnableCaching; import org.springframework.cache.caffeine.CaffeineCacheManager; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; +import org.springframework.scheduling.annotation.Scheduled; @Configuration @EnableCaching +@Slf4j public class CachingConfig { @Bean public CacheManager cacheManager() { - CaffeineCacheManager cacheManager = new CaffeineCacheManager("license"); + /** + * Creating some cache : - license is for the EE license - global for global settings that do + * not need to be fetched from the DB everytime we need it (like features flags) - adminUsers is + * a low retention cache for users that are admin. This is useful when receiving a lot of calls. + * Execution traces for instance can receive several thousands a sec and not fetching the user + * everytime helps for the RBAC + */ + CaffeineCacheManager cacheManager = new CaffeineCacheManager("license", "global", "adminUsers"); cacheManager.setCaffeine( Caffeine.newBuilder().expireAfterWrite(Duration.ofDays(1)).maximumSize(100)); return cacheManager; } + + /** Emptying the cache every second to avoid old data on the admin users being persisted */ + @CacheEvict(value = "adminUsers", allEntries = true) + @Scheduled(fixedRateString = "1000") + public void emptyAdminUsersCache() { + log.info("emptying admin users cache"); + } } diff --git a/openaev-api/src/main/java/io/openaev/config/ThreadPoolTaskSchedulerConfig.java b/openaev-api/src/main/java/io/openaev/config/ThreadPoolTaskSchedulerConfig.java index d3667d3ded0..9b4cc860802 100644 --- a/openaev-api/src/main/java/io/openaev/config/ThreadPoolTaskSchedulerConfig.java +++ b/openaev-api/src/main/java/io/openaev/config/ThreadPoolTaskSchedulerConfig.java @@ -1,10 +1,15 @@ package io.openaev.config; +import java.util.concurrent.Executor; +import java.util.concurrent.ThreadPoolExecutor; +import lombok.extern.slf4j.Slf4j; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; +import org.springframework.scheduling.concurrent.ThreadPoolTaskExecutor; import org.springframework.scheduling.concurrent.ThreadPoolTaskScheduler; @Configuration +@Slf4j public class ThreadPoolTaskSchedulerConfig { @Bean @@ -12,6 +17,26 @@ public ThreadPoolTaskScheduler threadPoolTaskScheduler() { ThreadPoolTaskScheduler threadPoolTaskScheduler = new ThreadPoolTaskScheduler(); threadPoolTaskScheduler.setPoolSize(20); threadPoolTaskScheduler.setThreadNamePrefix("ThreadPoolTaskScheduler"); + threadPoolTaskScheduler.setErrorHandler( + t -> log.error("Error during scheduled task : {}", t.getMessage(), t)); return threadPoolTaskScheduler; } + + /** Dedicated executor for stream events */ + @Bean(name = "streamExecutor") + public Executor streamExecutor() { + ThreadPoolTaskExecutor executor = new ThreadPoolTaskExecutor(); + executor.setCorePoolSize(5); + executor.setMaxPoolSize(10); + executor.setQueueCapacity(100); + executor.setThreadNamePrefix("Stream-"); + + // If we have more event to deal with than the available size in the waiting queue, we discard + // the oldest to prevent overloading the stream. This also helps a little preventing + // overloading the tab of a user connected when having a lot of events + executor.setRejectedExecutionHandler(new ThreadPoolExecutor.DiscardOldestPolicy()); + + executor.initialize(); + return executor; + } } diff --git a/openaev-api/src/main/java/io/openaev/migration/V4_53__Convert_expectations_to_jsonb.java b/openaev-api/src/main/java/io/openaev/migration/V4_53__Convert_expectations_to_jsonb.java new file mode 100644 index 00000000000..7602ba2758a --- /dev/null +++ b/openaev-api/src/main/java/io/openaev/migration/V4_53__Convert_expectations_to_jsonb.java @@ -0,0 +1,57 @@ +package io.openaev.migration; + +import java.sql.Statement; +import org.flywaydb.core.api.migration.BaseJavaMigration; +import org.flywaydb.core.api.migration.Context; +import org.springframework.stereotype.Component; + +@Component +public class V4_53__Convert_expectations_to_jsonb extends BaseJavaMigration { + + @Override + public boolean canExecuteInTransaction() { + return false; + } + + @Override + public void migrate(Context context) throws Exception { + try (Statement select = context.getConnection().createStatement()) { + select.execute( + """ + ALTER TABLE injects_expectations + ALTER COLUMN inject_expectation_signatures + TYPE jsonb + USING inject_expectation_signatures::jsonb; + -- Transform Foreign Key to deferrable key + ALTER TABLE execution_traces + DROP CONSTRAINT execution_traces_execution_agent_id_fkey, + DROP CONSTRAINT execution_traces_execution_inject_status_id_fkey, + DROP CONSTRAINT execution_traces_execution_inject_test_status_id_fkey; + + ALTER TABLE execution_traces + ADD CONSTRAINT execution_traces_execution_inject_status_id_fkey + FOREIGN KEY (execution_inject_status_id) + REFERENCES injects_statuses(status_id) + ON DELETE CASCADE + DEFERRABLE INITIALLY DEFERRED, + + ADD CONSTRAINT execution_traces_execution_inject_test_status_id_fkey + FOREIGN KEY (execution_inject_test_status_id) + REFERENCES injects_tests_statuses(status_id) + ON DELETE CASCADE + DEFERRABLE INITIALLY DEFERRED, + + ADD CONSTRAINT execution_traces_execution_agent_id_fkey + FOREIGN KEY (execution_agent_id) + REFERENCES agents(agent_id) + ON DELETE CASCADE + DEFERRABLE INITIALLY DEFERRED; + """); + select.execute( + """ + CREATE INDEX CONCURRENTLY idx_injects_expectations_inject_agent + ON injects_expectations(inject_id, agent_id); + """); + } + } +} diff --git a/openaev-api/src/main/java/io/openaev/rest/finding/FindingService.java b/openaev-api/src/main/java/io/openaev/rest/finding/FindingService.java index b8eee82d75c..8ca3d7718c6 100644 --- a/openaev-api/src/main/java/io/openaev/rest/finding/FindingService.java +++ b/openaev-api/src/main/java/io/openaev/rest/finding/FindingService.java @@ -20,7 +20,6 @@ import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import org.jetbrains.annotations.NotNull; -import org.springframework.dao.DataIntegrityViolationException; import org.springframework.stereotype.Service; import org.springframework.transaction.annotation.Transactional; @@ -91,71 +90,25 @@ public void deleteFinding(@NotNull final String id) { */ public void buildFinding( Inject inject, Asset asset, ContractOutputElement contractOutputElement, String finalValue) { - try { - Optional optionalFinding = - findingRepository.findByInjectIdAndValueAndTypeAndKey( - inject.getId(), - finalValue, - contractOutputElement.getType(), - contractOutputElement.getKey()); - - Finding finding = - optionalFinding.orElseGet( - () -> { - Finding newFinding = new Finding(); - newFinding.setInject(inject); - newFinding.setField(contractOutputElement.getKey()); - newFinding.setType(contractOutputElement.getType()); - newFinding.setValue(finalValue); - newFinding.setName(contractOutputElement.getName()); - newFinding.setTags(new HashSet<>(contractOutputElement.getTags())); - return newFinding; - }); - - boolean isNewAsset = - finding.getAssets().stream().noneMatch(a -> a.getId().equals(asset.getId())); - - if (isNewAsset) { - finding.getAssets().add(asset); - } - - if (optionalFinding.isEmpty() || isNewAsset) { - findingRepository.save(finding); - } - - } catch (DataIntegrityViolationException ex) { - log.info( - String.format( - "Race condition: finding already exists. Retrying ... %s ", ex.getMessage()), - ex); - // Re-fetch and try to add the asset - handleRaceCondition(inject, asset, contractOutputElement, finalValue); - } - } - - private void handleRaceCondition( - Inject inject, Asset asset, ContractOutputElement contractOutputElement, String finalValue) { - Optional retryFinding = - findingRepository.findByInjectIdAndValueAndTypeAndKey( - inject.getId(), - finalValue, - contractOutputElement.getType(), - contractOutputElement.getKey()); - - if (retryFinding.isPresent()) { - Finding existingFinding = retryFinding.get(); - boolean isNewAsset = - existingFinding.getAssets().stream().noneMatch(a -> a.getId().equals(asset.getId())); - if (isNewAsset) { - existingFinding.getAssets().add(asset); - findingRepository.save(existingFinding); - } - } else { - log.warn("Retry failed: Finding still not found after race condition."); - } + String[] tagIds = + contractOutputElement.getTags().isEmpty() + ? new String[0] + : contractOutputElement.getTags().stream().map(Tag::getId).toArray(String[]::new); + + // Save or update the finding and add or update the list of assets and/or tags + findingRepository.saveCompleteFinding( + contractOutputElement.getKey(), + contractOutputElement.getType().name(), + finalValue, + new String[0], + inject.getId(), + contractOutputElement.getName(), + asset.getId(), + tagIds); } - // -- Extract findings from strctured output : Here we compute the findings from structured output + // -- Extract findings from structured output : Here we compute the findings from structured + // output // from ExecutionInjectInput sent by injectors // This structured output is generated based on injectorcontract where we can find the node // Outputs and with that the injector generate this structure output-- diff --git a/openaev-api/src/main/java/io/openaev/rest/helper/queue/BatchQueueService.java b/openaev-api/src/main/java/io/openaev/rest/helper/queue/BatchQueueService.java new file mode 100644 index 00000000000..9abe406f031 --- /dev/null +++ b/openaev-api/src/main/java/io/openaev/rest/helper/queue/BatchQueueService.java @@ -0,0 +1,395 @@ +package io.openaev.rest.helper.queue; + +import com.fasterxml.jackson.databind.ObjectMapper; +import com.rabbitmq.client.*; +import io.openaev.config.QueueConfig; +import io.openaev.config.RabbitmqConfig; +import jakarta.annotation.PreDestroy; +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.util.*; +import java.util.concurrent.*; +import java.util.concurrent.atomic.AtomicBoolean; +import lombok.extern.slf4j.Slf4j; + +@Slf4j +public class BatchQueueService { + + private final Class clazz; + private final QueueExecution queueExecution; + + public static final String ROUTING_KEY = "_push_routing_%s"; + public static final String EXCHANGE_KEY = "_amqp.%s.exchange"; + public static final String QUEUE_NAME = "_execution_%s"; + + protected ObjectMapper mapper; + + private final RabbitmqConfig rabbitmqConfig; + + private Connection connection; + private List publisherChannels = new ArrayList<>(); + private final String routingKey; + private final String exchangeName; + private final String queueName; + + private final Map> queue; + + private final Map deliveryTable = new ConcurrentHashMap<>(); + + private final QueueConfig queueConfig; + private final ScheduledExecutorService reconnectionExecutor; + private final ShutdownListener shutdownListener; + + private final List consumerChannels = new ArrayList<>(); + private final Map insertInProgress = new HashMap<>(); + private final ExecutorService executor; + + /** + * Public constructor of the BatchQueueService + * + * @param clazz the class of element that will be processed + * @param queueExecution the method to handle a list of the class element + * @param rabbitmqConfig the rabbitmq config object + * @param mapper the mapper to use + * @param queueConfig the queue config to use + * @throws IOException In case of issue when communicating with rabbitMQ + * @throws TimeoutException In case of a non responding rabbitMQ + */ + public BatchQueueService( + Class clazz, + QueueExecution queueExecution, + RabbitmqConfig rabbitmqConfig, + ObjectMapper mapper, + QueueConfig queueConfig) + throws IOException, TimeoutException { + this.clazz = clazz; + this.queueExecution = queueExecution; + this.mapper = mapper; + this.queueConfig = queueConfig; + this.rabbitmqConfig = rabbitmqConfig; + + executor = Executors.newFixedThreadPool(queueConfig.getWorkerNumber()); + shutdownListener = this::handleConnectionShutdown; + exchangeName = + rabbitmqConfig.getPrefix() + + String.format(BatchQueueService.EXCHANGE_KEY, queueConfig.getQueueName()); + routingKey = + rabbitmqConfig.getPrefix() + + String.format(BatchQueueService.ROUTING_KEY, queueConfig.getQueueName()); + queueName = + rabbitmqConfig.getPrefix() + + String.format(BatchQueueService.QUEUE_NAME, queueConfig.getQueueName()); + + // The queue that will contain the object we need to process + queue = new HashMap<>(); + for (int i = 0; i < queueConfig.getWorkerNumber(); i++) { + queue.put(i, new LinkedBlockingQueue<>()); + } + + establishConnection(); + + // A scheduler to handle batches that did not reached the critical mass + ScheduledExecutorService scheduledExecutor = Executors.newSingleThreadScheduledExecutor(); + scheduledExecutor.scheduleAtFixedRate( + () -> queue.keySet().forEach(this::processBufferedBatch), + this.queueConfig.getWorkerFrequency(), + this.queueConfig.getWorkerFrequency(), + TimeUnit.MILLISECONDS); + + // Reconnection executor that we will start if we ever lose connection + this.reconnectionExecutor = Executors.newScheduledThreadPool(1); + } + + /** + * Method to establish connection with rabbitMQ + * + * @throws IOException in case of issue while connecting to the server + * @throws TimeoutException in case of issue while connecting to the server + */ + private void establishConnection() throws IOException, TimeoutException { + // Init a Connection factory + ConnectionFactory factory = new ConnectionFactory(); + factory.setHost(rabbitmqConfig.getHostname()); + factory.setPort(rabbitmqConfig.getPort()); + factory.setUsername(rabbitmqConfig.getUser()); + factory.setPassword(rabbitmqConfig.getPass()); + factory.setVirtualHost(rabbitmqConfig.getVhost()); + factory.setAutomaticRecoveryEnabled(false); + factory.setNetworkRecoveryInterval(5000); + factory.setRequestedHeartbeat(30); + factory.setConnectionTimeout(10000); + factory.setSharedExecutor( + Executors.newFixedThreadPool( + queueConfig.getConsumerNumber() + queueConfig.getPublisherNumber())); + + connection = factory.newConnection(); + + // Handle shutdown + connection.addShutdownListener(shutdownListener); + + // Create consumers that will handle the processing + createChannels(); + } + + /** + * Creates a consumer for the queue + * + * @throws IOException In case of issue when communicating with rabbitMQ + */ + private void createChannels() throws IOException { + try { + for (int i = 0; i < queueConfig.getPublisherNumber(); ++i) { + // Creation of the channels, exchange and queue + Channel publisherChannel = connection.createChannel(); + publisherChannel.basicQos(queueConfig.getPublisherQos()); // Per publisher limit + publisherChannel.exchangeDeclare(exchangeName, "topic", true); + Map arguments = new HashMap<>(); + arguments.put("x-queue-type", "quorum"); + publisherChannel.queueDeclare(queueName, true, false, false, arguments); + publisherChannel.queueBind(queueName, exchangeName, routingKey); + publisherChannels.add(publisherChannel); + } + + consumerChannels.clear(); + + for (int i = 0; i < queueConfig.getConsumerNumber(); ++i) { + Channel consumerChannel = connection.createChannel(); + consumerChannels.add(consumerChannel); + consumerChannel.basicQos(queueConfig.getConsumerQos()); + + // What to do when a message is consumed + DeliverCallback deliverCallback = + (consumerTag, delivery) -> { + // We get the object to process + String message = new String(delivery.getBody(), StandardCharsets.UTF_8); + log.trace("Received message from queue {} : '{}'", queueName, message); + + // Unmarshalling of our object and setting it in the queue for processing + T element = mapper.readValue(message, clazz); + int elementKey = groupByKey(element); + queue + .computeIfAbsent(elementKey, integer -> new LinkedBlockingQueue<>()) + .add(element); + + // Add the message and delivery tag into a hashmap that will allow us to ack when + // we've inserted in base + deliveryTable.put( + element, + DeliveryContext.builder() + .tag(delivery.getEnvelope().getDeliveryTag()) + .deliveryChannel(consumerChannel) + .build()); + + // If we reach a critical mass, we take care of it immediately + if (queue.get(elementKey).size() > this.queueConfig.getMaxSize()) { + processBufferedBatch(elementKey); + } + }; + + CancelCallback cancelCallback = + consumerTag -> log.warn("Consumer {} was cancelled", consumerTag); + + // Setting up the consumer itself + consumerChannel.basicConsume( + queueName, + false, + String.format("consumer-%s-%d", queueConfig.getQueueName(), i), + false, + false, + null, + deliverCallback, + cancelCallback); + } + } catch (IOException e) { + log.error("Error creating consumer: {}", e.getMessage(), e); + throw e; + } + } + + /** + * Handle the connection shutdown + * + * @param cause the cause of the shutdown + */ + private void handleConnectionShutdown(ShutdownSignalException cause) { + // If we're just closing openaev, all is good + if (cause.isInitiatedByApplication()) { + log.info("Connection shut down by application"); + return; + } + + // Otherwise, we lost the connection to the server + log.error("Connection lost unexpectedly: {}", cause.getMessage(), cause); + connection.removeShutdownListener(shutdownListener); + + // Start trying to reconnect + reconnectionExecutor.schedule(this::attemptReconnection, 10, TimeUnit.SECONDS); + } + + /** Reconnection attempt */ + private void attemptReconnection() { + log.info("Attempting RabbitMQ reconnection"); + + try { + // Close the resources + closeResources(); + + // Trying to reestablish connection + establishConnection(); + + log.info("Reconnection successful"); + + } catch (Exception e) { + log.error(String.format("Reconnection attempt failed: %s", e.getMessage()), e); + + // We failed. We schedule a new try ... + reconnectionExecutor.schedule(this::attemptReconnection, 10, TimeUnit.SECONDS); + } + } + + /** Close the resources */ + private void closeResources() throws IOException, TimeoutException { + try { + // Close consumer channels + for (Channel channel : consumerChannels) { + if (channel != null && channel.isOpen()) { + channel.close(); + } + } + + // Closing the publishing channel + for (Channel channel : publisherChannels) { + if (channel != null && channel.isOpen()) { + channel.close(); + } + } + + // Close the connection if it's open + if (connection != null && connection.isOpen()) { + connection.close(); + } + } catch (Exception e) { + log.warn("Error closing resources: {}", e.getMessage()); + throw e; + } finally { + publisherChannels.clear(); + consumerChannels.clear(); + } + } + + @PreDestroy + public void stop() throws IOException, TimeoutException { + closeResources(); + } + + /** + * Process messages in the queue buffer. It will only process as many messages as what's + * configures in openbas.queue-config..max-size + */ + public void processBufferedBatch(int workerId) { + if (insertInProgress + .computeIfAbsent(workerId, integer -> new AtomicBoolean(false)) + .compareAndSet(false, true)) { + executor.execute( + () -> { + do { + // Draining the queue into the list with a max size + List currentBatch = new ArrayList<>(); + queue.get(workerId).drainTo(currentBatch); + + // If the list is not empty, we process it + List processedElement = new ArrayList<>(); + if (!currentBatch.isEmpty()) { + log.info("Processing batch of {}", currentBatch.size()); + try { + processedElement.addAll(queueExecution.perform(currentBatch)); + } catch (Exception e) { + log.error("Error processing batch - Error during ingestion", e); + } + } + + // Sending Ack for all the processed element in the batch + for (T element : processedElement) { + try { + DeliveryContext elementToAck = deliveryTable.remove(element); + if (elementToAck != null) { + elementToAck.getDeliveryChannel().basicAck(elementToAck.getTag(), false); + currentBatch.remove(element); + } + } catch (IOException e) { + log.error( + String.format( + "Error processing batch - Cannot Ack the message: %s", e.getMessage()), + e); + } + } + + // The elements that were not successfully processed are rejected + for (T element : currentBatch) { + try { + DeliveryContext elementToReject = deliveryTable.remove(element); + if (elementToReject != null) { + // To avoid having elements that are not properly processed but can never be, + // we're not requeueing them. + elementToReject + .getDeliveryChannel() + .basicReject(elementToReject.getTag(), false); + } + } catch (IOException e) { + log.error( + String.format( + "Error processing batch - Cannot Nack the message: %s", e.getMessage()), + e); + } + } + } while (queue.get(workerId).size() > (queueConfig.getMaxSize() * 0.75)); + insertInProgress.get(workerId).set(false); + }); + } + } + + /** + * Publish a stringified object of type T into the queue + * + * @param element the T object to publish + * @throws IOException in case of error during the publish + */ + public void publish(T element) throws IOException { + try { + publisherChannels + .get(element.hashCode() % publisherChannels.size()) + .basicPublish( + exchangeName, routingKey, null, mapper.writeValueAsString(element).getBytes()); + } catch (IOException e) { + log.error(String.format("Error publishing batch: %s", e.getMessage()), e); + throw e; + } + } + + /** + * Purge a queue + * + * @throws IOException in case of error during the publish + */ + public void forcePurge() throws IOException { + try { + publisherChannels.getFirst().queuePurge(queueName); + } catch (IOException e) { + log.error(String.format("Error publishing batch: %s", e.getMessage()), e); + throw e; + } + } + + /** + * Get the id of the worker depending on the key of the element and the number of workers + * + * @param element the element that we need to process + * @return the id of the worker + */ + private int groupByKey(T element) { + if (element.getUniqueElementKey() != null && !element.getUniqueElementKey().isEmpty()) { + return element.getUniqueElementKey().hashCode() % queueConfig.getWorkerNumber(); + } + return 0; + } +} diff --git a/openaev-api/src/main/java/io/openaev/rest/helper/queue/DeliveryContext.java b/openaev-api/src/main/java/io/openaev/rest/helper/queue/DeliveryContext.java new file mode 100644 index 00000000000..87e265d0711 --- /dev/null +++ b/openaev-api/src/main/java/io/openaev/rest/helper/queue/DeliveryContext.java @@ -0,0 +1,14 @@ +package io.openaev.rest.helper.queue; + +import com.rabbitmq.client.Channel; +import lombok.Builder; +import lombok.Data; + +@Data +@Builder +public class DeliveryContext { + + private long tag; + + private Channel deliveryChannel; +} diff --git a/openaev-api/src/main/java/io/openaev/rest/helper/queue/QueueExecution.java b/openaev-api/src/main/java/io/openaev/rest/helper/queue/QueueExecution.java new file mode 100644 index 00000000000..984e1976fe0 --- /dev/null +++ b/openaev-api/src/main/java/io/openaev/rest/helper/queue/QueueExecution.java @@ -0,0 +1,13 @@ +package io.openaev.rest.helper.queue; + +import java.util.List; + +public interface QueueExecution { + /** + * Function that process a list of elements and return the list of successfully processed elements + * + * @param elements the elements to process + * @return the successfully processed elements + */ + List perform(List elements); +} diff --git a/openaev-api/src/main/java/io/openaev/rest/helper/queue/Queueable.java b/openaev-api/src/main/java/io/openaev/rest/helper/queue/Queueable.java new file mode 100644 index 00000000000..d9d35e574d2 --- /dev/null +++ b/openaev-api/src/main/java/io/openaev/rest/helper/queue/Queueable.java @@ -0,0 +1,5 @@ +package io.openaev.rest.helper.queue; + +public interface Queueable { + String getUniqueElementKey(); +} diff --git a/openaev-api/src/main/java/io/openaev/rest/helper/queue/executor/BatchExecutionTraceExecutor.java b/openaev-api/src/main/java/io/openaev/rest/helper/queue/executor/BatchExecutionTraceExecutor.java new file mode 100644 index 00000000000..50a60f54811 --- /dev/null +++ b/openaev-api/src/main/java/io/openaev/rest/helper/queue/executor/BatchExecutionTraceExecutor.java @@ -0,0 +1,19 @@ +package io.openaev.rest.helper.queue.executor; + +import io.openaev.rest.inject.form.InjectExecutionCallback; +import io.openaev.rest.inject.service.BatchingInjectStatusService; +import java.util.List; +import lombok.RequiredArgsConstructor; +import org.springframework.stereotype.Component; + +@Component +@RequiredArgsConstructor +public class BatchExecutionTraceExecutor { + + private final BatchingInjectStatusService batchingInjectStatusService; + + public List handleInjectExecutionCallbackList( + List injectExecutionCallbacks) { + return batchingInjectStatusService.handleInjectExecutionCallback(injectExecutionCallbacks); + } +} diff --git a/openaev-api/src/main/java/io/openaev/rest/inject/InjectApi.java b/openaev-api/src/main/java/io/openaev/rest/inject/InjectApi.java index ac253e0855a..f2875a11048 100644 --- a/openaev-api/src/main/java/io/openaev/rest/inject/InjectApi.java +++ b/openaev-api/src/main/java/io/openaev/rest/inject/InjectApi.java @@ -3,10 +3,14 @@ import static io.openaev.config.SessionHelper.currentUser; import static io.openaev.helper.StreamHelper.fromIterable; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.google.common.annotations.VisibleForTesting; import io.openaev.aop.LogExecutionTime; import io.openaev.aop.RBAC; import io.openaev.aop.lock.Lock; import io.openaev.aop.lock.LockResourceType; +import io.openaev.config.OpenAEVConfig; +import io.openaev.config.RabbitmqConfig; import io.openaev.database.model.*; import io.openaev.database.raw.RawDocument; import io.openaev.database.repository.ExerciseRepository; @@ -22,13 +26,17 @@ import io.openaev.rest.exception.ElementNotFoundException; import io.openaev.rest.exercise.exports.ExportOptions; import io.openaev.rest.helper.RestBehavior; +import io.openaev.rest.helper.queue.BatchQueueService; +import io.openaev.rest.helper.queue.executor.BatchExecutionTraceExecutor; import io.openaev.rest.inject.form.*; import io.openaev.rest.inject.service.ExecutableInjectService; import io.openaev.rest.inject.service.InjectExecutionService; import io.openaev.rest.inject.service.InjectExportService; import io.openaev.rest.inject.service.InjectService; import io.openaev.rest.payload.form.DetectionRemediationOutput; +import io.openaev.rest.settings.PreviewFeature; import io.openaev.service.InjectImportService; +import io.openaev.service.PreviewFeatureService; import io.openaev.service.UserService; import io.openaev.service.targets.TargetService; import io.openaev.utils.FilterUtilsJpa; @@ -38,15 +46,19 @@ import io.swagger.v3.oas.annotations.Operation; import io.swagger.v3.oas.annotations.responses.ApiResponse; import io.swagger.v3.oas.annotations.responses.ApiResponses; +import jakarta.annotation.PostConstruct; import jakarta.servlet.ServletOutputStream; import jakarta.servlet.http.HttpServletResponse; import jakarta.validation.Valid; import jakarta.validation.constraints.NotBlank; import java.io.IOException; +import java.time.Instant; import java.util.ArrayList; import java.util.List; import java.util.Optional; +import java.util.concurrent.TimeoutException; import lombok.RequiredArgsConstructor; +import lombok.Setter; import lombok.extern.slf4j.Slf4j; import org.apache.commons.lang3.StringUtils; import org.springframework.data.domain.Page; @@ -60,6 +72,7 @@ @Slf4j @RestController @RequiredArgsConstructor +@Setter public class InjectApi extends RestBehavior { public static final String INJECT_URI = "/api/injects"; @@ -79,6 +92,30 @@ public class InjectApi extends RestBehavior { private final UserService userService; private final DocumentService documentService; private final GrantRepository grantRepository; + private final BatchExecutionTraceExecutor batchExecutionTraceExecutor; + + private final RabbitmqConfig rabbitmqConfig; + private final OpenAEVConfig openAEVConfig; + private final ObjectMapper objectMapper; + + private final PreviewFeatureService previewFeatureService; + + // For testing purpose, we add a setter + @Setter private BatchQueueService injectTraceQueueService; + + @PostConstruct + public void init() throws IOException, TimeoutException { + if (openAEVConfig.getQueueConfig().get("inject-trace") != null) { + // Initializing the queue for batching the inject execution trace + injectTraceQueueService = + new BatchQueueService<>( + InjectExecutionCallback.class, + batchExecutionTraceExecutor::handleInjectExecutionCallbackList, + rabbitmqConfig, + objectMapper, + openAEVConfig.getQueueConfig().get("inject-trace")); + } + } // -- INJECTS -- @@ -304,7 +341,8 @@ public Inject injectExecutionReception( actionPerformed = Action.WRITE, resourceType = ResourceType.INJECT) public void injectExecutionCallback( - @PathVariable String injectId, @Valid @RequestBody InjectExecutionInput input) { + @PathVariable String injectId, @Valid @RequestBody InjectExecutionInput input) + throws IOException { injectExecutionCallback(null, injectId, input); } @@ -333,8 +371,23 @@ public void injectExecutionCallback( @PathVariable String agentId, // must allow null because http injector used also this method to work. @PathVariable String injectId, - @Valid @RequestBody InjectExecutionInput input) { - injectExecutionService.handleInjectExecutionCallback(injectId, agentId, input); + @Valid @RequestBody InjectExecutionInput input) + throws IOException { + if (!previewFeatureService.isFeatureEnabled(PreviewFeature.LEGACY_INGESTION_EXECUTION_TRACE) + && injectTraceQueueService != null) { + InjectExecutionCallback injectExecutionCallback = + InjectExecutionCallback.builder() + .injectExecutionInput(input) + .agentId(agentId) + .injectId(injectId) + .emissionDate(Instant.now().toEpochMilli()) + .build(); + + // Publishing the parameters into a queue for later ingestion + injectTraceQueueService.publish(injectExecutionCallback); + } else { + injectExecutionService.handleInjectExecutionCallback(injectId, agentId, input); + } } @GetMapping(INJECT_URI + "/{injectId}/{agentId}/executable-payload") @@ -533,4 +586,9 @@ public List getPayloadDocumentsByInjectIdAndPayloadId( return documentService.documentsForPayload(payloadId); } + + @VisibleForTesting + public BatchQueueService getInjectTraceQueueService() { + return injectTraceQueueService; + } } diff --git a/openaev-api/src/main/java/io/openaev/rest/inject/form/InjectExecutionCallback.java b/openaev-api/src/main/java/io/openaev/rest/inject/form/InjectExecutionCallback.java new file mode 100644 index 00000000000..5e2f65d8784 --- /dev/null +++ b/openaev-api/src/main/java/io/openaev/rest/inject/form/InjectExecutionCallback.java @@ -0,0 +1,43 @@ +package io.openaev.rest.inject.form; + +import com.fasterxml.jackson.annotation.JsonProperty; +import io.openaev.rest.helper.queue.Queueable; +import java.util.UUID; +import lombok.AllArgsConstructor; +import lombok.Builder; +import lombok.Data; +import lombok.NoArgsConstructor; + +@Data +@Builder +@NoArgsConstructor +@AllArgsConstructor +public class InjectExecutionCallback implements Queueable { + + private String id = UUID.randomUUID().toString(); + + @JsonProperty("agent_id") + private String agentId; + + @JsonProperty("inject_id") + private String injectId; + + @JsonProperty("inject_execution_input") + private InjectExecutionInput injectExecutionInput; + + @JsonProperty("execution_emission_date") + private long emissionDate; + + @Override + public boolean equals(Object o) { + if (o instanceof InjectExecutionCallback) { + return id != null && id.equals(((InjectExecutionCallback) o).getId()); + } + return false; + } + + @Override + public String getUniqueElementKey() { + return injectId + agentId; + } +} diff --git a/openaev-api/src/main/java/io/openaev/rest/inject/service/BatchingInjectStatusService.java b/openaev-api/src/main/java/io/openaev/rest/inject/service/BatchingInjectStatusService.java new file mode 100644 index 00000000000..7fa2e4d2615 --- /dev/null +++ b/openaev-api/src/main/java/io/openaev/rest/inject/service/BatchingInjectStatusService.java @@ -0,0 +1,146 @@ +package io.openaev.rest.inject.service; + +import com.fasterxml.jackson.databind.ObjectMapper; +import io.openaev.aop.LogExecutionTime; +import io.openaev.database.model.*; +import io.openaev.database.repository.*; +import io.openaev.rest.exception.ElementNotFoundException; +import io.openaev.rest.inject.form.InjectExecutionAction; +import io.openaev.rest.inject.form.InjectExecutionCallback; +import jakarta.annotation.Resource; +import jakarta.transaction.Transactional; +import java.util.*; +import java.util.function.Function; +import java.util.stream.Collectors; +import java.util.stream.Stream; +import java.util.stream.StreamSupport; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.dao.DataIntegrityViolationException; +import org.springframework.stereotype.Service; + +@RequiredArgsConstructor +@Service +@Slf4j +@Transactional +public class BatchingInjectStatusService { + + private final InjectRepository injectRepository; + private final AgentRepository agentRepository; + private final StructuredOutputUtils structuredOutputUtils; + private final InjectExecutionService injectExecutionService; + + @Resource protected ObjectMapper mapper; + + /** + * Handle the list of inject execution callbacks + * + * @param injectExecutionCallbacks the inject execution callbacks + */ + @LogExecutionTime + @Transactional(Transactional.TxType.REQUIRES_NEW) + public List handleInjectExecutionCallback( + List injectExecutionCallbacks) { + + List successfullyProcessedCallbacks = new ArrayList<>(); + + // Getting all the injects needed all at once + Map mapInjectsById = + injectRepository + .findAllByIdWithExpectations( + injectExecutionCallbacks.stream() + .map(InjectExecutionCallback::getInjectId) + .toList()) + .stream() + .collect(Collectors.toMap(Inject::getId, Function.identity())); + + // Getting all the agents all at once + Map mapAgentsById = + StreamSupport.stream( + agentRepository + .findAllById( + injectExecutionCallbacks.stream() + .map(InjectExecutionCallback::getAgentId) + .toList()) + .spliterator(), + false) + .collect(Collectors.toMap(Agent::getId, Function.identity())); + + // Sorting the inject execution callbacks to make sure we handle them in chronological order + Stream sortedInjectExecutionCallbacks = + injectExecutionCallbacks.stream() + .sorted(Comparator.comparing(InjectExecutionCallback::getEmissionDate)); + + // For each of the callback + sortedInjectExecutionCallbacks.forEach( + callback -> { + Inject inject = null; + + try { + // Get the inject or throw if not found + inject = + Optional.ofNullable(mapInjectsById.get(callback.getInjectId())) + .orElseThrow( + () -> + new ElementNotFoundException( + "Inject not found: " + callback.getInjectId())); + // issue/3550: added this condition to ensure we only update statuses if the inject is + // in a + // coherent state. + // This prevents issues where the PENDING status took more time to persist than it took + // for + // the agent to send the complete action. + // FIXME: At the moment, this whole function is only called by our implant. These + // implant are + // launched with the async value to true, which force the implant to go from EXECUTING + // to + // PENDING, before going to EXECUTED. + // So if in the future, this function is called to update a synchronous inject, we will + // need + // to find a way to get the async boolean somehow and add it to this condition. + if (callback + .getInjectExecutionInput() + .getAction() + .equals(InjectExecutionAction.complete) + && (inject.getStatus().isEmpty() + || !inject.getStatus().get().getName().equals(ExecutionStatus.PENDING))) { + // If we receive a status update with a terminal state status, we must first check + // that the + // current status is in the PENDING state + log.warn( + String.format( + "Received a complete action for inject %s with status %s, but current status is not PENDING", + callback.getInjectId(), + inject.getStatus().map(is -> is.getName().toString()).orElse("unknown"))); + throw new DataIntegrityViolationException( + "Cannot complete inject that is not in PENDING state"); + } + // Get the agent or throw if not found + Agent agent = + Optional.ofNullable(mapAgentsById.get(callback.getAgentId())) + .orElseThrow( + () -> + new ElementNotFoundException( + "Agent not found: " + callback.getAgentId())); + + // Extract the output parsers + Set outputParsers = structuredOutputUtils.extractOutputParsers(inject); + + // Process the execution trace + injectExecutionService.processInjectExecution( + inject, agent, callback.getInjectExecutionInput(), outputParsers); + successfullyProcessedCallbacks.add(callback); + } catch (ElementNotFoundException e) { + injectExecutionService.handleInjectExecutionError(inject, e); + successfullyProcessedCallbacks.add(callback); + } catch (Exception e) { + log.warn( + "The was a problem processing the element for the inject {} and agent {}", + callback.getInjectId(), + callback.getAgentId(), + e); + } + }); + return successfullyProcessedCallbacks; + } +} diff --git a/openaev-api/src/main/java/io/openaev/rest/inject/service/InjectExecutionService.java b/openaev-api/src/main/java/io/openaev/rest/inject/service/InjectExecutionService.java index e9d104c3488..55edb301a1f 100644 --- a/openaev-api/src/main/java/io/openaev/rest/inject/service/InjectExecutionService.java +++ b/openaev-api/src/main/java/io/openaev/rest/inject/service/InjectExecutionService.java @@ -28,6 +28,7 @@ import lombok.extern.slf4j.Slf4j; import org.springframework.dao.DataIntegrityViolationException; import org.springframework.stereotype.Service; +import org.springframework.transaction.annotation.Transactional; @RequiredArgsConstructor @Service @@ -44,6 +45,7 @@ public class InjectExecutionService { @Resource protected ObjectMapper mapper; + @Transactional public void handleInjectExecutionCallback( String injectId, String agentId, InjectExecutionInput input) { Inject inject = null; @@ -83,7 +85,6 @@ public void handleInjectExecutionCallback( } /** Processes the execution of an inject by updating its status and extracting findings. */ - @VisibleForTesting public void processInjectExecution( Inject inject, @Nullable Agent agent, @@ -237,7 +238,7 @@ private Inject loadInjectOrThrow(String injectId) { .orElseThrow(() -> new ElementNotFoundException("Inject not found: " + injectId)); } - private void handleInjectExecutionError(Inject inject, Exception e) { + public void handleInjectExecutionError(Inject inject, Exception e) { log.error(e.getMessage(), e); if (inject != null) { inject diff --git a/openaev-api/src/main/java/io/openaev/rest/inject/service/InjectStatusService.java b/openaev-api/src/main/java/io/openaev/rest/inject/service/InjectStatusService.java index 25f814b4173..e1830a8b2f2 100644 --- a/openaev-api/src/main/java/io/openaev/rest/inject/service/InjectStatusService.java +++ b/openaev-api/src/main/java/io/openaev/rest/inject/service/InjectStatusService.java @@ -7,6 +7,7 @@ import com.google.common.annotations.VisibleForTesting; import io.openaev.aop.lock.Lock; import io.openaev.aop.lock.LockResourceType; +import io.openaev.database.helper.ExecutionTraceRepositoryHelper; import io.openaev.database.model.*; import io.openaev.database.repository.AgentRepository; import io.openaev.database.repository.InjectRepository; @@ -17,6 +18,7 @@ import io.openaev.rest.inject.form.InjectUpdateStatusInput; import io.openaev.utils.InjectUtils; import jakarta.annotation.Nullable; +import jakarta.persistence.EntityManager; import jakarta.transaction.Transactional; import jakarta.validation.constraints.NotNull; import java.time.Instant; @@ -36,6 +38,9 @@ public class InjectStatusService { private final InjectService injectService; private final InjectUtils injectUtils; private final InjectStatusRepository injectStatusRepository; + private final ExecutionTraceRepositoryHelper executionTraceRepositoryHelper; + + private final EntityManager entityManager; public List findPendingInjectStatusByType(String injectType) { return this.injectStatusRepository.pendingForInjectType(injectType); @@ -179,19 +184,30 @@ public void updateInjectStatus( Agent agent, Inject inject, InjectExecutionInput input, ObjectNode structuredOutput) { InjectStatus injectStatus = inject.getStatus().orElseThrow(ElementNotFoundException::new); + // Creating the Execution Trace ExecutionTrace executionTrace = createExecutionTrace(injectStatus, input, agent, structuredOutput); + // Update the status of the execution trace if needed computeExecutionTraceStatusIfNeeded(injectStatus, executionTrace, agent); injectStatus.addTrace(executionTrace); + // Save the trace using a low level call to the database + String executionTraceId = executionTraceRepositoryHelper.saveExecutionTrace(executionTrace); + executionTrace.setId(executionTraceId); + entityManager.merge(injectStatus); + // If the trace is complete if (executionTrace.getAction().equals(ExecutionTraceAction.COMPLETE) && (agent == null || isAllInjectAgentsExecuted(inject))) { + // We update the status of the inject updateFinalInjectStatus(injectStatus); - log.debug("Successfully updated inject final status: " + inject.getId()); + executionTraceRepositoryHelper.updateInjectUpdateDate( + injectStatus.getInject().getId(), injectStatus.getInject().getUpdatedAt()); + executionTraceRepositoryHelper.updateInjectStatus( + injectStatus.getId(), injectStatus.getName().name(), injectStatus.getTrackingEndDate()); + log.debug("Successfully updated inject final status: {}", inject.getId()); } - injectRepository.save(inject); - log.debug("Successfully updated inject: " + inject.getId()); + log.debug("Successfully updated inject: {}", inject.getId()); } public ExecutionStatus computeStatus(List traces) { diff --git a/openaev-api/src/main/java/io/openaev/rest/settings/PreviewFeature.java b/openaev-api/src/main/java/io/openaev/rest/settings/PreviewFeature.java index 4f024d8b25c..d4f2eaf5ff4 100644 --- a/openaev-api/src/main/java/io/openaev/rest/settings/PreviewFeature.java +++ b/openaev-api/src/main/java/io/openaev/rest/settings/PreviewFeature.java @@ -12,7 +12,8 @@ public enum PreviewFeature { // Reserved for internal use. _RESERVED, - STIX_SECURITY_COVERAGE_FOR_VULNERABILITIES; + STIX_SECURITY_COVERAGE_FOR_VULNERABILITIES, + LEGACY_INGESTION_EXECUTION_TRACE; public static PreviewFeature fromStringIgnoreCase(String str) { for (PreviewFeature feature : PreviewFeature.values()) { diff --git a/openaev-api/src/main/java/io/openaev/rest/stream/StreamApi.java b/openaev-api/src/main/java/io/openaev/rest/stream/StreamApi.java index 9726c851619..08451d07fa7 100644 --- a/openaev-api/src/main/java/io/openaev/rest/stream/StreamApi.java +++ b/openaev-api/src/main/java/io/openaev/rest/stream/StreamApi.java @@ -70,7 +70,7 @@ private void sendStreamEvent(FluxSink flux, BaseEvent event) { private static final EnumSet RESOURCES_STREAM_BLACKLIST = EnumSet.of(ResourceType.VULNERABILITY, ResourceType.PAYLOAD); - @Async + @Async("streamExecutor") @Transactional @TransactionalEventListener public void listenDatabaseUpdate(BaseEvent event) { diff --git a/openaev-api/src/main/java/io/openaev/service/InjectExpectationService.java b/openaev-api/src/main/java/io/openaev/service/InjectExpectationService.java index 112ecb55c82..69171fbae14 100644 --- a/openaev-api/src/main/java/io/openaev/service/InjectExpectationService.java +++ b/openaev-api/src/main/java/io/openaev/service/InjectExpectationService.java @@ -11,6 +11,7 @@ import static io.openaev.utils.inject_expectation_result.InjectExpectationResultUtils.computeScore; import com.fasterxml.jackson.databind.ObjectMapper; +import io.openaev.database.helper.InjectExpectationRepositoryHelper; import io.openaev.database.model.*; import io.openaev.database.repository.InjectExpectationRepository; import io.openaev.database.specification.InjectExpectationSpecification; @@ -55,6 +56,7 @@ public class InjectExpectationService { public static final String PENDING = "Pending"; public static final String COLLECTOR = "collector"; private final InjectExpectationRepository injectExpectationRepository; + private final InjectExpectationRepositoryHelper injectExpectationRepositoryHelper; private final CollectorService collectorService; @Resource private ExpectationPropertiesConfig expectationPropertiesConfig; private final SecurityCoverageSendJobService securityCoverageSendJobService; @@ -527,16 +529,9 @@ private void addDateSignatureToInjectExpectationsByAgent( @NotBlank final String agentId, @NotBlank final Instant date, @NotBlank final String signatureType) { - List injectExpectations = - this.injectExpectationRepository.findAllByInjectAndAgent(injectId, agentId); - - injectExpectations.forEach( - expectation -> { - List signatures = expectation.getSignatures(); - signatures.add(new InjectExpectationSignature(signatureType, date.toString())); - }); - - injectExpectationRepository.saveAll(injectExpectations); + // Insert the signature for all agent and inject in one query + injectExpectationRepositoryHelper.insertSignatureForAgentAndInject( + injectId, agentId, signatureType, date.toString()); } /** diff --git a/openaev-api/src/main/java/io/openaev/service/PreviewFeatureService.java b/openaev-api/src/main/java/io/openaev/service/PreviewFeatureService.java index d91dee772a0..7a71562b080 100644 --- a/openaev-api/src/main/java/io/openaev/service/PreviewFeatureService.java +++ b/openaev-api/src/main/java/io/openaev/service/PreviewFeatureService.java @@ -3,6 +3,7 @@ import io.openaev.rest.settings.PreviewFeature; import java.util.List; import lombok.RequiredArgsConstructor; +import org.springframework.cache.annotation.Cacheable; import org.springframework.stereotype.Service; @Service @@ -10,6 +11,7 @@ public class PreviewFeatureService { private final PlatformSettingsService platformSettingsService; + @Cacheable("global") public boolean isFeatureEnabled(PreviewFeature feature) { List enabledFeatures = platformSettingsService.findSettings().getEnabledDevFeatures(); diff --git a/openaev-api/src/main/java/io/openaev/service/UserService.java b/openaev-api/src/main/java/io/openaev/service/UserService.java index d1390c9609e..8318e57be74 100644 --- a/openaev-api/src/main/java/io/openaev/service/UserService.java +++ b/openaev-api/src/main/java/io/openaev/service/UserService.java @@ -21,7 +21,10 @@ import jakarta.validation.constraints.NotBlank; import jakarta.validation.constraints.NotNull; import java.util.*; +import lombok.RequiredArgsConstructor; import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.cache.Cache; +import org.springframework.cache.CacheManager; import org.springframework.security.core.Authentication; import org.springframework.security.core.GrantedAuthority; import org.springframework.security.core.authority.SimpleGrantedAuthority; @@ -33,6 +36,7 @@ import org.springframework.util.StringUtils; @Service +@RequiredArgsConstructor public class UserService { @Resource private SessionManager sessionManager; private final Argon2PasswordEncoder passwordEncoder = @@ -42,12 +46,19 @@ public class UserService { private TagRepository tagRepository; private GroupRepository groupRepository; private OrganizationRepository organizationRepository; + private CacheManager cacheManager; + private Cache adminCache; @Autowired public void setOrganizationRepository(OrganizationRepository organizationRepository) { this.organizationRepository = organizationRepository; } + @Autowired + public void setCacheManager(CacheManager cacheManager) { + this.cacheManager = cacheManager; + } + @Autowired public void setTagRepository(TagRepository tagRepository) { this.tagRepository = tagRepository; @@ -152,9 +163,35 @@ public List users() { } public User currentUser() { - return this.userRepository - .findById(SessionHelper.currentUser().getId()) - .orElseThrow(() -> new ElementNotFoundException("Current user not found")); + User user; + // If we don't have the cache, we get it + if (adminCache == null) { + adminCache = cacheManager.getCache("adminUsers"); + } + // If the cache is available + if (adminCache != null) { + // We try to check if the user is in the cache + user = adminCache.get(SessionHelper.currentUser().getId(), User.class); + // If not, we get it + if (user == null) { + user = + this.userRepository + .findById(SessionHelper.currentUser().getId()) + .orElseThrow(() -> new ElementNotFoundException("Current user not found")); + + // If the user is admin, we put him in cache + if (user.isAdmin()) { + adminCache.put(SessionHelper.currentUser().getId(), user); + } + } + } else { + // If for some reason, the cache is unavailable, we just get the user and return it + user = + this.userRepository + .findById(SessionHelper.currentUser().getId()) + .orElseThrow(() -> new ElementNotFoundException("Current user not found")); + } + return user; } // endregion diff --git a/openaev-api/src/main/java/io/openaev/utils/ExpectationUtils.java b/openaev-api/src/main/java/io/openaev/utils/ExpectationUtils.java index 2f24d08297e..b67999ebe0e 100644 --- a/openaev-api/src/main/java/io/openaev/utils/ExpectationUtils.java +++ b/openaev-api/src/main/java/io/openaev/utils/ExpectationUtils.java @@ -436,7 +436,11 @@ public static List getExpectationsAgentsForAsset( @NotNull final InjectExpectation injectExpectation) { return injectExpectation.getInject().getExpectations().stream() .filter(ExpectationUtils::isAgentExpectation) - .filter(e -> e.getAsset().getId().equals(injectExpectation.getAsset().getId())) + .filter( + e -> + e.getAsset() != null + && injectExpectation.getAsset() != null + && e.getAsset().getId().equals(injectExpectation.getAsset().getId())) .filter(e -> e.getType().equals(injectExpectation.getType())) .toList(); } diff --git a/openaev-api/src/main/resources/application.properties b/openaev-api/src/main/resources/application.properties index 27e80ca823c..998bd9b1be2 100644 --- a/openaev-api/src/main/resources/application.properties +++ b/openaev-api/src/main/resources/application.properties @@ -35,7 +35,15 @@ spring.datasource.url= spring.datasource.username=openaev # Password: the password for the username set above spring.datasource.password= +spring.jpa.properties.hibernate.order_inserts=true +spring.jpa.properties.hibernate.order_updates=true +spring.jpa.properties.hibernate.jdbc.batch_size=30 +spring.datasource.hikari.maximum-pool-size=20 spring.datasource.hikari.leak-detection-threshold=30000 +spring.datasource.hikari.data-source-properties.cachePrepStmts=true +spring.datasource.hikari.data-source-properties.prepStmtCacheSize=250 +spring.datasource.hikari.data-source-properties.prepStmtCacheSqlLimit=2048 +spring.datasource.hikari.data-source-properties.useServerPrepStmts=true ### ENGINE Configuration # selector can be elk or opensearch @@ -81,11 +89,14 @@ openaev.rabbitmq.management-insecure=true openaev.rabbitmq.trust-store-password= openaev.rabbitmq.trust.store= - - - -#Feature Flags -openaev.enabled-dev-features= +openaev.queue-config.inject-trace.publisher-number=2 +openaev.queue-config.inject-trace.consumer-number=8 +openaev.queue-config.inject-trace.worker-number=8 +openaev.queue-config.inject-trace.worker-frequency=10000 +openaev.queue-config.inject-trace.queue-name=inject-trace +openaev.queue-config.inject-trace.max-size=200 +openaev.queue-config.inject-trace.consumer-qos=1000 +openaev.queue-config.inject-trace.publisher-qos=0 # Web server configuration server.address=0.0.0.0 diff --git a/openaev-api/src/test/java/io/openaev/IntegrationTest.java b/openaev-api/src/test/java/io/openaev/IntegrationTest.java index 9866cd8d09b..d15e2ee66d2 100644 --- a/openaev-api/src/test/java/io/openaev/IntegrationTest.java +++ b/openaev-api/src/test/java/io/openaev/IntegrationTest.java @@ -7,6 +7,7 @@ import io.openaev.utils.fixtures.composers.GrantComposer; import io.openaev.utils.mockUser.TestUserHolder; import io.openaev.utils.mockUser.WithMockUserTestExecutionListener; +import io.openaev.utilstest.RabbitMQTestListener; import io.openaev.utilstest.StartupSnapshotTestListener; import jakarta.persistence.EntityManager; import org.springframework.beans.factory.annotation.Autowired; @@ -18,7 +19,11 @@ @AutoConfigureMockMvc(print = MockMvcPrint.SYSTEM_ERR) @SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT) @TestExecutionListeners( - value = {StartupSnapshotTestListener.class, WithMockUserTestExecutionListener.class}, + value = { + StartupSnapshotTestListener.class, + WithMockUserTestExecutionListener.class, + RabbitMQTestListener.class + }, mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public abstract class IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/config/OpenCTIConfigTest.java b/openaev-api/src/test/java/io/openaev/config/OpenCTIConfigTest.java index f37de03fe28..e912888e33a 100644 --- a/openaev-api/src/test/java/io/openaev/config/OpenCTIConfigTest.java +++ b/openaev-api/src/test/java/io/openaev/config/OpenCTIConfigTest.java @@ -5,13 +5,18 @@ import io.openaev.IntegrationTest; import io.openaev.opencti.config.OpenCTIConfig; import io.openaev.utils.mockConfig.WithMockOpenCTIConfig; +import io.openaev.utilstest.RabbitMQTestListener; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @DisplayName("OpenCTIConfig tests") public class OpenCTIConfigTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/config/XtmHubConfigTest.java b/openaev-api/src/test/java/io/openaev/config/XtmHubConfigTest.java index 414a9304890..d00c74663e6 100644 --- a/openaev-api/src/test/java/io/openaev/config/XtmHubConfigTest.java +++ b/openaev-api/src/test/java/io/openaev/config/XtmHubConfigTest.java @@ -4,14 +4,19 @@ import io.openaev.IntegrationTest; import io.openaev.utils.mockConfig.WithMockXtmHubConfig; +import io.openaev.utilstest.RabbitMQTestListener; import io.openaev.xtmhub.config.XtmHubConfig; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @DisplayName("XtmHubConfig tests") public class XtmHubConfigTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/database/model/InjectTest.java b/openaev-api/src/test/java/io/openaev/database/model/InjectTest.java index 9b12c42d5fc..f636d9480f5 100644 --- a/openaev-api/src/test/java/io/openaev/database/model/InjectTest.java +++ b/openaev-api/src/test/java/io/openaev/database/model/InjectTest.java @@ -9,13 +9,18 @@ import io.openaev.utils.fixtures.composers.ExerciseComposer; import io.openaev.utils.fixtures.composers.InjectComposer; import io.openaev.utils.fixtures.composers.PauseComposer; +import io.openaev.utilstest.RabbitMQTestListener; import java.time.Instant; import org.junit.jupiter.api.*; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.transaction.annotation.Transactional; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(PER_CLASS) @Transactional public class InjectTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/database/model/SimulationTest.java b/openaev-api/src/test/java/io/openaev/database/model/SimulationTest.java index c1619c528aa..9ec0fa124cb 100644 --- a/openaev-api/src/test/java/io/openaev/database/model/SimulationTest.java +++ b/openaev-api/src/test/java/io/openaev/database/model/SimulationTest.java @@ -6,14 +6,19 @@ import io.openaev.database.repository.ExerciseRepository; import io.openaev.utils.fixtures.ExerciseFixture; import io.openaev.utils.fixtures.composers.ExerciseComposer; +import io.openaev.utilstest.RabbitMQTestListener; import jakarta.persistence.EntityManager; import java.time.Instant; import org.junit.jupiter.api.*; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.transaction.annotation.Transactional; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(PER_CLASS) @Transactional class SimulationTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/injector_contract/InjectorContractContentUtilsTest.java b/openaev-api/src/test/java/io/openaev/injector_contract/InjectorContractContentUtilsTest.java index e01eb23dc32..02602411007 100644 --- a/openaev-api/src/test/java/io/openaev/injector_contract/InjectorContractContentUtilsTest.java +++ b/openaev-api/src/test/java/io/openaev/injector_contract/InjectorContractContentUtilsTest.java @@ -26,14 +26,19 @@ import io.openaev.database.model.InjectorContract; import io.openaev.rest.injector_contract.InjectorContractContentUtils; import io.openaev.utils.fixtures.InjectorContractFixture; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.Set; import java.util.stream.Collectors; import java.util.stream.StreamSupport; import org.junit.jupiter.api.Test; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class InjectorContractContentUtilsTest { public static final String EXPECTATION_NAME = "expectation_name"; diff --git a/openaev-api/src/test/java/io/openaev/injects/InjectCrudTest.java b/openaev-api/src/test/java/io/openaev/injects/InjectCrudTest.java index b78d66aeb9f..10337e20b78 100644 --- a/openaev-api/src/test/java/io/openaev/injects/InjectCrudTest.java +++ b/openaev-api/src/test/java/io/openaev/injects/InjectCrudTest.java @@ -9,13 +9,18 @@ import io.openaev.database.repository.ExerciseRepository; import io.openaev.database.repository.InjectRepository; import io.openaev.database.repository.InjectorContractRepository; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.List; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) class InjectCrudTest extends IntegrationTest { @Autowired private InjectRepository injectRepository; diff --git a/openaev-api/src/test/java/io/openaev/injects/email/EmailExecutorTest.java b/openaev-api/src/test/java/io/openaev/injects/email/EmailExecutorTest.java index f7a1754da1a..980864c3d2c 100644 --- a/openaev-api/src/test/java/io/openaev/injects/email/EmailExecutorTest.java +++ b/openaev-api/src/test/java/io/openaev/injects/email/EmailExecutorTest.java @@ -19,14 +19,19 @@ import io.openaev.injectors.email.EmailExecutor; import io.openaev.injectors.email.model.EmailContent; import io.openaev.model.inject.form.Expectation; +import io.openaev.utilstest.RabbitMQTestListener; import jakarta.annotation.Resource; import java.util.Collections; import java.util.List; import org.junit.jupiter.api.Test; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class EmailExecutorTest extends IntegrationTest { @Autowired private EmailExecutor emailExecutor; diff --git a/openaev-api/src/test/java/io/openaev/injects/manual/ManualExecutorTest.java b/openaev-api/src/test/java/io/openaev/injects/manual/ManualExecutorTest.java index c01d323dce1..7790bfe24fc 100644 --- a/openaev-api/src/test/java/io/openaev/injects/manual/ManualExecutorTest.java +++ b/openaev-api/src/test/java/io/openaev/injects/manual/ManualExecutorTest.java @@ -15,6 +15,7 @@ import io.openaev.model.expectation.ManualExpectation; import io.openaev.model.inject.form.Expectation; import io.openaev.service.InjectExpectationService; +import io.openaev.utilstest.RabbitMQTestListener; import java.time.Instant; import java.util.List; import org.junit.jupiter.api.BeforeEach; @@ -22,9 +23,13 @@ import org.mockito.InjectMocks; import org.mockito.Mock; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.util.ReflectionTestUtils; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class ManualExecutorTest extends IntegrationTest { @Mock InjectExpectationService injectExpectationService; diff --git a/openaev-api/src/test/java/io/openaev/opencti/client/mutation/PingTest.java b/openaev-api/src/test/java/io/openaev/opencti/client/mutation/PingTest.java index de56ee21578..97f305b8700 100644 --- a/openaev-api/src/test/java/io/openaev/opencti/client/mutation/PingTest.java +++ b/openaev-api/src/test/java/io/openaev/opencti/client/mutation/PingTest.java @@ -6,11 +6,16 @@ import io.openaev.opencti.client.mutations.Ping; import io.openaev.opencti.connectors.ConnectorBase; import io.openaev.utils.fixtures.opencti.ConnectorFixture; +import io.openaev.utilstest.RabbitMQTestListener; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class PingTest { @Test @DisplayName("When Ping mutation is passed a connector, variables are correctly interpolated") diff --git a/openaev-api/src/test/java/io/openaev/opencti/client/mutation/PushStixBundleTest.java b/openaev-api/src/test/java/io/openaev/opencti/client/mutation/PushStixBundleTest.java index f908832c5c5..5b66a854549 100644 --- a/openaev-api/src/test/java/io/openaev/opencti/client/mutation/PushStixBundleTest.java +++ b/openaev-api/src/test/java/io/openaev/opencti/client/mutation/PushStixBundleTest.java @@ -9,13 +9,18 @@ import io.openaev.stix.objects.Bundle; import io.openaev.stix.types.Identifier; import io.openaev.utils.fixtures.opencti.ConnectorFixture; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.List; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class PushStixBundleTest { @Autowired private ObjectMapper mapper; diff --git a/openaev-api/src/test/java/io/openaev/opencti/client/mutation/RegisterConnectorTest.java b/openaev-api/src/test/java/io/openaev/opencti/client/mutation/RegisterConnectorTest.java index 139f05cf7e6..15b3cd6fb6b 100644 --- a/openaev-api/src/test/java/io/openaev/opencti/client/mutation/RegisterConnectorTest.java +++ b/openaev-api/src/test/java/io/openaev/opencti/client/mutation/RegisterConnectorTest.java @@ -6,12 +6,17 @@ import io.openaev.opencti.client.mutations.RegisterConnector; import io.openaev.opencti.connectors.ConnectorBase; import io.openaev.utils.fixtures.opencti.ConnectorFixture; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.stream.Collectors; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class RegisterConnectorTest { @Test @DisplayName( diff --git a/openaev-api/src/test/java/io/openaev/opencti/connectors/impl/SecurityCoverageConnectorTest.java b/openaev-api/src/test/java/io/openaev/opencti/connectors/impl/SecurityCoverageConnectorTest.java index d5b6e3c9ce1..b927ed98877 100644 --- a/openaev-api/src/test/java/io/openaev/opencti/connectors/impl/SecurityCoverageConnectorTest.java +++ b/openaev-api/src/test/java/io/openaev/opencti/connectors/impl/SecurityCoverageConnectorTest.java @@ -6,11 +6,13 @@ import io.openaev.api.stix_process.StixApi; import io.openaev.config.OpenAEVConfig; import io.openaev.utils.mockConfig.WithMockSecurityCoverageConnectorConfig; +import io.openaev.utilstest.RabbitMQTestListener; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; public class SecurityCoverageConnectorTest extends IntegrationTest { @Nested @@ -19,6 +21,9 @@ public class RemoteUrlOverride { @Nested @WithMockSecurityCoverageConnectorConfig(url = "https://opencti") @SpringBootTest + @TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @DisplayName("With only OpenCTI URL defined as FQDN") public class WithOnlyOpenCTIURLDefinedAsFQDN { @Autowired private SecurityCoverageConnector connector; @@ -33,6 +38,9 @@ public void itAppendsTheGraphQLEndpointToTheURL() { @Nested @WithMockSecurityCoverageConnectorConfig(url = "https://opencti/") @SpringBootTest + @TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @DisplayName("With only OpenCTI URL defined as FQDN with trailing slash") public class WithOnlyOpenCTIURLDefinedAsFQDNWithTrailingSlash { @Autowired private SecurityCoverageConnector connector; @@ -47,6 +55,9 @@ public void itAppendsTheGraphQLEndpointToTheURL() { @Nested @WithMockSecurityCoverageConnectorConfig(url = "https://opencti/graphql") @SpringBootTest + @TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @DisplayName("With only OpenCTI URL defined as FQDN with graphql endpoint set") public class WithOnlyOpenCTIURLDefinedAsFQDNWithGraphqlEndpointSet { @Autowired private SecurityCoverageConnector connector; @@ -61,6 +72,9 @@ public void itAppendsTheGraphQLEndpointToTheURL() { @Nested @WithMockSecurityCoverageConnectorConfig(url = "") @SpringBootTest + @TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @DisplayName("With OpenCTI URL not defined") public class WithNullUrl { @Autowired private SecurityCoverageConnector connector; @@ -78,6 +92,9 @@ public void itAppendsTheGraphQLEndpointToTheURL() { public class ListenCallbackURIOverride { @Nested @SpringBootTest + @TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @WithMockSecurityCoverageConnectorConfig(listenCallbackURI = "some_url") @DisplayName("When listen callback URI is set") public class WhenListenCallbackURIIsSet { @@ -94,6 +111,9 @@ public void itIgnoresAndOverridesIt() { @Nested @SpringBootTest + @TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @WithMockSecurityCoverageConnectorConfig @DisplayName("When listen callback URI is NOT set") public class WhenListenCallbackURIIsNOTSet { diff --git a/openaev-api/src/test/java/io/openaev/rest/ChannelApiTest.java b/openaev-api/src/test/java/io/openaev/rest/ChannelApiTest.java index d8193c2690f..2a1d57b8dfb 100644 --- a/openaev-api/src/test/java/io/openaev/rest/ChannelApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/ChannelApiTest.java @@ -19,14 +19,19 @@ import io.openaev.rest.channel.form.ArticleUpdateInput; import io.openaev.service.ScenarioService; import io.openaev.utils.mockUser.WithMockUser; +import io.openaev.utilstest.RabbitMQTestListener; import org.junit.jupiter.api.*; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.autoconfigure.web.servlet.AutoConfigureMockMvc; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.http.MediaType; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @AutoConfigureMockMvc @TestMethodOrder(MethodOrderer.OrderAnnotation.class) @TestInstance(PER_CLASS) diff --git a/openaev-api/src/test/java/io/openaev/rest/ExerciseLessonsApiTest.java b/openaev-api/src/test/java/io/openaev/rest/ExerciseLessonsApiTest.java index 0b049fb1f92..e00b5af91aa 100644 --- a/openaev-api/src/test/java/io/openaev/rest/ExerciseLessonsApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/ExerciseLessonsApiTest.java @@ -23,6 +23,7 @@ import io.openaev.rest.lessons.form.LessonsSendInput; import io.openaev.service.MailingService; import io.openaev.utils.mockUser.WithMockUser; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.List; import org.junit.jupiter.api.*; import org.springframework.beans.factory.annotation.Autowired; @@ -30,9 +31,13 @@ import org.springframework.boot.test.context.SpringBootTest; import org.springframework.boot.test.mock.mockito.SpyBean; import org.springframework.http.MediaType; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @AutoConfigureMockMvc @TestInstance(TestInstance.Lifecycle.PER_CLASS) public class ExerciseLessonsApiTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/rest/HealthCheckApiTest.java b/openaev-api/src/test/java/io/openaev/rest/HealthCheckApiTest.java index 777375a361b..c3491639c25 100644 --- a/openaev-api/src/test/java/io/openaev/rest/HealthCheckApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/HealthCheckApiTest.java @@ -9,6 +9,7 @@ import io.openaev.rest.health_check.HealthCheckApi; import io.openaev.service.HealthCheckService; import io.openaev.service.exception.HealthCheckFailureException; +import io.openaev.utilstest.RabbitMQTestListener; import org.junit.jupiter.api.*; import org.mockito.InjectMocks; import org.mockito.Mock; @@ -16,9 +17,13 @@ import org.springframework.http.HttpStatus; import org.springframework.http.HttpStatusCode; import org.springframework.http.ResponseEntity; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.web.server.ResponseStatusException; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(PER_CLASS) public class HealthCheckApiTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/rest/MapperApiTest.java b/openaev-api/src/test/java/io/openaev/rest/MapperApiTest.java index 7d1a11bb321..d6c2ba2eaaa 100644 --- a/openaev-api/src/test/java/io/openaev/rest/MapperApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/MapperApiTest.java @@ -22,6 +22,7 @@ import io.openaev.service.MapperService; import io.openaev.utils.fixtures.PaginationFixture; import io.openaev.utils.mockMapper.MockMapperUtils; +import io.openaev.utilstest.RabbitMQTestListener; import java.io.File; import java.io.FileInputStream; import java.io.InputStream; @@ -42,12 +43,16 @@ import org.springframework.data.domain.Pageable; import org.springframework.http.MediaType; import org.springframework.mock.web.MockMultipartFile; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; import org.springframework.test.web.servlet.request.MockMvcRequestBuilders; import org.springframework.test.web.servlet.setup.MockMvcBuilders; import org.springframework.util.ResourceUtils; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @ExtendWith(MockitoExtension.class) public class MapperApiTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/rest/ReportApiTest.java b/openaev-api/src/test/java/io/openaev/rest/ReportApiTest.java index 09285c27efd..2d399a2311d 100644 --- a/openaev-api/src/test/java/io/openaev/rest/ReportApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/ReportApiTest.java @@ -25,6 +25,7 @@ import io.openaev.rest.report.service.ReportService; import io.openaev.utils.fixtures.PaginationFixture; import io.openaev.utils.mockUser.WithMockUser; +import io.openaev.utilstest.RabbitMQTestListener; import java.lang.reflect.Field; import java.util.List; import java.util.UUID; @@ -34,11 +35,15 @@ import org.springframework.boot.test.autoconfigure.web.servlet.AutoConfigureMockMvc; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.http.MediaType; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; import org.springframework.test.web.servlet.request.MockMvcRequestBuilders; import org.springframework.test.web.servlet.setup.MockMvcBuilders; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @AutoConfigureMockMvc @TestMethodOrder(MethodOrderer.OrderAnnotation.class) @TestInstance(PER_CLASS) diff --git a/openaev-api/src/test/java/io/openaev/rest/VariableApiTest.java b/openaev-api/src/test/java/io/openaev/rest/VariableApiTest.java index 3e95213330c..0544854e8c5 100644 --- a/openaev-api/src/test/java/io/openaev/rest/VariableApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/VariableApiTest.java @@ -17,14 +17,19 @@ import io.openaev.database.repository.VariableRepository; import io.openaev.service.ScenarioService; import io.openaev.utils.mockUser.WithMockUser; +import io.openaev.utilstest.RabbitMQTestListener; import org.junit.jupiter.api.*; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.autoconfigure.web.servlet.AutoConfigureMockMvc; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.http.MediaType; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @AutoConfigureMockMvc @TestMethodOrder(MethodOrderer.OrderAnnotation.class) @TestInstance(PER_CLASS) diff --git a/openaev-api/src/test/java/io/openaev/rest/exercise/ExerciseApiStatusTest.java b/openaev-api/src/test/java/io/openaev/rest/exercise/ExerciseApiStatusTest.java index 776e4622d11..2dc34b56459 100644 --- a/openaev-api/src/test/java/io/openaev/rest/exercise/ExerciseApiStatusTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/exercise/ExerciseApiStatusTest.java @@ -24,6 +24,7 @@ import io.openaev.rest.exercise.form.ExerciseUpdateStatusInput; import io.openaev.utils.fixtures.*; import io.openaev.utils.mockUser.WithMockUser; +import io.openaev.utilstest.RabbitMQTestListener; import jakarta.annotation.Resource; import jakarta.servlet.ServletException; import java.time.Clock; @@ -41,10 +42,14 @@ import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.http.MediaType; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; import org.springframework.transaction.annotation.Transactional; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(PER_CLASS) @Transactional public class ExerciseApiStatusTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/rest/exercise/ExerciseApiTest.java b/openaev-api/src/test/java/io/openaev/rest/exercise/ExerciseApiTest.java index 7b5917a5231..0540ba93421 100644 --- a/openaev-api/src/test/java/io/openaev/rest/exercise/ExerciseApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/exercise/ExerciseApiTest.java @@ -21,6 +21,7 @@ import io.openaev.utils.fixtures.*; import io.openaev.utils.fixtures.composers.*; import io.openaev.utils.mockUser.WithMockUser; +import io.openaev.utilstest.RabbitMQTestListener; import jakarta.annotation.Nullable; import jakarta.transaction.Transactional; import java.time.Instant; @@ -33,9 +34,13 @@ import org.springframework.boot.test.autoconfigure.web.servlet.AutoConfigureMockMvc; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.http.MediaType; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @AutoConfigureMockMvc @TestInstance(PER_CLASS) public class ExerciseApiTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/rest/finding/FindingServiceTest.java b/openaev-api/src/test/java/io/openaev/rest/finding/FindingServiceTest.java index 280b431f474..2b41338f46a 100644 --- a/openaev-api/src/test/java/io/openaev/rest/finding/FindingServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/finding/FindingServiceTest.java @@ -5,31 +5,23 @@ import static io.openaev.utils.fixtures.OutputParserFixture.getDefaultContractOutputElement; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertTrue; -import static org.mockito.Mockito.*; import io.openaev.IntegrationTest; import io.openaev.database.model.Asset; import io.openaev.database.model.ContractOutputElement; import io.openaev.database.model.Finding; import io.openaev.database.model.Inject; -import io.openaev.database.repository.AssetRepository; import io.openaev.database.repository.FindingRepository; -import io.openaev.database.repository.TeamRepository; -import io.openaev.database.repository.UserRepository; -import io.openaev.rest.inject.service.InjectService; -import io.openaev.rest.injector_contract.InjectorContractContentUtils; +import io.openaev.utils.helpers.InjectTestHelper; import java.util.ArrayList; import java.util.Arrays; -import java.util.Optional; import java.util.Set; import java.util.stream.Collectors; import org.junit.jupiter.api.*; import org.junit.jupiter.api.extension.ExtendWith; -import org.mockito.ArgumentCaptor; -import org.mockito.InjectMocks; -import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.transaction.annotation.Transactional; @ExtendWith(MockitoExtension.class) class FindingServiceTest extends IntegrationTest { @@ -37,34 +29,19 @@ class FindingServiceTest extends IntegrationTest { public static final String ASSET_1 = "asset1"; public static final String ASSET_2 = "asset2"; - @Autowired private InjectorContractContentUtils injectorContractContentUtils; - @Mock private InjectService injectService; - @Mock private FindingRepository findingRepository; - @Mock private AssetRepository assetRepository; - @Mock private TeamRepository teamRepository; - @Mock private UserRepository userRepository; - @InjectMocks private FindingService findingService; - - @BeforeEach - void setUp() { - findingService = - new FindingService( - injectService, - findingRepository, - assetRepository, - teamRepository, - userRepository, - injectorContractContentUtils); - } + @Autowired private InjectTestHelper injectTestHelper; + @Autowired private FindingService findingService; + @Autowired private FindingRepository findingRepository; @Test @DisplayName("Should have two assets for a finding") + @Transactional void given_a_finding_already_existent_with_one_asset_should_have_two_assets() { Inject inject = getDefaultInject(); Asset asset1 = createDefaultAsset(ASSET_1); - asset1.setId(ASSET_1); + asset1 = injectTestHelper.forceSaveAsset(asset1); Asset asset2 = createDefaultAsset(ASSET_2); - asset2.setId(ASSET_2); + asset2 = injectTestHelper.forceSaveAsset(asset2); String value = "value-already-existent"; ContractOutputElement contractOutputElement = getDefaultContractOutputElement(); @@ -75,29 +52,34 @@ void given_a_finding_already_existent_with_one_asset_should_have_two_assets() { finding1.setType(contractOutputElement.getType()); finding1.setAssets(new ArrayList<>(Arrays.asList(asset1))); - when(findingRepository.findByInjectIdAndValueAndTypeAndKey( - inject.getId(), value, contractOutputElement.getType(), contractOutputElement.getKey())) - .thenReturn(Optional.of(finding1)); + injectTestHelper.forceSaveInject(inject); + injectTestHelper.forceSaveFinding(finding1); findingService.buildFinding(inject, asset2, contractOutputElement, value); - ArgumentCaptor findingCaptor = ArgumentCaptor.forClass(Finding.class); - verify(findingRepository).save(findingCaptor.capture()); - Finding capturedFinding = findingCaptor.getValue(); + Finding capturedFinding = + findingRepository + .findByInjectIdAndValueAndTypeAndKey( + finding1.getInject().getId(), + finding1.getValue(), + finding1.getType(), + finding1.getField()) + .orElseThrow(); assertEquals(2, capturedFinding.getAssets().size()); Set assetIds = capturedFinding.getAssets().stream().map(Asset::getId).collect(Collectors.toSet()); - assertTrue(assetIds.contains(ASSET_1)); - assertTrue(assetIds.contains(ASSET_2)); + assertTrue(assetIds.contains(asset1.getId())); + assertTrue(assetIds.contains(asset2.getId())); } @Test @DisplayName("Should have one asset for a finding") + @Transactional void given_a_finding_already_existent_with_same_asset_should_have_one_assets() { Inject inject = getDefaultInject(); Asset asset1 = createDefaultAsset(ASSET_1); - asset1.setId(ASSET_1); + asset1 = injectTestHelper.forceSaveAsset(asset1); String value = "value-already-existent"; ContractOutputElement contractOutputElement = getDefaultContractOutputElement(); @@ -108,12 +90,23 @@ void given_a_finding_already_existent_with_same_asset_should_have_one_assets() { finding1.setType(contractOutputElement.getType()); finding1.setAssets(new ArrayList<>(Arrays.asList(asset1))); - when(findingRepository.findByInjectIdAndValueAndTypeAndKey( - inject.getId(), value, contractOutputElement.getType(), contractOutputElement.getKey())) - .thenReturn(Optional.of(finding1)); + injectTestHelper.forceSaveInject(inject); + injectTestHelper.forceSaveFinding(finding1); findingService.buildFinding(inject, asset1, contractOutputElement, value); - verify(findingRepository, never()).save(any()); + Finding capturedFinding = + findingRepository + .findByInjectIdAndValueAndTypeAndKey( + finding1.getInject().getId(), + finding1.getValue(), + finding1.getType(), + finding1.getField()) + .orElseThrow(); + + assertEquals(1, capturedFinding.getAssets().size()); + Set assetIds = + capturedFinding.getAssets().stream().map(Asset::getId).collect(Collectors.toSet()); + assertTrue(assetIds.contains(asset1.getId())); } } diff --git a/openaev-api/src/test/java/io/openaev/rest/inject/InjectApiTest.java b/openaev-api/src/test/java/io/openaev/rest/inject/InjectApiTest.java index d0488ee1ce9..c7d4c4361bc 100644 --- a/openaev-api/src/test/java/io/openaev/rest/inject/InjectApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/inject/InjectApiTest.java @@ -40,10 +40,13 @@ import io.openaev.utils.TargetType; import io.openaev.utils.fixtures.*; import io.openaev.utils.fixtures.composers.*; +import io.openaev.utils.helpers.InjectTestHelper; import io.openaev.utils.mockUser.WithMockUser; +import io.openaev.utilstest.KeepRabbit; import jakarta.annotation.Resource; import jakarta.mail.Session; import jakarta.mail.internet.MimeMessage; +import jakarta.persistence.EntityManager; import jakarta.transaction.Transactional; import java.io.File; import java.io.FileInputStream; @@ -51,7 +54,9 @@ import java.nio.charset.StandardCharsets; import java.time.Instant; import java.util.*; +import java.util.concurrent.TimeUnit; import net.javacrumbs.jsonunit.core.Option; +import org.awaitility.Awaitility; import org.junit.jupiter.api.*; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.ArgumentCaptor; @@ -60,6 +65,7 @@ import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.mock.mockito.MockBean; import org.springframework.boot.test.mock.mockito.SpyBean; +import org.springframework.context.ApplicationContext; import org.springframework.http.MediaType; import org.springframework.mail.SimpleMailMessage; import org.springframework.mail.javamail.JavaMailSender; @@ -82,6 +88,7 @@ class InjectApiTest extends IntegrationTest { static Agent AGENT; @Resource protected ObjectMapper mapper; @Autowired private MockMvc mvc; + @Autowired private ApplicationContext applicationContext; @Autowired private ScenarioService scenarioService; @Autowired private ExerciseService exerciseService; @SpyBean private InjectStatusService injectStatusService; @@ -105,6 +112,7 @@ class InjectApiTest extends IntegrationTest { @Autowired private EndpointRepository endpointRepository; @Autowired private ScenarioRepository scenarioRepository; @Autowired private InjectRepository injectRepository; + @Autowired private InjectStatusRepository injectStatusRepository; @Autowired private DocumentRepository documentRepository; @Autowired private CommunicationRepository communicationRepository; @Autowired private InjectExpectationRepository injectExpectationRepository; @@ -117,6 +125,9 @@ class InjectApiTest extends IntegrationTest { @Resource private ObjectMapper objectMapper; @MockBean private JavaMailSender javaMailSender; + @Autowired private EntityManager entityManager; + @Autowired private InjectTestHelper injectTestHelper; + @BeforeAll void beforeAll() { Scenario scenario = new Scenario(); @@ -535,6 +546,7 @@ void deleteInjectsForExerciseTest() throws Exception { @WithMockUser(isAdmin = true) @Transactional @DisplayName("Retrieving executable payloads injects") + @KeepRabbit class RetrievingExecutablePayloadInject { @DisplayName("Get encoded command payload with arguments") @@ -693,27 +705,28 @@ void calling_RetrievingExecutablePayload_should_setStartDateSignature() throws E Command payloadCommand = PayloadFixture.createCommand( "bash", "echo command name #{arg_value}", List.of(), "echo cleanup cmd", domains); - Payload payloadSaved = payloadRepository.save(payloadCommand); + Payload payloadSaved = injectTestHelper.forceSavePayload(payloadCommand); Injector injector = injectorRepository.findByType("openaev_implant").orElseThrow(); InjectorContract injectorContract = InjectorContractFixture.createPayloadInjectorContract(injector, payloadSaved); - InjectorContract injectorContractSaved = injectorContractRepository.save(injectorContract); + InjectorContract injectorContractSaved = + injectTestHelper.forceSaveInjectorContract(injectorContract); Inject inject = InjectFixture.createInjectWithPayloadArg(injectorContractSaved, new HashMap<>()); - Inject injectSaved = injectRepository.save(inject); + Inject injectSaved = injectTestHelper.forceSaveInject(inject); // Prepare injectExpectation on specific agent Endpoint endpoint = EndpointFixture.createEndpoint(); endpoint.setSeenIp("seen-ip-endpoint"); - Endpoint endpointSaved = endpointRepository.save(endpoint); + Endpoint endpointSaved = injectTestHelper.forceSaveEndpoint(endpoint); Agent agent = AgentFixture.createDefaultAgentService(); agent.setAsset(endpointSaved); - Agent agentSaved = agentRepository.save(agent); + Agent agentSaved = injectTestHelper.forceSaveAgent(agent); InjectExpectation detectionExpectation = InjectExpectationFixture.createDetectionInjectExpectation(injectSaved, agentSaved); - injectExpectationRepository.save(detectionExpectation); + injectTestHelper.forceSaveInjectExpectation(detectionExpectation); doNothing() .when(injectStatusService) @@ -731,14 +744,23 @@ void calling_RetrievingExecutablePayload_should_setStartDateSignature() throws E .andExpect(status().is2xxSuccessful()); // -- ASSERT -- - List injectExpectationSaved = - injectExpectationRepository.findAllByInjectAndAgent(injectSaved.getId(), agent.getId()); - assertEquals(1, injectExpectationSaved.size()); - assertEquals( - 1, - injectExpectationSaved.getFirst().getSignatures().stream() - .filter(s -> EXPECTATION_SIGNATURE_TYPE_START_DATE.equals(s.getType())) - .count()); + Awaitility.await() + .atMost(15, TimeUnit.SECONDS) + .with() + .pollInterval(1, TimeUnit.SECONDS) + .until( + () -> { + List injectExpectationSaved = + injectExpectationRepository.findAllByInjectAndAgent( + injectSaved.getId(), agent.getId()); + if (injectExpectationSaved.isEmpty()) { + return false; + } + return injectExpectationSaved.getFirst().getSignatures().stream() + .filter(s -> EXPECTATION_SIGNATURE_TYPE_START_DATE.equals(s.getType())) + .count() + > 0; + }); } @DisplayName("Get obfuscate command") @@ -795,20 +817,12 @@ void getExecutableObfuscatePayloadInject() throws Exception { @Transactional @WithMockUser(isAdmin = true) @DisplayName("Inject Execution Callback Handling (simulating a request from an implant)") + @KeepRabbit class handleInjectExecutionCallback { private Inject getPendingInjectWithAssets() { - return injectComposer - .forInject(InjectFixture.getDefaultInject()) - .withEndpoint( - endpointComposer - .forEndpoint(EndpointFixture.createEndpoint()) - .withAgent(agentComposer.forAgent(AgentFixture.createDefaultAgentService())) - .withAgent(agentComposer.forAgent(AgentFixture.createDefaultAgentSession()))) - .withInjectStatus( - injectStatusComposer.forInjectStatus(InjectStatusFixture.createPendingInjectStatus())) - .persist() - .get(); + return injectTestHelper.getPendingInjectWithAssets( + injectComposer, endpointComposer, agentComposer, injectStatusComposer); } private void performCallbackRequest(String agentId, String injectId, InjectExecutionInput input) @@ -826,6 +840,7 @@ private void performCallbackRequest(String agentId, String injectId, InjectExecu @Nested @DisplayName("Action Handling:") + @KeepRabbit class ActionHandlingTest { @DisplayName("Should add trace when process is not finished") @@ -839,10 +854,28 @@ void shouldAddTraceWhenProcessNotFinished() throws Exception { input.setStatus("SUCCESS"); Inject inject = getPendingInjectWithAssets(); + entityManager.flush(); + // -- EXECUTE -- String agentId = ((Endpoint) inject.getAssets().getFirst()).getAgents().getFirst().getId(); performCallbackRequest(agentId, inject.getId(), input); + Awaitility.await() + .atMost(15, TimeUnit.SECONDS) + .with() + .pollInterval(1, TimeUnit.SECONDS) + .until( + () -> { + Optional injectSaved = injectRepository.findById(inject.getId()); + if (injectSaved.isEmpty()) { + return false; + } + Optional injectStatusSaved = injectSaved.get().getStatus(); + return injectStatusSaved + .filter(injectStatus -> !injectStatus.getTraces().isEmpty()) + .isPresent(); + }); + // -- ASSERT -- Inject injectSaved = injectRepository.findById(inject.getId()).orElseThrow(); InjectStatus injectStatusSaved = injectSaved.getStatus().orElseThrow(); @@ -877,6 +910,22 @@ void shouldAddTraceAndComputeAgentStatusWhenOneAgentFinishes() throws Exception input2.setStatus("INFO"); performCallbackRequest(agentId, inject.getId(), input2); + Awaitility.await() + .atMost(180, TimeUnit.SECONDS) + .with() + .pollInterval(1, TimeUnit.SECONDS) + .until( + () -> { + Optional injectSaved = injectRepository.findById(inject.getId()); + if (injectSaved.isEmpty()) { + return false; + } + Optional injectStatusSaved = injectSaved.get().getStatus(); + return injectStatusSaved + .filter(injectStatus -> injectStatus.getTraces().size() > 1) + .isPresent(); + }); + // -- ASSERT -- Inject injectSaved = injectRepository.findById(inject.getId()).orElseThrow(); InjectStatus injectStatusSaved = injectSaved.getStatus().orElseThrow(); @@ -922,6 +971,22 @@ void shouldAddTraceComputeAgentStatusAndUpdateInjectStatusWhenAllAgentsFinish() performCallbackRequest(firstAgentId, inject.getId(), input2); performCallbackRequest(secondAgentId, inject.getId(), input2); + Awaitility.await() + .atMost(15, TimeUnit.SECONDS) + .with() + .pollInterval(1, TimeUnit.SECONDS) + .until( + () -> { + Optional injectSaved = injectRepository.findById(inject.getId()); + if (injectSaved.isEmpty()) { + return false; + } + Optional injectStatusSaved = injectSaved.get().getStatus(); + return injectStatusSaved + .filter(injectStatus -> injectStatus.getTraces().size() > 1) + .isPresent(); + }); + // -- ASSERT -- Inject injectSaved = injectRepository.findById(inject.getId()).orElseThrow(); InjectStatus injectStatusSaved = injectSaved.getStatus().orElseThrow(); @@ -940,7 +1005,7 @@ void given_completeTrace_should_setEndDateSignature() throws Exception { // create expectation InjectExpectation detectionExpectation = InjectExpectationFixture.createDetectionInjectExpectation(inject, agent); - injectExpectationRepository.save(detectionExpectation); + injectTestHelper.forceSaveInjectExpectation(detectionExpectation); InjectExecutionInput input = new InjectExecutionInput(); input.setMessage("Complete log received"); @@ -949,6 +1014,23 @@ void given_completeTrace_should_setEndDateSignature() throws Exception { input.setDuration(1000); performCallbackRequest(agent.getId(), inject.getId(), input); + Awaitility.await() + .atMost(15, TimeUnit.SECONDS) + .with() + .pollInterval(1, TimeUnit.SECONDS) + .until( + () -> { + List injectExpectationSaved = + injectExpectationRepository.findAllByInjectAndAgent( + inject.getId(), agent.getId()); + List endDatesignatures = + injectExpectationSaved.getFirst().getSignatures().stream() + .filter(s -> EXPECTATION_SIGNATURE_TYPE_END_DATE.equals(s.getType())) + .toList(); + return endDatesignatures.size() > 0; + }); + + // -- ASSERT -- List injectExpectationSaved = injectExpectationRepository.findAllByInjectAndAgent(inject.getId(), agent.getId()); assertEquals(1, injectExpectationSaved.size()); @@ -962,6 +1044,7 @@ void given_completeTrace_should_setEndDateSignature() throws Exception { @Nested @DisplayName("Agent Status Computation") + @KeepRabbit class AgentStatusComputationTest { private void testAgentStatusFunction( @@ -987,6 +1070,22 @@ private void testAgentStatusFunction( input.setStatus("INFO"); performCallbackRequest(firstAgentId, inject.getId(), input); + Awaitility.await() + .atMost(15, TimeUnit.SECONDS) + .with() + .pollInterval(1, TimeUnit.SECONDS) + .until( + () -> { + Optional injectSaved = injectRepository.findById(inject.getId()); + if (injectSaved.isEmpty()) { + return false; + } + Optional injectStatusSaved = injectSaved.get().getStatus(); + return injectStatusSaved + .filter(injectStatus -> injectStatus.getTraces().size() > 2) + .isPresent(); + }); + // -- ASSERT -- Inject injectSaved = injectRepository.findById(inject.getId()).orElseThrow(); InjectStatus injectStatusSaved = injectSaved.getStatus().orElseThrow(); @@ -1027,6 +1126,7 @@ void shouldComputeAgentStatusAsMayBePrevented() throws Exception { @Nested @Transactional @DisplayName("Finding Handling") + @KeepRabbit class FindingHandlingTest { @Test @DisplayName("Should link finding to targeted asset") @@ -1049,7 +1149,7 @@ void given_targetedAsset_should_linkFindingToIt() throws Exception { Command payloadCommand = PayloadFixture.createCommand("bash", "command", null, null, domains); payloadCommand.setOutputParsers(Set.of(outputParser)); - Payload payloadSaved = payloadRepository.save(payloadCommand); + Payload payloadSaved = injectTestHelper.forceSavePayload(payloadCommand); // Create injectorContract with targeted asset field Injector injector = injectorRepository.findByType("openaev_implant").orElseThrow(); @@ -1058,23 +1158,35 @@ void given_targetedAsset_should_linkFindingToIt() throws Exception { injector, payloadSaved, List.of()); InjectorContractFixture.addTargetedAssetFields( injectorContract, "asset-key", ContractTargetedProperty.seen_ip); - InjectorContract injectorContractSaved = injectorContractRepository.save(injectorContract); + injectorContract.setContent(injectorContract.getConvertedContent().toString()); + InjectorContract injectorContractSaved = + injectTestHelper.forceSaveInjectorContract(injectorContract); inject.setInjectorContract(injectorContractSaved); // Set targeted inject on inject Endpoint endpoint = EndpointFixture.createEndpoint(); endpoint.setSeenIp("seen-ip-endpoint"); - Endpoint endpointSaved = endpointRepository.save(endpoint); + Endpoint endpointSaved = injectTestHelper.forceSaveEndpoint(endpoint); ObjectNode content = objectMapper.createObjectNode(); content.set( "asset-key", objectMapper.convertValue(List.of(endpointSaved.getId()), JsonNode.class)); inject.setContent(content); - injectRepository.save(inject); + injectTestHelper.forceSaveInject(inject); // -- EXECUTE -- String agentId = ((Endpoint) inject.getAssets().getFirst()).getAgents().getFirst().getId(); performCallbackRequest(agentId, inject.getId(), input); + Awaitility.await() + .atMost(15, TimeUnit.SECONDS) + .with() + .pollInterval(1, TimeUnit.SECONDS) + .until( + () -> { + List findings = findingRepository.findAllByInjectId(inject.getId()); + return findings.size() > 1; + }); + List findings = findingRepository.findAllByInjectId(inject.getId()); assertEquals(2, findings.size()); assertEquals(1, findings.getFirst().getAssets().size()); @@ -1088,6 +1200,7 @@ void given_targetedAsset_should_linkFindingToIt() throws Exception { @Nested @WithMockUser(isAdmin = true) @DisplayName("Fetch execution traces for inject/atomic overview") + @KeepRabbit class ShouldFetchExecutionTracesForInjectOverview { private Inject buildInjectWithTraces(List traces) { @@ -1310,6 +1423,7 @@ void shouldReturn400WhenTargetTypeIsUnsupported() throws Exception { @Nested @WithMockUser(isAdmin = true) @DisplayName("Fetch documents for inject by payload") + @KeepRabbit class ShouldFetchDocuments { private Inject getInjectWithPayloadAndFileDropDocumentsLinkedOnIt() { diff --git a/openaev-api/src/test/java/io/openaev/rest/notification_rule/NotificationRuleApiTest.java b/openaev-api/src/test/java/io/openaev/rest/notification_rule/NotificationRuleApiTest.java index 6fc77240b40..20904c3479f 100644 --- a/openaev-api/src/test/java/io/openaev/rest/notification_rule/NotificationRuleApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/notification_rule/NotificationRuleApiTest.java @@ -16,6 +16,7 @@ import io.openaev.rest.notification_rule.form.UpdateNotificationRuleInput; import io.openaev.utils.mockUser.WithMockUser; import io.openaev.utils.pagination.SearchPaginationInput; +import io.openaev.utilstest.RabbitMQTestListener; import jakarta.transaction.Transactional; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; @@ -23,9 +24,13 @@ import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.http.MediaType; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(PER_CLASS) public class NotificationRuleApiTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/rest/role/RoleApiTest.java b/openaev-api/src/test/java/io/openaev/rest/role/RoleApiTest.java index 3c6910df684..8bd5157d3b2 100644 --- a/openaev-api/src/test/java/io/openaev/rest/role/RoleApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/role/RoleApiTest.java @@ -15,6 +15,7 @@ import io.openaev.utils.fixtures.RoleFixture; import io.openaev.utils.mockUser.WithMockUser; import io.openaev.utils.pagination.SearchPaginationInput; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.*; import java.util.stream.Collectors; import org.junit.jupiter.api.AfterEach; @@ -23,9 +24,13 @@ import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.http.MediaType; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(PER_CLASS) public class RoleApiTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/rest/scenario/ScenarioImportApiTest.java b/openaev-api/src/test/java/io/openaev/rest/scenario/ScenarioImportApiTest.java index 409f1d4df4c..eec574abdfe 100644 --- a/openaev-api/src/test/java/io/openaev/rest/scenario/ScenarioImportApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/scenario/ScenarioImportApiTest.java @@ -14,6 +14,7 @@ import io.openaev.rest.scenario.response.ImportTestSummary; import io.openaev.service.InjectImportService; import io.openaev.service.ScenarioService; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.Optional; import java.util.UUID; import org.junit.jupiter.api.BeforeEach; @@ -24,11 +25,15 @@ import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.http.MediaType; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; import org.springframework.test.web.servlet.request.MockMvcRequestBuilders; import org.springframework.test.web.servlet.setup.MockMvcBuilders; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @ExtendWith(MockitoExtension.class) public class ScenarioImportApiTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/rest/scenario/ScenarioToExerciseServiceTest.java b/openaev-api/src/test/java/io/openaev/rest/scenario/ScenarioToExerciseServiceTest.java index 5c7c9964034..1d28dd74283 100644 --- a/openaev-api/src/test/java/io/openaev/rest/scenario/ScenarioToExerciseServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/scenario/ScenarioToExerciseServiceTest.java @@ -20,6 +20,7 @@ import io.openaev.service.ScenarioService; import io.openaev.service.ScenarioToExerciseService; import io.openaev.utils.fixtures.ScenarioFixture; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.ArrayList; import java.util.HashSet; import java.util.List; @@ -29,8 +30,12 @@ import org.junit.jupiter.api.TestInstance; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(TestInstance.Lifecycle.PER_CLASS) class ScenarioToExerciseServiceTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/rest/tag_rule/TagRuleApiTest.java b/openaev-api/src/test/java/io/openaev/rest/tag_rule/TagRuleApiTest.java index 684eadc4f40..08998a4ff03 100644 --- a/openaev-api/src/test/java/io/openaev/rest/tag_rule/TagRuleApiTest.java +++ b/openaev-api/src/test/java/io/openaev/rest/tag_rule/TagRuleApiTest.java @@ -18,6 +18,7 @@ import io.openaev.utils.fixtures.AssetGroupFixture; import io.openaev.utils.mockUser.WithMockUser; import io.openaev.utils.pagination.SearchPaginationInput; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.List; import java.util.Map; import org.junit.jupiter.api.AfterEach; @@ -26,9 +27,13 @@ import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.http.MediaType; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.web.servlet.MockMvc; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(PER_CLASS) public class TagRuleApiTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/runner/InitAdminCommandLineRunnerTest.java b/openaev-api/src/test/java/io/openaev/runner/InitAdminCommandLineRunnerTest.java index 8bed8f9d532..37539389549 100644 --- a/openaev-api/src/test/java/io/openaev/runner/InitAdminCommandLineRunnerTest.java +++ b/openaev-api/src/test/java/io/openaev/runner/InitAdminCommandLineRunnerTest.java @@ -9,14 +9,19 @@ import io.openaev.database.model.User; import io.openaev.database.repository.TokenRepository; import io.openaev.database.repository.UserRepository; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.Optional; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.TestInstance; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(TestInstance.Lifecycle.PER_CLASS) public class InitAdminCommandLineRunnerTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/runner/InitStarterPackCommandLineRunnerTest.java b/openaev-api/src/test/java/io/openaev/runner/InitStarterPackCommandLineRunnerTest.java index 5b1f546a0f0..c8339685f17 100644 --- a/openaev-api/src/test/java/io/openaev/runner/InitStarterPackCommandLineRunnerTest.java +++ b/openaev-api/src/test/java/io/openaev/runner/InitStarterPackCommandLineRunnerTest.java @@ -20,6 +20,7 @@ import io.openaev.utils.fixtures.composers.DomainComposer; import io.openaev.utils.fixtures.composers.InjectorContractComposer; import io.openaev.utils.fixtures.composers.PayloadComposer; +import io.openaev.utilstest.RabbitMQTestListener; import java.io.IOException; import java.util.List; import java.util.Optional; @@ -30,10 +31,14 @@ import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.core.io.support.ResourcePatternResolver; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.util.ReflectionTestUtils; import org.springframework.transaction.annotation.Transactional; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(TestInstance.Lifecycle.PER_CLASS) @DisplayName("StarterPack process tests") @Transactional diff --git a/openaev-api/src/test/java/io/openaev/scheduler/jobs/ScenarioExecutionJobTest.java b/openaev-api/src/test/java/io/openaev/scheduler/jobs/ScenarioExecutionJobTest.java index 3a47e8794ea..64aafbf2d8c 100644 --- a/openaev-api/src/test/java/io/openaev/scheduler/jobs/ScenarioExecutionJobTest.java +++ b/openaev-api/src/test/java/io/openaev/scheduler/jobs/ScenarioExecutionJobTest.java @@ -10,6 +10,7 @@ import io.openaev.database.repository.ExerciseRepository; import io.openaev.service.ScenarioService; import io.openaev.utils.fixtures.ScenarioFixture; +import io.openaev.utilstest.RabbitMQTestListener; import java.time.Instant; import java.time.ZoneId; import java.time.ZonedDateTime; @@ -19,8 +20,12 @@ import org.quartz.JobExecutionException; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestMethodOrder(MethodOrderer.OrderAnnotation.class) @TestInstance(TestInstance.Lifecycle.PER_CLASS) class ScenarioExecutionJobTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/service/ExerciseServiceIntegrationTest.java b/openaev-api/src/test/java/io/openaev/service/ExerciseServiceIntegrationTest.java index 4803ada0698..fc0851cc75c 100644 --- a/openaev-api/src/test/java/io/openaev/service/ExerciseServiceIntegrationTest.java +++ b/openaev-api/src/test/java/io/openaev/service/ExerciseServiceIntegrationTest.java @@ -25,6 +25,7 @@ import io.openaev.utils.mapper.ExerciseMapper; import io.openaev.utils.mapper.InjectExpectationMapper; import io.openaev.utils.mapper.InjectMapper; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.ArrayList; import java.util.List; import org.junit.jupiter.api.*; @@ -32,9 +33,13 @@ import org.mockito.Mock; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.transaction.annotation.Transactional; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(TestInstance.Lifecycle.PER_CLASS) class ExerciseServiceIntegrationTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/service/GrantServiceTest.java b/openaev-api/src/test/java/io/openaev/service/GrantServiceTest.java index 903a495f07c..04a619a247e 100644 --- a/openaev-api/src/test/java/io/openaev/service/GrantServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/GrantServiceTest.java @@ -9,12 +9,17 @@ import io.openaev.database.model.User; import io.openaev.database.repository.GrantRepository; import io.openaev.utils.fixtures.UserFixture; +import io.openaev.utilstest.RabbitMQTestListener; import org.junit.jupiter.api.Test; import org.mockito.InjectMocks; import org.mockito.Mock; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class GrantServiceTest extends IntegrationTest { private static final String USER_ID = "userid"; diff --git a/openaev-api/src/test/java/io/openaev/service/HealthCheckServiceTest.java b/openaev-api/src/test/java/io/openaev/service/HealthCheckServiceTest.java index 9dca023c4b2..65f2eaab8cc 100644 --- a/openaev-api/src/test/java/io/openaev/service/HealthCheckServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/HealthCheckServiceTest.java @@ -14,6 +14,7 @@ import io.openaev.database.repository.*; import io.openaev.driver.MinioDriver; import io.openaev.service.exception.HealthCheckFailureException; +import io.openaev.utilstest.RabbitMQTestListener; import java.io.IOException; import java.security.InvalidKeyException; import java.security.NoSuchAlgorithmException; @@ -22,8 +23,12 @@ import org.mockito.InjectMocks; import org.mockito.Mock; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(TestInstance.Lifecycle.PER_CLASS) class HealthCheckServiceTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/service/MapperServiceTest.java b/openaev-api/src/test/java/io/openaev/service/MapperServiceTest.java index 847e3c92952..3952fab74b4 100644 --- a/openaev-api/src/test/java/io/openaev/service/MapperServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/MapperServiceTest.java @@ -17,6 +17,7 @@ import io.openaev.rest.mapper.form.*; import io.openaev.rest.tag.TagService; import io.openaev.utils.mockMapper.MockMapperUtils; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.Optional; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.DisplayName; @@ -26,8 +27,12 @@ import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @ExtendWith(MockitoExtension.class) public class MapperServiceTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/service/NotificationEvenServiceTest.java b/openaev-api/src/test/java/io/openaev/service/NotificationEvenServiceTest.java index fba816784eb..6460fee3ddf 100644 --- a/openaev-api/src/test/java/io/openaev/service/NotificationEvenServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/NotificationEvenServiceTest.java @@ -8,6 +8,7 @@ import io.openaev.notification.handler.ScenarioNotificationEventHandler; import io.openaev.notification.model.NotificationEvent; import io.openaev.notification.model.NotificationEventType; +import io.openaev.utilstest.RabbitMQTestListener; import java.time.Instant; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -15,8 +16,12 @@ import org.springframework.boot.test.context.SpringBootTest; import org.springframework.context.ApplicationEventPublisher; import org.springframework.scheduling.concurrent.ThreadPoolTaskScheduler; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class NotificationEvenServiceTest extends IntegrationTest { @Mock private ApplicationEventPublisher appPublisher; diff --git a/openaev-api/src/test/java/io/openaev/service/NotificationRuleServiceTest.java b/openaev-api/src/test/java/io/openaev/service/NotificationRuleServiceTest.java index fbc53522ed2..1454ea4bb6c 100644 --- a/openaev-api/src/test/java/io/openaev/service/NotificationRuleServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/NotificationRuleServiceTest.java @@ -9,6 +9,7 @@ import io.openaev.database.model.NotificationRuleTrigger; import io.openaev.database.model.NotificationRuleType; import io.openaev.database.repository.NotificationRuleRepository; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.HashMap; import java.util.List; import java.util.Map; @@ -16,8 +17,12 @@ import org.mockito.InjectMocks; import org.mockito.Mock; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class NotificationRuleServiceTest extends IntegrationTest { @Mock private NotificationRuleRepository notificationRuleRepository; diff --git a/openaev-api/src/test/java/io/openaev/service/PermissionServiceTest.java b/openaev-api/src/test/java/io/openaev/service/PermissionServiceTest.java index df77f0f0993..29f789725ed 100644 --- a/openaev-api/src/test/java/io/openaev/service/PermissionServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/PermissionServiceTest.java @@ -12,14 +12,19 @@ import io.openaev.database.repository.ObjectiveRepository; import io.openaev.rest.inject.service.InjectService; import io.openaev.utils.fixtures.UserFixture; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.*; import org.junit.jupiter.api.Test; import org.mockito.InjectMocks; import org.mockito.Mock; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.web.bind.annotation.RequestMethod; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class PermissionServiceTest extends IntegrationTest { private static final String RESOURCE_ID = "resourceid"; private static final String USER_ID = "userid"; diff --git a/openaev-api/src/test/java/io/openaev/service/PlatformServiceSettingsTest.java b/openaev-api/src/test/java/io/openaev/service/PlatformServiceSettingsTest.java index 21254be344a..2b957fcf2bc 100644 --- a/openaev-api/src/test/java/io/openaev/service/PlatformServiceSettingsTest.java +++ b/openaev-api/src/test/java/io/openaev/service/PlatformServiceSettingsTest.java @@ -10,6 +10,7 @@ import io.openaev.rest.settings.PreviewFeature; import io.openaev.rest.settings.response.PlatformSettings; import io.openaev.utils.mockUser.WithMockUser; +import io.openaev.utilstest.RabbitMQTestListener; import jakarta.annotation.Resource; import java.util.List; import org.junit.jupiter.api.BeforeAll; @@ -19,10 +20,14 @@ import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.transaction.annotation.Transactional; @Transactional @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @ExtendWith(MockitoExtension.class) @TestInstance(PER_CLASS) public class PlatformServiceSettingsTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/service/ReportServiceTest.java b/openaev-api/src/test/java/io/openaev/service/ReportServiceTest.java index 54aca29ca23..4c4ec226ba7 100644 --- a/openaev-api/src/test/java/io/openaev/service/ReportServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/ReportServiceTest.java @@ -16,6 +16,7 @@ import io.openaev.rest.report.model.ReportInjectComment; import io.openaev.rest.report.repository.ReportRepository; import io.openaev.rest.report.service.ReportService; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.List; import java.util.Optional; import java.util.UUID; @@ -25,8 +26,12 @@ import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @ExtendWith(MockitoExtension.class) public class ReportServiceTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/service/RoleServiceTest.java b/openaev-api/src/test/java/io/openaev/service/RoleServiceTest.java index 63e66383de3..ee7b42cacd5 100644 --- a/openaev-api/src/test/java/io/openaev/service/RoleServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/RoleServiceTest.java @@ -6,13 +6,18 @@ import io.openaev.IntegrationTest; import io.openaev.database.model.Capability; import io.openaev.database.repository.RoleRepository; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.Set; import org.junit.jupiter.api.Test; import org.mockito.InjectMocks; import org.mockito.Mock; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class RoleServiceTest extends IntegrationTest { @Mock RoleRepository roleRepository; diff --git a/openaev-api/src/test/java/io/openaev/service/ScenarioServiceTest.java b/openaev-api/src/test/java/io/openaev/service/ScenarioServiceTest.java index e727a7a1564..2b3020affe5 100644 --- a/openaev-api/src/test/java/io/openaev/service/ScenarioServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/ScenarioServiceTest.java @@ -26,15 +26,20 @@ import io.openaev.utils.fixtures.*; import io.openaev.utils.mapper.ExerciseMapper; import io.openaev.utils.mapper.ScenarioMapper; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.*; import org.junit.jupiter.api.*; import org.mockito.InjectMocks; import org.mockito.Mock; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.transaction.annotation.Transactional; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(TestInstance.Lifecycle.PER_CLASS) class ScenarioServiceTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/service/SmtpServiceTest.java b/openaev-api/src/test/java/io/openaev/service/SmtpServiceTest.java index 85fb41f272c..c45503207ee 100644 --- a/openaev-api/src/test/java/io/openaev/service/SmtpServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/SmtpServiceTest.java @@ -6,6 +6,7 @@ import io.openaev.database.repository.SettingRepository; import io.openaev.injectors.email.service.SmtpService; +import io.openaev.utilstest.RabbitMQTestListener; import jakarta.mail.MessagingException; import jakarta.mail.Session; import jakarta.mail.internet.MimeMessage; @@ -19,10 +20,14 @@ import org.springframework.boot.test.context.SpringBootTest; import org.springframework.boot.test.mock.mockito.MockBean; import org.springframework.mail.javamail.JavaMailSenderImpl; +import org.springframework.test.context.TestExecutionListeners; import org.springframework.test.util.ReflectionTestUtils; import org.springframework.transaction.annotation.Transactional; @SpringBootTest(properties = "spring.task.scheduling.enabled=false") +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(TestInstance.Lifecycle.PER_CLASS) @Transactional public class SmtpServiceTest { diff --git a/openaev-api/src/test/java/io/openaev/service/TagRuleServiceTest.java b/openaev-api/src/test/java/io/openaev/service/TagRuleServiceTest.java index 98067121841..7b80155aee8 100644 --- a/openaev-api/src/test/java/io/openaev/service/TagRuleServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/TagRuleServiceTest.java @@ -17,6 +17,7 @@ import io.openaev.utils.fixtures.AssetGroupFixture; import io.openaev.utils.fixtures.TagFixture; import io.openaev.utils.fixtures.TagRuleFixture; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.HashSet; import java.util.List; import java.util.Optional; @@ -24,8 +25,12 @@ import org.mockito.InjectMocks; import org.mockito.Mock; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) public class TagRuleServiceTest extends IntegrationTest { private static final String TAG_RULE_ID = "tagruleid"; private static final String TAG_RULE_ID_2 = "tagruleid2"; diff --git a/openaev-api/src/test/java/io/openaev/service/VariableServiceTest.java b/openaev-api/src/test/java/io/openaev/service/VariableServiceTest.java index 68b4c2f6046..f95100faf78 100644 --- a/openaev-api/src/test/java/io/openaev/service/VariableServiceTest.java +++ b/openaev-api/src/test/java/io/openaev/service/VariableServiceTest.java @@ -8,13 +8,18 @@ import io.openaev.database.model.Variable; import io.openaev.database.model.Variable.VariableType; import io.openaev.database.repository.ExerciseRepository; +import io.openaev.utilstest.RabbitMQTestListener; import java.util.List; import java.util.NoSuchElementException; import org.junit.jupiter.api.*; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.test.context.TestExecutionListeners; @SpringBootTest +@TestExecutionListeners( + value = {RabbitMQTestListener.class}, + mergeMode = TestExecutionListeners.MergeMode.MERGE_WITH_DEFAULTS) @TestInstance(PER_CLASS) @TestMethodOrder(MethodOrderer.OrderAnnotation.class) public class VariableServiceTest extends IntegrationTest { diff --git a/openaev-api/src/test/java/io/openaev/utils/helpers/InjectTestHelper.java b/openaev-api/src/test/java/io/openaev/utils/helpers/InjectTestHelper.java new file mode 100644 index 00000000000..4d5ac82077c --- /dev/null +++ b/openaev-api/src/test/java/io/openaev/utils/helpers/InjectTestHelper.java @@ -0,0 +1,89 @@ +package io.openaev.utils.helpers; + +import io.openaev.database.model.*; +import io.openaev.database.repository.*; +import io.openaev.utils.fixtures.AgentFixture; +import io.openaev.utils.fixtures.EndpointFixture; +import io.openaev.utils.fixtures.InjectFixture; +import io.openaev.utils.fixtures.InjectStatusFixture; +import io.openaev.utils.fixtures.composers.AgentComposer; +import io.openaev.utils.fixtures.composers.EndpointComposer; +import io.openaev.utils.fixtures.composers.InjectComposer; +import io.openaev.utils.fixtures.composers.InjectStatusComposer; +import lombok.RequiredArgsConstructor; +import org.springframework.stereotype.Component; +import org.springframework.transaction.annotation.Propagation; +import org.springframework.transaction.annotation.Transactional; + +@Component +@RequiredArgsConstructor +public class InjectTestHelper { + + private final InjectExpectationRepository injectExpectationRepository; + private final PayloadRepository payloadRepository; + private final InjectorContractRepository injectorContractRepository; + private final AgentRepository agentRepository; + private final EndpointRepository endpointRepository; + private final InjectRepository injectRepository; + private final FindingRepository findingRepository; + private final AssetRepository assetRepository; + + @Transactional(propagation = Propagation.REQUIRES_NEW) + public Inject getPendingInjectWithAssets( + InjectComposer injectComposer, + EndpointComposer endpointComposer, + AgentComposer agentComposer, + InjectStatusComposer injectStatusComposer) { + return injectComposer + .forInject(InjectFixture.getDefaultInject()) + .withEndpoint( + endpointComposer + .forEndpoint(EndpointFixture.createEndpoint()) + .withAgent(agentComposer.forAgent(AgentFixture.createDefaultAgentService())) + .withAgent(agentComposer.forAgent(AgentFixture.createDefaultAgentSession()))) + .withInjectStatus( + injectStatusComposer.forInjectStatus(InjectStatusFixture.createPendingInjectStatus())) + .persist() + .get(); + } + + @Transactional(propagation = Propagation.REQUIRES_NEW) + public InjectExpectation forceSaveInjectExpectation(InjectExpectation expectation) { + return injectExpectationRepository.save(expectation); + } + + @Transactional(propagation = Propagation.REQUIRES_NEW) + public Payload forceSavePayload(Payload payload) { + return payloadRepository.save(payload); + } + + @Transactional(propagation = Propagation.REQUIRES_NEW) + public InjectorContract forceSaveInjectorContract(InjectorContract injectorContract) { + return injectorContractRepository.save(injectorContract); + } + + @Transactional(propagation = Propagation.REQUIRES_NEW) + public Inject forceSaveInject(Inject inject) { + return injectRepository.save(inject); + } + + @Transactional(propagation = Propagation.REQUIRES_NEW) + public Agent forceSaveAgent(Agent agent) { + return agentRepository.save(agent); + } + + @Transactional(propagation = Propagation.REQUIRES_NEW) + public Endpoint forceSaveEndpoint(Endpoint endpoint) { + return endpointRepository.save(endpoint); + } + + @Transactional(propagation = Propagation.REQUIRES_NEW) + public Finding forceSaveFinding(Finding finding) { + return findingRepository.save(finding); + } + + @Transactional(propagation = Propagation.REQUIRES_NEW) + public Asset forceSaveAsset(Asset asset) { + return assetRepository.save(asset); + } +} diff --git a/openaev-api/src/test/java/io/openaev/utilstest/KeepRabbit.java b/openaev-api/src/test/java/io/openaev/utilstest/KeepRabbit.java new file mode 100644 index 00000000000..b5b25d93c54 --- /dev/null +++ b/openaev-api/src/test/java/io/openaev/utilstest/KeepRabbit.java @@ -0,0 +1,10 @@ +package io.openaev.utilstest; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +@Retention(RetentionPolicy.RUNTIME) +@Target({ElementType.TYPE}) +public @interface KeepRabbit {} diff --git a/openaev-api/src/test/java/io/openaev/utilstest/RabbitMQTestListener.java b/openaev-api/src/test/java/io/openaev/utilstest/RabbitMQTestListener.java new file mode 100644 index 00000000000..a189f505d10 --- /dev/null +++ b/openaev-api/src/test/java/io/openaev/utilstest/RabbitMQTestListener.java @@ -0,0 +1,28 @@ +package io.openaev.utilstest; + +import io.openaev.rest.inject.InjectApi; +import lombok.extern.slf4j.Slf4j; +import org.springframework.context.ApplicationContext; +import org.springframework.test.context.TestContext; +import org.springframework.test.context.TestExecutionListener; + +@Slf4j +public class RabbitMQTestListener implements TestExecutionListener { + + @Override + public void afterTestClass(TestContext testContext) throws Exception { + Class testClass = testContext.getTestClass(); + + // Ignoring nested classes + if (testClass.isAnnotationPresent(KeepRabbit.class)) { + log.info("Skipping restore for @Nested class: {}", testClass.getSimpleName()); + return; + } + + // Closing RabbitMQ consumers + ApplicationContext context = testContext.getApplicationContext(); + context.getBean(InjectApi.class).getInjectTraceQueueService().stop(); + + log.info("RabbitMQ consumers closed for class : {}", testClass.getSimpleName()); + } +} diff --git a/openaev-api/src/test/resources/application.properties b/openaev-api/src/test/resources/application.properties index 09d9cace66f..f0fd6b7bae7 100644 --- a/openaev-api/src/test/resources/application.properties +++ b/openaev-api/src/test/resources/application.properties @@ -18,10 +18,27 @@ logging.level.root=ERROR logging.level.org=ERROR # rabbit mq +openaev.rabbitmq.hostname=localhost +openaev.rabbitmq.port=5672 +openaev.rabbitmq.prefix=openbas +openaev.rabbitmq.user=guest +openaev.rabbitmq.pass=guest +openaev.rabbitmq.vhost=/ +openaev.rabbitmq.ssl=false +openaev.rabbitmq.management-port=15672 +openaev.rabbitmq.queue-type=classic openaev.rabbitmq.management-insecure=true openaev.rabbitmq.trust-store-password= openaev.rabbitmq.trust.store= +openaev.queue-config.inject-trace.publisher-number=1 +openaev.queue-config.inject-trace.consumer-number=1 +openaev.queue-config.inject-trace.worker-number=1 +openaev.queue-config.inject-trace.worker-frequency=2000 +openaev.queue-config.inject-trace.queue-name=inject-trace +openaev.queue-config.inject-trace.max-size=100 +openaev.queue-config.inject-trace.consumer-qos=30 +openaev.queue-config.inject-trace.publisher-qos=30 # Authenticators ## Local @@ -41,7 +58,7 @@ spring.flyway.baseline-on-migrate=true spring.flyway.baseline-version=0 spring.flyway.postgresql.transactional-lock=false spring.profiles.active=test -spring.jpa.properties.hibernate.jdbc.batch_size=250 +spring.jpa.properties.hibernate.jdbc.batch_size=30 spring.jpa.properties.hibernate.order_inserts=true ### ENGINE Configuration @@ -145,4 +162,4 @@ executor.crowdstrike.id=2a16dcc4-55ac-40fc-8110-d5968a46cdd1 executor.caldera.id=2a16dcc4-55ac-40fc-8110-d5968a46cdd1 injector.caldera.id=2a16dcc4-55ac-40fc-8110-d5968a46cdd1 -spring.datasource.hikari.maximum-pool-size=2 \ No newline at end of file +spring.datasource.hikari.maximum-pool-size=5 \ No newline at end of file diff --git a/openaev-framework/src/main/java/io/openaev/config/OpenAEVConfig.java b/openaev-framework/src/main/java/io/openaev/config/OpenAEVConfig.java index fec6962542f..7e3ea608945 100644 --- a/openaev-framework/src/main/java/io/openaev/config/OpenAEVConfig.java +++ b/openaev-framework/src/main/java/io/openaev/config/OpenAEVConfig.java @@ -5,12 +5,15 @@ import com.fasterxml.jackson.annotation.JsonIgnore; import com.fasterxml.jackson.annotation.JsonProperty; import jakarta.validation.constraints.NotBlank; +import java.util.Map; import lombok.Data; import org.springframework.beans.factory.annotation.Value; +import org.springframework.boot.context.properties.ConfigurationProperties; import org.springframework.stereotype.Component; @Component @Data +@ConfigurationProperties(prefix = "openaev") public class OpenAEVConfig { @JsonProperty("parameters_id") @@ -105,6 +108,10 @@ public class OpenAEVConfig { @Value("${openbas.extra-trusted-certs-dir:${openaev.extra-trusted-certs-dir:#{null}}}") private String extraTrustedCertsDir; + @JsonProperty("queue-config") + @Value("${openbas.queue-config:${openaev.queue-config:#{null}}}") + private Map queueConfig; + public String getBaseUrl() { return url(baseUrl); } diff --git a/openaev-framework/src/main/java/io/openaev/config/QueueConfig.java b/openaev-framework/src/main/java/io/openaev/config/QueueConfig.java new file mode 100644 index 00000000000..780e881e083 --- /dev/null +++ b/openaev-framework/src/main/java/io/openaev/config/QueueConfig.java @@ -0,0 +1,31 @@ +package io.openaev.config; + +import com.fasterxml.jackson.annotation.JsonProperty; +import lombok.Data; + +@Data +public class QueueConfig { + @JsonProperty("publisher-number") + private int publisherNumber = 1; + + @JsonProperty("consumer-number") + private int consumerNumber = 1; + + @JsonProperty("worker-number") + private int workerNumber = 1; + + @JsonProperty("worker-frequency") + private int workerFrequency = 10000; + + @JsonProperty("queue-name") + private String queueName = "openaev-queue"; + + @JsonProperty("max-size") + private int maxSize = 100; + + @JsonProperty("consumer-qos") + private int consumerQos = 30; + + @JsonProperty("publisher-qos") + private int publisherQos = 30; +} diff --git a/openaev-front/src/utils/api-types.d.ts b/openaev-front/src/utils/api-types.d.ts index fc4574505ba..85e196c34f3 100644 --- a/openaev-front/src/utils/api-types.d.ts +++ b/openaev-front/src/utils/api-types.d.ts @@ -4770,6 +4770,7 @@ export interface PlatformSettings { enabled_dev_features?: ( | "_RESERVED" | "STIX_SECURITY_COVERAGE_FOR_VULNERABILITIES" + | "LEGACY_INGESTION_EXECUTION_TRACE" )[]; /** True if the Caldera Executor is enabled */ executor_caldera_enable?: boolean; diff --git a/openaev-model/src/main/java/io/openaev/database/helper/ExecutionTraceRepositoryHelper.java b/openaev-model/src/main/java/io/openaev/database/helper/ExecutionTraceRepositoryHelper.java new file mode 100644 index 00000000000..a5247f8e030 --- /dev/null +++ b/openaev-model/src/main/java/io/openaev/database/helper/ExecutionTraceRepositoryHelper.java @@ -0,0 +1,140 @@ +package io.openaev.database.helper; + +import io.openaev.database.model.ExecutionTrace; +import java.sql.Connection; +import java.sql.PreparedStatement; +import java.sql.SQLException; +import java.sql.Timestamp; +import java.time.Instant; +import java.util.UUID; +import javax.sql.DataSource; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.stereotype.Repository; + +@Repository +public class ExecutionTraceRepositoryHelper { + + @Autowired private DataSource dataSource; + + private static final String INSERT_EXECUTION_TRACE = + """ + INSERT INTO execution_traces ( + execution_trace_id, + execution_inject_status_id, + execution_inject_test_status_id, + execution_agent_id, + execution_message, + execution_structured_output, + execution_action, + execution_status, + execution_time, + execution_context_identifiers, + execution_created_at, + execution_updated_at + ) VALUES ( + ?, + ?, + ?, + ?, + ?, + ?, + ?, + ?, + ?, + ?, + ?, + ? + )"""; + + /** + * Save execution trace with a low level database call + * + * @param executionTrace the execution trace + * @return the id of the new trace + */ + public String saveExecutionTrace(ExecutionTrace executionTrace) { + try (Connection conn = dataSource.getConnection()) { + + try (PreparedStatement ps = conn.prepareStatement(INSERT_EXECUTION_TRACE)) { + + String injectStatusId = null; + if (executionTrace.getInjectStatus() != null) { + injectStatusId = executionTrace.getInjectStatus().getId(); + } + String injectTestStatusId = null; + if (executionTrace.getInjectTestStatus() != null) { + injectTestStatusId = executionTrace.getInjectTestStatus().getId(); + } + String structuredOutputAsText = null; + if (executionTrace.getStructuredOutput() != null) { + structuredOutputAsText = executionTrace.getStructuredOutput().asText(); + } + String id = UUID.randomUUID().toString(); + + ps.setString(1, id); + ps.setString(2, injectStatusId); + ps.setString(3, injectTestStatusId); + ps.setString(4, executionTrace.getAgent().getId()); + ps.setString(5, executionTrace.getMessage()); + ps.setString(6, structuredOutputAsText); + ps.setString(7, executionTrace.getAction().name()); + ps.setString(8, executionTrace.getStatus().name()); + ps.setTimestamp(9, Timestamp.from(executionTrace.getTime())); + ps.setArray(10, conn.createArrayOf("text", executionTrace.getIdentifiers().toArray())); + ps.setTimestamp(11, Timestamp.from(executionTrace.getCreationDate())); + ps.setTimestamp(12, Timestamp.from(executionTrace.getUpdateDate())); + + ps.executeUpdate(); + + return id; + } + } catch (SQLException e) { + throw new RuntimeException("Failed to insert execution trace", e); + } + } + + /** + * Update an inject status with a new status name and end_date with a low level database call + * + * @param injectStatusId the id of the inject status to update + * @param name the name of the new status + * @param endDate the end date + */ + public void updateInjectStatus(String injectStatusId, String name, Instant endDate) { + String sql = + "UPDATE injects_statuses SET status_name = ?, tracking_end_date = ? WHERE status_id = ?"; + + try (Connection conn = dataSource.getConnection(); + PreparedStatement ps = conn.prepareStatement(sql)) { + + ps.setString(1, name); + ps.setTimestamp(2, endDate != null ? Timestamp.from(endDate) : null); + ps.setString(3, injectStatusId); + ps.executeUpdate(); + + } catch (SQLException e) { + throw new RuntimeException("Failed to update inject status", e); + } + } + + /** + * Update the update date of an injects with a low level database call + * + * @param id the id of the inject + * @param updatedAt the update date + */ + public void updateInjectUpdateDate(String id, Instant updatedAt) { + String sql = "UPDATE injects SET inject_updated_at = ? WHERE inject_id = ?"; + + try (Connection conn = dataSource.getConnection(); + PreparedStatement ps = conn.prepareStatement(sql)) { + + ps.setTimestamp(1, updatedAt != null ? Timestamp.from(updatedAt) : null); + ps.setString(2, id); + ps.executeUpdate(); + + } catch (SQLException e) { + throw new RuntimeException("Failed to update inject update date", e); + } + } +} diff --git a/openaev-model/src/main/java/io/openaev/database/helper/InjectExpectationRepositoryHelper.java b/openaev-model/src/main/java/io/openaev/database/helper/InjectExpectationRepositoryHelper.java new file mode 100644 index 00000000000..0ca99e51008 --- /dev/null +++ b/openaev-model/src/main/java/io/openaev/database/helper/InjectExpectationRepositoryHelper.java @@ -0,0 +1,49 @@ +package io.openaev.database.helper; + +import java.sql.Connection; +import java.sql.PreparedStatement; +import java.sql.SQLException; +import javax.sql.DataSource; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.stereotype.Repository; + +@Repository +public class InjectExpectationRepositoryHelper { + + @Autowired private DataSource dataSource; + + /** + * Update the signature of the expectation with a new type/value tuple passed in parameter for an + * inject and agent + * + * @param injectId the id of the inject + * @param agentId the id of the agent + * @param type the type of the element + * @param value the value of the element + */ + public void insertSignatureForAgentAndInject( + String injectId, String agentId, String type, String value) { + try (Connection conn = dataSource.getConnection()) { + + try (PreparedStatement ps = + conn.prepareStatement( + """ + UPDATE injects_expectations + SET inject_expectation_signatures = + COALESCE(inject_expectation_signatures, '[]'::jsonb) || + jsonb_build_array(jsonb_build_object('type', ?, 'value', ?)) + WHERE inject_id = ? AND agent_id = ? + """)) { + + ps.setString(1, type); + ps.setString(2, value); + ps.setString(3, injectId); + ps.setString(4, agentId); + + ps.executeUpdate(); + } + } catch (SQLException e) { + throw new RuntimeException(e); + } + } +} diff --git a/openaev-model/src/main/java/io/openaev/database/repository/FindingRepository.java b/openaev-model/src/main/java/io/openaev/database/repository/FindingRepository.java index 36cc7fff93a..7b40c31dfd5 100644 --- a/openaev-model/src/main/java/io/openaev/database/repository/FindingRepository.java +++ b/openaev-model/src/main/java/io/openaev/database/repository/FindingRepository.java @@ -14,6 +14,8 @@ import org.springframework.data.repository.CrudRepository; import org.springframework.data.repository.query.Param; import org.springframework.stereotype.Repository; +import org.springframework.transaction.annotation.Propagation; +import org.springframework.transaction.annotation.Transactional; @Repository public interface FindingRepository @@ -45,4 +47,45 @@ Optional findByInjectIdAndValueAndTypeAndKey( + ";", nativeQuery = true) List findForIndexing(@Param("from") Instant from); + + @Query( + value = + """ + WITH inserted_finding AS ( + INSERT INTO findings + (finding_id, finding_field, finding_type, finding_value, + finding_labels, finding_inject_id, finding_name) + VALUES + (gen_random_uuid(), :findingField, :findingType, :findingValue, + :findingLabels, :findingInjectId, :findingName) + ON CONFLICT (finding_inject_id, finding_field, finding_type, finding_value) + DO UPDATE SET finding_name = EXCLUDED.finding_name + RETURNING finding_id + ), + inserted_asset AS ( + INSERT INTO findings_assets (finding_id, asset_id) + SELECT finding_id, :assetId + FROM inserted_finding + ON CONFLICT DO NOTHING + ), + inserted_tags AS ( + INSERT INTO findings_tags (finding_id, tag_id) + SELECT finding_id, tag_id + FROM inserted_finding + CROSS JOIN unnest(CAST(:tagIds AS varchar[])) AS tag_id + ON CONFLICT DO NOTHING + ) + SELECT finding_id FROM inserted_finding + """, + nativeQuery = true) + @Transactional(propagation = Propagation.REQUIRES_NEW) + String saveCompleteFinding( + @Param("findingField") String findingField, + @Param("findingType") String findingType, + @Param("findingValue") String findingValue, + @Param("findingLabels") String[] findingLabels, + @Param("findingInjectId") String injectId, + @Param("findingName") String name, + @Param("assetId") String assetId, + @Param("tagIds") String[] tagIds); } diff --git a/openaev-model/src/main/java/io/openaev/database/repository/InjectRepository.java b/openaev-model/src/main/java/io/openaev/database/repository/InjectRepository.java index acf67723c3e..930d82e6fd6 100644 --- a/openaev-model/src/main/java/io/openaev/database/repository/InjectRepository.java +++ b/openaev-model/src/main/java/io/openaev/database/repository/InjectRepository.java @@ -12,10 +12,7 @@ import java.util.Set; import org.jetbrains.annotations.NotNull; import org.springframework.data.domain.Pageable; -import org.springframework.data.jpa.repository.JpaRepository; -import org.springframework.data.jpa.repository.JpaSpecificationExecutor; -import org.springframework.data.jpa.repository.Modifying; -import org.springframework.data.jpa.repository.Query; +import org.springframework.data.jpa.repository.*; import org.springframework.data.repository.query.Param; import org.springframework.stereotype.Repository; import org.springframework.transaction.annotation.Transactional; @@ -226,28 +223,82 @@ List userCountGroupByAttackPatternInExercise( @Query( value = - " SELECT injects.inject_id, ins.status_name, injects.inject_scenario, " - + "coalesce(array_agg(it.team_id) FILTER ( WHERE it.team_id IS NOT NULL ), '{}') as inject_teams, " - + "coalesce(array_agg(assets.asset_id) FILTER ( WHERE assets.asset_id IS NOT NULL ), '{}') as inject_assets, " - + "coalesce(array_agg(iag.asset_group_id) FILTER ( WHERE iag.asset_group_id IS NOT NULL ), '{}') as inject_asset_groups, " - + "coalesce(array_agg(ie.inject_expectation_id) FILTER ( WHERE ie.inject_expectation_id IS NOT NULL ), '{}') as inject_expectations, " - + "coalesce(array_agg(com.communication_id) FILTER ( WHERE com.communication_id IS NOT NULL ), '{}') as inject_communications, " - + "coalesce(array_agg(apkcp.phase_id) FILTER ( WHERE apkcp.phase_id IS NOT NULL ), '{}') as inject_kill_chain_phases, " - + "coalesce(array_union_agg(injcon.injector_contract_platforms) FILTER ( WHERE injcon.injector_contract_platforms IS NOT NULL ), '{}') as inject_platforms " - + "FROM injects " - + "LEFT JOIN injects_teams it ON injects.inject_id = it.inject_id " - + "LEFT JOIN injects_assets ia ON injects.inject_id = ia.inject_id " - + "LEFT JOIN injects_asset_groups iag ON injects.inject_id = iag.inject_id " - + "LEFT JOIN asset_groups_assets aga ON aga.asset_group_id = iag.asset_group_id " - + "LEFT JOIN assets ON assets.asset_id = ia.asset_id OR aga.asset_id = assets.asset_id " - + "LEFT JOIN communications com ON com.communication_inject = injects.inject_id " - + "LEFT JOIN injects_expectations ie ON injects.inject_id = ie.inject_id " - + "LEFT JOIN injectors_contracts_attack_patterns icap ON icap.injector_contract_id = injects.inject_injector_contract " - + "LEFT JOIN injectors_contracts injcon ON injcon.injector_contract_id = injects.inject_injector_contract " - + "LEFT JOIN attack_patterns_kill_chain_phases apkcp ON apkcp.attack_pattern_id = icap.attack_pattern_id " - + "LEFT JOIN injects_statuses ins ON ins.status_inject = injects.inject_id " - + "WHERE injects.inject_id IN :ids " - + "GROUP BY injects.inject_id, ins.status_name;", + "WITH inject_teams AS ( " + + " SELECT inject_id, array_agg(team_id) as team_ids " + + " FROM injects_teams " + + " WHERE inject_id IN (:ids) " + + " GROUP BY inject_id " + + "), " + + "inject_assets AS ( " + + " SELECT " + + " i.inject_id, " + + " array_agg(DISTINCT a.asset_id) as asset_ids " + + " FROM injects i " + + " LEFT JOIN injects_assets ia ON i.inject_id = ia.inject_id " + + " LEFT JOIN injects_asset_groups iag ON i.inject_id = iag.inject_id " + + " LEFT JOIN asset_groups_assets aga ON aga.asset_group_id = iag.asset_group_id " + + " LEFT JOIN assets a ON a.asset_id = ia.asset_id OR aga.asset_id = a.asset_id " + + " WHERE i.inject_id IN (:ids) " + + " GROUP BY i.inject_id " + + "), " + + "inject_asset_groups AS ( " + + " SELECT inject_id, array_agg(asset_group_id) as asset_group_ids " + + " FROM injects_asset_groups " + + " WHERE inject_id IN (:ids) " + + " GROUP BY inject_id " + + "), " + + "inject_expectations AS ( " + + " SELECT inject_id, array_agg(inject_expectation_id) as expectation_ids " + + " FROM injects_expectations " + + " WHERE inject_id IN (:ids) " + + " GROUP BY inject_id " + + "), " + + "inject_communications AS ( " + + " SELECT communication_inject as inject_id, array_agg(communication_id) as communication_ids " + + " FROM communications " + + " WHERE communication_inject IN (:ids) " + + " GROUP BY communication_inject " + + "), " + + "inject_kill_chains AS ( " + + " SELECT " + + " i.inject_id, " + + " array_agg(DISTINCT apkcp.phase_id) as phase_ids " + + " FROM injects i " + + " JOIN injectors_contracts_attack_patterns icap ON icap.injector_contract_id = i.inject_injector_contract " + + " JOIN attack_patterns_kill_chain_phases apkcp ON apkcp.attack_pattern_id = icap.attack_pattern_id " + + " WHERE i.inject_id IN (:ids) " + + " GROUP BY i.inject_id " + + "), " + + "inject_platforms AS ( " + + " SELECT " + + " i.inject_id, " + + " array_union_agg(injcon.injector_contract_platforms) as platform_ids " + + " FROM injects i " + + " JOIN injectors_contracts injcon ON injcon.injector_contract_id = i.inject_injector_contract " + + " WHERE i.inject_id IN (:ids) " + + " GROUP BY i.inject_id " + + ") " + + "SELECT " + + " i.inject_id, " + + " ins.status_name, " + + " i.inject_scenario, " + + " COALESCE(it.team_ids, '{}') as inject_teams, " + + " COALESCE(ia.asset_ids, '{}') as inject_assets, " + + " COALESCE(iag.asset_group_ids, '{}') as inject_asset_groups, " + + " COALESCE(ie.expectation_ids, '{}') as inject_expectations, " + + " COALESCE(ic.communication_ids, '{}') as inject_communications, " + + " COALESCE(ikc.phase_ids, '{}') as inject_kill_chain_phases, " + + " COALESCE(ip.platform_ids, '{}') as inject_platforms " + + "FROM injects i " + + "LEFT JOIN injects_statuses ins ON ins.status_inject = i.inject_id " + + "LEFT JOIN inject_teams it ON it.inject_id = i.inject_id " + + "LEFT JOIN inject_assets ia ON ia.inject_id = i.inject_id " + + "LEFT JOIN inject_asset_groups iag ON iag.inject_id = i.inject_id " + + "LEFT JOIN inject_expectations ie ON ie.inject_id = i.inject_id " + + "LEFT JOIN inject_communications ic ON ic.inject_id = i.inject_id " + + "LEFT JOIN inject_kill_chains ikc ON ikc.inject_id = i.inject_id " + + "LEFT JOIN inject_platforms ip ON ip.inject_id = i.inject_id " + + "WHERE i.inject_id IN (:ids);", nativeQuery = true) List findRawByIds(@Param("ids") List ids); @@ -384,6 +435,10 @@ void deleteAllInjectsWithAttackPatternContractsByScenarioId( nativeQuery = true) void deleteAllByScenarioIdAndInjectorContract(String injectorContract, String scenarioId); + @EntityGraph(attributePaths = {"expectations", "injectorContract"}) + @Query("SELECT i FROM Inject i WHERE i.id IN :ids") + List findAllByIdWithExpectations(@Param("ids") List ids); + @Modifying @Query(value = "DELETE FROM injects WHERE inject_id = :id", nativeQuery = true) void deleteByIdNative(@Param("id") String id);