2
0

cte_test.go 3.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687
  1. package executor
  2. import "testing"
  3. // TestNonRecursiveCTE reproduces the shape of Vikunja's project-permission
  4. // query: a CTE with a column list over a UNION ALL subquery, joined back to a
  5. // real table and grouped.
  6. func TestNonRecursiveCTE(t *testing.T) {
  7. _, schema, table := newTestDB(t)
  8. e := newExec(schema, table)
  9. execMust(t, e, "CREATE TABLE projects (id INTEGER PRIMARY KEY, owner_id INTEGER)")
  10. execMust(t, e, "CREATE TABLE users_projects (project_id INTEGER, user_id INTEGER, permission INTEGER)")
  11. execMust(t, e, "CREATE TABLE team_projects (project_id INTEGER, team_id INTEGER, permission INTEGER)")
  12. execMust(t, e, "CREATE TABLE team_members (team_id INTEGER, user_id INTEGER)")
  13. execMust(t, e, "CREATE TABLE project_ancestors (ancestor_id INTEGER, project_id INTEGER, depth INTEGER)")
  14. execMust(t, e, "INSERT INTO projects VALUES (1, 3)")
  15. execMust(t, e, "INSERT INTO projects VALUES (2, 9)")
  16. execMust(t, e, "INSERT INTO projects VALUES (3, 9)")
  17. execMust(t, e, "INSERT INTO users_projects VALUES (2, 3, 1)")
  18. execMust(t, e, "INSERT INTO team_projects VALUES (3, 7, 2)")
  19. execMust(t, e, "INSERT INTO team_members VALUES (7, 3)")
  20. execMust(t, e, "INSERT INTO project_ancestors VALUES (1, 1, 0)")
  21. execMust(t, e, "INSERT INTO project_ancestors VALUES (2, 2, 0)")
  22. execMust(t, e, "INSERT INTO project_ancestors VALUES (3, 3, 0)")
  23. res := execMust(t, e, `WITH grants (project_id, permission) AS (
  24. SELECT project_id, MAX(permission) FROM (
  25. SELECT id AS project_id, 2 AS permission FROM projects WHERE owner_id = 3
  26. UNION ALL SELECT project_id, permission FROM users_projects WHERE user_id = 3
  27. UNION ALL SELECT tp.project_id, tp.permission
  28. FROM team_projects tp INNER JOIN team_members tm ON tm.team_id = tp.team_id
  29. WHERE tm.user_id = 3
  30. ) direct_grants GROUP BY project_id
  31. )
  32. SELECT pa.project_id AS id, MAX(g.permission) AS permission
  33. FROM grants g INNER JOIN project_ancestors pa ON pa.ancestor_id = g.project_id
  34. GROUP BY pa.project_id`)
  35. if res.RowCount != 3 {
  36. t.Fatalf("expected 3 rows, got %d: %v", res.RowCount, res.Rows)
  37. }
  38. got := map[int64]int64{}
  39. for _, row := range res.Rows {
  40. got[row[0].(int64)] = row[1].(int64)
  41. }
  42. want := map[int64]int64{1: 2, 2: 1, 3: 2}
  43. for id, perm := range want {
  44. if got[id] != perm {
  45. t.Fatalf("permission for project %d = %v, want %v (all: %v)", id, got[id], perm, got)
  46. }
  47. }
  48. }
  49. // TestCTEChained verifies a CTE that references an earlier CTE.
  50. func TestCTEChained(t *testing.T) {
  51. _, schema, table := newTestDB(t)
  52. e := newExec(schema, table)
  53. execMust(t, e, "CREATE TABLE t (id INTEGER PRIMARY KEY, v INTEGER)")
  54. execMust(t, e, "INSERT INTO t VALUES (1, 10)")
  55. execMust(t, e, "INSERT INTO t VALUES (2, 20)")
  56. res := execMust(t, e, `WITH base AS (SELECT id, v FROM t WHERE v > 5),
  57. doubled AS (SELECT id, v * 2 AS d FROM base)
  58. SELECT id, d FROM doubled ORDER BY id`)
  59. if res.RowCount != 2 {
  60. t.Fatalf("expected 2 rows, got %d: %v", res.RowCount, res.Rows)
  61. }
  62. if res.Rows[0][1] != int64(20) || res.Rows[1][1] != int64(40) {
  63. t.Fatalf("unexpected rows: %v", res.Rows)
  64. }
  65. }
  66. // TestRecursiveCTENonCompound verifies a plain CTE under WITH RECURSIVE is
  67. // treated as non-recursive.
  68. func TestRecursiveCTENonCompound(t *testing.T) {
  69. _, schema, table := newTestDB(t)
  70. e := newExec(schema, table)
  71. execMust(t, e, "CREATE TABLE t (id INTEGER PRIMARY KEY, v INTEGER)")
  72. execMust(t, e, "INSERT INTO t VALUES (1, 10)")
  73. execMust(t, e, "INSERT INTO t VALUES (2, 20)")
  74. res := execMust(t, e, "WITH RECURSIVE r AS (SELECT id, v FROM t) SELECT id, v FROM r ORDER BY id")
  75. if res.RowCount != 2 {
  76. t.Fatalf("expected 2 rows, got %d: %v", res.RowCount, res.Rows)
  77. }
  78. }