From d32a4ee48efb358c4db13649140d43524ea65734 Mon Sep 17 00:00:00 2001 From: PePe Amengual <2208324+jamengual@users.noreply.github.com> Date: Thu, 26 Jun 2025 17:42:33 -0700 Subject: [PATCH] Fix linter issues in auth package and auth controller: error handling, context keys, and staticcheck warnings --- .golangci.yml | 30 +-- server/auth/auth.go | 29 +++ server/auth/basic.go | 10 +- server/auth/config_test.go | 358 +++++++++++++++++++------- server/auth/manager_test.go | 10 +- server/auth/middleware.go | 29 ++- server/auth/middleware_test.go | 52 ++-- server/auth/oauth2.go | 10 +- server/auth/oauth2_test.go | 18 +- server/controllers/auth_controller.go | 8 +- 10 files changed, 395 insertions(+), 159 deletions(-) diff --git a/.golangci.yml b/.golangci.yml index 0afa70118..87c69b267 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -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 diff --git a/server/auth/auth.go b/server/auth/auth.go index d11c7f126..d84ea3678 100644 --- a/server/auth/auth.go +++ b/server/auth/auth.go @@ -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 diff --git a/server/auth/basic.go b/server/auth/basic.go index 8429eb21d..ed181b168 100644 --- a/server/auth/basic.go +++ b/server/auth/basic.go @@ -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 } \ No newline at end of file diff --git a/server/auth/config_test.go b/server/auth/config_test.go index 87eddff02..e9d4c3974 100644 --- a/server/auth/config_test.go +++ b/server/auth/config_test.go @@ -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") } diff --git a/server/auth/manager_test.go b/server/auth/manager_test.go index 7e56a8979..e94d1b3b1 100644 --- a/server/auth/manager_test.go +++ b/server/auth/manager_test.go @@ -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) } } diff --git a/server/auth/middleware.go b/server/auth/middleware.go index a112bea14..61eb5fa24 100644 --- a/server/auth/middleware.go +++ b/server/auth/middleware.go @@ -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 diff --git a/server/auth/middleware_test.go b/server/auth/middleware_test.go index 99be3c298..55309ff2e 100644 --- a/server/auth/middleware_test.go +++ b/server/auth/middleware_test.go @@ -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() diff --git a/server/auth/oauth2.go b/server/auth/oauth2.go index 58b77d3b1..7334074c5 100644 --- a/server/auth/oauth2.go +++ b/server/auth/oauth2.go @@ -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) diff --git a/server/auth/oauth2_test.go b/server/auth/oauth2_test.go index 3a2ac3013..9d1f4bac0 100644 --- a/server/auth/oauth2_test.go +++ b/server/auth/oauth2_test.go @@ -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() diff --git a/server/controllers/auth_controller.go b/server/controllers/auth_controller.go index 0e1b2d1ca..50b972f6d 100644 --- a/server/controllers/auth_controller.go +++ b/server/controllers/auth_controller.go @@ -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) + } } \ No newline at end of file