Source file
src/database/sql/fakedb_test.go
1
2
3
4
5 package sql
6
7 import (
8 "bytes"
9 "context"
10 "database/sql/driver"
11 "errors"
12 "fmt"
13 "io"
14 "reflect"
15 "slices"
16 "strconv"
17 "strings"
18 "sync"
19 "testing"
20 "time"
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 type fakeDriver struct {
48 mu sync.Mutex
49 openCount int
50 closeCount int
51 waitCh chan struct{}
52 waitingCh chan struct{}
53 dbs map[string]*fakeDB
54 }
55
56 type fakeConnector struct {
57 name string
58
59 waiter func(context.Context)
60 closed bool
61 }
62
63 func (c *fakeConnector) Connect(context.Context) (driver.Conn, error) {
64 conn, err := fdriver.Open(c.name)
65 if err != nil {
66 return nil, err
67 }
68 conn.(*fakeConn).waiter = c.waiter
69 return conn, nil
70 }
71
72 func getFakeConn(c driver.Conn) *fakeConn {
73 return c.(interface {
74 getFakeConn() *fakeConn
75 }).getFakeConn()
76 }
77
78 func (c *fakeConn) getFakeConn() *fakeConn {
79 return c
80 }
81
82 func getRowsCursor(rs *Rows) *rowsCursor {
83 return rs.rowsi.(interface {
84 getRowsCursor() *rowsCursor
85 }).getRowsCursor()
86 }
87
88 func (rc *rowsCursor) getRowsCursor() *rowsCursor {
89 return rc
90 }
91
92 func (c *fakeConnector) Driver() driver.Driver {
93 return fdriver
94 }
95
96 func (c *fakeConnector) Close() error {
97 if c.closed {
98 return errors.New("fakedb: connector is closed")
99 }
100 c.closed = true
101 return nil
102 }
103
104 type fakeDriverCtx struct {
105 fakeDriver
106 }
107
108 var _ driver.DriverContext = &fakeDriverCtx{}
109
110 func (cc *fakeDriverCtx) OpenConnector(name string) (driver.Connector, error) {
111 return &fakeConnector{name: name}, nil
112 }
113
114 type fakeDB struct {
115 name string
116
117 mu sync.Mutex
118 tables map[string]*table
119 badConn bool
120 allowAny bool
121 }
122
123 type fakeError struct {
124 Message string
125 Wrapped error
126 }
127
128 func (err fakeError) Error() string {
129 return err.Message
130 }
131
132 func (err fakeError) Unwrap() error {
133 return err.Wrapped
134 }
135
136 type table struct {
137 mu sync.Mutex
138 colname []string
139 coltype []string
140 rows []*row
141 }
142
143 func (t *table) columnIndex(name string) int {
144 return slices.Index(t.colname, name)
145 }
146
147 type row struct {
148 cols []any
149 }
150
151 type memToucher interface {
152
153 touchMem()
154 }
155
156 type fakeConn struct {
157 db *fakeDB
158
159 currTx *fakeTx
160
161
162
163 line int64
164
165
166 mu sync.Mutex
167 stmtsMade int
168 stmtsClosed int
169 numPrepare int
170
171
172 bad bool
173 stickyBad bool
174
175 skipDirtySession bool
176
177
178
179 dirtySession bool
180
181
182
183 waiter func(context.Context)
184 }
185
186 func (c *fakeConn) touchMem() {
187 c.line++
188 }
189
190 func (c *fakeConn) incrStat(v *int) {
191 c.mu.Lock()
192 *v++
193 c.mu.Unlock()
194 }
195
196 type fakeTx struct {
197 c *fakeConn
198 }
199
200 type boundCol struct {
201 Column string
202 Placeholder string
203 Ordinal int
204 }
205
206 type fakeStmt struct {
207 memToucher
208 c *fakeConn
209 q string
210
211 cmd string
212 table string
213 panic string
214 wait time.Duration
215
216 next *fakeStmt
217
218 closed bool
219
220 colName []string
221 colType []string
222 colValue []any
223 placeholders int
224
225 whereCol []boundCol
226
227 placeholderConverter []driver.ValueConverter
228 }
229
230 var fdriver driver.Driver = &fakeDriver{}
231
232 func init() {
233 Register("test", fdriver)
234 }
235
236 type Dummy struct {
237 driver.Driver
238 }
239
240 func TestDrivers(t *testing.T) {
241 unregisterAllDrivers()
242 Register("test", fdriver)
243 Register("invalid", Dummy{})
244 all := Drivers()
245 if len(all) < 2 || !slices.IsSorted(all) || !slices.Contains(all, "test") || !slices.Contains(all, "invalid") {
246 t.Fatalf("Drivers = %v, want sorted list with at least [invalid, test]", all)
247 }
248 }
249
250
251 var hookOpenErr struct {
252 sync.Mutex
253 fn func() error
254 }
255
256 func setHookOpenErr(fn func() error) {
257 hookOpenErr.Lock()
258 defer hookOpenErr.Unlock()
259 hookOpenErr.fn = fn
260 }
261
262
263
264
265
266
267
268 func (d *fakeDriver) Open(dsn string) (driver.Conn, error) {
269 hookOpenErr.Lock()
270 fn := hookOpenErr.fn
271 hookOpenErr.Unlock()
272 if fn != nil {
273 if err := fn(); err != nil {
274 return nil, err
275 }
276 }
277 parts := strings.Split(dsn, ";")
278 if len(parts) < 1 {
279 return nil, errors.New("fakedb: no database name")
280 }
281 name := parts[0]
282
283 db := d.getDB(name)
284
285 d.mu.Lock()
286 d.openCount++
287 d.mu.Unlock()
288 conn := &fakeConn{db: db}
289
290 if len(parts) >= 2 && parts[1] == "badConn" {
291 conn.bad = true
292 }
293 if d.waitCh != nil {
294 d.waitingCh <- struct{}{}
295 <-d.waitCh
296 d.waitCh = nil
297 d.waitingCh = nil
298 }
299 return conn, nil
300 }
301
302 func (d *fakeDriver) getDB(name string) *fakeDB {
303 d.mu.Lock()
304 defer d.mu.Unlock()
305 if d.dbs == nil {
306 d.dbs = make(map[string]*fakeDB)
307 }
308 db, ok := d.dbs[name]
309 if !ok {
310 db = &fakeDB{name: name}
311 d.dbs[name] = db
312 }
313 return db
314 }
315
316 func (db *fakeDB) wipe() {
317 db.mu.Lock()
318 defer db.mu.Unlock()
319 db.tables = nil
320 }
321
322 func (db *fakeDB) createTable(name string, columnNames, columnTypes []string) error {
323 db.mu.Lock()
324 defer db.mu.Unlock()
325 if db.tables == nil {
326 db.tables = make(map[string]*table)
327 }
328 if _, exist := db.tables[name]; exist {
329 return fmt.Errorf("fakedb: table %q already exists", name)
330 }
331 if len(columnNames) != len(columnTypes) {
332 return fmt.Errorf("fakedb: create table of %q len(names) != len(types): %d vs %d",
333 name, len(columnNames), len(columnTypes))
334 }
335 db.tables[name] = &table{colname: columnNames, coltype: columnTypes}
336 return nil
337 }
338
339
340 func (db *fakeDB) table(table string) (*table, bool) {
341 if db.tables == nil {
342 return nil, false
343 }
344 t, ok := db.tables[table]
345 return t, ok
346 }
347
348 func (db *fakeDB) columnType(table, column string) (typ string, ok bool) {
349 db.mu.Lock()
350 defer db.mu.Unlock()
351 t, ok := db.table(table)
352 if !ok {
353 return
354 }
355 if i := slices.Index(t.colname, column); i != -1 {
356 return t.coltype[i], true
357 }
358 return "", false
359 }
360
361 func (c *fakeConn) isBad() bool {
362 if c.stickyBad {
363 return true
364 } else if c.bad {
365 if c.db == nil {
366 return false
367 }
368
369 c.db.badConn = !c.db.badConn
370 return c.db.badConn
371 } else {
372 return false
373 }
374 }
375
376 func (c *fakeConn) isDirtyAndMark() bool {
377 if c.skipDirtySession {
378 return false
379 }
380 if c.currTx != nil {
381 c.dirtySession = true
382 return false
383 }
384 if c.dirtySession {
385 return true
386 }
387 c.dirtySession = true
388 return false
389 }
390
391 func (c *fakeConn) Begin() (driver.Tx, error) {
392 if c.isBad() {
393 return nil, fakeError{Wrapped: driver.ErrBadConn}
394 }
395 if c.currTx != nil {
396 return nil, errors.New("fakedb: already in a transaction")
397 }
398 c.touchMem()
399 c.currTx = &fakeTx{c: c}
400 return c.currTx, nil
401 }
402
403 var hookPostCloseConn struct {
404 sync.Mutex
405 fn func(*fakeConn, error)
406 }
407
408 func setHookpostCloseConn(fn func(*fakeConn, error)) {
409 hookPostCloseConn.Lock()
410 defer hookPostCloseConn.Unlock()
411 hookPostCloseConn.fn = fn
412 }
413
414 var testStrictClose *testing.T
415
416
417
418 func setStrictFakeConnClose(t *testing.T) {
419 testStrictClose = t
420 }
421
422 func (c *fakeConn) ResetSession(ctx context.Context) error {
423 c.dirtySession = false
424 c.currTx = nil
425 if c.isBad() {
426 return fakeError{Message: "Reset Session: bad conn", Wrapped: driver.ErrBadConn}
427 }
428 return nil
429 }
430
431 var _ driver.Validator = (*fakeConn)(nil)
432
433 func (c *fakeConn) IsValid() bool {
434 return !c.isBad()
435 }
436
437 func (c *fakeConn) Close() (err error) {
438 drv := fdriver.(*fakeDriver)
439 defer func() {
440 if err != nil && testStrictClose != nil {
441 testStrictClose.Errorf("failed to close a test fakeConn: %v", err)
442 }
443 hookPostCloseConn.Lock()
444 fn := hookPostCloseConn.fn
445 hookPostCloseConn.Unlock()
446 if fn != nil {
447 fn(c, err)
448 }
449 if err == nil {
450 drv.mu.Lock()
451 drv.closeCount++
452 drv.mu.Unlock()
453 }
454 }()
455 c.touchMem()
456 if c.currTx != nil {
457 return errors.New("fakedb: can't close fakeConn; in a Transaction")
458 }
459 if c.db == nil {
460 return errors.New("fakedb: can't close fakeConn; already closed")
461 }
462 if c.stmtsMade > c.stmtsClosed {
463 return errors.New("fakedb: can't close; dangling statement(s)")
464 }
465 c.db = nil
466 return nil
467 }
468
469 func checkSubsetTypes(allowAny bool, args []driver.NamedValue) error {
470 for _, arg := range args {
471 switch arg.Value.(type) {
472 case int64, float64, bool, nil, []byte, string, time.Time:
473 default:
474 if !allowAny {
475 return fmt.Errorf("fakedb: invalid argument ordinal %[1]d: %[2]v, type %[2]T", arg.Ordinal, arg.Value)
476 }
477 }
478 }
479 return nil
480 }
481
482 func (c *fakeConn) Exec(query string, args []driver.Value) (driver.Result, error) {
483
484 panic("ExecContext was not called.")
485 }
486
487 func (c *fakeConn) ExecContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Result, error) {
488
489
490
491
492 err := checkSubsetTypes(c.db.allowAny, args)
493 if err != nil {
494 return nil, err
495 }
496 return nil, driver.ErrSkip
497 }
498
499 func (c *fakeConn) Query(query string, args []driver.Value) (driver.Rows, error) {
500
501 panic("QueryContext was not called.")
502 }
503
504 func (c *fakeConn) QueryContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Rows, error) {
505
506
507
508
509 err := checkSubsetTypes(c.db.allowAny, args)
510 if err != nil {
511 return nil, err
512 }
513 return nil, driver.ErrSkip
514 }
515
516 func errf(msg string, args ...any) error {
517 return errors.New("fakedb: " + fmt.Sprintf(msg, args...))
518 }
519
520
521
522
523 func (c *fakeConn) prepareSelect(stmt *fakeStmt, parts []string) (*fakeStmt, error) {
524 if len(parts) != 3 {
525 stmt.Close()
526 return nil, errf("invalid SELECT syntax with %d parts; want 3", len(parts))
527 }
528 stmt.table = parts[0]
529
530 stmt.colName = strings.Split(parts[1], ",")
531 for n, colspec := range strings.Split(parts[2], ",") {
532 if colspec == "" {
533 continue
534 }
535 nameVal := strings.Split(colspec, "=")
536 if len(nameVal) != 2 {
537 stmt.Close()
538 return nil, errf("SELECT on table %q has invalid column spec of %q (index %d)", stmt.table, colspec, n)
539 }
540 column, value := nameVal[0], nameVal[1]
541 _, ok := c.db.columnType(stmt.table, column)
542 if !ok {
543 stmt.Close()
544 return nil, errf("SELECT on table %q references non-existent column %q", stmt.table, column)
545 }
546 if !strings.HasPrefix(value, "?") {
547 stmt.Close()
548 return nil, errf("SELECT on table %q has pre-bound value for where column %q; need a question mark",
549 stmt.table, column)
550 }
551 stmt.placeholders++
552 stmt.whereCol = append(stmt.whereCol, boundCol{Column: column, Placeholder: value, Ordinal: stmt.placeholders})
553 }
554 return stmt, nil
555 }
556
557
558 func (c *fakeConn) prepareCreate(stmt *fakeStmt, parts []string) (*fakeStmt, error) {
559 if len(parts) != 2 {
560 stmt.Close()
561 return nil, errf("invalid CREATE syntax with %d parts; want 2", len(parts))
562 }
563 stmt.table = parts[0]
564 for n, colspec := range strings.Split(parts[1], ",") {
565 nameType := strings.Split(colspec, "=")
566 if len(nameType) != 2 {
567 stmt.Close()
568 return nil, errf("CREATE table %q has invalid column spec of %q (index %d)", stmt.table, colspec, n)
569 }
570 stmt.colName = append(stmt.colName, nameType[0])
571 stmt.colType = append(stmt.colType, nameType[1])
572 }
573 return stmt, nil
574 }
575
576
577 func (c *fakeConn) prepareInsert(ctx context.Context, stmt *fakeStmt, parts []string) (*fakeStmt, error) {
578 if len(parts) != 2 {
579 stmt.Close()
580 return nil, errf("invalid INSERT syntax with %d parts; want 2", len(parts))
581 }
582 stmt.table = parts[0]
583 for n, colspec := range strings.Split(parts[1], ",") {
584 nameVal := strings.Split(colspec, "=")
585 if len(nameVal) != 2 {
586 stmt.Close()
587 return nil, errf("INSERT table %q has invalid column spec of %q (index %d)", stmt.table, colspec, n)
588 }
589 column, value := nameVal[0], nameVal[1]
590 ctype, ok := c.db.columnType(stmt.table, column)
591 if !ok {
592 stmt.Close()
593 return nil, errf("INSERT table %q references non-existent column %q", stmt.table, column)
594 }
595 stmt.colName = append(stmt.colName, column)
596
597 if !strings.HasPrefix(value, "?") {
598 var subsetVal any
599
600 switch ctype {
601 case "string":
602 subsetVal = []byte(value)
603 case "blob":
604 subsetVal = []byte(value)
605 case "int32":
606 i, err := strconv.Atoi(value)
607 if err != nil {
608 stmt.Close()
609 return nil, errf("invalid conversion to int32 from %q", value)
610 }
611 subsetVal = int64(i)
612 case "table":
613 c.skipDirtySession = true
614 vparts := strings.Split(value, "!")
615
616 substmt, err := c.PrepareContext(ctx, fmt.Sprintf("SELECT|%s|%s|", vparts[0], strings.Join(vparts[1:], ",")))
617 if err != nil {
618 return nil, err
619 }
620 cursor, err := (substmt.(driver.StmtQueryContext)).QueryContext(ctx, []driver.NamedValue{})
621 substmt.Close()
622 if err != nil {
623 return nil, err
624 }
625 subsetVal = cursor
626 default:
627 stmt.Close()
628 return nil, errf("unsupported conversion for pre-bound parameter %q to type %q", value, ctype)
629 }
630 stmt.colValue = append(stmt.colValue, subsetVal)
631 } else {
632 stmt.placeholders++
633 stmt.placeholderConverter = append(stmt.placeholderConverter, converterForType(ctype))
634 stmt.colValue = append(stmt.colValue, value)
635 }
636 }
637 return stmt, nil
638 }
639
640
641 var hookPrepareBadConn func() bool
642
643 func (c *fakeConn) Prepare(query string) (driver.Stmt, error) {
644 panic("use PrepareContext")
645 }
646
647 func (c *fakeConn) PrepareContext(ctx context.Context, query string) (driver.Stmt, error) {
648 c.numPrepare++
649 if c.db == nil {
650 panic("nil c.db; conn = " + fmt.Sprintf("%#v", c))
651 }
652
653 if c.stickyBad || (hookPrepareBadConn != nil && hookPrepareBadConn()) {
654 return nil, fakeError{Message: "Prepare: Sticky Bad", Wrapped: driver.ErrBadConn}
655 }
656
657 c.touchMem()
658 var firstStmt, prev *fakeStmt
659 for _, query := range strings.Split(query, ";") {
660 parts := strings.Split(query, "|")
661 if len(parts) < 1 {
662 return nil, errf("empty query")
663 }
664 stmt := &fakeStmt{q: query, c: c, memToucher: c}
665 if firstStmt == nil {
666 firstStmt = stmt
667 }
668 if len(parts) >= 3 {
669 switch parts[0] {
670 case "PANIC":
671 stmt.panic = parts[1]
672 parts = parts[2:]
673 case "WAIT":
674 wait, err := time.ParseDuration(parts[1])
675 if err != nil {
676 return nil, errf("expected section after WAIT to be a duration, got %q %v", parts[1], err)
677 }
678 parts = parts[2:]
679 stmt.wait = wait
680 }
681 }
682 cmd := parts[0]
683 stmt.cmd = cmd
684 parts = parts[1:]
685
686 if c.waiter != nil {
687 c.waiter(ctx)
688 if err := ctx.Err(); err != nil {
689 return nil, err
690 }
691 }
692
693 if stmt.wait > 0 {
694 wait := time.NewTimer(stmt.wait)
695 select {
696 case <-wait.C:
697 case <-ctx.Done():
698 wait.Stop()
699 return nil, ctx.Err()
700 }
701 }
702
703 c.incrStat(&c.stmtsMade)
704 var err error
705 switch cmd {
706 case "WIPE":
707
708 case "SELECT":
709 stmt, err = c.prepareSelect(stmt, parts)
710 case "CREATE":
711 stmt, err = c.prepareCreate(stmt, parts)
712 case "INSERT":
713 stmt, err = c.prepareInsert(ctx, stmt, parts)
714 case "NOSERT":
715
716
717 stmt, err = c.prepareInsert(ctx, stmt, parts)
718 default:
719 stmt.Close()
720 return nil, errf("unsupported command type %q", cmd)
721 }
722 if err != nil {
723 return nil, err
724 }
725 if prev != nil {
726 prev.next = stmt
727 }
728 prev = stmt
729 }
730 return firstStmt, nil
731 }
732
733 func (s *fakeStmt) ColumnConverter(idx int) driver.ValueConverter {
734 if s.panic == "ColumnConverter" {
735 panic(s.panic)
736 }
737 if len(s.placeholderConverter) == 0 {
738 return driver.DefaultParameterConverter
739 }
740 return s.placeholderConverter[idx]
741 }
742
743 func (s *fakeStmt) Close() error {
744 if s.panic == "Close" {
745 panic(s.panic)
746 }
747 if s.c == nil {
748 panic("nil conn in fakeStmt.Close")
749 }
750 if s.c.db == nil {
751 panic("in fakeStmt.Close, conn's db is nil (already closed)")
752 }
753 s.touchMem()
754 if !s.closed {
755 s.c.incrStat(&s.c.stmtsClosed)
756 s.closed = true
757 }
758 if s.next != nil {
759 s.next.Close()
760 }
761 return nil
762 }
763
764 var errClosed = errors.New("fakedb: statement has been closed")
765
766
767 var hookExecBadConn func() bool
768
769 func (s *fakeStmt) Exec(args []driver.Value) (driver.Result, error) {
770 panic("Using ExecContext")
771 }
772
773 var errFakeConnSessionDirty = errors.New("fakedb: session is dirty")
774
775 func (s *fakeStmt) ExecContext(ctx context.Context, args []driver.NamedValue) (driver.Result, error) {
776 if s.panic == "Exec" {
777 panic(s.panic)
778 }
779 if s.closed {
780 return nil, errClosed
781 }
782
783 if s.c.stickyBad || (hookExecBadConn != nil && hookExecBadConn()) {
784 return nil, fakeError{Message: "Exec: Sticky Bad", Wrapped: driver.ErrBadConn}
785 }
786 if s.c.isDirtyAndMark() {
787 return nil, errFakeConnSessionDirty
788 }
789
790 err := checkSubsetTypes(s.c.db.allowAny, args)
791 if err != nil {
792 return nil, err
793 }
794 s.touchMem()
795
796 if s.wait > 0 {
797 time.Sleep(s.wait)
798 }
799
800 select {
801 default:
802 case <-ctx.Done():
803 return nil, ctx.Err()
804 }
805
806 db := s.c.db
807 switch s.cmd {
808 case "WIPE":
809 db.wipe()
810 return driver.ResultNoRows, nil
811 case "CREATE":
812 if err := db.createTable(s.table, s.colName, s.colType); err != nil {
813 return nil, err
814 }
815 return driver.ResultNoRows, nil
816 case "INSERT":
817 return s.execInsert(args, true)
818 case "NOSERT":
819
820
821 return s.execInsert(args, false)
822 }
823 return nil, fmt.Errorf("fakedb: unimplemented statement Exec command type of %q", s.cmd)
824 }
825
826 func valueFromPlaceholderName(args []driver.NamedValue, name string) driver.Value {
827 for i := range args {
828 if args[i].Name == name {
829 return args[i].Value
830 }
831 }
832 return nil
833 }
834
835
836
837
838 func (s *fakeStmt) execInsert(args []driver.NamedValue, doInsert bool) (driver.Result, error) {
839 db := s.c.db
840 if len(args) != s.placeholders {
841 panic("error in pkg db; should only get here if size is correct")
842 }
843 db.mu.Lock()
844 t, ok := db.table(s.table)
845 db.mu.Unlock()
846 if !ok {
847 return nil, fmt.Errorf("fakedb: table %q doesn't exist", s.table)
848 }
849
850 t.mu.Lock()
851 defer t.mu.Unlock()
852
853 var cols []any
854 if doInsert {
855 cols = make([]any, len(t.colname))
856 }
857 argPos := 0
858 for n, colname := range s.colName {
859 colidx := t.columnIndex(colname)
860 if colidx == -1 {
861 return nil, fmt.Errorf("fakedb: column %q doesn't exist or dropped since prepared statement was created", colname)
862 }
863 var val any
864 if strvalue, ok := s.colValue[n].(string); ok && strings.HasPrefix(strvalue, "?") {
865 if strvalue == "?" {
866 val = args[argPos].Value
867 } else {
868
869 if v := valueFromPlaceholderName(args, strvalue[1:]); v != nil {
870 val = v
871 }
872 }
873 argPos++
874 } else {
875 val = s.colValue[n]
876 }
877 if doInsert {
878 cols[colidx] = val
879 }
880 }
881
882 if doInsert {
883 t.rows = append(t.rows, &row{cols: cols})
884 }
885 return driver.RowsAffected(1), nil
886 }
887
888
889 var hookQueryBadConn func() bool
890
891 func (s *fakeStmt) Query(args []driver.Value) (driver.Rows, error) {
892 panic("Use QueryContext")
893 }
894
895 func (s *fakeStmt) QueryContext(ctx context.Context, args []driver.NamedValue) (driver.Rows, error) {
896 if s.panic == "Query" {
897 panic(s.panic)
898 }
899 if s.closed {
900 return nil, errClosed
901 }
902
903 if s.c.stickyBad || (hookQueryBadConn != nil && hookQueryBadConn()) {
904 return nil, fakeError{Message: "Query: Sticky Bad", Wrapped: driver.ErrBadConn}
905 }
906 if s.c.isDirtyAndMark() {
907 return nil, errFakeConnSessionDirty
908 }
909
910 err := checkSubsetTypes(s.c.db.allowAny, args)
911 if err != nil {
912 return nil, err
913 }
914
915 s.touchMem()
916 db := s.c.db
917 if len(args) != s.placeholders {
918 panic("error in pkg db; should only get here if size is correct")
919 }
920
921 setMRows := make([][]*row, 0, 1)
922 setColumns := make([][]string, 0, 1)
923 setColType := make([][]string, 0, 1)
924
925 for {
926 db.mu.Lock()
927 t, ok := db.table(s.table)
928 db.mu.Unlock()
929 if !ok {
930 return nil, fmt.Errorf("fakedb: table %q doesn't exist", s.table)
931 }
932
933 if s.table == "magicquery" {
934 if len(s.whereCol) == 2 && s.whereCol[0].Column == "op" && s.whereCol[1].Column == "millis" {
935 if args[0].Value == "sleep" {
936 time.Sleep(time.Duration(args[1].Value.(int64)) * time.Millisecond)
937 }
938 }
939 }
940 if s.table == "tx_status" && s.colName[0] == "tx_status" {
941 txStatus := "autocommit"
942 if s.c.currTx != nil {
943 txStatus = "transaction"
944 }
945 cursor := &rowsCursor{
946 db: s.c.db,
947 parentMem: s.c,
948 posRow: -1,
949 rows: [][]*row{
950 {
951 {
952 cols: []any{
953 txStatus,
954 },
955 },
956 },
957 },
958 cols: [][]string{
959 {
960 "tx_status",
961 },
962 },
963 colType: [][]string{
964 {
965 "string",
966 },
967 },
968 errPos: -1,
969 }
970 return cursor, nil
971 }
972
973 t.mu.Lock()
974
975 colIdx := make(map[string]int)
976 for _, name := range s.colName {
977 idx := t.columnIndex(name)
978 if idx == -1 {
979 t.mu.Unlock()
980 return nil, fmt.Errorf("fakedb: unknown column name %q", name)
981 }
982 colIdx[name] = idx
983 }
984
985 mrows := []*row{}
986 rows:
987 for _, trow := range t.rows {
988
989
990
991 for _, wcol := range s.whereCol {
992 idx := t.columnIndex(wcol.Column)
993 if idx == -1 {
994 t.mu.Unlock()
995 return nil, fmt.Errorf("fakedb: invalid where clause column %v", wcol)
996 }
997 tcol := trow.cols[idx]
998 if bs, ok := tcol.([]byte); ok {
999
1000 tcol = string(bs)
1001 }
1002 var argValue any
1003 if wcol.Placeholder == "?" {
1004 argValue = args[wcol.Ordinal-1].Value
1005 } else {
1006 if v := valueFromPlaceholderName(args, wcol.Placeholder[1:]); v != nil {
1007 argValue = v
1008 }
1009 }
1010 if fmt.Sprintf("%v", tcol) != fmt.Sprintf("%v", argValue) {
1011 continue rows
1012 }
1013 }
1014 mrow := &row{cols: make([]any, len(s.colName))}
1015 for seli, name := range s.colName {
1016 mrow.cols[seli] = trow.cols[colIdx[name]]
1017 }
1018 mrows = append(mrows, mrow)
1019 }
1020
1021 var colType []string
1022 for _, column := range s.colName {
1023 colType = append(colType, t.coltype[t.columnIndex(column)])
1024 }
1025
1026 t.mu.Unlock()
1027
1028 setMRows = append(setMRows, mrows)
1029 setColumns = append(setColumns, s.colName)
1030 setColType = append(setColType, colType)
1031
1032 if s.next == nil {
1033 break
1034 }
1035 s = s.next
1036 }
1037
1038 cursor := &rowsCursor{
1039 db: s.c.db,
1040 parentMem: s.c,
1041 posRow: -1,
1042 rows: setMRows,
1043 cols: setColumns,
1044 colType: setColType,
1045 errPos: -1,
1046 }
1047 return cursor, nil
1048 }
1049
1050 func (s *fakeStmt) NumInput() int {
1051 if s.panic == "NumInput" {
1052 panic(s.panic)
1053 }
1054 return s.placeholders
1055 }
1056
1057
1058 var hookCommitBadConn func() bool
1059
1060 func (tx *fakeTx) Commit() error {
1061 tx.c.currTx = nil
1062 if hookCommitBadConn != nil && hookCommitBadConn() {
1063 return fakeError{Message: "Commit: Hook Bad Conn", Wrapped: driver.ErrBadConn}
1064 }
1065 tx.c.touchMem()
1066 return nil
1067 }
1068
1069
1070 var hookRollbackBadConn func() bool
1071
1072 func (tx *fakeTx) Rollback() error {
1073 tx.c.currTx = nil
1074 if hookRollbackBadConn != nil && hookRollbackBadConn() {
1075 return fakeError{Message: "Rollback: Hook Bad Conn", Wrapped: driver.ErrBadConn}
1076 }
1077 tx.c.touchMem()
1078 return nil
1079 }
1080
1081 type rowsCursor struct {
1082 db *fakeDB
1083 parentMem memToucher
1084 cols [][]string
1085 colType [][]string
1086 posSet int
1087 posRow int
1088 rows [][]*row
1089 closed bool
1090
1091
1092 errPos int
1093 err error
1094
1095
1096
1097 driverOwnedMemory [][]byte
1098
1099
1100
1101
1102
1103 line int64
1104
1105
1106 closeErr error
1107 }
1108
1109 func (rc *rowsCursor) touchMem() {
1110 rc.parentMem.touchMem()
1111 rc.line++
1112 }
1113
1114 func (rc *rowsCursor) invalidateDriverOwnedMemory() {
1115 for _, buf := range rc.driverOwnedMemory {
1116 for i := range buf {
1117 buf[i] = 'x'
1118 }
1119 }
1120 rc.driverOwnedMemory = nil
1121 }
1122
1123 func (rc *rowsCursor) Close() error {
1124 rc.touchMem()
1125 rc.parentMem.touchMem()
1126 rc.invalidateDriverOwnedMemory()
1127 rc.closed = true
1128 return rc.closeErr
1129 }
1130
1131 func (rc *rowsCursor) Columns() []string {
1132 return rc.cols[rc.posSet]
1133 }
1134
1135 func (rc *rowsCursor) ColumnTypeScanType(index int) reflect.Type {
1136 return colTypeToReflectType(rc.colType[rc.posSet][index])
1137 }
1138
1139 var rowsCursorNextHook func(dest []driver.Value) error
1140
1141 func (rc *rowsCursor) Next(dest []driver.Value) error {
1142 if rowsCursorNextHook != nil {
1143 return rowsCursorNextHook(dest)
1144 }
1145
1146 if rc.closed {
1147 return errors.New("fakedb: cursor is closed")
1148 }
1149 rc.touchMem()
1150 rc.posRow++
1151 if rc.posRow == rc.errPos {
1152 return rc.err
1153 }
1154 if rc.posRow >= len(rc.rows[rc.posSet]) {
1155 return io.EOF
1156 }
1157
1158 rc.invalidateDriverOwnedMemory()
1159 for i, v := range rc.rows[rc.posSet][rc.posRow].cols {
1160
1161
1162
1163
1164
1165
1166 if bs, ok := v.([]byte); ok {
1167
1168 bs = bytes.Clone(bs)
1169 rc.driverOwnedMemory = append(rc.driverOwnedMemory, bs)
1170 v = bs
1171 }
1172 dest[i] = v
1173 }
1174 return nil
1175 }
1176
1177 func (rc *rowsCursor) HasNextResultSet() bool {
1178 rc.touchMem()
1179 return rc.posSet < len(rc.rows)-1
1180 }
1181
1182 func (rc *rowsCursor) NextResultSet() error {
1183 rc.touchMem()
1184 if rc.HasNextResultSet() {
1185 rc.posSet++
1186 rc.posRow = -1
1187 return nil
1188 }
1189 return io.EOF
1190 }
1191
1192
1193
1194
1195
1196
1197
1198 type fakeDriverString struct{}
1199
1200 func (fakeDriverString) ConvertValue(v any) (driver.Value, error) {
1201 switch c := v.(type) {
1202 case string, []byte:
1203 return v, nil
1204 case *string:
1205 if c == nil {
1206 return nil, nil
1207 }
1208 return *c, nil
1209 }
1210 return fmt.Sprintf("%v", v), nil
1211 }
1212
1213 type anyTypeConverter struct{}
1214
1215 func (anyTypeConverter) ConvertValue(v any) (driver.Value, error) {
1216 return v, nil
1217 }
1218
1219 func converterForType(typ string) driver.ValueConverter {
1220 switch typ {
1221 case "bool":
1222 return driver.Bool
1223 case "nullbool":
1224 return driver.Null{Converter: driver.Bool}
1225 case "byte", "int16":
1226 return driver.NotNull{Converter: driver.DefaultParameterConverter}
1227 case "int32":
1228 return driver.Int32
1229 case "nullbyte", "nullint32", "nullint16":
1230 return driver.Null{Converter: driver.DefaultParameterConverter}
1231 case "string":
1232 return driver.NotNull{Converter: fakeDriverString{}}
1233 case "nullstring":
1234 return driver.Null{Converter: fakeDriverString{}}
1235 case "int64":
1236
1237 return driver.NotNull{Converter: driver.DefaultParameterConverter}
1238 case "nullint64":
1239
1240 return driver.Null{Converter: driver.DefaultParameterConverter}
1241 case "float64":
1242
1243 return driver.NotNull{Converter: driver.DefaultParameterConverter}
1244 case "nullfloat64":
1245
1246 return driver.Null{Converter: driver.DefaultParameterConverter}
1247 case "datetime":
1248 return driver.NotNull{Converter: driver.DefaultParameterConverter}
1249 case "nulldatetime":
1250 return driver.Null{Converter: driver.DefaultParameterConverter}
1251 case "uuid":
1252 return driver.NotNull{Converter: driver.DefaultParameterConverter}
1253 case "nulluuid":
1254 return driver.Null{Converter: driver.DefaultParameterConverter}
1255 case "nulltable":
1256 return driver.Null{Converter: driver.DefaultParameterConverter}
1257 case "any":
1258 return anyTypeConverter{}
1259 }
1260 panic("invalid fakedb column type of " + typ)
1261 }
1262
1263 func colTypeToReflectType(typ string) reflect.Type {
1264 switch typ {
1265 case "bool":
1266 return reflect.TypeFor[bool]()
1267 case "nullbool":
1268 return reflect.TypeFor[NullBool]()
1269 case "int16":
1270 return reflect.TypeFor[int16]()
1271 case "nullint16":
1272 return reflect.TypeFor[NullInt16]()
1273 case "int32":
1274 return reflect.TypeFor[int32]()
1275 case "nullint32":
1276 return reflect.TypeFor[NullInt32]()
1277 case "string":
1278 return reflect.TypeFor[string]()
1279 case "nullstring":
1280 return reflect.TypeFor[NullString]()
1281 case "int64":
1282 return reflect.TypeFor[int64]()
1283 case "nullint64":
1284 return reflect.TypeFor[NullInt64]()
1285 case "float64":
1286 return reflect.TypeFor[float64]()
1287 case "nullfloat64":
1288 return reflect.TypeFor[NullFloat64]()
1289 case "datetime":
1290 return reflect.TypeFor[time.Time]()
1291 case "any":
1292 return reflect.TypeFor[any]()
1293 }
1294 panic("invalid fakedb column type of " + typ)
1295 }
1296
View as plain text