Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 44 additions & 4 deletions memcache/memcache.go
Original file line number Diff line number Diff line change
Expand Up @@ -75,18 +75,39 @@ const (

const buffered = 8 // arbitrary buffered channel size, for readability

// resumableError returns true if err is only a protocol-level cache error.
// resumableError returns true if err is only a protocol-level error.
// This is used to determine whether or not a server connection should
// be re-used or not. If an error occurs, by default we don't reuse the
// connection, unless it was just a cache error.
// connection, unless it was just a protocol-level error.
//
// SERVER_ERROR replies are resumable: memcached sends them as complete,
// well-formed lines and keeps the connection in sync (on storage errors it
// even swallows the request's data block; see process_update_command in
// memcached's proto_text.c). The known exception is "out of memory reading
// request", which closes the connection server-side; reusing it costs one
// failed request, the same as any stale pooled connection.
func resumableError(err error) bool {
switch err {
case ErrCacheMiss, ErrCASConflict, ErrNotStored, ErrMalformedKey:
switch {
case errors.Is(err, ErrCacheMiss),
errors.Is(err, ErrCASConflict),
errors.Is(err, ErrNotStored),
errors.Is(err, ErrMalformedKey),
errors.Is(err, ErrServerError):
return true
}
return false
}

// serverErrorFromLine returns an error wrapping ErrServerError if line is a
// SERVER_ERROR response, or nil otherwise.
func serverErrorFromLine(line []byte) error {
if !bytes.HasPrefix(line, resultServerErrorPrefix) {
return nil
}
msg := bytes.TrimSuffix(line[len(resultServerErrorPrefix):], crlf)
return fmt.Errorf("%w: %s", ErrServerError, msg)
}

func legalKey(key string) bool {
if len(key) > 250 {
return false
Expand All @@ -113,6 +134,7 @@ var (
resultTouched = []byte("TOUCHED\r\n")

resultClientErrorPrefix = []byte("CLIENT_ERROR ")
resultServerErrorPrefix = []byte("SERVER_ERROR ")
versionPrefix = []byte("VERSION")
)

Expand Down Expand Up @@ -464,6 +486,9 @@ func (c *Client) touchFromAddr(addr net.Addr, keys []string, expiration int32) e
case bytes.Equal(line, resultNotFound):
return ErrCacheMiss
default:
if err := serverErrorFromLine(line); err != nil {
return err
}
return fmt.Errorf("memcache: unexpected response line from touch: %q", string(line))
}
}
Expand Down Expand Up @@ -530,6 +555,11 @@ func parseGetResponse(r *bufio.Reader, conn *conn, cb func(*Item)) error {
it := new(Item)
size, err := scanGetResponseLine(line, it)
if err != nil {
// The line is not a VALUE line: check whether it's a protocol-level
// server error before reporting it as unexpected.
if serverErr := serverErrorFromLine(line); serverErr != nil {
return serverErr
}
return err
}
it.Value = make([]byte, size+2)
Expand Down Expand Up @@ -699,6 +729,10 @@ func (c *Client) populateOne(rw *bufio.ReadWriter, verb string, item *Item) erro
return ErrCASConflict
case bytes.Equal(line, resultNotFound):
return ErrCacheMiss
default:
if err := serverErrorFromLine(line); err != nil {
return err
}
}
return fmt.Errorf("memcache: unexpected response line from %q: %q", verb, string(line))
}
Expand Down Expand Up @@ -731,6 +765,10 @@ func writeExpectf(rw *bufio.ReadWriter, expect []byte, format string, args ...in
return ErrCASConflict
case bytes.Equal(line, resultNotFound):
return ErrCacheMiss
default:
if err := serverErrorFromLine(line); err != nil {
return err
}
}
return fmt.Errorf("memcache: unexpected response line: %q", string(line))
}
Expand Down Expand Up @@ -817,6 +855,8 @@ func (c *Client) incrDecr(verb, key string, delta uint64) (uint64, error) {
case bytes.HasPrefix(line, resultClientErrorPrefix):
errMsg := line[len(resultClientErrorPrefix) : len(line)-2]
return errors.New("memcache: client error: " + string(errMsg))
case bytes.HasPrefix(line, resultServerErrorPrefix):
return serverErrorFromLine(line)
}
val, err = strconv.ParseUint(string(line[:len(line)-2]), 10, 64)
if err != nil {
Expand Down
197 changes: 197 additions & 0 deletions memcache/memcache_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ import (
"bytes"
"context"
"crypto/tls"
"errors"
"flag"
"fmt"
"io"
Expand Down Expand Up @@ -502,3 +503,199 @@ func TestScanGetResponseLine(t *testing.T) {
})
}
}

