diff --git a/bifromq-session-dict/bifromq-session-dict-server/src/main/java/org/apache/bifromq/sessiondict/server/SessionRegistry.java b/bifromq-session-dict/bifromq-session-dict-server/src/main/java/org/apache/bifromq/sessiondict/server/SessionRegistry.java index ef761d6cd..80c7704af 100644 --- a/bifromq-session-dict/bifromq-session-dict-server/src/main/java/org/apache/bifromq/sessiondict/server/SessionRegistry.java +++ b/bifromq-session-dict/bifromq-session-dict-server/src/main/java/org/apache/bifromq/sessiondict/server/SessionRegistry.java @@ -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(); @@ -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); @@ -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); @@ -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 @@ -98,6 +109,14 @@ 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 @@ -105,7 +124,11 @@ 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); diff --git a/bifromq-session-dict/bifromq-session-dict-server/src/test/java/org/apache/bifromq/sessiondict/server/SessionRegistryTest.java b/bifromq-session-dict/bifromq-session-dict-server/src/test/java/org/apache/bifromq/sessiondict/server/SessionRegistryTest.java index cdc15d26a..75caf449f 100644 --- a/bifromq-session-dict/bifromq-session-dict-server/src/test/java/org/apache/bifromq/sessiondict/server/SessionRegistryTest.java +++ b/bifromq-session-dict/bifromq-session-dict-server/src/test/java/org/apache/bifromq/sessiondict/server/SessionRegistryTest.java @@ -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 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 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 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);