From 534095396658e2bc71ff9869950f4e25b78dda6c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aur=C3=A9lien=20Sibiril?= <81782+aureliensibiril@users.noreply.github.com> Date: Thu, 9 Apr 2026 01:04:24 +0200 Subject: [PATCH] feat(probo): validate reconnects and preserve dropped token fields MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Reconnect now takes a ReconnectConnectorRequest carrying the expected OrganizationID and Provider. It validates inside the same transaction that the loaded connector belongs to the requested org, provider and OAUTH2 protocol before mutating the row. This blocks cross-org and cross-provider corruption via a crafted connector_id reaching the OAuth callback through the HMAC-signed state token. preserveConnectionFields copies fields from the existing connection onto the new one when the new one omits them: - OAuth2 refresh_token: Google drops it on incremental-auth reuse when prompt=consent is skipped. - Slack webhook URL, channel and channel ID: access review Slack reconnects without the incoming-webhook scope return a token response with no incoming_webhook field. GetByOrganizationIDAndProvider now routes through the widest-scope coredata loader, and GetWithConnection exposes a by-ID load that returns the fully decrypted connector so the initiate handler can read the stored scope set. Signed-off-by: Aurélien Sibiril <81782+aureliensibiril@users.noreply.github.com> --- pkg/probo/connector_service.go | 125 +++++++++++++++++++++----- pkg/server/api/console/v1/resolver.go | 10 ++- 2 files changed, 112 insertions(+), 23 deletions(-) diff --git a/pkg/probo/connector_service.go b/pkg/probo/connector_service.go index 8b1f6c6ed..1a1814d5d 100644 --- a/pkg/probo/connector_service.go +++ b/pkg/probo/connector_service.go @@ -60,6 +60,13 @@ type ( GitHubSettings *coredata.GitHubConnectorSettings OnePasswordUsersAPISettings *coredata.OnePasswordUsersAPISettings } + + ReconnectConnectorRequest struct { + ConnectorID gid.GID + OrganizationID gid.GID + Provider coredata.ConnectorProvider + Connection connector.Connection + } ) func (car *CreateConnectorRequest) Validate() error { @@ -71,6 +78,15 @@ func (car *CreateConnectorRequest) Validate() error { return v.Error() } +func (rcr *ReconnectConnectorRequest) Validate() error { + v := validator.New() + v.Check(rcr.ConnectorID, "connector_id", validator.Required(), validator.GID(coredata.ConnectorEntityType)) + v.Check(rcr.OrganizationID, "organization_id", validator.Required(), validator.GID(coredata.OrganizationEntityType)) + v.Check(rcr.Provider, "provider", validator.Required(), validator.OneOfSlice(coredata.ConnectorProviders())) + v.Check(rcr.Connection, "connection", validator.Required()) + return v.Error() +} + func (s *ConnectorService) ListForOrganizationID( ctx context.Context, organizationID gid.GID, @@ -129,32 +145,51 @@ func (s *ConnectorService) GetByOrganizationIDAndProvider( organizationID gid.GID, provider coredata.ConnectorProvider, ) (*coredata.Connector, error) { - var connectors coredata.Connectors + cnnctr := &coredata.Connector{} err := s.svc.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) error { - return connectors.LoadAllByOrganizationIDProtocolAndProvider( + return cnnctr.LoadOneByOrganizationIDAndProvider( ctx, conn, s.svc.scope, - organizationID, - coredata.ConnectorProtocolOAuth2, - provider, s.svc.encryptionKey, + organizationID, + provider, ) }, ) - if err != nil { return nil, fmt.Errorf("cannot get connector: %w", err) } - if len(connectors) == 0 { - return nil, coredata.ErrResourceNotFound + return cnnctr, nil +} + +// GetWithConnection loads a specific connector by ID and returns the +// full *coredata.Connector with Connection populated. Used by the +// initiate handler's explicit reconnect path (?connector_id=), +// which needs to read the stored scope set to compute the union. +// Contrast with Get, which uses LoadMetadataByID and returns a +// connector with Connection == nil. +func (s *ConnectorService) GetWithConnection( + ctx context.Context, + connectorID gid.GID, +) (*coredata.Connector, error) { + cnnctr := &coredata.Connector{} + + err := s.svc.pg.WithConn( + ctx, + func(ctx context.Context, conn pg.Querier) error { + return cnnctr.LoadByID(ctx, conn, s.svc.scope, connectorID, s.svc.encryptionKey) + }, + ) + if err != nil { + return nil, fmt.Errorf("cannot get connector: %w", err) } - return connectors[0], nil + return cnnctr, nil } func (s *ConnectorService) Get( @@ -288,29 +323,75 @@ func (s *ConnectorService) Create( return newConnector, nil } -// Reconnect updates an existing connector's connection (token) without -// changing its settings or identity. Used when an OAuth token expires -// and the user re-authenticates. +// Reconnect updates an existing OAuth2 connector's connection (token) +// in place. It validates that the loaded connector belongs to the +// expected org and provider inside the same transaction, blocking +// cross-org and cross-provider corruption via a crafted connector_id +// in the initiate URL. Refresh tokens and Slack webhook settings are +// preserved from the existing connection when the new one omits them. func (s *ConnectorService) Reconnect( ctx context.Context, - connectorID gid.GID, - connection connector.Connection, + req ReconnectConnectorRequest, ) (*coredata.Connector, error) { + if err := req.Validate(); err != nil { + return nil, fmt.Errorf("cannot reconnect connector: %w", err) + } + cnnctr := &coredata.Connector{} - err := s.svc.pg.WithTx(ctx, func(ctx context.Context, conn pg.Tx) error { - if err := cnnctr.LoadMetadataByID(ctx, conn, s.svc.scope, connectorID); err != nil { - return fmt.Errorf("cannot load connector: %w", err) - } + err := s.svc.pg.WithTx( + ctx, + func(ctx context.Context, conn pg.Tx) error { + if err := cnnctr.LoadByID(ctx, conn, s.svc.scope, req.ConnectorID, s.svc.encryptionKey); err != nil { + return fmt.Errorf("cannot load connector: %w", err) + } - cnnctr.Connection = connection - cnnctr.UpdatedAt = time.Now() + if cnnctr.OrganizationID != req.OrganizationID { + return fmt.Errorf("cannot reconnect connector: organization mismatch") + } + if cnnctr.Provider != req.Provider { + return fmt.Errorf("cannot reconnect connector: provider mismatch") + } + if cnnctr.Protocol != coredata.ConnectorProtocolOAuth2 { + return fmt.Errorf("cannot reconnect connector: not an OAuth2 connector") + } - return cnnctr.Update(ctx, conn, s.svc.scope, s.svc.encryptionKey) - }) + cnnctr.Connection = preserveConnectionFields(req.Connection, cnnctr.Connection) + cnnctr.UpdatedAt = time.Now() + + return cnnctr.Update(ctx, conn, s.svc.scope, s.svc.encryptionKey) + }, + ) if err != nil { return nil, fmt.Errorf("cannot reconnect connector: %w", err) } return cnnctr, nil } + +// preserveConnectionFields returns newConn with refresh token and Slack +// webhook settings copied from oldConn when newConn omits them. Google +// omits refresh_token on incremental-auth reuse; a Slack access-review +// reconnect with no incoming-webhook scope omits the webhook settings. +func preserveConnectionFields(newConn, oldConn connector.Connection) connector.Connection { + switch n := newConn.(type) { + case *connector.OAuth2Connection: + if o, ok := oldConn.(*connector.OAuth2Connection); ok { + if n.RefreshToken == "" { + n.RefreshToken = o.RefreshToken + } + } + case *connector.SlackConnection: + if o, ok := oldConn.(*connector.SlackConnection); ok { + if n.RefreshToken == "" { + n.RefreshToken = o.RefreshToken + } + if n.Settings.WebhookURL == "" { + n.Settings.WebhookURL = o.Settings.WebhookURL + n.Settings.Channel = o.Settings.Channel + n.Settings.ChannelID = o.Settings.ChannelID + } + } + } + return newConn +} diff --git a/pkg/server/api/console/v1/resolver.go b/pkg/server/api/console/v1/resolver.go index 44e4c264f..3b000b830 100644 --- a/pkg/server/api/console/v1/resolver.go +++ b/pkg/server/api/console/v1/resolver.go @@ -225,7 +225,15 @@ func handleConnectorComplete( return } - cnnctr, err = svc.Connectors.Reconnect(r.Context(), connectorID, connection) + cnnctr, err = svc.Connectors.Reconnect( + r.Context(), + probo.ReconnectConnectorRequest{ + ConnectorID: connectorID, + OrganizationID: organizationID, + Provider: connectorProvider, + Connection: connection, + }, + ) if err != nil { panic(fmt.Errorf("cannot reconnect connector: %w", err)) }