diff --git a/go/tunnels/manager.go b/go/tunnels/manager.go index 263c3f6b..ce30e667 100644 --- a/go/tunnels/manager.go +++ b/go/tunnels/manager.go @@ -14,6 +14,7 @@ import ( "net/http" "net/url" "reflect" + "sort" "strings" ) @@ -112,6 +113,55 @@ type Manager struct { OnClusterSelected func(ClusterSelection) } +func withPreconditionHeader( + options *TunnelRequestOptions, + name string, + value string, + preserveExisting bool, +) *TunnelRequestOptions { + if options == nil { + options = &TunnelRequestOptions{} + } + + result := *options + result.AdditionalHeaders = make(map[string]string, len(options.AdditionalHeaders)+1) + headers := sortedHeaderNames(options.AdditionalHeaders) + existingValue := "" + existingValueFound := false + for _, header := range headers { + headerValue := options.AdditionalHeaders[header] + if strings.EqualFold(header, name) && preserveExisting { + existingValue = headerValue + existingValueFound = true + } else if !strings.EqualFold(header, "If-Match") && + !strings.EqualFold(header, "If-None-Match") && + !strings.EqualFold(header, "If-Not-Match") { + result.AdditionalHeaders[header] = headerValue + } + } + if existingValueFound { + result.AdditionalHeaders[name] = existingValue + } else { + result.AdditionalHeaders[name] = value + } + return &result +} + +func sortedHeaderNames(headers map[string]string) []string { + names := make([]string, 0, len(headers)) + for name := range headers { + names = append(names, name) + } + sort.Strings(names) + return names +} + +func setHeaders(destination http.Header, headers map[string]string) { + for _, name := range sortedHeaderNames(headers) { + destination.Set(name, headers[name]) + } +} + // Creates a new Manager used for interacting with the Tunnels APIs. // tokenProvider is an optional paramater containing a function that returns the access token to use for the request. // If no tunnelServiceUrl or httpClient is provided, the default values will be used. @@ -233,13 +283,7 @@ func (m *Manager) CreateTunnel(ctx context.Context, tunnel *Tunnel, options *Tun idGenerated = true } - if options == nil { - options = &TunnelRequestOptions{} - } - if options.AdditionalHeaders == nil { - options.AdditionalHeaders = map[string]string{} - } - options.AdditionalHeaders["If-Not-Match"] = "*" + options = withPreconditionHeader(options, "If-None-Match", "*", false) // If the caller didn't specify a cluster, auto-select one via the // recommendations API. Failures fall back to global routing. @@ -279,8 +323,8 @@ func (m *Manager) CreateTunnel(ctx context.Context, tunnel *Tunnel, options *Tun // otherwise invisible in service telemetry: a create that fell back looks identical to // one that was never recommended at all. // - // This is passed explicitly rather than via options.AdditionalHeaders because that map - // is never read when the request is built. + // This is passed explicitly so SDK-generated telemetry takes precedence over a + // caller-provided header with the same name. clusterSourceHeader := map[string]string{clusterSourceHeaderName: string(clusterSource)} convertedTunnel, err := tunnel.requestObject() @@ -349,13 +393,7 @@ func (m *Manager) UpdateTunnel(ctx context.Context, tunnel *Tunnel, updateFields } } - if options == nil { - options = &TunnelRequestOptions{} - } - if options.AdditionalHeaders == nil { - options.AdditionalHeaders = map[string]string{} - } - options.AdditionalHeaders["If-Match"] = "*" + options = withPreconditionHeader(options, "If-Match", "*", true) url, err := m.buildTunnelSpecificUri(tunnel, "", options, "", false) if err != nil { @@ -560,13 +598,7 @@ func (m *Manager) GetTunnelPort( func (m *Manager) CreateTunnelPort( ctx context.Context, tunnel *Tunnel, port *TunnelPort, options *TunnelRequestOptions, ) (tp *TunnelPort, err error) { - if options == nil { - options = &TunnelRequestOptions{} - } - if options.AdditionalHeaders == nil { - options.AdditionalHeaders = map[string]string{} - } - options.AdditionalHeaders["If-Not-Match"] = "*" + options = withPreconditionHeader(options, "If-None-Match", "*", false) path := fmt.Sprintf("%s/%d", portsApiSubPath, port.PortNumber) url, err := m.buildTunnelSpecificUri(tunnel, path, options, "", false) @@ -612,13 +644,7 @@ func (m *Manager) UpdateTunnelPort( return nil, fmt.Errorf("cluster ids do not match") } - if options == nil { - options = &TunnelRequestOptions{} - } - if options.AdditionalHeaders == nil { - options.AdditionalHeaders = map[string]string{} - } - options.AdditionalHeaders["If-Match"] = "*" + options = withPreconditionHeader(options, "If-Match", "*", true) path := fmt.Sprintf("%s/%d", portsApiSubPath, port.PortNumber) url, err := m.buildTunnelSpecificUri(tunnel, path, options, "", false) @@ -852,11 +878,17 @@ func (m *Manager) sendTunnelRequest( extraHeaders ...map[string]string, ) ([]byte, error) { authHeaderValue := m.getAccessToken(tunnel, tunnelRequestOptions, accessTokenScopes) - return m.sendRequest(ctx, method, uri, requestObject, partialFields, authHeaderValue, allowNotFound, extraHeaders...) + requestHeaders := make([]map[string]string, 0, len(extraHeaders)+1) + if tunnelRequestOptions != nil { + requestHeaders = append(requestHeaders, tunnelRequestOptions.AdditionalHeaders) + } + requestHeaders = append(requestHeaders, extraHeaders...) + return m.sendRequest( + ctx, method, uri, requestObject, partialFields, authHeaderValue, allowNotFound, requestHeaders...) } -// sendRequest sends a request to the service. extraHeaders is optional and adds -// per-request headers on top of the manager-wide ones. +// sendRequest sends a request to the service. requestHeaders are applied in order +// on top of manager-wide headers. func (m *Manager) sendRequest( ctx context.Context, method string, @@ -865,19 +897,26 @@ func (m *Manager) sendRequest( partialFields []string, authHeaderValue string, allowNotFound bool, - extraHeaders ...map[string]string, + requestHeaders ...map[string]string, ) ([]byte, error) { request, err := m.createRequest(ctx, method, uri, requestObject, partialFields) if err != nil { return nil, fmt.Errorf("error creating request: %w", err) } - // Add authorization header + // Add manager-level and per-request headers before SDK-owned headers so callers + // cannot replace authentication or SDK identity values. + setHeaders(request.Header, m.additionalHeaders) + for _, headers := range requestHeaders { + setHeaders(request.Header, headers) + } + if authHeaderValue != "" { - request.Header.Add("Authorization", authHeaderValue) + request.Header.Set("Authorization", authHeaderValue) + } else { + request.Header.Del("Authorization") } - // Add user agent header userAgentString := "" for _, userAgent := range m.userAgents { if len(userAgent.Version) == 0 { @@ -889,19 +928,8 @@ func (m *Manager) sendRequest( userAgentString = fmt.Sprintf("%s%s/%s ", userAgentString, userAgent.Name, userAgent.Version) } userAgentString = strings.TrimSpace(userAgentString) - request.Header.Add("User-Agent", fmt.Sprintf("%s %s", goUserAgent, userAgentString)) - request.Header.Add("Content-Type", "application/json;charset=UTF-8") - - // Add additional headers - for header, headerValue := range m.additionalHeaders { - request.Header.Add(header, headerValue) - } - - for _, headers := range extraHeaders { - for header, headerValue := range headers { - request.Header.Set(header, headerValue) - } - } + request.Header.Set("User-Agent", fmt.Sprintf("%s %s", goUserAgent, userAgentString)) + request.Header.Set("Content-Type", "application/json;charset=UTF-8") result, err := m.httpClient.Do(request) if err != nil { diff --git a/go/tunnels/manager_test.go b/go/tunnels/manager_test.go index 0a537419..0867cc15 100644 --- a/go/tunnels/manager_test.go +++ b/go/tunnels/manager_test.go @@ -131,6 +131,12 @@ func TestCreateTunnelRetriesOnConflictForGeneratedID(t *testing.T) { if strings.Contains(r.URL.Path, "recommendations") { return responseWithStatus(http.StatusOK, "{}"), nil } + if r.Header.Get("If-None-Match") != "*" { + t.Fatalf("expected If-None-Match header on every create attempt") + } + if r.Header.Get("If-Not-Match") != "" { + t.Fatalf("did not expect If-Not-Match header") + } callCount++ pathIDs = append(pathIDs, tunnelIDFromPath(r.URL.Path)) bodyBytes, readErr := io.ReadAll(r.Body) @@ -180,6 +186,254 @@ func TestCreateTunnelRetriesOnConflictForGeneratedID(t *testing.T) { } } +func TestTunnelRequestHeadersAndPreconditions(t *testing.T) { + url, err := url.Parse("https://example.test/") + if err != nil { + t.Fatalf("error parsing url: %v", err) + } + + var requests []*http.Request + client := &http.Client{Transport: roundTripperFunc(func(r *http.Request) (*http.Response, error) { + requests = append(requests, r.Clone(r.Context())) + switch r.Method { + case http.MethodPut: + return responseWithStatus( + http.StatusOK, + "{\"tunnelId\":\"tunnel-id\",\"clusterId\":\"usw2\"}"), nil + case http.MethodDelete: + return responseWithStatus(http.StatusOK, "true"), nil + default: + return responseWithStatus(http.StatusMethodNotAllowed, ""), nil + } + })} + + managementClient, err := NewManager( + userAgentManagerTest, + getUserToken, + url, + client, + "2023-09-27-preview") + if err != nil { + t.Fatalf("error creating manager: %v", err) + } + managementClient.additionalHeaders = map[string]string{ + "X-Manager": "manager", + "X-Override": "manager", + } + + options := &TunnelRequestOptions{ + AccessToken: "expected-token", + AdditionalHeaders: map[string]string{ + "Authorization": "caller-token", + "Content-Type": "caller-content-type", + "User-Agent": "caller-user-agent", + "X-Override": "request", + "X-Request": "request", + "X-Case-Duplicate": "canonical", + "x-case-duplicate": "lowercase", + "if-match": "caller-match", + "IF-NONE-MATCH": "caller-none-match", + "If-Not-Match": "caller-nonstandard", + clusterSourceHeaderName: "caller-source", + }, + } + tunnel := &Tunnel{TunnelID: "tunnel-id", ClusterID: "usw2"} + + if _, err = managementClient.CreateTunnel(context.Background(), tunnel, options); err != nil { + t.Fatalf("unexpected create error: %v", err) + } + if _, err = managementClient.UpdateTunnel(context.Background(), tunnel, nil, options); err != nil { + t.Fatalf("unexpected update error: %v", err) + } + if err = managementClient.DeleteTunnel(context.Background(), tunnel, options); err != nil { + t.Fatalf("unexpected delete error: %v", err) + } + + if len(requests) != 3 { + t.Fatalf("expected 3 requests, got %d", len(requests)) + } + for _, request := range requests { + if request.Header.Get("X-Manager") != "manager" { + t.Fatalf("expected manager-level header") + } + if request.Header.Get("X-Override") != "request" { + t.Fatalf("expected per-request header to override manager-level header") + } + if request.Header.Get("X-Request") != "request" { + t.Fatalf("expected per-request header") + } + if request.Header.Get("X-Case-Duplicate") != "lowercase" { + t.Fatalf("expected deterministic resolution of case-only duplicate headers") + } + if request.Header.Get("Authorization") != "Tunnel expected-token" { + t.Fatalf("expected SDK authorization header") + } + if request.Header.Get("Content-Type") != "application/json;charset=UTF-8" { + t.Fatalf("expected SDK content type") + } + if request.Header.Get("User-Agent") == "caller-user-agent" { + t.Fatalf("expected SDK user agent") + } + } + + if requests[0].Header.Get("If-None-Match") != "*" || + requests[0].Header.Get("If-Match") != "" || + requests[0].Header.Get("If-Not-Match") != "" { + t.Fatalf("expected create-only precondition on create") + } + if requests[0].Header.Get(clusterSourceHeaderName) != string(ClusterSourceExplicit) { + t.Fatalf("expected SDK cluster-source header") + } + if requests[1].Header.Get("If-Match") != "caller-match" || + requests[1].Header.Get("If-None-Match") != "" || + requests[1].Header.Get("If-Not-Match") != "" { + t.Fatalf("expected update-only precondition on update") + } + if requests[2].Header.Get("If-Match") != "caller-match" || + requests[2].Header.Get("If-None-Match") != "caller-none-match" || + requests[2].Header.Get("If-Not-Match") != "caller-nonstandard" { + t.Fatalf("expected delete to preserve caller-provided preconditions") + } + if options.AdditionalHeaders["IF-NONE-MATCH"] != "caller-none-match" { + t.Fatalf("create mutated caller-owned If-None-Match header") + } + if options.AdditionalHeaders["if-match"] != "caller-match" { + t.Fatalf("update mutated caller-owned If-Match header") + } +} + +func TestCreateOrUpdateTunnelDoesNotAddPrecondition(t *testing.T) { + url, err := url.Parse("https://example.test/") + if err != nil { + t.Fatalf("error parsing url: %v", err) + } + + var requests []*http.Request + client := &http.Client{Transport: roundTripperFunc(func(r *http.Request) (*http.Response, error) { + requests = append(requests, r.Clone(r.Context())) + if len(requests) < createNameRetries { + return responseWithStatus(http.StatusConflict, ""), nil + } + return responseWithStatus( + http.StatusOK, + "{\"tunnelId\":\"tunnel-id\",\"clusterId\":\"usw2\"}"), nil + })} + + managementClient, err := NewManager( + userAgentManagerTest, + getUserToken, + url, + client, + "2023-09-27-preview") + if err != nil { + t.Fatalf("error creating manager: %v", err) + } + + options := &TunnelRequestOptions{ + AdditionalHeaders: map[string]string{"X-Request": "request"}, + } + tunnel := &Tunnel{TunnelID: "tunnel-id", ClusterID: "usw2"} + if _, err = managementClient.CreateOrUpdateTunnel( + context.Background(), tunnel, options); err != nil { + t.Fatalf("unexpected create-or-update error: %v", err) + } + + if len(requests) != createNameRetries { + t.Fatalf("expected %d requests, got %d", createNameRetries, len(requests)) + } + for _, request := range requests { + if request.Header.Get("If-Match") != "" || + request.Header.Get("If-None-Match") != "" || + request.Header.Get("If-Not-Match") != "" { + t.Fatalf("did not expect an automatic precondition on create-or-update") + } + if request.Header.Get("X-Request") != "request" { + t.Fatalf("expected per-request header on create-or-update") + } + } +} + +func TestTunnelPortRequestPreconditions(t *testing.T) { + url, err := url.Parse("https://example.test/") + if err != nil { + t.Fatalf("error parsing url: %v", err) + } + + var requests []*http.Request + client := &http.Client{Transport: roundTripperFunc(func(r *http.Request) (*http.Response, error) { + requests = append(requests, r.Clone(r.Context())) + return responseWithStatus(http.StatusOK, "{\"portNumber\":3000}"), nil + })} + + managementClient, err := NewManager( + userAgentManagerTest, + getUserToken, + url, + client, + "2023-09-27-preview") + if err != nil { + t.Fatalf("error creating manager: %v", err) + } + + options := &TunnelRequestOptions{ + AdditionalHeaders: map[string]string{ + "Authorization": "caller-token", + "X-Request": "request", + }, + } + tunnel := &Tunnel{TunnelID: "tunnel-id", ClusterID: "usw2"} + port := &TunnelPort{PortNumber: 3000} + + if _, err = managementClient.CreateTunnelPort( + context.Background(), tunnel, port, options); err != nil { + t.Fatalf("unexpected create port error: %v", err) + } + if _, err = managementClient.UpdateTunnelPort( + context.Background(), tunnel, port, nil, options); err != nil { + t.Fatalf("unexpected update port error: %v", err) + } + if _, err = managementClient.CreateOrUpdateTunnelPort( + context.Background(), tunnel, port, options); err != nil { + t.Fatalf("unexpected create-or-update port error: %v", err) + } + if err = managementClient.DeleteTunnelPort( + context.Background(), tunnel, port.PortNumber, options); err != nil { + t.Fatalf("unexpected delete port error: %v", err) + } + + if len(requests) != 4 { + t.Fatalf("expected 4 requests, got %d", len(requests)) + } + if requests[0].Header.Get("If-None-Match") != "*" || + requests[0].Header.Get("If-Not-Match") != "" { + t.Fatalf("expected create-only precondition on create port") + } + if requests[1].Header.Get("If-Match") != "*" || + requests[1].Header.Get("If-None-Match") != "" { + t.Fatalf("expected update-only precondition on update port") + } + for i, request := range requests[2:] { + if request.Header.Get("If-Match") != "" || + request.Header.Get("If-None-Match") != "" { + t.Fatalf("did not expect an automatic precondition on request %d", i+2) + } + } + for _, request := range requests { + if request.Header.Get("X-Request") != "request" { + t.Fatalf("expected per-request header on port requests") + } + if request.Header.Get("Authorization") != "" { + t.Fatalf("expected caller authorization header to be removed") + } + } + if _, ok := options.AdditionalHeaders["If-None-Match"]; ok { + t.Fatalf("create port mutated caller-owned headers") + } + if _, ok := options.AdditionalHeaders["If-Match"]; ok { + t.Fatalf("update port mutated caller-owned headers") + } +} + func TestCreateTunnelDoesNotRetryOnNonConflict(t *testing.T) { url, err := url.Parse("https://example.test/") if err != nil { diff --git a/go/tunnels/request_options.go b/go/tunnels/request_options.go index 4882d179..09a6bd7c 100644 --- a/go/tunnels/request_options.go +++ b/go/tunnels/request_options.go @@ -13,7 +13,8 @@ type TunnelRequestOptions struct { // Token used for authentication for service. AccessToken string - // Additional headers to be included in the request. + // Additional headers to be included in the request. SDK-owned Authorization, + // User-Agent, and Content-Type headers cannot be overridden. AdditionalHeaders map[string]string // Additional qurey parameters to be included in the request. diff --git a/go/tunnels/tunnels.go b/go/tunnels/tunnels.go index 8ecb62f3..5d89436a 100644 --- a/go/tunnels/tunnels.go +++ b/go/tunnels/tunnels.go @@ -10,7 +10,7 @@ import ( "github.com/rodaine/table" ) -const PackageVersion = "0.1.28" +const PackageVersion = "0.2.0" func (tunnel *Tunnel) requestObject() (*Tunnel, error) { convertedTunnel := &Tunnel{