sdlan/src/ssl/sdlan_ssl_transport.erl
2026-05-03 13:44:34 +08:00

262 lines
10 KiB
Erlang
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

%%%-------------------------------------------------------------------
%%% @author anlicheng
%%% @copyright (C) 2026, <COMPANY>
%%% @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, <<?PACKET_WELCOME, WelcomePkt/binary>>),
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(<<Buf/binary, Data/binary>>, 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(<<Len:16, _/binary>>, MaxPacketSize, _Frames) when Len > MaxPacketSize ->
{error, frame_too_large};
decode_frames0(<<Len:16, Frame:Len/binary, Rest/binary>>, 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, <<Len:16, Packet/binary>>) 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.