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
3 changes: 3 additions & 0 deletions internal/config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@ import (
"go.uber.org/zap"
"go.yaml.in/yaml/v4"
"golang.org/x/crypto/ssh"

"gateway/internal/util/useragent"
)

var (
Expand Down Expand Up @@ -251,6 +253,7 @@ func resolveTwingateHostname(targetURL, defaultHost string, retryMax int, logger
logger = logger.With(zap.String("url", targetURL), zap.String("defaultHost", defaultHost))

client := retryablehttp.NewClient()
client.HTTPClient.Transport = useragent.Transport{Base: client.HTTPClient.Transport}
client.HTTPClient.Timeout = 1 * time.Second
client.HTTPClient.CheckRedirect = func(_ *http.Request, _ []*http.Request) error {
return http.ErrUseLastResponse
Expand Down
20 changes: 20 additions & 0 deletions internal/config/config_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,26 @@ func TestResolveTwingateHostname(t *testing.T) {
assert.Equal(t, "twingate.com", result)
})

t.Run("identifies the gateway with a User-Agent header", func(t *testing.T) {
userAgents := make(chan string, 1)

server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
userAgents <- r.UserAgent()

w.WriteHeader(http.StatusOK)
}))
Comment thread
minhtule marked this conversation as resolved.
t.Cleanup(server.Close)

resolveTwingateHostname(server.URL+"/api/v1/jwk/ec", "twingate.com", 0, zap.NewNop())

select {
case userAgent := <-userAgents:
assert.Equal(t, "Twingate-Gateway/dev", userAgent)
default:
t.Fatal("hostname resolution endpoint was not requested")
}
})

t.Run("does not follow redirect", func(t *testing.T) {
shardServerCalled := make(chan struct{}, 1)

Expand Down
8 changes: 7 additions & 1 deletion internal/token/parser.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,14 @@
package token

import (
"context"
"fmt"
"net/http"

"github.com/MicahParks/keyfunc/v3"
"github.com/golang-jwt/jwt/v5"

"gateway/internal/util/useragent"
)

var allowedSigningMethods = []string{jwt.SigningMethodES256.Alg()}
Comment thread
minhtule marked this conversation as resolved.
Expand All @@ -30,7 +34,9 @@ type Parser struct {

func NewParser(config ParserConfig) (*Parser, error) {
if config.Keyfunc == nil {
jwks, err := keyfunc.NewDefault([]string{config.JWKSURL})
jwks, err := keyfunc.NewDefaultOverrideCtx(context.Background(), []string{config.JWKSURL}, keyfunc.Override{
Client: &http.Client{Transport: useragent.Transport{}},
})
Comment on lines +37 to +39

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I wonder if we should pass the ctx from the NewListener to here?

if err != nil {
return nil, fmt.Errorf("failed to create JWKS store: %w", err)
}
Expand Down
32 changes: 32 additions & 0 deletions internal/token/parser_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,13 @@ import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"net/http"
"net/http/httptest"
"testing"
"time"

"github.com/golang-jwt/jwt/v5"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

Expand Down Expand Up @@ -142,3 +145,32 @@ func TestNewParser(t *testing.T) {
})
}
}

func TestNewParser_RemoteJWKS(t *testing.T) {
userAgents := make(chan string, 1)

server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
select {
case userAgents <- r.UserAgent():
default:
}

_, _ = w.Write([]byte(`{"keys":[]}`))
}))
t.Cleanup(server.Close)

parser, err := NewParser(ParserConfig{
Issuer: "twingate",
Audience: "acme",
JWKSURL: server.URL + "/api/v1/jwk/ec",
})
require.NoError(t, err)
require.NotNil(t, parser)

select {
case userAgent := <-userAgents:
assert.Equal(t, "Twingate-Gateway/dev", userAgent)
default:
t.Fatal("JWKS endpoint was not requested")
}
}
Comment thread
minhtule marked this conversation as resolved.
32 changes: 32 additions & 0 deletions internal/util/useragent/useragent.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
// Copyright (c) Twingate Inc.
// SPDX-License-Identifier: MPL-2.0

package useragent

import (
"net/http"

"gateway/internal/version"
)

// String returns the User-Agent identifying this build in outbound HTTP requests.
func String() string {
return "Twingate-Gateway/" + version.Version
}

type Transport struct {
// Base is the transport the request is handed to. If nil, http.DefaultTransport is used.
Base http.RoundTripper
}

func (t Transport) RoundTrip(req *http.Request) (*http.Response, error) {
base := t.Base
if base == nil {
base = http.DefaultTransport
}

req = req.Clone(req.Context())
req.Header.Set("User-Agent", String())

return base.RoundTrip(req)
Comment thread
minhtule marked this conversation as resolved.
}
22 changes: 22 additions & 0 deletions internal/util/useragent/useragent_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
// Copyright (c) Twingate Inc.
// SPDX-License-Identifier: MPL-2.0

package useragent

import (
"testing"

"github.com/stretchr/testify/assert"

"gateway/internal/version"
)

func TestString(t *testing.T) {
original := version.Version

t.Cleanup(func() { version.Version = original })

version.Version = "1.2.3"

assert.Equal(t, "Twingate-Gateway/1.2.3", String())
}