Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 69 additions & 20 deletions src/hackney_conn.erl
Original file line number Diff line number Diff line change
Expand Up @@ -234,6 +234,7 @@
%% {stream, headers, Status, Headers, Buffer, Pending}
%% {stream, body_full, Status, Headers, Acc, From}
%% {stream, done, Status, Headers, Buffer}
%% {stream, refused, ErrorCode} (refused by a GOAWAY, caller not yet told)
h2_streams = #{} :: #{pos_integer() => {term(), tuple()}},
%% Shared through the pool (share_h2/1): no owner, each stream monitors
%% its caller, and the connection closes itself once idle.
Expand Down Expand Up @@ -1033,6 +1034,17 @@ connected({call, From}, get_state, #conn_data{h2_goaway = ErrorCode})
connected({call, From}, get_state, _Data) ->
{keep_state_and_data, [{reply, From, {ok, connected}}]};

%% The streaming-body calls only make sense in streaming_body, except for a
%% request a GOAWAY refused while its caller was between them: the connection
%% came back here then (h2_on_goaway/3), and the caller's next call is the
%% first chance to tell it.
connected({call, From}, {send_body_chunk, _}, #conn_data{protocol = http2} = Data) ->
h2_answer_refused_request(From, Data);
connected({call, From}, finish_send_body, #conn_data{protocol = http2} = Data) ->
h2_answer_refused_request(From, Data);
connected({call, From}, start_response, #conn_data{protocol = http2} = Data) ->
h2_answer_refused_request(From, Data);

connected({call, From}, verify_socket, #conn_data{socket = undefined} = Data) ->
%% Socket not connected
{next_state, closed, Data, [{reply, From, {error, closed}}]};
Expand Down Expand Up @@ -3309,9 +3321,16 @@ handle_h2_stream_owner_down(Ref, State,
h2_stream_owner_down_result(streaming_body, StreamId,
#conn_data{h2_stream_id = StreamId} = Data) ->
%% The caller streaming a request body died: its stream is gone, so the
%% connection can take requests again.
{next_state, connected,
Data#conn_data{h2_stream_id = undefined, request_from = undefined}};
%% connection can take requests again, unless that stream was the last
%% one of a GOAWAY drain.
case h2_stream_result(Data#conn_data{h2_stream_id = undefined,
request_from = undefined}, []) of
{keep_state, Data1, Actions} ->
{next_state, connected, Data1,
[{state_timeout, infinity, idle_timeout} | Actions]};
Closed ->
Closed
end;
h2_stream_owner_down_result(_State, _StreamId, Data) ->
h2_stream_result(Data, []).

Expand Down Expand Up @@ -3991,24 +4010,42 @@ h2_idle_actions(_Data) ->
%% here, and close once the accepted streams end. Aborting those too reported
%% requests the server went on to complete as failed.
%%
%% A streamed request or response body (a `stream' entry) drives the
%% connection through states of its own, so with one in flight the connection
%% still closes at once.
%% A streaming-body request the GOAWAY refuses has no caller parked on it, so
%% it is kept as a `refused' entry and the connection leaves streaming_body:
%% the caller learns of it on its next call (h2_answer_refused_request/2).
h2_on_goaway(LastStreamId, ErrorCode, #conn_data{h2_streams = Streams} = Data) ->
{Accepted, Refused} = lists:partition(fun(SId) -> SId =< LastStreamId end,
maps:keys(Streams)),
case Accepted =/= [] andalso not h2_streaming_in_flight(Streams) of
Refused = [SId || SId <- maps:keys(Streams), SId > LastStreamId],
{Replies, Data1} = abort_h2_streams(Refused, ErrorCode, Data),
#conn_data{h2_streams = Streams1, h2_stream_id = StreamId} = Data1,
case map_size(Streams1) =:= 0 of
true ->
{Replies, Data1} = abort_h2_streams(Refused, {goaway, ErrorCode}, Data),
ok = leave_h2_pool(Data1),
{keep_state, Data1#conn_data{h2_goaway = ErrorCode}, Replies};
{next_state, closed, Data2, CloseReplies} = h2_close_on_goaway(ErrorCode, Data1),
{next_state, closed, Data2, Replies ++ CloseReplies};
false ->
h2_close_on_goaway(ErrorCode, Data)
ok = leave_h2_pool(Data1),
Data3 = Data1#conn_data{h2_goaway = ErrorCode},
case maps:get(StreamId, Streams1, undefined) of
{_, {stream, refused, _}} ->
{next_state, connected, Data3,
[{state_timeout, infinity, idle_timeout} | Replies]};
_ ->
{keep_state, Data3, Replies}
end
end.

h2_streaming_in_flight(Streams) ->
lists:any(fun({_Owner, Inner}) -> element(1, Inner) =:= stream end,
maps:values(Streams)).
%% @private Answer a streaming-body call made after the request's stream was
%% refused by a GOAWAY (see h2_on_goaway/3). Any other such call in
%% `connected' is a misuse, as before.
h2_answer_refused_request({Caller, _} = From, #conn_data{h2_streams = Streams} = Data) ->
case [{Id, Code} || {Id, {Owner, {stream, refused, Code}}} <- maps:to_list(Streams),
h2_stream_owner_pid(Owner) =:= Caller] of
[{StreamId, ErrorCode}] ->
Data1 = Data#conn_data{h2_stream_id = undefined, request_from = undefined},
h2_stream_result(drop_h2_stream(StreamId, Data1),
[{reply, From, {error, {goaway, ErrorCode}}}]);
[] ->
{keep_state_and_data, [{reply, From, {error, invalid_state}}]}
end.

leave_h2_pool(#conn_data{pool_pid = PoolPid}) when is_pid(PoolPid) ->
gen_server:cast(PoolPid, {unregister_h2, self()});
Expand Down Expand Up @@ -4051,10 +4088,22 @@ collect_h2_aborts(Err, #conn_data{h2_streams = Streams} = Data) ->
Data#conn_data{h2_streams = #{}, request_from = undefined}),
{Replies, Data1}.

%% @private Fail some streams and keep the rest of the connection running.
abort_h2_streams(StreamIds, Err, #conn_data{h2_streams = Streams} = Data) ->
Replies = h2_abort_replies(Err, maps:with(StreamIds, Streams)),
{Replies, lists:foldl(fun drop_h2_stream/2, Data, StreamIds)}.
%% @private Fail the refused streams and keep the rest of the connection
%% running. A request body still being sent has no caller parked on it, so
%% its stream is kept as a `refused' entry for the caller's next call to find.
abort_h2_streams(StreamIds, ErrorCode, #conn_data{h2_streams = Streams} = Data) ->
Refused = maps:with(StreamIds, Streams),
Replies = h2_abort_replies({goaway, ErrorCode}, Refused),
Data1 = maps:fold(fun
(StreamId, {Owner, {stream, sending}}, D) ->
D#conn_data{h2_streams = maps:put(StreamId, {Owner, {stream, refused, ErrorCode}},
D#conn_data.h2_streams)};
(_StreamId, {_Owner, {stream, refused, _}}, D) ->
D;
(StreamId, _Entry, D) ->
drop_h2_stream(StreamId, D)
end, Data, Refused),
{Replies, Data1}.

h2_abort_replies(Err, Streams) ->
maps:fold(fun
Expand Down
161 changes: 161 additions & 0 deletions test/hackney_h2_goaway_server.erl
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
%%% Frame-level HTTP/2 server for the GOAWAY tests, built on the h2 dep's
%%% h2_frame / h2_hpack.
%%%
%%% Its first connection holds every request until it has sent its one GOAWAY,
%%% then answers only the streams at or below last_stream_id, a request with a
%%% streamed body once that body is complete. Later connections answer every
%%% request as soon as it is complete. Each response body is the stream id.
%%%
%%% The GOAWAY is sent on a trigger:
%%% {second_stream, Pick} when the second stream opens, with
%%% Pick(FirstStreamId, SecondStreamId) as last_stream_id
%%% first_data when a stream's first body chunk arrives, with that
%%% stream as last_stream_id
%%% and shaped by options:
%%% goaway => direct | two_step two_step first sends GOAWAY(2^31-1), the
%%% graceful shutdown of RFC 9113 6.8 (default direct)
%%% answer => boolean() false answers nothing after the GOAWAY, not
%%% even the accepted streams (default true)
-module(hackney_h2_goaway_server).

-export([start/1, start/2, stop/1]).

-define(PREFACE, <<"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n">>).

start(Trigger) ->
start(Trigger, #{}).

start(Trigger, Opts) ->
Certs = cert_dir(),
{ok, LSock} = ssl:listen(0,
[{certfile, filename:join(Certs, "server.pem")},
{keyfile, filename:join(Certs, "server.key")},
{alpn_preferred_protocols, [<<"h2">>]},
{versions, ['tlsv1.2', 'tlsv1.3']},
{active, false}, {mode, binary}, {reuseaddr, true}]),
{ok, {_, Port}} = ssl:sockname(LSock),
Options = maps:merge(#{goaway => direct, answer => true}, Opts),
Pid = spawn(fun() -> accept_loop(LSock, Trigger, Options) end),
Url = iolist_to_binary([<<"https://localhost:">>, integer_to_list(Port), <<"/">>]),
{Pid, Url}.

stop(Pid) ->
exit(Pid, kill).

accept_loop(LSock, Trigger, Opts) ->
case ssl:transport_accept(LSock, 2000) of
{ok, TSock} ->
spawn(fun() -> serve(TSock, Trigger, Opts) end),
accept_loop(LSock, none, Opts);
{error, timeout} -> accept_loop(LSock, Trigger, Opts);
{error, closed} -> ok
end.

serve(TSock, Trigger, Opts) ->
case ssl:handshake(TSock, 5000) of
{ok, Sock} ->
case recv_preface(Sock, <<>>) of
{ok, Rest} ->
send(Sock, h2_frame:settings([])),
%% seen: stream ids in order of arrival. complete: streams
%% whose request has fully arrived and is not answered yet.
%% accepted: after the GOAWAY, the streams it let through.
loop(Sock, Rest, Opts#{enc => h2_hpack:new_context(), trigger => Trigger,
seen => [], complete => [], accepted => all});
_ -> ok
end;
_ -> ok
end.

recv_preface(_Sock, Acc) when byte_size(Acc) >= 24 ->
<<Pre:24/binary, Rest/binary>> = Acc,
case Pre of ?PREFACE -> {ok, Rest}; _ -> {error, bad_preface} end;
recv_preface(Sock, Acc) ->
case ssl:recv(Sock, 0, 5000) of
{ok, Data} -> recv_preface(Sock, <<Acc/binary, Data/binary>>);
{error, R} -> {error, R}
end.

loop(Sock, Buf, St) ->
case h2_frame:decode(Buf) of
{ok, Frame, Rest} ->
case handle(Sock, Frame, St) of
{continue, St2} -> loop(Sock, Rest, St2);
stop -> ok
end;
{more, _} ->
case ssl:recv(Sock, 0, 30000) of
{ok, Data} -> loop(Sock, <<Buf/binary, Data/binary>>, St);
{error, _} -> ok
end;
{error, _, Rest} -> loop(Sock, Rest, St);
{error, _} -> ok
end.

handle(Sock, {settings, _}, St) -> send(Sock, h2_frame:settings_ack()), {continue, St};
handle(Sock, {ping, D}, St) -> send(Sock, h2_frame:ping_ack(D)), {continue, St};
handle(_Sock, {goaway, _, _, _}, _St) -> stop;
handle(Sock, {headers, Sid, _Block, EndStream, _EndHeaders}, St) ->
{continue, on_request_frame(Sock, headers, Sid, EndStream, St)};
%% decode/1 adds the flow-controlled size as a fifth element.
handle(Sock, {data, Sid, _Bin, EndStream, _FlowControlled}, St) ->
{continue, on_request_frame(Sock, data, Sid, EndStream, St)};
handle(_Sock, _Other, St) -> {continue, St}.

on_request_frame(Sock, Kind, Sid, EndStream, #{seen := Seen, complete := Complete} = St) ->
Seen2 = case lists:member(Sid, Seen) of true -> Seen; false -> Seen ++ [Sid] end,
Complete2 = case EndStream of true -> [Sid | Complete]; false -> Complete end,
St1 = St#{seen := Seen2, complete := Complete2},
St2 = case goaway_for(Kind, Sid, EndStream, St1) of
none -> St1;
LastStreamId -> send_goaway(Sock, LastStreamId, St1)
end,
answer_ready(Sock, St2).

goaway_for(headers, _Sid, _EndStream, #{trigger := {second_stream, Pick}, seen := [First, Second]}) ->
Pick(First, Second);
goaway_for(data, Sid, false, #{trigger := first_data}) ->
Sid;
goaway_for(_Kind, _Sid, _EndStream, _St) ->
none.

%% One GOAWAY per connection, then a short drain before anything is answered.
send_goaway(Sock, LastStreamId, #{seen := Seen, goaway := Shape, answer := Answer} = St) ->
case Shape of
two_step ->
send(Sock, h2_frame:goaway(16#7fffffff, no_error, <<>>)),
timer:sleep(50);
direct ->
ok
end,
send(Sock, h2_frame:goaway(LastStreamId, no_error, <<>>)),
timer:sleep(200),
Accepted = case Answer of
true -> [S || S <- Seen, S =< LastStreamId];
false -> []
end,
St#{trigger := none, accepted := Accepted}.

%% Answer every complete request the connection may still answer. Nothing is
%% answered while a GOAWAY is still to come, so the streams it covers are in
%% flight when it does.
answer_ready(_Sock, #{trigger := Trigger} = St) when Trigger =/= none ->
St;
answer_ready(Sock, #{complete := Complete, accepted := Accepted} = St) ->
Ready = [S || S <- lists:reverse(Complete),
Accepted =:= all orelse lists:member(S, Accepted)],
lists:foldl(fun(S, Acc) -> respond(Sock, S, Acc) end,
St#{complete := Complete -- Ready}, Ready).

respond(Sock, Sid, #{enc := Enc} = St) ->
{HBlock, Enc2} = h2_hpack:encode([{<<":status">>, <<"200">>}], Enc),
send(Sock, h2_frame:headers(Sid, HBlock, false)),
send(Sock, h2_frame:data(Sid, integer_to_binary(Sid), true)),
St#{enc := Enc2}.

send(Sock, FrameData) -> ssl:send(Sock, h2_frame:encode(FrameData)).

cert_dir() ->
BeamDir = filename:dirname(code:which(?MODULE)),
Root = filename:join([BeamDir, "..", "..", "..", "..", ".."]),
filename:join([filename:absname(Root), "test", "certs"]).
Loading
Loading