Files
probo/pkg/accessreview/source_service.go
Aurélien Sibiril f711e9d816 Reset access source name sync on connection change
Once the source-name worker sets name_synced_at, nothing cleared it,
so a source that hit a terminal failure (Brex 403, wrong Sentry org
slug) kept its generic name forever — even after the user reconnected
with the right scope or picked the correct org, contradicting the
Brex guidance to reconnect.

Clear name_synced_at whenever the connection changes: on UpdateSource
when a connector is (re)set, inside ConfigureAccessReviewSource when
the org is (re)selected, and on the OAuth reconnect path via the new
ResetSourceNameSyncForConnector service method. The worker then
re-claims the row and re-resolves the display name.

Signed-off-by: Aurélien Sibiril <81782+aureliensibiril@users.noreply.github.com>
2026-07-22 17:27:55 +02:00

490 lines
14 KiB
Go

// Copyright (c) 2026 Probo Inc <hello@probo.com>.
//
// Permission is hereby granted, free of charge, to any person obtaining a copy
// of this software and associated documentation files (the "Software"), to deal
// in the Software without restriction, including without limitation the rights
// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
// copies of the Software, and to permit persons to whom the Software is
// furnished to do so, subject to the following conditions:
//
// The above copyright notice and this permission notice shall be included in
// all copies or substantial portions of the Software.
//
// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
// SOFTWARE.
package accessreview
import (
"context"
"fmt"
"net/http"
"time"
"go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/connector"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
"go.probo.inc/probo/pkg/validator"
)
const (
NameMaxLength = 1000
)
type (
CreateAccessReviewSourceRequest struct {
OrganizationID gid.GID
ConnectorID *gid.GID
Name string
CsvData *string
}
UpdateAccessReviewSourceRequest struct {
AccessReviewSourceID gid.GID
Name **string
ConnectorID **gid.GID
CsvData **string
}
ConfigureAccessReviewSourceRequest struct {
AccessReviewSourceID gid.GID
OrganizationSlug string
}
)
func (r *CreateAccessReviewSourceRequest) Validate() error {
v := validator.New()
v.Check(r.OrganizationID, "organization_id", validator.Required(), validator.GID(coredata.OrganizationEntityType))
v.Check(r.Name, "name", validator.SafeTextNoNewLine(NameMaxLength))
return v.Error()
}
func (r *ConfigureAccessReviewSourceRequest) Validate() error {
v := validator.New()
v.Check(r.AccessReviewSourceID, "access_review_source_id", validator.Required(), validator.GID(coredata.AccessReviewSourceEntityType))
v.Check(r.OrganizationSlug, "organization_slug", validator.Required())
return v.Error()
}
func (r *UpdateAccessReviewSourceRequest) Validate() error {
v := validator.New()
v.Check(r.AccessReviewSourceID, "access_review_source_id", validator.Required(), validator.GID(coredata.AccessReviewSourceEntityType))
v.Check(r.Name, "name", validator.SafeTextNoNewLine(NameMaxLength))
return v.Error()
}
func (s *Service) CreateSource(
ctx context.Context,
scope coredata.Scoper,
req CreateAccessReviewSourceRequest,
) (*coredata.AccessReviewSource, error) {
if err := req.Validate(); err != nil {
return nil, err
}
now := time.Now()
source := &coredata.AccessReviewSource{
ID: gid.New(scope.GetTenantID(), coredata.AccessReviewSourceEntityType),
OrganizationID: req.OrganizationID,
ConnectorID: req.ConnectorID,
Name: req.Name,
CsvData: req.CsvData,
CreatedAt: now,
UpdatedAt: now,
}
err := s.pg.WithTx(
ctx,
func(ctx context.Context, conn pg.Tx) error {
// Validate connector exists if provided
if req.ConnectorID != nil {
connector := &coredata.Connector{}
if err := connector.LoadMetadataByID(ctx, conn, scope, *req.ConnectorID); err != nil {
return fmt.Errorf("cannot load connector: %w", err)
}
}
if err := source.Insert(ctx, conn, scope); err != nil {
return fmt.Errorf("cannot insert access source: %w", err)
}
return nil
},
)
if err != nil {
return nil, fmt.Errorf("cannot create access source: %w", err)
}
return source, nil
}
func (s *Service) GetSource(
ctx context.Context,
scope coredata.Scoper,
accessSourceID gid.GID,
) (*coredata.AccessReviewSource, error) {
source := &coredata.AccessReviewSource{}
err := s.pg.WithConn(
ctx,
func(ctx context.Context, conn pg.Querier) error {
return source.LoadByID(ctx, conn, scope, accessSourceID)
},
)
if err != nil {
return nil, fmt.Errorf("cannot get access source: %w", err)
}
return source, nil
}
func (s *Service) UpdateSource(
ctx context.Context,
scope coredata.Scoper,
req UpdateAccessReviewSourceRequest,
) (*coredata.AccessReviewSource, error) {
if err := req.Validate(); err != nil {
return nil, err
}
source := &coredata.AccessReviewSource{}
err := s.pg.WithTx(
ctx,
func(ctx context.Context, conn pg.Tx) error {
if err := source.LoadByID(ctx, conn, scope, req.AccessReviewSourceID); err != nil {
return fmt.Errorf("cannot load access source: %w", err)
}
if req.Name != nil {
if *req.Name != nil {
source.Name = **req.Name
}
}
if req.ConnectorID != nil {
if *req.ConnectorID != nil {
connector := &coredata.Connector{}
if err := connector.LoadMetadataByID(ctx, conn, scope, **req.ConnectorID); err != nil {
return fmt.Errorf("cannot load connector: %w", err)
}
}
source.ConnectorID = *req.ConnectorID
// A (re)linked connector may resolve to a different instance
// name; clear the synced flag so the source-name worker picks
// the row up and re-resolves it.
source.NameSyncedAt = nil
}
if req.CsvData != nil {
source.CsvData = *req.CsvData
}
source.UpdatedAt = time.Now()
if err := source.Update(ctx, conn, scope); err != nil {
return fmt.Errorf("cannot update access source: %w", err)
}
return nil
},
)
if err != nil {
return nil, fmt.Errorf("cannot update access source: %w", err)
}
return source, nil
}
func (s *Service) DeleteSource(
ctx context.Context,
scope coredata.Scoper,
accessSourceID gid.GID,
) error {
source := &coredata.AccessReviewSource{}
return s.pg.WithTx(
ctx,
func(ctx context.Context, conn pg.Tx) error {
if err := source.LoadByID(ctx, conn, scope, accessSourceID); err != nil {
return fmt.Errorf("cannot load access source: %w", err)
}
if err := source.Delete(ctx, conn, scope); err != nil {
return fmt.Errorf("cannot delete access source: %w", err)
}
// Garbage-collect the underlying connector once nothing else
// references it. The connectors table is unique per
// (organization_id, provider), so leaving an orphaned connector
// behind would block re-adding a source for the same provider.
if source.ConnectorID == nil {
return nil
}
accessSources := &coredata.AccessReviewSources{}
sourceCount, err := accessSources.CountByConnectorID(ctx, conn, scope, *source.ConnectorID)
if err != nil {
return fmt.Errorf("cannot count access sources for connector: %w", err)
}
if sourceCount > 0 {
return nil
}
bridges := &coredata.SCIMBridges{}
bridgeCount, err := bridges.CountByConnectorID(ctx, conn, scope, *source.ConnectorID)
if err != nil {
return fmt.Errorf("cannot count scim bridges for connector: %w", err)
}
if bridgeCount > 0 {
return nil
}
// Garbage-collecting the connector is best-effort. A
// concurrent transaction may insert a new access source or
// SCIM bridge referencing this connector between the counts
// above and the DELETE, producing a foreign-key violation.
// Run the delete inside a savepoint so such a failure rolls
// back only the GC attempt and still commits the access
// source deletion instead of aborting the whole transaction.
if err := conn.Savepoint(
ctx,
func(ctx context.Context, conn pg.Tx) error {
cnnctr := &coredata.Connector{ID: *source.ConnectorID}
if err := cnnctr.Delete(ctx, conn, scope); err != nil {
return fmt.Errorf("cannot delete connector: %w", err)
}
return nil
},
); err != nil {
return err
}
return nil
},
)
}
func (s *Service) ListSourcesForOrganizationID(
ctx context.Context,
scope coredata.Scoper,
organizationID gid.GID,
cursor *page.Cursor[coredata.AccessReviewSourceOrderField],
) (*page.Page[*coredata.AccessReviewSource, coredata.AccessReviewSourceOrderField], error) {
var sources coredata.AccessReviewSources
err := s.pg.WithConn(
ctx,
func(ctx context.Context, conn pg.Querier) error {
return sources.LoadByOrganizationID(ctx, conn, scope, organizationID, cursor)
},
)
if err != nil {
return nil, fmt.Errorf("cannot list access sources: %w", err)
}
return page.NewPage(sources, cursor), nil
}
func (s *Service) CountSourcesForOrganizationID(
ctx context.Context,
scope coredata.Scoper,
organizationID gid.GID,
) (int, error) {
var count int
err := s.pg.WithConn(
ctx,
func(ctx context.Context, conn pg.Querier) (err error) {
sources := coredata.AccessReviewSources{}
count, err = sources.CountByOrganizationID(ctx, conn, scope, organizationID)
return err
},
)
if err != nil {
return 0, fmt.Errorf("cannot count access sources: %w", err)
}
return count, nil
}
// ConnectorHTTPClient loads a connector by ID with decrypted credentials
// and returns an HTTP client with token refresh support. If the token was
// refreshed during client creation, the updated credentials are persisted.
func (s *Service) ConnectorHTTPClient(
ctx context.Context,
scope coredata.Scoper,
connectorID gid.GID,
) (*http.Client, *coredata.Connector, error) {
var dbConnector coredata.Connector
err := s.pg.WithConn(
ctx,
func(ctx context.Context, conn pg.Querier) error {
if err := dbConnector.LoadByID(ctx, conn, scope, connectorID, s.encryptionKey); err != nil {
return fmt.Errorf("cannot load connector: %w", err)
}
return nil
},
)
if err != nil {
return nil, nil, err
}
var tokenBefore string
oauth2Conn, isOAuth2 := dbConnector.Connection.(*connector.OAuth2Connection)
if isOAuth2 {
tokenBefore = oauth2Conn.AccessToken
}
var httpClient *http.Client
if isOAuth2 && s.connectorRegistry != nil {
refreshCfg := s.connectorRegistry.GetOAuth2RefreshConfig(string(dbConnector.Provider))
if refreshCfg != nil {
var err error
httpClient, err = oauth2Conn.RefreshableClient(ctx, *refreshCfg)
if err != nil {
return nil, nil, fmt.Errorf("cannot create refreshable HTTP client: %w", err)
}
}
}
if httpClient == nil {
// Inject the Probo-held key for ManagedAPIKey providers (no-op
// otherwise), resolving it fresh at use time rather than from the
// connection row.
if err := s.providerRegistry.ApplyManagedAPIKey(&dbConnector); err != nil {
return nil, nil, err
}
var err error
httpClient, err = dbConnector.Connection.Client(ctx)
if err != nil {
return nil, nil, fmt.Errorf("cannot create HTTP client: %w", err)
}
}
// Persist refreshed token if it changed.
if isOAuth2 && oauth2Conn.AccessToken != tokenBefore {
dbConnector.UpdatedAt = time.Now()
if err := s.pg.WithTx(
ctx,
func(ctx context.Context, tx pg.Tx) error {
return dbConnector.Update(ctx, tx, scope, s.encryptionKey)
},
); err != nil {
return nil, nil, fmt.Errorf("cannot persist refreshed token: %w", err)
}
}
return httpClient, &dbConnector, nil
}
func (s *Service) ConfigureAccessReviewSource(
ctx context.Context,
scope coredata.Scoper,
req ConfigureAccessReviewSourceRequest,
) (*coredata.AccessReviewSource, error) {
if err := req.Validate(); err != nil {
return nil, err
}
source := &coredata.AccessReviewSource{}
err := s.pg.WithTx(
ctx,
func(ctx context.Context, conn pg.Tx) error {
if err := source.LoadByID(ctx, conn, scope, req.AccessReviewSourceID); err != nil {
return fmt.Errorf("cannot load access source: %w", err)
}
if source.ConnectorID == nil {
return fmt.Errorf("cannot configure access source: no connector attached")
}
dbConnector := &coredata.Connector{}
if err := dbConnector.LoadByID(ctx, conn, scope, *source.ConnectorID, s.encryptionKey); err != nil {
return fmt.Errorf("cannot load connector: %w", err)
}
reg, ok := s.providerRegistry.Get(dbConnector.Provider)
if !ok || reg.SetOrganizationSettings == nil {
return fmt.Errorf("cannot configure access source: provider %s does not support organization configuration", dbConnector.Provider)
}
if err := reg.SetOrganizationSettings(dbConnector, req.OrganizationSlug); err != nil {
return fmt.Errorf("cannot set %s settings: %w", dbConnector.Provider, err)
}
dbConnector.UpdatedAt = time.Now()
if err := dbConnector.Update(ctx, conn, scope, s.encryptionKey); err != nil {
return fmt.Errorf("cannot update connector: %w", err)
}
// The selected org changed, so the resolvable instance name may
// have too; clear the synced flag so the source-name worker
// re-resolves the display name.
source.NameSyncedAt = nil
source.UpdatedAt = time.Now()
if err := source.Update(ctx, conn, scope); err != nil {
return fmt.Errorf("cannot reset access source name sync: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return source, nil
}
// ResetSourceNameSyncForConnector clears the synced-name flag on every access
// source backed by connectorID so the source-name worker re-resolves the
// display name. Called after a connector is reconnected — the new grant may
// scope a different org/workspace, changing the resolvable name.
func (s *Service) ResetSourceNameSyncForConnector(
ctx context.Context,
scope coredata.Scoper,
connectorID gid.GID,
) error {
return s.pg.WithTx(
ctx,
func(ctx context.Context, conn pg.Tx) error {
sources := &coredata.AccessReviewSources{}
return sources.ClearNameSyncedAtByConnectorID(ctx, conn, scope, connectorID)
},
)
}