-module(my_request). -author('Manuel Rubio '). -behaviour(gen_fsm). -define(SERVER, ?MODULE). -include("myproto.hrl"). -export([ start/4, check_clean_pass/2, check_sha1_pass/2, sha1_hex/1, to_hex/1, % states auth/2, normal/2, % FSM callbacks init/1, handle_sync_event/4, handle_event/3, handle_info/3, terminate/3, code_change/4 ]). -record(state, { socket :: gen_tcp:socket(), %% TCP connection id :: integer(), %% connection id hash :: binary(), %% hash for auth handler :: atom(), %% Handler for auth/queries msg = <<>> :: binary(), %% Received sql query when partial packet = <<>> :: binary(), %% Received packet parse_query = false :: boolean(), %% parse the received string or not handler_state }). %% API -spec start(Socket :: gen_tcp:socket(), Id :: pos_integer(), Handler :: atom(), ParseQuery :: boolean()) -> {ok, pid()}. start(Socket, Id, Handler, ParseQuery) -> {ok, Pid} = gen_fsm:start(?MODULE, [Socket, Id, Handler, ParseQuery], []), gen_tcp:controlling_process(Socket, Pid), inet:setopts(Socket, [{active, true}]), {ok, Pid}. -spec sha1_hex(Data :: binary()) -> binary(). sha1_hex(Data) -> to_hex(crypto:hash(sha, Data)). -spec to_hex(Hash :: binary() | undefined) -> binary(). to_hex(<<>>) -> <<"undefined">>; to_hex(undefined) -> <<"undefined">>; to_hex(<>) -> list_to_binary(io_lib:format("~40.16.0b", [X])). -spec check_sha1_pass(Pass::binary(), Salt::binary()) -> binary(). check_sha1_pass(Stage, Salt) -> Res = crypto:hash_final( crypto:hash_update( crypto:hash_update(crypto:hash_init(sha), Salt), crypto:hash(sha, Stage) ) ), crypto:exor(Stage, Res). -spec check_clean_pass(Pass::binary(), Salt::binary()) -> binary(). check_clean_pass(Pass, Salt) -> check_sha1_pass(crypto:hash(sha, Pass), Salt). %% callbacks init([Socket, Id, Handler, ParseQuery]) -> Hash = list_to_binary( lists:map(fun (0) -> 1; (X) -> X end, binary_to_list( crypto:strong_rand_bytes(20) )) ), Hello = #response{ id = Id, status = ?STATUS_HELLO, info = Hash }, ok = my_response:send_or_reply(Hello, Socket), {ok, auth, #state{ socket = Socket, id = Id, hash = Hash, handler = Handler, parse_query = ParseQuery}}. auth(#request{info = #user{password = Password} = User}, #state{hash = Hash, socket = Socket, handler = Handler} = StateData) -> ?DEBUG("Hash=~p; Pass=~p~n", [to_hex(Hash), to_hex(Password)]), case Handler:check_pass(User#user{server_hash = Hash}) of {ok, Password, HandlerState} -> Response = #response{status = ?STATUS_OK, status_flags = ?SERVER_STATUS_AUTOCOMMIT, id = 2}, ok = my_response:send_or_reply(Response, Socket), {next_state, normal, StateData#state{handler_state = HandlerState}}; {error, Reason} -> Response = #response{status = ?STATUS_ERR, error_code = 1047, info = Reason, id = 2}, ok = my_response:send_or_reply(Response, Socket), gen_tcp:close(Socket), {stop, normal, StateData}; {error, Code, Reason} -> Response = #response{status = ?STATUS_ERR, error_code = Code, info = Reason, id = 2}, ok = my_response:send_or_reply(Response, Socket), gen_tcp:close(Socket), {stop, normal, StateData}; {error, Code, SQLState, Reason} -> Response = #response{status = ?STATUS_ERR, error_code = Code, error_info = SQLState, info = Reason, id = 2}, ok = my_response:send_or_reply(Response, Socket), gen_tcp:close(Socket), {stop, normal, StateData} end. normal(#request{id = Id, info = Info, command = Command} = Request, #state{socket = Socket, handler = Handler, packet = Packet, handler_state = HandlerState} = StateData) -> ?DEBUG("Received: ~p~n", [Request]), FullPacket = <>, ParsedRequest = case StateData#state.parse_query andalso Command =:= ?COM_QUERY of false -> Request#request{info = FullPacket}; true -> case mysql_parser:parse(FullPacket) of {fail,Expected} -> ?ERROR_MSG("SQL invalid: ~p~n", [Expected]), Request#request{info = FullPacket}; {_, Extra, Where} -> ?ERROR_MSG("SQL error: ~p ~p~n", [Extra, Where]), Request#request{info = FullPacket}; Parsed -> Request#request{info = Parsed} end end, NewStateData = StateData#state{packet = <<>>}, case Handler:execute(ParsedRequest, HandlerState) of {noreply, NewHandlerState} -> {next_state, normal, NewStateData#state{handler_state = NewHandlerState}}; {reply, #response{} = Response, NewHandlerState} -> ok = my_response:send_or_reply(Response#response{id = Id + 1}, Socket), {next_state, normal, NewStateData#state{handler_state = NewHandlerState}}; {reply, default, NewHandlerState} -> {reply, Response, ModHandlerState} = my_response:default_reply(ParsedRequest, Handler, NewHandlerState), error_logger:info_msg("default reply: ~p~n", [Response]), ok = my_response:send_or_reply(Response#response{id = Id + 1}, Socket), {next_state, normal, NewStateData#state{handler_state = ModHandlerState}}; {stop, Reason, NewHandlerState} -> {stop, Reason, NewStateData#state{handler_state = NewHandlerState}} end. handle_info({tcp, _Port, Msg}, auth, #state{msg = PrevMsg} = StateData) -> Msg2 = <>, process_packet(Msg2, my_packet:decode_auth(Msg2), auth, StateData); handle_info({tcp, _Port, Msg}, normal, #state{msg = PrevMsg} = StateData) -> Msg2 = <>, process_packet(Msg2, my_packet:decode(Msg2), normal, StateData); handle_info({tcp_closed, _Socket}, _StateName, #state{id = Id} = StateData) -> ?INFO_MSG("Connection ID#~w closed~n", [Id]), {stop, normal, StateData}; handle_info(Info, _StateName, StateData = #state{socket = Socket}) -> ?ERROR_MSG("unknown message: ~p~n", [Info]), gen_tcp:close(Socket), {stop, normal, StateData}. handle_event(_Event, StateName, StateData) -> {next_state, StateName, StateData}. handle_sync_event(_Event, _From, StateName, StateData) -> {reply, ok, StateName, StateData}. terminate(Reason, _StateName, #state{handler = Handler, handler_state = HandlerState}) -> Handler:terminate(Reason, HandlerState), ok. code_change(_OldVsn, StateName, StateData, _Extra) -> {ok, StateName, StateData}. process_packet(_, {ok, #request{continue = true, info = Info} = Request, <<>>}, _StateName, StateData) -> ?DEBUG("Received (partial): ~p~n", [Request]), Packet = StateData#state.packet, {next_state, normal, StateData#state{packet = <>}}; process_packet(Msg, {more, _NumBytes}, _StateName, StateData) -> ?DEBUG("Received (partial): bytes remaining = ~w~n", [_NumBytes]), {next_state, normal, StateData#state{msg = Msg}}; process_packet(_Msg, {ok, #request{} = Request, <<>>}, normal, StateData) -> normal(Request, StateData); process_packet(_Msg, {ok, #request{} = Request, <<>>}, auth, StateData) -> auth(Request, StateData).