From a809c579200de4cb628fefbce565d98417425afb Mon Sep 17 00:00:00 2001 From: Giacomo Cavalieri Date: Mon, 18 Nov 2024 21:36:40 +0100 Subject: [PATCH] :bug: Fix nullability inference from query plan --- .../left_join_nullability_inference.accepted | 47 ++++++ src/squirrel/internal/database/postgres.gleam | 140 +++++++++--------- src/squirrel/internal/error.gleam | 17 +++ test/squirrel_test.gleam | 75 +++++++--- 4 files changed, 190 insertions(+), 89 deletions(-) create mode 100644 birdie_snapshots/left_join_nullability_inference.accepted diff --git a/birdie_snapshots/left_join_nullability_inference.accepted b/birdie_snapshots/left_join_nullability_inference.accepted new file mode 100644 index 0000000..5bdaf7c --- /dev/null +++ b/birdie_snapshots/left_join_nullability_inference.accepted @@ -0,0 +1,47 @@ +--- +version: 1.2.3 +title: left join nullability inference +file: ./test/squirrel_test.gleam +test_name: left_join_nullability_inference_test +--- +import decode/zero +import gleam/option.{type Option} +import pog + +/// A row you get from running the `query` query +/// defined in `query.sql`. +/// +/// > 🐿️ This type definition was generated automatically using v-test of the +/// > [squirrel package](https://github.com/giacomocavalieri/squirrel). +/// +pub type QueryRow { + QueryRow(user_id: Int, roles: Option(String)) +} + +/// Runs the `query` query +/// defined in `query.sql`. +/// +/// > 🐿️ This function was generated automatically using v-test of +/// > the [squirrel package](https://github.com/giacomocavalieri/squirrel). +/// +pub fn query(db) { + let decoder = { + use user_id <- zero.field(0, zero.int) + use roles <- zero.field(1, zero.optional(zero.string)) + zero.success(QueryRow(user_id:, roles:)) + } + + let query = " +select + users_issue41.user_id, + profile_issue41.roles +from + users_issue41 + left join profile_issue41 + on profile_issue41.user_id = users_issue41.user_id; +" + + pog.query(query) + |> pog.returning(zero.run(_, decoder)) + |> pog.execute(db) +} diff --git a/src/squirrel/internal/database/postgres.gleam b/src/squirrel/internal/database/postgres.gleam index 7bbd178..2bdb610 100644 --- a/src/squirrel/internal/database/postgres.gleam +++ b/src/squirrel/internal/database/postgres.gleam @@ -13,11 +13,11 @@ //// > that is a bug! Please do reach out, I'd love to hear your feedback. //// +import decode/zero import eval import gleam/bit_array import gleam/bool import gleam/dict.{type Dict} -import gleam/dynamic.{type DecodeErrors, type Dynamic} as d import gleam/int import gleam/json import gleam/list @@ -210,24 +210,14 @@ type Nullability { /// A query plan produced by Postgres when we ask it to `explain` a query. /// type Plan { - Plan( - join_type: Option(JoinType), - parent_relation: Option(ParentRelation), - output: Option(List(String)), - plans: Option(List(Plan)), - ) + Plan(join_type: Option(JoinType), output: List(String), plans: List(Plan)) } type JoinType { - Full - Left - Right - Other -} - -type ParentRelation { - Inner - NotInner + FullJoin + LeftJoin + RightJoin + InnerJoin } /// This is the type of a database-related action. @@ -741,19 +731,19 @@ fn query_plan(query: UntypedQuery, parameters: Int) -> Db(Plan) { // We know the output will only contain a single row that is the json string // containing the query plan. let assert [[plan]] = res - let assert Ok([plan, ..]) = json.decode_bits(plan, json_plans_decoder) - eval.return(plan) + case json.decode_bits(plan, zero.run(_, json_plans_decoder())) { + Ok([plan, ..]) -> eval.return(plan) + Ok([]) -> panic as "unreachable: no query plan" + Error(reason) -> + eval.throw(error.CannotParsePlanForQuery(file: query.file, reason:)) + } } /// Given a query plan, returns a set with the indices of the output columns /// that can contain null values. /// fn nullables_from_plan(plan: Plan) -> Set(Int) { - let outputs = case plan.output { - Some(outputs) -> list.index_fold(outputs, dict.new(), dict.insert) - None -> dict.new() - } - + let outputs = list.index_fold(plan.output, dict.new(), dict.insert) do_nullables_from_plan(plan, outputs, set.new()) } @@ -763,27 +753,57 @@ fn do_nullables_from_plan( query_outputs: Dict(String, Int), nullables: Set(Int), ) -> Set(Int) { - let nullables = case plan.output, plan.join_type, plan.parent_relation { - // - All the outputs of a full join must be marked as nullable - // - All the outputs of an inner half join must be marked as nullable - Some(outputs), Some(Full), _ | Some(outputs), _, Some(Inner) -> { - use nullables, output <- list.fold(outputs, from: nullables) - case dict.get(query_outputs, output) { - Ok(i) -> set.insert(nullables, i) - Error(_) -> nullables - } + case plan.join_type, plan.plans { + // If this is a full join then all its outputs could be optional!! + Some(FullJoin), _ -> + plan_outputs_indices(plan, query_outputs) + |> set.union(nullables) + + // If this is a right join then we must mark the outputs of its left part as + // nullable! + Some(RightJoin), [left, right] -> { + let nullables = + plan_outputs_indices(left, query_outputs) + |> set.union(nullables) + + do_nullables_from_plan(right, query_outputs, nullables) + } + + // If this is a left join then we must mark the outputs of its right part as + // nullable! + Some(LeftJoin), [left, right] -> { + let nullables = + plan_outputs_indices(right, query_outputs) + |> set.union(nullables) + + do_nullables_from_plan(left, query_outputs, nullables) + } + + // This should never happen in theory (a join with 0, 1, or more than two + // childs), so we just inspect their plans as a safe bet. + Some(RightJoin), plans | Some(LeftJoin), plans | None, plans -> { + use nullables, plan <- list.fold(plans, nullables) + do_nullables_from_plan(plan, query_outputs, nullables) } - _, _, _ -> nullables - } - case plan.plans, plan.join_type { - // If this is an inner half join we keep inspecting the children to mark - // their outputs as nullable. - Some(plans), Some(Left) | Some(plans), Some(Right) -> { - use nullables, plan <- list.fold(plans, from: nullables) + // If this is an inner join then it's outputs are not necessarily nullable, + // we inspect the children's plans to see if they do have some nullable + // columns. + Some(InnerJoin), plans -> { + use nullables, plan <- list.fold(plans, nullables) do_nullables_from_plan(plan, query_outputs, nullables) } - _, _ -> nullables + } +} + +fn plan_outputs_indices( + plan: Plan, + query_outputs: Dict(String, Int), +) -> Set(Int) { + use nullables, output <- list.fold(plan.output, from: set.new()) + case dict.get(query_outputs, output) { + Ok(i) -> set.insert(nullables, i) + Error(_) -> nullables } } @@ -1263,37 +1283,25 @@ fn adjust_parse_error_for_explain(error: Error) -> Error { // --- DECODERS ---------------------------------------------------------------- -fn json_plans_decoder(data: Dynamic) -> Result(List(Plan), DecodeErrors) { - d.list(d.field("Plan", plan_decoder))(data) -} - -fn plan_decoder(data: Dynamic) -> Result(Plan, DecodeErrors) { - d.decode4( - Plan, - d.optional_field("Join Type", join_type_decoder), - d.optional_field("Parent Relationship", parent_relation_decoder), - d.optional_field("Output", d.list(d.string)), - d.optional_field("Plans", d.list(plan_decoder)), - )(data) +fn json_plans_decoder() { + zero.list(zero.at(["Plan"], plan_decoder())) } -fn join_type_decoder(data: Dynamic) -> Result(JoinType, DecodeErrors) { - use data <- result.map(d.string(data)) - case data { - "Full" -> Full - "Left" -> Left - "Right" -> Right - _ -> Other - } +fn plan_decoder() { + use join_type <- zero.optional_field("Join Type", None, join_type_decoder()) + use output <- zero.optional_field("Output", [], zero.list(zero.string)) + use plans <- zero.optional_field("Plans", [], zero.list(plan_decoder())) + zero.success(Plan(join_type:, output:, plans:)) } -fn parent_relation_decoder( - data: Dynamic, -) -> Result(ParentRelation, DecodeErrors) { - use data <- result.map(d.string(data)) +fn join_type_decoder() { + use data <- zero.then(zero.string) case data { - "Inner" -> Inner - _ -> NotInner + "Full" -> zero.success(Some(FullJoin)) + "Left" -> zero.success(Some(LeftJoin)) + "Right" -> zero.success(Some(RightJoin)) + "Inner" -> zero.success(Some(InnerJoin)) + _ -> zero.failure(None, "a join type") } } diff --git a/src/squirrel/internal/error.gleam b/src/squirrel/internal/error.gleam index 3ec5abc..579ec8e 100644 --- a/src/squirrel/internal/error.gleam +++ b/src/squirrel/internal/error.gleam @@ -1,5 +1,6 @@ import glam/doc.{type Document} import gleam/int +import gleam/json import gleam/list import gleam/option.{type Option, None, Some} import gleam/regex @@ -166,6 +167,14 @@ pub type Error { starting_line: Int, names: List(String), ) + + /// If the postgres server sends in a query explanantion in a format that I + /// cannot parse. + /// + /// This should never happen, and if it does it means I grossly forgot about + /// a possible value I shouldn't so I have to ask to open an issue. + /// + CannotParsePlanForQuery(file: String, reason: json.DecodeError) } pub type ValueIdentifierError { @@ -533,6 +542,14 @@ Gleam type!", <> ".", ) } + + CannotParsePlanForQuery(file:, reason:) -> + printable_error("Cannot decode query plan") + |> add_paragraph( + "I ran into an unexpected error while trying to figure out how to +generate code for query " <> style_file(file) <> ".", + ) + |> report_bug(string.inspect(reason)) } printable_error_to_doc(printable_error) diff --git a/test/squirrel_test.gleam b/test/squirrel_test.gleam index 0b5db26..d5aba79 100644 --- a/test/squirrel_test.gleam +++ b/test/squirrel_test.gleam @@ -42,8 +42,7 @@ fn setup_database() { create table if not exists squirrel( name text primary key, acorns int -); -" +)" |> pog.query |> pog.execute(db) @@ -53,8 +52,7 @@ create table if not exists jsons( id bigserial primary key, json json, jsonb jsonb -) -" +)" |> pog.query |> pog.execute(db) @@ -64,41 +62,57 @@ do $$ begin if not exists (select * from pg_type where typname = 'squirrel_colour') then create type squirrel_colour as enum ('red', 'grey', 'light brown'); end if; -end $$; - " +end $$;" + |> pog.query + |> pog.execute(db) + + let assert Ok(_) = + " +do $$ begin + if not exists (select * from pg_type where typname = '1 invalid enum') then + create type \"1 invalid enum\" as enum ('value'); + end if; +end $$;" |> pog.query |> pog.execute(db) let assert Ok(_) = " - do $$ begin - if not exists (select * from pg_type where typname = '1 invalid enum') then - create type \"1 invalid enum\" as enum ('value'); - end if; - end $$; +do $$ begin + if not exists (select * from pg_type where typname = 'invalid_variant') then + create type invalid_variant as enum ('1 invalid value'); + end if; +end $$;" + |> pog.query + |> pog.execute(db) + + let assert Ok(_) = " +do $$ begin + if not exists (select * from pg_type where typname = 'no_variants') then + create type no_variants as enum (); + end if; +end $$;" |> pog.query |> pog.execute(db) + // https://github.com/giacomocavalieri/squirrel/issues/41 let assert Ok(_) = " - do $$ begin - if not exists (select * from pg_type where typname = 'invalid_variant') then - create type invalid_variant as enum ('1 invalid value'); - end if; - end $$; - " +create table if not exists users_issue41( + user_id bigserial primary key +) + " |> pog.query |> pog.execute(db) let assert Ok(_) = " - do $$ begin - if not exists (select * from pg_type where typname = 'no_variants') then - create type no_variants as enum (); - end if; - end $$; - " +create table if not exists profile_issue41( + profile_id bigserial primary key, + user_id bigserial not null, + roles text not null +);" |> pog.query |> pog.execute(db) @@ -629,3 +643,18 @@ pub fn a_query_failing_does_not_change_the_other_query_error_2_test() { title: "a query failing does not change the other query's error 2", ) } + +// https://github.com/giacomocavalieri/squirrel/issues/41 +pub fn left_join_nullability_inference_test() { + " +select + users_issue41.user_id, + profile_issue41.roles +from + users_issue41 + left join profile_issue41 + on profile_issue41.user_id = users_issue41.user_id; +" + |> should_codegen + |> birdie.snap(title: "left join nullability inference") +} -- 2.51.2