Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@

import com.fasterxml.jackson.databind.ObjectMapper;
import com.openframe.client.service.rmm.ScriptExecutionAcknowledgeService;
import com.openframe.data.document.delivery.DeliveryType;
import com.openframe.delivery.DeliveryTracker;
import com.openframe.data.nats.listener.AbstractJetStreamPushListener;
import com.openframe.data.nats.rmm.model.ScriptExecutionAcknowledgeMessage;
import io.nats.client.Connection;
Expand All @@ -17,15 +19,18 @@ public class ScriptExecutionAcknowledgeListener extends AbstractJetStreamPushLis

private final ObjectMapper objectMapper;
private final ScriptExecutionAcknowledgeService acknowledgeService;
private final DeliveryTracker deliveryTracker;

public ScriptExecutionAcknowledgeListener(
Connection natsConnection,
ObjectMapper objectMapper,
ScriptExecutionAcknowledgeService acknowledgeService
ScriptExecutionAcknowledgeService acknowledgeService,
DeliveryTracker deliveryTracker
) {
super(natsConnection);
this.objectMapper = objectMapper;
this.acknowledgeService = acknowledgeService;
this.deliveryTracker = deliveryTracker;
}

@Override
Expand Down Expand Up @@ -58,10 +63,23 @@ protected void handleMessage(Message message) {
String payload = new String(message.getData(), StandardCharsets.UTF_8);
try {
ScriptExecutionAcknowledgeMessage ack = objectMapper.readValue(payload, ScriptExecutionAcknowledgeMessage.class);
acknowledgeService.acknowledge(ack);
if (isDeliveryAck(ack)) {
deliveryTracker.acknowledge(ack.getType(), ack.getTargetId(), ack.getMachineId());
}
if (isScriptAck(ack)) {
acknowledgeService.acknowledge(ack);
}
message.ack();
} catch (Exception e) {
log.error("Unexpected error processing execution ack: {}", payload, e);
}
}

private static boolean isDeliveryAck(ScriptExecutionAcknowledgeMessage ack) {
return ack.getType() != null && ack.getTargetId() != null;
}

