| 1 | package provider |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "crypto/tls" |
| 6 | "errors" |
| 7 | "fmt" |
| 8 | "io" |
| 9 | "net/http" |
| 10 | "net/http/httptest" |
| 11 | "net/url" |
| 12 | "sync" |
| 13 | "testing" |
| 14 | "time" |
| 15 | |
| 16 | "golang.org/x/net/http2" |
| 17 | ) |
| 18 | |
| 19 | func TestHTTP2FailureClassificationDoesNotEnableRetries(t *testing.T) { |
| 20 | for _, cause := range []error{ |
| 21 | http2.ConnectionError(http2.ErrCodeProtocol), |
| 22 | http2.StreamError{StreamID: 3, Code: http2.ErrCodeProtocol}, |
| 23 | http2.GoAwayError{LastStreamID: 3, ErrCode: http2.ErrCodeProtocol, DebugData: "private"}, |
| 24 | } { |
| 25 | t.Run(cause.Error(), func(t *testing.T) { |
| 26 | calls := 0 |
| 27 | client := &http.Client{Transport: rtFunc(func(*http.Request) (*http.Response, error) { calls++; return nil, cause })} |
| 28 | _, err := SendWithRetry(t.Context(), client, SendOptions{Provider: "test", Protocol: "openai"}, newDummyReq) |
| 29 | diagnostic := DiagnoseFailure(fmt.Errorf("outer: %w", err)) |
| 30 | if calls != 1 || diagnostic.Kind != FailureKindTransportProtocol || diagnostic.TransportCode != "PROTOCOL_ERROR" || !errors.Is(err, cause) { |
| 31 | t.Fatalf("calls=%d diagnostic=%+v err=%v", calls, diagnostic, err) |
| 32 | } |
| 33 | if ClassifyRecovery(err).Retryable { |
| 34 | t.Fatal("classification changed the retry policy") |
| 35 | } |
| 36 | }) |
| 37 | } |
| 38 | for _, cause := range []error{ |
| 39 | errors.New("connection error: PROTOCOL_ERROR"), |
| 40 | errors.New("stream error: stream ID 3; PROTOCOL_ERROR"), |
| 41 | errors.New("http2: server sent GOAWAY and closed the connection; LastStreamID=3, ErrCode=PROTOCOL_ERROR"), |
| 42 | errors.New("invalid header field value: PROTOCOL_ERROR"), |
| 43 | &APIError{Body: "connection error: PROTOCOL_ERROR", Status: 400}, |
| 44 | fmt.Errorf("connection error: PROTOCOL_ERROR: %w", &APIError{Status: 400}), |
| 45 | context.Canceled, |
| 46 | } { |
| 47 | if got := DiagnoseFailure(&url.Error{Op: "Post", URL: "https://example.test", Err: cause}); got.Kind == FailureKindTransportProtocol { |
| 48 | t.Fatalf("misclassified caller/provider error: %+v", got) |
| 49 | } |
| 50 | } |
| 51 | } |
| 52 | |
| 53 | func TestHTTP2StdlibProtocolFailureRecordsNegotiatedTransport(t *testing.T) { |
| 54 | server := httptest.NewUnstartedServer(http.NotFoundHandler()) |
| 55 | server.EnableHTTP2 = true |
| 56 | server.Config.TLSNextProto = map[string]func(*http.Server, *tls.Conn, http.Handler){ |
| 57 | "h2": func(_ *http.Server, conn *tls.Conn, _ http.Handler) { |
| 58 | defer conn.Close() |
| 59 | _ = conn.SetDeadline(time.Now().Add(5 * time.Second)) |
| 60 | if _, err := io.ReadFull(conn, make([]byte, len(http2.ClientPreface))); err != nil { |
| 61 | return |
| 62 | } |
| 63 | // A valid initial SETTINGS frame, followed by an illegal DATA frame on |
| 64 | // stream zero, deterministically exercises net/http's private error type. |
| 65 | _, _ = conn.Write([]byte{0, 0, 0, 4, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}) |
| 66 | _, _ = io.Copy(io.Discard, conn) |
| 67 | }, |
| 68 | } |
| 69 | server.StartTLS() |
| 70 | defer server.Close() |
| 71 | var mu sync.Mutex |
| 72 | var last RequestObservation |
| 73 | ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) |
| 74 | defer cancel() |
| 75 | ctx = WithRequestObserver(ctx, func(v RequestObservation) { mu.Lock(); last = v; mu.Unlock() }) |
| 76 | _, err := SendWithRetry(ctx, server.Client(), SendOptions{}, func(ctx context.Context) (*http.Request, error) { |
| 77 | return http.NewRequestWithContext(ctx, http.MethodPost, server.URL+"/v1/chat/completions?credential=private-query", nil) |
| 78 | }) |
| 79 | if err == nil || DiagnoseFailure(err).Kind != FailureKindTransportProtocol { |
| 80 | t.Fatalf("stdlib error: %T %v", err, err) |
| 81 | } |
| 82 | mu.Lock() |
| 83 | defer mu.Unlock() |
| 84 | if last.TransportCode != "PROTOCOL_ERROR" || last.HTTPProtocol != "HTTP/2.0" || last.RemoteAddress == "" || last.Phase != "request_error" || last.ConnectedAt.IsZero() || !last.HeadersAt.IsZero() { |
| 85 | t.Fatalf("lost negotiated transport evidence: %+v", last) |
| 86 | } |
| 87 | } |
| 88 |