diff --git a/src/cmd/turnkey/pkg/auth.go b/src/cmd/turnkey/pkg/auth.go index 0e34c70..07d191c 100644 --- a/src/cmd/turnkey/pkg/auth.go +++ b/src/cmd/turnkey/pkg/auth.go @@ -3,6 +3,8 @@ package pkg import ( "regexp" + "github.com/go-openapi/runtime" + httptransport "github.com/go-openapi/runtime/client" "github.com/rotisserie/eris" "github.com/tkhq/go-sdk" @@ -59,11 +61,26 @@ func LoadClient() { if pattern := regexp.MustCompile(`^localhost:\d+$`); pattern.MatchString(apiHost) { scheme = "http" } - transportConfig := client.DefaultTransportConfig().WithHost(apiHost).WithSchemes([]string{scheme}) - APIClient = &sdk.Client{ - Client: client.NewHTTPClientWithConfig(nil, transportConfig), + Client: newAPIClient(scheme, apiHost), Authenticator: &sdk.Authenticator{Key: APIKeypair}, APIKey: APIKeypair, } } + +func newAPIClient(scheme string, host string) *client.TurnkeyAPI { + transport := httptransport.NewWithClient(host, client.DefaultBasePath, []string{scheme}, newHTTPClient()) + + return client.New(redirectTransport{transport}, nil) +} + +type redirectTransport struct { + runtime.ClientTransport +} + +func (t redirectTransport) Submit(operation *runtime.ClientOperation) (interface{}, error) { + request := *operation + request.Client = nil + + return t.ClientTransport.Submit(&request) +} diff --git a/src/cmd/turnkey/pkg/auth_test.go b/src/cmd/turnkey/pkg/auth_test.go new file mode 100644 index 0000000..1f05963 --- /dev/null +++ b/src/cmd/turnkey/pkg/auth_test.go @@ -0,0 +1,36 @@ +package pkg + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/tkhq/go-sdk/pkg/api/client/sessions" +) + +func TestAPIClientDoesNotFollowRedirectToOtherHost(t *testing.T) { + var otherRequests atomic.Int32 + + other := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + otherRequests.Add(1) + })) + defer other.Close() + + origin := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, other.URL, http.StatusTemporaryRedirect) + })) + defer origin.Close() + + apiClient := newAPIClient("http", strings.TrimPrefix(origin.URL, "http://")) + + _, err := apiClient.Sessions.GetWhoami(&sessions.GetWhoamiParams{ + Context: context.Background(), + HTTPClient: &http.Client{}, + }, nil) + assert.Error(t, err) + assert.Equal(t, int32(0), otherRequests.Load()) +} diff --git a/src/cmd/turnkey/pkg/redirect.go b/src/cmd/turnkey/pkg/redirect.go new file mode 100644 index 0000000..9b680e9 --- /dev/null +++ b/src/cmd/turnkey/pkg/redirect.go @@ -0,0 +1,78 @@ +package pkg + +import ( + "net/http" + "net/netip" + "net/url" + "strings" + + "github.com/rotisserie/eris" + "golang.org/x/net/idna" +) + +func newHTTPClient() *http.Client { + return &http.Client{CheckRedirect: checkRedirect} +} + +func checkRedirect(req *http.Request, via []*http.Request) error { + if len(via) > 10 { + return eris.New("stopped after 10 redirects") + } + + status := 0 + if req.Response != nil { + status = req.Response.StatusCode + } + + if status != http.StatusTemporaryRedirect && status != http.StatusPermanentRedirect { + return eris.Errorf("refusing to follow redirect with status %d; only 307 and 308 are supported", status) + } + + origin, err := effectiveOrigin(via[0].URL) + if err != nil { + return err + } + + target, err := effectiveOrigin(req.URL) + if err != nil { + return err + } + + if target != origin { + return eris.Errorf("refusing to follow redirect from %s to %s", origin, target) + } + + return nil +} + +func effectiveOrigin(u *url.URL) (string, error) { + scheme := strings.ToLower(u.Scheme) + + port := u.Port() + if port == "" { + switch scheme { + case "http": + port = "80" + case "https": + port = "443" + default: + return "", eris.Errorf("no default port for scheme %q", scheme) + } + } + + host := strings.ToLower(u.Hostname()) + if addr, err := netip.ParseAddr(host); err == nil { + host = addr.Unmap().String() + if strings.Contains(host, ":") { + host = "[" + host + "]" + } + } else { + ascii, err := idna.Lookup.ToASCII(host) + if err != nil { + return "", eris.Wrapf(err, "failed to normalize host %q", u.Hostname()) + } + host = ascii + } + + return scheme + "://" + host + ":" + port, nil +} diff --git a/src/cmd/turnkey/pkg/redirect_test.go b/src/cmd/turnkey/pkg/redirect_test.go new file mode 100644 index 0000000..c9e065d --- /dev/null +++ b/src/cmd/turnkey/pkg/redirect_test.go @@ -0,0 +1,41 @@ +package pkg + +import ( + "net/http" + "net/url" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestCheckRedirectLimit(t *testing.T) { + origin, err := url.Parse("https://example.com") + require.NoError(t, err) + + redirect := &http.Request{ + URL: origin, + Response: &http.Response{StatusCode: http.StatusTemporaryRedirect}, + } + previous := make([]*http.Request, 11) + previous[0] = &http.Request{URL: origin} + + require.NoError(t, checkRedirect(redirect, previous[:10])) + assert.EqualError(t, checkRedirect(redirect, previous), "stopped after 10 redirects") +} + +func TestEffectiveOrigin(t *testing.T) { + for rawURL, expected := range map[string]string{ + "https://EXAMPLE.com/path": "https://example.com:443", + "http://example.com:80": "http://example.com:80", + "http://[0:0:0:0:0:0:0:1]:8080/x": "http://[::1]:8080", + "https://BÜCHER.example": "https://xn--bcher-kva.example:443", + } { + parsed, err := url.Parse(rawURL) + require.NoError(t, err) + + origin, err := effectiveOrigin(parsed) + require.NoError(t, err) + assert.Equal(t, expected, origin, rawURL) + } +} diff --git a/src/cmd/turnkey/pkg/request.go b/src/cmd/turnkey/pkg/request.go index 6d9f640..bed774a 100644 --- a/src/cmd/turnkey/pkg/request.go +++ b/src/cmd/turnkey/pkg/request.go @@ -111,9 +111,7 @@ func post(ctx context.Context, protocol string, host string, path string, body [ req.Header.Set("X-Stamp", stamp) - client := http.Client{} - - response, err := client.Do(req) + response, err := newHTTPClient().Do(req) if err != nil { return nil, eris.Wrap(err, "error while sending HTTP POST request") } diff --git a/src/cmd/turnkey/pkg/request_test.go b/src/cmd/turnkey/pkg/request_test.go new file mode 100644 index 0000000..5d620f9 --- /dev/null +++ b/src/cmd/turnkey/pkg/request_test.go @@ -0,0 +1,71 @@ +package pkg + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestPostFollowsRedirectOnSameHost(t *testing.T) { + mux := http.NewServeMux() + mux.HandleFunc("/initial", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "/final", http.StatusTemporaryRedirect) + }) + mux.HandleFunc("/final", func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + require.NoError(t, err) + assert.Equal(t, "test-stamp", r.Header.Get("X-Stamp")) + assert.Equal(t, []byte(`{"a":1}`), body) + }) + + server := httptest.NewServer(mux) + defer server.Close() + + response, err := post(context.Background(), "http", strings.TrimPrefix(server.URL, "http://"), "/initial", []byte(`{"a":1}`), "test-stamp") + require.NoError(t, err) + require.NoError(t, response.Body.Close()) +} + +func TestPostDoesNotFollowMethodChangingRedirect(t *testing.T) { + var finalRequests atomic.Int32 + + mux := http.NewServeMux() + mux.HandleFunc("/initial", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "/final", http.StatusFound) + }) + mux.HandleFunc("/final", func(w http.ResponseWriter, r *http.Request) { + finalRequests.Add(1) + }) + + server := httptest.NewServer(mux) + defer server.Close() + + _, err := post(context.Background(), "http", strings.TrimPrefix(server.URL, "http://"), "/initial", []byte(`{"a":1}`), "test-stamp") + assert.Error(t, err) + assert.Equal(t, int32(0), finalRequests.Load()) +} + +func TestPostDoesNotFollowRedirectToOtherHost(t *testing.T) { + var otherRequests atomic.Int32 + + other := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + otherRequests.Add(1) + })) + defer other.Close() + + origin := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, other.URL, http.StatusTemporaryRedirect) + })) + defer origin.Close() + + _, err := post(context.Background(), "http", strings.TrimPrefix(origin.URL, "http://"), "/", []byte(`{"a":1}`), "test-stamp") + assert.Error(t, err) + assert.Equal(t, int32(0), otherRequests.Load()) +} diff --git a/src/go.mod b/src/go.mod index 77079dd..41cb69b 100644 --- a/src/go.mod +++ b/src/go.mod @@ -6,12 +6,14 @@ toolchain go1.21.0 require ( github.com/btcsuite/btcutil v1.0.2 + github.com/go-openapi/runtime v0.26.0 github.com/google/uuid v1.3.1 github.com/rotisserie/eris v0.5.4 github.com/spf13/cobra v1.7.0 github.com/stretchr/testify v1.8.4 github.com/tkhq/go-sdk v0.0.0-20240813203011-ed45fe0d5c27 github.com/tkhq/go-sdk/pkg/enclave_encrypt v0.0.0-20240513225018-5ebfb539ec1e + golang.org/x/net v0.24.0 gopkg.in/yaml.v3 v3.0.1 ) @@ -28,7 +30,6 @@ require ( github.com/go-openapi/jsonpointer v0.20.0 // indirect github.com/go-openapi/jsonreference v0.20.2 // indirect github.com/go-openapi/loads v0.21.2 // indirect - github.com/go-openapi/runtime v0.26.0 // indirect github.com/go-openapi/spec v0.20.9 // indirect github.com/go-openapi/strfmt v0.21.7 // indirect github.com/go-openapi/swag v0.22.4 // indirect @@ -48,5 +49,6 @@ require ( go.opentelemetry.io/otel/trace v1.19.0 // indirect golang.org/x/crypto v0.22.0 // indirect golang.org/x/sys v0.20.0 // indirect + golang.org/x/text v0.14.0 // indirect gopkg.in/yaml.v2 v2.4.0 // indirect ) diff --git a/src/go.sum b/src/go.sum index bab0f42..10c3131 100644 --- a/src/go.sum +++ b/src/go.sum @@ -189,8 +189,6 @@ github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o github.com/stretchr/testify v1.8.4 h1:CcVxjf3Q8PM0mHUKJCdn+eZZtm5yQwehR5yeSVQQcUk= github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= github.com/tidwall/pretty v1.0.0/go.mod h1:XNkn88O1ChpSDQmQeStsy+sBenx6DDtFZJxhVysOjyk= -github.com/tkhq/go-sdk v0.0.0-20240813182504-228a50933080 h1:Yhc2J2GCB0SDbLBVwK1ZlrYNiHVuwHGCU+N9CdJz4WQ= -github.com/tkhq/go-sdk v0.0.0-20240813182504-228a50933080/go.mod h1:NgCPbnpGdhx+31NLwmK3iC6UftT7I70dbKXVbblVpjk= github.com/tkhq/go-sdk v0.0.0-20240813203011-ed45fe0d5c27 h1:1Tm6Z2uD9THuycnXtkNbTMf07Owdm071fV5JcKLsAQE= github.com/tkhq/go-sdk v0.0.0-20240813203011-ed45fe0d5c27/go.mod h1:2372WQ2x5SWlXmFBygP8PaNcR225Pn8Nd2WmzT9E35Y= github.com/tkhq/go-sdk/pkg/enclave_encrypt v0.0.0-20240513225018-5ebfb539ec1e h1:6TQn08QGF615Bt2LRNv1MwlI5qL9NlpO2A/DIKX8MUo= @@ -236,6 +234,8 @@ golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v golang.org/x/net v0.0.0-20210421230115-4e50805a0758/go.mod h1:72T/g9IO56b78aLF+1Kcs5dz7/ng1VjMUvfKvpfy+jM= golang.org/x/net v0.0.0-20211112202133-69e39bad7dc2/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= +golang.org/x/net v0.24.0 h1:1PcaxkF854Fu3+lvBIx5SYn9wRlBzzcnHZSiaFFAb0w= +golang.org/x/net v0.24.0/go.mod h1:2Q7sJY5mzlzWjKtYUEXSlBWCdyaioyXzRB2RtU8KVE8= golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190227155943-e225da77a7e6/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190412183630-56d357773e84/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= @@ -268,6 +268,8 @@ golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= golang.org/x/text v0.3.8/go.mod h1:E6s5w1FMmriuDzIBO73fBruAKo1PCIq6d2Q6DHfQ8WQ= golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= +golang.org/x/text v0.14.0 h1:ScX5w1eTa3QqT8oi6+ziP7dTV1S2+ALU0bI+0zXKWiQ= +golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20190329151228-23e29df326fe/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= golang.org/x/tools v0.0.0-20190416151739-9c9e1878f421/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs=