mirror of
https://github.com/wahyd4/go-mocket.git
synced 2026-08-09 04:46:54 +10:00
First commit.
This commit is contained in:
+17
@@ -0,0 +1,17 @@
|
|||||||
|
# Created by .ignore support plugin (hsz.mobi)
|
||||||
|
### Go template
|
||||||
|
# Binaries for programs and plugins
|
||||||
|
*.exe
|
||||||
|
*.dll
|
||||||
|
*.so
|
||||||
|
*.dylib
|
||||||
|
|
||||||
|
# Test binary, build with `go test -c`
|
||||||
|
*.test
|
||||||
|
|
||||||
|
# Output of the go coverage tool, specifically when used with LiteIDE
|
||||||
|
*.out
|
||||||
|
|
||||||
|
# Project-local glide cache, RE: https://github.com/Masterminds/glide/issues/736
|
||||||
|
.glide/
|
||||||
|
|
||||||
@@ -0,0 +1,84 @@
|
|||||||
|
###Go-Mocket
|
||||||
|
|
||||||
|
Go-Mocket is library inspired by [DATA-DOG/go-sqlmock](https://github.com/DATA-DOG/go-sqlmock)
|
||||||
|
As inspiration library it is implementation of [sql/driver](https://godoc.org/database/sql/driver) interface but at the same time follows different approaches and has only similar API.
|
||||||
|
This library helps to mock any DB connection also with [jinzhu/gorm](https://github.com/jinzhu/gorm) and it was main goal to create it
|
||||||
|
|
||||||
|
List of features in the library:
|
||||||
|
|
||||||
|
* Mock `INSERT`, `UPDATE`, `SELECT`, `DELETE`
|
||||||
|
* Support of transactions
|
||||||
|
* 2 API's to use - `chaining` and via specifying whole mock object
|
||||||
|
* Matching by prepared statements arguments
|
||||||
|
* You will not require to change anything inside you code to start using this library
|
||||||
|
* Ability to trigger exceptions
|
||||||
|
* Attach callbacks to mocked response to add additional check or modify response
|
||||||
|
|
||||||
|
**NOTE** Please be aware that driver catches SQL without DB specifics. Generating of queries is done by *sql* package
|
||||||
|
|
||||||
|
####Install
|
||||||
|
```
|
||||||
|
go get github.com/selvatico/go-mocket
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Usage
|
||||||
|
There are two possible ways to use `mocket`:
|
||||||
|
|
||||||
|
* Chaining API
|
||||||
|
* Specifying `FakeResponse` object with all fields manually. Could be useful for
|
||||||
|
|
||||||
|
##### Enabling driver
|
||||||
|
|
||||||
|
Somewhere in you code to setup a tests
|
||||||
|
```
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
mocket "github.com/selvatico/go-mocket"
|
||||||
|
"github.com/jinzhu/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
func SetupTests() {
|
||||||
|
sql.Register("fake_test", mocket.FakeDriver{})
|
||||||
|
// GORM
|
||||||
|
db, err := gorm.Open("fake_test", "connection_string") // Could be any connection string
|
||||||
|
app.DB = db
|
||||||
|
|
||||||
|
// Regular sql package usage
|
||||||
|
db, err := sql.Open(driver, source)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
Now if use singleton instance of DB it will use everywhere mocked connection.
|
||||||
|
|
||||||
|
##### Chain usage
|
||||||
|
###### Example of mocking by pattern
|
||||||
|
|
||||||
|
```go
|
||||||
|
import mocket "github.com/selvatico/go-mocket"
|
||||||
|
import "net/http/httptest"
|
||||||
|
|
||||||
|
func TestHandler(t *testing.T) {
|
||||||
|
request := httptest.NewRequest("POST", "/application", postBody)
|
||||||
|
recorder := httptest.NewRecorder()
|
||||||
|
|
||||||
|
GlobalMock := mocket.Catcher
|
||||||
|
GlobalMock.Logging = true // log mocket behavior
|
||||||
|
|
||||||
|
commonReply := []map[string]interface{}{{"id": "2", "field": "value"}}
|
||||||
|
// Mock only by query pattern
|
||||||
|
GlobalMock.NewMock().WithQuery(`"campaigns".name IS NULL AND (("uuid" = test_uuid))`).WithReply(commonReply)
|
||||||
|
Post(recorder, request) // call handler
|
||||||
|
|
||||||
|
r := recorder.Result()
|
||||||
|
body, _ := ioutil.ReadAll(r.Body)
|
||||||
|
|
||||||
|
// some assertion about results
|
||||||
|
//...
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
```
|
||||||
|
|
||||||
|
|
||||||
|
Documentation in progress....
|
||||||
|
|
||||||
|
|
||||||
@@ -0,0 +1,66 @@
|
|||||||
|
package go_mocket
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql/driver"
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
type FakeConn struct {
|
||||||
|
db *FakeDB
|
||||||
|
currTx *FakeTx // Transaction pointer
|
||||||
|
mu sync.Mutex
|
||||||
|
bad bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *FakeConn) isBad() bool {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *FakeConn) Begin() (driver.Tx, error) {
|
||||||
|
if c.isBad() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
if c.currTx != nil {
|
||||||
|
return nil, errors.New("already in a transaction")
|
||||||
|
}
|
||||||
|
c.currTx = &FakeTx{c: c}
|
||||||
|
return c.currTx, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *FakeConn) Close() (err error) {
|
||||||
|
c.db = nil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *FakeConn) Exec(query string, args []driver.Value) (driver.Result, error) {
|
||||||
|
panic("ExecContext was not called.")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *FakeConn) ExecContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Result, error) {
|
||||||
|
return nil, driver.ErrSkip
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *FakeConn) Query(query string, args []driver.Value) (driver.Rows, error) {
|
||||||
|
panic("QueryContext was not called.")
|
||||||
|
}
|
||||||
|
|
||||||
|
// We do
|
||||||
|
func (c *FakeConn) QueryContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Rows, error) {
|
||||||
|
return nil, driver.ErrSkip
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should not be called
|
||||||
|
func (c *FakeConn) Prepare(query string) (driver.Stmt, error) {
|
||||||
|
panic("use Prepare")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *FakeConn) PrepareContext(ctx context.Context, query string) (driver.Stmt, error) {
|
||||||
|
var firstStmt = &FakeStmt{q: query, connection: c} // Create statement
|
||||||
|
firstStmt.placeholders = len(strings.Split(query, "?")) - 1 // Checking how many placeholders do we have
|
||||||
|
queryParts := strings.Split(query, " ") // By First statement define the query type
|
||||||
|
firstStmt.command = strings.ToUpper(queryParts[0])
|
||||||
|
return firstStmt, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,59 @@
|
|||||||
|
package go_mocket
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql/driver"
|
||||||
|
"log"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
var _ = log.Printf
|
||||||
|
|
||||||
|
type FakeDriver struct {
|
||||||
|
mu sync.Mutex // guards 3 following fields
|
||||||
|
openCount int // conn opens
|
||||||
|
closeCount int // conn closes
|
||||||
|
waitCh chan struct{}
|
||||||
|
waitingCh chan struct{}
|
||||||
|
dbs map[string]*FakeDB
|
||||||
|
}
|
||||||
|
|
||||||
|
type FakeDB struct {
|
||||||
|
name string
|
||||||
|
mu sync.Mutex
|
||||||
|
tables map[string]*table
|
||||||
|
badConn bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type table struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
colname []string
|
||||||
|
coltype []string
|
||||||
|
rows []*row
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *table) columnIndex(name string) int {
|
||||||
|
for n, name := range t.colname {
|
||||||
|
if name == name {
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d FakeDriver) Open(database string) (driver.Conn, error) {
|
||||||
|
return &FakeConn{db: d.getDB(database)}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *FakeDriver) getDB(name string) *FakeDB {
|
||||||
|
d.mu.Lock()
|
||||||
|
defer d.mu.Unlock()
|
||||||
|
if d.dbs == nil {
|
||||||
|
d.dbs = make(map[string]*FakeDB)
|
||||||
|
}
|
||||||
|
db, ok := d.dbs[name]
|
||||||
|
if !ok {
|
||||||
|
db = &FakeDB{name: name}
|
||||||
|
d.dbs[name] = db
|
||||||
|
}
|
||||||
|
return db
|
||||||
|
}
|
||||||
+165
@@ -0,0 +1,165 @@
|
|||||||
|
package go_mocket
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql/driver"
|
||||||
|
"fmt"
|
||||||
|
"log"
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
var Catcher *MockCatcher
|
||||||
|
|
||||||
|
type MockCatcher struct {
|
||||||
|
Mocks []*FakeResponse
|
||||||
|
Logging bool
|
||||||
|
PanicOnEmptyResponse bool // If not response matches - do we need to panic?
|
||||||
|
}
|
||||||
|
|
||||||
|
func (this *MockCatcher) Attach(fr []*FakeResponse) {
|
||||||
|
this.Mocks = append(this.Mocks, fr...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Find suitable response by provided
|
||||||
|
func (this *MockCatcher) FindResponse(query string, args []driver.NamedValue) *FakeResponse {
|
||||||
|
if this.Logging {
|
||||||
|
log.Printf("mock_catcher: check query: %s", query)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, resp := range this.Mocks {
|
||||||
|
if resp.IsMatch(query, args) {
|
||||||
|
resp.MarkAsTriggered()
|
||||||
|
return resp
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if this.PanicOnEmptyResponse {
|
||||||
|
panic(fmt.Sprintf("No responses matches query %s ", query))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Let's have always dummy version of response
|
||||||
|
return &FakeResponse{
|
||||||
|
Response: make([]map[string]interface{}, 0),
|
||||||
|
Exceptions: &Exceptions{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create new FakeResponse and return for chains of attachments
|
||||||
|
func (this *MockCatcher) NewMock() *FakeResponse {
|
||||||
|
fr := &FakeResponse{Exceptions: &Exceptions{}, Response: make([]map[string]interface{}, 0)}
|
||||||
|
this.Mocks = append(this.Mocks, fr)
|
||||||
|
return fr
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove all Mocks to start process again
|
||||||
|
func (this *MockCatcher) Reset() *MockCatcher {
|
||||||
|
this.Mocks = make([]*FakeResponse, 0)
|
||||||
|
return this
|
||||||
|
}
|
||||||
|
|
||||||
|
// Possible exceptions during query executions
|
||||||
|
type Exceptions struct {
|
||||||
|
HookQueryBadConnection func() bool
|
||||||
|
HookExecBadConnection func() bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// Represents mock of response with holding all required values to return mocked response
|
||||||
|
type FakeResponse struct {
|
||||||
|
Pattern string // SQL query pattern to match with
|
||||||
|
Args []interface{} // List args to be matched with
|
||||||
|
Response []map[string]interface{} // Array of rows to be parsed as result
|
||||||
|
Once bool // To trigger only once
|
||||||
|
Triggered bool // If it was triggered at least once
|
||||||
|
Callback func(string, []driver.NamedValue) // Callback to execute when response triggered
|
||||||
|
RowsAffected int64 // Defines affected rows count
|
||||||
|
LastInsertId int64 // ID to be returned for INSERT queries
|
||||||
|
*Exceptions
|
||||||
|
}
|
||||||
|
|
||||||
|
// Return true either when nothing to compare or deep equal check passed
|
||||||
|
func (fr *FakeResponse) isArgsMatch(args []driver.NamedValue) bool {
|
||||||
|
arguments := make([]interface{}, len(args))
|
||||||
|
if len(args) > 0 {
|
||||||
|
for index, arg := range args {
|
||||||
|
arguments[index] = arg.Value
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fr.Args == nil || reflect.DeepEqual(fr.Args, arguments)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResponse) isQueryMatch(query string) bool {
|
||||||
|
return fr.Pattern == "" || strings.Contains(query, fr.Pattern)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResponse) IsMatch(query string, args []driver.NamedValue) bool {
|
||||||
|
if fr.Once && fr.Triggered {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return fr.isQueryMatch(query) && fr.isArgsMatch(args)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResponse) MarkAsTriggered() {
|
||||||
|
fr.Triggered = true
|
||||||
|
}
|
||||||
|
|
||||||
|
// For chaining init
|
||||||
|
func (fr *FakeResponse) WithQuery(query string) *FakeResponse {
|
||||||
|
fr.Pattern = query
|
||||||
|
return fr
|
||||||
|
}
|
||||||
|
|
||||||
|
// Attach Args check for prepared statements
|
||||||
|
func (fr *FakeResponse) WithArgs(vars ...interface{}) *FakeResponse {
|
||||||
|
if len(vars) > 0 {
|
||||||
|
fr.Args = make([]interface{}, len(vars))
|
||||||
|
for index, v := range vars {
|
||||||
|
fr.Args[index] = v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fr
|
||||||
|
}
|
||||||
|
|
||||||
|
// Methods to chain and assign some parts of response
|
||||||
|
|
||||||
|
func (fr *FakeResponse) WithReply(response []map[string]interface{}) *FakeResponse {
|
||||||
|
fr.Response = response
|
||||||
|
return fr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResponse) OneTime() *FakeResponse {
|
||||||
|
fr.Once = true
|
||||||
|
return fr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResponse) WithExecException() *FakeResponse {
|
||||||
|
fr.Exceptions.HookExecBadConnection = func() bool {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return fr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResponse) WithQueryException() *FakeResponse {
|
||||||
|
fr.Exceptions.HookQueryBadConnection = func() bool {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return fr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResponse) WithCallback(f func(string, []driver.NamedValue)) *FakeResponse {
|
||||||
|
fr.Callback = f
|
||||||
|
return fr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResponse) WithRowsNum(num int64) *FakeResponse {
|
||||||
|
fr.RowsAffected = num
|
||||||
|
return fr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResponse) WithId(id int64) *FakeResponse {
|
||||||
|
fr.LastInsertId = id
|
||||||
|
return fr
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
Catcher = &MockCatcher{}
|
||||||
|
}
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
package go_mocket
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql/driver"
|
||||||
|
)
|
||||||
|
|
||||||
|
type FakeResult struct {
|
||||||
|
insertID int64
|
||||||
|
rowsAffected int64
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewFakeResult(insertId int64, rowsAffected int64) driver.Result {
|
||||||
|
return &FakeResult{insertId, rowsAffected}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResult) LastInsertId() (int64, error) {
|
||||||
|
return fr.insertID, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (fr *FakeResult) RowsAffected() (int64, error) {
|
||||||
|
return fr.rowsAffected, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,115 @@
|
|||||||
|
package go_mocket
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"database/sql/driver"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"reflect"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type RowsCursor struct {
|
||||||
|
cols []string
|
||||||
|
colType [][]string
|
||||||
|
posSet int
|
||||||
|
posRow int
|
||||||
|
rows [][]*row
|
||||||
|
closed bool
|
||||||
|
|
||||||
|
// errPos and err are for making Next return early with error.
|
||||||
|
errPos int
|
||||||
|
err error
|
||||||
|
|
||||||
|
bytesClone map[*byte][]byte
|
||||||
|
}
|
||||||
|
|
||||||
|
type row struct {
|
||||||
|
cols []interface{} // must be same size as its table colname + coltype
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rc *RowsCursor) Close() error {
|
||||||
|
if !rc.closed {
|
||||||
|
for _, bs := range rc.bytesClone {
|
||||||
|
bs[0] = 255 // first byte corrupted
|
||||||
|
}
|
||||||
|
}
|
||||||
|
rc.closed = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rc *RowsCursor) Columns() []string {
|
||||||
|
return rc.cols
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rc *RowsCursor) ColumnTypeScanType(index int) reflect.Type {
|
||||||
|
return colTypeToReflectType(rc.colType[rc.posSet][index])
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rc *RowsCursor) Next(accumulator []driver.Value) error {
|
||||||
|
if rc.closed {
|
||||||
|
return errors.New("fake_db_driver: cursor is closed")
|
||||||
|
}
|
||||||
|
rc.posRow++
|
||||||
|
if rc.posRow == rc.errPos {
|
||||||
|
return rc.err
|
||||||
|
}
|
||||||
|
if rc.posRow >= len(rc.rows[rc.posSet]) {
|
||||||
|
return io.EOF // per interface spec
|
||||||
|
}
|
||||||
|
for i, v := range rc.rows[rc.posSet][rc.posRow].cols {
|
||||||
|
accumulator[i] = v
|
||||||
|
if bs, ok := v.([]byte); ok {
|
||||||
|
if rc.bytesClone == nil {
|
||||||
|
rc.bytesClone = make(map[*byte][]byte)
|
||||||
|
}
|
||||||
|
clone, ok := rc.bytesClone[&bs[0]]
|
||||||
|
if !ok {
|
||||||
|
clone = make([]byte, len(bs))
|
||||||
|
copy(clone, bs)
|
||||||
|
rc.bytesClone[&bs[0]] = clone
|
||||||
|
}
|
||||||
|
accumulator[i] = clone
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rc *RowsCursor) HasNextResultSet() bool {
|
||||||
|
return rc.posSet < len(rc.rows)-1
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rc *RowsCursor) NextResultSet() error {
|
||||||
|
if rc.HasNextResultSet() {
|
||||||
|
rc.posSet++
|
||||||
|
rc.posRow = -1
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return io.EOF // Per interface spec.
|
||||||
|
}
|
||||||
|
|
||||||
|
func colTypeToReflectType(typ string) reflect.Type {
|
||||||
|
switch typ {
|
||||||
|
case "bool":
|
||||||
|
return reflect.TypeOf(false)
|
||||||
|
case "nullbool":
|
||||||
|
return reflect.TypeOf(sql.NullBool{})
|
||||||
|
case "int32":
|
||||||
|
return reflect.TypeOf(int32(0))
|
||||||
|
case "string":
|
||||||
|
return reflect.TypeOf("")
|
||||||
|
case "nullstring":
|
||||||
|
return reflect.TypeOf(sql.NullString{})
|
||||||
|
case "int64":
|
||||||
|
return reflect.TypeOf(int64(0))
|
||||||
|
case "nullint64":
|
||||||
|
return reflect.TypeOf(sql.NullInt64{})
|
||||||
|
case "float64":
|
||||||
|
return reflect.TypeOf(float64(0))
|
||||||
|
case "nullfloat64":
|
||||||
|
return reflect.TypeOf(sql.NullFloat64{})
|
||||||
|
case "datetime":
|
||||||
|
return reflect.TypeOf(time.Time{})
|
||||||
|
}
|
||||||
|
panic("invalid fakedb column type of " + typ)
|
||||||
|
}
|
||||||
@@ -0,0 +1,175 @@
|
|||||||
|
package go_mocket
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql/driver"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"math/rand"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
type FakeStmt struct {
|
||||||
|
connection *FakeConn
|
||||||
|
q string // just for debugging SQL query generated by sql package
|
||||||
|
command string // String name of the command SELECT etc, taken as first word in the query
|
||||||
|
next *FakeStmt // used for returning multiple results.
|
||||||
|
closed bool // If connection closed already
|
||||||
|
colName []string //Names of columns in response
|
||||||
|
colType []string // Not used for now
|
||||||
|
placeholders int // Amount of passed args
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *FakeStmt) ColumnConverter(idx int) driver.ValueConverter {
|
||||||
|
return driver.DefaultParameterConverter
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *FakeStmt) Close() error {
|
||||||
|
// No connection added
|
||||||
|
if s.connection == nil {
|
||||||
|
panic("nil conn in FakeStmt.Close")
|
||||||
|
}
|
||||||
|
if s.connection.db == nil {
|
||||||
|
panic("in FakeStmt.Close, conn's db is nil (already closed)")
|
||||||
|
}
|
||||||
|
if !s.closed {
|
||||||
|
s.closed = true
|
||||||
|
}
|
||||||
|
if s.next != nil {
|
||||||
|
s.next.Close()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var errClosed = errors.New("fake_db_driver: statement has been closed")
|
||||||
|
|
||||||
|
func (smt *FakeStmt) Exec(args []driver.Value) (driver.Result, error) {
|
||||||
|
panic("Using ExecContext")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (smt *FakeStmt) ExecContext(ctx context.Context, args []driver.NamedValue) (driver.Result, error) {
|
||||||
|
if smt.closed {
|
||||||
|
return nil, errClosed
|
||||||
|
}
|
||||||
|
|
||||||
|
fResp := Catcher.FindResponse(smt.q, args)
|
||||||
|
|
||||||
|
// To emulate any exception during query which returns rows
|
||||||
|
if fResp.Exceptions != nil && fResp.Exceptions.HookExecBadConnection != nil && fResp.Exceptions.HookExecBadConnection() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
if fResp.Callback != nil {
|
||||||
|
fResp.Callback(smt.q, args)
|
||||||
|
}
|
||||||
|
|
||||||
|
switch smt.command {
|
||||||
|
case "INSERT":
|
||||||
|
id := fResp.LastInsertId
|
||||||
|
if id == 0 {
|
||||||
|
id = rand.Int63()
|
||||||
|
}
|
||||||
|
res := NewFakeResult(id, 1)
|
||||||
|
return res, nil
|
||||||
|
case "UPDATE":
|
||||||
|
return driver.RowsAffected(fResp.RowsAffected), nil
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("unimplemented statement Exec command type of %q", smt.command)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *FakeStmt) Query(args []driver.Value) (driver.Rows, error) {
|
||||||
|
panic("Use QueryContext")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (smt *FakeStmt) QueryContext(ctx context.Context, args []driver.NamedValue) (driver.Rows, error) {
|
||||||
|
|
||||||
|
if smt.closed {
|
||||||
|
return nil, errClosed
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(args) > 0 {
|
||||||
|
// Replace all "?" to "%v" and replace them with the values after
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
smt.q = strings.Replace(smt.q, "?", "%v", 1)
|
||||||
|
smt.q = fmt.Sprintf(smt.q, args[i].Value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fResp := Catcher.FindResponse(smt.q, args)
|
||||||
|
|
||||||
|
if fResp.Exceptions != nil && fResp.Exceptions.HookQueryBadConnection != nil && fResp.Exceptions.HookQueryBadConnection() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
resultRows := make([][]*row, 0, 1)
|
||||||
|
columnNames := make([]string, 0, 1)
|
||||||
|
columnTypes := make([][]string, 0, 1)
|
||||||
|
rows := []*row{}
|
||||||
|
|
||||||
|
// Check if we have such query in the map
|
||||||
|
colIndexes := make(map[string]int)
|
||||||
|
|
||||||
|
// Collecting column names from first record
|
||||||
|
if len(fResp.Response) > 0 {
|
||||||
|
for colName, _ := range fResp.Response[0] {
|
||||||
|
columnNames = append(columnNames, colName)
|
||||||
|
colIndexes[colName] = len(columnNames) - 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extracting values from result according columns
|
||||||
|
for _, record := range fResp.Response {
|
||||||
|
oneRow := &row{cols: make([]interface{}, len(columnNames))}
|
||||||
|
for _, col := range columnNames {
|
||||||
|
oneRow.cols[colIndexes[col]] = []byte(record[col].(string))
|
||||||
|
}
|
||||||
|
rows = append(rows, oneRow)
|
||||||
|
}
|
||||||
|
resultRows = append(resultRows, rows)
|
||||||
|
|
||||||
|
cursor := &RowsCursor{
|
||||||
|
posRow: -1,
|
||||||
|
rows: resultRows,
|
||||||
|
cols: columnNames,
|
||||||
|
colType: columnTypes, // TODO: implement support of that
|
||||||
|
errPos: -1,
|
||||||
|
closed: false,
|
||||||
|
}
|
||||||
|
|
||||||
|
if fResp.Callback != nil {
|
||||||
|
fResp.Callback(smt.q, args)
|
||||||
|
}
|
||||||
|
|
||||||
|
return cursor, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Returns number of args passed to query
|
||||||
|
func (s *FakeStmt) NumInput() int {
|
||||||
|
return s.placeholders
|
||||||
|
}
|
||||||
|
|
||||||
|
type FakeTx struct {
|
||||||
|
c *FakeConn
|
||||||
|
}
|
||||||
|
|
||||||
|
// hook to simulate broken connections
|
||||||
|
var HookBadCommit func() bool
|
||||||
|
|
||||||
|
func (tx *FakeTx) Commit() error {
|
||||||
|
tx.c.currTx = nil
|
||||||
|
if HookBadCommit != nil && HookBadCommit() {
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// hook to simulate broken connections
|
||||||
|
var HookBadRollback func() bool
|
||||||
|
|
||||||
|
func (tx *FakeTx) Rollback() error {
|
||||||
|
tx.c.currTx = nil
|
||||||
|
if HookBadRollback != nil && HookBadRollback() {
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user