private static boolean isScriptAck(ScriptExecutionAcknowledgeMessage ack) {
return ack.getType() == null || ack.getType() == DeliveryType.SCRIPT_SCHEDULE;
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@
import com.openframe.client.exception.MachineNotFoundException;
import com.openframe.data.document.installedagents.InstalledAgent;
import com.openframe.data.document.tool.ConnectionStatus;
import com.openframe.data.document.delivery.DeliveryType;
import com.openframe.delivery.DeliveryTracker;
import com.openframe.data.repository.device.MachineRepository;
import com.openframe.data.repository.installedagents.InstalledAgentRepository;
import lombok.RequiredArgsConstructor;
Expand All @@ -21,6 +23,7 @@ public class InstalledAgentService {

private final InstalledAgentRepository installedAgentRepository;
private final MachineRepository machineRepository;
private final DeliveryTracker deliveryTracker;

@Transactional
public void addInstalledAgent(String machineId, String agentType, String version, boolean lastAttempt) {
Expand All @@ -36,6 +39,7 @@ public void addInstalledAgent(String machineId, String agentType, String version
installedAgent -> updateExistingInstalledAgent(installedAgent, version, machineId, agentType),
() -> addNewInstalledAgent(machineId, agentType, version)
);
deliveryTracker.complete(DeliveryType.TOOL_INSTALLATION, agentType, machineId);
}

@Transactional
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,114 @@
package com.openframe.client.listener.rmm;

import com.fasterxml.jackson.databind.ObjectMapper;
import com.openframe.client.service.rmm.ScriptExecutionAcknowledgeService;
import com.openframe.data.document.delivery.DeliveryType;
import com.openframe.delivery.DeliveryTracker;
import com.openframe.data.nats.rmm.model.ScriptExecutionAcknowledgeMessage;
import io.nats.client.Connection;
import io.nats.client.Message;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.ArgumentCaptor;
import org.mockito.Captor;
import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;

import static java.nio.charset.StandardCharsets.UTF_8;
import static org.assertj.core.api.Assertions.assertThat;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.verifyNoInteractions;
import static org.mockito.Mockito.when;

@ExtendWith(MockitoExtension.class)
class ScriptExecutionAcknowledgeListenerTest {

private static final String MACHINE_ID = "mach-42";
private static final String EXECUTION_ID = "exec-1";
private static final String TOOL_AGENT_ID = "tactical-agent";
private static final String LEGACY_SCRIPT_ACK =
"{\"executionId\":\"exec-1\",\"machineId\":\"mach-42\",\"scriptIds\":[\"s1\"]}";
private static final String TOOL_INSTALLATION_ACK =
"{\"type\":\"TOOL_INSTALLATION\",\"targetId\":\"tactical-agent\",\"machineId\":\"mach-42\"}";
private static final String SCRIPT_SCHEDULE_ACK =
"{\"type\":\"SCRIPT_SCHEDULE\",\"targetId\":\"exec-1\",\"executionId\":\"exec-1\",\"machineId\":\"mach-42\",\"scriptIds\":[\"s1\"]}";
private static final String MALFORMED = "not json";

@Mock private Connection natsConnection;
@Mock private ScriptExecutionAcknowledgeService acknowledgeService;
@Mock private DeliveryTracker deliveryTracker;
@Mock private Message message;

@Captor private ArgumentCaptor<ScriptExecutionAcknowledgeMessage> ackCaptor;

private ScriptExecutionAcknowledgeListener listener;

@BeforeEach
void setUp() {
listener = new ScriptExecutionAcknowledgeListener(natsConnection, new ObjectMapper(), acknowledgeService, deliveryTracker);
}

@Test
void handleMessage_legacyScriptAck_scriptServiceOnly() {
// setup
stubPayload(LEGACY_SCRIPT_ACK);

// execution
listener.handleMessage(message);

// verifications
verify(acknowledgeService).acknowledge(ackCaptor.capture());
assertThat(ackCaptor.getValue().getExecutionId()).isEqualTo(EXECUTION_ID);
verifyNoInteractions(deliveryTracker);
verify(message).ack();
}

@Test
void handleMessage_toolInstallationAck_trackerOnly() {
// setup
stubPayload(TOOL_INSTALLATION_ACK);

// execution
listener.handleMessage(message);

// verifications
verify(deliveryTracker).acknowledge(DeliveryType.TOOL_INSTALLATION, TOOL_AGENT_ID, MACHINE_ID);
verifyNoInteractions(acknowledgeService);
verify(message).ack();
}

@Test
void handleMessage_scriptScheduleAckWithType_trackerAndScriptService() {
// setup
stubPayload(SCRIPT_SCHEDULE_ACK);

// execution
listener.handleMessage(message);

// verifications
verify(deliveryTracker).acknowledge(DeliveryType.SCRIPT_SCHEDULE, EXECUTION_ID, MACHINE_ID);
verify(acknowledgeService).acknowledge(ackCaptor.capture());
assertThat(ackCaptor.getValue().getExecutionId()).isEqualTo(EXECUTION_ID);
verify(message).ack();
}

@Test
void handleMessage_malformedPayload_leftUnacked() {
// setup
stubPayload(MALFORMED);

// execution
listener.handleMessage(message);

// verifications
verify(message, never()).ack();
verifyNoInteractions(acknowledgeService);
verifyNoInteractions(deliveryTracker);
}

private void stubPayload(String json) {
when(message.getData()).thenReturn(json.getBytes(UTF_8));
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
package com.openframe.client.service;

import com.openframe.client.exception.MachineNotFoundException;
import com.openframe.delivery.DeliveryTracker;
import com.openframe.data.document.device.Machine;
import com.openframe.data.document.installedagents.InstalledAgent;
import com.openframe.data.document.delivery.DeliveryType;
import com.openframe.data.document.tool.ConnectionStatus;
import com.openframe.data.repository.device.MachineRepository;
import com.openframe.data.repository.installedagents.InstalledAgentRepository;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.ArgumentCaptor;
import org.mockito.Captor;
import org.mockito.InjectMocks;
import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;

import java.util.Optional;

import static org.assertj.core.api.Assertions.assertThat;
import static org.junit.jupiter.api.Assertions.assertThrows;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.verifyNoInteractions;
import static org.mockito.Mockito.when;

@ExtendWith(MockitoExtension.class)
class InstalledAgentServiceTest {

private static final String MACHINE_ID = "mach-42";
private static final String AGENT_TYPE = "tactical-agent";
private static final String OLD_VERSION = "1.0.0";
private static final String VERSION = "1.2.3";

@Mock private InstalledAgentRepository installedAgentRepository;
@Mock private MachineRepository machineRepository;
@Mock private DeliveryTracker deliveryTracker;

@Captor private ArgumentCaptor<InstalledAgent> installedAgentCaptor;

@InjectMocks private InstalledAgentService service;

private Machine machine;
private InstalledAgent existing;

@BeforeEach
void setUp() {
machine = new Machine();
machine.setMachineId(MACHINE_ID);
existing = new InstalledAgent();
existing.setMachineId(MACHINE_ID);
existing.setAgentType(AGENT_TYPE);
existing.setVersion(OLD_VERSION);
existing.setStatus(ConnectionStatus.DISCONNECTED);
}

@Test
void addInstalledAgent_newAgent_savedAndDeliveryCompleted() {
// setup
when(machineRepository.findByMachineId(MACHINE_ID)).thenReturn(Optional.of(machine));
when(installedAgentRepository.findByMachineIdAndAgentType(MACHINE_ID, AGENT_TYPE)).thenReturn(Optional.empty());

// execution
service.addInstalledAgent(MACHINE_ID, AGENT_TYPE, VERSION, false);

// verifications
verify(installedAgentRepository).save(installedAgentCaptor.capture());
assertThat(installedAgentCaptor.getValue().getVersion()).isEqualTo(VERSION);
assertThat(installedAgentCaptor.getValue().getStatus()).isEqualTo(ConnectionStatus.CONNECTED);
verify(deliveryTracker).complete(DeliveryType.TOOL_INSTALLATION, AGENT_TYPE, MACHINE_ID);
}

@Test
void addInstalledAgent_existingAgent_versionUpdatedAndDeliveryCompleted() {
// setup
when(machineRepository.findByMachineId(MACHINE_ID)).thenReturn(Optional.of(machine));
when(installedAgentRepository.findByMachineIdAndAgentType(MACHINE_ID, AGENT_TYPE)).thenReturn(Optional.of(existing));

// execution
service.addInstalledAgent(MACHINE_ID, AGENT_TYPE, VERSION, false);

// verifications
assertThat(existing.getVersion()).isEqualTo(VERSION);
assertThat(existing.getStatus()).isEqualTo(ConnectionStatus.CONNECTED);
verify(installedAgentRepository).save(existing);
verify(deliveryTracker).complete(DeliveryType.TOOL_INSTALLATION, AGENT_TYPE, MACHINE_ID);
}

@Test
void addInstalledAgent_unknownMachine_throwsAndDeliveryUntouched() {
// setup
when(machineRepository.findByMachineId(MACHINE_ID)).thenReturn(Optional.empty());

// execution
MachineNotFoundException ex = assertThrows(MachineNotFoundException.class,
() -> service.addInstalledAgent(MACHINE_ID, AGENT_TYPE, VERSION, false));

// verifications
assertThat(ex.getMessage()).contains(MACHINE_ID);
verifyNoInteractions(deliveryTracker);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,7 @@ public DeliveryRequest<ToolInstallationMessage> request(Seed seed) {
@Override
public void publish(String machineId, ToolInstallationMessage payload) {
String subject = format(SUBJECT_TEMPLATE, machineId);
natsMessagePublisher.publishPersistent(subject, payload);
natsMessagePublisher.publish(subject, payload);
}

@Override
Expand Down
Loading