func TestServerError_KeepsConnectionOpen(t *testing.T) {
// countFreeConns returns the number of connections in the client's free pool.
countFreeConns := func(c *Client) int {
c.mu.Lock()
defer c.mu.Unlock()
return len(c.freeconn[dummyAddr{}.String()])
}

// assertReusable verifies that the connection went back to the free pool and
// that a subsequent request over the very same connection still works.
assertReusable := func(t *testing.T, c *Client, srvw io.Writer, srvReqs <-chan string) {
t.Helper()

if got := countFreeConns(c); got != 1 {
t.Fatalf("free conns after SERVER_ERROR: got %d, want 1", got)
}

errCh := make(chan error)
go func() {
_, err := c.Get("foo")
errCh <- err
}()

if req := <-srvReqs; req != "gets foo\r\n" {
t.Fatalf("unexpected follow-up request: %q", req)
}
if _, err := io.WriteString(srvw, "END\r\n"); err != nil {
t.Fatalf("write follow-up response: %v", err)
}

if err := <-errCh; !errors.Is(err, ErrCacheMiss) {
t.Fatalf("follow-up Get on reused connection: got err=%v, want ErrCacheMiss", err)
}
}

t.Run("Set", func(t *testing.T) {
c, srvw, srvReqs := getClientWithFakeServer(t)

errCh := make(chan error)
go func() {
errCh <- c.Set(&Item{Key: "foo", Value: []byte("hello")})
}()

if req := <-srvReqs; req != "set foo 0 0 5\r\n" {
t.Fatalf("unexpected request: %q", req)
}
<-srvReqs // data block
if _, err := io.WriteString(srvw, "SERVER_ERROR object too large for cache\r\n"); err != nil {
t.Fatalf("write response: %v", err)
}

err := <-errCh
if !errors.Is(err, ErrServerError) {
t.Fatalf("Set: got err=%v, want ErrServerError", err)
}
if want := "memcache: server error: object too large for cache"; err.Error() != want {
t.Fatalf("Set: got err=%q, want %q", err.Error(), want)
}

assertReusable(t, c, srvw, srvReqs)
})

t.Run("Get", func(t *testing.T) {
c, srvw, srvReqs := getClientWithFakeServer(t)

errCh := make(chan error)
go func() {
_, err := c.Get("foo")
errCh <- err
}()

if req := <-srvReqs; req != "gets foo\r\n" {
t.Fatalf("unexpected request: %q", req)
}
if _, err := io.WriteString(srvw, "SERVER_ERROR out of memory writing get response\r\n"); err != nil {
t.Fatalf("write response: %v", err)
}

if err := <-errCh; !errors.Is(err, ErrServerError) {
t.Fatalf("Get: got err=%v, want ErrServerError", err)
}

assertReusable(t, c, srvw, srvReqs)
})

t.Run("Delete", func(t *testing.T) {
c, srvw, srvReqs := getClientWithFakeServer(t)

errCh := make(chan error)
go func() {
errCh <- c.Delete("foo")
}()

if req := <-srvReqs; req != "delete foo\r\n" {
t.Fatalf("unexpected request: %q", req)
}
if _, err := io.WriteString(srvw, "SERVER_ERROR temporary failure\r\n"); err != nil {
t.Fatalf("write response: %v", err)
}

if err := <-errCh; !errors.Is(err, ErrServerError) {
t.Fatalf("Delete: got err=%v, want ErrServerError", err)
}

assertReusable(t, c, srvw, srvReqs)
})

t.Run("Increment", func(t *testing.T) {
c, srvw, srvReqs := getClientWithFakeServer(t)

errCh := make(chan error)
go func() {
_, err := c.Increment("foo", 1)
errCh <- err
}()

if req := <-srvReqs; req != "incr foo 1\r\n" {
t.Fatalf("unexpected request: %q", req)
}
if _, err := io.WriteString(srvw, "SERVER_ERROR temporary failure\r\n"); err != nil {
t.Fatalf("write response: %v", err)
}

if err := <-errCh; !errors.Is(err, ErrServerError) {
t.Fatalf("Increment: got err=%v, want ErrServerError", err)
}

assertReusable(t, c, srvw, srvReqs)
})

t.Run("non-protocol errors still close the connection", func(t *testing.T) {
c, srvw, srvReqs := getClientWithFakeServer(t)

errCh := make(chan error)
go func() {
errCh <- c.Set(&Item{Key: "foo", Value: []byte("hello")})
}()

<-srvReqs // command line
<-srvReqs // data block
if _, err := io.WriteString(srvw, "BOGUS RESPONSE\r\n"); err != nil {
t.Fatalf("write response: %v", err)
}

err := <-errCh
if err == nil || errors.Is(err, ErrServerError) {
t.Fatalf("Set: got err=%v, want a non-ErrServerError error", err)
}
if got := countFreeConns(c); got != 0 {
t.Fatalf("free conns after unexpected response: got %d, want 0", got)
}
})
}

// getClientWithFakeServer creates a new client, whose dial function is bound to an in-memory synchronous server.
func getClientWithFakeServer(t *testing.T) (*Client, io.Writer, <-chan string) {
cliConn, srvConn := net.Pipe()
t.Cleanup(func() {
_ = srvConn.Close()
_ = cliConn.Close()
})

c := NewFromSelector(dummySelector{})
c.DialContext = func(context.Context, string, string) (net.Conn, error) {
return cliConn, nil
}

srvReqs := make(chan string)
go func() {
br := bufio.NewReader(srvConn)
for {
// ignores error for simplicity
line, _ := br.ReadSlice('\n')
srvReqs <- string(line)
}
}()

return c, srvConn, srvReqs
}

type dummyAddr struct{}

func (dummyAddr) Network() string { return "dummy" }

func (dummyAddr) String() string { return "dummy" }

type dummySelector struct{}

func (dummySelector) PickServer(string) (net.Addr, error) {
return dummyAddr{}, nil
}

func (dummySelector) Each(f func(net.Addr) error) error {
return f(dummyAddr{})
}