2
0

indexexpr.go 9.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273
  1. package executor
  2. import (
  3. "fmt"
  4. "strings"
  5. "sync"
  6. "github.com/danfragoso/pizzasql-next/pkg/parser"
  7. "github.com/danfragoso/pizzasql-next/pkg/storage"
  8. )
  9. // statelessEval is a shared, immutable Executor used only to evaluate validated,
  10. // side-effect-free expressions (expression indexes and generated columns) from
  11. // the storage layer. It holds no session, catalog, or per-connection state, so
  12. // concurrent evaluations from different connections share no mutable data. The
  13. // expression validators below guarantee that only subquery-free, deterministic
  14. // expressions reach it.
  15. var statelessEval = &Executor{}
  16. // storedExprCache memoizes parsed stored expressions across all connections. A
  17. // stored expression is immutable text, so the cache is safe to share.
  18. var storedExprCache sync.Map // string -> parser.Expr
  19. // parseStoredExpr parses (once) an expression persisted in the durable schema
  20. // (generated column or expression index) and rejects anything that is not
  21. // deterministic, subquery-free SQL. The same text always parses to the same
  22. // expression, so the result is cached.
  23. func parseStoredExpr(text string) (parser.Expr, error) {
  24. if cached, ok := storedExprCache.Load(text); ok {
  25. return cached.(parser.Expr), nil
  26. }
  27. expr, err := parser.ParseExpr(text)
  28. if err != nil {
  29. return nil, fmt.Errorf("invalid stored expression %q: %w", text, err)
  30. }
  31. if err := ValidateDeterministicExpr(expr); err != nil {
  32. return nil, fmt.Errorf("invalid stored expression %q: %w", text, err)
  33. }
  34. storedExprCache.Store(text, expr)
  35. return expr, nil
  36. }
  37. // EvalStoredExpression evaluates a persisted, validated expression against a
  38. // row. It is the storage layer's expression-index evaluator and is stateless and
  39. // concurrency-safe.
  40. func EvalStoredExpression(text string, row storage.Row) (interface{}, error) {
  41. expr, err := parseStoredExpr(text)
  42. if err != nil {
  43. return nil, err
  44. }
  45. return statelessEval.evalExpr(expr, row)
  46. }
  47. // ValidateDeterministicExpr rejects an expression that SQLite would not allow in
  48. // an index or generated column: subqueries, window functions, aggregates, and
  49. // non-deterministic or unknown functions.
  50. func ValidateDeterministicExpr(expr parser.Expr) error {
  51. switch e := expr.(type) {
  52. case nil:
  53. return nil
  54. case *parser.LiteralExpr:
  55. return nil
  56. case *parser.ColumnRef:
  57. return nil
  58. case *parser.ParenExpr:
  59. return ValidateDeterministicExpr(e.Expr)
  60. case *parser.UnaryExpr:
  61. return ValidateDeterministicExpr(e.Operand)
  62. case *parser.BinaryExpr:
  63. if err := ValidateDeterministicExpr(e.Left); err != nil {
  64. return err
  65. }
  66. return ValidateDeterministicExpr(e.Right)
  67. case *parser.IsNullExpr:
  68. return ValidateDeterministicExpr(e.Left)
  69. case *parser.IsDistinctExpr:
  70. if err := ValidateDeterministicExpr(e.Left); err != nil {
  71. return err
  72. }
  73. return ValidateDeterministicExpr(e.Right)
  74. case *parser.BetweenExpr:
  75. if err := ValidateDeterministicExpr(e.Left); err != nil {
  76. return err
  77. }
  78. if err := ValidateDeterministicExpr(e.Low); err != nil {
  79. return err
  80. }
  81. return ValidateDeterministicExpr(e.High)
  82. case *parser.LikeExpr:
  83. if err := ValidateDeterministicExpr(e.Left); err != nil {
  84. return err
  85. }
  86. if err := ValidateDeterministicExpr(e.Pattern); err != nil {
  87. return err
  88. }
  89. return ValidateDeterministicExpr(e.Escape)
  90. case *parser.CaseExpr:
  91. if err := ValidateDeterministicExpr(e.Operand); err != nil {
  92. return err
  93. }
  94. for _, w := range e.Whens {
  95. if err := ValidateDeterministicExpr(w.Condition); err != nil {
  96. return err
  97. }
  98. if err := ValidateDeterministicExpr(w.Result); err != nil {
  99. return err
  100. }
  101. }
  102. return ValidateDeterministicExpr(e.Else)
  103. case *parser.CastExpr:
  104. return ValidateDeterministicExpr(e.Expr)
  105. case *parser.InExpr:
  106. if e.Subquery != nil {
  107. return fmt.Errorf("subqueries are not allowed in index expressions")
  108. }
  109. if err := ValidateDeterministicExpr(e.Left); err != nil {
  110. return err
  111. }
  112. for _, v := range e.Values {
  113. if err := ValidateDeterministicExpr(v); err != nil {
  114. return err
  115. }
  116. }
  117. return nil
  118. case *parser.FunctionCall:
  119. if e.Star {
  120. return fmt.Errorf("function %s(*) is not allowed in index expressions", e.Name)
  121. }
  122. if isAggregateFunctionName(e.Name, len(e.Args)) {
  123. return fmt.Errorf("aggregate function %s() is not allowed in index expressions", e.Name)
  124. }
  125. if !isDeterministicFunction(e.Name) {
  126. return fmt.Errorf("non-deterministic or unsupported function %s() is not allowed in index expressions", e.Name)
  127. }
  128. for _, a := range e.Args {
  129. if err := ValidateDeterministicExpr(a); err != nil {
  130. return err
  131. }
  132. }
  133. return nil
  134. case *parser.SubqueryExpr:
  135. return fmt.Errorf("subqueries are not allowed in index expressions")
  136. case *parser.ExistsExpr:
  137. return fmt.Errorf("subqueries are not allowed in index expressions")
  138. case *parser.WindowExpr:
  139. return fmt.Errorf("window functions are not allowed in index expressions")
  140. default:
  141. return fmt.Errorf("unsupported expression in index definition: %T", expr)
  142. }
  143. }
  144. // isAggregateFunctionName reports whether a function name is an aggregate in the
  145. // given call shape. MIN/MAX are scalar with two or more arguments.
  146. func isAggregateFunctionName(name string, argCount int) bool {
  147. switch strings.ToUpper(name) {
  148. case "COUNT", "SUM", "AVG", "TOTAL", "GROUP_CONCAT",
  149. "JSON_GROUP_ARRAY", "JSONB_GROUP_ARRAY", "JSON_GROUP_OBJECT", "JSONB_GROUP_OBJECT":
  150. return true
  151. case "MIN", "MAX":
  152. return argCount < 2
  153. }
  154. return false
  155. }
  156. // deterministicFunctions is the allowlist of scalar functions that may appear in
  157. // a persisted expression. It intentionally excludes date/time functions (which
  158. // are non-deterministic when using "now") and every random/session function.
  159. var deterministicFunctions = map[string]bool{
  160. "LOWER": true, "UPPER": true, "LENGTH": true, "ABS": true,
  161. "COALESCE": true, "NULLIF": true, "IFNULL": true, "NVL": true,
  162. "TYPEOF": true, "SUBSTR": true, "SUBSTRING": true, "TRIM": true,
  163. "REPLACE": true, "PRINTF": true, "HEX": true, "UNHEX": true,
  164. "ZEROBLOB": true, "INSTR": true, "GLOB": true, "ROUND": true,
  165. "MAX": true, "MIN": true, "CONCAT": true, "PERCENT_DIFF": true,
  166. // JSON1 scalar functions.
  167. "JSON": true, "JSONB": true, "JSON_VALID": true, "JSONB_VALID": true,
  168. "JSON_TYPE": true, "JSON_EXTRACT": true, "JSONB_EXTRACT": true,
  169. "JSON_SET": true, "JSONB_SET": true, "JSON_INSERT": true, "JSONB_INSERT": true,
  170. "JSON_REPLACE": true, "JSONB_REPLACE": true, "JSON_REMOVE": true, "JSONB_REMOVE": true,
  171. "JSON_ARRAY": true, "JSONB_ARRAY": true, "JSON_OBJECT": true, "JSONB_OBJECT": true,
  172. "JSON_QUOTE": true, "JSONB_QUOTE": true,
  173. }
  174. func isDeterministicFunction(name string) bool {
  175. return deterministicFunctions[strings.ToUpper(name)]
  176. }
  177. // ValidateIndexColumns checks that every column referenced by a stored
  178. // expression exists on the table (or is a hidden rowid alias). It is applied to
  179. // expression indexes and generated columns at creation time.
  180. func ValidateIndexColumns(expr parser.Expr, schema *storage.Schema) error {
  181. switch e := expr.(type) {
  182. case nil:
  183. return nil
  184. case *parser.ColumnRef:
  185. if e.Column == "*" {
  186. return fmt.Errorf("wildcards are not allowed in index expressions")
  187. }
  188. if storage.IsRowIDColumn(e.Column) {
  189. return nil
  190. }
  191. if _, ok := schema.GetColumn(e.Column); !ok {
  192. return fmt.Errorf("column not found in index expression: %s", e.Column)
  193. }
  194. return nil
  195. case *parser.ParenExpr:
  196. return ValidateIndexColumns(e.Expr, schema)
  197. case *parser.UnaryExpr:
  198. return ValidateIndexColumns(e.Operand, schema)
  199. case *parser.BinaryExpr:
  200. if err := ValidateIndexColumns(e.Left, schema); err != nil {
  201. return err
  202. }
  203. return ValidateIndexColumns(e.Right, schema)
  204. case *parser.IsNullExpr:
  205. return ValidateIndexColumns(e.Left, schema)
  206. case *parser.IsDistinctExpr:
  207. if err := ValidateIndexColumns(e.Left, schema); err != nil {
  208. return err
  209. }
  210. return ValidateIndexColumns(e.Right, schema)
  211. case *parser.BetweenExpr:
  212. if err := ValidateIndexColumns(e.Left, schema); err != nil {
  213. return err
  214. }
  215. if err := ValidateIndexColumns(e.Low, schema); err != nil {
  216. return err
  217. }
  218. return ValidateIndexColumns(e.High, schema)
  219. case *parser.LikeExpr:
  220. if err := ValidateIndexColumns(e.Left, schema); err != nil {
  221. return err
  222. }
  223. if err := ValidateIndexColumns(e.Pattern, schema); err != nil {
  224. return err
  225. }
  226. return ValidateIndexColumns(e.Escape, schema)
  227. case *parser.CaseExpr:
  228. if err := ValidateIndexColumns(e.Operand, schema); err != nil {
  229. return err
  230. }
  231. for _, w := range e.Whens {
  232. if err := ValidateIndexColumns(w.Condition, schema); err != nil {
  233. return err
  234. }
  235. if err := ValidateIndexColumns(w.Result, schema); err != nil {
  236. return err
  237. }
  238. }
  239. return ValidateIndexColumns(e.Else, schema)
  240. case *parser.CastExpr:
  241. return ValidateIndexColumns(e.Expr, schema)
  242. case *parser.InExpr:
  243. if err := ValidateIndexColumns(e.Left, schema); err != nil {
  244. return err
  245. }
  246. for _, v := range e.Values {
  247. if err := ValidateIndexColumns(v, schema); err != nil {
  248. return err
  249. }
  250. }
  251. return nil
  252. case *parser.FunctionCall:
  253. for _, a := range e.Args {
  254. if err := ValidateIndexColumns(a, schema); err != nil {
  255. return err
  256. }
  257. }
  258. return nil
  259. default:
  260. return nil
  261. }
  262. }