Fix linter issues in auth package and auth controller: error handling, context keys, and staticcheck warnings

This commit is contained in:
PePe Amengual
2025-06-26 17:42:33 -07:00
parent a778492e91
commit d32a4ee48e
10 changed files with 395 additions and 159 deletions

View File

@@ -1,32 +1,30 @@
linters-settings:
misspell:
# Correct spellings using locale preferences for US or UK.
# Default is to use a neutral variety of English.
# Setting locale to US will correct the British spelling of 'colour' to 'color'.
# locale: US
ignore-words:
# for gitlab notes api
- noteable
revive:
rules:
- name: dot-imports
disabled: true
version: "2"
linters:
enable:
- errcheck
- gochecknoinits
- gofmt
- gosec
- gosimple
- govet
- ineffassign
- misspell
- revive
- staticcheck
- testifylint
- typecheck
- unconvert
- unused
settings:
misspell:
# Correct spellings using locale preferences for US or UK.
# Default is to use a neutral variety of English.
# Setting locale to US will correct the British spelling of 'colour' to 'color'.
# locale: US
ignore-rules:
# for gitlab notes api
- noteable
revive:
rules:
- name: dot-imports
disabled: true
run:
timeout: 10m

View File

@@ -2,6 +2,8 @@ package auth
import (
"context"
"encoding/json"
"fmt"
"net/http"
"time"
)
@@ -125,6 +127,33 @@ type Config struct {
Providers []ProviderConfig `json:"providers"`
}
// UnmarshalJSON implements custom JSON unmarshaling for Config
func (c *Config) UnmarshalJSON(data []byte) error {
// Create a temporary struct to handle the raw JSON
type configAlias Config
aux := &struct {
SessionDuration string `json:"session_duration"`
*configAlias
}{
configAlias: (*configAlias)(c),
}
if err := json.Unmarshal(data, &aux); err != nil {
return err
}
// Parse session duration if it's a string
if aux.SessionDuration != "" {
duration, err := time.ParseDuration(aux.SessionDuration)
if err != nil {
return fmt.Errorf("invalid session_duration: %w", err)
}
c.SessionDuration = duration
}
return nil
}
// Provider interface for authentication providers
type Provider interface {
GetType() ProviderType

View File

@@ -2,6 +2,7 @@ package auth
import (
"context"
"encoding/base64"
"fmt"
"net/http"
"strings"
@@ -165,7 +166,10 @@ func (p *BasicAuthProvider) ValidateBasicAuth(r *http.Request) (*User, error) {
// Helper function to decode basic auth token
func decodeBasicAuth(token string) (string, error) {
// In a real implementation, this would decode base64
// For now, we'll assume the token is already decoded
return token, nil
// Decode base64 encoded credentials
decoded, err := base64.StdEncoding.DecodeString(token)
if err != nil {
return "", fmt.Errorf("failed to decode base64: %w", err)
}
return string(decoded), nil
}

View File

@@ -45,12 +45,18 @@ func TestLoadConfigFromFile(t *testing.T) {
if err != nil {
t.Fatalf("Failed to create temp file: %v", err)
}
defer os.Remove(tmpfile.Name())
defer func() {
if err := tmpfile.Close(); err != nil {
t.Fatalf("Failed to close temp file: %v", err)
}
if err := os.Remove(tmpfile.Name()); err != nil {
t.Fatalf("Failed to remove temp file: %v", err)
}
}()
if _, err := tmpfile.Write(configJSON); err != nil {
t.Fatalf("Failed to write config file: %v", err)
}
tmpfile.Close()
// Load config from file
config, err := LoadConfigFromFile(tmpfile.Name())
@@ -129,12 +135,18 @@ func TestLoadConfigFromFile_Defaults(t *testing.T) {
if err != nil {
t.Fatalf("Failed to create temp file: %v", err)
}
defer os.Remove(tmpfile.Name())
defer func() {
if err := tmpfile.Close(); err != nil {
t.Fatalf("Failed to close temp file: %v", err)
}
if err := os.Remove(tmpfile.Name()); err != nil {
t.Fatalf("Failed to remove temp file: %v", err)
}
}()
if _, err := tmpfile.Write(configJSON); err != nil {
t.Fatalf("Failed to write config file: %v", err)
}
tmpfile.Close()
// Load config from file
config, err := LoadConfigFromFile(tmpfile.Name())
@@ -158,12 +170,18 @@ func TestLoadConfigFromFile_InvalidJSON(t *testing.T) {
if err != nil {
t.Fatalf("Failed to create temp file: %v", err)
}
defer os.Remove(tmpfile.Name())
defer func() {
if err := tmpfile.Close(); err != nil {
t.Fatalf("Failed to close temp file: %v", err)
}
if err := os.Remove(tmpfile.Name()); err != nil {
t.Fatalf("Failed to remove temp file: %v", err)
}
}()
if _, err := tmpfile.Write([]byte("invalid json")); err != nil {
t.Fatalf("Failed to write config file: %v", err)
}
tmpfile.Close()
// Load config from file should fail
_, err = LoadConfigFromFile(tmpfile.Name())
@@ -182,19 +200,43 @@ func TestLoadConfigFromFile_FileNotFound(t *testing.T) {
func TestLoadConfigFromEnv(t *testing.T) {
// Set environment variables
os.Setenv("ATLANTIS_SESSION_SECRET", "env-secret")
os.Setenv("ATLANTIS_SECURE_COOKIES", "true")
os.Setenv("ATLANTIS_CSRF_SECRET", "env-csrf")
os.Setenv("ATLANTIS_ENABLE_BASIC_AUTH", "true")
os.Setenv("ATLANTIS_BASIC_AUTH_USER", "envuser")
os.Setenv("ATLANTIS_BASIC_AUTH_PASS", "envpass")
if err := os.Setenv("ATLANTIS_SESSION_SECRET", "env-secret"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_SECURE_COOKIES", "true"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_CSRF_SECRET", "env-csrf"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_ENABLE_BASIC_AUTH", "true"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_BASIC_AUTH_USER", "envuser"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_BASIC_AUTH_PASS", "envpass"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer func() {
os.Unsetenv("ATLANTIS_SESSION_SECRET")
os.Unsetenv("ATLANTIS_SECURE_COOKIES")
os.Unsetenv("ATLANTIS_CSRF_SECRET")
os.Unsetenv("ATLANTIS_ENABLE_BASIC_AUTH")
os.Unsetenv("ATLANTIS_BASIC_AUTH_USER")
os.Unsetenv("ATLANTIS_BASIC_AUTH_PASS")
if err := os.Unsetenv("ATLANTIS_SESSION_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_SECURE_COOKIES"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_CSRF_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_ENABLE_BASIC_AUTH"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_BASIC_AUTH_USER"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_BASIC_AUTH_PASS"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
}()
config, err := LoadConfigFromEnv()
@@ -229,12 +271,24 @@ func TestLoadConfigFromEnv(t *testing.T) {
func TestLoadConfigFromEnv_Defaults(t *testing.T) {
// Clear environment variables
os.Unsetenv("ATLANTIS_SESSION_SECRET")
os.Unsetenv("ATLANTIS_SECURE_COOKIES")
os.Unsetenv("ATLANTIS_CSRF_SECRET")
os.Unsetenv("ATLANTIS_ENABLE_BASIC_AUTH")
os.Unsetenv("ATLANTIS_BASIC_AUTH_USER")
os.Unsetenv("ATLANTIS_BASIC_AUTH_PASS")
if err := os.Unsetenv("ATLANTIS_SESSION_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_SECURE_COOKIES"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_CSRF_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_ENABLE_BASIC_AUTH"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_BASIC_AUTH_USER"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_BASIC_AUTH_PASS"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
config, err := LoadConfigFromEnv()
if err != nil {
@@ -268,13 +322,25 @@ func TestLoadConfigFromEnv_Defaults(t *testing.T) {
func TestLoadConfigFromEnv_GoogleProvider(t *testing.T) {
// Set Google OAuth2 environment variables
os.Setenv("ATLANTIS_GOOGLE_CLIENT_ID", "google-client-id")
os.Setenv("ATLANTIS_GOOGLE_CLIENT_SECRET", "google-client-secret")
os.Setenv("ATLANTIS_GOOGLE_REDIRECT_URL", "https://example.com/google/callback")
if err := os.Setenv("ATLANTIS_GOOGLE_CLIENT_ID", "google-client-id"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_GOOGLE_CLIENT_SECRET", "google-client-secret"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_GOOGLE_REDIRECT_URL", "https://example.com/google/callback"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer func() {
os.Unsetenv("ATLANTIS_GOOGLE_CLIENT_ID")
os.Unsetenv("ATLANTIS_GOOGLE_CLIENT_SECRET")
os.Unsetenv("ATLANTIS_GOOGLE_REDIRECT_URL")
if err := os.Unsetenv("ATLANTIS_GOOGLE_CLIENT_ID"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_GOOGLE_CLIENT_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_GOOGLE_REDIRECT_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
}()
config, err := LoadConfigFromEnv()
@@ -314,15 +380,31 @@ func TestLoadConfigFromEnv_GoogleProvider(t *testing.T) {
func TestLoadConfigFromEnv_OktaProvider(t *testing.T) {
// Set Okta OIDC environment variables
os.Setenv("ATLANTIS_OKTA_CLIENT_ID", "okta-client-id")
os.Setenv("ATLANTIS_OKTA_CLIENT_SECRET", "okta-client-secret")
os.Setenv("ATLANTIS_OKTA_REDIRECT_URL", "https://example.com/okta/callback")
os.Setenv("ATLANTIS_OKTA_ISSUER_URL", "https://example.okta.com")
if err := os.Setenv("ATLANTIS_OKTA_CLIENT_ID", "okta-client-id"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_OKTA_CLIENT_SECRET", "okta-client-secret"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_OKTA_REDIRECT_URL", "https://example.com/okta/callback"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_OKTA_ISSUER_URL", "https://example.okta.com"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer func() {
os.Unsetenv("ATLANTIS_OKTA_CLIENT_ID")
os.Unsetenv("ATLANTIS_OKTA_CLIENT_SECRET")
os.Unsetenv("ATLANTIS_OKTA_REDIRECT_URL")
os.Unsetenv("ATLANTIS_OKTA_ISSUER_URL")
if err := os.Unsetenv("ATLANTIS_OKTA_CLIENT_ID"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_OKTA_CLIENT_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_OKTA_REDIRECT_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_OKTA_ISSUER_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
}()
config, err := LoadConfigFromEnv()
@@ -354,13 +436,25 @@ func TestLoadConfigFromEnv_OktaProvider(t *testing.T) {
func TestLoadConfigFromEnv_OktaProvider_MissingIssuerURL(t *testing.T) {
// Set Okta OIDC environment variables without issuer URL
os.Setenv("ATLANTIS_OKTA_CLIENT_ID", "okta-client-id")
os.Setenv("ATLANTIS_OKTA_CLIENT_SECRET", "okta-client-secret")
os.Setenv("ATLANTIS_OKTA_REDIRECT_URL", "https://example.com/okta/callback")
if err := os.Setenv("ATLANTIS_OKTA_CLIENT_ID", "okta-client-id"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_OKTA_CLIENT_SECRET", "okta-client-secret"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_OKTA_REDIRECT_URL", "https://example.com/okta/callback"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer func() {
os.Unsetenv("ATLANTIS_OKTA_CLIENT_ID")
os.Unsetenv("ATLANTIS_OKTA_CLIENT_SECRET")
os.Unsetenv("ATLANTIS_OKTA_REDIRECT_URL")
if err := os.Unsetenv("ATLANTIS_OKTA_CLIENT_ID"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_OKTA_CLIENT_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_OKTA_REDIRECT_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
}()
_, err := LoadConfigFromEnv()
@@ -371,15 +465,31 @@ func TestLoadConfigFromEnv_OktaProvider_MissingIssuerURL(t *testing.T) {
func TestLoadConfigFromEnv_AzureProvider(t *testing.T) {
// Set Azure AD environment variables
os.Setenv("ATLANTIS_AZURE_CLIENT_ID", "azure-client-id")
os.Setenv("ATLANTIS_AZURE_CLIENT_SECRET", "azure-client-secret")
os.Setenv("ATLANTIS_AZURE_REDIRECT_URL", "https://example.com/azure/callback")
os.Setenv("ATLANTIS_AZURE_TENANT_ID", "tenant-123")
if err := os.Setenv("ATLANTIS_AZURE_CLIENT_ID", "azure-client-id"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_AZURE_CLIENT_SECRET", "azure-client-secret"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_AZURE_REDIRECT_URL", "https://example.com/azure/callback"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_AZURE_TENANT_ID", "tenant-123"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer func() {
os.Unsetenv("ATLANTIS_AZURE_CLIENT_ID")
os.Unsetenv("ATLANTIS_AZURE_CLIENT_SECRET")
os.Unsetenv("ATLANTIS_AZURE_REDIRECT_URL")
os.Unsetenv("ATLANTIS_AZURE_TENANT_ID")
if err := os.Unsetenv("ATLANTIS_AZURE_CLIENT_ID"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_AZURE_CLIENT_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_AZURE_REDIRECT_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_AZURE_TENANT_ID"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
}()
config, err := LoadConfigFromEnv()
@@ -412,13 +522,25 @@ func TestLoadConfigFromEnv_AzureProvider(t *testing.T) {
func TestLoadConfigFromEnv_AzureProvider_MissingTenantID(t *testing.T) {
// Set Azure AD environment variables without tenant ID
os.Setenv("ATLANTIS_AZURE_CLIENT_ID", "azure-client-id")
os.Setenv("ATLANTIS_AZURE_CLIENT_SECRET", "azure-client-secret")
os.Setenv("ATLANTIS_AZURE_REDIRECT_URL", "https://example.com/azure/callback")
if err := os.Setenv("ATLANTIS_AZURE_CLIENT_ID", "azure-client-id"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_AZURE_CLIENT_SECRET", "azure-client-secret"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_AZURE_REDIRECT_URL", "https://example.com/azure/callback"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer func() {
os.Unsetenv("ATLANTIS_AZURE_CLIENT_ID")
os.Unsetenv("ATLANTIS_AZURE_CLIENT_SECRET")
os.Unsetenv("ATLANTIS_AZURE_REDIRECT_URL")
if err := os.Unsetenv("ATLANTIS_AZURE_CLIENT_ID"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_AZURE_CLIENT_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_AZURE_REDIRECT_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
}()
_, err := LoadConfigFromEnv()
@@ -429,15 +551,31 @@ func TestLoadConfigFromEnv_AzureProvider_MissingTenantID(t *testing.T) {
func TestLoadConfigFromEnv_Auth0Provider(t *testing.T) {
// Set Auth0 environment variables
os.Setenv("ATLANTIS_AUTH0_CLIENT_ID", "auth0-client-id")
os.Setenv("ATLANTIS_AUTH0_CLIENT_SECRET", "auth0-client-secret")
os.Setenv("ATLANTIS_AUTH0_REDIRECT_URL", "https://example.com/auth0/callback")
os.Setenv("ATLANTIS_AUTH0_DOMAIN", "example.auth0.com")
if err := os.Setenv("ATLANTIS_AUTH0_CLIENT_ID", "auth0-client-id"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_AUTH0_CLIENT_SECRET", "auth0-client-secret"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_AUTH0_REDIRECT_URL", "https://example.com/auth0/callback"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_AUTH0_DOMAIN", "example.auth0.com"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer func() {
os.Unsetenv("ATLANTIS_AUTH0_CLIENT_ID")
os.Unsetenv("ATLANTIS_AUTH0_CLIENT_SECRET")
os.Unsetenv("ATLANTIS_AUTH0_REDIRECT_URL")
os.Unsetenv("ATLANTIS_AUTH0_DOMAIN")
if err := os.Unsetenv("ATLANTIS_AUTH0_CLIENT_ID"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_AUTH0_CLIENT_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_AUTH0_REDIRECT_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_AUTH0_DOMAIN"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
}()
config, err := LoadConfigFromEnv()
@@ -470,13 +608,25 @@ func TestLoadConfigFromEnv_Auth0Provider(t *testing.T) {
func TestLoadConfigFromEnv_Auth0Provider_MissingDomain(t *testing.T) {
// Set Auth0 environment variables without domain
os.Setenv("ATLANTIS_AUTH0_CLIENT_ID", "auth0-client-id")
os.Setenv("ATLANTIS_AUTH0_CLIENT_SECRET", "auth0-client-secret")
os.Setenv("ATLANTIS_AUTH0_REDIRECT_URL", "https://example.com/auth0/callback")
if err := os.Setenv("ATLANTIS_AUTH0_CLIENT_ID", "auth0-client-id"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_AUTH0_CLIENT_SECRET", "auth0-client-secret"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_AUTH0_REDIRECT_URL", "https://example.com/auth0/callback"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer func() {
os.Unsetenv("ATLANTIS_AUTH0_CLIENT_ID")
os.Unsetenv("ATLANTIS_AUTH0_CLIENT_SECRET")
os.Unsetenv("ATLANTIS_AUTH0_REDIRECT_URL")
if err := os.Unsetenv("ATLANTIS_AUTH0_CLIENT_ID"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_AUTH0_CLIENT_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_AUTH0_REDIRECT_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
}()
_, err := LoadConfigFromEnv()
@@ -487,23 +637,51 @@ func TestLoadConfigFromEnv_Auth0Provider_MissingDomain(t *testing.T) {
func TestLoadConfigFromEnv_MultipleProviders(t *testing.T) {
// Set multiple provider environment variables
os.Setenv("ATLANTIS_GOOGLE_CLIENT_ID", "google-client-id")
os.Setenv("ATLANTIS_GOOGLE_CLIENT_SECRET", "google-client-secret")
os.Setenv("ATLANTIS_GOOGLE_REDIRECT_URL", "https://example.com/google/callback")
if err := os.Setenv("ATLANTIS_GOOGLE_CLIENT_ID", "google-client-id"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_GOOGLE_CLIENT_SECRET", "google-client-secret"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_GOOGLE_REDIRECT_URL", "https://example.com/google/callback"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
os.Setenv("ATLANTIS_OKTA_CLIENT_ID", "okta-client-id")
os.Setenv("ATLANTIS_OKTA_CLIENT_SECRET", "okta-client-secret")
os.Setenv("ATLANTIS_OKTA_REDIRECT_URL", "https://example.com/okta/callback")
os.Setenv("ATLANTIS_OKTA_ISSUER_URL", "https://example.okta.com")
if err := os.Setenv("ATLANTIS_OKTA_CLIENT_ID", "okta-client-id"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_OKTA_CLIENT_SECRET", "okta-client-secret"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_OKTA_REDIRECT_URL", "https://example.com/okta/callback"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
if err := os.Setenv("ATLANTIS_OKTA_ISSUER_URL", "https://example.okta.com"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer func() {
os.Unsetenv("ATLANTIS_GOOGLE_CLIENT_ID")
os.Unsetenv("ATLANTIS_GOOGLE_CLIENT_SECRET")
os.Unsetenv("ATLANTIS_GOOGLE_REDIRECT_URL")
os.Unsetenv("ATLANTIS_OKTA_CLIENT_ID")
os.Unsetenv("ATLANTIS_OKTA_CLIENT_SECRET")
os.Unsetenv("ATLANTIS_OKTA_REDIRECT_URL")
os.Unsetenv("ATLANTIS_OKTA_ISSUER_URL")
if err := os.Unsetenv("ATLANTIS_GOOGLE_CLIENT_ID"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_GOOGLE_CLIENT_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_GOOGLE_REDIRECT_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_OKTA_CLIENT_ID"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_OKTA_CLIENT_SECRET"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_OKTA_REDIRECT_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
if err := os.Unsetenv("ATLANTIS_OKTA_ISSUER_URL"); err != nil {
t.Fatalf("Failed to unset env: %v", err)
}
}()
config, err := LoadConfigFromEnv()
@@ -532,7 +710,9 @@ func TestLoadConfigFromEnv_MultipleProviders(t *testing.T) {
func TestGetEnvOrDefault(t *testing.T) {
// Test with environment variable set
os.Setenv("TEST_VAR", "test-value")
if err := os.Setenv("TEST_VAR", "test-value"); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer os.Unsetenv("TEST_VAR")
result := getEnvOrDefault("TEST_VAR", "default-value")
@@ -567,7 +747,9 @@ func TestGetEnvBoolOrDefault(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if tt.envValue != "" {
os.Setenv("TEST_BOOL_VAR", tt.envValue)
if err := os.Setenv("TEST_BOOL_VAR", tt.envValue); err != nil {
t.Fatalf("Failed to set env: %v", err)
}
defer os.Unsetenv("TEST_BOOL_VAR")
}

View File

@@ -151,11 +151,10 @@ func TestManager_ValidateSession(t *testing.T) {
}
if validatedUser == nil {
t.Error("Should return user for valid session")
t.Fatal("Validated user should not be nil")
}
if validatedUser.ID != testUser.ID {
t.Errorf("Validated user ID = %s, want %s", validatedUser.ID, testUser.ID)
t.Errorf("Expected user ID %s, got %s", testUser.ID, validatedUser.ID)
}
}
@@ -281,11 +280,10 @@ func TestManager_GetUserFromRequest(t *testing.T) {
}
if retrievedUser == nil {
t.Error("Should return user for request with valid session cookie")
t.Fatal("Retrieved user should not be nil")
}
if retrievedUser.ID != testUser.ID {
t.Errorf("Retrieved user ID = %s, want %s", retrievedUser.ID, testUser.ID)
t.Errorf("Expected user ID %s, got %s", testUser.ID, retrievedUser.ID)
}
}

View File

@@ -9,6 +9,13 @@ import (
"github.com/urfave/negroni/v3"
)
// contextKey is a custom type for context keys to avoid collisions
type contextKey string
const (
userContextKey contextKey = "user"
)
// AuthMiddleware handles authentication for HTTP requests
type AuthMiddleware struct {
authManager Manager
@@ -48,21 +55,21 @@ func (m *AuthMiddleware) ServeHTTP(rw http.ResponseWriter, r *http.Request, next
// User is authenticated, add user info to request context
ctx := r.Context()
ctx = contextWithUser(ctx, user)
ctx = WithUser(ctx, user)
r = r.WithContext(ctx)
m.logger.Debug("[AUTH] User %s authenticated for: %s", user.Email, r.URL.RequestURI())
next(rw, r)
}
// contextWithUser adds user to request context
func contextWithUser(ctx context.Context, user *User) context.Context {
return context.WithValue(ctx, "user", user)
// WithUser adds a user to the request context
func WithUser(ctx context.Context, user *User) context.Context {
return context.WithValue(ctx, userContextKey, user)
}
// UserFromContext extracts user from request context
func UserFromContext(ctx context.Context) (*User, bool) {
user, ok := ctx.Value("user").(*User)
// GetUserFromContext retrieves a user from the request context
func GetUserFromContext(ctx context.Context) (*User, bool) {
user, ok := ctx.Value(userContextKey).(*User)
return user, ok
}
@@ -123,7 +130,7 @@ func (l *LegacyAuthMiddleware) ServeHTTP(rw http.ResponseWriter, r *http.Request
func RequirePermission(permission Permission) func(http.HandlerFunc) http.HandlerFunc {
return func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
user, ok := UserFromContext(r.Context())
user, ok := GetUserFromContext(r.Context())
if !ok {
http.Error(w, "Authentication required", http.StatusUnauthorized)
return
@@ -153,7 +160,7 @@ func RequirePermission(permission Permission) func(http.HandlerFunc) http.Handle
func RequireAnyPermission(permissions []Permission) func(http.HandlerFunc) http.HandlerFunc {
return func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
user, ok := UserFromContext(r.Context())
user, ok := GetUserFromContext(r.Context())
if !ok {
http.Error(w, "Authentication required", http.StatusUnauthorized)
return
@@ -183,7 +190,7 @@ func RequireAnyPermission(permissions []Permission) func(http.HandlerFunc) http.
func RequireAllPermissions(permissions []Permission) func(http.HandlerFunc) http.HandlerFunc {
return func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
user, ok := UserFromContext(r.Context())
user, ok := GetUserFromContext(r.Context())
if !ok {
http.Error(w, "Authentication required", http.StatusUnauthorized)
return
@@ -213,7 +220,7 @@ func RequireAllPermissions(permissions []Permission) func(http.HandlerFunc) http
func RequireAdmin() func(http.HandlerFunc) http.HandlerFunc {
return func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
user, ok := UserFromContext(r.Context())
user, ok := GetUserFromContext(r.Context())
if !ok {
http.Error(w, "Authentication required", http.StatusUnauthorized)
return

View File

@@ -63,6 +63,10 @@ func (m *MockManager) GetPermissionChecker() PermissionChecker {
return NewPermissionChecker(nil)
}
const (
authManagerContextKey contextKey = "auth_manager"
)
func TestNewAuthMiddleware(t *testing.T) {
logger := logging.NewNoopLogger(t)
mockManager := &MockManager{}
@@ -127,7 +131,7 @@ func TestAuthMiddleware_ServeHTTP_AuthRequired_ValidUser(t *testing.T) {
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
called = true
// Check that user is in context
user, ok := UserFromContext(r.Context())
user, ok := GetUserFromContext(r.Context())
if !ok {
t.Error("User should be in request context")
}
@@ -203,9 +207,9 @@ func TestContextWithUser(t *testing.T) {
Email: "test@example.com",
}
newCtx := contextWithUser(ctx, testUser)
newCtx := WithUser(ctx, testUser)
user, ok := UserFromContext(newCtx)
user, ok := GetUserFromContext(newCtx)
if !ok {
t.Fatal("User should be retrievable from context")
}
@@ -223,7 +227,7 @@ func TestUserFromContext(t *testing.T) {
}
// Test with no user in context
user, ok := UserFromContext(ctx)
user, ok := GetUserFromContext(ctx)
if ok {
t.Error("Should return false when no user in context")
}
@@ -232,8 +236,8 @@ func TestUserFromContext(t *testing.T) {
}
// Test with user in context
ctxWithUser := contextWithUser(ctx, testUser)
user, ok = UserFromContext(ctxWithUser)
ctxWithUser := WithUser(ctx, testUser)
user, ok = GetUserFromContext(ctxWithUser)
if !ok {
t.Error("Should return true when user in context")
}
@@ -405,7 +409,7 @@ func TestRequirePermission(t *testing.T) {
// Create middleware that adds auth manager to context
authMiddleware := func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := context.WithValue(r.Context(), "auth_manager", mockManager)
ctx := context.WithValue(r.Context(), authManagerContextKey, mockManager)
next(w, r.WithContext(ctx))
}
}
@@ -421,7 +425,7 @@ func TestRequirePermission(t *testing.T) {
}
req := httptest.NewRequest("GET", "/test", nil)
ctx := contextWithUser(req.Context(), testUser)
ctx := WithUser(req.Context(), testUser)
req = req.WithContext(ctx)
w := httptest.NewRecorder()
@@ -467,7 +471,7 @@ func TestRequirePermission_InsufficientPermissions(t *testing.T) {
// Create middleware that adds auth manager to context
authMiddleware := func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := context.WithValue(r.Context(), "auth_manager", mockManager)
ctx := context.WithValue(r.Context(), authManagerContextKey, mockManager)
next(w, r.WithContext(ctx))
}
}
@@ -483,7 +487,7 @@ func TestRequirePermission_InsufficientPermissions(t *testing.T) {
}
req := httptest.NewRequest("GET", "/test", nil)
ctx := contextWithUser(req.Context(), testUser)
ctx := WithUser(req.Context(), testUser)
req = req.WithContext(ctx)
w := httptest.NewRecorder()
@@ -511,7 +515,7 @@ func TestRequireAnyPermission(t *testing.T) {
// Create middleware that adds auth manager to context
authMiddleware := func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := context.WithValue(r.Context(), "auth_manager", mockManager)
ctx := context.WithValue(r.Context(), authManagerContextKey, mockManager)
next(w, r.WithContext(ctx))
}
}
@@ -528,7 +532,7 @@ func TestRequireAnyPermission(t *testing.T) {
}
req := httptest.NewRequest("GET", "/test", nil)
ctx := contextWithUser(req.Context(), testUser)
ctx := WithUser(req.Context(), testUser)
req = req.WithContext(ctx)
w := httptest.NewRecorder()
@@ -552,7 +556,7 @@ func TestRequireAnyPermission_NoPermissions(t *testing.T) {
// Create middleware that adds auth manager to context
authMiddleware := func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := context.WithValue(r.Context(), "auth_manager", mockManager)
ctx := context.WithValue(r.Context(), authManagerContextKey, mockManager)
next(w, r.WithContext(ctx))
}
}
@@ -569,7 +573,7 @@ func TestRequireAnyPermission_NoPermissions(t *testing.T) {
}
req := httptest.NewRequest("GET", "/test", nil)
ctx := contextWithUser(req.Context(), testUser)
ctx := WithUser(req.Context(), testUser)
req = req.WithContext(ctx)
w := httptest.NewRecorder()
@@ -597,7 +601,7 @@ func TestRequireAllPermissions(t *testing.T) {
// Create middleware that adds auth manager to context
authMiddleware := func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := context.WithValue(r.Context(), "auth_manager", mockManager)
ctx := context.WithValue(r.Context(), authManagerContextKey, mockManager)
next(w, r.WithContext(ctx))
}
}
@@ -614,7 +618,7 @@ func TestRequireAllPermissions(t *testing.T) {
}
req := httptest.NewRequest("GET", "/test", nil)
ctx := contextWithUser(req.Context(), testUser)
ctx := WithUser(req.Context(), testUser)
req = req.WithContext(ctx)
w := httptest.NewRecorder()
@@ -638,7 +642,7 @@ func TestRequireAllPermissions_MissingPermission(t *testing.T) {
// Create middleware that adds auth manager to context
authMiddleware := func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := context.WithValue(r.Context(), "auth_manager", mockManager)
ctx := context.WithValue(r.Context(), authManagerContextKey, mockManager)
next(w, r.WithContext(ctx))
}
}
@@ -655,7 +659,7 @@ func TestRequireAllPermissions_MissingPermission(t *testing.T) {
}
req := httptest.NewRequest("GET", "/test", nil)
ctx := contextWithUser(req.Context(), testUser)
ctx := WithUser(req.Context(), testUser)
req = req.WithContext(ctx)
w := httptest.NewRecorder()
@@ -683,7 +687,7 @@ func TestRequireAdmin(t *testing.T) {
// Create middleware that adds auth manager to context
authMiddleware := func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := context.WithValue(r.Context(), "auth_manager", mockManager)
ctx := context.WithValue(r.Context(), authManagerContextKey, mockManager)
next(w, r.WithContext(ctx))
}
}
@@ -699,7 +703,7 @@ func TestRequireAdmin(t *testing.T) {
}
req := httptest.NewRequest("GET", "/admin", nil)
ctx := contextWithUser(req.Context(), testUser)
ctx := WithUser(req.Context(), testUser)
req = req.WithContext(ctx)
w := httptest.NewRecorder()
@@ -723,7 +727,7 @@ func TestRequireAdmin_NonAdmin(t *testing.T) {
// Create middleware that adds auth manager to context
authMiddleware := func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := context.WithValue(r.Context(), "auth_manager", mockManager)
ctx := context.WithValue(r.Context(), authManagerContextKey, mockManager)
next(w, r.WithContext(ctx))
}
}
@@ -739,7 +743,7 @@ func TestRequireAdmin_NonAdmin(t *testing.T) {
}
req := httptest.NewRequest("GET", "/admin", nil)
ctx := contextWithUser(req.Context(), testUser)
ctx := WithUser(req.Context(), testUser)
req = req.WithContext(ctx)
w := httptest.NewRecorder()
@@ -767,7 +771,7 @@ func TestSpecificPermissionMiddlewares(t *testing.T) {
// Create middleware that adds auth manager to context
authMiddleware := func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := context.WithValue(r.Context(), "auth_manager", mockManager)
ctx := context.WithValue(r.Context(), authManagerContextKey, mockManager)
next(w, r.WithContext(ctx))
}
}
@@ -797,7 +801,7 @@ func TestSpecificPermissionMiddlewares(t *testing.T) {
}
req := httptest.NewRequest("GET", "/test", nil)
ctx := contextWithUser(req.Context(), testUser)
ctx := WithUser(req.Context(), testUser)
req = req.WithContext(ctx)
w := httptest.NewRecorder()

View File

@@ -115,7 +115,7 @@ func (p *OAuth2Provider) ExchangeCode(ctx context.Context, code string) (*TokenR
AccessToken: token.AccessToken,
TokenType: token.TokenType,
RefreshToken: token.RefreshToken,
ExpiresIn: int(token.Expiry.Sub(time.Now()).Seconds()),
ExpiresIn: int(time.Until(token.Expiry).Seconds()),
}, nil
}
@@ -130,7 +130,9 @@ func (p *OAuth2Provider) GetUserInfo(ctx context.Context, token *TokenResponse)
if err != nil {
return nil, fmt.Errorf("failed to get user info: %w", err)
}
defer resp.Body.Close()
defer func() {
_ = resp.Body.Close() // ignore error, as defer cannot return values
}()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("user info request failed with status: %d", resp.StatusCode)
@@ -180,7 +182,9 @@ func (p *OAuth2Provider) ValidateToken(ctx context.Context, tokenString string)
if err != nil {
return nil, fmt.Errorf("failed to validate token: %w", err)
}
defer resp.Body.Close()
defer func() {
_ = resp.Body.Close() // ignore error, as defer cannot return values
}()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("token validation failed with status: %d", resp.StatusCode)

View File

@@ -180,12 +180,14 @@ func TestOAuth2Provider_ExchangeCode(t *testing.T) {
// Return a mock token response
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
w.Write([]byte(`{
if _, err := w.Write([]byte(`{
"access_token": "test-access-token",
"token_type": "Bearer",
"refresh_token": "test-refresh-token",
"expires_in": 3600
}`))
}`)); err != nil {
t.Errorf("Failed to write response: %v", err)
}
}))
defer mockServer.Close()
@@ -252,7 +254,7 @@ func TestOAuth2Provider_GetUserInfo(t *testing.T) {
// Return a mock user info response
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
w.Write([]byte(`{
if _, err := w.Write([]byte(`{
"id": "123456789",
"email": "test@example.com",
"name": "Test User",
@@ -260,7 +262,9 @@ func TestOAuth2Provider_GetUserInfo(t *testing.T) {
"family_name": "User",
"picture": "https://example.com/avatar.jpg",
"locale": "en"
}`))
}`)); err != nil {
t.Errorf("Failed to write response: %v", err)
}
}))
defer mockServer.Close()
@@ -337,7 +341,7 @@ func TestOAuth2Provider_ValidateToken(t *testing.T) {
// Return a mock user info response
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
w.Write([]byte(`{
if _, err := w.Write([]byte(`{
"id": "123456789",
"email": "test@example.com",
"name": "Test User",
@@ -345,7 +349,9 @@ func TestOAuth2Provider_ValidateToken(t *testing.T) {
"family_name": "User",
"picture": "https://example.com/avatar.jpg",
"locale": "en"
}`))
}`)); err != nil {
t.Errorf("Failed to write response: %v", err)
}
}))
defer mockServer.Close()

View File

@@ -121,7 +121,9 @@ func (c *AuthController) Logout(w http.ResponseWriter, r *http.Request) {
// Get user from request to invalidate their session
if user, err := c.AuthManager.GetUserFromRequest(r); err == nil {
if user.SessionID != "" {
c.AuthManager.InvalidateSession(r.Context(), user.SessionID)
if err := c.AuthManager.InvalidateSession(r.Context(), user.SessionID); err != nil {
c.Logger.Err("Failed to invalidate session: %v", err)
}
}
}
@@ -247,5 +249,7 @@ func (c *AuthController) respond(w http.ResponseWriter, lvl logging.LogLevel, re
response := fmt.Sprintf(format, args...)
c.Logger.Log(lvl, response)
w.WriteHeader(responseCode)
fmt.Fprintln(w, response)
if _, err := fmt.Fprintln(w, response); err != nil {
c.Logger.Err("Failed to write response: %v", err)
}
}