Packages

Generates typed Elixir query modules from plain SQL files using Postgres inference or static metadata.

Current section

Files

Jump to
squirr_elix lib squirrelix postgres.ex
Raw

lib/squirrelix/postgres.ex

defmodule Squirrelix.Postgres do
@moduledoc """
Postgrex-backed query inferrer for Squirrelix inference.
"""
alias Squirrelix.Column
alias Squirrelix.ConnectionOptions
alias Squirrelix.Error
alias Squirrelix.Error.MissingPostgresColumn
alias Squirrelix.Error.MissingPostgresTable
alias Squirrelix.Error.PostgresInferenceError
alias Squirrelix.Error.PostgresSyntaxError
alias Squirrelix.Error.QueryHasInvalidEnum
alias Squirrelix.Query
alias Squirrelix.SQL
alias Squirrelix.TypeMapper
require Logger
# $1 = relation name, $2 = column name, $3 = schema (NULL = search_path / temp).
@column_nullability_query """
select a.attnotnull
from pg_attribute a
join pg_class c on a.attrelid = c.oid
join pg_namespace n on c.relnamespace = n.oid
where c.relname = $1
and a.attname = $2
and a.attnum > 0
and not a.attisdropped
and (
case
when $3::text is null then n.nspname = any(current_schemas(true))
when $3::text = 'pg_temp' or $3::text like 'pg_temp_%' then
n.oid = pg_my_temp_schema()
else n.nspname = $3
end
)
limit 1
"""
@type_lookup_query """
with recursive types(oid, name, elem, kind, base, array_dimensions, jumps) as (
select
pg_type.oid as oid,
pg_type.typname as name,
pg_type.typelem as elem,
pg_type.typtype as kind,
pg_type.typbasetype as base,
0 as array_dimensions,
0 as jumps
from pg_type
where pg_type.oid = $1::oid
union all
select
pg_type.oid as oid,
pg_type.typname as name,
pg_type.typelem as elem,
pg_type.typtype as kind,
pg_type.typbasetype as base,
next_type.array_dimensions as array_dimensions,
types.jumps + 1 as jumps
from types
join lateral (
values
(case when types.elem != 0 and types.name not in ('name', 'point') then types.elem end, types.array_dimensions + 1),
(case when types.kind = 'd' then types.base end, types.array_dimensions)
) as next_type(oid, array_dimensions)
on next_type.oid is not null
join pg_type
on pg_type.oid = next_type.oid
)
select types.name, types.kind, types.array_dimensions, types.oid
from types
order by types.jumps desc
limit 1
"""
@enum_variants_query """
select enumlabel
from pg_enum
where enumtypid = $1::oid
order by enumsortorder asc
"""
@doc """
Opens a Postgrex connection after a synchronous probe so connection failures
and timeouts become structured Squirrelix errors.
"""
@spec connect(ConnectionOptions.t()) :: {:ok, pid()} | {:error, struct()}
def connect(%ConnectionOptions{} = connection_options) do
postgrex_opts = postgrex_opts(connection_options)
case probe_connection(postgrex_opts) do
:ok ->
case Postgrex.start_link(postgrex_opts) do
{:ok, conn} ->
{:ok, conn}
{:error, reason} ->
{:error, Error.connection_error(reason, connection_options)}
end
{:error, reason} ->
{:error, Error.connection_error(reason, connection_options)}
end
end
@spec inferrer(Postgrex.conn()) :: Squirrelix.Inference.inferrer()
def inferrer(conn) do
&infer(conn, &1)
end
@spec infer(Postgrex.conn(), Query.t()) :: {:ok, keyword()} | {:error, struct()}
def infer(conn, %Query{} = query) do
with {:ok, prepared_query} <- prepare(conn, query),
{:ok, params} <- describe_oids(conn, prepared_query.param_oids || [], query),
{:ok, returns} <- describe_returns(conn, prepared_query, query) do
{:ok, [params: params, returns: returns]}
end
end
@doc """
Converts Squirrelix connection options into a Postgrex `start_link/1` keyword list.
"""
@spec postgrex_opts(ConnectionOptions.t()) :: keyword()
def postgrex_opts(%ConnectionOptions{} = connection_options) do
[
hostname: connection_options.host,
port: connection_options.port,
username: connection_options.user,
password: connection_options.password,
database: connection_options.database,
timeout: connection_options.timeout_seconds * 1000,
connect_timeout: connection_options.timeout_seconds * 1000,
ssl: connection_options.ssl || false,
types: Postgrex.DefaultTypes
]
|> Enum.reject(fn {_key, value} -> is_nil(value) end)
end
defp probe_connection(opts) do
case Postgrex.Protocol.connect(opts) do
{:ok, state} ->
_ = Postgrex.Protocol.disconnect(:normal, state)
:ok
{:error, reason} ->
{:error, reason}
end
end
defp prepare(conn, query) do
case Postgrex.prepare(conn, "", query.content) do
{:ok, prepared_query} -> {:ok, prepared_query}
{:error, %Postgrex.Error{} = error} -> {:error, postgres_error(query, error)}
end
end
defp postgres_error(query, %Postgrex.Error{postgres: %{code: :syntax_error} = postgres}) do
%PostgresSyntaxError{
file: query.file,
starting_line: query.starting_line,
content: query.content,
message: postgres.message,
position: parse_position(postgres)
}
end
defp postgres_error(query, %Postgrex.Error{postgres: %{code: :undefined_table} = postgres}) do
%MissingPostgresTable{
file: query.file,
starting_line: query.starting_line,
content: query.content,
message: postgres.message,
table: quoted_identifier(postgres.message),
position: parse_position(postgres)
}
end
defp postgres_error(query, %Postgrex.Error{postgres: %{code: :undefined_column} = postgres}) do
%MissingPostgresColumn{
file: query.file,
starting_line: query.starting_line,
content: query.content,
message: postgres.message,
column: quoted_identifier(postgres.message),
position: parse_position(postgres)
}
end
defp postgres_error(query, %Postgrex.Error{postgres: postgres}) do
%PostgresInferenceError{
file: query.file,
starting_line: query.starting_line,
content: query.content,
message: Map.get(postgres, :message, "Postgres rejected query inference"),
code: Map.get(postgres, :code),
position: parse_position(postgres)
}
end
defp parse_position(%{position: position}) when is_binary(position) do
case Integer.parse(position) do
{value, ""} -> value
_invalid -> nil
end
end
defp parse_position(_postgres), do: nil
defp quoted_identifier(message) when is_binary(message) do
case Regex.run(~r/"([^"]+)"/, message) do
[_match, identifier] -> identifier
nil -> nil
end
end
defp quoted_identifier(_message), do: nil
defp describe_returns(conn, prepared_query, query) do
columns = prepared_query.columns || []
result_oids = prepared_query.result_oids || []
{plan_available?, plan_nullables, column_sources} =
infer_nullability(conn, query, prepared_query)
with {:ok, types} <- describe_oids(conn, result_oids, query) do
returns =
columns
|> Enum.zip(types)
|> Enum.with_index()
|> Enum.map(fn {{name, type}, index} ->
nullable? =
column_nullable?(
conn,
name,
index,
plan_nullables,
column_sources,
plan_available?
)
%Column{name: name, type: type, nullable?: nullable?}
end)
{:ok, returns}
end
end
# -- Plan-based nullability inference (mirrors upstream Gleam squirrel) --
defp infer_nullability(conn, query, _prepared_query) do
case query_plan(conn, query) do
{:ok, plan} ->
{true, nullables_from_plan(plan), column_sources_from_plan(plan)}
:error ->
{false, MapSet.new(), []}
end
end
defp query_plan(conn, query) do
content = explainable_query_content(query.content)
if SQL.single_statement?(content) do
# Simple protocol is required so `$n` placeholders work with
# `generic_plan` without binding params. Multi-statement SQL is rejected
# above; prepare/2 already validated the query as a single statement.
explain_query = "explain (format json, verbose, generic_plan) " <> content
try do
with {:ok, %Postgrex.Result{rows: [[plan_json]]}} <-
Postgrex.query(conn, explain_query, [], query_type: :text),
{:ok, root_plan} <- decode_plan_json(plan_json) do
{:ok, parse_plan(root_plan)}
else
_ ->
warn_explain_unavailable(query.file)
:error
end
rescue
_ ->
warn_explain_unavailable(query.file)
:error
end
else
Logger.warning(
"Squirrelix: skipping EXPLAIN nullability for #{query.file} because the SQL is not a single statement"
)
:error
end
end
defp warn_explain_unavailable(file) do
Logger.warning(
"Squirrelix: EXPLAIN nullability unavailable for #{file}; treating unknown columns as nullable"
)
end
defp explainable_query_content(content) do
content
|> String.trim_leading()
|> String.replace(~r/^;+\s*/, "")
end
defp decode_plan_json(plan_json) when is_binary(plan_json) do
with {:ok, data} <- JSON.decode(plan_json) do
decode_plan_json(data)
end
end
defp decode_plan_json([%{"Plan" => root_plan} | _]), do: {:ok, root_plan}
defp decode_plan_json(%{"Plan" => root_plan}), do: {:ok, root_plan}
defp decode_plan_json(_), do: :error
defp parse_plan(plan_map) do
%{
join_type: Map.get(plan_map, "Join Type"),
output: Map.get(plan_map, "Output", []),
relation: Map.get(plan_map, "Relation Name"),
schema: Map.get(plan_map, "Schema"),
plans: plan_map |> Map.get("Plans", []) |> Enum.map(&parse_plan/1)
}
end
defp nullables_from_plan(plan) do
outputs =
plan.output
|> Enum.with_index()
|> Map.new(fn {expr, idx} -> {expr, idx} end)
do_nullables_from_plan(plan, outputs, MapSet.new())
end
defp do_nullables_from_plan(plan, query_outputs, nullables) do
case {plan.join_type, plan.plans} do
{"Full", _} ->
plan_outputs_indices(plan, query_outputs)
|> MapSet.union(nullables)
{"Right", [left, right]} ->
nullables =
plan_outputs_indices(left, query_outputs)
|> MapSet.union(nullables)
do_nullables_from_plan(right, query_outputs, nullables)
{"Left", [left, right]} ->
nullables =
plan_outputs_indices(right, query_outputs)
|> MapSet.union(nullables)
do_nullables_from_plan(left, query_outputs, nullables)
{"Semi", [left, right]} ->
nullables =
plan_outputs_indices(right, query_outputs)
|> MapSet.union(nullables)
do_nullables_from_plan(left, query_outputs, nullables)
{"Inner", plans} ->
Enum.reduce(plans, nullables, fn child, acc ->
do_nullables_from_plan(child, query_outputs, acc)
end)
{_, plans} ->
Enum.reduce(plans, nullables, fn child, acc ->
do_nullables_from_plan(child, query_outputs, acc)
end)
end
end
defp plan_outputs_indices(plan, query_outputs) do
Enum.reduce(plan.output, MapSet.new(), fn output, acc ->
case Map.fetch(query_outputs, output) do
{:ok, idx} -> MapSet.put(acc, idx)
:error -> acc
end
end)
end
defp column_sources_from_plan(plan) do
expr_to_source = collect_expr_sources(plan)
Enum.map(plan.output, fn expr ->
Map.get(expr_to_source, expr) || classify_output_expr(expr, nil, nil)
end)
end
defp collect_expr_sources(%{relation: relation, schema: schema, output: output, plans: plans})
when is_binary(relation) and relation != "" do
own =
output
|> Enum.map(fn expr -> {expr, classify_output_expr(expr, schema, relation)} end)
|> Map.new()
Enum.reduce(plans, own, fn child, acc ->
Map.merge(acc, collect_expr_sources(child))
end)
end
defp collect_expr_sources(%{plans: plans}) do
Enum.reduce(plans, %{}, fn child, acc ->
Map.merge(acc, collect_expr_sources(child))
end)
end
# Classify EXPLAIN "Output" entries into table columns, scalar subplans, or
# expression-derived values. Matches Gleam squirrel's table_oid=0 behaviour
# for expressions (non-nullable) while treating SubPlan outputs as nullable.
defp classify_output_expr(expr, schema, relation) when is_binary(expr) do
if Regex.match?(~r/^\(SubPlan \d+\)$/, expr) do
:subquery
else
table_column_source(expr, schema, relation) || :expression
end
end
defp classify_output_expr(_expr, _schema, _relation), do: :expression
defp table_column_source(expr, schema, relation) do
with {:ok, column} <- parse_column_ref(expr),
table when is_binary(table) <- relation || table_from_qualified_expr(expr) do
{:table_column, schema, table, column}
else
_ -> nil
end
end
defp parse_column_ref(expr) do
# Optional alias/table qualifier, then a single identifier (quoted or bare).
# Anything else (operators, function calls, casts) is an expression.
case Regex.run(
~r/^((?:[A-Za-z_][A-Za-z0-9_]*|"[^"]+")\.)?([A-Za-z_][A-Za-z0-9_]*|"[^"]+")$/,
expr
) do
[_match, _qualifier, column] -> {:ok, unquote_ident(column)}
nil -> :error
end
end
defp table_from_qualified_expr(expr) do
case String.split(expr, ".", parts: 2) do
[table, _column] -> unquote_ident(table)
_ -> nil
end
end
defp unquote_ident(<<"\"", rest::binary>>) do
String.trim_trailing(rest, "\"")
end
defp unquote_ident(ident), do: ident
defp column_nullable?(conn, name, index, plan_nullables, column_sources, plan_available?) do
cond do
String.ends_with?(name, "!") ->
false
String.ends_with?(name, "?") ->
true
MapSet.member?(plan_nullables, index) ->
true
true ->
source_nullable?(conn, Enum.at(column_sources, index), plan_available?)
end
end
defp source_nullable?(_conn, :subquery, _plan_available?), do: true
defp source_nullable?(_conn, :expression, _plan_available?), do: false
defp source_nullable?(conn, {:table_column, schema, table, column}, _plan_available?)
when is_binary(table) and is_binary(column) do
!column_has_not_null_constraint?(conn, schema, table, column)
end
defp source_nullable?(_conn, _other, plan_available?), do: !plan_available?
defp column_has_not_null_constraint?(conn, schema, table, column) do
case Postgrex.query(conn, @column_nullability_query, [table, column, schema]) do
{:ok, %Postgrex.Result{rows: [[true]]}} -> true
{:ok, %Postgrex.Result{rows: [[false]]}} -> false
_ -> false
end
end
defp describe_oids(conn, oids, query) do
Enum.reduce_while(oids, {:ok, []}, fn oid, {:ok, types} ->
case describe_oid(conn, oid, query) do
{:ok, type} -> {:cont, {:ok, [type | types]}}
{:error, error} -> {:halt, {:error, error}}
end
end)
|> case do
{:ok, types} -> {:ok, Enum.reverse(types)}
{:error, error} -> {:error, error}
end
end
defp describe_oid(conn, oid, query) do
with {:ok, %Postgrex.Result{rows: [[name, kind, array_dimensions, type_oid]]}} <-
Postgrex.query(conn, @type_lookup_query, [oid]) do
resolve_postgres_type(conn, type_oid, name, kind, array_dimensions, query)
end
end
defp resolve_postgres_type(conn, oid, name, "e", array_dimensions, query) do
with {:ok, variants} <- enum_variants(conn, oid),
:ok <- TypeMapper.validate_enum(name, variants),
{:ok, type} <-
TypeMapper.from_postgres(name, kind: "e", array_dimensions: array_dimensions) do
{:ok, type}
else
{:error, :no_variants} ->
{:error, invalid_enum_error(query, name, :no_variants)}
{:error, error} ->
{:error, error}
end
end
defp resolve_postgres_type(_conn, _oid, name, kind, array_dimensions, _query) do
TypeMapper.from_postgres(name, kind: kind, array_dimensions: array_dimensions)
end
defp enum_variants(conn, oid) do
case Postgrex.query(conn, @enum_variants_query, [oid]) do
{:ok, %Postgrex.Result{rows: rows}} ->
{:ok, Enum.map(rows, fn [variant] -> variant end)}
{:error, error} ->
{:error, error}
end
end
defp invalid_enum_error(query, enum_name, reason) do
%QueryHasInvalidEnum{
file: query.file,
starting_line: query.starting_line,
content: query.content,
enum_name: enum_name,
reason: reason
}
end
end