Skip to content

Commit 4852cfd

Browse files
feat: support provider-specific OAuth whitelists
1 parent 3194f4b commit 4852cfd

7 files changed

Lines changed: 57 additions & 5 deletions

File tree

‎.env.example‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -101,6 +101,10 @@ TINYAUTH_OAUTH_PROVIDERS_name_CLIENTID=
101101
TINYAUTH_OAUTH_PROVIDERS_name_CLIENTSECRET=
102102
# Path to the file containing the OAuth client secret.
103103
TINYAUTH_OAUTH_PROVIDERS_name_CLIENTSECRETFILE=
104+
# Comma-separated list of allowed OAuth domains for this provider.
105+
TINYAUTH_OAUTH_PROVIDERS_name_WHITELIST=
106+
# Path to the OAuth whitelist file for this provider.
107+
TINYAUTH_OAUTH_PROVIDERS_name_WHITELISTFILE=
104108
# OAuth scopes.
105109
TINYAUTH_OAUTH_PROVIDERS_name_SCOPES=
106110
# OAuth redirect URL.

‎internal/bootstrap/app_bootstrap.go‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -112,6 +112,13 @@ func (app *BootstrapApp) Setup() error {
112112
app.runtime.OAuthProviders = app.config.OAuth.Providers
113113

114114
for id, provider := range app.runtime.OAuthProviders {
115+
providerWhitelist, err := utils.GetStringList(provider.Whitelist, provider.WhitelistFile)
116+
if err != nil {
117+
return fmt.Errorf("failed to load oauth whitelist for provider %s: %w", id, err)
118+
}
119+
120+
provider.Whitelist = providerWhitelist
121+
115122
secret := utils.GetSecret(provider.ClientSecret, provider.ClientSecretFile)
116123
provider.ClientSecret = secret
117124
provider.ClientSecretFile = ""

‎internal/controller/oauth_controller.go‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -183,7 +183,7 @@ func (controller *OAuthController) oauthCallbackHandler(c *gin.Context) {
183183
return
184184
}
185185

186-
if !controller.auth.IsEmailWhitelisted(user.Email) {
186+
if !controller.auth.IsEmailWhitelisted(req.Provider, user.Email) {
187187
controller.log.App.Warn().Str("email", user.Email).Msg("Email not whitelisted, denying access")
188188
controller.log.AuditLoginFailure(user.Email, req.Provider, c.ClientIP(), "email not whitelisted")
189189

‎internal/middleware/context_middleware.go‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -205,7 +205,7 @@ func (m *ContextMiddleware) cookieAuth(ctx context.Context, uuid string, ip stri
205205
return nil, nil, fmt.Errorf("oauth provider from session cookie not found: %s", userContext.OAuth.ID)
206206
}
207207

208-
if !m.auth.IsEmailWhitelisted(userContext.OAuth.Email) {
208+
if !m.auth.IsEmailWhitelisted(userContext.OAuth.ID, userContext.OAuth.Email) {
209209
m.auth.DeleteSession(ctx, uuid)
210210
return nil, nil, fmt.Errorf("email from session cookie not whitelisted: %s", userContext.OAuth.Email)
211211
}

‎internal/model/config.go‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -225,6 +225,8 @@ type OAuthServiceConfig struct {
225225
ClientID string `description:"OAuth client ID." yaml:"clientId"`
226226
ClientSecret string `description:"OAuth client secret." yaml:"clientSecret"`
227227
ClientSecretFile string `description:"Path to the file containing the OAuth client secret." yaml:"clientSecretFile"`
228+
Whitelist []string `description:"Comma-separated list of allowed OAuth domains for this provider." yaml:"whitelist"`
229+
WhitelistFile string `description:"Path to the OAuth whitelist file for this provider." yaml:"whitelistFile"`
228230
Scopes []string `description:"OAuth scopes." yaml:"scopes"`
229231
RedirectURL string `description:"OAuth redirect URL." yaml:"redirectUrl"`
230232
AuthURL string `description:"OAuth authorization URL." yaml:"authUrl"`

‎internal/service/auth_service.go‎

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -285,10 +285,15 @@ func (auth *AuthService) RecordLoginAttempt(identifier string, success bool) {
285285
}
286286
}
287287

288-
func (auth *AuthService) IsEmailWhitelisted(email string) bool {
289-
match, err := utils.CheckFilter(strings.Join(auth.runtime.OAuthWhitelist, ","), email)
288+
func (auth *AuthService) IsEmailWhitelisted(provider string, email string) bool {
289+
whitelist := auth.runtime.OAuthWhitelist
290+
if providerConfig, ok := auth.runtime.OAuthProviders[provider]; ok && len(providerConfig.Whitelist) > 0 {
291+
whitelist = providerConfig.Whitelist
292+
}
293+
294+
match, err := utils.CheckFilter(strings.Join(whitelist, ","), email)
290295
if err != nil {
291-
auth.log.App.Warn().Err(err).Str("email", email).Msg("Invalid email filter pattern")
296+
auth.log.App.Warn().Err(err).Str("provider", provider).Str("email", email).Msg("Invalid email filter pattern")
292297
return false
293298
}
294299
return match
Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,34 @@
1+
package service
2+
3+
import (
4+
"testing"
5+
6+
"github.com/stretchr/testify/assert"
7+
"github.com/tinyauthapp/tinyauth/internal/model"
8+
"github.com/tinyauthapp/tinyauth/internal/utils/logger"
9+
)
10+
11+
func TestIsEmailWhitelistedUsesProviderSpecificList(t *testing.T) {
12+
log := logger.NewLogger().WithTestConfig()
13+
log.Init()
14+
15+
auth := &AuthService{
16+
log: log,
17+
runtime: model.RuntimeConfig{
18+
OAuthWhitelist: []string{"global@example.com"},
19+
OAuthProviders: map[string]model.OAuthServiceConfig{
20+
"github": {
21+
Whitelist: []string{"github@example.com"},
22+
},
23+
"pocketid": {
24+
Whitelist: []string{"pocket@example.com"},
25+
},
26+
},
27+
},
28+
}
29+
30+
assert.True(t, auth.IsEmailWhitelisted("github", "github@example.com"))
31+
assert.False(t, auth.IsEmailWhitelisted("github", "pocket@example.com"))
32+
assert.True(t, auth.IsEmailWhitelisted("pocketid", "pocket@example.com"))
33+
assert.True(t, auth.IsEmailWhitelisted("google", "global@example.com"))
34+
}

0 commit comments

Comments
 (0)