diff --git a/memcache/memcache.go b/memcache/memcache.go index 6f48caa..a9897b0 100644 --- a/memcache/memcache.go +++ b/memcache/memcache.go @@ -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 @@ -113,6 +134,7 @@ var ( resultTouched = []byte("TOUCHED\r\n") resultClientErrorPrefix = []byte("CLIENT_ERROR ") + resultServerErrorPrefix = []byte("SERVER_ERROR ") versionPrefix = []byte("VERSION") ) @@ -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)) } } @@ -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) @@ -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)) } @@ -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)) } @@ -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 { diff --git a/memcache/memcache_test.go b/memcache/memcache_test.go index a0fa746..fb49c2a 100644 --- a/memcache/memcache_test.go +++ b/memcache/memcache_test.go @@ -22,6 +22,7 @@ import ( "bytes" "context" "crypto/tls" + "errors" "flag" "fmt" "io" @@ -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{}) +}