From 0ba1992456098a74f8d780957fd0670e16bc5613 Mon Sep 17 00:00:00 2001 From: RKS Date: Sat, 12 Sep 2026 17:43:46 -0400 Subject: [PATCH] fix(go): propagate context to management HTTP requests --- go/tunnels/manager.go | 4 +- go/tunnels/manager_test.go | 103 +++++++++++++++++++++++++++++++++++++ go/tunnels/tunnels.go | 2 +- 3 files changed, 106 insertions(+), 3 deletions(-) diff --git a/go/tunnels/manager.go b/go/tunnels/manager.go index ce30e667..52a42b39 100644 --- a/go/tunnels/manager.go +++ b/go/tunnels/manager.go @@ -967,13 +967,13 @@ func (m *Manager) createRequest( partialFields []string, ) (*http.Request, error) { if requestObject == nil { - return http.NewRequest(method, uri.String(), nil) + return http.NewRequestWithContext(ctx, method, uri.String(), nil) } requestJson, err := partialMarshal(requestObject, partialFields) if err != nil { return nil, fmt.Errorf("error converting request object to json: %w", err) } - return http.NewRequest(method, uri.String(), bytes.NewBuffer(requestJson)) + return http.NewRequestWithContext(ctx, method, uri.String(), bytes.NewBuffer(requestJson)) } func (m *Manager) readProblemDetails(response *http.Response) (*string, error) { diff --git a/go/tunnels/manager_test.go b/go/tunnels/manager_test.go index 0867cc15..72e8b97a 100644 --- a/go/tunnels/manager_test.go +++ b/go/tunnels/manager_test.go @@ -6,14 +6,17 @@ package tunnels import ( "context" "encoding/json" + "errors" "fmt" "io" "log" "math/rand" "net/http" + "net/http/httptest" "net/url" "os" "strings" + "sync/atomic" "testing" "time" ) @@ -100,6 +103,106 @@ func responseWithStatus(status int, body string) *http.Response { } } +func TestManagerRequestContext(t *testing.T) { + for _, operation := range []struct { + name string + call func(context.Context, *Manager) error + }{ + {"list", func(ctx context.Context, manager *Manager) error { + _, err := manager.ListClusters(ctx) + return err + }}, + {"update", func(ctx context.Context, manager *Manager) error { + _, err := manager.UpdateTunnel(ctx, &Tunnel{Name: "test-tunnel"}, nil, nil) + return err + }}, + } { + t.Run(operation.name, func(t *testing.T) { + for _, test := range []struct { + name string + wantErr error + }{ + {"active", nil}, + {"canceled", context.Canceled}, + {"deadline", context.DeadlineExceeded}, + {"in-flight", context.Canceled}, + } { + t.Run(test.name, func(t *testing.T) { + var requests int32 + received := make(chan struct{}, 1) + release := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if _, err := io.Copy(io.Discard, r.Body); err != nil { + t.Errorf("error reading request body: %v", err) + return + } + atomic.AddInt32(&requests, 1) + received <- struct{}{} + if test.name == "in-flight" { + select { + case <-r.Context().Done(): + return + case <-release: + } + } + if r.Method == http.MethodGet { + fmt.Fprint(w, "[]") + } else { + fmt.Fprint(w, `{"name":"test-tunnel"}`) + } + })) + defer server.Close() + defer close(release) + uri, err := url.Parse(server.URL) + if err != nil { + t.Fatal(err) + } + client := server.Client() + client.Timeout = 5 * time.Second + manager, err := NewManager(userAgentManagerTest, nil, uri, client, "2023-09-27-preview") + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + if test.name == "deadline" { + cancel() + ctx, cancel = context.WithDeadline(context.Background(), time.Now().Add(-time.Second)) + } else if test.name == "canceled" { + cancel() + } + defer cancel() + + result := make(chan error, 1) + go func() { result <- operation.call(ctx, manager) }() + if test.name == "in-flight" { + select { + case <-received: + cancel() + case <-time.After(5 * time.Second): + t.Fatal("request did not reach server") + } + } + select { + case err = <-result: + if !errors.Is(err, test.wantErr) { + t.Fatalf("expected %v, got %v", test.wantErr, err) + } + case <-time.After(5 * time.Second): + t.Fatal("request did not complete") + } + wantRequests := int32(1) + if test.name == "canceled" || test.name == "deadline" { + wantRequests = 0 + } + if got := atomic.LoadInt32(&requests); got != wantRequests { + t.Fatalf("expected %d requests, got %d", wantRequests, got) + } + }) + } + }) + } +} + func tunnelIDFromPath(path string) string { trimmed := strings.Trim(path, "/") parts := strings.Split(trimmed, "/") diff --git a/go/tunnels/tunnels.go b/go/tunnels/tunnels.go index 5d89436a..61c4f390 100644 --- a/go/tunnels/tunnels.go +++ b/go/tunnels/tunnels.go @@ -10,7 +10,7 @@ import ( "github.com/rodaine/table" ) -const PackageVersion = "0.2.0" +const PackageVersion = "0.2.1" func (tunnel *Tunnel) requestObject() (*Tunnel, error) { convertedTunnel := &Tunnel{