服务器管理组

EventLoopGroup bossGroup = new NioEventLoopGroup(1);
EventLoopGroup workerGroup = new NioEventLoopGroup();

ServerBootstrap bootstrap = new ServerBootstrap();

bootstrap.group(bossGroup, workerGroup)
	.channel(NioServerSocketChannel.class)
	.childHandler(new MqttBrokerChannelInitializer())
	.childOption(ChannelOption.TCP_NODELAY, true)
	.childOption(ChannelOption.SO_KEEPALIVE, true);

Channel serverChannel = bootstrap.bind(port).sync().channel();
serverChannel.closeFuture().sync();

MqttBrokerChannelInitializer

/** 为每一条设备连接创建独立的MQTT Pipeline和连接上下文。 */
public final class MqttBrokerChannelInitializer
        extends ChannelInitializer<SocketChannel> {

    /** Lesson11暂时把单个MQTT报文限制为1 MiB。 */
    public static final int MAX_MQTT_MESSAGE_SIZE = 1024 * 1024;

    @Override
    protected void initChannel(SocketChannel channel) {
        MqttConnectionContext.attach(channel);//这里创建的是“TCP连接上下文”,不是MQTT持久会话。

        ChannelPipeline pipeline = channel.pipeline();
        pipeline.addLast(
                "mqttDecoder",
                new MqttDecoder(MAX_MQTT_MESSAGE_SIZE)); //mqtt解码器,解析TCP中的mqtt内容
        pipeline.addLast("mqttEncoder", MqttEncoder.INSTANCE);//mqtt编码器,java方便使用
        pipeline.addLast(
                "mqttMessageHandler",
                new MqttBrokerMessageHandler());//自定义的流水线节点
    }
}

MQTT上下文

package psn.wyl.mqtt.broker;

import io.netty.channel.Channel;
import io.netty.handler.codec.mqtt.MqttConnectMessage;
import io.netty.util.AttributeKey;

/**
 * 一条Channel独享的MQTT连接状态。
 *
 * <p>它先放在Channel Attribute中,因此不会被其他连接共享。后续课程会在
 * 这个对象上继续加入认证状态、订阅、会话和QoS状态机。</p>
 */
public final class MqttConnectionContext {

    public static final AttributeKey<MqttConnectionContext> ATTRIBUTE_KEY =
            AttributeKey.valueOf(
                    MqttConnectionContext.class,
                    "connectionContext");

    private final String channelId;
    private final String remoteAddress;
    private final long tcpConnectedAtMillis;

    private long lastPacketAtMillis;
    private long receivedPacketCount;
    private boolean mqttConnected;
    private boolean disconnected;
    private String clientId;
    private String protocolName;
    private int protocolLevel;
    private int keepAliveSeconds;
    private boolean cleanSession;

    private MqttConnectionContext(Channel channel) {
        this.channelId = channel.id().asLongText();
        this.remoteAddress = String.valueOf(channel.remoteAddress());
        this.tcpConnectedAtMillis = System.currentTimeMillis();
        this.lastPacketAtMillis = tcpConnectedAtMillis;
    }

    public static MqttConnectionContext attach(Channel channel) {
        MqttConnectionContext created =
                new MqttConnectionContext(channel);
        MqttConnectionContext existing =
                channel.attr(ATTRIBUTE_KEY).setIfAbsent(created);
        return existing == null ? created : existing;
    }

    public static MqttConnectionContext get(Channel channel) {
        MqttConnectionContext context =
                channel.attr(ATTRIBUTE_KEY).get();
        return context == null ? attach(channel) : context;
    }

    public void recordPacket() {
        receivedPacketCount++;
        lastPacketAtMillis = System.currentTimeMillis();
    }

    public void recordConnect(MqttConnectMessage message) {
        clientId = message.payload().clientIdentifier();
        protocolName = message.variableHeader().name();
        protocolLevel = message.variableHeader().version();
        keepAliveSeconds =
                message.variableHeader().keepAliveTimeSeconds();
        cleanSession = message.variableHeader().isCleanSession();
        mqttConnected = true;
        disconnected = false;
    }

    public void recordDisconnect() {
        mqttConnected = false;
        disconnected = true;
    }

    public String channelId() {
        return channelId;
    }

    public String remoteAddress() {
        return remoteAddress;
    }

    public long tcpConnectedAtMillis() {
        return tcpConnectedAtMillis;
    }

    public long lastPacketAtMillis() {
        return lastPacketAtMillis;
    }

    public long receivedPacketCount() {
        return receivedPacketCount;
    }

    public boolean mqttConnected() {
        return mqttConnected;
    }

    public boolean disconnected() {
        return disconnected;
    }

    public String clientId() {
        return clientId;
    }

