iam: Refresh live connection when ID tokens expire (#107209)

* iam: refresh live connection when ID tokens expire

* add coverage for the handler functions

* reinstate inadvertently broken unit test
This commit is contained in:
Victor Cinaglia
2025-07-03 10:16:24 -03:00
committed by GitHub
parent 8d8b824f73
commit 4f66c4a2a1
6 changed files with 337 additions and 14 deletions
+177 -13
View File
@@ -7,15 +7,20 @@ import (
"testing"
"time"
"github.com/go-jose/go-jose/v3"
"github.com/go-jose/go-jose/v3/jwt"
"github.com/stretchr/testify/require"
"github.com/centrifugal/centrifuge"
"github.com/grafana/grafana/pkg/api/routing"
"github.com/grafana/grafana/pkg/apimachinery/identity"
"github.com/grafana/grafana/pkg/infra/db"
"github.com/grafana/grafana/pkg/infra/usagestats"
"github.com/grafana/grafana/pkg/services/accesscontrol/acimpl"
"github.com/grafana/grafana/pkg/services/annotations/annotationstest"
"github.com/grafana/grafana/pkg/services/dashboards"
"github.com/grafana/grafana/pkg/services/featuremgmt"
"github.com/grafana/grafana/pkg/services/live/livecontext"
"github.com/grafana/grafana/pkg/setting"
"github.com/grafana/grafana/pkg/tests/testsuite"
)
@@ -29,20 +34,9 @@ func TestIntegration_provideLiveService_RedisUnavailable(t *testing.T) {
cfg.LiveHAEngine = "testredisunavailable"
_, err := ProvideService(nil, cfg,
routing.NewRouteRegister(),
nil, nil, nil, nil,
db.InitTestDB(t),
nil,
&usagestats.UsageStatsMock{T: t},
nil,
featuremgmt.WithFeatures(),
acimpl.ProvideAccessControl(featuremgmt.WithFeatures()),
&dashboards.FakeDashboardService{},
annotationstest.NewFakeAnnotationsRepo(),
nil, nil)
_, err := setupLiveService(cfg, t)
// Proceeds without live HA if redis is unavaialble
// Proceeds without live HA if redis is unavailable
require.NoError(t, err)
}
@@ -233,3 +227,173 @@ func Test_getHistogramMetric(t *testing.T) {
})
}
}
func Test_handleOnPublish_IDTokenExpiration(t *testing.T) {
g, err := setupLiveService(nil, t)
require.NoError(t, err)
client, _, err := centrifuge.NewClient(context.Background(), g.node, newDummyTransport("test"))
require.NoError(t, err)
t.Run("expired token", func(t *testing.T) {
expiration := time.Now().Add(-time.Hour)
token := createToken(t, &expiration)
ctx := livecontext.SetContextSignedUser(context.Background(), &identity.StaticRequester{IDToken: token})
reply, err := g.handleOnPublish(ctx, client, centrifuge.PublishEvent{
Channel: "test",
Data: []byte("test"),
})
require.ErrorIs(t, err, centrifuge.ErrorExpired)
require.Empty(t, reply)
})
t.Run("unexpired token", func(t *testing.T) {
expiration := time.Now().Add(time.Hour)
token := createToken(t, &expiration)
ctx := livecontext.SetContextSignedUser(context.Background(), &identity.StaticRequester{IDToken: token})
reply, err := g.handleOnPublish(ctx, client, centrifuge.PublishEvent{
Channel: "test",
Data: []byte("test"),
})
// Another error is returned if the token is not expired but the refresh fails.
// That happens because we're providing an invalid orgID as the channel.
require.NotErrorIs(t, err, centrifuge.ErrorExpired)
require.Empty(t, reply)
})
}
func Test_handleOnRPC_IDTokenExpiration(t *testing.T) {
g, err := setupLiveService(nil, t)
require.NoError(t, err)
client, _, err := centrifuge.NewClient(context.Background(), g.node, newDummyTransport("test"))
require.NoError(t, err)
t.Run("expired token", func(t *testing.T) {
expiration := time.Now().Add(-time.Hour)
token := createToken(t, &expiration)
ctx := livecontext.SetContextSignedUser(context.Background(), &identity.StaticRequester{IDToken: token})
reply, err := g.handleOnRPC(ctx, client, centrifuge.RPCEvent{
Method: "grafana.query",
Data: []byte("test"),
})
require.ErrorIs(t, err, centrifuge.ErrorExpired)
require.Empty(t, reply)
})
t.Run("unexpired token", func(t *testing.T) {
expiration := time.Now().Add(time.Hour)
token := createToken(t, &expiration)
ctx := livecontext.SetContextSignedUser(context.Background(), &identity.StaticRequester{IDToken: token})
reply, err := g.handleOnRPC(ctx, client, centrifuge.RPCEvent{
Method: "grafana.query",
Data: []byte("test"),
})
// Another error is returned if the token is not expired but the refresh fails.
// That happens because we're providing an invalid orgID as the channel.
require.NotErrorIs(t, err, centrifuge.ErrorExpired)
require.Empty(t, reply)
})
}
func Test_handleOnSubscribe_IDTokenExpiration(t *testing.T) {
g, err := setupLiveService(nil, t)
require.NoError(t, err)
client, _, err := centrifuge.NewClient(context.Background(), g.node, newDummyTransport("test"))
require.NoError(t, err)
t.Run("expired token", func(t *testing.T) {
expiration := time.Now().Add(-time.Hour)
token := createToken(t, &expiration)
ctx := livecontext.SetContextSignedUser(context.Background(), &identity.StaticRequester{IDToken: token})
reply, err := g.handleOnSubscribe(ctx, client, centrifuge.SubscribeEvent{
Channel: "test",
})
require.ErrorIs(t, err, centrifuge.ErrorExpired)
require.Empty(t, reply)
})
t.Run("unexpired token", func(t *testing.T) {
expiration := time.Now().Add(time.Hour)
token := createToken(t, &expiration)
ctx := livecontext.SetContextSignedUser(context.Background(), &identity.StaticRequester{IDToken: token})
reply, err := g.handleOnSubscribe(ctx, client, centrifuge.SubscribeEvent{
Channel: "test",
})
// Another error is returned if the token is not expired but the refresh fails.
// That happens because we're providing an invalid orgID as the channel.
require.NotErrorIs(t, err, centrifuge.ErrorExpired)
require.Empty(t, reply)
})
}
func setupLiveService(cfg *setting.Cfg, t *testing.T) (*GrafanaLive, error) {
if cfg == nil {
cfg = setting.NewCfg()
}
return ProvideService(nil,
cfg,
routing.NewRouteRegister(),
nil, nil, nil, nil,
db.InitTestDB(t),
nil,
&usagestats.UsageStatsMock{T: t},
nil,
featuremgmt.WithFeatures(),
acimpl.ProvideAccessControl(featuremgmt.WithFeatures()),
&dashboards.FakeDashboardService{},
annotationstest.NewFakeAnnotationsRepo(),
nil, nil)
}
type dummyTransport struct {
name string
}
func (t *dummyTransport) Name() string { return t.name }
func (t *dummyTransport) Protocol() centrifuge.ProtocolType { return centrifuge.ProtocolTypeJSON }
func (t *dummyTransport) ProtocolVersion() centrifuge.ProtocolVersion {
return centrifuge.ProtocolVersion2
}
func (t *dummyTransport) Emulation() bool { return false }
func (t *dummyTransport) Unidirectional() bool { return false }
func (t *dummyTransport) DisabledPushFlags() uint64 { return 0 }
func (t *dummyTransport) PingPongConfig() centrifuge.PingPongConfig {
return centrifuge.PingPongConfig{}
}
func (t *dummyTransport) Write(data []byte) error { return nil }
func (t *dummyTransport) WriteMany(d ...[]byte) error { return nil }
func (t *dummyTransport) Close(disconnect centrifuge.Disconnect) error {
return nil
}
func newDummyTransport(name string) *dummyTransport {
return &dummyTransport{name: name}
}
func createToken(t *testing.T, exp *time.Time) string {
key := []byte("test-secret-key")
signer, err := jose.NewSigner(jose.SigningKey{Algorithm: jose.HS256, Key: key}, nil)
require.NoError(t, err)
claims := struct {
jwt.Claims
}{
Claims: jwt.Claims{
Subject: "test-user",
},
}
if exp != nil {
claims.Expiry = jwt.NewNumericDate(*exp)
}
token, err := jwt.Signed(signer).Claims(claims).CompactSerialize()
require.NoError(t, err)
return token
}