Current section

Files

Jump to
postgrex lib postgrex type_module.ex
Raw

lib/postgrex/type_module.ex

defmodule Postgrex.TypeModule do
@moduledoc false
alias Postgrex.TypeInfo
@doc """
Creates a module to aid in type processing.
The resulting module has three main purposes:
1. Associate the type information from the bootstrap query to extensions
2. Encode Elixir data types into binaries to send to the Postgres server
3. Decode binaries sent from the Postgres server into Elixir data types
The first point is handled by creating a `find/2` function that accepts a `%Postgrex.TypeInfo{}`
struct and a format such as `:binary`, `:text`, `:any`. It returns either a 2-tuple for
regular extensions containing the format and the extension name or a 3-tuple for super
extensions containing the format, the extension name and the sub-oids. Some important points:
- The type infos are associated to extensions using the `matching/1` callback
- The sub-oids for super extensions are created using the `oids/2` callback
The second point is handled by creating encoding functions for each extension. These functions
are created by iterating over each extension, taking the pattern returned by the `encode/1`
callback and creating a function that either:
- Accepts the pattern for regular extensions
- Accepts the pattern, sub-oids and sub-tupes for super extensions
Each function is given the same name as the extension module and simply returns the body of the
`encode/1` callback. Several helper functions that help encode parameters, lists, tuples and
values are also exposed and can be called either locally or remotely. These helper functions
eventually call the encoding functions specific to each extension.
The third point is handled by creating decoding functions for each extension. These functions
are created by iterating over each extension, taking the pattern returned by the `decode/1`
callback and creating a function that either:
- Accepts the pattern for regular extensions
- Accepts the pattern, sub-oids and sub-tupes for super extensions
Each function also accepts several variables that help accumulate the results and a special
variable `mod` that allows the extensions to access the type modifier of the column. This
allows them, for example, to return the correct precision for timestamps.
Each function is given the same name as the extension module and adds the body of the `decode/1`
callback to the accumulated results. Helper functions that help decode lists and tuples are also
exposed. These helper functions eventually call the decoding functions specific to each extension.
"""
def define(module, extensions, opts) do
opts =
opts
|> Keyword.put_new(:decode_binary, :copy)
config = configure(extensions, opts)
define_inline(module, config, opts)
end
## Helpers
defp directives(config, opts) do
requires =
for {extension, _} <- config do
quote do: require(unquote(extension))
end
preludes =
for {extension, {state, _, _}} <- config,
function_exported?(extension, :prelude, 1),
do: extension.prelude(state)
null = Keyword.get(opts, :null)
moduledoc = Keyword.get(opts, :moduledoc, false)
quote do
@moduledoc unquote(moduledoc)
import Postgrex.BinaryUtils
require unquote(__MODULE__)
unquote(requires)
unquote(preludes)
unquote(bin_opt_info(opts))
@compile {:inline, [encode_value: 2]}
@dialyzer {:no_opaque, [decode_tuple: 5]}
@null unquote(Macro.escape(null))
end
end
defp bin_opt_info(opts) do
if Keyword.get(opts, :bin_opt_info) do
quote do: @compile(:bin_opt_info)
else
[]
end
end
@anno [generated: true]
defp find(config) do
clauses = Enum.flat_map(config, &find_clauses/1)
clauses = clauses ++ quote do: (_ -> nil)
quote @anno do
@doc false
def find(type_info, formats) do
case {type_info, formats} do
unquote(clauses)
end
end
end
end
defp find_clauses({extension, {opts, matching, format}}) do
for {key, value} <- matching do
[clause] = find_clause(extension, opts, key, value, format)
clause
end
end
defp find_clause(extension, opts, key, value, :super_binary) do
quote do
{%{unquote(key) => unquote(value)} = type_info, formats}
when formats in [:any, :binary] ->
oids = unquote(extension).oids(type_info, unquote(opts))
{:super_binary, unquote(extension), oids}
end
end
defp find_clause(extension, _opts, key, value, format) do
quote do
{%{unquote(key) => unquote(value)}, formats}
when formats in [:any, unquote(format)] ->
{unquote(format), unquote(extension)}
end
end
defp maybe_rewrite(ast, extension, cases, opts) do
if Postgrex.Utils.default_extension?(extension) and
not Keyword.get(opts, :debug_defaults, false) do
ast
else
rewrite(ast, cases)
end
end
defp rewrite(ast, [{:->, clause_meta, _} | _original]) do
Macro.prewalk(ast, fn
{kind, meta, [{fun, _, args}, block]} when kind in [:def, :defp] and is_list(args) ->
{kind, meta, [{fun, clause_meta, args}, block]}
other ->
other
end)
end
defp encode(config, define_opts) do
encodes =
for {extension, {opts, [_ | _], format}} <- config do
encode = extension.encode(opts)
clauses =
for clause <- encode do
encode_type(extension, format, clause)
end
clauses = [encode_null(extension, format) | clauses]
quote do
unquote(encode_value(extension, format))
unquote(encode_inline(extension, format))
unquote(clauses |> maybe_rewrite(extension, encode, define_opts))
end
end
quote location: :keep do
unquote(encodes)
@doc false
def encode_params(params, types) do
encode_params(params, types, [])
end
defp encode_params([param | params], [type | types], encoded) do
encode_params(params, types, [encode_value(param, type) | encoded])
end
defp encode_params([], [], encoded), do: Enum.reverse(encoded)
defp encode_params(params, _, _) when is_list(params), do: :error
@doc false
def encode_tuple(tuple, nil, _types) do
raise DBConnection.EncodeError, """
cannot encode anonymous tuple #{inspect(tuple)}. \
Please define a custom Postgrex extension that matches on its underlying type:
use Postgrex.BinaryExtension, type: "typeinthedb"
"""
end
def encode_tuple(tuple, oids, types) do
encode_tuple(tuple, 1, oids, types, [])
end
defp encode_tuple(tuple, n, [oid | oids], [type | types], acc) do
param = :erlang.element(n, tuple)
acc = [acc, <<oid::uint32()>> | encode_value(param, type)]
encode_tuple(tuple, n + 1, oids, types, acc)
end
defp encode_tuple(tuple, n, [], [], acc) when tuple_size(tuple) < n do
acc
end
defp encode_tuple(tuple, n, [], [], _) when is_tuple(tuple) do
raise DBConnection.EncodeError,
"expected a tuple of size #{n - 1}, got: #{inspect(tuple)}"
end
@doc false
def encode_list(list, type) do
encode_list(list, type, [])
end
defp encode_list([value | rest], type, acc) do
encode_list(rest, type, [acc | encode_value(value, type)])
end
defp encode_list([], _, acc) do
acc
end
end
end
defp encode_type(extension, :super_binary, clause) do
encode_super(extension, clause)
end
defp encode_type(extension, _, clause) do
encode_extension(extension, clause)
end
defp encode_extension(extension, clause) do
case split_extension(clause) do
{pattern, guard, body} ->
encode_extension(extension, pattern, guard, body)
{pattern, body} ->
encode_extension(extension, pattern, body)
end
end
defp encode_extension(extension, pattern, guard, body) do
quote do
defp unquote(extension)(unquote(pattern)) when unquote(guard) do
unquote(body)
end
end
end
defp encode_extension(extension, pattern, body) do
quote do
defp unquote(extension)(unquote(pattern)) do
unquote(body)
end
end
end
defp encode_super(extension, clause) do
case split_super(clause) do
{pattern, sub_oids, sub_types, guard, body} ->
encode_super(extension, pattern, sub_oids, sub_types, guard, body)
{pattern, sub_oids, sub_types, body} ->
encode_super(extension, pattern, sub_oids, sub_types, body)
end
end
defp encode_super(extension, pattern, sub_oids, sub_types, guard, body) do
quote do
defp unquote(extension)(unquote(pattern), unquote(sub_oids), unquote(sub_types))
when unquote(guard) do
unquote(body)
end
end
end
defp encode_super(extension, pattern, sub_oids, sub_types, body) do
quote do
defp unquote(extension)(unquote(pattern), unquote(sub_oids), unquote(sub_types)) do
unquote(body)
end
end
end
defp encode_inline(extension, :super_binary) do
quote do
@compile {:inline, [{unquote(extension), 3}]}
end
end
defp encode_inline(extension, _) do
quote do
@compile {:inline, [{unquote(extension), 1}]}
end
end
defp encode_null(extension, :super_binary) do
quote do
defp unquote(extension)(@null, _sub_oids, _sub_types), do: <<-1::int32()>>
end
end
defp encode_null(extension, _) do
quote do
defp unquote(extension)(@null), do: <<-1::int32()>>
end
end
defp encode_value(extension, :super_binary) do
quote do
@doc false
def encode_value(value, {unquote(extension), sub_oids, sub_types}) do
unquote(extension)(value, sub_oids, sub_types)
end
end
end
defp encode_value(extension, _) do
quote do
@doc false
def encode_value(value, unquote(extension)) do
unquote(extension)(value)
end
end
end
defp decode(config, define_opts) do
rest = quote do: rest
acc = quote do: acc
rem = quote do: rem
full = quote do: full
rows = quote do: rows
row_dispatch =
for {extension, {_, [_ | _], format}} <- config do
decode_row_dispatch(extension, format, rest, acc, rem, full, rows)
end
next_dispatch = decode_rows_dispatch(rest, acc, rem, full, rows)
row_dispatch = row_dispatch ++ next_dispatch
decodes =
for {extension, {opts, [_ | _], format}} <- config do
decode = extension.decode(opts)
clauses =
for clause <- decode do
decode_type(extension, format, clause, row_dispatch, rest, acc, rem, full, rows)
end
null_clauses = decode_null(extension, format, row_dispatch, rest, acc, rem, full, rows)
quote location: :keep do
unquote(clauses |> maybe_rewrite(extension, decode, define_opts))
unquote(null_clauses)
end
end
quote location: :keep do
unquote(decode_rows(row_dispatch, rest, acc, rem, full, rows))
unquote(decode_simple())
unquote(decode_list(config))
unquote(decode_tuple(config))
unquote(decodes)
end
end
defp decode_rows(dispatch, rest, acc, rem, full, rows) do
quote location: :keep, generated: true do
@doc false
def decode_rows(binary, types, rows) do
decode_rows(binary, byte_size(binary), types, rows)
end
defp decode_rows(
<<?D, size::int32(), _::int16(), unquote(rest)::binary>>,
rem,
unquote(full),
unquote(rows)
)
when rem > size do
unquote(rem) = rem - (1 + size)
unquote(acc) = []
case unquote(full) do
unquote(dispatch)
end
end
defp decode_rows(<<?D, size::int32(), rest::binary>>, rem, _, rows) do
more = size + 1 - rem
{:more, [?D, <<size::int32()>> | rest], rows, more}
end
defp decode_rows(<<?D, rest::binary>>, _, _, rows) do
{:more, [?D | rest], rows, 0}
end
defp decode_rows(<<rest::binary-size(0)>>, _, _, rows) do
{:more, [], rows, 0}
end
defp decode_rows(<<rest::binary>>, _, _, rows) do
{:ok, rows, rest}
end
end
end
defp decode_row_dispatch(extension, :super_binary, rest, acc, rem, full, rows) do
[clause] =
quote do
[{unquote(extension), sub_oids, sub_types, mod} | types] ->
unquote(extension)(
unquote(rest),
sub_oids,
sub_types,
mod,
types,
unquote(acc),
unquote(rem),
unquote(full),
unquote(rows)
)
end
clause
end
defp decode_row_dispatch(extension, _, rest, acc, rem, full, rows) do
[clause] =
quote do
[{unquote(extension), mod} | types2] ->
unquote(extension)(
unquote(rest),
mod,
types2,
unquote(acc),
unquote(rem),
unquote(full),
unquote(rows)
)
end
clause
end
defp decode_rows_dispatch(rest, acc, rem, full, rows) do
quote do
[] ->
rows = [Enum.reverse(unquote(acc)) | unquote(rows)]
decode_rows(unquote(rest), unquote(rem), unquote(full), rows)
end
end
defp decode_simple() do
quote do
@doc false
def decode_simple(<<>>),
do: []
def decode_simple(<<-1::int32(), rest::binary>>),
do: [@null | decode_simple(rest)]
def decode_simple(<<len::int32(), value::binary-size(len), rest::binary>>),
do: [:binary.copy(value) | decode_simple(rest)]
end
end
defp decode_list(config) do
rest = quote do: rest
dispatch =
for {extension, {_, [_ | _], format}} <- config do
decode_list_dispatch(extension, format, rest)
end
quote do
@doc false
def decode_list(<<unquote(rest)::binary>>, type) do
case type do
unquote(dispatch)
end
end
end
end
defp decode_list_dispatch(extension, :super_binary, rest) do
[clause] =
quote do
{unquote(extension), sub_oids, sub_types, mod} ->
unquote(extension)(unquote(rest), sub_oids, sub_types, mod, [])
end
clause
end
defp decode_list_dispatch(extension, _, rest) do
[clause] =
quote do
{unquote(extension), mod} ->
unquote(extension)(unquote(rest), mod, [])
end
clause
end
defp decode_tuple(config) do
rest = quote do: rest
oids = quote do: oids
n = quote do: n
acc = quote do: acc
dispatch =
for {extension, {_, [_ | _], format}} <- config do
decode_tuple_dispatch(extension, format, rest, oids, n, acc)
end
quote generated: true do
@doc false
def decode_tuple(<<rest::binary>>, count, types) when is_integer(count) do
decode_tuple(rest, count, types, 0, [])
end
def decode_tuple(<<rest::binary>>, oids, types) do
decode_tuple(rest, oids, types, 0, [])
end
defp decode_tuple(
<<oid::int32(), unquote(rest)::binary>>,
[oid | unquote(oids)],
types,
unquote(n),
unquote(acc)
) do
case types do
unquote(dispatch)
end
end
defp decode_tuple(<<>>, [], [], n, acc) do
:erlang.make_tuple(n, @null, acc)
end
defp decode_tuple(
<<oid::int32(), unquote(rest)::binary>>,
rem,
types,
unquote(n),
unquote(acc)
)
when rem > 0 do
case Postgrex.Types.fetch(oid, types) do
{:ok, {:binary, type}} ->
unquote(oids) = rem - 1
case [type | types] do
unquote(dispatch)
end
{:ok, {:text, _}} ->
msg =
"oid `#{oid}` was bootstrapped in text format and can not " <>
"be decoded inside an anonymous record"
raise RuntimeError, msg
{:error, %TypeInfo{type: pg_type}, _mod} ->
msg = "type `#{pg_type}` can not be handled by the configured extensions"
raise RuntimeError, msg
{:error, nil, _mod} ->
msg = "oid `#{oid}` was not bootstrapped and lacks type information"
raise RuntimeError, msg
end
end
defp decode_tuple(<<>>, 0, _types, n, acc) do
:erlang.make_tuple(n, @null, acc)
end
end
end
defp decode_tuple_dispatch(extension, :super_binary, rest, oids, n, acc) do
[clause] =
quote do
[{unquote(extension), sub_oids, sub_types} | types] ->
unquote(extension)(
unquote(rest),
sub_oids,
sub_types,
nil,
unquote(oids),
types,
unquote(n) + 1,
unquote(acc)
)
end
clause
end
defp decode_tuple_dispatch(extension, _, rest, oids, n, acc) do
[clause] =
quote do
[unquote(extension) | types] ->
unquote(extension)(
unquote(rest),
nil,
unquote(oids),
types,
unquote(n) + 1,
unquote(acc)
)
end
clause
end
defp decode_type(extension, :super_binary, clause, dispatch, rest, acc, rem, full, rows) do
decode_super(extension, clause, dispatch, rest, acc, rem, full, rows)
end
defp decode_type(extension, _, clause, dispatch, rest, acc, rem, full, rows) do
decode_extension(extension, clause, dispatch, rest, acc, rem, full, rows)
end
defp decode_null(extension, :super_binary, dispatch, rest, acc, rem, full, rows) do
decode_super_null(extension, dispatch, rest, acc, rem, full, rows)
end
defp decode_null(extension, _, dispatch, rest, acc, rem, full, rows) do
decode_extension_null(extension, dispatch, rest, acc, rem, full, rows)
end
defp decode_extension(extension, clause, dispatch, rest, acc, rem, full, rows) do
case split_extension(clause) do
{pattern, guard, body} ->
decode_extension(
extension,
pattern,
guard,
body,
dispatch,
rest,
acc,
rem,
full,
rows
)
{pattern, body} ->
decode_extension(extension, pattern, body, dispatch, rest, acc, rem, full, rows)
end
end
defp decode_extension(
extension,
pattern,
guard,
body,
dispatch,
rest,
acc,
rem,
full,
rows
) do
quote do
defp unquote(extension)(
<<unquote(pattern), unquote(rest)::binary>>,
var!(mod),
types,
acc,
unquote(rem),
unquote(full),
unquote(rows)
)
when unquote(guard) do
_ = var!(mod)
unquote(acc) = [unquote(body) | acc]
case types do
unquote(dispatch)
end
end
defp unquote(extension)(<<unquote(pattern), rest::binary>>, var!(mod), acc)
when unquote(guard) do
_ = var!(mod)
unquote(extension)(rest, var!(mod), [unquote(body) | acc])
end
defp unquote(extension)(
<<unquote(pattern), rest::binary>>,
var!(mod),
oids,
types,
n,
acc
)
when unquote(guard) do
_ = var!(mod)
decode_tuple(rest, oids, types, n, [{n, unquote(body)} | acc])
end
end
end
defp decode_extension(extension, pattern, body, dispatch, rest, acc, rem, full, rows) do
quote do
defp unquote(extension)(
<<unquote(pattern), unquote(rest)::binary>>,
var!(mod),
types,
acc,
unquote(rem),
unquote(full),
unquote(rows)
) do
_ = var!(mod)
unquote(acc) = [unquote(body) | acc]
case types do
unquote(dispatch)
end
end
defp unquote(extension)(<<unquote(pattern), rest::binary>>, var!(mod), acc) do
_ = var!(mod)
decoded = unquote(body)
unquote(extension)(rest, var!(mod), [decoded | acc])
end
defp unquote(extension)(
<<unquote(pattern), rest::binary>>,
var!(mod),
oids,
types,
n,
acc
) do
_ = var!(mod)
decode_tuple(rest, oids, types, n, [{n, unquote(body)} | acc])
end
end
end
defp decode_extension_null(extension, dispatch, rest, acc, rem, full, rows) do
quote do
defp unquote(extension)(
<<-1::int32(), unquote(rest)::binary>>,
_mod,
types,
acc,
unquote(rem),
unquote(full),
unquote(rows)
) do
unquote(acc) = [@null | acc]
case types do
unquote(dispatch)
end
end
defp unquote(extension)(<<-1::int32(), rest::binary>>, var!(mod), acc) do
unquote(extension)(rest, var!(mod), [@null | acc])
end
defp unquote(extension)(<<>>, _, acc) do
acc
end
defp unquote(extension)(<<-1::int32(), rest::binary>>, _mod, oids, types, n, acc) do
decode_tuple(rest, oids, types, n, acc)
end
end
end
defp split_extension({:->, _, [head, body]}) do
case head do
[{:when, _, [pattern, guard]}] ->
{pattern, guard, body}
[pattern] ->
{pattern, body}
end
end
defp decode_super(extension, clause, dispatch, rest, acc, rem, full, rows) do
case split_super(clause) do
{pattern, oids, types, guard, body} ->
decode_super(
extension,
pattern,
oids,
types,
guard,
body,
dispatch,
rest,
acc,
rem,
full,
rows
)
{pattern, oids, types, body} ->
decode_super(
extension,
pattern,
oids,
types,
body,
dispatch,
rest,
acc,
rem,
full,
rows
)
end
end
defp decode_super(
extension,
pattern,
sub_oids,
sub_types,
guard,
body,
dispatch,
rest,
acc,
rem,
full,
rows
) do
quote do
defp unquote(extension)(
<<unquote(pattern), unquote(rest)::binary>>,
unquote(sub_oids),
unquote(sub_types),
var!(mod),
types,
acc,
unquote(rem),
unquote(full),
unquote(rows)
)
when unquote(guard) do
_ = var!(mod)
unquote(acc) = [unquote(body) | acc]
case types do
unquote(dispatch)
end
end
defp unquote(extension)(
<<unquote(pattern), rest::binary>>,
unquote(sub_oids),
unquote(sub_types),
var!(mod),
acc
)
when unquote(guard) do
_ = var!(mod)
acc = [unquote(body) | acc]
unquote(extension)(rest, unquote(sub_oids), unquote(sub_types), var!(mod), acc)
end
defp unquote(extension)(
<<unquote(pattern), rest::binary>>,
unquote(sub_oids),
unquote(sub_types),
var!(mod),
oids,
types,
n,
acc
)
when unquote(guard) do
_ = var!(mod)
decode_tuple(rest, oids, types, n, [{n, unquote(body)} | acc])
end
end
end
defp decode_super(
extension,
pattern,
sub_oids,
sub_types,
body,
dispatch,
rest,
acc,
rem,
full,
rows
) do
quote do
defp unquote(extension)(
<<unquote(pattern), unquote(rest)::binary>>,
unquote(sub_oids),
unquote(sub_types),
var!(mod),
types,
acc,
unquote(rem),
unquote(full),
unquote(rows)
) do
_ = var!(mod)
unquote(acc) = [unquote(body) | acc]
case types do
unquote(dispatch)
end
end
defp unquote(extension)(
<<unquote(pattern), rest::binary>>,
unquote(sub_oids),
unquote(sub_types),
var!(mod),
acc
) do
_ = var!(mod)
acc = [unquote(body) | acc]
unquote(extension)(rest, unquote(sub_oids), unquote(sub_types), var!(mod), acc)
end
defp unquote(extension)(
<<unquote(pattern), rest::binary>>,
unquote(sub_oids),
unquote(sub_types),
var!(mod),
oids,
types,
n,
acc
) do
_ = var!(mod)
acc = [{n, unquote(body)} | acc]
decode_tuple(rest, oids, types, n, acc)
end
end
end
defp decode_super_null(extension, dispatch, rest, acc, rem, full, rows) do
quote do
defp unquote(extension)(
<<-1::int32(), unquote(rest)::binary>>,
_sub_oids,
_sub_types,
_mod,
types,
acc,
unquote(rem),
unquote(full),
unquote(rows)
) do
unquote(acc) = [@null | acc]
case types do
unquote(dispatch)
end
end
defp unquote(extension)(<<-1::int32(), rest::binary>>, sub_oids, sub_types, var!(mod), acc) do
unquote(extension)(rest, sub_oids, sub_types, var!(mod), [@null | acc])
end
defp unquote(extension)(<<>>, _sub_oid, _sub_types, _mod, acc) do
acc
end
defp unquote(extension)(
<<-1::int32(), rest::binary>>,
_sub_oids,
_sub_types,
_mod,
oids,
types,
n,
acc
) do
decode_tuple(rest, oids, types, n, acc)
end
end
end
defp split_super({:->, _, [head, body]}) do
case head do
[{:when, _, [pattern, sub_oids, sub_types, guard]}] ->
{pattern, sub_oids, sub_types, guard, body}
[pattern, sub_oids, sub_types] ->
{pattern, sub_oids, sub_types, body}
end
end
defp configure(extensions, opts) do
defaults = Postgrex.Utils.default_extensions(opts)
Enum.map(extensions ++ defaults, &configure/1)
end
defp configure({extension, opts}) do
state = extension.init(opts)
matching = extension.matching(state)
format = extension.format(state)
{extension, {state, matching, format}}
end
defp configure(extension) do
configure({extension, []})
end
defp define_inline(module, config, opts) do
quoted = [
directives(config, opts),
find(config),
encode(config, opts),
decode(config, opts)
]
Module.create(module, quoted, Macro.Env.location(__ENV__))
end
end