2
0

cte.go 9.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373
  1. package parser
  2. import (
  3. "fmt"
  4. "strings"
  5. "github.com/danfragoso/pizzasql-next/pkg/lexer"
  6. )
  7. // cteDef is a single common table expression parsed from a WITH clause.
  8. type cteDef struct {
  9. name string
  10. cols []string
  11. query *SelectStmt
  12. }
  13. // isWithStart reports whether the current token begins a WITH clause. WITH is
  14. // not a lexer keyword (it can be a column or table name), so it is recognized
  15. // by its literal at statement start.
  16. func (p *Parser) isWithStart() bool {
  17. return p.curTokenIs(lexer.TokenIdent) && strings.EqualFold(p.curToken.Literal, "WITH")
  18. }
  19. // parseWithStatement parses a WITH clause followed by a SELECT and desugars each
  20. // non-recursive CTE into a derived table. Downstream stages therefore only ever
  21. // see regular SELECTs. Recursive CTEs are rejected explicitly rather than being
  22. // silently mis-executed.
  23. func (p *Parser) parseWithStatement() (Statement, error) {
  24. p.nextToken() // consume WITH
  25. recursive := false
  26. if p.curTokenIs(lexer.TokenIdent) && strings.EqualFold(p.curToken.Literal, "RECURSIVE") {
  27. recursive = true
  28. p.nextToken()
  29. }
  30. var ctes []*cteDef
  31. for {
  32. if !p.curTokenIs(lexer.TokenIdent) {
  33. return nil, p.curError("expected CTE name")
  34. }
  35. cte := &cteDef{name: p.curToken.Literal}
  36. p.nextToken()
  37. if p.curTokenIs(lexer.TokenLParen) {
  38. p.nextToken()
  39. for {
  40. if !p.curTokenIs(lexer.TokenIdent) {
  41. return nil, p.curError("expected column name in CTE column list")
  42. }
  43. cte.cols = append(cte.cols, p.curToken.Literal)
  44. p.nextToken()
  45. if p.curTokenIs(lexer.TokenComma) {
  46. p.nextToken()
  47. continue
  48. }
  49. break
  50. }
  51. if !p.curTokenIs(lexer.TokenRParen) {
  52. return nil, p.curError("expected ) after CTE column list")
  53. }
  54. p.nextToken()
  55. }
  56. if !p.curTokenIs(lexer.TokenAS) {
  57. return nil, p.curError("expected AS in CTE definition")
  58. }
  59. p.nextToken()
  60. if !p.curTokenIs(lexer.TokenLParen) {
  61. return nil, p.curError("expected ( before CTE query")
  62. }
  63. p.nextToken()
  64. query, err := p.parseSelect()
  65. if err != nil {
  66. return nil, err
  67. }
  68. if !p.curTokenIs(lexer.TokenRParen) {
  69. return nil, p.curError("expected ) after CTE query")
  70. }
  71. p.nextToken()
  72. // A recursive CTE's query is a compound (anchor UNION recursive), so its
  73. // column names are applied at materialization time instead.
  74. if !recursive {
  75. if err := applyCTEColumnNames(cte, query); err != nil {
  76. return nil, err
  77. }
  78. }
  79. cte.query = query
  80. ctes = append(ctes, cte)
  81. if p.curTokenIs(lexer.TokenComma) {
  82. p.nextToken()
  83. continue
  84. }
  85. break
  86. }
  87. stmt, err := p.parseStatement()
  88. if err != nil {
  89. return nil, err
  90. }
  91. if recursive {
  92. // Recursive CTEs are materialized by the executor, which needs the
  93. // definitions; desugaring cannot express self-reference. Only SELECT
  94. // carries that machinery today.
  95. sel, ok := stmt.(*SelectStmt)
  96. if !ok {
  97. return nil, p.curError("WITH RECURSIVE is only supported before a SELECT statement")
  98. }
  99. sel.With = make([]*CTE, 0, len(ctes))
  100. for _, c := range ctes {
  101. sel.With = append(sel.With, &CTE{
  102. Name: c.name,
  103. Columns: c.cols,
  104. Recursive: true,
  105. Query: c.query,
  106. })
  107. }
  108. return sel, nil
  109. }
  110. // Each CTE may reference the CTEs defined before it.
  111. for i := range ctes {
  112. if err := substituteSelectCTEs(ctes[i].query, ctes[:i]); err != nil {
  113. return nil, err
  114. }
  115. }
  116. if err := substituteStatementCTEs(stmt, ctes); err != nil {
  117. return nil, err
  118. }
  119. return stmt, nil
  120. }
  121. // substituteStatementCTEs desugars named CTE references inside any DML
  122. // statement. UPDATE reads them from its FROM clause (and expressions); DELETE
  123. // and INSERT read them from subqueries.
  124. func substituteStatementCTEs(stmt Statement, ctes []*cteDef) error {
  125. switch s := stmt.(type) {
  126. case *SelectStmt:
  127. return substituteSelectCTEs(s, ctes)
  128. case *UpdateStmt:
  129. for i := range s.From {
  130. if err := substituteTableRefCTEs(&s.From[i], ctes); err != nil {
  131. return err
  132. }
  133. }
  134. for i := range s.Set {
  135. if err := substituteExprCTEs(s.Set[i].Value, ctes); err != nil {
  136. return err
  137. }
  138. }
  139. if err := substituteExprCTEs(s.Where, ctes); err != nil {
  140. return err
  141. }
  142. for i := range s.Returning {
  143. if err := substituteExprCTEs(s.Returning[i].Expr, ctes); err != nil {
  144. return err
  145. }
  146. }
  147. return nil
  148. case *DeleteStmt:
  149. if err := substituteExprCTEs(s.Where, ctes); err != nil {
  150. return err
  151. }
  152. for i := range s.Returning {
  153. if err := substituteExprCTEs(s.Returning[i].Expr, ctes); err != nil {
  154. return err
  155. }
  156. }
  157. return nil
  158. case *InsertStmt:
  159. if s.Select != nil {
  160. return substituteSelectCTEs(s.Select, ctes)
  161. }
  162. for _, row := range s.Values {
  163. for _, v := range row {
  164. if err := substituteExprCTEs(v, ctes); err != nil {
  165. return err
  166. }
  167. }
  168. }
  169. return nil
  170. default:
  171. return fmt.Errorf("WITH is not supported before this statement")
  172. }
  173. }
  174. // applyCTEColumnNames aliases the CTE query's projection columns with the names
  175. // given in the CTE column list so a derived-table reference exposes them.
  176. func applyCTEColumnNames(cte *cteDef, query *SelectStmt) error {
  177. if len(cte.cols) == 0 {
  178. return nil
  179. }
  180. if query.Compound != nil {
  181. return fmt.Errorf("CTE %q: column list on a compound query is not supported", cte.name)
  182. }
  183. if len(cte.cols) > len(query.Columns) {
  184. return fmt.Errorf("CTE %q: %d column names for %d columns", cte.name, len(cte.cols), len(query.Columns))
  185. }
  186. for i, name := range cte.cols {
  187. query.Columns[i].Alias = name
  188. }
  189. return nil
  190. }
  191. // substituteSelectCTEs replaces every reference to a named CTE with a derived
  192. // table carrying that CTE's query. Substitution recurses through set operations,
  193. // derived tables, joins, and subquery expressions.
  194. func substituteSelectCTEs(sel *SelectStmt, ctes []*cteDef) error {
  195. if sel == nil {
  196. return nil
  197. }
  198. if sel.Compound != nil {
  199. if err := substituteSelectCTEs(sel.Compound.Left, ctes); err != nil {
  200. return err
  201. }
  202. if err := substituteSelectCTEs(sel.Compound.Right, ctes); err != nil {
  203. return err
  204. }
  205. }
  206. for i := range sel.From {
  207. if err := substituteTableRefCTEs(&sel.From[i], ctes); err != nil {
  208. return err
  209. }
  210. }
  211. if err := substituteExprCTEs(sel.Where, ctes); err != nil {
  212. return err
  213. }
  214. for i := range sel.Columns {
  215. if err := substituteExprCTEs(sel.Columns[i].Expr, ctes); err != nil {
  216. return err
  217. }
  218. }
  219. for i := range sel.GroupBy {
  220. if err := substituteExprCTEs(sel.GroupBy[i], ctes); err != nil {
  221. return err
  222. }
  223. }
  224. if err := substituteExprCTEs(sel.Having, ctes); err != nil {
  225. return err
  226. }
  227. for i := range sel.OrderBy {
  228. if err := substituteExprCTEs(sel.OrderBy[i].Expr, ctes); err != nil {
  229. return err
  230. }
  231. }
  232. if err := substituteExprCTEs(sel.Limit, ctes); err != nil {
  233. return err
  234. }
  235. return substituteExprCTEs(sel.Offset, ctes)
  236. }
  237. func lookupCTE(name string, ctes []*cteDef) *cteDef {
  238. for _, cte := range ctes {
  239. if strings.EqualFold(cte.name, name) {
  240. return cte
  241. }
  242. }
  243. return nil
  244. }
  245. func substituteTableRefCTEs(ref *TableRef, ctes []*cteDef) error {
  246. if ref == nil {
  247. return nil
  248. }
  249. if ref.Subquery != nil {
  250. if err := substituteSelectCTEs(ref.Subquery, ctes); err != nil {
  251. return err
  252. }
  253. } else if cte := lookupCTE(ref.Name, ctes); cte != nil {
  254. alias := ref.Alias
  255. if alias == "" {
  256. alias = cte.name
  257. }
  258. ref.Subquery = cte.query
  259. ref.Name = ""
  260. ref.Alias = alias
  261. }
  262. if ref.Join != nil {
  263. return substituteJoinCTEs(ref.Join, ctes)
  264. }
  265. return nil
  266. }
  267. func substituteJoinCTEs(join *JoinClause, ctes []*cteDef) error {
  268. if join == nil {
  269. return nil
  270. }
  271. if err := substituteTableRefCTEs(join.Table, ctes); err != nil {
  272. return err
  273. }
  274. return substituteExprCTEs(join.Condition, ctes)
  275. }
  276. func substituteExprCTEs(expr Expr, ctes []*cteDef) error {
  277. if expr == nil {
  278. return nil
  279. }
  280. switch e := expr.(type) {
  281. case *InExpr:
  282. if err := substituteExprCTEs(e.Left, ctes); err != nil {
  283. return err
  284. }
  285. for _, v := range e.Values {
  286. if err := substituteExprCTEs(v, ctes); err != nil {
  287. return err
  288. }
  289. }
  290. return substituteSelectCTEs(e.Subquery, ctes)
  291. case *SubqueryExpr:
  292. return substituteSelectCTEs(e.Query, ctes)
  293. case *ExistsExpr:
  294. return substituteSelectCTEs(e.Subquery, ctes)
  295. case *BinaryExpr:
  296. if err := substituteExprCTEs(e.Left, ctes); err != nil {
  297. return err
  298. }
  299. return substituteExprCTEs(e.Right, ctes)
  300. case *UnaryExpr:
  301. return substituteExprCTEs(e.Operand, ctes)
  302. case *BetweenExpr:
  303. if err := substituteExprCTEs(e.Left, ctes); err != nil {
  304. return err
  305. }
  306. if err := substituteExprCTEs(e.Low, ctes); err != nil {
  307. return err
  308. }
  309. return substituteExprCTEs(e.High, ctes)
  310. case *LikeExpr:
  311. if err := substituteExprCTEs(e.Left, ctes); err != nil {
  312. return err
  313. }
  314. if err := substituteExprCTEs(e.Pattern, ctes); err != nil {
  315. return err
  316. }
  317. return substituteExprCTEs(e.Escape, ctes)
  318. case *IsNullExpr:
  319. return substituteExprCTEs(e.Left, ctes)
  320. case *IsDistinctExpr:
  321. if err := substituteExprCTEs(e.Left, ctes); err != nil {
  322. return err
  323. }
  324. return substituteExprCTEs(e.Right, ctes)
  325. case *CaseExpr:
  326. if err := substituteExprCTEs(e.Operand, ctes); err != nil {
  327. return err
  328. }
  329. for _, w := range e.Whens {
  330. if err := substituteExprCTEs(w.Condition, ctes); err != nil {
  331. return err
  332. }
  333. if err := substituteExprCTEs(w.Result, ctes); err != nil {
  334. return err
  335. }
  336. }
  337. return substituteExprCTEs(e.Else, ctes)
  338. case *FunctionCall:
  339. for _, a := range e.Args {
  340. if err := substituteExprCTEs(a, ctes); err != nil {
  341. return err
  342. }
  343. }
  344. return nil
  345. case *ParenExpr:
  346. return substituteExprCTEs(e.Expr, ctes)
  347. case *CastExpr:
  348. return substituteExprCTEs(e.Expr, ctes)
  349. }
  350. return nil
  351. }