From 4791b66086531dca0f7980c681b09b2cb34bdfde Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Wed, 1 Jul 2026 17:42:04 +0800 Subject: [PATCH] =?UTF-8?q?=E5=8D=8F=E8=AE=AE=E8=B0=83=E6=95=B4=EF=BC=8C?= =?UTF-8?q?=E9=81=BF=E5=85=8D=E9=98=BB=E5=A1=9E?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../src/relay_server_udp_listener.erl | 2 +- ..._handler.erl => relay_server_udp_peer.erl} | 224 ++++++++---------- .../src/relay_server_udp_protocol.erl | 7 +- .../src/relay_server_udp_stream.erl | 124 ++++++++++ 4 files changed, 234 insertions(+), 123 deletions(-) rename apps/relay_server/src/{relay_server_udp_handler.erl => relay_server_udp_peer.erl} (55%) create mode 100644 apps/relay_server/src/relay_server_udp_stream.erl diff --git a/apps/relay_server/src/relay_server_udp_listener.erl b/apps/relay_server/src/relay_server_udp_listener.erl index 9488b1f..c92b74b 100644 --- a/apps/relay_server/src/relay_server_udp_listener.erl +++ b/apps/relay_server/src/relay_server_udp_listener.erl @@ -55,7 +55,7 @@ init([]) -> ], Opts = maybe_add(max_conn_rate, Props, BaseOpts), - MFA = {relay_server_udp_handler, start_link, [IdleTimeout, ConnectTimeout, Users]}, + MFA = {relay_server_udp_peer, start_link, [IdleTimeout, ConnectTimeout, Users]}, case esockd:open_udp(Listener, ListenOn, Opts, MFA) of {ok, Pid} -> logger:info("UDP listener ~p started on ~s", diff --git a/apps/relay_server/src/relay_server_udp_handler.erl b/apps/relay_server/src/relay_server_udp_peer.erl similarity index 55% rename from apps/relay_server/src/relay_server_udp_handler.erl rename to apps/relay_server/src/relay_server_udp_peer.erl index 89835af..975f708 100644 --- a/apps/relay_server/src/relay_server_udp_handler.erl +++ b/apps/relay_server/src/relay_server_udp_peer.erl @@ -1,9 +1,9 @@ %%%------------------------------------------------------------------- -%% @doc Per-peer UDP packet handler. +%% @doc Per-peer UDP relay router. %% @end %%%------------------------------------------------------------------- --module(relay_server_udp_handler). +-module(relay_server_udp_peer). -behaviour(gen_server). @@ -11,25 +11,14 @@ -export([init/1, handle_call/3, handle_cast/2, handle_info/2, terminate/2, code_change/3]). --define(TCP_OPTIONS, [binary, {packet, raw}, {active, once}, {nodelay, true}]). - --record(stream, { - id :: non_neg_integer(), - socket :: inet:socket(), - host :: binary(), - port :: inet:port_number() -}). - --type stream() :: #stream{}. - -record(state, { server :: pid(), peer :: {inet:ip_address(), inet:port_number()}, idle_timeout :: timeout(), connect_timeout :: timeout(), users = #{} :: #{binary() => binary()}, - streams = #{} :: #{non_neg_integer() => stream()}, - sockets = #{} :: #{inet:socket() => non_neg_integer()} + streams = #{} :: #{non_neg_integer() => pid()}, + stream_ids = #{} :: #{pid() => non_neg_integer()} }). start_link(Transport, Peer, IdleTimeout, ConnectTimeout) -> @@ -39,6 +28,7 @@ start_link(Transport, Peer, IdleTimeout, ConnectTimeout, Users) -> gen_server:start_link(?MODULE, [Transport, Peer, IdleTimeout, ConnectTimeout, Users], []). init([{udp, Server, _Sock}, Peer, IdleTimeout, ConnectTimeout, Users]) -> + process_flag(trap_exit, true), logger:debug("UDP peer connected: ~s", [esockd:format(Peer)]), {ok, #state{ server = Server, @@ -55,14 +45,57 @@ handle_cast(_Request, State = #state{idle_timeout = IdleTimeout}) -> {noreply, State, IdleTimeout}. handle_info({datagram, Server, Packet}, State = #state{server = Server}) -> - handle_encrypted_datagram(Packet, State); + handle_datagram(Packet, State); -handle_info({tcp, Socket, Data}, State) -> - handle_remote_data(Socket, Data, State); -handle_info({tcp_closed, Socket}, State) -> - handle_remote_closed(Socket, State); -handle_info({tcp_error, Socket, Reason}, State) -> - handle_remote_error(Socket, Reason, State); +handle_info({stream_opened, StreamId, StreamPid}, State = #state{idle_timeout = IdleTimeout}) -> + NState = case stream_pid(StreamId, State) of + StreamPid -> + send_frame(open_ack, StreamId, <<>>, State), + State; + _Other -> + relay_server_udp_stream:close(StreamPid), + State + end, + {noreply, NState, IdleTimeout}; + +handle_info({stream_data, StreamId, StreamPid, Data}, State = #state{idle_timeout = IdleTimeout}) -> + case stream_pid(StreamId, State) of + StreamPid -> + send_frame(data, StreamId, Data, State), + {noreply, State, IdleTimeout}; + _Other -> + {noreply, State, IdleTimeout} + end; + +handle_info({stream_closed, StreamId, StreamPid}, State = #state{idle_timeout = IdleTimeout}) -> + NState = case stream_pid(StreamId, State) of + StreamPid -> + send_frame(close, StreamId, <<>>, State), + remove_stream(StreamId, StreamPid, State); + _Other -> + State + end, + {noreply, NState, IdleTimeout}; + +handle_info({stream_error, StreamId, StreamPid, Reason}, State = #state{idle_timeout = IdleTimeout}) -> + NState = case stream_pid(StreamId, State) of + StreamPid -> + send_error(StreamId, format_stream_error(Reason), State), + remove_stream(StreamId, StreamPid, State); + _Other -> + State + end, + {noreply, NState, IdleTimeout}; + +handle_info({'EXIT', StreamPid, Reason}, State = #state{idle_timeout = IdleTimeout}) -> + NState = case stream_id(StreamPid, State) of + undefined -> + State; + StreamId -> + logger:debug("UDP relay stream ~p worker exited: ~p", [StreamId, Reason]), + remove_stream(StreamId, StreamPid, State) + end, + {noreply, NState, IdleTimeout}; handle_info(timeout, State = #state{peer = Peer}) -> logger:debug("UDP peer idle timeout: ~s", [esockd:format(Peer)]), @@ -71,19 +104,18 @@ handle_info(_Info, State = #state{idle_timeout = IdleTimeout}) -> {noreply, State, IdleTimeout}. terminate(_Reason, #state{streams = Streams}) -> - maps:foreach(fun(_StreamId, Stream) -> close_remote(Stream) end, Streams), + maps:foreach(fun(_StreamId, StreamPid) -> relay_server_udp_stream:close(StreamPid) end, Streams), ok. code_change(_OldVsn, State, _Extra) -> {ok, State}. -handle_encrypted_datagram(Packet, State = #state{idle_timeout = IdleTimeout}) -> +handle_datagram(Packet, State = #state{idle_timeout = IdleTimeout}) -> case relay_server_udp_protocol:decode(Packet) of {ok, connection_close} -> {stop, normal, State}; - {ok, {open, StreamId, Payload}} -> - logger:debug("UDP handler open stream: ~p", [StreamId]), + logger:debug("UDP peer open stream: ~p", [StreamId]), {noreply, handle_open(StreamId, Payload, State), IdleTimeout}; {ok, {data, StreamId, Payload}} -> {noreply, handle_data(StreamId, Payload, State), IdleTimeout}; @@ -107,9 +139,10 @@ handle_open(StreamId, Payload, State = #state{streams = Streams, users = Users}) {ok, Request} -> case authenticate(Request, Users) of ok -> - open_stream(StreamId, Request, State); + start_stream(StreamId, Request, State); error -> - logger:warning("UDP relay stream ~p authentication failed from ~s", [StreamId, esockd:format(State#state.peer)]), + logger:warning("UDP relay stream ~p authentication failed from ~s", + [StreamId, esockd:format(State#state.peer)]), send_error(StreamId, <<"authentication failed">>, State), State end; @@ -119,110 +152,51 @@ handle_open(StreamId, Payload, State = #state{streams = Streams, users = Users}) end end. -open_stream(StreamId, #{host := Host, port := Port}, State = #state{streams = Streams, sockets = Sockets}) -> - case connect_remote(Host, Port, State#state.connect_timeout) of - {ok, Socket} -> - Stream = #stream{id = StreamId, socket = Socket, host = Host, port = Port}, - logger:debug("UDP relay stream ~p opened to ~ts:~p", [StreamId, Host, Port]), - send_frame(open_ack, StreamId, <<>>, State), - State#state{streams = maps:put(StreamId, Stream, Streams), sockets = maps:put(Socket, StreamId, Sockets)}; +start_stream(StreamId, Request, State = #state{connect_timeout = ConnectTimeout}) -> + case relay_server_udp_stream:start_link(self(), StreamId, Request, ConnectTimeout) of + {ok, StreamPid} -> + put_stream(StreamId, StreamPid, State); {error, Reason} -> - Message = iolist_to_binary(io_lib:format("connect failed: ~p", [Reason])), - send_error(StreamId, Message, State), + send_error(StreamId, format_stream_error({start_failed, Reason}), State), State end. -handle_data(StreamId, Payload, State = #state{streams = Streams}) -> - case maps:get(StreamId, Streams, undefined) of +handle_data(StreamId, Payload, State) -> + case stream_pid(StreamId, State) of undefined -> send_error(StreamId, <<"unknown stream">>, State), State; - Stream -> - case send_remote(Stream, Payload) of - ok -> - State; - {error, Reason} -> - send_error(StreamId, format_error(Reason), State), - close_stream(StreamId, State) - end + StreamPid -> + relay_server_udp_stream:send(StreamPid, Payload), + State end. -handle_remote_data(Socket, Data, State = #state{sockets = Sockets, streams = Streams, idle_timeout = IdleTimeout}) -> - case maps:get(Socket, Sockets, undefined) of - undefined -> - {noreply, State, State#state.idle_timeout}; - StreamId -> - send_frame(data, StreamId, Data, State), - NState = case maps:get(StreamId, Streams, undefined) of - undefined -> - State; - Stream -> - case set_active_once(Stream) of - ok -> - State; - {error, _Reason} -> - close_stream(StreamId, State) - end - end, - {noreply, NState, IdleTimeout} - end. - -handle_remote_closed(Socket, State = #state{sockets = Sockets, idle_timeout = IdleTimeout}) -> - case maps:get(Socket, Sockets, undefined) of - undefined -> - {noreply, State, IdleTimeout}; - StreamId -> - send_frame(close, StreamId, <<>>, State), - NState = remove_stream(StreamId, State), - {noreply, NState, IdleTimeout} - end. - -handle_remote_error(Socket, Reason, State = #state{sockets = Sockets, idle_timeout = IdleTimeout}) -> - case maps:get(Socket, Sockets, undefined) of - undefined -> - {noreply, State, IdleTimeout}; - StreamId -> - send_error(StreamId, format_error(Reason), State), - NState = close_stream(StreamId, State), - {noreply, NState, IdleTimeout} - end. - -connect_remote(Host, Port, Timeout) -> - case gen_tcp:connect(binary_to_list(Host), Port, ?TCP_OPTIONS, Timeout) of - {ok, Socket} -> - {ok, Socket}; - {error, Reason} -> - {error, Reason} - end. - -send_remote(#stream{socket = Socket}, Payload) -> - gen_tcp:send(Socket, Payload). - -set_active_once(#stream{socket = Socket}) -> - inet:setopts(Socket, [{active, once}]). - -close_stream(StreamId, State = #state{streams = Streams}) -> - case maps:get(StreamId, Streams, undefined) of +close_stream(StreamId, State) -> + case stream_pid(StreamId, State) of undefined -> State; - Stream -> - close_remote(Stream), - remove_stream(StreamId, State) + StreamPid -> + relay_server_udp_stream:close(StreamPid), + remove_stream(StreamId, StreamPid, State) end. -remove_stream(StreamId, State = #state{streams = Streams, sockets = Sockets}) -> - case maps:get(StreamId, Streams, undefined) of - undefined -> - State; - #stream{socket = Socket} -> - State#state{ - streams = maps:remove(StreamId, Streams), - sockets = maps:remove(Socket, Sockets) - } - end. +stream_pid(StreamId, #state{streams = Streams}) -> + maps:get(StreamId, Streams, undefined). -close_remote(#stream{socket = Socket}) -> - gen_tcp:close(Socket). +stream_id(StreamPid, #state{stream_ids = StreamIds}) -> + maps:get(StreamPid, StreamIds, undefined). + +put_stream(StreamId, StreamPid, State = #state{streams = Streams, stream_ids = StreamIds}) -> + State#state{ + streams = maps:put(StreamId, StreamPid, Streams), + stream_ids = maps:put(StreamPid, StreamId, StreamIds) + }. + +remove_stream(StreamId, StreamPid, State = #state{streams = Streams, stream_ids = StreamIds}) -> + State#state{ + streams = maps:remove(StreamId, Streams), + stream_ids = maps:remove(StreamPid, StreamIds) + }. send_frame(Type, StreamId, <>, State) when byte_size(Rest) > 0 -> @@ -238,17 +212,27 @@ send_frame(Type, StreamId, Payload, State) -> send_frame0(Type, StreamId, Part, #state{server = Server, peer = Peer}) -> case relay_server_udp_protocol:encode(Type, StreamId, Part) of {ok, Packet} -> - logger:error("UDP relay encrypt packet size: ~p", [byte_size(Packet)]), + logger:debug("UDP relay encrypted packet size: ~p", [byte_size(Packet)]), Server ! {datagram, Peer, Packet}, ok; {error, Reason} -> - logger:error("UDP relay failed to encrypt ~p frame for stream ~p: ~p", [Type, StreamId, Reason]), + logger:error("UDP relay failed to encode ~p frame for stream ~p: ~p", + [Type, StreamId, Reason]), {error, Reason} end. send_error(StreamId, Message, State) -> send_frame(error, StreamId, Message, State). +format_stream_error({connect_failed, Reason}) -> + iolist_to_binary(io_lib:format("connect failed: ~p", [Reason])); +format_stream_error({send_failed, Reason}) -> + iolist_to_binary(io_lib:format("send failed: ~p", [Reason])); +format_stream_error({start_failed, Reason}) -> + iolist_to_binary(io_lib:format("stream start failed: ~p", [Reason])); +format_stream_error(Reason) -> + format_error(Reason). + format_error(Reason) when is_binary(Reason) -> Reason; format_error(Reason) -> diff --git a/apps/relay_server/src/relay_server_udp_protocol.erl b/apps/relay_server/src/relay_server_udp_protocol.erl index 2a13630..557fddd 100644 --- a/apps/relay_server/src/relay_server_udp_protocol.erl +++ b/apps/relay_server/src/relay_server_udp_protocol.erl @@ -31,7 +31,10 @@ decode(<>) -> decode0(PlainText, Body); Error -> Error - end. + end; +decode(_Packet) -> + {error, invalid_packet}. + %% 全局控制指令也要加密处理 decode0(<<"RKP1", ?VERSION:8, ?CONNECTION_CLOSE:8, 0:16, 0:64, 0:32>>, _Body) -> {ok, connection_close}; @@ -56,7 +59,7 @@ encode(Type, StreamId, Payload) when is_atom(Type), is_integer(StreamId), is_bin Head = <<"RKP1", ?VERSION:8, TypeNo:8, 0:16, StreamId:64, PayloadLen:32>>, case chacha20_cipher:encrypt(Head, ?CHACHA20_KEY) of {ok, EncryptHead} -> - <>; + {ok, <>}; Error -> Error end. diff --git a/apps/relay_server/src/relay_server_udp_stream.erl b/apps/relay_server/src/relay_server_udp_stream.erl new file mode 100644 index 0000000..f3a01a1 --- /dev/null +++ b/apps/relay_server/src/relay_server_udp_stream.erl @@ -0,0 +1,124 @@ +%%%------------------------------------------------------------------- +%% @doc Per-stream TCP relay worker. +%% @end +%%%------------------------------------------------------------------- + +-module(relay_server_udp_stream). + +-behaviour(gen_server). + +-export([start_link/4, send/2, close/1]). + +-export([init/1, handle_continue/2, handle_call/3, handle_cast/2, handle_info/2, terminate/2, code_change/3]). + +-define(TCP_OPTIONS(Timeout), [ + binary, + {packet, raw}, + {active, false}, + {nodelay, true}, + {send_timeout, Timeout}, + {send_timeout_close, true} +]). + +-record(state, { + peer :: pid(), + stream_id :: non_neg_integer(), + host :: binary(), + port :: inet:port_number(), + connect_timeout :: timeout(), + socket = undefined :: undefined | inet:socket() +}). + +start_link(Peer, StreamId, #{host := Host, port := Port}, ConnectTimeout) -> + gen_server:start_link(?MODULE, [Peer, StreamId, Host, Port, ConnectTimeout], []). + +send(StreamPid, Payload) when is_pid(StreamPid), is_binary(Payload) -> + gen_server:cast(StreamPid, {send, Payload}). + +close(StreamPid) when is_pid(StreamPid) -> + gen_server:cast(StreamPid, close). + +init([Peer, StreamId, Host, Port, ConnectTimeout]) -> + {ok, #state{ + peer = Peer, + stream_id = StreamId, + host = Host, + port = Port, + connect_timeout = ConnectTimeout + }, {continue, connect}}. + +handle_call(_Request, _From, State) -> + {reply, {error, bad_request}, State}. + +handle_cast({send, _Payload}, State = #state{socket = undefined}) -> + {noreply, State}; +handle_cast({send, Payload}, State = #state{socket = Socket}) -> + case gen_tcp:send(Socket, Payload) of + ok -> + {noreply, State}; + {error, Reason} -> + notify_error({send_failed, Reason}, State), + {stop, normal, State} + end; +handle_cast(close, State) -> + {stop, normal, State}; +handle_cast(_Request, State) -> + {noreply, State}. + +handle_info({tcp, Socket, Data}, State = #state{socket = Socket}) -> + notify_data(Data, State), + case inet:setopts(Socket, [{active, once}]) of + ok -> + {noreply, State}; + {error, Reason} -> + notify_error(Reason, State), + {stop, normal, State} + end; +handle_info({tcp_closed, Socket}, State = #state{socket = Socket}) -> + notify_closed(State), + {stop, normal, State}; +handle_info({tcp_error, Socket, Reason}, State = #state{socket = Socket}) -> + notify_error(Reason, State), + {stop, normal, State}; +handle_info(_Info, State) -> + {noreply, State}. + +terminate(_Reason, #state{socket = undefined}) -> + ok; +terminate(_Reason, #state{socket = Socket}) -> + gen_tcp:close(Socket), + ok. + +code_change(_OldVsn, State, _Extra) -> + {ok, State}. + +handle_continue(connect, State = #state{host = Host, port = Port, connect_timeout = ConnectTimeout}) -> + case gen_tcp:connect(binary_to_list(Host), Port, ?TCP_OPTIONS(ConnectTimeout), ConnectTimeout) of + {ok, Socket} -> + logger:debug("UDP relay stream ~p opened to ~ts:~p", + [State#state.stream_id, Host, Port]), + case inet:setopts(Socket, [{active, once}]) of + ok -> + notify_opened(State), + {noreply, State#state{socket = Socket}}; + {error, Reason} -> + gen_tcp:close(Socket), + notify_error(Reason, State), + {stop, normal, State} + end; + {error, Reason} -> + notify_error({connect_failed, Reason}, State), + {stop, normal, State} + end. + +notify_opened(#state{peer = Peer, stream_id = StreamId}) -> + Peer ! {stream_opened, StreamId, self()}. + +notify_data(Data, #state{peer = Peer, stream_id = StreamId}) -> + Peer ! {stream_data, StreamId, self(), Data}. + +notify_closed(#state{peer = Peer, stream_id = StreamId}) -> + Peer ! {stream_closed, StreamId, self()}. + +notify_error(Reason, #state{peer = Peer, stream_id = StreamId}) -> + Peer ! {stream_error, StreamId, self(), Reason}.