Source file src/database/sql/fakedb_test.go

     1  // Copyright 2011 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     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  // fakeDriver is a fake database that implements Go's driver.Driver
    24  // interface, just for testing.
    25  //
    26  // It speaks a query language that's semantically similar to but
    27  // syntactically different and simpler than SQL.  The syntax is as
    28  // follows:
    29  //
    30  //	WIPE
    31  //	CREATE|<tablename>|<col>=<type>,<col>=<type>,...
    32  //	  where types are: "string", [u]int{8,16,32,64}, "bool"
    33  //	INSERT|<tablename>|col=val,col2=val2,col3=?
    34  //	SELECT|<tablename>|projectcol1,projectcol2|filtercol=?,filtercol2=?
    35  //	SELECT|<tablename>|projectcol1,projectcol2|filtercol=?param1,filtercol2=?param2
    36  //
    37  // Any of these can be preceded by PANIC|<method>|, to cause the
    38  // named method on fakeStmt to panic.
    39  //
    40  // Any of these can be proceeded by WAIT|<duration>|, to cause the
    41  // named method on fakeStmt to sleep for the specified duration.
    42  //
    43  // Multiple of these can be combined when separated with a semicolon.
    44  //
    45  // When opening a fakeDriver's database, it starts empty with no
    46  // tables. All tables and data are stored in memory only.
    47  type fakeDriver struct {
    48  	mu         sync.Mutex // guards 3 following fields
    49  	openCount  int        // conn opens
    50  	closeCount int        // conn closes
    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 // must be same size as its table colname + coltype
   149  }
   150  
   151  type memToucher interface {
   152  	// touchMem reads & writes some memory, to help find data races.
   153  	touchMem()
   154  }
   155  
   156  type fakeConn struct {
   157  	db *fakeDB // where to return ourselves to
   158  
   159  	currTx *fakeTx
   160  
   161  	// Every operation writes to line to enable the race detector
   162  	// check for data races.
   163  	line int64
   164  
   165  	// Stats for tests:
   166  	mu          sync.Mutex
   167  	stmtsMade   int
   168  	stmtsClosed int
   169  	numPrepare  int
   170  
   171  	// bad connection tests; see isBad()
   172  	bad       bool
   173  	stickyBad bool
   174  
   175  	skipDirtySession bool // tests that use Conn should set this to true.
   176  
   177  	// dirtySession tests ResetSession, true if a query has executed
   178  	// until ResetSession is called.
   179  	dirtySession bool
   180  
   181  	// The waiter is called before each query. May be used in place of the "WAIT"
   182  	// directive.
   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 // just for debugging
   210  
   211  	cmd   string
   212  	table string
   213  	panic string
   214  	wait  time.Duration
   215  
   216  	next *fakeStmt // used for returning multiple results.
   217  
   218  	closed bool
   219  
   220  	colName      []string // used by CREATE, INSERT, SELECT (selected columns)
   221  	colType      []string // used by CREATE
   222  	colValue     []any    // used by INSERT (mix of strings and "?" for bound params)
   223  	placeholders int      // used by INSERT/SELECT: number of ? params
   224  
   225  	whereCol []boundCol // used by SELECT (all placeholders)
   226  
   227  	placeholderConverter []driver.ValueConverter // used by INSERT
   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  // hook to simulate connection failures
   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  // Supports dsn forms:
   263  //
   264  //	<dbname>
   265  //	<dbname>;<opts>  (only currently supported option is `badConn`,
   266  //	                  which causes driver.ErrBadConn to be returned on
   267  //	                  every other conn.Begin())
   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  // must be called with db.mu lock held
   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  		// alternate between bad conn and not bad conn
   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  // setStrictFakeConnClose sets the t to Errorf on when fakeConn.Close
   417  // fails to close. If nil, the check is disabled.
   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  	// Ensure that ExecContext is called if available.
   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  	// This is an optional interface, but it's implemented here
   489  	// just to check that all the args are of the proper types.
   490  	// ErrSkip is returned so the caller acts as if we didn't
   491  	// implement this at all.
   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  	// Ensure that ExecContext is called if available.
   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  	// This is an optional interface, but it's implemented here
   506  	// just to check that all the args are of the proper types.
   507  	// ErrSkip is returned so the caller acts as if we didn't
   508  	// implement this at all.
   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  // parts are table|selectCol1,selectCol2|whereCol=?,whereCol2=?
   521  // (note that where columns must always contain ? marks,
   522  // just a limitation for fakedb)
   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  // parts are table|col=type,col2=type2
   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  // parts are table|col=?,col2=val
   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  			// Convert to driver subset type
   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) // int64 is a subset type, but not int32
   612  			case "table": // For testing cursor reads.
   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  // hook to simulate broken connections
   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  			// Nothing
   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  			// Do all the prep-work like for an INSERT but don't actually insert the row.
   716  			// Used for some of the concurrent tests.
   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  // hook to simulate broken connections
   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  		// Do all the prep-work like for an INSERT but don't actually insert the row.
   820  		// Used for some of the concurrent tests.
   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  // When doInsert is true, add the row to the table.
   836  // When doInsert is false do prep-work and error checking, but don't
   837  // actually add the row to the table.
   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  				// Assign value from argument placeholder name.
   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  // hook to simulate broken connections
   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) // select column name -> column index in table
   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  			// Process the where clause, skipping non-match rows. This is lazy
   989  			// and just uses fmt.Sprintf("%v") to test equality. Good enough
   990  			// for test code.
   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  					// lazy hack to avoid sprintf %v on a []byte
  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  // hook to simulate broken connections
  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  // hook to simulate broken connections
  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  	// errPos and err are for making Next return early with error.
  1092  	errPos int
  1093  	err    error
  1094  
  1095  	// Data returned to clients.
  1096  	// We clone and stash it here so it can be invalidated by Close and Next.
  1097  	driverOwnedMemory [][]byte
  1098  
  1099  	// Every operation writes to line to enable the race detector
  1100  	// check for data races.
  1101  	// This is separate from the fakeConn.line to allow for drivers that
  1102  	// can start multiple queries on the same transaction at the same time.
  1103  	line int64
  1104  
  1105  	// closeErr is returned when rowsCursor.Close
  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 // per interface spec
  1156  	}
  1157  	// Corrupt any previously returned bytes.
  1158  	rc.invalidateDriverOwnedMemory()
  1159  	for i, v := range rc.rows[rc.posSet][rc.posRow].cols {
  1160  		// TODO(bradfitz): convert to subset types? naah, I
  1161  		// think the subset types should only be input to
  1162  		// driver, but the sql package should be able to handle
  1163  		// a wider range of types coming out of drivers. all
  1164  		// for ease of drivers, and to prevent drivers from
  1165  		// messing up conversions or doing them differently.
  1166  		if bs, ok := v.([]byte); ok {
  1167  			// Clone []bytes and stash for later invalidation.
  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 // Per interface spec.
  1190  }
  1191  
  1192  // fakeDriverString is like driver.String, but indirects pointers like
  1193  // DefaultValueConverter.
  1194  //
  1195  // This could be surprising behavior to retroactively apply to
  1196  // driver.String now that Go1 is out, but this is convenient for
  1197  // our TestPointerParamsAndScans.
  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  		// TODO(coopernurse): add type-specific converter
  1237  		return driver.NotNull{Converter: driver.DefaultParameterConverter}
  1238  	case "nullint64":
  1239  		// TODO(coopernurse): add type-specific converter
  1240  		return driver.Null{Converter: driver.DefaultParameterConverter}
  1241  	case "float64":
  1242  		// TODO(coopernurse): add type-specific converter
  1243  		return driver.NotNull{Converter: driver.DefaultParameterConverter}
  1244  	case "nullfloat64":
  1245  		// TODO(coopernurse): add type-specific converter
  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