fix udp_handler

This commit is contained in:
anlicheng 2026-06-28 21:58:18 +08:00
parent 77df81f8c3
commit df3696cd50
2 changed files with 43 additions and 37 deletions

View File

@ -67,9 +67,6 @@ handle_call(_Request, _From, State = #state{idle_timeout = IdleTimeout}) ->
handle_cast(_Request, State = #state{idle_timeout = IdleTimeout}) -> handle_cast(_Request, State = #state{idle_timeout = IdleTimeout}) ->
{noreply, State, IdleTimeout}. {noreply, State, IdleTimeout}.
handle_info({datagram, Server, <<"stop">>}, State = #state{server = Server, peer = Peer}) ->
logger:debug("UDP peer stopped: ~s", [esockd:format(Peer)]),
{stop, normal, State};
handle_info({datagram, Server, Packet}, State = #state{server = Server}) -> handle_info({datagram, Server, Packet}, State = #state{server = Server}) ->
handle_encrypted_datagram(Packet, State); handle_encrypted_datagram(Packet, State);
@ -104,6 +101,9 @@ handle_encrypted_datagram(Packet, State = #state{idle_timeout = IdleTimeout, cip
handle_datagram(Packet, State = #state{idle_timeout = IdleTimeout}) -> handle_datagram(Packet, State = #state{idle_timeout = IdleTimeout}) ->
case relay_server_udp_protocol:decode(Packet) of case relay_server_udp_protocol:decode(Packet) of
{ok, connection_close} ->
{stop, normal, State};
{ok, {open, StreamId, Payload}} -> {ok, {open, StreamId, Payload}} ->
logger:debug("UDP handler open stream: ~p", [StreamId]), logger:debug("UDP handler open stream: ~p", [StreamId]),
{noreply, handle_open(StreamId, Payload, State), IdleTimeout}; {noreply, handle_open(StreamId, Payload, State), IdleTimeout};
@ -131,8 +131,7 @@ handle_open(StreamId, Payload, State = #state{streams = Streams, users = Users})
ok -> ok ->
open_stream(StreamId, Request, State); open_stream(StreamId, Request, State);
error -> error ->
logger:warning("UDP relay stream ~p authentication failed from ~s", logger:warning("UDP relay stream ~p authentication failed from ~s", [StreamId, esockd:format(State#state.peer)]),
[StreamId, esockd:format(State#state.peer)]),
send_error(StreamId, <<"authentication failed">>, State), send_error(StreamId, <<"authentication failed">>, State),
State State
end; end;
@ -145,13 +144,9 @@ handle_open(StreamId, Payload, State = #state{streams = Streams, users = Users})
open_stream(StreamId, #{host := Host, port := Port}, State = #state{streams = Streams, sockets = Sockets}) -> open_stream(StreamId, #{host := Host, port := Port}, State = #state{streams = Streams, sockets = Sockets}) ->
case connect_remote(Host, Port, State#state.connect_timeout) of case connect_remote(Host, Port, State#state.connect_timeout) of
{ok, Socket} -> {ok, Socket} ->
Stream = #stream{ Stream = #stream{id = StreamId, socket = Socket, host = Host, port = Port},
id = StreamId,
socket = Socket,
host = Host,
port = Port
},
logger:debug("UDP relay stream ~p opened to ~ts:~p", [StreamId, Host, 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)}; State#state{streams = maps:put(StreamId, Stream, Streams), sockets = maps:put(Socket, StreamId, Sockets)};
{error, Reason} -> {error, Reason} ->
Message = iolist_to_binary(io_lib:format("connect failed: ~p", [Reason])), Message = iolist_to_binary(io_lib:format("connect failed: ~p", [Reason])),
@ -174,48 +169,52 @@ handle_data(StreamId, Payload, State = #state{streams = Streams}) ->
end end
end. end.
handle_remote_data(Socket, Data, State = #state{sockets = Sockets}) -> handle_remote_data(Socket, Data, State = #state{sockets = Sockets, streams = Streams, idle_timeout = IdleTimeout}) ->
case maps:get(Socket, Sockets, undefined) of case maps:get(Socket, Sockets, undefined) of
undefined -> undefined ->
{noreply, State, State#state.idle_timeout}; {noreply, State, State#state.idle_timeout};
StreamId -> StreamId ->
send_frame(data, StreamId, Data, State), send_frame(data, StreamId, Data, State),
NState = case maps:get(StreamId, State#state.streams, undefined) of NState = case maps:get(StreamId, Streams, undefined) of
undefined -> undefined ->
State; State;
Stream -> Stream ->
case set_active_once(Stream) of case set_active_once(Stream) of
ok -> State; ok ->
{error, _Reason} -> close_stream(StreamId, State) State;
{error, _Reason} ->
close_stream(StreamId, State)
end end
end, end,
{noreply, NState, NState#state.idle_timeout} {noreply, NState, IdleTimeout}
end. end.
handle_remote_closed(Socket, State = #state{sockets = Sockets}) -> handle_remote_closed(Socket, State = #state{sockets = Sockets, idle_timeout = IdleTimeout}) ->
case maps:get(Socket, Sockets, undefined) of case maps:get(Socket, Sockets, undefined) of
undefined -> undefined ->
{noreply, State, State#state.idle_timeout}; {noreply, State, IdleTimeout};
StreamId -> StreamId ->
send_frame(close, StreamId, <<>>, State), send_frame(close, StreamId, <<>>, State),
NState = remove_stream(StreamId, State), NState = remove_stream(StreamId, State),
{noreply, NState, NState#state.idle_timeout} {noreply, NState, IdleTimeout}
end. end.
handle_remote_error(Socket, Reason, State = #state{sockets = Sockets}) -> handle_remote_error(Socket, Reason, State = #state{sockets = Sockets, idle_timeout = IdleTimeout}) ->
case maps:get(Socket, Sockets, undefined) of case maps:get(Socket, Sockets, undefined) of
undefined -> undefined ->
{noreply, State, State#state.idle_timeout}; {noreply, State, IdleTimeout};
StreamId -> StreamId ->
send_error(StreamId, format_error(Reason), State), send_error(StreamId, format_error(Reason), State),
NState = close_stream(StreamId, State), NState = close_stream(StreamId, State),
{noreply, NState, NState#state.idle_timeout} {noreply, NState, IdleTimeout}
end. end.
connect_remote(Host, Port, Timeout) -> connect_remote(Host, Port, Timeout) ->
case gen_tcp:connect(binary_to_list(Host), Port, ?TCP_OPTIONS, Timeout) of case gen_tcp:connect(binary_to_list(Host), Port, ?TCP_OPTIONS, Timeout) of
{ok, Socket} -> {ok, Socket}; {ok, Socket} ->
{error, Reason} -> {error, Reason} {ok, Socket};
{error, Reason} ->
{error, Reason}
end. end.
send_remote(#stream{socket = Socket}, Payload) -> send_remote(#stream{socket = Socket}, Payload) ->

View File

@ -12,7 +12,12 @@
-define(VERSION, 1). -define(VERSION, 1).
-define(HEADER_LENGTH, 20). -define(HEADER_LENGTH, 20).
decode(<<"RKP1", ?VERSION:8, TypeNo:8, _Reserved:16, StreamId:64, PayloadLen:32, Rest/binary>>) -> -define(CONNECTION_CLOSE, 16#FF).
%%
decode(<<"RKP1", ?VERSION:8, ?CONNECTION_CLOSE:8, 0:16, 0:64, 0:32, _Rest/binary>>) ->
{ok, connection_close};
decode(<<"RKP1", ?VERSION:8, TypeNo:8, 0:16, StreamId:64, PayloadLen:32, Rest/binary>>) ->
case {type(TypeNo), byte_size(Rest) >= PayloadLen} of case {type(TypeNo), byte_size(Rest) >= PayloadLen} of
{undefined, _} -> {undefined, _} ->
{error, unknown_frame_type}; {error, unknown_frame_type};
@ -67,12 +72,14 @@ encode_open_request(Host, Port, Username, Password) when
PasswordLen:16, Password/binary>>. PasswordLen:16, Password/binary>>.
type(1) -> open; type(1) -> open;
type(2) -> data; type(2) -> open_ack;
type(3) -> close; type(3) -> data;
type(4) -> error; type(4) -> close;
type(5) -> error;
type(_) -> undefined. type(_) -> undefined.
type_no(open) -> 1; type_no(open) -> 1;
type_no(data) -> 2; type_no(open_ack) -> 2;
type_no(close) -> 3; type_no(data) -> 3;
type_no(error) -> 4. type_no(close) -> 4;
type_no(error) -> 5.