mirror of https://github.com/mattn/go-sqlite3.git
410 lines
8.5 KiB
Go
410 lines
8.5 KiB
Go
package sqlite3_test
|
|
|
|
import (
|
|
"database/sql"
|
|
"fmt"
|
|
"math/rand"
|
|
"regexp"
|
|
"strconv"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
type Dialect int
|
|
|
|
const (
|
|
SQLITE Dialect = iota
|
|
POSTGRESQL
|
|
MYSQL
|
|
)
|
|
|
|
type DB struct {
|
|
*testing.T
|
|
*sql.DB
|
|
dialect Dialect
|
|
once sync.Once
|
|
}
|
|
|
|
var db *DB
|
|
|
|
// the following tables will be created and dropped during the test
|
|
var testTables = []string{"foo", "bar", "t", "bench"}
|
|
|
|
var tests = []testing.InternalTest{
|
|
{"TestBlobs", TestBlobs},
|
|
{"TestManyQueryRow", TestManyQueryRow},
|
|
{"TestTxQuery", TestTxQuery},
|
|
{"TestPreparedStmt", TestPreparedStmt},
|
|
}
|
|
|
|
var benchmarks = []testing.InternalBenchmark{
|
|
{"BenchmarkExec", BenchmarkExec},
|
|
{"BenchmarkQuery", BenchmarkQuery},
|
|
{"BenchmarkParams", BenchmarkParams},
|
|
{"BenchmarkStmt", BenchmarkStmt},
|
|
{"BenchmarkRows", BenchmarkRows},
|
|
{"BenchmarkStmtRows", BenchmarkStmtRows},
|
|
}
|
|
|
|
// RunTests runs the SQL test suite
|
|
func RunTests(t *testing.T, d *sql.DB, dialect Dialect) {
|
|
db = &DB{t, d, dialect, sync.Once{}}
|
|
testing.RunTests(func(string, string) (bool, error) { return true, nil }, tests)
|
|
|
|
if !testing.Short() {
|
|
for _, b := range benchmarks {
|
|
fmt.Printf("%-20s", b.Name)
|
|
r := testing.Benchmark(b.F)
|
|
fmt.Printf("%10d %10.0f req/s\n", r.N, float64(r.N)/r.T.Seconds())
|
|
}
|
|
}
|
|
db.tearDown()
|
|
}
|
|
|
|
func (db *DB) mustExec(sql string, args ...interface{}) sql.Result {
|
|
res, err := db.Exec(sql, args...)
|
|
if err != nil {
|
|
db.Fatalf("Error running %q: %v", sql, err)
|
|
}
|
|
return res
|
|
}
|
|
|
|
func (db *DB) tearDown() {
|
|
for _, tbl := range testTables {
|
|
switch db.dialect {
|
|
case SQLITE:
|
|
db.mustExec("drop table if exists " + tbl)
|
|
case MYSQL, POSTGRESQL:
|
|
db.mustExec("drop table if exists " + tbl)
|
|
default:
|
|
db.Fatal("unkown dialect")
|
|
}
|
|
}
|
|
}
|
|
|
|
// q replaces ? parameters if needed
|
|
func (db *DB) q(sql string) string {
|
|
switch db.dialect {
|
|
case POSTGRESQL: // repace with $1, $2, ..
|
|
qrx := regexp.MustCompile(`\?`)
|
|
n := 0
|
|
return qrx.ReplaceAllStringFunc(sql, func(string) string {
|
|
n++
|
|
return "$" + strconv.Itoa(n)
|
|
})
|
|
}
|
|
return sql
|
|
}
|
|
|
|
func (db *DB) blobType(size int) string {
|
|
switch db.dialect {
|
|
case SQLITE:
|
|
return fmt.Sprintf("blob[%d]", size)
|
|
case POSTGRESQL:
|
|
return "bytea"
|
|
case MYSQL:
|
|
return fmt.Sprintf("VARBINARY(%d)", size)
|
|
}
|
|
panic("unkown dialect")
|
|
}
|
|
|
|
func (db *DB) serialPK() string {
|
|
switch db.dialect {
|
|
case SQLITE:
|
|
return "integer primary key autoincrement"
|
|
case POSTGRESQL:
|
|
return "serial primary key"
|
|
case MYSQL:
|
|
return "integer primary key auto_increment"
|
|
}
|
|
panic("unkown dialect")
|
|
}
|
|
|
|
func (db *DB) now() string {
|
|
switch db.dialect {
|
|
case SQLITE:
|
|
return "datetime('now')"
|
|
case POSTGRESQL:
|
|
return "now()"
|
|
case MYSQL:
|
|
return "now()"
|
|
}
|
|
panic("unkown dialect")
|
|
}
|
|
|
|
func makeBench() {
|
|
if _, err := db.Exec("create table bench (n varchar(32), i integer, d double, s varchar(32), t datetime)"); err != nil {
|
|
panic(err)
|
|
}
|
|
st, err := db.Prepare("insert into bench values (?, ?, ?, ?, ?)")
|
|
if err != nil {
|
|
panic(err)
|
|
}
|
|
defer st.Close()
|
|
for i := 0; i < 100; i++ {
|
|
if _, err = st.Exec(nil, i, float64(i), fmt.Sprintf("%d", i), time.Now()); err != nil {
|
|
panic(err)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestResult(t *testing.T) {
|
|
db.tearDown()
|
|
db.mustExec("create temporary table test (id " + db.serialPK() + ", name varchar(10))")
|
|
|
|
for i := 1; i < 3; i++ {
|
|
r := db.mustExec(db.q("insert into test (name) values (?)"), fmt.Sprintf("row %d", i))
|
|
n, err := r.RowsAffected()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if n != 1 {
|
|
t.Errorf("got %v, want %v", n, 1)
|
|
}
|
|
n, err = r.LastInsertId()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if n != int64(i) {
|
|
t.Errorf("got %v, want %v", n, i)
|
|
}
|
|
}
|
|
if _, err := db.Exec("error!"); err == nil {
|
|
t.Fatalf("expected error")
|
|
}
|
|
}
|
|
|
|