| 1 | package provider |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "errors" |
| 7 | "io" |
| 8 | "net/http" |
| 9 | "net/http/httptest" |
| 10 | "strings" |
| 11 | "sync" |
| 12 | "testing" |
| 13 | "time" |
| 14 | ) |
| 15 | |
| 16 | func TestRequestObservationDistinguishesHeaderAndBodyWaitCancellation(t *testing.T) { |
| 17 | for _, headers := range []bool{false, true} { |
| 18 | t.Run(map[bool]string{false: "headers", true: "body"}[headers], func(t *testing.T) { |
| 19 | arrived := make(chan struct{}) |
| 20 | server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 21 | if headers { |
| 22 | w.Header().Set("Content-Type", "text/event-stream") |
| 23 | _, _ = io.WriteString(w, ": ping\n\n") |
| 24 | w.(http.Flusher).Flush() |
| 25 | } |
| 26 | close(arrived) |
| 27 | <-r.Context().Done() |
| 28 | })) |
| 29 | defer server.Close() |
| 30 | ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) |
| 31 | defer cancel() |
| 32 | var mu sync.Mutex |
| 33 | var observations []RequestObservation |
| 34 | ctx = WithRequestObserver(ctx, func(v RequestObservation) { mu.Lock(); observations = append(observations, v); mu.Unlock() }) |
| 35 | snapshot := func() RequestObservation { mu.Lock(); defer mu.Unlock(); return observations[len(observations)-1] } |
| 36 | done := make(chan error, 1) |
| 37 | bodyRead := make(chan struct{}) |
| 38 | go func() { |
| 39 | resp, err := SendWithRetry(ctx, server.Client(), SendOptions{}, func(ctx context.Context) (*http.Request, error) { |
| 40 | return http.NewRequestWithContext(ctx, http.MethodGet, server.URL+"/?secret=private-token", nil) |
| 41 | }) |
| 42 | if err == nil { |
| 43 | defer resp.Body.Close() |
| 44 | _, err = io.ReadFull(resp.Body, make([]byte, len(": ping\n\n"))) |
| 45 | close(bodyRead) |
| 46 | if err == nil { |
| 47 | _, err = resp.Body.Read(make([]byte, 1)) |
| 48 | } |
| 49 | } |
| 50 | done <- err |
| 51 | }() |
| 52 | select { |
| 53 | case <-arrived: |
| 54 | case <-ctx.Done(): |
| 55 | t.Fatal(ctx.Err()) |
| 56 | } |
| 57 | if headers { |
| 58 | select { |
| 59 | case <-bodyRead: |
| 60 | case <-ctx.Done(): |
| 61 | t.Fatal(ctx.Err()) |
| 62 | } |
| 63 | } |
| 64 | before := snapshot() |
| 65 | if headers { |
| 66 | if before.HeadersAt.IsZero() || before.FirstBodyAt.IsZero() || before.BodyBytes != 8 || !before.FinishedAt.IsZero() { |
| 67 | t.Fatalf("body wait: %+v", before) |
| 68 | } |
| 69 | } else if !before.HeadersAt.IsZero() || before.BodyBytes != 0 || !before.FinishedAt.IsZero() { |
| 70 | t.Fatalf("header wait: %+v", before) |
| 71 | } |
| 72 | cancel() |
| 73 | select { |
| 74 | case err := <-done: |
| 75 | if !errors.Is(err, context.Canceled) { |
| 76 | t.Fatalf("cancel: %v", err) |
| 77 | } |
| 78 | case <-time.After(5 * time.Second): |
| 79 | t.Fatal("cancellation stalled") |
| 80 | } |
| 81 | after := snapshot() |
| 82 | if after.Phase != "canceled" || after.FinishedAt.IsZero() || before.ID != after.ID { |
| 83 | t.Fatalf("terminal observation: %+v", after) |
| 84 | } |
| 85 | encoded, _ := json.Marshal(after) |
| 86 | if strings.Contains(string(encoded), "private-token") || strings.Contains(string(encoded), server.URL) { |
| 87 | t.Fatal("diagnostics leaked transport content") |
| 88 | } |
| 89 | }) |
| 90 | } |
| 91 | } |
| 92 | |
| 93 | func TestRequestObservationPreservesResponseBytesAndEOF(t *testing.T) { |
| 94 | const body = "data: {\"content\":\"hello\"}\n\ndata: [DONE]\n\n" |
| 95 | server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _, _ = io.WriteString(w, body) })) |
| 96 | defer server.Close() |
| 97 | var mu sync.Mutex |
| 98 | var last RequestObservation |
| 99 | ctx := WithRequestObserver(t.Context(), func(v RequestObservation) { mu.Lock(); last = v; mu.Unlock() }) |
| 100 | resp, err := SendWithRetry(ctx, server.Client(), SendOptions{}, func(ctx context.Context) (*http.Request, error) { |
| 101 | return http.NewRequestWithContext(ctx, http.MethodGet, server.URL, nil) |
| 102 | }) |
| 103 | if err != nil { |
| 104 | t.Fatal(err) |
| 105 | } |
| 106 | got, err := io.ReadAll(resp.Body) |
| 107 | _ = resp.Body.Close() |
| 108 | if err != nil || string(got) != body { |
| 109 | t.Fatalf("body changed: %q %v", got, err) |
| 110 | } |
| 111 | mu.Lock() |
| 112 | defer mu.Unlock() |
| 113 | if last.Phase != "body_eof" || last.BodyBytes != int64(len(body)) { |
| 114 | t.Fatalf("observation: %+v", last) |
| 115 | } |
| 116 | } |
| 117 |