From e8a85d0d93e866bf74a37d0fe3076117268cf7a7 Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Sun, 3 May 2026 13:24:25 +0800 Subject: [PATCH] fix transport --- src/quic/sdlan_quic_channel.erl | 515 +--------------------------- src/quic/sdlan_quic_channel_sup.erl | 8 +- src/quic/sdlan_quic_server.erl | 4 +- src/quic/sdlan_quic_transport.erl | 323 +++++++++++++++++ src/quic/sdlan_session.erl | 349 +++++++++++++++++++ 5 files changed, 687 insertions(+), 512 deletions(-) create mode 100644 src/quic/sdlan_quic_transport.erl create mode 100644 src/quic/sdlan_session.erl diff --git a/src/quic/sdlan_quic_channel.erl b/src/quic/sdlan_quic_channel.erl index 894363c..7b933b9 100644 --- a/src/quic/sdlan_quic_channel.erl +++ b/src/quic/sdlan_quic_channel.erl @@ -2,540 +2,43 @@ %%% @author anlicheng %%% @copyright (C) 2026, %%% @doc -%%% +%%% Compatibility facade for the QUIC channel API. %%% @end -%%% Created : 11. 2月 2026 23:00 %%%------------------------------------------------------------------- -module(sdlan_quic_channel). -author("anlicheng"). --include("sdlan.hrl"). --include("sdlan_pb.hrl"). - --behaviour(gen_statem). - -%% 心跳包监测机制 --define(PING_TICKER, 15000). --define(STREAM_ACTIVE_N, 100). --define(METRICS_TICKER, 60000). - -%% 注册失败的的错误码 - -%% 网络错误 --define(NAK_NETWORK_FAULT, 4). -%% 内部错误 --define(NAK_INTERNAL_FAULT, 5). %% API -export([start_link/2]). -export([accept_stream/1, send_event/2, command/4, stop/2, debug_info/1]). -export([test_rules/2]). -%% gen_statem callbacks --export([init/1, handle_event/4, terminate/3, code_change/4, callback_mode/0]). - --record(state, { - conn :: quicer:connection_handle(), - %% 最大包大小 - max_packet_size = 16384, - %% 心跳间隔 - heartbeat_sec = 15, - ping_timer :: undefined | reference(), - stream_active_n = ?STREAM_ACTIVE_N, - - stream :: undefined | quicer:stream_handle(), - %% 累积器,用于处理协议framing的解析 - buf = <<>>, - - client_id :: undefined | binary(), - network_id = 0 :: integer(), - %% 网络相关信息id - network_pid :: undefined | pid(), - %% mac地址 - mac :: undefined | binary(), - ip = 0 :: integer(), - - %% 建立请求和响应的对应关系 - pkt_id = 1, - %% #{pkt_id => {Ref, ReceiverPid}} - pending_commands = #{}, - - ping_counter = 0, - frames_recv = 0, - bytes_recv = 0, - - %% 离线回调函数 - offline_cb :: undefined | fun(), - close_reason = undefined -}). - %%%=================================================================== %%% API %%%=================================================================== %% 测试规则函数 test_rules(SrcIdentityId, DstIdentityId) when is_integer(SrcIdentityId), is_integer(DstIdentityId) -> - {ok, Rules} = get_rules(SrcIdentityId, DstIdentityId), - logger:debug("[sdlan_channel] test_rules policy_request src_identity_id: ~p, dst_identity_id: ~p, rules: ~p", [SrcIdentityId, DstIdentityId, Rules]), - iolist_to_binary(lists:map(fun({Proto, Port}) -> <> end, Rules)). + sdlan_session:test_rules(SrcIdentityId, DstIdentityId). -spec send_event(Pid :: pid(), Event :: binary()) -> no_return(). send_event(Pid, ProtobufEvent) when is_pid(Pid), is_binary(ProtobufEvent) -> - gen_statem:cast(Pid, {send_event, ProtobufEvent}). + sdlan_quic_transport:send_event(Pid, ProtobufEvent). -spec command(Pid :: pid(), Ref :: reference(), ReceiverPid :: pid(), {Tag :: atom(), SubCommand :: any()}) -> no_return(). command(Pid, Ref, ReceiverPid, SubCommand) when is_pid(Pid), is_pid(ReceiverPid) -> - gen_statem:cast(Pid, {command, Ref, ReceiverPid, SubCommand}). + sdlan_quic_transport:command(Pid, Ref, ReceiverPid, SubCommand). accept_stream(Pid) when is_pid(Pid) -> - gen_statem:cast(Pid, accept_stream). + sdlan_quic_transport:accept_stream(Pid). -spec stop(Pid :: pid(), Reason :: term()) -> ok. stop(Pid, Reason) when is_pid(Pid) -> - gen_statem:stop(Pid, Reason, 2000). + sdlan_quic_transport:stop(Pid, Reason). debug_info(Pid) when is_pid(Pid) -> - gen_statem:call(Pid, debug_info). + sdlan_quic_transport:debug_info(Pid). -%% @doc Creates a gen_statem process which calls Module:init/1 to -%% initialize. To ensure a synchronized start-up procedure, this -%% function does not return until Module:init/1 has returned. +%% @doc Creates a transport process. Kept for existing supervisor specs. start_link(Conn, Limits) when is_list(Limits) -> - gen_statem:start_link(?MODULE, [Conn, Limits], []). - -%%%=================================================================== -%%% gen_statem callbacks -%%%=================================================================== - -%% @private -%% @doc Whenever a gen_statem is started using gen_statem:start/[3,4] or -%% gen_statem:start_link/[3,4], this function is called by the new -%% process to initialize. -init([Conn, Limits]) -> - MaxPacketSize = proplists:get_value(max_packet_size, Limits, 16384), - HeartbeatSec = proplists:get_value(heartbeat_sec, Limits, 15), - StreamActiveN = proplists:get_value(stream_active_n, Limits, ?STREAM_ACTIVE_N), - {ok, initializing, #state{conn = Conn, max_packet_size = MaxPacketSize, heartbeat_sec = HeartbeatSec, stream_active_n = StreamActiveN}}. - -%% @private -%% @doc This function is called by a gen_statem when it needs to find out -%% the callback mode of the callback module. -callback_mode() -> - handle_event_function. - -%% @private -%% @doc If callback_mode is handle_event_function, then whenever a -%% gen_statem receives an event from call/2, cast/2, or as a normal -%% process message, this function is called. - -handle_event(cast, accept_stream, initializing, State=#state{conn = Conn, stream_active_n = StreamActiveN}) -> - logger:debug("[sdlan_quic_channel] call do_init of conn: ~p", [Conn]), - case quicer:async_accept_stream(Conn, #{active => StreamActiveN}) of - {ok, _} -> - {next_state, waiting_stream, State}; - {error, Reason} -> - {stop, {accept_stream_failed, Reason}, State} - end; - -%% 处理收到的quic消息 -handle_event(info, {quic, dgram_state_changed, Conn, Opts = #{dgram_send_enabled := true}}, _, State=#state{conn = Conn}) -> - logger:debug("[sdlan_quic_channel] dgram_state_changed, opts: ~p", [Opts]), - {keep_state, State}; - -handle_event(info, {quic, new_stream, Stream, Opts}, waiting_stream, State=#state{max_packet_size = MaxPacketSize, heartbeat_sec = HeartbeatSec}) -> - logger:debug("[sdlan_quic_channel] call new_stream: ~p, opts: ~p", [Stream, Opts]), - Ipv6Assist = case application:get_env(sdlan, ipv6_assist_info) of - {ok, {V6Bytes, Port}} -> - #'SDLV6Info' { - v6 = V6Bytes, - port = Port - }; - _ -> - undefined - end, - %% 发送欢迎消息 - WelcomePkt = sdlan_pb:encode_msg(#'SDLWelcome'{ - version = 1, - max_bidi_streams = 1, - max_packet_size = MaxPacketSize, - heartbeat_sec = HeartbeatSec, - ipv6_assist = Ipv6Assist - }), - quic_send(Stream, <>), - logger:debug("[sdlan_quic_channel] get stream: ~p, send welcome", [Stream]), - - {next_state, initialized, State#state{stream = Stream}}; - -handle_event(info, {quic, new_stream, Stream, Opts}, _StateName, State) -> - logger:warning("[sdlan_quic_channel] reject unexpected stream: ~p, opts: ~p", [Stream, Opts]), - quicer:close_stream(Stream, 1000), - {keep_state, State}; - -handle_event(info, {quic, stream_closed, Stream, Props}, _StateName, State = #state{stream = Stream}) -> - expected_stop({stream_closed, Props}, State); - -handle_event(info, {quic, peer_send_shutdown, Stream, _Props}, _StateName, State = #state{stream = Stream}) -> - expected_stop(peer_send_shutdown, State); - -handle_event(info, {quic, peer_send_aborted, Stream, ErrorCode}, _StateName, State = #state{stream = Stream}) -> - expected_stop({peer_send_aborted, ErrorCode}, State); - -handle_event(info, {quic, peer_receive_aborted, Stream, ErrorCode}, _StateName, State = #state{stream = Stream}) -> - expected_stop({peer_receive_aborted, ErrorCode}, State); - -handle_event(info, {quic, send_shutdown_complete, Stream, _Props}, _StateName, State = #state{stream = Stream}) -> - expected_stop(connection_shutdown, State); - -handle_event(info, {quic, passive, Stream, _Props}, _StateName, State = #state{stream = Stream, stream_active_n = StreamActiveN}) -> - ok = quicer:setopt(Stream, active, StreamActiveN), - {keep_state, State}; - -handle_event(info, {quic, closed, Conn, Props}, _StateName, State = #state{conn = Conn}) -> - expected_stop({connection_closed, Props}, State); - -handle_event(info, {quic, transport_shutdown, Conn, Props}, _StateName, State = #state{conn = Conn}) -> - expected_stop({transport_shutdown, Props}, State); - -handle_event(info, {quic, shutdown, Conn, ErrorCode}, _StateName, State = #state{conn = Conn}) -> - expected_stop({connection_shutdown_by_peer, ErrorCode}, State); - -%% 处理quicer相关的信息, 需要转换成内部能够识别的frame消息 -handle_event(info, {quic, Data, Stream, _Props}, _StateName, State = #state{stream = Stream, 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], - %logger:debug("[sdlan_quic_channel] get frames: ~p", [Frames]), - {keep_state, State#state{buf = NBuf, bytes_recv = BytesRecv + byte_size(Data), frames_recv = FramesRecv + length(Frames)}, Actions} - end; - -%% 处理内部的包消息 -handle_event(internal, {frame, <>}, initialized, State=#state{stream = Stream}) -> - #'SDLRegisterSuper'{ - client_id = ClientId, network_id = NetworkId, mac = Mac, ip = Ip, mask_len = MaskLen, - hostname = HostName, pub_key = PubKey, access_token = AccessToken} = sdlan_pb:decode_msg(Body, 'SDLRegisterSuper'), - - true = (Mac =/= <<>> andalso PubKey =/= <<>> andalso ClientId =/= <<>>), - %% Mac地址不能是广播地址 - true = not (sdlan_util:is_multicast_mac(Mac) orelse sdlan_util:is_broadcast_mac(Mac)), - - MacBinStr = sdlan_util:format_mac(Mac), - IpAddr = sdlan_util:int_to_ipv4(Ip), - Params = #{ - <<"network_id">> => NetworkId, - <<"client_id">> => ClientId, - <<"mac">> => MacBinStr, - <<"ip">> => IpAddr, - <<"mask_len">> => MaskLen, - <<"hostname">> => HostName, - <<"access_token">> => AccessToken - }, - %% 参数检查 - logger:debug("[sdlan_quic_channel] client_id: ~p, ip: ~p, mac: ~p, host_name: ~p, access_token: ~p, network_id: ~p", - [ClientId, Ip, Mac, HostName, AccessToken, NetworkId]), - - case sdlan_api:auth_access_token(Params) of - {ok, #{<<"result">> := <<"ok">>}} -> - %% 建立到network的对应关系 - case sdlan_network:get_pid(NetworkId) of - NetworkPid when is_pid(NetworkPid) -> - {ok, Algorithm, Key, RegionId, SessionToken} = sdlan_network:attach(NetworkPid, self(), ClientId, Mac, Ip, HostName), - RsaPubKey = sdlan_cipher:rsa_pem_decode(PubKey), - RegisterSuperAck = sdlan_pb:encode_msg(#'SDLRegisterSuperAck'{ - algorithm = Algorithm, - key = rsa_encode(Key, RsaPubKey), - region_id = RegionId, - session_token = SessionToken - }), - - %% 发送确认信息 - quic_send(Stream, <>), - %% 设置节点的在线状态 - Result = sdlan_api:set_node_status(#{ - <<"network_id">> => NetworkId, - <<"client_id">> => ClientId, - <<"access_token">> => AccessToken, - <<"status">> => 1 - }), - logger:debug("[sdlan_quic_channel] client_id: ~p, set none online result is: ~p", [ClientId, Result]), - - OfflineCb = fun() -> - Result = sdlan_api:set_node_status(#{ - <<"network_id">> => NetworkId, - <<"client_id">> => ClientId, - <<"access_token">> => AccessToken, - <<"status">> => 0 - }) - end, - {next_state, registered, schedule_ping(State#state{network_id = NetworkId, network_pid = NetworkPid, client_id = ClientId, mac = Mac, ip = Ip, offline_cb = OfflineCb})}; - undefined -> - logger:warning("[sdlan_quic_channel] client_id: ~p, register get error: network not found", [ClientId]), - quic_send(Stream, register_nak_reply(?NAK_INTERNAL_FAULT, <<"Internal Error">>)), - {stop, normal, State} - end; - {ok, #{<<"error">> := #{<<"code">> := Code, <<"message">> := Message}}} -> - logger:warning("[sdlan_quic_channel] network_id: ~p, client_id: ~p, register get error: ~ts, error_code: ~p", [NetworkId, ClientId, Message, Code]), - quic_send(Stream, register_nak_reply(Code, Message)), - {stop, normal, State}; - {error, Reason} -> - logger:warning("[sdlan_quic_channel] network_id: ~p, client_id: ~p, register get error: ~p", [NetworkId, ClientId, Reason]), - quic_send(Stream, register_nak_reply(?NAK_NETWORK_FAULT, <<"Network Error">>)), - {stop, normal, State} - end; - -handle_event(internal, {frame, <>}, registered, #state{stream = Stream, network_pid = NetworkPid, mac = SrcMac}) when is_pid(NetworkPid) -> - #'SDLQueryInfo'{dst_mac = DstMac} = sdlan_pb:decode_msg(Body, 'SDLQueryInfo'), - case sdlan_network:peer_info(NetworkPid, SrcMac, DstMac) of - error -> - logger:debug("[sdlan_channel] query_info src_mac is: ~p, dst_mac: ~p, nat_peer not found", - [sdlan_util:format_mac(SrcMac), sdlan_util:format_mac(DstMac)]), - - EmptyResponse = sdlan_pb:encode_msg(#'SDLPeerInfo'{ - dst_mac = DstMac, - v4_info = undefined, - v6_info = undefined - }), - quic_send(Stream, <>), - keep_state_and_data; - {ok, {NatPeer = {{Ip0, Ip1, Ip2, Ip3}, NatPort}, NatType}, V6Info} -> - logger:debug("[sdlan_channel] query_info src_mac is: ~p, dst_mac: ~p, nat_peer: ~p", - [sdlan_util:format_mac(SrcMac), sdlan_util:format_mac(DstMac), NatPeer]), - - PeerInfo = sdlan_pb:encode_msg(#'SDLPeerInfo'{ - dst_mac = DstMac, - v4_info = #'SDLV4Info' { - port = NatPort, - v4 = <>, - nat_type = NatType - }, - v6_info = V6Info - }), - quic_send(Stream, <>), - keep_state_and_data - end; - -%% arp查询 -handle_event(internal, {frame, <>}, registered, #state{stream = Stream, network_id = NetworkId, network_pid = NetworkPid}) when is_pid(NetworkPid) -> - #'SDLArpRequest'{target_ip = TargetIp, origin_ip = OriginIp, context = Context} = sdlan_pb:decode_msg(Body, 'SDLArpRequest'), - case sdlan_network:arp_request(NetworkPid, TargetIp) of - error -> - logger:debug("[sdlan_channel] network: ~p, arp_request target_ip: ~p, mac not found", [NetworkId, sdlan_util:int_to_ipv4(TargetIp)]), - EmptyArpResponsePkt = sdlan_pb:encode_msg(#'SDLArpResponse'{ - target_ip = TargetIp, - target_mac = <<>>, - origin_ip = OriginIp, - context = Context - }), - quic_send(Stream, <>), - keep_state_and_data; - {ok, Mac} -> - logger:debug("[sdlan_channel] network: ~p, arp_request target_ip: ~p, mac: ~p", [NetworkId, sdlan_util:int_to_ipv4(TargetIp), sdlan_util:format_mac(Mac)]), - ArpResponsePkt = sdlan_pb:encode_msg(#'SDLArpResponse'{ - target_ip = TargetIp, - target_mac = Mac, - origin_ip = OriginIp, - context = Context - }), - quic_send(Stream, <>), - keep_state_and_data - end; - -handle_event(internal, {frame, <>}, registered, #state{stream = Stream, network_pid = NetworkPid}) when is_pid(NetworkPid) -> - maybe - #'SDLPolicyRequest'{src_identity_id = SrcIdentityId, dst_identity_id = DstIdentityId, version = Version} ?= sdlan_pb:decode_msg(Body, 'SDLPolicyRequest'), - - {ok, Rules} = get_rules(SrcIdentityId, DstIdentityId), - logger:debug("[sdlan_channel] policy_request src_identity_id: ~p, dst_identity_id: ~p, rules: ~p", [SrcIdentityId, DstIdentityId, Rules]), - - RuleBin = iolist_to_binary(lists:map(fun({Proto, Port}) -> <> end, Rules)), - PolicyResponsePkt = sdlan_pb:encode_msg(#'SDLPolicyResponse'{ - src_identity_id = SrcIdentityId, - dst_identity_id = DstIdentityId, - version = Version, - rules = RuleBin - }), - quic_send(Stream, <>) - end, - keep_state_and_data; - -%% 处理命令的响应逻辑 -handle_event(internal, {frame, <>}, registered, State=#state{pending_commands = PendingCommands}) -> - maybe - CommandAck = sdlan_pb:decode_msg(Body, 'SDLCommandAck'), - #'SDLCommandAck'{pkt_id = PktId} ?= CommandAck, - - {{Ref, ReceiverPid}, RestPendingCommands} ?= maps:take(PktId, PendingCommands), - case is_process_alive(ReceiverPid) of - true -> - ReceiverPid ! {quic_command_ack, Ref, CommandAck}; - false -> - ok - end, - {keep_state, State#state{pending_commands = RestPendingCommands}} - else _ -> - keep_state_and_data - end; - -handle_event(internal, {frame, <>}, _StateName, State = #state{stream = Stream, ping_counter = PingCounter, client_id = ClientId}) -> - quic_send(Stream, <>), - logger:warning("[sdlan_channel] get ping: ~p", [ClientId]), - {keep_state, State#state{ping_counter = PingCounter + 1}}; - -%% 取消注册 -handle_event(internal, {frame, <>}, registered, State=#state{client_id = ClientId, mac = Mac, network_pid = NetworkPid}) when is_pid(NetworkPid) -> - logger:warning("[sdlan_channel] unregister client_id: ~p", [ClientId]), - sdlan_network:unregister(NetworkPid, ClientId, Mac), - {stop, normal, State}; - -handle_event(info, {timeout, TimerRef, ping_ticker}, _, State = #state{ping_timer = TimerRef, client_id = ClientId, ping_counter = PingCounter}) -> - case PingCounter > 0 of - true -> - {keep_state, schedule_ping(State#state{ping_counter = 0, ping_timer = undefined})}; - false -> - logger:debug("[sdlan_channel] client_id: ~p, ping losted", [ClientId]), - expected_stop(heartbeat_timeout, State#state{ping_counter = 0, ping_timer = undefined}) - end; - -handle_event(info, {timeout, _TimerRef, ping_ticker}, _, State) -> - {keep_state, State}; - -%% 发送指令信息 -handle_event(cast, {send_event, Event}, registered, #state{stream = Stream}) -> - quic_send(Stream, <>), - keep_state_and_data; - -%% 发送命令信息 -handle_event(cast, {command, Ref, ReceiverPid, SubCommand}, registered, State=#state{stream = Stream, pkt_id = PktId, pending_commands = PendingCommands, client_id = ClientId}) -> - CommandPkt = sdlan_pb:encode_msg(#'SDLCommand'{ - pkt_id = PktId, - command = SubCommand - }), - logger:debug("[sdlan_channel] client_id: ~p, will send Command: ~p", [ClientId, SubCommand]), - - quic_send(Stream, <>), - {keep_state, State#state{pkt_id = PktId + 1, pending_commands = maps:put(PktId, {Ref, ReceiverPid}, PendingCommands)}}; - -handle_event({call, From}, debug_info, StateName, State) -> - {keep_state, State, [{reply, From, debug_info(StateName, State)}]}; - -handle_event(info, {'EXIT', _, _}, _StateName, State) -> - expected_stop(connection_closed, State); - -handle_event(EventType, Info, StateName, State) -> - logger:notice("[sdlan_quic_channel] state: ~p, state_name: ~p, event_type: ~p, info: ~p", [State, StateName, EventType, Info]), - keep_state_and_data. - -%% @private -%% @doc This function is called by a gen_statem when it is about to -%% terminate. It should be the opposite of Module:init/1 and do any -%% necessary cleaning up. When it returns, the gen_statem terminates with -%% Reason. The return value is ignored. -terminate(Reason, _StateName, _State = #state{conn = Conn, stream = Stream, client_id = ClientId, offline_cb = OfflineCb, close_reason = CloseReason}) -> - Stream /= undefined andalso quicer:close_stream(Stream, 1000), - quicer:close_connection(Conn), - logger:notice("[sdlan_quic_conn] client_id: ~p, terminate closed with reason: ~p, close_reason: ~p", [ClientId, Reason, CloseReason]), - %% 触发客户端的离线逻辑 - is_function(OfflineCb) andalso OfflineCb(), - ok. - -%% @private -%% @doc Convert process state when code is changed -code_change(_OldVsn, StateName, State = #state{}, _Extra) -> - {ok, StateName, State}. - -%%%=================================================================== -%%% Internal functions -%%%=================================================================== - -%% 有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)}. - --spec register_nak_reply(ErrorCode :: integer(), ErrorMsg :: binary()) -> binary(). -register_nak_reply(ErrorCode, ErrorMsg) when is_integer(ErrorCode), is_binary(ErrorMsg) -> - RegisterNakReply = sdlan_pb:encode_msg(#'SDLRegisterSuperNak'{ - error_code = ErrorCode, - error_message = ErrorMsg - }), - <>. - -rsa_encode(PlainText, RsaPubKey) when is_binary(PlainText) -> - iolist_to_binary(sdlan_cipher:rsa_encrypt(PlainText, RsaPubKey)). - --spec quic_send(Stream :: quicer:stream_handle(), Packet :: binary()) -> no_return(). -quic_send(Stream, Packet) when is_binary(Packet) -> - Len = byte_size(Packet), - true = Len =< 65535, - case quicer:send(Stream, <>) of - {ok, _} -> - incr_counter(quic_frames_sent, 1), - incr_counter(quic_bytes_sent, Len + 2), - ok; - {error, Reason} -> - exit({quic_send_failed, Reason}) - end. - --spec get_rules(SrcIdentityId :: integer(), DstIdentityId :: integer()) -> {ok, [{Proto :: integer(), Port :: integer()}]}. -get_rules(SrcIdentityId, DstIdentityId) when is_integer(SrcIdentityId), is_integer(DstIdentityId) -> - SrcPolicyIds = identity_policy_ets:get_policies(SrcIdentityId), - DstPolicyIds = identity_policy_ets:get_policies(DstIdentityId), - rule_ets:get_rules(SrcPolicyIds, DstPolicyIds). - -expected_stop(Reason, State) -> - logger:notice("[sdlan_quic_channel] expected close: ~p", [Reason]), - {stop, normal, State#state{close_reason = Reason}}. - -schedule_ping(State = #state{heartbeat_sec = HeartbeatSec}) -> - State#state{ping_timer = erlang:start_timer(heartbeat_ms(HeartbeatSec), self(), ping_ticker)}. - -heartbeat_ms(HeartbeatSec) when is_integer(HeartbeatSec), HeartbeatSec > 0 -> - HeartbeatSec * 1000; -heartbeat_ms(_) -> - ?PING_TICKER. - -debug_info(StateName, #state{ - client_id = ClientId, - network_id = NetworkId, - mac = Mac, - ip = Ip, - pending_commands = PendingCommands, - frames_recv = FramesRecv, - bytes_recv = BytesRecv, - stream_active_n = StreamActiveN, - heartbeat_sec = HeartbeatSec -}) -> - ProcInfo = maps:from_list(process_info(self(), [message_queue_len, memory, reductions])), - ProcInfo#{ - state => StateName, - client_id => ClientId, - network_id => NetworkId, - mac => Mac, - ip => Ip, - pending_commands => maps:size(PendingCommands), - frames_recv => FramesRecv, - frames_sent => get_counter(quic_frames_sent), - bytes_recv => BytesRecv, - bytes_sent => get_counter(quic_bytes_sent), - stream_active_n => StreamActiveN, - 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. + sdlan_quic_transport:start_link(Conn, Limits). diff --git a/src/quic/sdlan_quic_channel_sup.erl b/src/quic/sdlan_quic_channel_sup.erl index 92e720b..f920070 100644 --- a/src/quic/sdlan_quic_channel_sup.erl +++ b/src/quic/sdlan_quic_channel_sup.erl @@ -47,12 +47,12 @@ init([]) -> SupFlags = #{strategy => simple_one_for_one, intensity => 0, period => 1}, AChild = #{ - id => sdlan_quic_channel, - start => {'sdlan_quic_channel', start_link, []}, + id => sdlan_quic_transport, + start => {'sdlan_quic_transport', start_link, []}, restart => temporary, shutdown => 2000, type => worker, - modules => ['sdlan_quic_channel'] + modules => ['sdlan_quic_transport'] }, {ok, {SupFlags, [AChild]}}. @@ -62,4 +62,4 @@ init([]) -> -spec start_channel(NConn :: quicer:connection_handle(), Limits :: proplists:proplist()) -> supervisor:startchild_ret(). start_channel(NConn, Limits) when is_list(Limits) -> - supervisor:start_child(?MODULE, [NConn, Limits]). \ No newline at end of file + supervisor:start_child(?MODULE, [NConn, Limits]). diff --git a/src/quic/sdlan_quic_server.erl b/src/quic/sdlan_quic_server.erl index d1b5ad7..41e458e 100644 --- a/src/quic/sdlan_quic_server.erl +++ b/src/quic/sdlan_quic_server.erl @@ -72,10 +72,10 @@ loop_accept(L, Limits, AcceptorId) -> logger:debug("[sdlan_quic_server] conn: ~p, handshake success, channel pid: ~p", [NConn, ChannelPid]), case quicer:controlling_process(NConn, ChannelPid) of ok -> - sdlan_quic_channel:accept_stream(ChannelPid); + sdlan_quic_transport:accept_stream(ChannelPid); {error, Reason} -> logger:warning("[sdlan_quic_server] conn: ~p, controlling_process failed: ~p", [NConn, Reason]), - sdlan_quic_channel:stop(ChannelPid, {controlling_process_failed, Reason}), + sdlan_quic_transport:stop(ChannelPid, {controlling_process_failed, Reason}), quicer:close_connection(NConn) end; Error -> diff --git a/src/quic/sdlan_quic_transport.erl b/src/quic/sdlan_quic_transport.erl new file mode 100644 index 0000000..595720c --- /dev/null +++ b/src/quic/sdlan_quic_transport.erl @@ -0,0 +1,323 @@ +%%%------------------------------------------------------------------- +%%% @author anlicheng +%%% @copyright (C) 2026, +%%% @doc +%%% QUIC transport for sdlan sessions. +%%% @end +%%%------------------------------------------------------------------- +-module(sdlan_quic_transport). +-author("anlicheng"). +-include("sdlan.hrl"). +-include("sdlan_pb.hrl"). + +-behaviour(gen_statem). + +%% 心跳包监测机制 +-define(STREAM_ACTIVE_N, 100). + +%% API +-export([start_link/2]). +-export([accept_stream/1, send_event/2, command/4, stop/2, debug_info/1]). + +%% gen_statem callbacks +-export([init/1, handle_event/4, terminate/3, code_change/4, callback_mode/0]). + +-record(state, { + conn :: quicer:connection_handle(), + %% 最大包大小 + max_packet_size = 16384, + %% 心跳间隔 + heartbeat_sec = 15, + stream_active_n = ?STREAM_ACTIVE_N, + + stream :: undefined | quicer:stream_handle(), + %% 累积器,用于处理协议framing的解析 + buf = <<>>, + + session :: sdlan_session:state(), + + frames_recv = 0, + bytes_recv = 0, + + close_reason = undefined +}). + +%%%=================================================================== +%%% API +%%%=================================================================== + +-spec send_event(Pid :: pid(), Event :: binary()) -> no_return(). +send_event(Pid, ProtobufEvent) when is_pid(Pid), is_binary(ProtobufEvent) -> + gen_statem:cast(Pid, {send_event, ProtobufEvent}). + +-spec command(Pid :: pid(), Ref :: reference(), ReceiverPid :: pid(), {Tag :: atom(), SubCommand :: any()}) -> no_return(). +command(Pid, Ref, ReceiverPid, SubCommand) when is_pid(Pid), is_pid(ReceiverPid) -> + gen_statem:cast(Pid, {command, Ref, ReceiverPid, SubCommand}). + +accept_stream(Pid) when is_pid(Pid) -> + gen_statem:cast(Pid, accept_stream). + +-spec stop(Pid :: pid(), Reason :: term()) -> ok. +stop(Pid, Reason) when is_pid(Pid) -> + gen_statem:stop(Pid, Reason, 2000). + +debug_info(Pid) when is_pid(Pid) -> + gen_statem:call(Pid, debug_info). + +%% @doc Creates a gen_statem process which calls Module:init/1 to +%% initialize. To ensure a synchronized start-up procedure, this +%% function does not return until Module:init/1 has returned. +start_link(Conn, Limits) when is_list(Limits) -> + gen_statem:start_link(?MODULE, [Conn, Limits], []). + +%%%=================================================================== +%%% gen_statem callbacks +%%%=================================================================== + +%% @private +%% @doc Whenever a gen_statem is started using gen_statem:start/[3,4] or +%% gen_statem:start_link/[3,4], this function is called by the new +%% process to initialize. +init([Conn, Limits]) -> + MaxPacketSize = proplists:get_value(max_packet_size, Limits, 16384), + HeartbeatSec = proplists:get_value(heartbeat_sec, Limits, 15), + StreamActiveN = proplists:get_value(stream_active_n, Limits, ?STREAM_ACTIVE_N), + Session = sdlan_session:new(HeartbeatSec), + {ok, initializing, #state{ + conn = Conn, + max_packet_size = MaxPacketSize, + heartbeat_sec = HeartbeatSec, + stream_active_n = StreamActiveN, + session = Session + }}. + +%% @private +%% @doc This function is called by a gen_statem when it needs to find out +%% the callback mode of the callback module. +callback_mode() -> + handle_event_function. + +%% @private +%% @doc If callback_mode is handle_event_function, then whenever a +%% gen_statem receives an event from call/2, cast/2, or as a normal +%% process message, this function is called. + +handle_event(cast, accept_stream, initializing, State = #state{conn = Conn, stream_active_n = StreamActiveN}) -> + logger:debug("[sdlan_quic_transport] call do_init of conn: ~p", [Conn]), + case quicer:async_accept_stream(Conn, #{active => StreamActiveN}) of + {ok, _} -> + {next_state, waiting_stream, State}; + {error, Reason} -> + {stop, {accept_stream_failed, Reason}, State} + end; + +%% 处理收到的quic消息 +handle_event(info, {quic, dgram_state_changed, Conn, Opts = #{dgram_send_enabled := true}}, _, State = #state{conn = Conn}) -> + logger:debug("[sdlan_quic_transport] dgram_state_changed, opts: ~p", [Opts]), + {keep_state, State}; + +handle_event(info, {quic, new_stream, Stream, Opts}, waiting_stream, State = #state{max_packet_size = MaxPacketSize, heartbeat_sec = HeartbeatSec}) -> + logger:debug("[sdlan_quic_transport] call new_stream: ~p, opts: ~p", [Stream, Opts]), + Ipv6Assist = case application:get_env(sdlan, ipv6_assist_info) of + {ok, {V6Bytes, Port}} -> + #'SDLV6Info' { + v6 = V6Bytes, + port = Port + }; + _ -> + undefined + end, + %% 发送欢迎消息 + WelcomePkt = sdlan_pb:encode_msg(#'SDLWelcome'{ + version = 1, + max_bidi_streams = 1, + max_packet_size = MaxPacketSize, + heartbeat_sec = HeartbeatSec, + ipv6_assist = Ipv6Assist + }), + quic_send(Stream, <>), + logger:debug("[sdlan_quic_transport] get stream: ~p, send welcome", [Stream]), + + {next_state, initialized, State#state{stream = Stream}}; + +handle_event(info, {quic, new_stream, Stream, Opts}, _StateName, State) -> + logger:warning("[sdlan_quic_transport] reject unexpected stream: ~p, opts: ~p", [Stream, Opts]), + quicer:close_stream(Stream, 1000), + {keep_state, State}; + +handle_event(info, {quic, stream_closed, Stream, Props}, _StateName, State = #state{stream = Stream}) -> + expected_stop({stream_closed, Props}, State); + +handle_event(info, {quic, peer_send_shutdown, Stream, _Props}, _StateName, State = #state{stream = Stream}) -> + expected_stop(peer_send_shutdown, State); + +handle_event(info, {quic, peer_send_aborted, Stream, ErrorCode}, _StateName, State = #state{stream = Stream}) -> + expected_stop({peer_send_aborted, ErrorCode}, State); + +handle_event(info, {quic, peer_receive_aborted, Stream, ErrorCode}, _StateName, State = #state{stream = Stream}) -> + expected_stop({peer_receive_aborted, ErrorCode}, State); + +handle_event(info, {quic, send_shutdown_complete, Stream, _Props}, _StateName, State = #state{stream = Stream}) -> + expected_stop(connection_shutdown, State); + +handle_event(info, {quic, passive, Stream, _Props}, _StateName, State = #state{stream = Stream, stream_active_n = StreamActiveN}) -> + ok = quicer:setopt(Stream, active, StreamActiveN), + {keep_state, State}; + +handle_event(info, {quic, closed, Conn, Props}, _StateName, State = #state{conn = Conn}) -> + expected_stop({connection_closed, Props}, State); + +handle_event(info, {quic, transport_shutdown, Conn, Props}, _StateName, State = #state{conn = Conn}) -> + expected_stop({transport_shutdown, Props}, State); + +handle_event(info, {quic, shutdown, Conn, ErrorCode}, _StateName, State = #state{conn = Conn}) -> + expected_stop({connection_shutdown_by_peer, ErrorCode}, State); + +%% 处理quicer相关的信息, 需要转换成内部能够识别的frame消息 +handle_event(info, {quic, Data, Stream, _Props}, _StateName, + State = #state{stream = Stream, 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(internal, {frame, Frame}, StateName, State = #state{stream = Stream, session = Session}) -> + case sdlan_session:handle_frame(Frame, Session) of + {ok, NSession, Packets} -> + send_packets(Stream, Packets), + {next_state, next_state_name(StateName, NSession), State#state{session = NSession}}; + {stop, Reason, NSession, Packets} -> + send_packets(Stream, 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{stream = Stream, session = Session}) -> + case sdlan_session:send_event(Event, Session) of + {ok, NSession, Packets} -> + send_packets(Stream, 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{stream = Stream, session = Session}) -> + case sdlan_session:command(Ref, ReceiverPid, SubCommand, Session) of + {ok, NSession, Packets} -> + send_packets(Stream, 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(info, {'EXIT', _, _}, _StateName, State) -> + expected_stop(connection_closed, State); + +handle_event(EventType, Info, StateName, State) -> + logger:notice("[sdlan_quic_transport] state: ~p, state_name: ~p, event_type: ~p, info: ~p", [State, StateName, EventType, Info]), + keep_state_and_data. + +%% @private +%% @doc This function is called by a gen_statem when it is about to +%% terminate. It should be the opposite of Module:init/1 and do any +%% necessary cleaning up. When it returns, the gen_statem terminates with +%% Reason. The return value is ignored. +terminate(Reason, _StateName, _State = #state{conn = Conn, stream = Stream, session = Session, close_reason = CloseReason}) -> + Stream /= undefined andalso quicer:close_stream(Stream, 1000), + quicer:close_connection(Conn), + logger:notice("[sdlan_quic_transport] terminate closed with reason: ~p, close_reason: ~p", [Reason, CloseReason]), + sdlan_session:close(Session), + ok. + +%% @private +%% @doc Convert process state when code is changed +code_change(_OldVsn, StateName, State = #state{}, _Extra) -> + {ok, StateName, State}. + +%%%=================================================================== +%%% Internal functions +%%%=================================================================== + +%% 有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)}. + +-spec quic_send(Stream :: quicer:stream_handle(), Packet :: binary()) -> no_return(). +quic_send(Stream, Packet) when is_binary(Packet) -> + Len = byte_size(Packet), + true = Len =< 65535, + case quicer:send(Stream, <>) of + {ok, _} -> + incr_counter(quic_frames_sent, 1), + incr_counter(quic_bytes_sent, Len + 2), + ok; + {error, Reason} -> + exit({quic_send_failed, Reason}) + end. + +send_packets(Stream, Packets) -> + lists:foreach(fun(Packet) -> quic_send(Stream, Packet) end, Packets). + +expected_stop(Reason, State) -> + logger:notice("[sdlan_quic_transport] expected close: ~p", [Reason]), + {stop, normal, State#state{close_reason = Reason}}. + +next_state_name(_StateName, Session) -> + sdlan_session:state_name(Session). + +debug_info(StateName, #state{ + session = Session, + frames_recv = FramesRecv, + bytes_recv = BytesRecv, + stream_active_n = StreamActiveN, + 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(quic_frames_sent), + bytes_recv => BytesRecv, + bytes_sent => get_counter(quic_bytes_sent), + stream_active_n => StreamActiveN, + 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. diff --git a/src/quic/sdlan_session.erl b/src/quic/sdlan_session.erl new file mode 100644 index 0000000..35b29b6 --- /dev/null +++ b/src/quic/sdlan_session.erl @@ -0,0 +1,349 @@ +%%%------------------------------------------------------------------- +%%% @author anlicheng +%%% @copyright (C) 2026, +%%% @doc +%%% Transport independent sdlan session logic. +%%% @end +%%%------------------------------------------------------------------- +-module(sdlan_session). +-author("anlicheng"). +-include("sdlan.hrl"). +-include("sdlan_pb.hrl"). + +%% API +-export([new/1, state_name/1, handle_frame/2, handle_timeout/2]). +-export([send_event/2, command/4, close/1, debug_info/1]). +-export([test_rules/2]). +-export_type([state/0]). + +%% 心跳包监测机制 +-define(PING_TICKER, 15000). + +%% 注册失败的的错误码 + +%% 网络错误 +-define(NAK_NETWORK_FAULT, 4). +%% 内部错误 +-define(NAK_INTERNAL_FAULT, 5). + +-record(state, { + status = initialized :: initialized | registered, + %% 心跳间隔 + heartbeat_sec = 15, + ping_timer :: undefined | reference(), + + client_id :: undefined | binary(), + network_id = 0 :: integer(), + %% 网络相关信息id + network_pid :: undefined | pid(), + %% mac地址 + mac :: undefined | binary(), + ip = 0 :: integer(), + + %% 建立请求和响应的对应关系 + pkt_id = 1, + %% #{pkt_id => {Ref, ReceiverPid}} + pending_commands = #{}, + + ping_counter = 0, + + %% 离线回调函数 + offline_cb :: undefined | fun() +}). + +-type state() :: #state{}. + +%%%=================================================================== +%%% API +%%%=================================================================== + +-spec new(HeartbeatSec :: integer()) -> #state{}. +new(HeartbeatSec) -> + #state{heartbeat_sec = HeartbeatSec}. + +%% 测试规则函数 +test_rules(SrcIdentityId, DstIdentityId) when is_integer(SrcIdentityId), is_integer(DstIdentityId) -> + {ok, Rules} = get_rules(SrcIdentityId, DstIdentityId), + logger:debug("[sdlan_session] test_rules policy_request src_identity_id: ~p, dst_identity_id: ~p, rules: ~p", [SrcIdentityId, DstIdentityId, Rules]), + iolist_to_binary(lists:map(fun({Proto, Port}) -> <> end, Rules)). + +-spec state_name(State :: #state{}) -> initialized | registered. +state_name(#state{status = Status}) -> + Status. + +-spec handle_frame(Frame :: binary(), State :: #state{}) -> + {ok, NewState :: #state{}, Packets :: [binary()]} | + {stop, Reason :: term(), NewState :: #state{}, Packets :: [binary()]}. +handle_frame(<>, State = #state{status = initialized}) -> + handle_register_super(Body, State); + +handle_frame(<>, State = #state{status = registered, network_pid = NetworkPid, mac = SrcMac}) when is_pid(NetworkPid) -> + #'SDLQueryInfo'{dst_mac = DstMac} = sdlan_pb:decode_msg(Body, 'SDLQueryInfo'), + case sdlan_network:peer_info(NetworkPid, SrcMac, DstMac) of + error -> + logger:debug("[sdlan_session] query_info src_mac is: ~p, dst_mac: ~p, nat_peer not found", + [sdlan_util:format_mac(SrcMac), sdlan_util:format_mac(DstMac)]), + + EmptyResponse = sdlan_pb:encode_msg(#'SDLPeerInfo'{ + dst_mac = DstMac, + v4_info = undefined, + v6_info = undefined + }), + {ok, State, [<>]}; + {ok, {NatPeer = {{Ip0, Ip1, Ip2, Ip3}, NatPort}, NatType}, V6Info} -> + logger:debug("[sdlan_session] query_info src_mac is: ~p, dst_mac: ~p, nat_peer: ~p", + [sdlan_util:format_mac(SrcMac), sdlan_util:format_mac(DstMac), NatPeer]), + + PeerInfo = sdlan_pb:encode_msg(#'SDLPeerInfo'{ + dst_mac = DstMac, + v4_info = #'SDLV4Info' { + port = NatPort, + v4 = <>, + nat_type = NatType + }, + v6_info = V6Info + }), + {ok, State, [<>]} + end; + +%% arp查询 +handle_frame(<>, State = #state{status = registered, network_id = NetworkId, network_pid = NetworkPid}) when is_pid(NetworkPid) -> + #'SDLArpRequest'{target_ip = TargetIp, origin_ip = OriginIp, context = Context} = sdlan_pb:decode_msg(Body, 'SDLArpRequest'), + case sdlan_network:arp_request(NetworkPid, TargetIp) of + error -> + logger:debug("[sdlan_session] network: ~p, arp_request target_ip: ~p, mac not found", [NetworkId, sdlan_util:int_to_ipv4(TargetIp)]), + EmptyArpResponsePkt = sdlan_pb:encode_msg(#'SDLArpResponse'{ + target_ip = TargetIp, + target_mac = <<>>, + origin_ip = OriginIp, + context = Context + }), + {ok, State, [<>]}; + {ok, Mac} -> + logger:debug("[sdlan_session] network: ~p, arp_request target_ip: ~p, mac: ~p", [NetworkId, sdlan_util:int_to_ipv4(TargetIp), sdlan_util:format_mac(Mac)]), + ArpResponsePkt = sdlan_pb:encode_msg(#'SDLArpResponse'{ + target_ip = TargetIp, + target_mac = Mac, + origin_ip = OriginIp, + context = Context + }), + {ok, State, [<>]} + end; + +handle_frame(<>, State = #state{status = registered, network_pid = NetworkPid}) when is_pid(NetworkPid) -> + Packets = maybe + #'SDLPolicyRequest'{src_identity_id = SrcIdentityId, dst_identity_id = DstIdentityId, version = Version} ?= sdlan_pb:decode_msg(Body, 'SDLPolicyRequest'), + + {ok, Rules} = get_rules(SrcIdentityId, DstIdentityId), + logger:debug("[sdlan_session] policy_request src_identity_id: ~p, dst_identity_id: ~p, rules: ~p", [SrcIdentityId, DstIdentityId, Rules]), + + RuleBin = iolist_to_binary(lists:map(fun({Proto, Port}) -> <> end, Rules)), + PolicyResponsePkt = sdlan_pb:encode_msg(#'SDLPolicyResponse'{ + src_identity_id = SrcIdentityId, + dst_identity_id = DstIdentityId, + version = Version, + rules = RuleBin + }), + [<>] + else _ -> + [] + end, + {ok, State, Packets}; + +%% 处理命令的响应逻辑 +handle_frame(<>, State = #state{status = registered, pending_commands = PendingCommands}) -> + maybe + CommandAck = sdlan_pb:decode_msg(Body, 'SDLCommandAck'), + #'SDLCommandAck'{pkt_id = PktId} ?= CommandAck, + + {{Ref, ReceiverPid}, RestPendingCommands} ?= maps:take(PktId, PendingCommands), + case is_process_alive(ReceiverPid) of + true -> + ReceiverPid ! {quic_command_ack, Ref, CommandAck}; + false -> + ok + end, + {ok, State#state{pending_commands = RestPendingCommands}, []} + else _ -> + {ok, State, []} + end; + +handle_frame(<>, State = #state{ping_counter = PingCounter, client_id = ClientId}) -> + logger:warning("[sdlan_session] get ping: ~p", [ClientId]), + {ok, State#state{ping_counter = PingCounter + 1}, [<>]}; + +%% 取消注册 +handle_frame(<>, State = #state{status = registered, client_id = ClientId, mac = Mac, network_pid = NetworkPid}) when is_pid(NetworkPid) -> + logger:warning("[sdlan_session] unregister client_id: ~p", [ClientId]), + sdlan_network:unregister(NetworkPid, ClientId, Mac), + {stop, normal, State, []}; + +handle_frame(Frame, State) -> + logger:notice("[sdlan_session] unexpected frame: ~p, status: ~p", [Frame, State#state.status]), + {ok, State, []}. + +-spec handle_timeout(TimerRef :: reference(), State :: #state{}) -> + {ok, NewState :: #state{}} | {stop, Reason :: term(), NewState :: #state{}}. +handle_timeout(TimerRef, State = #state{ping_timer = TimerRef, client_id = ClientId, ping_counter = PingCounter}) -> + case PingCounter > 0 of + true -> + {ok, schedule_ping(State#state{ping_counter = 0, ping_timer = undefined})}; + false -> + logger:debug("[sdlan_session] client_id: ~p, ping losted", [ClientId]), + {stop, heartbeat_timeout, State#state{ping_counter = 0, ping_timer = undefined}} + end; +handle_timeout(_TimerRef, State) -> + {ok, State}. + +-spec send_event(Event :: binary(), State :: #state{}) -> {ok, NewState :: #state{}, Packets :: [binary()]} | {error, not_registered}. +send_event(Event, State = #state{status = registered}) when is_binary(Event) -> + {ok, State, [<>]}; +send_event(_Event, _State) -> + {error, not_registered}. + +-spec command(Ref :: reference(), ReceiverPid :: pid(), {Tag :: atom(), SubCommand :: any()}, State :: #state{}) -> + {ok, NewState :: #state{}, Packets :: [binary()]} | {error, not_registered}. +command(Ref, ReceiverPid, SubCommand, State = #state{status = registered, pkt_id = PktId, pending_commands = PendingCommands, client_id = ClientId}) + when is_reference(Ref), is_pid(ReceiverPid) -> + CommandPkt = sdlan_pb:encode_msg(#'SDLCommand'{ + pkt_id = PktId, + command = SubCommand + }), + logger:debug("[sdlan_session] client_id: ~p, will send Command: ~p", [ClientId, SubCommand]), + + {ok, State#state{pkt_id = PktId + 1, pending_commands = maps:put(PktId, {Ref, ReceiverPid}, PendingCommands)}, + [<>]}; +command(_Ref, _ReceiverPid, _SubCommand, _State) -> + {error, not_registered}. + +-spec close(State :: #state{}) -> ok. +close(#state{offline_cb = OfflineCb}) -> + %% 触发客户端的离线逻辑 + is_function(OfflineCb) andalso OfflineCb(), + ok. + +-spec debug_info(State :: #state{}) -> map(). +debug_info(#state{ + status = Status, + client_id = ClientId, + network_id = NetworkId, + mac = Mac, + ip = Ip, + pending_commands = PendingCommands, + heartbeat_sec = HeartbeatSec +}) -> + #{ + session_state => Status, + client_id => ClientId, + network_id => NetworkId, + mac => Mac, + ip => Ip, + pending_commands => maps:size(PendingCommands), + heartbeat_sec => HeartbeatSec + }. + +%%%=================================================================== +%%% Internal functions +%%%=================================================================== + +handle_register_super(Body, State) -> + #'SDLRegisterSuper'{ + client_id = ClientId, network_id = NetworkId, mac = Mac, ip = Ip, mask_len = MaskLen, + hostname = HostName, pub_key = PubKey, access_token = AccessToken} = sdlan_pb:decode_msg(Body, 'SDLRegisterSuper'), + + true = (Mac =/= <<>> andalso PubKey =/= <<>> andalso ClientId =/= <<>>), + %% Mac地址不能是广播地址 + true = not (sdlan_util:is_multicast_mac(Mac) orelse sdlan_util:is_broadcast_mac(Mac)), + + MacBinStr = sdlan_util:format_mac(Mac), + IpAddr = sdlan_util:int_to_ipv4(Ip), + Params = #{ + <<"network_id">> => NetworkId, + <<"client_id">> => ClientId, + <<"mac">> => MacBinStr, + <<"ip">> => IpAddr, + <<"mask_len">> => MaskLen, + <<"hostname">> => HostName, + <<"access_token">> => AccessToken + }, + %% 参数检查 + logger:debug("[sdlan_session] client_id: ~p, ip: ~p, mac: ~p, host_name: ~p, access_token: ~p, network_id: ~p", + [ClientId, Ip, Mac, HostName, AccessToken, NetworkId]), + + case sdlan_api:auth_access_token(Params) of + {ok, #{<<"result">> := <<"ok">>}} -> + %% 建立到network的对应关系 + case sdlan_network:get_pid(NetworkId) of + NetworkPid when is_pid(NetworkPid) -> + {ok, Algorithm, Key, RegionId, SessionToken} = sdlan_network:attach(NetworkPid, self(), ClientId, Mac, Ip, HostName), + RsaPubKey = sdlan_cipher:rsa_pem_decode(PubKey), + RegisterSuperAck = sdlan_pb:encode_msg(#'SDLRegisterSuperAck'{ + algorithm = Algorithm, + key = rsa_encode(Key, RsaPubKey), + region_id = RegionId, + session_token = SessionToken + }), + + %% 设置节点的在线状态 + Result = sdlan_api:set_node_status(#{ + <<"network_id">> => NetworkId, + <<"client_id">> => ClientId, + <<"access_token">> => AccessToken, + <<"status">> => 1 + }), + logger:debug("[sdlan_session] client_id: ~p, set none online result is: ~p", [ClientId, Result]), + + OfflineCb = fun() -> + sdlan_api:set_node_status(#{ + <<"network_id">> => NetworkId, + <<"client_id">> => ClientId, + <<"access_token">> => AccessToken, + <<"status">> => 0 + }) + end, + NState = schedule_ping(State#state{ + status = registered, + network_id = NetworkId, + network_pid = NetworkPid, + client_id = ClientId, + mac = Mac, + ip = Ip, + offline_cb = OfflineCb + }), + {ok, NState, [<>]}; + undefined -> + logger:warning("[sdlan_session] client_id: ~p, register get error: network not found", [ClientId]), + {stop, normal, State, [register_nak_reply(?NAK_INTERNAL_FAULT, <<"Internal Error">>)]} + end; + {ok, #{<<"error">> := #{<<"code">> := Code, <<"message">> := Message}}} -> + logger:warning("[sdlan_session] network_id: ~p, client_id: ~p, register get error: ~ts, error_code: ~p", [NetworkId, ClientId, Message, Code]), + {stop, normal, State, [register_nak_reply(Code, Message)]}; + {error, Reason} -> + logger:warning("[sdlan_session] network_id: ~p, client_id: ~p, register get error: ~p", [NetworkId, ClientId, Reason]), + {stop, normal, State, [register_nak_reply(?NAK_NETWORK_FAULT, <<"Network Error">>)]} + end. + +-spec register_nak_reply(ErrorCode :: integer(), ErrorMsg :: binary()) -> binary(). +register_nak_reply(ErrorCode, ErrorMsg) when is_integer(ErrorCode), is_binary(ErrorMsg) -> + RegisterNakReply = sdlan_pb:encode_msg(#'SDLRegisterSuperNak'{ + error_code = ErrorCode, + error_message = ErrorMsg + }), + <>. + +rsa_encode(PlainText, RsaPubKey) when is_binary(PlainText) -> + iolist_to_binary(sdlan_cipher:rsa_encrypt(PlainText, RsaPubKey)). + +-spec get_rules(SrcIdentityId :: integer(), DstIdentityId :: integer()) -> {ok, [{Proto :: integer(), Port :: integer()}]}. +get_rules(SrcIdentityId, DstIdentityId) when is_integer(SrcIdentityId), is_integer(DstIdentityId) -> + SrcPolicyIds = identity_policy_ets:get_policies(SrcIdentityId), + DstPolicyIds = identity_policy_ets:get_policies(DstIdentityId), + rule_ets:get_rules(SrcPolicyIds, DstPolicyIds). + +schedule_ping(State = #state{heartbeat_sec = HeartbeatSec}) -> + State#state{ping_timer = erlang:start_timer(heartbeat_ms(HeartbeatSec), self(), ping_ticker)}. + +heartbeat_ms(HeartbeatSec) when is_integer(HeartbeatSec), HeartbeatSec > 0 -> + HeartbeatSec * 1000; +heartbeat_ms(_) -> + ?PING_TICKER.