Skip to content
Open
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 @@ -37,9 +37,11 @@
import java.util.TreeMap;
import java.util.concurrent.ConcurrentSkipListMap;
import java.util.concurrent.atomic.AtomicInteger;
import lombok.extern.slf4j.Slf4j;
import org.apache.bifromq.sessiondict.rpc.proto.ServerRedirection;
import org.apache.bifromq.type.ClientInfo;

@Slf4j
class SessionRegistry implements ISessionRegistry {
private static final ServerRedirection NO_MOVE =
ServerRedirection.newBuilder().setType(ServerRedirection.Type.NO_MOVE).build();
Expand All @@ -56,6 +58,13 @@ class SessionRegistry implements ISessionRegistry {
public void add(ClientInfo sessionOwner, ISessionRegister register) {
String tenantId = sessionOwner.getTenantId();
MqttClientKey clientKey = MqttClientKey.from(sessionOwner);
// The kick callback synchronously calls back into remove() on the same thread (the kicked
// register unregisters itself from its Quit path), and ConcurrentHashMap.compute must not
// be re-entered for the same key — the nested remove could delete the mapping this add
// has just installed, losing the new session from the dictionary. Collect the kick and
// perform it only after leaving the compute block.
ISessionRegister[] kickedRegister = new ISessionRegister[1];
ClientInfo[] kickedOwner = new ClientInfo[1];
tenantSessions.compute(tenantId, (k, v) -> {
if (v == null) {
v = new ConcurrentSkipListMap<>(ClientKeyComparator);
Expand All @@ -69,9 +78,10 @@ public void add(ClientInfo sessionOwner, ISessionRegister register) {
if (!prevSessionOwner.equals(sessionOwner)) {
ISessionRegister prevSessionRegister = clientRegisterMap.remove(prevSessionOwner);
clientRegisterMap.put(sessionOwner, register);
// kick previous session owner
// kick previous session owner after leaving the compute block
assert prevSessionRegister != null;
prevSessionRegister.kick(tenantId, prevSessionOwner, sessionOwner, NO_MOVE);
kickedRegister[0] = prevSessionRegister;
kickedOwner[0] = prevSessionOwner;
if (isPersistent(sessionOwner) && !isPersistent(prevSessionOwner)) {
// kicked by a persistent session
SessionCounter sessionCounter = sessionCounters.get(tenantId);
Expand All @@ -83,7 +93,8 @@ public void add(ClientInfo sessionOwner, ISessionRegister register) {
} else {
ISessionRegister prevSessionRegister = clientRegisterMap.put(sessionOwner, register);
if (prevSessionRegister != null && prevSessionRegister != register) {
prevSessionRegister.kick(tenantId, prevSessionOwner, sessionOwner, NO_MOVE);
kickedRegister[0] = prevSessionRegister;
kickedOwner[0] = prevSessionOwner;
}
}
// ignore duplicated add
Expand All @@ -98,14 +109,26 @@ public void add(ClientInfo sessionOwner, ISessionRegister register) {
}
return v;
});
if (kickedRegister[0] != null) {
try {
kickedRegister[0].kick(tenantId, kickedOwner[0], sessionOwner, NO_MOVE);
} catch (RuntimeException e) {
// best-effort notification: the previous register's stream may already be closed
log.debug("Kicked session register ignored exception: {}", e.getMessage());
}
}
}

@Override
public void remove(ClientInfo sessionOwner, ISessionRegister register) {
String tenantId = sessionOwner.getTenantId();
MqttClientKey clientKey = MqttClientKey.from(sessionOwner);
tenantSessions.computeIfPresent(tenantId, (k, v) -> {
boolean s1 = v.remove(clientKey, sessionOwner);
// Only drop the dict entry if the register being removed is still the live one —
// after a same-owner takeover the old register's (possibly delayed) teardown must
// not delete the mapping installed by the new register, nor drift the counters.
boolean s1 = register == clientRegisterMap.get(sessionOwner)
&& v.remove(clientKey, sessionOwner);
boolean s2 = clientRegisterMap.remove(sessionOwner, register);
if (s1) {
SessionCounter sessionCounter = sessionCounters.get(tenantId);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -218,6 +218,117 @@ public void testTransientKicksPersistent() {
assertGaugeValue(tenantId1, MqttLivePersistentSessionGauge, 0.0);
}

@Test
public void testKickCallbackReenteringRemoveKeepsNewSession() {
// The real register's kick() synchronously unregisters itself via regListener,
// which calls remove() back on the same thread — the compute block must not be
// re-entered, and the new session must survive in the dictionary.
ISessionRegister registerMock1 = Mockito.mock(ISessionRegister.class);
ISessionRegister registerMock2 = Mockito.mock(ISessionRegister.class);

ClientInfo sessionOwner1 = ClientInfo.newBuilder()
.setTenantId(tenantId1)
.putMetadata(MQTT_USER_ID_KEY, "user1")
.putMetadata(MQTT_CLIENT_ID_KEY, "client1")
.putMetadata(MQTT_CLIENT_SESSION_TYPE, MQTT_CLIENT_SESSION_TYPE_T_VALUE)
.putMetadata(MQTT_CHANNEL_ID_KEY, "channel1")
.build();

ClientInfo sessionOwner2 = ClientInfo.newBuilder()
.setTenantId(tenantId1)
.putMetadata(MQTT_USER_ID_KEY, "user1")
.putMetadata(MQTT_CLIENT_ID_KEY, "client1")
.putMetadata(MQTT_CLIENT_SESSION_TYPE, MQTT_CLIENT_SESSION_TYPE_T_VALUE)
.putMetadata(MQTT_CHANNEL_ID_KEY, "channel2")
.build();

Mockito.doAnswer(invocation -> {
sessionRegistry.remove(sessionOwner1, registerMock1);
return null;
}).when(registerMock1).kick(org.mockito.ArgumentMatchers.any(),
org.mockito.ArgumentMatchers.any(), org.mockito.ArgumentMatchers.any(),
org.mockito.ArgumentMatchers.any());

sessionRegistry.add(sessionOwner1, registerMock1);
sessionRegistry.add(sessionOwner2, registerMock2);

Mockito.verify(registerMock1).kick(eq(tenantId1), eq(sessionOwner1), eq(sessionOwner2),
eq(ServerRedirection.newBuilder().setType(ServerRedirection.Type.NO_MOVE).build()));

Optional<ClientInfo> retrieved = sessionRegistry.get(tenantId1, "user1", "client1");
assertTrue(retrieved.isPresent());
assertEquals(sessionOwner2, retrieved.get());
assertGaugeValue(tenantId1, MqttConnectionGauge, 1.0);
}

@Test
public void testReAddSameOwnerWithReenteringKickKeepsNewSession() {
// Same owner re-registered with a different register: the nested remove() from the
// kick callback used to delete the mapping this add had just installed, silently
// losing the session from the dictionary (subsequent kill could not find it).
ISessionRegister registerMock1 = Mockito.mock(ISessionRegister.class);
ISessionRegister registerMock2 = Mockito.mock(ISessionRegister.class);

ClientInfo sessionOwner = ClientInfo.newBuilder()
.setTenantId(tenantId1)
.putMetadata(MQTT_USER_ID_KEY, "user1")
.putMetadata(MQTT_CLIENT_ID_KEY, "client1")
.putMetadata(MQTT_CLIENT_SESSION_TYPE, MQTT_CLIENT_SESSION_TYPE_T_VALUE)
.putMetadata(MQTT_CHANNEL_ID_KEY, "channel1")
.build();

Mockito.doAnswer(invocation -> {
sessionRegistry.remove(sessionOwner, registerMock1);
return null;
}).when(registerMock1).kick(org.mockito.ArgumentMatchers.any(),
org.mockito.ArgumentMatchers.any(), org.mockito.ArgumentMatchers.any(),
org.mockito.ArgumentMatchers.any());

sessionRegistry.add(sessionOwner, registerMock1);
sessionRegistry.add(sessionOwner, registerMock2);

Mockito.verify(registerMock1).kick(eq(tenantId1), eq(sessionOwner), eq(sessionOwner),
eq(ServerRedirection.newBuilder().setType(ServerRedirection.Type.NO_MOVE).build()));

Optional<ClientInfo> retrieved = sessionRegistry.get(tenantId1, "user1", "client1");
assertTrue(retrieved.isPresent());
assertEquals(sessionOwner, retrieved.get());
assertGaugeValue(tenantId1, MqttConnectionGauge, 1.0);
}

@Test
public void testDelayedTeardownOfTakenOverRegisterKeepsNewSession() {
// Production sequence after a same-owner takeover: the old register's stream is closed
// by the client only after it receives the Quit, so the old register's teardown
// (remove) can land well after the new register has been added. The remove must not
// delete the dictionary entry installed by the new register, nor drift the counters.
ISessionRegister registerMock1 = Mockito.mock(ISessionRegister.class);
ISessionRegister registerMock2 = Mockito.mock(ISessionRegister.class);

ClientInfo sessionOwner = ClientInfo.newBuilder()
.setTenantId(tenantId1)
.putMetadata(MQTT_USER_ID_KEY, "user1")
.putMetadata(MQTT_CLIENT_ID_KEY, "client1")
.putMetadata(MQTT_CLIENT_SESSION_TYPE, MQTT_CLIENT_SESSION_TYPE_T_VALUE)
.putMetadata(MQTT_CHANNEL_ID_KEY, "channel1")
.build();

sessionRegistry.add(sessionOwner, registerMock1);
sessionRegistry.add(sessionOwner, registerMock2); // takeover, kicks registerMock1

// delayed teardown of the OLD register (what doFinally does when its stream closes)
sessionRegistry.remove(sessionOwner, registerMock1);

Optional<ClientInfo> retrieved = sessionRegistry.get(tenantId1, "user1", "client1");
assertTrue(retrieved.isPresent());
assertEquals(sessionOwner, retrieved.get());
assertGaugeValue(tenantId1, MqttConnectionGauge, 1.0);

// teardown of the LIVE register still removes the entry
sessionRegistry.remove(sessionOwner, registerMock2);
assertTrue(sessionRegistry.get(tenantId1, "user1", "client1").isEmpty());
}

@Test
public void testRemoveSession() {
ISessionRegister registerMock = Mockito.mock(ISessionRegister.class);
Expand Down