diff --git a/src/pot_base32.erl b/src/pot_base32.erl index 408c6da..dd454c8 100644 --- a/src/pot_base32.erl +++ b/src/pot_base32.erl @@ -89,25 +89,46 @@ decode(Bin, Opts) when is_binary(Bin) andalso is_list(Opts) -> decode(List, Opts) when is_list(List) andalso is_list(Opts) -> decode(list_to_binary(List), Opts). -decode(Fun, <>, Bits) -> - <>; -decode(Fun, <>, Bits) -> - <>; -decode(Fun, <>, Bits) -> - <>; -decode(Fun, <>, Bits) -> + decode_final(Fun, [A, B], Bits, 2); +decode(Fun, <>, Bits) -> + decode_final(Fun, [A, B, C, D], Bits, 4); +decode(Fun, <>, Bits) -> + decode_final(Fun, [A, B, C, D, E], Bits, 1); +decode(Fun, <>, _Bits) -> + error(badarg); decode(Fun, <>, Bits) -> decode(Fun, Rest, <>); -decode(_Fun, <<>>, Bin) -> Bin. +decode(_Fun, <<>>, Bits) -> + trim_padding(Bits, bit_size(Bits) rem 8). + +decode_final(Fun, [X | Rest], Bits, UnusedBits) -> + decode_final(Fun, Rest, <>, UnusedBits); +decode_final(_Fun, [], Bits, UnusedBits) -> + trim_padding(Bits, UnusedBits). + +trim_padding(Bits, 0) -> Bits; +trim_padding(Bits, UnusedBits) when UnusedBits >= 1, UnusedBits =< 4 -> + BitSize = bit_size(Bits) - UnusedBits, + <> = Bits, + case Unused of + 0 -> Decoded; + _ -> error(badarg) + end; +trim_padding(_Bits, _UnusedBits) -> + error(badarg). std_dec(I) when I >= $2 andalso I =< $7 -> I - 24; std_dec(I) when I >= $a andalso I =< $z -> I - $a; -std_dec(I) when I >= $A andalso I =< $Z -> I - $A. +std_dec(I) when I >= $A andalso I =< $Z -> I - $A; +std_dec(_I) -> error(badarg). hex_dec(I) when I >= $0 andalso I =< $9 -> I - 48; -hex_dec(I) when I >= $a andalso I =< $z -> I - 87; -hex_dec(I) when I >= $A andalso I =< $Z -> I - 55. +hex_dec(I) when I >= $a andalso I =< $v -> I - 87; +hex_dec(I) when I >= $A andalso I =< $V -> I - 55; +hex_dec(_I) -> error(badarg). -ifdef(TEST). @@ -147,10 +168,36 @@ std_encode_nopad_test_() -> [ ?_assertEqual(Out, encode(In, [nopad])) || {In, Out} <- nopad_cases(std_cases()) ]. +std_decode_nopad_test_() -> + [ ?_assertEqual(Out, decode(In)) + || {Out, In} <- nopad_cases(std_cases()) ]. + +std_decode_nopad_byte_aligned_test_() -> + [ ?_assertEqual(8 * byte_size(Out), bit_size(decode(In))) + || {Out, In} <- nopad_cases(std_cases()) ]. + std_encode_lower_nopad_test_() -> [ ?_assertEqual(Out, encode(In, [lower,nopad])) || {In, Out} <- nopad_cases(lower_cases(std_cases())) ]. +std_decode_malformed_padding_test_() -> + [ ?_assertError(badarg, decode(<<"M=">>)), + ?_assertError(badarg, decode(<<"MY=">>)), + ?_assertError(badarg, decode(<<"MY=====">>)), + ?_assertError(badarg, decode(<<"M=Y=====">>)) ]. + +std_decode_non_zero_unused_bits_test_() -> + [ ?_assertError(badarg, decode(<<"MZ======">>)), + ?_assertError(badarg, decode(<<"MZXZ====">>)), + ?_assertError(badarg, decode(<<"MZXW7===">>)), + ?_assertError(badarg, decode(<<"MZXW6YZ=">>)) ]. + +std_decode_invalid_alphabet_test_() -> + [ ?_assertError(badarg, decode(<<"M0======">>)), + ?_assertError(badarg, decode(<<"M1======">>)), + ?_assertError(badarg, decode(<<"M8======">>)), + ?_assertError(badarg, decode(<<"M+======">>)) ]. + std_encode_string_test_() -> [ ?_assertEqual(Out, encode(In)) || {In, Out} <- stringinput_cases(std_cases()) ]. @@ -186,10 +233,36 @@ hex_encode_nopad_test_() -> [ ?_assertEqual(Out, encode(In, [hex,nopad])) || {In, Out} <- nopad_cases(hex_cases()) ]. +hex_decode_nopad_test_() -> + [ ?_assertEqual(Out, decode(In, [hex])) + || {Out, In} <- nopad_cases(hex_cases()) ]. + +hex_decode_nopad_byte_aligned_test_() -> + [ ?_assertEqual(8 * byte_size(Out), bit_size(decode(In, [hex]))) + || {Out, In} <- nopad_cases(hex_cases()) ]. + hex_encode_lower_nopad_test_() -> [ ?_assertEqual(Out, encode(In, [hex,lower,nopad])) || {In, Out} <- nopad_cases(lower_cases(hex_cases())) ]. +hex_decode_malformed_padding_test_() -> + [ ?_assertError(badarg, decode(<<"C=">>, [hex])), + ?_assertError(badarg, decode(<<"CO=">>, [hex])), + ?_assertError(badarg, decode(<<"CO=====">>, [hex])), + ?_assertError(badarg, decode(<<"C=O=====">>, [hex])) ]. + +hex_decode_non_zero_unused_bits_test_() -> + [ ?_assertError(badarg, decode(<<"CP======">>, [hex])), + ?_assertError(badarg, decode(<<"CPNH====">>, [hex])), + ?_assertError(badarg, decode(<<"CPNMV===">>, [hex])), + ?_assertError(badarg, decode(<<"CPNMUOH=">>, [hex])) ]. + +hex_decode_invalid_alphabet_test_() -> + [ ?_assertError(badarg, decode(<<"CW======">>, [hex])), + ?_assertError(badarg, decode(<<"CX======">>, [hex])), + ?_assertError(badarg, decode(<<"CY======">>, [hex])), + ?_assertError(badarg, decode(<<"CZ======">>, [hex])) ]. + hex_encode_string_test_() -> [ ?_assertEqual(Out, encode(In, [hex])) || {In, Out} <- stringinput_cases(hex_cases()) ].