@@ -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
1817type 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.
3031func 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
8899func (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
99111func (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