Skip to content

Commit fa0d31a

Browse files
authored
Oauth/manager/introduce ticks (#160)
* security/oauth/manager: introduce ticks to check expiration of tokens
1 parent 1be6312 commit fa0d31a

2 files changed

Lines changed: 50 additions & 33 deletions

File tree

security/oauth/manager/config.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@ type Config struct {
1818
Endpoint Endpoint `envconfig:"ENDPOINT" env:"ENDPOINT"`
1919
Audience string `envconfig:"AUDIENCE" env:"AUDIENCE"`
2020
RequestTimeout time.Duration `envconfig:"REQUEST_TIMEOUT" env:"REQUEST_TIMEOUT" default:"10s"`
21+
TickFrequency time.Duration `envconfig:"TICK_FREQUENCY" env:"TICK_FREQUENCY" long:"tick-frequency" description:"how frequently we should check whether our token needs renewal" default:"15s"`
2122
}
2223

2324
// ToClientCrendtials converts to clientcredentials.Config

security/oauth/manager/manager.go

Lines changed: 49 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -10,36 +10,50 @@ import (
1010
"golang.org/x/oauth2/clientcredentials"
1111

1212
"github.com/plgd-dev/kit/log"
13-
"github.com/plgd-dev/kit/net/http/transport"
1413
"golang.org/x/oauth2"
1514
)
1615

1716
// Manager holds certificates from filesystem watched for changes
1817
type Manager struct {
19-
mutex sync.Mutex
20-
config clientcredentials.Config
21-
tlsCfg *tls.Config
22-
requestTimeout time.Duration
23-
token *oauth2.Token
24-
tokenErr error
25-
doneWg sync.WaitGroup
26-
done chan struct{}
18+
mutex sync.Mutex
19+
config clientcredentials.Config
20+
requestTimeout time.Duration
21+
tickFrequency time.Duration
22+
startRefreshToken time.Time
23+
token *oauth2.Token
24+
httpClient *http.Client
25+
tokenErr error
26+
doneWg sync.WaitGroup
27+
done chan struct{}
2728
}
2829

2930
// NewManagerFromConfiguration creates a new oauth manager which refreshing token.
3031
func NewManagerFromConfiguration(config Config, tlsCfg *tls.Config) (*Manager, error) {
3132
cfg := config.ToClientCrendtials()
32-
token, err := getToken(cfg, tlsCfg, config.RequestTimeout)
33+
t := http.DefaultTransport.(*http.Transport).Clone()
34+
t.MaxIdleConns = 1
35+
t.MaxConnsPerHost = 1
36+
t.MaxIdleConnsPerHost = 1
37+
t.IdleConnTimeout = time.Second * 30
38+
t.TLSClientConfig = tlsCfg
39+
httpClient := &http.Client{
40+
Transport: t,
41+
Timeout: config.RequestTimeout,
42+
}
43+
token, startRefreshToken, err := getToken(cfg, httpClient, config.RequestTimeout)
3344
if err != nil {
3445
return nil, err
3546
}
36-
mgr := &Manager{
37-
config: cfg,
38-
token: token,
39-
tlsCfg: tlsCfg,
4047

41-
requestTimeout: config.RequestTimeout,
42-
done: make(chan struct{}),
48+
mgr := &Manager{
49+
config: cfg,
50+
token: token,
51+
startRefreshToken: startRefreshToken,
52+
requestTimeout: config.RequestTimeout,
53+
httpClient: httpClient,
54+
tickFrequency: config.TickFrequency,
55+
56+
done: make(chan struct{}),
4357
}
4458
mgr.doneWg.Add(1)
4559

@@ -63,48 +77,50 @@ func (a *Manager) Close() {
6377
}
6478
}
6579

66-
func (a *Manager) nextRenewal() time.Duration {
67-
t, _ := a.GetToken(context.Background())
68-
now := time.Now()
69-
lifetime := t.Expiry.Sub(now) * 2 / 3
70-
if lifetime < a.requestTimeout {
71-
lifetime = a.requestTimeout
72-
}
73-
return lifetime
80+
func (a *Manager) shouldRefresh() bool {
81+
return time.Now().After(a.startRefreshToken)
7482
}
7583

76-
func getToken(cfg clientcredentials.Config, tlsCfg *tls.Config, requestTimeout time.Duration) (*oauth2.Token, error) {
84+
func getToken(cfg clientcredentials.Config, httpClient *http.Client, requestTimeout time.Duration) (*oauth2.Token, time.Time, error) {
7785
ctx, cancel := context.WithTimeout(context.Background(), requestTimeout)
7886
defer cancel()
7987

80-
t := transport.NewDefaultTransport()
81-
t.TLSClientConfig = tlsCfg
82-
httpClient := &http.Client{Transport: t}
8388
ctx = context.WithValue(ctx, oauth2.HTTPClient, httpClient)
8489

85-
return cfg.Token(ctx)
90+
token, err := cfg.Token(ctx)
91+
var startRefreshToken time.Time
92+
if err == nil {
93+
now := time.Now()
94+
startRefreshToken = now.Add(token.Expiry.Sub(now) * 2 / 3)
95+
}
96+
return token, startRefreshToken, err
8697
}
8798

8899
func (a *Manager) refreshToken() {
89-
token, err := getToken(a.config, a.tlsCfg, a.requestTimeout)
100+
token, startRefreshToken, err := getToken(a.config, a.httpClient, a.requestTimeout)
90101
if err != nil {
91102
log.Errorf("cannot refresh token: %v", err)
92103
}
93104
a.mutex.Lock()
94105
defer a.mutex.Unlock()
95106
a.token = token
96107
a.tokenErr = err
108+
a.startRefreshToken = startRefreshToken
97109
}
98110

99111
func (a *Manager) watchToken() {
100112
defer a.doneWg.Done()
113+
t := time.NewTicker(a.tickFrequency)
114+
defer t.Stop()
115+
101116
for {
102-
nextRenewal := a.nextRenewal()
103117
select {
104118
case <-a.done:
105119
return
106-
case <-time.After(nextRenewal):
107-
a.refreshToken()
120+
case <-t.C:
121+
if a.shouldRefresh() {
122+
a.refreshToken()
123+
}
108124
}
109125
}
110126
}

0 commit comments

Comments
 (0)