协议调整,避免阻塞

This commit is contained in:
anlicheng 2026-07-01 17:42:04 +08:00
parent 22db3e1f46
commit 4791b66086
4 changed files with 234 additions and 123 deletions

View File

@ -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",

View File

@ -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, <<Part:1200/binary, Rest/binary>>, 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) ->

View File

@ -31,7 +31,10 @@ decode(<<Nonce:12/binary, CipherHead:20/binary, Tag:16/binary, Body/binary>>) ->
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} ->
<<EncryptHead/binary, Payload/binary>>;
{ok, <<EncryptHead/binary, Payload/binary>>};
Error ->
Error
end.

View File

@ -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}.