%%%------------------------------------------------------------------- %%% @author anlicheng %%% @copyright (C) 2026, %%% @doc %%% SSL transport for sdlan sessions. %%% @end %%%------------------------------------------------------------------- -module(sdlan_ssl_transport). -author("anlicheng"). -include("sdlan.hrl"). -include("sdlan_pb.hrl"). -behaviour(gen_statem). -behaviour(ranch_protocol). -define(SOCKET_ACTIVE_N, 100). %% Ranch protocol callback -export([start_link/4]). %% gen_statem callbacks -export([init/1, handle_event/4, terminate/3, code_change/4, callback_mode/0]). -record(state, { ref :: ranch:ref(), socket :: undefined | ssl:sslsocket(), transport :: module(), ok_msg :: atom(), closed_msg :: atom(), error_msg :: atom(), passive_msg :: atom(), max_packet_size = 16384, heartbeat_sec = 15, socket_active_n = ?SOCKET_ACTIVE_N, %% 累积器,用于处理协议framing的解析 buf = <<>>, session :: sdlan_session:state(), frames_recv = 0, bytes_recv = 0, close_reason = undefined }). %%%=================================================================== %%% Ranch protocol callback %%%=================================================================== start_link(Ref, Socket, Transport, Limits) -> gen_statem:start_link(?MODULE, [Ref, Socket, Transport, Limits], []). %%%=================================================================== %%% gen_statem callbacks %%%=================================================================== init([Ref, Socket, Transport, Limits]) -> MaxPacketSize = proplists:get_value(max_packet_size, Limits, 16384), HeartbeatSec = proplists:get_value(heartbeat_sec, Limits, 15), SocketActiveN = proplists:get_value(socket_active_n, Limits, proplists:get_value(stream_active_n, Limits, ?SOCKET_ACTIVE_N)), {OkMsg, ClosedMsg, ErrorMsg} = Transport:messages(), Session = sdlan_session:new(HeartbeatSec), {ok, handshaking, #state{ ref = Ref, socket = Socket, transport = Transport, ok_msg = OkMsg, closed_msg = ClosedMsg, error_msg = ErrorMsg, passive_msg = passive_msg(OkMsg), max_packet_size = MaxPacketSize, heartbeat_sec = HeartbeatSec, socket_active_n = SocketActiveN, session = Session }}. callback_mode() -> handle_event_function. handle_event(info, {handshake, Ref, Transport, Socket, Timeout}, handshaking, State = #state{ref = Ref, transport = Transport, max_packet_size = MaxPacketSize, heartbeat_sec = HeartbeatSec, socket_active_n = SocketActiveN}) -> case Transport:handshake(Socket, [], Timeout) of {ok, SslSocket} -> ok = Transport:setopts(SslSocket, [{mode, binary}, {active, SocketActiveN}]), WelcomePkt = welcome_packet(MaxPacketSize, HeartbeatSec), ssl_send(Transport, SslSocket, <>), logger:debug("[sdlan_ssl_transport] ssl handshake ok, send welcome"), {next_state, initialized, State#state{socket = SslSocket}}; {error, Reason} -> {stop, {ssl_handshake_failed, Reason}, State} end; handle_event(info, {OkMsg, Socket, Data}, _StateName, State = #state{socket = Socket, ok_msg = OkMsg, buf = Buf, max_packet_size = MaxPacketSize, bytes_recv = BytesRecv, frames_recv = FramesRecv}) when is_binary(Data) -> case decode_frames(<>, MaxPacketSize) of {error, Reason} -> {stop, Reason, State}; {ok, NBuf, Frames} -> Actions = [{next_event, internal, {frame, Frame}} || Frame <- Frames], {keep_state, State#state{buf = NBuf, bytes_recv = BytesRecv + byte_size(Data), frames_recv = FramesRecv + length(Frames)}, Actions} end; handle_event(info, {PassiveMsg, Socket}, _StateName, State = #state{socket = Socket, transport = Transport, passive_msg = PassiveMsg, socket_active_n = SocketActiveN}) -> ok = Transport:setopts(Socket, [{active, SocketActiveN}]), {keep_state, State}; handle_event(info, {ClosedMsg, Socket}, _StateName, State = #state{socket = Socket, closed_msg = ClosedMsg}) -> expected_stop(socket_closed, State); handle_event(info, {ErrorMsg, Socket, Reason}, _StateName, State = #state{socket = Socket, error_msg = ErrorMsg}) -> expected_stop({socket_error, Reason}, State); %% 处理内部的包消息 handle_event(internal, {frame, Frame}, StateName, State = #state{socket = Socket, transport = Transport, session = Session}) -> case sdlan_session:handle_frame(Frame, Session) of {ok, NSession, Packets} -> send_packets(Transport, Socket, Packets), {next_state, next_state_name(StateName, NSession), State#state{session = NSession}}; {stop, Reason, NSession, Packets} -> send_packets(Transport, Socket, Packets), {stop, Reason, State#state{session = NSession}} end; handle_event(info, {timeout, TimerRef, ping_ticker}, _StateName, State = #state{session = Session}) -> case sdlan_session:handle_timeout(TimerRef, Session) of {ok, NSession} -> {next_state, next_state_name(registered, NSession), State#state{session = NSession}}; {stop, Reason, NSession} -> expected_stop(Reason, State#state{session = NSession}) end; %% 发送事件信息 handle_event(cast, {send_event, Event}, _StateName, State = #state{socket = Socket, transport = Transport, session = Session}) -> case sdlan_session:send_event(Event, Session) of {ok, NSession, Packets} -> send_packets(Transport, Socket, Packets), {keep_state, State#state{session = NSession}}; {error, not_registered} -> keep_state_and_data end; %% 发送命令信息 handle_event(cast, {command, Ref, ReceiverPid, SubCommand}, _StateName, State = #state{socket = Socket, transport = Transport, session = Session}) -> case sdlan_session:command(Ref, ReceiverPid, SubCommand, Session) of {ok, NSession, Packets} -> send_packets(Transport, Socket, Packets), {keep_state, State#state{session = NSession}}; {error, not_registered} -> keep_state_and_data end; handle_event({call, From}, debug_info, StateName, State) -> {keep_state, State, [{reply, From, debug_info(StateName, State)}]}; handle_event(EventType, Info, StateName, State) -> logger:notice("[sdlan_ssl_transport] state: ~p, state_name: ~p, event_type: ~p, info: ~p", [State, StateName, EventType, Info]), keep_state_and_data. terminate(Reason, _StateName, #state{socket = Socket, transport = Transport, session = Session, close_reason = CloseReason}) -> Socket =/= undefined andalso catch Transport:close(Socket), logger:notice("[sdlan_ssl_transport] terminate closed with reason: ~p, close_reason: ~p", [Reason, CloseReason]), sdlan_session:close(Session), ok. code_change(_OldVsn, StateName, State = #state{}, _Extra) -> {ok, StateName, State}. %%%=================================================================== %%% Internal functions %%%=================================================================== welcome_packet(MaxPacketSize, HeartbeatSec) -> Ipv6Assist = case application:get_env(sdlan, ipv6_assist_info) of {ok, {V6Bytes, Port}} -> #'SDLV6Info' { v6 = V6Bytes, port = Port }; _ -> undefined end, sdlan_pb:encode_msg(#'SDLWelcome'{ version = 1, max_bidi_streams = 1, max_packet_size = MaxPacketSize, heartbeat_sec = HeartbeatSec, ipv6_assist = Ipv6Assist }). %% 有2种情况 %% 1. 收到了多个完整的请求 %% 2. 不完整,则不处理 -spec decode_frames(Buf :: binary(), MaxPacketSize :: integer()) -> {ok, RestBin::binary(), Frames :: list()} | {error, Reason :: any()}. decode_frames(Buf, MaxPacketSize) when is_binary(Buf) -> decode_frames0(Buf, MaxPacketSize, []). decode_frames0(<>, MaxPacketSize, _Frames) when Len > MaxPacketSize -> {error, frame_too_large}; decode_frames0(<>, MaxPacketSize, Frames) -> decode_frames0(Rest, MaxPacketSize, [Frame|Frames]); decode_frames0(Rest, _MaxPacketSize, Frames) -> {ok, Rest, lists:reverse(Frames)}. ssl_send(Transport, Socket, Packet) when is_binary(Packet) -> Len = byte_size(Packet), true = Len =< 65535, case Transport:send(Socket, <>) of ok -> incr_counter(ssl_frames_sent, 1), incr_counter(ssl_bytes_sent, Len + 2), ok; {error, Reason} -> exit({ssl_send_failed, Reason}) end. send_packets(Transport, Socket, Packets) -> lists:foreach(fun(Packet) -> ssl_send(Transport, Socket, Packet) end, Packets). expected_stop(Reason, State) -> logger:notice("[sdlan_ssl_transport] expected close: ~p", [Reason]), {stop, normal, State#state{close_reason = Reason}}. next_state_name(_StateName, Session) -> sdlan_session:state_name(Session). passive_msg(OkMsg) -> list_to_atom(atom_to_list(OkMsg) ++ "_passive"). debug_info(StateName, #state{ session = Session, frames_recv = FramesRecv, bytes_recv = BytesRecv, socket_active_n = SocketActiveN, heartbeat_sec = HeartbeatSec }) -> ProcInfo = maps:from_list(process_info(self(), [message_queue_len, memory, reductions])), SessionInfo = sdlan_session:debug_info(Session), maps:merge(SessionInfo, ProcInfo#{ state => StateName, session => SessionInfo, frames_recv => FramesRecv, frames_sent => get_counter(ssl_frames_sent), bytes_recv => BytesRecv, bytes_sent => get_counter(ssl_bytes_sent), socket_active_n => SocketActiveN, heartbeat_sec => HeartbeatSec }). incr_counter(Key, Inc) -> erlang:put(Key, get_counter(Key) + Inc). get_counter(Key) -> case erlang:get(Key) of undefined -> 0; Value when is_integer(Value) -> Value end.