    public String protocolName() {
        return protocolName;
    }

    public int protocolLevel() {
        return protocolLevel;
    }

    public int keepAliveSeconds() {
        return keepAliveSeconds;
    }

    public boolean cleanSession() {
        return cleanSession;
    }
}

MqttBrokerMessageHandler

package psn.wyl.mqtt.broker;

import io.netty.channel.ChannelHandlerContext;
import io.netty.channel.SimpleChannelInboundHandler;
import io.netty.handler.codec.DecoderResult;
import io.netty.handler.codec.mqtt.MqttConnectMessage;
import io.netty.handler.codec.mqtt.MqttConnectReturnCode;
import io.netty.handler.codec.mqtt.MqttMessage;
import io.netty.handler.codec.mqtt.MqttMessageBuilders;
import io.netty.handler.codec.mqtt.MqttMessageType;

/**
 *
 * <p>MqttDecoder只负责把ByteBuf转换为MqttMessage;真正的协议状态检查、
 * 登录认证、主题路由和QoS必须由Broker Handler实现。本课只建立分发骨架。</p>
 */
public final class MqttBrokerMessageHandler
        extends SimpleChannelInboundHandler<MqttMessage> {

    @Override
    public void channelActive(ChannelHandlerContext ctx) {
        MqttConnectionContext connection =
                MqttConnectionContext.get(ctx.channel());
        ctx.fireChannelActive();
    }

    @Override
    protected void channelRead0(
            ChannelHandlerContext ctx,
            MqttMessage message) {

        MqttConnectionContext connection =
                MqttConnectionContext.get(ctx.channel());
        connection.recordPacket();

        DecoderResult decoderResult = message.decoderResult();
        if (decoderResult.isFailure()) {
            Throwable cause = decoderResult.cause();
            System.err.printf(
                    "[MQTT] decode failed channel=%s cause=%s%n",
                    connection.channelId(),
                    cause == null ? "unknown" : cause.getMessage());
            connection.recordDisconnect();
            ctx.close();
            return;
        }

        MqttMessageType messageType =
                message.fixedHeader().messageType();
        switch (messageType) {
            case CONNECT -> handleConnect(
                    ctx,
                    connection,
                    (MqttConnectMessage) message);
            case PINGREQ -> handlePingReq(ctx, connection);
            case DISCONNECT -> handleDisconnect(ctx, connection);
            default -> System.out.printf(
                    "[MQTT] received type=%s channel=%s; handler will be implemented later%n",
                    messageType,
                    connection.channelId());
        }
    }

    private void handleConnect(
            ChannelHandlerContext ctx,
            MqttConnectionContext connection,
            MqttConnectMessage message) {

        connection.recordConnect(message);
        System.out.printf(
                "[MQTT] CONNECT channel=%s clientId=%s protocol=%s/%d keepAlive=%ds cleanSession=%s%n",
                connection.channelId(),
                connection.clientId(),
                connection.protocolName(),
                connection.protocolLevel(),
                connection.keepAliveSeconds(),
                connection.cleanSession());

        ctx.writeAndFlush(
                MqttMessageBuilders.connAck()
                        .returnCode(
                                MqttConnectReturnCode.CONNECTION_ACCEPTED)
                        .sessionPresent(false)
                        .build());
    }

    private void handlePingReq(
            ChannelHandlerContext ctx,
            MqttConnectionContext connection) {

        System.out.printf(
                "[MQTT] PINGREQ channel=%s%n",
                connection.channelId());
        ctx.writeAndFlush(MqttMessage.PINGRESP);
    }

    private void handleDisconnect(
            ChannelHandlerContext ctx,
            MqttConnectionContext connection) {

        System.out.printf(
                "[MQTT] DISCONNECT channel=%s clientId=%s%n",
                connection.channelId(),
                connection.clientId());
        connection.recordDisconnect();
        ctx.close();
    }

    @Override
    public void channelInactive(ChannelHandlerContext ctx) {
        MqttConnectionContext connection =
                MqttConnectionContext.get(ctx.channel());
        connection.recordDisconnect();
        System.out.printf(
                "[MQTT] TCP disconnected channel=%s clientId=%s packets=%d%n",
                connection.channelId(),
                connection.clientId(),
                connection.receivedPacketCount());
        ctx.fireChannelInactive();
    }

    @Override
    public void exceptionCaught(
            ChannelHandlerContext ctx,
            Throwable cause) {

        MqttConnectionContext connection =
                MqttConnectionContext.get(ctx.channel());
        connection.recordDisconnect();
        System.err.printf(
                "[MQTT] exception channel=%s cause=%s%n",
                connection.channelId(),
                cause.getMessage());
        ctx.close();
    }
}