Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions go/tunnels/manager.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
103 changes: 103 additions & 0 deletions go/tunnels/manager_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Expand Down Expand Up @@ -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, "/")
Expand Down
2 changes: 1 addition & 1 deletion go/tunnels/tunnels.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{
Expand Down