Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 38 additions & 15 deletions coderd/exp_chats.go
Original file line number Diff line number Diff line change
Expand Up @@ -3995,7 +3995,7 @@ func (api *API) resolveChatDiffContents(
return result, nil
}

gp := api.resolveGitProvider(ctx, reference.RepositoryRef.RemoteOrigin)
gp, _ := api.resolveGitProvider(ctx, reference.RepositoryRef.RemoteOrigin)
if gp == nil {
return result, nil
}
Expand Down Expand Up @@ -4059,7 +4059,7 @@ func (api *API) resolveChatDiffReference(
// current open PR. This picks up new PRs after the previous
// one was closed.
if reference.RepositoryRef != nil && reference.RepositoryRef.Owner != "" {
gp := api.resolveGitProvider(ctx, reference.RepositoryRef.RemoteOrigin)
gp, _ := api.resolveGitProvider(ctx, reference.RepositoryRef.RemoteOrigin)
if gp != nil {
token, err := api.resolveChatGitAccessToken(ctx, chat.OwnerID, reference.RepositoryRef.RemoteOrigin)
if token == nil || errors.Is(err, gitsync.ErrNoTokenAvailable) {
Expand Down Expand Up @@ -4122,7 +4122,7 @@ func (api *API) buildChatRepositoryRefFromStatus(ctx context.Context, status dat
return nil
}

providerType, gp := api.resolveExternalAuth(ctx, origin)
providerType, gp, _ := api.resolveExternalAuth(ctx, origin)
repoRef := &chatRepositoryRef{
Provider: providerType,
RemoteOrigin: origin,
Expand Down Expand Up @@ -4189,40 +4189,63 @@ func (api *API) getCachedChatDiffStatus(

// resolveExternalAuth finds the external auth config matching the
// given remote origin URL and returns both the provider type string
// (e.g. "github") and the gitprovider.Provider. Returns ("", nil)
// if no matching config is found or no provider could be constructed.
func (api *API) resolveExternalAuth(ctx context.Context, origin string) (providerType string, gp gitprovider.Provider) {
// (e.g. "github") and the gitprovider.Provider. Returns ("", nil, nil)
// if no matching config is found, and ("", nil, err) when a matching
// config could not yield a provider: the construction error if one
// failed to build, otherwise gitsync.ErrProviderUnimplemented if a
// matching config names a git type that has no implementation yet.
// A construction error wins because an operator can fix it.
func (api *API) resolveExternalAuth(ctx context.Context, origin string) (providerType string, gp gitprovider.Provider, err error) {
origin = strings.TrimSpace(origin)
if origin == "" {
return "", nil
return "", nil, nil
}
var constructErr error
unimplemented := false
for _, extAuth := range api.ExternalAuthConfigs {
if extAuth.Regex == nil || !extAuth.Regex.MatchString(origin) {
continue
}
p, err := extAuth.Git()
if err != nil {
normalizedType := strings.ToLower(strings.TrimSpace(extAuth.Type))
p, gitErr := extAuth.Git()
if gitErr != nil {
api.Logger.Warn(ctx, "failed to construct git provider",
slog.F("provider_id", extAuth.ID),
slog.F("provider_type", extAuth.Type),
slog.Error(err),
slog.Error(gitErr),
)
if constructErr == nil {
constructErr = xerrors.Errorf("construct git provider %q: %w", extAuth.ID, gitErr)
}
continue
}
if p == nil {
// Config.Git() returns a nil provider both for non-git
// types and for git types Coder has not implemented.
// Only the latter is worth reporting, and only if no
// later config matches with a working provider.
if codersdk.EnhancedExternalAuthProvider(normalizedType).Git() {
unimplemented = true
}
continue
}
return strings.ToLower(strings.TrimSpace(extAuth.Type)), p
return normalizedType, p, nil
}
if constructErr != nil {
return "", nil, constructErr
}
if unimplemented {
return "", nil, gitsync.ErrProviderUnimplemented
}
return "", nil
return "", nil, nil
}

// resolveGitProvider finds the external auth config matching the
// given remote origin URL and returns its git provider. Returns
// nil if no matching git provider is configured.
func (api *API) resolveGitProvider(ctx context.Context, origin string) gitprovider.Provider {
_, gp := api.resolveExternalAuth(ctx, origin)
return gp
func (api *API) resolveGitProvider(ctx context.Context, origin string) (gitprovider.Provider, error) {
_, gp, err := api.resolveExternalAuth(ctx, origin)
return gp, err
}

func (api *API) resolveChatGitAccessToken(
Expand Down
109 changes: 109 additions & 0 deletions coderd/exp_chats_internal_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import (
"net/http"
"net/http/httptest"
"reflect"
"regexp"
"strings"
"testing"
"testing/iotest"
Expand All @@ -26,10 +27,12 @@ import (
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbauthz"
"github.com/coder/coder/v2/coderd/database/dbmock"
"github.com/coder/coder/v2/coderd/externalauth"
"github.com/coder/coder/v2/coderd/httpmw"
"github.com/coder/coder/v2/coderd/util/ptr"
"github.com/coder/coder/v2/coderd/x/chatd"
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
"github.com/coder/coder/v2/coderd/x/gitsync"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/testutil"
"github.com/coder/quartz"
Expand Down Expand Up @@ -839,3 +842,109 @@ func TestWorkspaceUsageReader(t *testing.T) {
require.ErrorIs(t, err, io.EOF)
require.Equal(t, 2, reports, "empty reads must not count as usage")
}

func TestResolveExternalAuthProviderErrors(t *testing.T) {
t.Parallel()

newConfig := func(id, providerType, pattern string) *externalauth.Config {
return &externalauth.Config{
ID: id,
Type: providerType,
Regex: regexp.MustCompile(pattern),
}
}
brokenConfig := func(id, providerType, pattern string) *externalauth.Config {
cfg := newConfig(id, providerType, pattern)
cfg.APIBaseURL = "://invalid"
return cfg
}

const origin = "https://bitbucket.org/owner/repo"

tests := []struct {
name string
configs []*externalauth.Config
origin string
wantType string
wantProvider bool
wantUnsupported bool
wantErrContains string
}{
{
name: "git type without an implementation",
configs: []*externalauth.Config{newConfig("bb", "bitbucket-cloud", `bitbucket\.org`)},
origin: origin,
wantUnsupported: true,
},
{
name: "implemented git type",
configs: []*externalauth.Config{newConfig("gh", "github", `github\.com`)},
origin: "https://github.com/owner/repo",
wantType: "github",
wantProvider: true,
},
{
name: "no config matches the origin",
configs: []*externalauth.Config{newConfig("gh", "github", `github\.com`)},
origin: origin,
},
{
// A non-git provider matching the origin is not a gap in
// Coder's git support, so it must not park the row.
name: "matching config is not a git provider",
configs: []*externalauth.Config{newConfig("jfrog", "jfrog", `bitbucket\.org`)},
origin: origin,
},
{
// A provider that failed to build is a deployment
// misconfiguration the operator can fix, so it must not
// be reported as an unimplemented git type.
name: "construction failure wins over an unimplemented type",
configs: []*externalauth.Config{
brokenConfig("gl", "gitlab", `bitbucket\.org`),
newConfig("bb", "bitbucket-cloud", `bitbucket\.org`),
},
origin: origin,
wantErrContains: "construct git provider \"gl\"",
},
{
name: "implemented config wins over an unimplemented one",
configs: []*externalauth.Config{
newConfig("bb", "bitbucket-cloud", `bitbucket\.org`),
newConfig("gh", "github", `bitbucket\.org`),
},
origin: origin,
wantType: "github",
wantProvider: true,
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()

api := &API{Options: &Options{
Logger: slogtest.Make(t, nil),
ExternalAuthConfigs: tt.configs,
}}

providerType, gp, err := api.resolveExternalAuth(context.Background(), tt.origin)

switch {
case tt.wantUnsupported:
require.ErrorIs(t, err, gitsync.ErrProviderUnimplemented)
case tt.wantErrContains != "":
require.ErrorContains(t, err, tt.wantErrContains)
require.NotErrorIs(t, err, gitsync.ErrProviderUnimplemented)
default:
require.NoError(t, err)
}
require.Equal(t, tt.wantType, providerType)
if tt.wantProvider {
require.NotNil(t, gp)
} else {
require.Nil(t, gp)
}
})
}
}
20 changes: 15 additions & 5 deletions coderd/x/gitsync/gitsync.go
Original file line number Diff line number Diff line change
Expand Up @@ -29,11 +29,19 @@ const (
)

// ProviderResolver maps a git remote origin to the gitprovider
// that handles it. Returns nil if no provider matches.
type ProviderResolver func(ctx context.Context, origin string) gitprovider.Provider
// that handles it. Returns (nil, nil) if no configured provider
// matches the origin, and ErrProviderUnimplemented if one matches
// but its git type has no implementation.
type ProviderResolver func(ctx context.Context, origin string) (gitprovider.Provider, error)

var ErrNoTokenAvailable error = errors.New("no token available")

// ErrProviderUnimplemented indicates the origin matched a configured
// provider whose git type Coder does not implement yet. Unlike a
// misconfigured origin, an operator cannot fix this, so the worker
// parks the row until an implementation ships.
var ErrProviderUnimplemented error = errors.New("git provider not implemented")

// ErrStalePullRequest indicates the row's stored PR belongs to a
// previous branch and the current branch has none.
var ErrStalePullRequest error = errors.New("stale pull request")
Expand Down Expand Up @@ -163,9 +171,11 @@ func (r *Refresher) Refresh(
// duplicate resolution for rows in the same group.
var resolved []resolvedGroup
for key, indices := range groups {
provider := r.providers(ctx, key.origin)
if provider == nil {
err := xerrors.Errorf("no provider for origin %q", key.origin)
provider, err := r.providers(ctx, key.origin)
if err != nil || provider == nil {
if err == nil {
err = xerrors.Errorf("no provider for origin %q", key.origin)
}
for _, i := range indices {
results[i].Error = err
}
Expand Down
28 changes: 14 additions & 14 deletions coderd/x/gitsync/gitsync_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -127,7 +127,7 @@ func TestRefresher_WithPRURL(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new("test-token"), nil
}
Expand Down Expand Up @@ -183,7 +183,7 @@ func TestRefresher_BranchResolvesToPR(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new("test-token"), nil
}
Expand Down Expand Up @@ -225,7 +225,7 @@ func TestRefresher_BranchNoPRYet(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new("test-token"), nil
}
Expand Down Expand Up @@ -254,7 +254,7 @@ func TestRefresher_BranchNoPRYet(t *testing.T) {
func TestRefresher_NoProviderForOrigin(t *testing.T) {
t.Parallel()

providers := func(_ context.Context, _ string) gitprovider.Provider { return nil }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return nil, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new("test-token"), nil
}
Expand Down Expand Up @@ -295,7 +295,7 @@ func TestRefresher_TokenResolutionFails(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return nil, errors.New("token lookup failed")
}
Expand Down Expand Up @@ -327,7 +327,7 @@ func TestRefresher_EmptyToken(t *testing.T) {

mp := &mockProvider{}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new(""), nil
}
Expand Down Expand Up @@ -365,7 +365,7 @@ func TestRefresher_ProviderFetchFails(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new("test-token"), nil
}
Expand Down Expand Up @@ -401,7 +401,7 @@ func TestRefresher_PRURLParseFailure(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new("test-token"), nil
}
Expand Down Expand Up @@ -439,7 +439,7 @@ func TestRefresher_BatchGroupsByOwnerAndOrigin(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }

var tokenCalls atomic.Int32
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
Expand Down Expand Up @@ -521,7 +521,7 @@ func TestRefresher_UsesInjectedClock(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new("test-token"), nil
}
Expand Down Expand Up @@ -573,7 +573,7 @@ func TestRefresher_RateLimitSkipsRemainingInGroup(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new("test-token"), nil
}
Expand Down Expand Up @@ -694,7 +694,7 @@ func TestRefresher_CorrectTokenPerOrigin(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }

r := gitsync.NewRefresher(providers, tokens, slogtest.Make(t, nil), quartz.NewReal())

Expand Down Expand Up @@ -779,7 +779,7 @@ func TestRefresher_ConcurrentProcessing(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new("test-token"), nil
}
Expand Down Expand Up @@ -901,7 +901,7 @@ func TestRefresher_PRURLBranchMismatch(t *testing.T) {
},
}

providers := func(_ context.Context, _ string) gitprovider.Provider { return mp }
providers := func(_ context.Context, _ string) (gitprovider.Provider, error) { return mp, nil }
tokens := func(_ context.Context, _ uuid.UUID, _ string) (*string, error) {
return new("test-token"), nil
}
Expand Down
Loading
Loading