From 50a6e0ba4c8f865cc5d12bffb06ae7e2c9696b25 Mon Sep 17 00:00:00 2001 From: Jianbo He Date: Wed, 2 Sep 2026 11:25:31 +0800 Subject: [PATCH 1/3] fix: isolate callback failures from HTTP workers --- src/ehttpc.erl | 188 ++++++++++++++++++++++++++++++------ test/ehttpc_async_tests.erl | 39 ++++++++ test/ehttpc_tests.erl | 135 ++++++++++++++++++++++++++ 3 files changed, 330 insertions(+), 32 deletions(-) diff --git a/src/ehttpc.erl b/src/ehttpc.erl index 5773c12..f4e9e0a 100644 --- a/src/ehttpc.erl +++ b/src/ehttpc.erl @@ -65,15 +65,23 @@ -type path() :: binary() | string(). -type headers() :: [{binary(), iodata()}]. -type body() :: iodata(). +-type callback_fun() :: {function(), [term()]}. -type callback() :: - {function(), list()} + callback_fun() | #{ %% where to send the final results - final_reply := {function(), list()}, + final_reply := callback_fun(), %% optional; sends the gun stream ref and the worker pid so the caller may cancel %% it. - stream_ref => {function(), list()} + stream_ref => callback_fun() }. +-type normalized_callback() :: #{ + final_reply := callback_fun(), + stream_ref => callback_fun() +}. +-type reply_target() :: + {sync, gen_server:from()} + | {async, normalized_callback()}. -type request() :: path() | {path(), headers()} | {path(), headers(), body()}. -include_lib("snabbkaffe/include/snabbkaffe.hrl"). @@ -226,9 +234,14 @@ mk_request(delete = Method, Req, ExpireAt) when ?IS_HEADERS_REQ(Req) -> %% response is received. -spec request_async(pid(), method(), request(), timeout(), callback()) -> ok. request_async(Worker, Method, Request, Timeout, ResultCallback) when is_pid(Worker) -> - ExpireAt = fresh_expire_at(Timeout), - _ = erlang:send(Worker, mk_async_request(Method, Request, ExpireAt, ResultCallback)), - ok. + case normalize_callback(ResultCallback) of + {ok, Callback} -> + ExpireAt = fresh_expire_at(Timeout), + _ = erlang:send(Worker, mk_async_request(Method, Request, ExpireAt, Callback)), + ok; + {error, Reason} -> + error({invalid_callback, Reason}) + end. mk_async_request(head = Method, Req, ExpireAt, RC) when ?IS_HEADERS_REQ(Req) -> ?ASYNC_REQ(Method, Req, ExpireAt, RC); @@ -323,7 +336,7 @@ handle_call({health_check, Timeout}, _From, State = #state{client = Client, gun_ end ); handle_call(?REQ(_Method, _Request, _ExpireAt) = Req, From, State0) -> - State1 = enqueue_req(From, Req, State0), + State1 = enqueue_req({sync, From}, Req, State0), State = maybe_shoot(State1), {noreply, State}; handle_call(Call, _From, State0) -> @@ -340,7 +353,7 @@ handle_cast(_Msg, State0) -> handle_info(?ASYNC_REQ(Method, Request, ExpireAt, ResultCallback), State0) -> Req = ?REQ(Method, Request, ExpireAt), - State1 = enqueue_req(ResultCallback, Req, State0), + State1 = enqueue_async_req(ResultCallback, Req, State0), State = maybe_shoot(State1), {noreply, State}; handle_info({suspend, Time}, State) -> @@ -664,13 +677,13 @@ take_sent_req(StreamRef, #{sent := Sent, max_sent_expire := T} = Requests) -> end end. -is_sent_req_expired(?SENT_REQ(_From, infinity = _ExpireAt, _), _Now) -> +is_sent_req_expired(?SENT_REQ({sync, _From}, infinity = _ExpireAt, _), _Now) -> false; -is_sent_req_expired(?SENT_REQ({Pid, _Ref}, ExpireAt, _), Now) when is_pid(Pid) -> +is_sent_req_expired(?SENT_REQ({sync, {Pid, _Ref}}, ExpireAt, _), Now) when is_pid(Pid) -> %% for gen_server:call, it is aborted after timeout, there is no need to send %% reply to the caller Now > ExpireAt orelse (not erlang:is_process_alive(Pid)); -is_sent_req_expired(?SENT_REQ(_, _, _), _) -> +is_sent_req_expired(?SENT_REQ({async, _Callback}, _, _), _) -> %% for async requests, there is no way to tell if the caller %% the provided result-callback should be evaluated or not, %% to be on the safe side, we never consider sent async-requests expired. @@ -718,11 +731,8 @@ drop_expired(#{pending := Pending, pending_count := PC} = Requests, Now) -> end. %% For async-request, we evaluate the result-callback with {error, timeout} -maybe_reply_timeout({F, A}) when is_function(F) -> - _ = erlang:apply(F, A ++ [{error, timeout}]), - ok; -maybe_reply_timeout(#{final_reply := {F, A}}) when is_function(F) -> - _ = erlang:apply(F, A ++ [{error, timeout}]), +maybe_reply_timeout({async, Callback}) -> + reply_async(Callback, {error, timeout}, final_reply), ok; maybe_reply_timeout(_) -> %% This is not a callback, but the gen_server:call's From @@ -731,6 +741,7 @@ maybe_reply_timeout(_) -> ok. %% enqueue the pending requests +-spec enqueue_req(reply_target(), term(), #state{}) -> #state{}. enqueue_req(ReplyTo, Req, #state{requests = Requests0} = State) -> #{ pending := Pending, @@ -741,6 +752,15 @@ enqueue_req(ReplyTo, Req, #state{requests = Requests0} = State) -> Requests = Requests0#{pending := NewPending, pending_count := PC + 1}, State#state{requests = drop_expired(Requests)}. +enqueue_async_req(Callback0, Req, State0) -> + case normalize_callback(Callback0) of + {ok, Callback} -> + enqueue_req({async, Callback}, Req, State0); + {error, Reason} -> + log_invalid_callback(Reason), + State0 + end. + %% call gun to shoot the request out maybe_shoot( #state{ @@ -952,12 +972,12 @@ gun_await_up(Pid, ExpireAt, Timeout, State0) -> {{error, Reason}, State0}; ?ASYNC_REQ(Method, Request, ExpireAt1, ResultCallback) -> Req = ?REQ(Method, Request, ExpireAt1), - State = enqueue_req(ResultCallback, Req, State0), + State = enqueue_async_req(ResultCallback, Req, State0), %% keep waiting NewTimeout = timeout(ExpireAt), gun_await_up(Pid, ExpireAt, NewTimeout, State); ?GEN_CALL_REQ(From, Call) -> - State = enqueue_req(From, Call, State0), + State = enqueue_req({sync, From}, Call, State0), %% keep waiting NewTimeout = timeout(ExpireAt), gun_await_up(Pid, ExpireAt, NewTimeout, State) @@ -983,12 +1003,12 @@ gun_await_tunnel(Pid, StreamRef, ExpireAt, Timeout, Headers, State0) -> {{error, {proxy_error, {StatusCode, Headers}}}, State0}; ?ASYNC_REQ(Method, Request, ExpireAt1, ResultCallback) -> Req = ?REQ(Method, Request, ExpireAt1), - State = enqueue_req(ResultCallback, Req, State0), + State = enqueue_async_req(ResultCallback, Req, State0), %% keep waiting NewTimeout = timeout(ExpireAt), gun_await_tunnel(Pid, StreamRef, ExpireAt, NewTimeout, Headers, State); ?GEN_CALL_REQ(From, Call) -> - State = enqueue_req(From, Call, State0), + State = enqueue_req({sync, From}, Call, State0), %% keep waiting NewTimeout = timeout(ExpireAt), gun_await_tunnel(Pid, StreamRef, ExpireAt, NewTimeout, Headers, State) @@ -1042,21 +1062,83 @@ handle_gun_reply(State, Client, StreamRef, IsFin, StatusCode, Headers, Data) -> end end. -reply({F, A}, Result) when is_function(F) -> - _ = erlang:apply(F, A ++ [Result]), +reply({sync, From}, Result) -> + gen_server:reply(From, Result); +reply({async, Callback}, Result) -> + reply_async(Callback, Result, final_reply), ok; -reply(#{final_reply := FinalReply}, Result) -> - %% assert - {F, A} = FinalReply, - _ = erlang:apply(F, A ++ [Result]), - ok; -reply(From, Result) -> - gen_server:reply(From, Result). +reply(_InvalidReplyTarget, _Result) -> + log_invalid_callback_target(), + ok. -maybe_send_stream_ref(#{stream_ref := {F, A}}, StreamRef) when is_function(F) -> - _ = erlang:apply(F, [StreamRef, self() | A]), +maybe_send_stream_ref({async, #{stream_ref := {F, A}}}, StreamRef) -> + safe_apply(stream_ref, F, [StreamRef, self() | A]), ok; -maybe_send_stream_ref(_ReplyTo, _StreamRef) -> +maybe_send_stream_ref(_, _) -> + ok. + +reply_async(#{final_reply := {F, A}}, Result, Kind) -> + safe_apply(Kind, F, A ++ [Result]); +reply_async(_Callback, _Result, _Kind) -> + log_invalid_callback_target(). + +safe_apply(Kind, F, Args) -> + try erlang:apply(F, Args) of + _ -> + ok + catch + Class:Reason:Stacktrace -> + logger:error(#{ + msg => "ehttpc_callback_failed", + callback_kind => Kind, + callback => callback_info(F), + class => Class, + reason_kind => callback_reason_kind(Class, Reason), + stacktrace => callback_stacktrace(Stacktrace) + }), + ok + end. + +callback_info(F) when is_function(F) -> + Info = erlang:fun_info(F), + #{ + module => proplists:get_value(module, Info), + name => proplists:get_value(name, Info), + arity => proplists:get_value(arity, Info) + }; +callback_info(_) -> + #{}. + +callback_reason_kind(error, {case_clause, _}) -> + case_clause; +callback_reason_kind(error, {badmatch, _}) -> + badmatch; +callback_reason_kind(error, function_clause) -> + function_clause; +callback_reason_kind(error, {badarity, _}) -> + badarity; +callback_reason_kind(error, undef) -> + undef; +callback_reason_kind(exit, _) -> + exit; +callback_reason_kind(throw, _) -> + throw; +callback_reason_kind(_, _) -> + other. + +callback_stacktrace([{M, F, A, _Info} | _]) when is_integer(A) -> + [{M, F, A}]; +callback_stacktrace([{M, F, Args, _Info} | _]) when is_list(Args) -> + [{M, F, length(Args)}]; +callback_stacktrace(_) -> + []. + +log_invalid_callback(Reason) -> + logger:error(#{msg => "ehttpc_invalid_callback", reason => Reason}), + ok. + +log_invalid_callback_target() -> + logger:error(#{msg => "ehttpc_invalid_callback_target"}), ok. peek_oldest_fn(#{prioritise_latest := true}) -> @@ -1120,6 +1202,48 @@ take_proplist(Key, Proplist0) -> {ValueFromProplist, Proplist1} end. +normalize_callback({F, A}) -> + case normalize_callback_fun({F, A}, 1) of + {ok, FinalReply} -> + {ok, #{final_reply => FinalReply}}; + {error, _} = Error -> + Error + end; +normalize_callback(#{final_reply := FinalReply} = Callback) -> + case normalize_callback_fun(FinalReply, 1) of + {ok, NormalizedFinalReply} -> + normalize_stream_ref_callback(Callback, NormalizedFinalReply); + {error, _} = Error -> + Error + end; +normalize_callback(_) -> + {error, invalid_callback}. + +normalize_stream_ref_callback(Callback, FinalReply) -> + case maps:find(stream_ref, Callback) of + error -> + {ok, #{final_reply => FinalReply}}; + {ok, StreamRef} -> + case normalize_callback_fun(StreamRef, 2) of + {ok, NormalizedStreamRef} -> + {ok, #{final_reply => FinalReply, stream_ref => NormalizedStreamRef}}; + {error, _} = Error -> + Error + end + end. + +normalize_callback_fun({F, A}, ExtraArity) when is_function(F), is_list(A) -> + case is_function(F, length(A) + ExtraArity) of + true -> + {ok, {F, A}}; + false -> + {error, invalid_callback_arity} + end; +normalize_callback_fun(_, 1) -> + {error, invalid_final_reply}; +normalize_callback_fun(_, 2) -> + {error, invalid_stream_ref}. + log(Level, Data, #state{host = Host, port = Port}) -> logger:log(Level, Data#{host => Host, port => Port}). diff --git a/test/ehttpc_async_tests.erl b/test/ehttpc_async_tests.erl index 4943494..5486387 100644 --- a/test/ehttpc_async_tests.erl +++ b/test/ehttpc_async_tests.erl @@ -45,6 +45,45 @@ send_10_async_test() -> PoolOpts = pool_opts(Port, false), true = ?WITH(ServerOpts, PoolOpts, req_async(10, 1000)). +timeout_callback_exception_does_not_kill_worker_test() -> + Port = ?PORT, + ServerOpts = #{port => Port, name => ?FUNCTION_NAME, delay => 3_000, oneoff => false}, + PoolOpts = pool_opts(Port, false), + ?WITH( + ServerOpts, + PoolOpts, + begin + Worker = ehttpc_pool:pick_worker(?POOL), + Tester = self(), + ok = ehttpc:request_async(Worker, get, req(), 5_000, {fun(_) -> ok end, []}), + %% Ensure the first request occupies the only in-flight slot. + _ = sys:get_state(Worker), + ok = ehttpc:request_async( + Worker, + get, + req(), + 10, + { + fun(Result) -> + Tester ! {timeout_callback, Result}, + error(timeout_callback_boom) + end, + [] + } + ), + timer:sleep(30), + ok = ehttpc:request_async(Worker, get, req(), 1_000, {fun(_) -> ok end, []}), + receive + {timeout_callback, {error, timeout}} -> ok + after 2_000 -> + error(timeout_callback_not_called) + end, + _ = sys:get_state(Worker), + ?assert(erlang:is_process_alive(Worker)), + ?assertEqual(1, length(ehttpc:workers(?POOL))) + end + ). + no_expired_req_send_test() -> Port = ?PORT, % infinity diff --git a/test/ehttpc_tests.erl b/test/ehttpc_tests.erl index 522e77f..98b1db8 100644 --- a/test/ehttpc_tests.erl +++ b/test/ehttpc_tests.erl @@ -104,6 +104,141 @@ send_100_test_() -> {"oneoff=false", fun() -> ?WITH(ServerOpts2, PoolOpts2, req_async(100)) end} ]. +callback_crash_does_not_kill_worker_test() -> + Port = ?PORT, + ServerOpts = #{port => Port, name => ?FUNCTION_NAME, delay => 0, oneoff => false}, + PoolOpts = pool_opts(Port, false), + ?WITH( + ServerOpts, + PoolOpts, + begin + Worker = ehttpc_pool:pick_worker(?POOL), + Tester = self(), + Callback = + { + fun(_Result) -> + Tester ! callback_called, + error(callback_boom) + end, + [] + }, + ok = ehttpc:request_async(Worker, get, req(), 1_000, Callback), + receive + callback_called -> ok + after 2_000 -> + error(callback_not_called) + end, + _ = sys:get_state(Worker), + ?assert(erlang:is_process_alive(Worker)), + ?assertMatch({ok, 200, _, _}, ehttpc:request(Worker, get, req(), 1_000, 0)) + end + ). + +structured_callback_exception_does_not_kill_worker_test() -> + Port = ?PORT, + ServerOpts = #{port => Port, name => ?FUNCTION_NAME, delay => 0, oneoff => false}, + PoolOpts = pool_opts(Port, false), + ?WITH( + ServerOpts, + PoolOpts, + begin + Worker = ehttpc_pool:pick_worker(?POOL), + Tester = self(), + Callback = #{ + final_reply => + { + fun(Result) -> + Tester ! {final_reply, Result}, + error(final_reply_callback_boom) + end, + [] + }, + stream_ref => + { + fun(_StreamRef, _Worker) -> + Tester ! stream_ref_called, + error(stream_ref_callback_boom) + end, + [] + } + }, + ok = ehttpc:request_async(Worker, get, req(), 1_000, Callback), + receive + stream_ref_called -> ok + after 2_000 -> + error(stream_ref_callback_not_called) + end, + receive + {final_reply, {ok, 200, _, _}} -> ok + after 2_000 -> + error(final_reply_not_called) + end, + _ = sys:get_state(Worker), + ?assertMatch({ok, 200, _, _}, ehttpc:request(Worker, get, req(), 1_000, 0)) + end + ). + +invalid_callback_does_not_kill_worker_test() -> + Port = ?PORT, + ServerOpts = #{port => Port, name => ?FUNCTION_NAME, delay => 0, oneoff => false}, + PoolOpts = pool_opts(Port, false), + ?WITH( + ServerOpts, + PoolOpts, + begin + Worker = ehttpc_pool:pick_worker(?POOL), + ?assertError( + {invalid_callback, invalid_callback}, + ehttpc:request_async(Worker, get, req(), 1_000, #{ + stream_ref => {fun() -> ok end, []} + }) + ), + _ = sys:get_state(Worker), + ?assertEqual(1, length(ehttpc:workers(?POOL))), + ?assertMatch({ok, 200, _, _}, ehttpc:request(Worker, get, req(), 1_000, 0)) + end + ). + +repeated_callback_exceptions_do_not_restart_worker_test() -> + Port = ?PORT, + ServerOpts = #{port => Port, name => ?FUNCTION_NAME, delay => 0, oneoff => false}, + PoolOpts = pool_opts(Port, false), + ?WITH( + ServerOpts, + PoolOpts, + begin + Worker = ehttpc_pool:pick_worker(?POOL), + Tester = self(), + lists:foreach( + fun(I) -> + Callback = + { + fun(_Result) -> + Tester ! {callback_called, I}, + error(callback_boom) + end, + [] + }, + ok = ehttpc:request_async(Worker, get, req(), 1_000, Callback) + end, + lists:seq(1, 20) + ), + lists:foreach( + fun(I) -> + receive + {callback_called, I} -> ok + after 2_000 -> + error({callback_not_called, I}) + end + end, + lists:seq(1, 20) + ), + _ = sys:get_state(Worker), + ?assertEqual([{{?POOL, 1}, Worker}], ehttpc:workers(?POOL)), + ?assertEqual(ok, ehttpc:check_pool_integrity(?POOL)) + end + ). + send_1000_async_pipeline_test_() -> TestTimeout = 30, Port = ?PORT, From 8a31d4b2cb8a1f55d4e3d045445308964ab3dc7e Mon Sep 17 00:00:00 2001 From: Jianbo He Date: Wed, 2 Sep 2026 15:57:34 +0800 Subject: [PATCH 2/3] perf: avoid reply target allocations --- src/ehttpc.erl | 121 ++++++++++++------------------------------ test/ehttpc_tests.erl | 28 +++++----- 2 files changed, 49 insertions(+), 100 deletions(-) diff --git a/src/ehttpc.erl b/src/ehttpc.erl index f4e9e0a..0843b3d 100644 --- a/src/ehttpc.erl +++ b/src/ehttpc.erl @@ -75,13 +75,6 @@ %% it. stream_ref => callback_fun() }. --type normalized_callback() :: #{ - final_reply := callback_fun(), - stream_ref => callback_fun() -}. --type reply_target() :: - {sync, gen_server:from()} - | {async, normalized_callback()}. -type request() :: path() | {path(), headers()} | {path(), headers(), body()}. -include_lib("snabbkaffe/include/snabbkaffe.hrl"). @@ -234,10 +227,10 @@ mk_request(delete = Method, Req, ExpireAt) when ?IS_HEADERS_REQ(Req) -> %% response is received. -spec request_async(pid(), method(), request(), timeout(), callback()) -> ok. request_async(Worker, Method, Request, Timeout, ResultCallback) when is_pid(Worker) -> - case normalize_callback(ResultCallback) of - {ok, Callback} -> + case validate_callback(ResultCallback) of + ok -> ExpireAt = fresh_expire_at(Timeout), - _ = erlang:send(Worker, mk_async_request(Method, Request, ExpireAt, Callback)), + _ = erlang:send(Worker, mk_async_request(Method, Request, ExpireAt, ResultCallback)), ok; {error, Reason} -> error({invalid_callback, Reason}) @@ -336,7 +329,7 @@ handle_call({health_check, Timeout}, _From, State = #state{client = Client, gun_ end ); handle_call(?REQ(_Method, _Request, _ExpireAt) = Req, From, State0) -> - State1 = enqueue_req({sync, From}, Req, State0), + State1 = enqueue_req(From, Req, State0), State = maybe_shoot(State1), {noreply, State}; handle_call(Call, _From, State0) -> @@ -353,7 +346,7 @@ handle_cast(_Msg, State0) -> handle_info(?ASYNC_REQ(Method, Request, ExpireAt, ResultCallback), State0) -> Req = ?REQ(Method, Request, ExpireAt), - State1 = enqueue_async_req(ResultCallback, Req, State0), + State1 = enqueue_req(ResultCallback, Req, State0), State = maybe_shoot(State1), {noreply, State}; handle_info({suspend, Time}, State) -> @@ -677,13 +670,13 @@ take_sent_req(StreamRef, #{sent := Sent, max_sent_expire := T} = Requests) -> end end. -is_sent_req_expired(?SENT_REQ({sync, _From}, infinity = _ExpireAt, _), _Now) -> +is_sent_req_expired(?SENT_REQ(_From, infinity = _ExpireAt, _), _Now) -> false; -is_sent_req_expired(?SENT_REQ({sync, {Pid, _Ref}}, ExpireAt, _), Now) when is_pid(Pid) -> +is_sent_req_expired(?SENT_REQ({Pid, _Ref}, ExpireAt, _), Now) when is_pid(Pid) -> %% for gen_server:call, it is aborted after timeout, there is no need to send %% reply to the caller Now > ExpireAt orelse (not erlang:is_process_alive(Pid)); -is_sent_req_expired(?SENT_REQ({async, _Callback}, _, _), _) -> +is_sent_req_expired(?SENT_REQ(_, _, _), _) -> %% for async requests, there is no way to tell if the caller %% the provided result-callback should be evaluated or not, %% to be on the safe side, we never consider sent async-requests expired. @@ -731,8 +724,11 @@ drop_expired(#{pending := Pending, pending_count := PC} = Requests, Now) -> end. %% For async-request, we evaluate the result-callback with {error, timeout} -maybe_reply_timeout({async, Callback}) -> - reply_async(Callback, {error, timeout}, final_reply), +maybe_reply_timeout({F, A}) when is_function(F), is_list(A) -> + safe_apply(final_reply, F, A ++ [{error, timeout}]), + ok; +maybe_reply_timeout(#{final_reply := {F, A}}) when is_function(F), is_list(A) -> + safe_apply(final_reply, F, A ++ [{error, timeout}]), ok; maybe_reply_timeout(_) -> %% This is not a callback, but the gen_server:call's From @@ -741,7 +737,6 @@ maybe_reply_timeout(_) -> ok. %% enqueue the pending requests --spec enqueue_req(reply_target(), term(), #state{}) -> #state{}. enqueue_req(ReplyTo, Req, #state{requests = Requests0} = State) -> #{ pending := Pending, @@ -752,15 +747,6 @@ enqueue_req(ReplyTo, Req, #state{requests = Requests0} = State) -> Requests = Requests0#{pending := NewPending, pending_count := PC + 1}, State#state{requests = drop_expired(Requests)}. -enqueue_async_req(Callback0, Req, State0) -> - case normalize_callback(Callback0) of - {ok, Callback} -> - enqueue_req({async, Callback}, Req, State0); - {error, Reason} -> - log_invalid_callback(Reason), - State0 - end. - %% call gun to shoot the request out maybe_shoot( #state{ @@ -972,12 +958,12 @@ gun_await_up(Pid, ExpireAt, Timeout, State0) -> {{error, Reason}, State0}; ?ASYNC_REQ(Method, Request, ExpireAt1, ResultCallback) -> Req = ?REQ(Method, Request, ExpireAt1), - State = enqueue_async_req(ResultCallback, Req, State0), + State = enqueue_req(ResultCallback, Req, State0), %% keep waiting NewTimeout = timeout(ExpireAt), gun_await_up(Pid, ExpireAt, NewTimeout, State); ?GEN_CALL_REQ(From, Call) -> - State = enqueue_req({sync, From}, Call, State0), + State = enqueue_req(From, Call, State0), %% keep waiting NewTimeout = timeout(ExpireAt), gun_await_up(Pid, ExpireAt, NewTimeout, State) @@ -1003,12 +989,12 @@ gun_await_tunnel(Pid, StreamRef, ExpireAt, Timeout, Headers, State0) -> {{error, {proxy_error, {StatusCode, Headers}}}, State0}; ?ASYNC_REQ(Method, Request, ExpireAt1, ResultCallback) -> Req = ?REQ(Method, Request, ExpireAt1), - State = enqueue_async_req(ResultCallback, Req, State0), + State = enqueue_req(ResultCallback, Req, State0), %% keep waiting NewTimeout = timeout(ExpireAt), gun_await_tunnel(Pid, StreamRef, ExpireAt, NewTimeout, Headers, State); ?GEN_CALL_REQ(From, Call) -> - State = enqueue_req({sync, From}, Call, State0), + State = enqueue_req(From, Call, State0), %% keep waiting NewTimeout = timeout(ExpireAt), gun_await_tunnel(Pid, StreamRef, ExpireAt, NewTimeout, Headers, State) @@ -1062,26 +1048,24 @@ handle_gun_reply(State, Client, StreamRef, IsFin, StatusCode, Headers, Data) -> end end. -reply({sync, From}, Result) -> +reply({F, A}, Result) when is_function(F), is_list(A) -> + safe_apply(final_reply, F, A ++ [Result]); +reply(#{final_reply := {F, A}}, Result) when is_function(F), is_list(A) -> + safe_apply(final_reply, F, A ++ [Result]); +reply({Pid, _Tag} = From, Result) when is_pid(Pid) -> gen_server:reply(From, Result); -reply({async, Callback}, Result) -> - reply_async(Callback, Result, final_reply), - ok; reply(_InvalidReplyTarget, _Result) -> log_invalid_callback_target(), ok. -maybe_send_stream_ref({async, #{stream_ref := {F, A}}}, StreamRef) -> +maybe_send_stream_ref(#{stream_ref := {F, A}}, StreamRef) when + is_function(F), is_list(A) +-> safe_apply(stream_ref, F, [StreamRef, self() | A]), ok; maybe_send_stream_ref(_, _) -> ok. -reply_async(#{final_reply := {F, A}}, Result, Kind) -> - safe_apply(Kind, F, A ++ [Result]); -reply_async(_Callback, _Result, _Kind) -> - log_invalid_callback_target(). - safe_apply(Kind, F, Args) -> try erlang:apply(F, Args) of _ -> @@ -1105,9 +1089,7 @@ callback_info(F) when is_function(F) -> module => proplists:get_value(module, Info), name => proplists:get_value(name, Info), arity => proplists:get_value(arity, Info) - }; -callback_info(_) -> - #{}. + }. callback_reason_kind(error, {case_clause, _}) -> case_clause; @@ -1133,10 +1115,6 @@ callback_stacktrace([{M, F, Args, _Info} | _]) when is_list(Args) -> callback_stacktrace(_) -> []. -log_invalid_callback(Reason) -> - logger:error(#{msg => "ehttpc_invalid_callback", reason => Reason}), - ok. - log_invalid_callback_target() -> logger:error(#{msg => "ehttpc_invalid_callback_target"}), ok. @@ -1202,48 +1180,15 @@ take_proplist(Key, Proplist0) -> {ValueFromProplist, Proplist1} end. -normalize_callback({F, A}) -> - case normalize_callback_fun({F, A}, 1) of - {ok, FinalReply} -> - {ok, #{final_reply => FinalReply}}; - {error, _} = Error -> - Error - end; -normalize_callback(#{final_reply := FinalReply} = Callback) -> - case normalize_callback_fun(FinalReply, 1) of - {ok, NormalizedFinalReply} -> - normalize_stream_ref_callback(Callback, NormalizedFinalReply); - {error, _} = Error -> - Error - end; -normalize_callback(_) -> +validate_callback({F, A}) when is_function(F), is_list(A) -> + ok; +validate_callback(#{final_reply := {F, A}}) when + is_function(F), is_list(A) +-> + ok; +validate_callback(_) -> {error, invalid_callback}. -normalize_stream_ref_callback(Callback, FinalReply) -> - case maps:find(stream_ref, Callback) of - error -> - {ok, #{final_reply => FinalReply}}; - {ok, StreamRef} -> - case normalize_callback_fun(StreamRef, 2) of - {ok, NormalizedStreamRef} -> - {ok, #{final_reply => FinalReply, stream_ref => NormalizedStreamRef}}; - {error, _} = Error -> - Error - end - end. - -normalize_callback_fun({F, A}, ExtraArity) when is_function(F), is_list(A) -> - case is_function(F, length(A) + ExtraArity) of - true -> - {ok, {F, A}}; - false -> - {error, invalid_callback_arity} - end; -normalize_callback_fun(_, 1) -> - {error, invalid_final_reply}; -normalize_callback_fun(_, 2) -> - {error, invalid_stream_ref}. - log(Level, Data, #state{host = Host, port = Port}) -> logger:log(Level, Data#{host => Host, port => Port}). diff --git a/test/ehttpc_tests.erl b/test/ehttpc_tests.erl index 98b1db8..f76c7cf 100644 --- a/test/ehttpc_tests.erl +++ b/test/ehttpc_tests.erl @@ -116,15 +116,15 @@ callback_crash_does_not_kill_worker_test() -> Tester = self(), Callback = { - fun(_Result) -> - Tester ! callback_called, + fun(Marker, _Result) -> + Tester ! {callback_called, Marker}, error(callback_boom) end, - [] + [legacy_marker] }, ok = ehttpc:request_async(Worker, get, req(), 1_000, Callback), receive - callback_called -> ok + {callback_called, legacy_marker} -> ok after 2_000 -> error(callback_not_called) end, @@ -147,29 +147,29 @@ structured_callback_exception_does_not_kill_worker_test() -> Callback = #{ final_reply => { - fun(Result) -> - Tester ! {final_reply, Result}, + fun(Marker, Result) -> + Tester ! {final_reply, Marker, Result}, error(final_reply_callback_boom) end, - [] + [final_marker] }, stream_ref => { - fun(_StreamRef, _Worker) -> - Tester ! stream_ref_called, + fun(StreamRef, WorkerPid, Marker) -> + Tester ! {stream_ref_called, StreamRef, WorkerPid, Marker}, error(stream_ref_callback_boom) end, - [] + [stream_marker] } }, ok = ehttpc:request_async(Worker, get, req(), 1_000, Callback), receive - stream_ref_called -> ok + {stream_ref_called, _, Worker, stream_marker} -> ok after 2_000 -> error(stream_ref_callback_not_called) end, receive - {final_reply, {ok, 200, _, _}} -> ok + {final_reply, final_marker, {ok, 200, _, _}} -> ok after 2_000 -> error(final_reply_not_called) end, @@ -193,6 +193,10 @@ invalid_callback_does_not_kill_worker_test() -> stream_ref => {fun() -> ok end, []} }) ), + ?assertError( + {invalid_callback, invalid_callback}, + ehttpc:request_async(Worker, get, req(), 1_000, {self(), make_ref()}) + ), _ = sys:get_state(Worker), ?assertEqual(1, length(ehttpc:workers(?POOL))), ?assertMatch({ok, 200, _, _}, ehttpc:request(Worker, get, req(), 1_000, 0)) From c7c88f797dc6ddad830ad593d0e9d8f51e4188c3 Mon Sep 17 00:00:00 2001 From: Jianbo He Date: Thu, 3 Sep 2026 10:03:27 +0800 Subject: [PATCH 3/3] fix: validate callback arity before dispatch --- src/ehttpc.erl | 26 +++++++++++++++++++++++--- test/ehttpc_tests.erl | 37 +++++++++++++++++++++++++++++++++++++ 2 files changed, 60 insertions(+), 3 deletions(-) diff --git a/src/ehttpc.erl b/src/ehttpc.erl index 0843b3d..57b7fb4 100644 --- a/src/ehttpc.erl +++ b/src/ehttpc.erl @@ -1181,14 +1181,34 @@ take_proplist(Key, Proplist0) -> end. validate_callback({F, A}) when is_function(F), is_list(A) -> - ok; -validate_callback(#{final_reply := {F, A}}) when + validate_callback_arity(final_reply, F, length(A) + 1); +validate_callback(#{final_reply := {F, A}} = Callback) when is_function(F), is_list(A) -> - ok; + case validate_callback_arity(final_reply, F, length(A) + 1) of + ok -> + validate_stream_ref_callback(Callback); + Error -> + Error + end; validate_callback(_) -> {error, invalid_callback}. +validate_stream_ref_callback(#{stream_ref := {F, A}}) when + is_function(F), is_list(A) +-> + validate_callback_arity(stream_ref, F, length(A) + 2); +validate_stream_ref_callback(_) -> + ok. + +validate_callback_arity(Kind, F, ExpectedArity) -> + case is_function(F, ExpectedArity) of + true -> + ok; + false -> + {error, {invalid_callback_arity, Kind, ExpectedArity}} + end. + log(Level, Data, #state{host = Host, port = Port}) -> logger:log(Level, Data#{host => Host, port => Port}). diff --git a/test/ehttpc_tests.erl b/test/ehttpc_tests.erl index f76c7cf..c309d4b 100644 --- a/test/ehttpc_tests.erl +++ b/test/ehttpc_tests.erl @@ -197,12 +197,49 @@ invalid_callback_does_not_kill_worker_test() -> {invalid_callback, invalid_callback}, ehttpc:request_async(Worker, get, req(), 1_000, {self(), make_ref()}) ), + ?assertError( + {invalid_callback, {invalid_callback_arity, final_reply, 1}}, + ehttpc:request_async(Worker, get, req(), 1_000, {fun() -> ok end, []}) + ), + ?assertError( + {invalid_callback, {invalid_callback_arity, stream_ref, 2}}, + ehttpc:request_async(Worker, get, req(), 1_000, #{ + final_reply => {fun(_) -> ok end, []}, + stream_ref => {fun() -> ok end, []} + }) + ), _ = sys:get_state(Worker), ?assertEqual(1, length(ehttpc:workers(?POOL))), ?assertMatch({ok, 200, _, _}, ehttpc:request(Worker, get, req(), 1_000, 0)) end ). +invalid_optional_stream_ref_is_ignored_test() -> + Port = ?PORT, + ServerOpts = #{port => Port, name => ?FUNCTION_NAME, delay => 0, oneoff => false}, + PoolOpts = pool_opts(Port, false), + ?WITH( + ServerOpts, + PoolOpts, + begin + Worker = ehttpc_pool:pick_worker(?POOL), + Tester = self(), + Callback = #{ + final_reply => + {fun(Result) -> Tester ! {final_reply, Result} end, []}, + stream_ref => undefined + }, + ok = ehttpc:request_async(Worker, get, req(), 1_000, Callback), + receive + {final_reply, {ok, 200, _, _}} -> ok + after 2_000 -> + error(final_reply_not_called) + end, + _ = sys:get_state(Worker), + ?assertMatch({ok, 200, _, _}, ehttpc:request(Worker, get, req(), 1_000, 0)) + end + ). + repeated_callback_exceptions_do_not_restart_worker_test() -> Port = ?PORT, ServerOpts = #{port => Port, name => ?FUNCTION_NAME, delay => 0, oneoff => false},