| 1 | package provider |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "crypto/x509" |
| 6 | "errors" |
| 7 | "fmt" |
| 8 | "io" |
| 9 | "net" |
| 10 | "net/http" |
| 11 | "net/url" |
| 12 | "os" |
| 13 | "syscall" |
| 14 | "testing" |
| 15 | ) |
| 16 | |
| 17 | func TestTransportRecoveryRequiresNetworkCause(t *testing.T) { |
| 18 | for _, tt := range []struct { |
| 19 | name string |
| 20 | err error |
| 21 | want bool |
| 22 | }{ |
| 23 | {"invalid URL", errors.New(`unsupported protocol scheme ""`), false}, |
| 24 | {"invalid header", errors.New(`net/http: invalid header field name "bad header"`), false}, |
| 25 | {"missing history file", &os.PathError{Op: "open", Path: "old.jsonl", Err: os.ErrNotExist}, false}, |
| 26 | {"storage permission", &os.PathError{Op: "open", Path: "old.jsonl", Err: syscall.EACCES}, false}, |
| 27 | {"storage full", &os.SyscallError{Syscall: "write", Err: syscall.ENOSPC}, false}, |
| 28 | {"invalid address", &net.AddrError{Err: "missing port", Addr: "provider"}, false}, |
| 29 | {"certificate", x509.UnknownAuthorityError{}, false}, |
| 30 | {"cancel", context.Canceled, false}, |
| 31 | {"deadline", context.DeadlineExceeded, false}, |
| 32 | {"EOF", io.EOF, true}, |
| 33 | {"truncated body", io.ErrUnexpectedEOF, true}, |
| 34 | {"closed socket", net.ErrClosed, true}, |
| 35 | {"reset", syscall.ECONNRESET, true}, |
| 36 | {"refused", &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED}, true}, |
| 37 | {"read failure", &net.OpError{Op: "read", Net: "tcp", Err: errors.New("wsarecv: forcibly closed")}, true}, |
| 38 | {"DNS", &net.DNSError{Err: "temporary failure", Name: "provider.invalid", IsTemporary: true}, true}, |
| 39 | {"timeout", os.ErrDeadlineExceeded, true}, |
| 40 | } { |
| 41 | t.Run(tt.name, func(t *testing.T) { |
| 42 | for _, err := range []error{tt.err, &url.Error{Op: "Post", URL: "https://provider.invalid", Err: tt.err}} { |
| 43 | err = fmt.Errorf("provider request: %w", err) |
| 44 | if got := IsConnReset(err); got != tt.want { |
| 45 | t.Errorf("IsConnReset(%v) = %v, want %v", err, got, tt.want) |
| 46 | } |
| 47 | if got := ClassifyRecovery(err); got.Retryable != tt.want { |
| 48 | t.Errorf("ClassifyRecovery(%v) = %+v, want retryable %v", err, got, tt.want) |
| 49 | } |
| 50 | } |
| 51 | }) |
| 52 | } |
| 53 | } |
| 54 | |
| 55 | func TestSendWithRetryRejectsPermanentTransportErrors(t *testing.T) { |
| 56 | for _, managed := range []bool{false, true} { |
| 57 | t.Run(fmt.Sprintf("managed=%v", managed), func(t *testing.T) { |
| 58 | cause := errors.New("invalid request configuration") |
| 59 | calls, retries := 0, 0 |
| 60 | client := &http.Client{Transport: rtFunc(func(*http.Request) (*http.Response, error) { |
| 61 | calls++ |
| 62 | return nil, cause |
| 63 | })} |
| 64 | ctx, cancel := context.WithCancel(t.Context()) |
| 65 | defer cancel() |
| 66 | ctx = WithRetryNotify(ctx, func(RetryInfo) { retries++; cancel() }) |
| 67 | if managed { |
| 68 | ctx = WithManagedRecovery(ctx) |
| 69 | } |
| 70 | _, err := SendWithRetry(ctx, client, SendOptions{Provider: "saved-provider"}, newDummyReq) |
| 71 | if calls != 1 || retries != 0 || !errors.Is(err, cause) { |
| 72 | t.Fatalf("calls=%d retries=%d err=%v", calls, retries, err) |
| 73 | } |
| 74 | if ClassifyRecovery(err).Retryable || DiagnoseFailure(err).Kind == "temporary" { |
| 75 | t.Fatalf("permanent transport failure entered network recovery: %v", err) |
| 76 | } |
| 77 | }) |
| 78 | } |
| 79 | } |
| 80 | |
| 81 | func TestSendWithRetryLeavesConnectionRetryToCaller(t *testing.T) { |
| 82 | calls := 0 |
| 83 | client := &http.Client{Transport: rtFunc(func(*http.Request) (*http.Response, error) { |
| 84 | calls++ |
| 85 | if calls == 1 { |
| 86 | return nil, &net.OpError{Op: "read", Net: "tcp", Err: syscall.ECONNRESET} |
| 87 | } |
| 88 | return statusResp(http.StatusOK, nil), nil |
| 89 | })} |
| 90 | _, err := SendWithRetry(t.Context(), client, SendOptions{}, newDummyReq) |
| 91 | if !errors.Is(err, syscall.ECONNRESET) || calls != 1 { |
| 92 | t.Fatalf("calls=%d err=%v", calls, err) |
| 93 | } |
| 94 | resp, err := SendWithRetry(t.Context(), client, SendOptions{}, newDummyReq) |
| 95 | if err != nil { |
| 96 | t.Fatal(err) |
| 97 | } |
| 98 | defer resp.Body.Close() |
| 99 | if calls != 2 { |
| 100 | t.Fatalf("connection reset made %d requests, want 2", calls) |
| 101 | } |
| 102 | } |
| 103 |