package main // A fake server: every session records who it connected as and what it was sent, and answers by the // first rule whose pattern matches the statement. That the server keeps data, refuses a reader's // write or installs an extension is proven against a real server (live_test.go), not here; this // holds the module to asking for the right things. import ( "context" "regexp" "strings" "sync" "testing" ) type call struct { Login SQL string } type rule struct { match *regexp.Regexp fields []string rows [][]string err error } type fakeServer struct { mu sync.Mutex calls []call rules []rule dialErr func(Login) error } func (f *fakeServer) answer(pattern string, fields []string, rows ...[]string) { f.rules = append(f.rules, rule{match: regexp.MustCompile(pattern), fields: fields, rows: rows}) } func (f *fakeServer) fail(pattern string, err error) { f.rules = append(f.rules, rule{match: regexp.MustCompile(pattern), err: err}) } type fakeSession struct { f *fakeServer l Login } func (s fakeSession) Run(_ context.Context, sql string) ([]Result, error) { s.f.mu.Lock() defer s.f.mu.Unlock() s.f.calls = append(s.f.calls, call{Login: s.l, SQL: sql}) for _, r := range s.f.rules { if r.match.MatchString(sql) { if r.err != nil { return nil, r.err } res := Result{Command: strings.Fields(sql)[0], Fields: r.fields} for _, row := range r.rows { cells := make([]*string, len(row)) for i := range row { v := row[i] cells[i] = &v } res.Rows = append(res.Rows, cells) } return []Result{res}, nil } } return []Result{{Command: strings.Fields(sql)[0]}}, nil } func (s fakeSession) Close() {} func (f *fakeServer) dial(_ context.Context, l Login) (Session, error) { if f.dialErr != nil { if err := f.dialErr(l); err != nil { return nil, err } } return fakeSession{f: f, l: l}, nil } func (f *fakeServer) statements() []string { f.mu.Lock() defer f.mu.Unlock() out := make([]string, len(f.calls)) for i, c := range f.calls { out[i] = c.SQL } return out } func (f *fakeServer) reset() { f.mu.Lock() defer f.mu.Unlock() f.calls = nil } var admin = Conn{Host: "127.0.0.1", Port: 5432, User: "postgres", Password: "admin-secret"} func newFake(conn Conn) (*fakeServer, *Client) { f := &fakeServer{} return f, NewClient(conn, f.dial) } func noDrop(t *testing.T, sql []string) { t.Helper() for _, s := range sql { if regexp.MustCompile(`(?i)\bDROP\b`).MatchString(s) { t.Fatalf("something was dropped:\n%s", strings.Join(sql, "\n")) } } } func has(sql []string, pattern string) bool { re := regexp.MustCompile(pattern) for _, s := range sql { if re.MatchString(s) { return true } } return false }