-
Notifications
You must be signed in to change notification settings - Fork 35
Expand file tree
/
Copy pathstatement_parser.go
More file actions
1050 lines (973 loc) · 37 KB
/
Copy pathstatement_parser.go
File metadata and controls
1050 lines (973 loc) · 37 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
// Copyright 2025 Google LLC
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// https://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package parser
import (
"fmt"
"maps"
"strconv"
"strings"
"sync"
"unicode/utf8"
"cloud.google.com/go/spanner"
"cloud.google.com/go/spanner/admin/database/apiv1/databasepb"
lru "github.com/hashicorp/golang-lru/v2"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)
var ddlStatements = map[string]bool{"CREATE": true, "DROP": true, "ALTER": true, "ANALYZE": true, "GRANT": true, "REVOKE": true, "RENAME": true}
var selectStatements = map[string]bool{"SELECT": true, "WITH": true, "GRAPH": true, "FROM": true}
var insertStatements = map[string]bool{"INSERT": true}
var updateStatements = map[string]bool{"UPDATE": true}
var deleteStatements = map[string]bool{"DELETE": true}
var dmlStatements = union(insertStatements, union(updateStatements, deleteStatements))
var clientSideKeywords = map[string]bool{
"SHOW": true,
"SET": true,
"RESET": true,
"START": true,
"RUN": true,
"ABORT": true,
"BEGIN": true,
"COMMIT": true,
"ROLLBACK": true,
}
var showStatements = map[string]bool{"SHOW": true}
var setStatements = map[string]bool{"SET": true}
var resetStatements = map[string]bool{"RESET": true}
var startStatements = map[string]bool{"START": true}
var runStatements = map[string]bool{"RUN": true}
var abortStatements = map[string]bool{"ABORT": true}
var beginStatements = map[string]bool{"BEGIN": true}
var commitStatements = map[string]bool{"COMMIT": true}
var rollbackStatements = map[string]bool{"ROLLBACK": true}
func union(m1 map[string]bool, m2 map[string]bool) map[string]bool {
res := make(map[string]bool, len(m1)+len(m2))
for k, v := range m1 {
res[k] = v
}
for k, v := range m2 {
res[k] = v
}
return res
}
type statementsCacheEntry struct {
sql string
params []string
namedParamsToIndexedParam map[string]int
info *StatementInfo
parsedStatement ParsedStatement
}
var createParserLock sync.Mutex
var statementParsers = sync.Map{}
func getStatementParser(dialect databasepb.DatabaseDialect, cacheSize int) (*StatementParser, error) {
key := fmt.Sprintf("%v-%v", dialect, cacheSize)
if val, ok := statementParsers.Load(key); ok {
return val.(*StatementParser), nil
} else {
// The statement parser map is thread-safe, so we don't need this lock for that,
// but we take this lock to be sure that we only *create* one instance of
// StatementParser per combination of dialect and cache size.
createParserLock.Lock()
defer createParserLock.Unlock()
parser, err := NewStatementParser(dialect, cacheSize)
if err != nil {
return nil, err
}
statementParsers.Store(key, parser)
return parser, nil
}
}
// StatementParser is a simple, dialect-aware SQL statement parser for Spanner.
// It can be used to determine the type of SQL statement (e.g. DQL/DML/DDL), and
// extract further information from the statement, such as the query parameters.
//
// This is an internal type that can receive breaking changes without prior notice.
type StatementParser struct {
Dialect databasepb.DatabaseDialect
useCache bool
statementsCache *lru.Cache[string, *statementsCacheEntry]
}
// NewStatementParser creates a new parser for the given SQL dialect and with the given
// cache size. Parsers can be shared among multiple database connections. The Spanner
// database/sql driver will only create one parser per database dialect and cache size combination.
//
// This is an internal function that can receive breaking changes without prior notice.
func NewStatementParser(dialect databasepb.DatabaseDialect, cacheSize int) (*StatementParser, error) {
if cacheSize > 0 {
cache, err := lru.New[string, *statementsCacheEntry](cacheSize)
if err != nil {
return nil, err
}
return &StatementParser{Dialect: dialect, statementsCache: cache, useCache: true}, nil
}
return &StatementParser{Dialect: dialect}, nil
}
// CacheSize returns the current size of the statement cache of this StatementParser.
func (p *StatementParser) CacheSize() int {
if p.useCache {
return p.statementsCache.Len()
}
return 0
}
// UseCache returns true if this StatementParser uses a cache.
func (p *StatementParser) UseCache() bool {
return p.useCache
}
// supportsHashSingleLineComments returns true if the database dialect of this parser supports
// comments of the following form:
//
// # This is a single-line comment.
//
// GoogleSQL supports this type of comment.
// PostgreSQL does not support this type of comment.
func (p *StatementParser) supportsHashSingleLineComments() bool {
return p.Dialect != databasepb.DatabaseDialect_POSTGRESQL
}
// supportsNestedComments returns true if the database dialect of this parser supports
// nested comments. Nested comments means that comments of this style are supported:
//
// /* This is a comment. /* This is a nested comment. */ This is still a comment. */
//
// GoogleSQL does not support nested comments.
// PostgreSQL supports nested comments.
func (p *StatementParser) supportsNestedComments() bool {
return p.Dialect == databasepb.DatabaseDialect_POSTGRESQL
}
// identifierQuoteToken returns the token that is used for quoted identifiers.
// GoogleSQL uses ` (backtick) for quoted identifiers.
// PostgreSQL uses " (double quotes) for quoted identifiers.
func (p *StatementParser) identifierQuoteToken() byte {
if p.Dialect == databasepb.DatabaseDialect_POSTGRESQL {
return '"'
}
return '`'
}
// supportsBacktickQuotes returns true if the dialect supports backticks as a valid quote token.
// GoogleSQL supports backtick quotes for identifiers.
// PostgreSQL does not support backticks as a valid quote token.
func (p *StatementParser) supportsBacktickQuotes() bool {
return p.Dialect != databasepb.DatabaseDialect_POSTGRESQL
}
// supportsDoubleQuotedStringLiterals returns true if the SQL dialect supports
// STRING literals starting with ".
func (p *StatementParser) supportsDoubleQuotedStringLiterals() bool {
return p.Dialect != databasepb.DatabaseDialect_POSTGRESQL
}
// supportsTripleQuotedLiterals returns true if the dialect supports quoted strings and identifiers
// that start with three occurrences of the same quote token. Triple-quoted strings and identifiers
// are allowed to contain linefeeds and single occurrences of the quote token. Example:
//
// ”'This is a string, and it's allowed to use a single quote inside the string”'
//
// GoogleSQL supports triple-quoted literals.
// PostgreSQL does not support triple-quoted literals.
func (p *StatementParser) supportsTripleQuotedLiterals() bool {
return p.Dialect != databasepb.DatabaseDialect_POSTGRESQL
}
// supportsDollarQuotedStrings returns true if the dialect supports strings that use double dollar signs
// to mark the start and end of a string. The two dollar signs can optionally contain a tag. Examples:
//
// $$ This is a dollar-quoted string without a tag $$
// $my_tag$ This is a dollar-quoted string with a tag $my_tag$
//
// GoogleSQL does not support dollar-quoted strings.
// PostgreSQL supports dollar-quoted strings.
func (p *StatementParser) supportsDollarQuotedStrings() bool {
return p.Dialect == databasepb.DatabaseDialect_POSTGRESQL
}
// supportsBackslashEscape returns true if the dialect supports escaping a quote within a quoted
// literal by prefixing it with a backslash. Example:
// PostgreSQL: 'It\'s true' => This is invalid. The first part is the string 'It\', and the rest is invalid syntax.
// GoogleSQL: 'It\'s true' => This is the string "It's true".
func (p *StatementParser) supportsBackslashEscape() bool {
return p.Dialect != databasepb.DatabaseDialect_POSTGRESQL
}
// supportsEscapeStrings returns true if the dialect supports enabling escaping using backslashes
// by prepending an e or E to the string. Example:
// e'It\'s a valid string' => This is the string "It's a valid string".
func (p *StatementParser) supportsEscapeStrings() bool {
return p.Dialect == databasepb.DatabaseDialect_POSTGRESQL
}
// supportsEscapeQuoteWithQuote returns true if the dialect supports escaping a quote within a quoted
// literal by repeating the quote twice. Example (note that the way that two single quotes are written in the following
// examples is something that is enforced by gofmt):
// PostgreSQL: 'It”s true' => This is the string "It's true"
// GoogleSQL: 'It”s true' => These are two strings: "It" and "s true".
func (p *StatementParser) supportsEscapeQuoteWithQuote() bool {
return p.Dialect == databasepb.DatabaseDialect_POSTGRESQL
}
var googleSqlReturningKeywords = []string{"then", "return"}
var postgreSqlReturningKeywords = []string{"returning"}
func (p *StatementParser) returningKeywords() []string {
if p.Dialect == databasepb.DatabaseDialect_POSTGRESQL {
return postgreSqlReturningKeywords
}
return googleSqlReturningKeywords
}
// supportsLinefeedInString returns true if the dialect allows linefeeds in standard string literals.
func (p *StatementParser) supportsLinefeedInString() bool {
return p.Dialect == databasepb.DatabaseDialect_POSTGRESQL
}
var googleSqlParamPrefixes = map[byte]bool{'@': true}
var postgresqlParamPrefixes = map[byte]bool{'$': true, '@': true}
func (p *StatementParser) paramPrefixes() map[byte]bool {
if p.Dialect == databasepb.DatabaseDialect_POSTGRESQL {
return postgresqlParamPrefixes
}
return googleSqlParamPrefixes
}
func (p *StatementParser) isValidParamsFirstChar(prefix, c byte) bool {
if p.Dialect == databasepb.DatabaseDialect_POSTGRESQL && prefix == '$' {
// PostgreSQL query parameters have the form $1, $2, ..., $10, ...
return isLatinDigit(c)
}
// GoogleSQL query parameters have the form @name or @_name123
return isLatinLetter(c) || c == '_'
}
func (p *StatementParser) isValidParamsChar(prefix, c byte) bool {
if p.Dialect == databasepb.DatabaseDialect_POSTGRESQL && prefix == '$' {
// PostgreSQL query parameters have the form $1, $2, ..., $10, ...
return isLatinDigit(c)
}
// GoogleSQL query parameters have the form @name or @_name123
return isLatinLetter(c) || isLatinDigit(c) || c == '_'
}
func (p *StatementParser) positionalParameterToNamed(positionalParameterIndex int) string {
if p.Dialect == databasepb.DatabaseDialect_POSTGRESQL {
return "$" + strconv.Itoa(positionalParameterIndex)
}
return "@p" + strconv.Itoa(positionalParameterIndex)
}
func (p *StatementParser) engineSupportsNamedParametersWithPrefix(prefix byte) bool {
if p.Dialect == databasepb.DatabaseDialect_POSTGRESQL {
return prefix == '$'
}
return prefix == '@'
}
// ParseParameters returns the parameters in the given sql string. If the input
// sql contains positional parameters, it returns the converted sql string with
// all positional parameters replaced with named parameters.
// The sql string must be a valid Cloud Spanner sql statement. It may contain
// comments and (string) literals without any restrictions. That is, string
// literals containing for example an email address ('test@test.com') will be
// recognized as a string literal and not returned as a named parameter.
//
// Returns:
// - The modified SQL string that should be sent to Spanner.
// - The list of query parameter names that were found in the SQL string.
// - Only for PostgreSQL: If named parameters are used with a PostgreSQL-dialect database,
// then the parser replaces the named parameters with PostgreSQL-style query parameters.
// The returned map contains a mapping from named query parameter to the PostgreSQL
// query parameter index (e.g. @id => $1, @value => $2).
func (p *StatementParser) ParseParameters(sql string) (string, []string, map[string]int, error) {
return p.findParams(sql)
}
// Skips all whitespaces from the given position and returns the
// position of the next non-whitespace character or len(sql) if
// the string does not contain any whitespaces after pos.
//
// PostgreSQL hints are encoded as comments in the following form:
// /*@ hint_key=hint_value[, hint_key2=hint_value2[,...]] */
// The skipPgHints argument indicates whether those comments should also be skipped or not.
func (p *StatementParser) skipWhitespacesAndComments(sql []byte, pos int, skipPgHints bool) int {
for pos < len(sql) {
c := sql[pos]
if isMultibyte(c) {
break
}
if c == '-' && len(sql) > pos+1 && sql[pos+1] == '-' {
// This is a single line comment starting with '--'.
pos = p.skipSingleLineComment(sql, pos+2)
} else if p.supportsHashSingleLineComments() && c == '#' {
// This is a single line comment starting with '#'.
pos = p.skipSingleLineComment(sql, pos+1)
} else if c == '/' && len(sql) > pos+1 && sql[pos+1] == '*' {
// This is a multi line comment starting with '/*'.
if !skipPgHints && len(sql) > pos+2 && sql[pos+2] == '@' {
// This is a PostgreSQL hint, and we should not skip it.
break
}
pos = p.skipMultiLineComment(sql, pos)
} else if !isSpace(c) {
break
} else {
pos++
}
}
return pos
}
// Skips the next character, quoted literal, quoted identifier or comment in
// the given sql string from the given position and returns the position of the
// next character.
func (p *StatementParser) skip(sql []byte, pos int) (int, error) {
if pos >= len(sql) {
return pos, nil
}
c := sql[pos]
if isMultibyte(c) {
_, size := utf8.DecodeRune(sql[pos:])
pos += size
return pos, nil
}
if c == '\'' || c == '"' || (c == '`' && p.supportsBacktickQuotes()) {
// This is a quoted string or quoted identifier.
pos, _, err := p.skipQuoted(sql, pos, c)
return pos, err
} else if c == '-' && len(sql) > pos+1 && sql[pos+1] == '-' {
// This is a single line comment starting with '--'.
return p.skipSingleLineComment(sql, pos+2), nil
} else if p.supportsHashSingleLineComments() && c == '#' {
// This is a single line comment starting with '#'.
return p.skipSingleLineComment(sql, pos+1), nil
} else if c == '/' && len(sql) > pos+1 && sql[pos+1] == '*' {
// This is a multi line comment starting with '/*'.
return p.skipMultiLineComment(sql, pos), nil
} else if c == '$' && p.supportsDollarQuotedStrings() {
return p.skipDollarQuotedString(sql, pos)
}
return pos + 1, nil
}
func (p *StatementParser) skipDollarQuotedString(sql []byte, pos int) (int, error) {
sp := &simpleParser{sql: sql, pos: pos, statementParser: p}
tag, ok := sp.eatDollarTag()
if !ok {
// Not a valid dollar tag, so only skip the current character.
return pos + 1, nil
}
if _, ok = sp.eatDollarQuotedString(tag); ok {
return sp.pos, nil
}
return 0, spanner.ToSpannerError(status.Errorf(codes.InvalidArgument, "SQL statement contains an unclosed literal: %s", string(sql)))
}
func (p *StatementParser) skipSingleLineComment(sql []byte, pos int) int {
for pos < len(sql) {
c := sql[pos]
if isMultibyte(c) {
_, size := utf8.DecodeRune(sql[pos:])
pos += size
continue
}
if c == '\n' {
break
}
pos++
}
return min(pos+1, len(sql))
}
func (p *StatementParser) skipMultiLineComment(sql []byte, pos int) int {
// Skip '/*'.
pos = pos + 2
level := 1
// Search for '*/' sequences. Depending on whether the dialect supports nested comments or not,
// the comment is considered terminated at either the first occurrence of '*/', or when we have
// seen as many '*/' as there have been occurrences of '/*'.
for pos < len(sql) {
if isMultibyte(sql[pos]) {
_, size := utf8.DecodeRune(sql[pos:])
pos += size
continue
}
if sql[pos] == '*' && len(sql) > pos+1 && sql[pos+1] == '/' {
if p.supportsNestedComments() {
level--
if level == 0 {
return pos + 2
}
} else {
return pos + 2
}
} else if p.supportsNestedComments() {
if sql[pos] == '/' && len(sql) > pos+1 && sql[pos+1] == '*' {
level++
}
}
pos++
}
return pos
}
// skipQuoted skips a quoted string at the given position in the sql string and
// returns the new position, the quote length, or an error if the quoted string
// could not be read.
// The quote length is either 1 for normal quoted strings, and 3 for triple-quoted string.
func (p *StatementParser) skipQuoted(sql []byte, pos int, quote byte) (int, int, error) {
isEscapeString := false
if p.supportsEscapeStrings() && pos > 0 {
// TODO: Also implement support for the standard_conforming_strings property in PostgreSQL.
// See https://www.postgresql.org/docs/current/runtime-config-compatible.html#GUC-STANDARD-CONFORMING-STRINGS
// Check if it is an escape-string. This enables the use of a backslash to start an escape sequence, even if
// the dialect normally does not support that. Escape strings start with an e or E, e.g. "e'It\'s valid'".
// The second part of the check is to verify that the e or E is not part of a keyword, e.g. WHERE.
// The following is valid SQL, but does not designate an escape-string:
// SELECT * FROM my_table WHERE'test'=col1;
isEscapeString = (sql[pos-1] == 'e' || sql[pos-1] == 'E') && (pos == 1 || !isLatinLetter(sql[pos-2]))
}
isTripleQuoted := p.supportsTripleQuotedLiterals() && len(sql) > pos+2 && sql[pos+1] == quote && sql[pos+2] == quote
if isTripleQuoted && (isMultibyte(sql[pos+1]) || isMultibyte(sql[pos+2])) {
isTripleQuoted = false
}
var quoteLength int
if isTripleQuoted {
quoteLength = 3
} else {
quoteLength = 1
}
pos += quoteLength
for pos < len(sql) {
c := sql[pos]
if isMultibyte(c) {
_, size := utf8.DecodeRune(sql[pos:])
pos += size
continue
}
if c == quote {
if isTripleQuoted {
// Check if this is the end of the triple-quoted string.
if len(sql) > pos+2 && sql[pos+1] == quote && sql[pos+2] == quote {
return pos + 3, quoteLength, nil
}
} else {
if p.supportsEscapeQuoteWithQuote() && len(sql) > pos+1 && sql[pos+1] == quote {
pos += 2
continue
}
// This was the end quote.
return pos + 1, quoteLength, nil
}
} else if (p.supportsBackslashEscape() || isEscapeString) && len(sql) > pos+1 && c == '\\' && (sql[pos+1] == quote || sql[pos+1] == '\\') {
// This is an escaped quote (e.g. 'foo\'bar') or an escaped backslash (e.g 'test\\').
// Note that in raw strings, the \ officially does not start an
// escape sequence, but the result is still the same, as in a raw
// string 'both characters are preserved'.
pos += 1
} else if !p.supportsLinefeedInString() && !isTripleQuoted && c == '\n' {
break
}
pos++
}
return 0, 0, spanner.ToSpannerError(status.Errorf(codes.InvalidArgument, "SQL statement contains an unclosed literal: %s", string(sql)))
}
// findParams finds all query parameters in the given SQL string.
// The SQL string may contain comments and statement hints.
func (p *StatementParser) findParams(sql string) (string, []string, map[string]int, error) {
if !p.useCache {
return p.calculateFindParamsResult(sql)
}
if val, ok := p.statementsCache.Get(sql); ok {
return val.sql, val.params, val.namedParamsToIndexedParam, nil
} else {
namedParamsSql, params, namedParamsToIndexedParam, err := p.calculateFindParamsResult(sql)
if err != nil {
return "", nil, nil, err
}
info := p.DetectStatementType(sql)
cachedParams := make([]string, len(params))
copy(cachedParams, params)
cachedNamedParamsToIndexedParam := make(map[string]int, len(namedParamsToIndexedParam))
maps.Copy(cachedNamedParamsToIndexedParam, namedParamsToIndexedParam)
p.statementsCache.Add(sql, &statementsCacheEntry{
sql: namedParamsSql,
params: cachedParams,
namedParamsToIndexedParam: cachedNamedParamsToIndexedParam,
info: info,
})
return namedParamsSql, params, namedParamsToIndexedParam, nil
}
}
const onlyOneParameterStyle = "statement must only use one parameter style (named or positional): %s"
func (p *StatementParser) calculateFindParamsResult(sql string) (string, []string, map[string]int, error) {
const positionalParamChar = '?'
paramPrefixes := p.paramPrefixes()
hasNamedParameter := false
hasIndexedParameter := false
hasPositionalParameter := false
numPotentialParams := strings.Count(sql, string(positionalParamChar))
for prefix := range paramPrefixes {
numPotentialParams += strings.Count(sql, string(prefix))
}
// Spanner does not support more than 950 query parameters in a query,
// so we limit the number of parameters that we pre-allocate room for
// here to that number. This prevents allocation of a large slice if
// the SQL string happens to contain a string literal that contains
// a lot of question marks.
namedParams := make([]string, 0, min(numPotentialParams, 950))
namedParamsToIndexedParam := make(map[string]int)
foundParams := make(map[string]bool)
parsedSQL := strings.Builder{}
parsedSQL.Grow(len(sql))
positionalParameterIndex := 1
parser := &simpleParser{sql: []byte(sql), statementParser: p}
addNamedParam := func(startIndex, endIndex int, paramPrefix byte, paramNamePrefix string) {
paramName := paramNamePrefix + string(parser.sql[startIndex:endIndex])
if p.engineSupportsNamedParametersWithPrefix(paramPrefix) {
if !foundParams[paramName] {
foundParams[paramName] = true
namedParams = append(namedParams, paramName)
}
} else {
index, ok := namedParamsToIndexedParam[paramName]
if !ok {
namedParamsToIndexedParam[paramName] = positionalParameterIndex
index = positionalParameterIndex
positionalParameterIndex++
}
translatedParamName := fmt.Sprintf("p%d", index)
if !foundParams[translatedParamName] {
foundParams[translatedParamName] = true
namedParams = append(namedParams, translatedParamName)
}
// Use $1, $2, etc. in the SQL string.
parsedSQL.Write([]byte(fmt.Sprintf("$%d", index)))
}
}
for parser.pos < len(parser.sql) {
startPos := parser.pos
parser.skipWhitespacesAndComments()
parsedSQL.Write(parser.sql[startPos:parser.pos])
if parser.pos >= len(parser.sql) {
break
}
if parser.isMultibyte() {
startPos = parser.pos
parser.nextChar()
parsedSQL.Write(parser.sql[startPos:parser.pos])
continue
}
c := parser.sql[parser.pos]
// We are not in a quoted string.
// GoogleSQL: It's a parameter if it is an '@' followed by a letter or an underscore.
// PostgreSQL: It's a parameter if it is a '$' followed by a digit.
// See https://cloud.google.com/spanner/docs/lexical#identifiers for identifier rules.
// Note that for PostgreSQL we support both styles of parameters (both $1 and @name).
// The reason for this is that other PostgreSQL drivers for Go also do this.
if _, ok := paramPrefixes[c]; ok && len(parser.sql) > parser.pos+1 && p.isValidParamsFirstChar(c, parser.sql[parser.pos+1]) {
if hasPositionalParameter {
return sql, nil, nil, spanner.ToSpannerError(status.Errorf(codes.InvalidArgument, onlyOneParameterStyle, sql))
}
paramPrefix := c
paramNamePrefix := ""
if paramPrefix == '$' {
paramNamePrefix = "p"
if hasNamedParameter {
return sql, nil, nil, spanner.ToSpannerError(status.Errorf(codes.InvalidArgument, onlyOneParameterStyle, sql))
}
hasIndexedParameter = true
} else {
if hasIndexedParameter {
return sql, nil, nil, spanner.ToSpannerError(status.Errorf(codes.InvalidArgument, onlyOneParameterStyle, sql))
}
hasNamedParameter = true
}
// Only write the parameter name to the SQL string if the backend engine
// actually supports named parameters. Otherwise, we replace it with a
// positional parameter.
if p.engineSupportsNamedParametersWithPrefix(paramPrefix) {
parsedSQL.WriteByte(paramPrefix)
}
parser.pos++
startIndex := parser.pos
for parser.pos < len(parser.sql) {
if parser.isMultibyte() || parser.pos == len(parser.sql)-1 || !p.isValidParamsChar(paramPrefix, parser.sql[parser.pos]) {
if parser.pos == len(parser.sql)-1 && p.isValidParamsChar(paramPrefix, parser.sql[parser.pos]) {
// Include the current character in the parameter name, as the reason that we are stopping
// the search here, is that we have reached the end of the statement.
if p.engineSupportsNamedParametersWithPrefix(paramPrefix) {
parsedSQL.WriteByte(parser.sql[parser.pos])
}
parser.pos++
}
addNamedParam(startIndex, parser.pos, paramPrefix, paramNamePrefix)
break
}
if p.engineSupportsNamedParametersWithPrefix(paramPrefix) {
parsedSQL.WriteByte(parser.sql[parser.pos])
}
parser.pos++
}
} else if c == positionalParamChar {
if hasNamedParameter || hasIndexedParameter {
return sql, nil, nil, spanner.ToSpannerError(status.Errorf(codes.InvalidArgument, onlyOneParameterStyle, sql))
}
hasPositionalParameter = true
parsedSQL.WriteString(p.positionalParameterToNamed(positionalParameterIndex))
namedParams = append(namedParams, "p"+strconv.Itoa(positionalParameterIndex))
positionalParameterIndex++
parser.pos++
} else {
startPos = parser.pos
newPos, err := p.skip(parser.sql, parser.pos)
if err != nil {
return sql, nil, nil, err
}
parsedSQL.Write(parser.sql[startPos:newPos])
parser.pos = newPos
}
}
if hasNamedParameter && len(namedParamsToIndexedParam) == 0 {
return sql, namedParams, nil, nil
}
return parsedSQL.String(), namedParams, namedParamsToIndexedParam, nil
}
// ParseClientSideStatement returns the executableClientSideStatement that
// corresponds with the given query string, or nil if it is not a valid client
// side statement.
func (p *StatementParser) ParseClientSideStatement(query string) (ParsedStatement, error) {
if p.useCache {
if val, ok := p.statementsCache.Get(query); ok {
if val.info.StatementType == StatementTypeClientSide {
if val.parsedStatement != nil {
return val.parsedStatement, nil
}
} else {
return nil, nil
}
}
}
// Determine whether it could be a valid client-side statement by looking at the first keyword.
sp := &simpleParser{sql: []byte(query), statementParser: p}
keyword := strings.ToUpper(sp.readKeyword())
if _, ok := clientSideKeywords[keyword]; !ok {
return nil, nil
}
if stmt, err := parseStatement(p, keyword, query); err != nil {
return nil, err
} else if stmt != nil {
if p.useCache {
cacheEntry := &statementsCacheEntry{
sql: query,
parsedStatement: stmt,
info: &StatementInfo{
StatementType: StatementTypeClientSide,
},
}
p.statementsCache.Add(query, cacheEntry)
}
return stmt, nil
}
return nil, nil
}
func isClientSideKeyword(keyword string) bool {
return isStatementKeyword(keyword, clientSideKeywords)
}
// isDDL returns true if the given sql string is a DDL statement.
// This function assumes that any comments and hints at the start
// of the sql string have been removed.
func (p *StatementParser) isDDL(query string) bool {
info := p.DetectStatementType(query)
return info.StatementType == StatementTypeDdl
}
func isDDLKeyword(keyword string) bool {
return isStatementKeyword(keyword, ddlStatements)
}
func isDmlKeyword(keyword string) bool {
return isStatementKeyword(keyword, dmlStatements)
}
func isQueryKeyword(keyword string) bool {
return isStatementKeyword(keyword, selectStatements)
}
func isShowStatementKeyword(keyword string) bool {
return isStatementKeyword(keyword, showStatements)
}
func isSetStatementKeyword(keyword string) bool {
return isStatementKeyword(keyword, setStatements)
}
func isResetStatementKeyword(keyword string) bool {
return isStatementKeyword(keyword, resetStatements)
}
func isStartStatementKeyword(keyword string) bool {
return isStatementKeyword(keyword, startStatements)
}
func isRunStatementKeyword(keyword string) bool {
return isStatementKeyword(keyword, runStatements)
}
func isAbortStatementKeyword(keyword string) bool {
return isStatementKeyword(keyword, abortStatements)
}
func isBeginStatementKeyword(keyword string) bool {
return isStatementKeyword(keyword, beginStatements)
}
func isCommitStatementKeyword(keyword string) bool {
return isStatementKeyword(keyword, commitStatements)
}
func isRollbackStatementKeyword(keyword string) bool {
return isStatementKeyword(keyword, rollbackStatements)
}
func isStatementKeyword(keyword string, keywords map[string]bool) bool {
_, ok := keywords[keyword]
return ok
}
// StatementType indicates the type of SQL statement.
type StatementType int
const (
// StatementTypeUnknown indicates that the parser was not able to determine the
// type of SQL statement. This could be an indication that the SQL string is invalid,
// or that it uses a syntax that is not (yet) supported by the parser.
StatementTypeUnknown StatementType = iota
// StatementTypeQuery indicates that the statement is a query that will return rows from
// Spanner, and that will not make any modifications to the database.
StatementTypeQuery
// StatementTypeDml indicates that the statement is a data modification language (DML)
// statement that will make modifications to the data in the database. It may or may not
// return rows, depending on whether it contains a THEN RETURN (GoogleSQL) or RETURNING
// (PostgreSQL) clause.
StatementTypeDml
// StatementTypeDdl indicates that the statement is a data definition language (DDL)
// statement that will modify the schema of the database. It will never return rows.
StatementTypeDdl
// StatementTypeClientSide indicates that the statement will be handled client-side in
// the database/sql driver, and not be sent to Spanner. Examples of this includes SHOW
// and SET statements.
StatementTypeClientSide
)
func (st StatementType) String() string {
switch st {
case StatementTypeQuery:
return "Query"
case StatementTypeDml:
return "DML"
case StatementTypeDdl:
return "DDL"
case StatementTypeClientSide:
return "ClientSide"
default:
return "Unknown"
}
}
// DmlType designates the type of modification that a DML statement will execute.
type DmlType int
const (
DmlTypeUnknown DmlType = iota
DmlTypeInsert
DmlTypeUpdate
DmlTypeDelete
)
// StatementInfo contains the type of SQL statement, and in case of a DML statement,
// the type of DML command.
type StatementInfo struct {
StatementType StatementType
DmlType DmlType
// TODO: Implement parsing of THEN RETURN / RETURNING clauses
HasThenReturn bool
}
// DetectStatementType returns the type of SQL statement based on the first
// keyword that is found in the SQL statement.
func (p *StatementParser) DetectStatementType(sql string) *StatementInfo {
if !p.useCache {
return p.calculateDetectStatementType(sql)
}
if val, ok := p.statementsCache.Get(sql); ok {
return val.info
} else {
info := p.calculateDetectStatementType(sql)
namedParamsSql, params, namedParamsToIndexedParam, err := p.calculateFindParamsResult(sql)
if err == nil {
p.statementsCache.Add(sql, &statementsCacheEntry{
sql: namedParamsSql,
params: params,
namedParamsToIndexedParam: namedParamsToIndexedParam,
info: info,
})
}
return info
}
}
func (p *StatementParser) calculateDetectStatementType(sql string) *StatementInfo {
parser := &simpleParser{sql: []byte(sql), statementParser: p}
_, _ = parser.skipStatementHint()
keyword := strings.ToUpper(parser.readKeyword())
if isQueryKeyword(keyword) {
return &StatementInfo{StatementType: StatementTypeQuery}
} else if isDmlKeyword(keyword) {
return &StatementInfo{
StatementType: StatementTypeDml,
DmlType: detectDmlKeyword(keyword),
HasThenReturn: p.detectHasThenReturnClause(sql),
}
} else if isDDLKeyword(keyword) {
return &StatementInfo{StatementType: StatementTypeDdl}
} else if isClientSideKeyword(keyword) {
return &StatementInfo{StatementType: StatementTypeClientSide}
}
return &StatementInfo{StatementType: StatementTypeUnknown}
}
func detectDmlKeyword(keyword string) DmlType {
if isStatementKeyword(keyword, insertStatements) {
return DmlTypeInsert
} else if isStatementKeyword(keyword, updateStatements) {
return DmlTypeUpdate
} else if isStatementKeyword(keyword, deleteStatements) {
return DmlTypeDelete
}
return DmlTypeUnknown
}
func (p *StatementParser) extractSetStatementsFromHints(sql string) (*ParsedSetStatement, error) {
sp := &simpleParser{sql: []byte(sql), statementParser: p}
if ok, startPos := sp.skipStatementHint(); ok {
// Mark the start and end of the statement hint and extract the values in the hint.
endPos := sp.pos
sp.pos = startPos
if p.Dialect == databasepb.DatabaseDialect_POSTGRESQL {
// eatTokensOnly will only look for the following character sequence: '/*@'
// It will not interpret it as a comment.
sp.eatTokensOnly(postgreSqlStatementHintPrefix)
} else {
sp.eatTokens(googleSqlStatementHintPrefix)
}
// The default is that the hint ends with a single '}'.
endIndex := endPos - 1
if p.Dialect == databasepb.DatabaseDialect_POSTGRESQL {
// The hint ends with '*/'
endIndex = endPos - 2
}
if endIndex > sp.pos {
return p.extractConnectionVariables(sql[sp.pos:endIndex])
}
}
return nil, nil
}
func (p *StatementParser) extractConnectionVariables(sql string) (*ParsedSetStatement, error) {
sp := &simpleParser{sql: []byte(sql), statementParser: p}
statement := &ParsedSetStatement{
Identifiers: make([]Identifier, 0, 2),
Literals: make([]Literal, 0, 2),
}
for {
if !sp.hasMoreTokens() {
break
}
identifier, err := sp.eatIdentifier()
if err != nil {
return nil, err
}
if !sp.eatToken('=') {
return nil, status.Errorf(codes.InvalidArgument, "missing '=' token after %s in hint", identifier)
}
literal, err := sp.eatLiteral()
if err != nil {
return nil, err
}
statement.Identifiers = append(statement.Identifiers, identifier)
statement.Literals = append(statement.Literals, literal)
if !sp.eatToken(',') {
break
}
}
if sp.hasMoreTokens() {
return nil, status.Errorf(codes.InvalidArgument, "unexpected tokens: %s", string(sp.sql[sp.pos:]))
}
return statement, nil
}
func (p *StatementParser) detectHasThenReturnClause(sql string) bool {
parser := &simpleParser{sql: []byte(sql), statementParser: p}
for parser.pos < len(parser.sql) {
parser.skipWhitespacesAndComments()
if parser.pos >= len(parser.sql) {
break
}
if parser.isMultibyte() {
parser.nextChar()
continue
}
c := parser.sql[parser.pos]
if c == '(' {
// THEN RETURN / RETURNING must be in the outermost expression.
if err := parser.skipExpressionInBrackets(); err != nil {
return false
}
continue
}
if parser.eatKeywords(p.returningKeywords()) {
return true
}
newPos, err := p.skip(parser.sql, parser.pos)
if err != nil {
return false
}
parser.pos = newPos
}
return false
}
func (p *StatementParser) IsCreateDatabaseStatement(sql string) bool {
return isCreateDatabase(p, sql)
}
func (p *StatementParser) IsDropDatabaseStatement(sql string) bool {
return isDropDatabase(p, sql)
}
// Split splits a SQL string that potentially contains multiple statements separated by
// semicolons into individual statements.
//
// Returns false if the SQL string only contains a single statement (including when the SQL string
// contains a single statement that is terminated by a semicolon). Whitespaces and comments after
// the last semicolon in the SQL string are ignored.
//
// Returns true and a slice containing the individual statements if the SQL string contains more
// than one statement. The slice always contains more than one element.
func (p *StatementParser) Split(sql string) (bool, []string, error) {
return p.split(sql, ';')
}
func (p *StatementParser) split(sql string, sep byte) (bool, []string, error) {
// Return early if the string does not contain the separator.
firstIndex := strings.IndexByte(sql, sep)
if firstIndex == -1 {
return false, nil, nil
}
tokens := []byte(sql)
// Also return early if it is a single statement that is just terminated by a semicolon.