mirror of
https://git.vectorsigma.ru/public/atlantis.git
synced 2026-08-06 10:58:35 +00:00
Fix linter issues in auth package and auth controller: error handling, context keys, and staticcheck warnings
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user