mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-12 17:59:06 +02:00
Apply the grant allowlist on the non-YAML provider startup path
This commit is contained in:
@@ -90,6 +90,11 @@ func NewProvider(ctx context.Context, config *Config) (*Provider, error) {
|
||||
return nil, fmt.Errorf("failed to ensure local connector: %w", err)
|
||||
}
|
||||
|
||||
if err := ensureConnectorGrantTypes(ctx, stor); err != nil {
|
||||
stor.Close()
|
||||
return nil, fmt.Errorf("failed to ensure connector grant types: %w", err)
|
||||
}
|
||||
|
||||
// Ensure issuer ends with /oauth2 for proper path mounting
|
||||
issuer := strings.TrimSuffix(config.Issuer, "/")
|
||||
if !strings.HasSuffix(issuer, "/oauth2") {
|
||||
@@ -118,6 +123,7 @@ func NewProvider(ctx context.Context, config *Config) (*Provider, error) {
|
||||
Storage: stor,
|
||||
SkipApprovalScreen: true,
|
||||
SupportedResponseTypes: []string{"code"},
|
||||
AllowedGrantTypes: DefaultGrantTypes,
|
||||
ContinueOnConnectorFailure: true,
|
||||
Logger: logger,
|
||||
PrometheusRegistry: prometheus.NewRegistry(),
|
||||
|
||||
@@ -756,3 +756,30 @@ connectors:
|
||||
assert.Equal(t, DefaultGrantTypes, conn.GrantTypes, "connector %s", conn.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewProvider_SetsGrantTypes(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
provider, err := NewProvider(ctx, &Config{
|
||||
Issuer: "https://example.com/oauth2",
|
||||
Port: 5556,
|
||||
DataDir: t.TempDir(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = provider.Stop(ctx) }()
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/oauth2/.well-known/openid-configuration", nil)
|
||||
rec := httptest.NewRecorder()
|
||||
provider.Handler().ServeHTTP(rec, req)
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
var discovery struct {
|
||||
GrantTypes []string `json:"grant_types_supported"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &discovery))
|
||||
assert.ElementsMatch(t, DefaultGrantTypes, discovery.GrantTypes)
|
||||
|
||||
local, err := provider.storage.GetConnector(ctx, "local")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, DefaultGrantTypes, local.GrantTypes)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user