package io.delphiplatform.api.security;

import com.nimbusds.jose.KeySourceException;
import com.nimbusds.jose.jwk.source.DefaultJWKSetCache;
import com.nimbusds.jose.jwk.source.JWKSetCache;
import com.nimbusds.jose.jwk.source.RemoteJWKSet;
import com.nimbusds.jose.proc.JWSAlgorithmFamilyJWSKeySelector;
import com.nimbusds.jose.proc.JWSKeySelector;
import com.nimbusds.jose.proc.SecurityContext;
import com.nimbusds.jose.util.DefaultResourceRetriever;
import com.nimbusds.jwt.proc.DefaultJWTProcessor;

import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.security.config.annotation.web.builders.HttpSecurity;
import org.springframework.security.config.annotation.web.configuration.EnableWebSecurity;
import org.springframework.security.config.annotation.web.configuration.WebSecurityConfigurerAdapter;
import org.springframework.security.config.http.SessionCreationPolicy;
import org.springframework.security.oauth2.core.DelegatingOAuth2TokenValidator;
import org.springframework.security.oauth2.core.OAuth2TokenValidator;
import org.springframework.security.oauth2.jwt.Jwt;
import org.springframework.security.oauth2.jwt.JwtDecoder;
import org.springframework.security.oauth2.jwt.JwtValidators;
import org.springframework.security.oauth2.jwt.NimbusJwtDecoder;

import java.net.MalformedURLException;
import java.net.URL;
import java.util.concurrent.TimeUnit;

import io.micrometer.core.instrument.config.InvalidConfigurationException;

@Configuration
@EnableWebSecurity
@ConditionalOnProperty(name="integration.tests.enable", havingValue = "false", matchIfMissing = true)
public class SecurityConfig extends WebSecurityConfigurerAdapter {

    private static final int JWKS_CONNECT_TIMEOUT_MS = 10000;
    private static final int JWKS_READ_TIMEOUT_MS = 10000;

    private final ErrorCodeAuthenticationEntryPoint authenticationEntryPoint;
    private final AuthConfiguration authConfiguration;

    public SecurityConfig(
        ErrorCodeAuthenticationEntryPoint authenticationEntryPoint, AuthConfiguration authConfiguration) {
        this.authenticationEntryPoint = authenticationEntryPoint;
        this.authConfiguration = authConfiguration;
    }

    @Override
    public void configure(HttpSecurity http) throws Exception {
        // @formatter:off
        http
            .sessionManagement()
                .sessionCreationPolicy(SessionCreationPolicy.STATELESS)
                .and()
            .authorizeRequests()
                .antMatchers("/**/health").permitAll()
                .antMatchers("/**/oauth/token").permitAll()
                // Swagger UI
                .antMatchers(
                    "/**/api-docs/**",
                    "/**/ui/**",
                    "/**/swagger-ui/**",
                    "/**/*.swagger.yaml").permitAll()
                .antMatchers("/**").authenticated()
                .anyRequest().permitAll()
                .and()
            .cors()
                .and()
            .csrf()
                .disable()
            .oauth2ResourceServer()
                .authenticationEntryPoint(authenticationEntryPoint)
                .jwt();
        // @formatter:on
    }

    @Bean
    JwtDecoder jwtDecoder() {
        JWSKeySelector<SecurityContext> jwsKeySelector = null;
        try {
            URL jwksUrl = new URL(authConfiguration.getJwksUrl());
            JWKSetCache jwkSetCache = new DefaultJWKSetCache(authConfiguration.getJwksCacheTtlMins(), -1,
                TimeUnit.MINUTES);
            RemoteJWKSet<SecurityContext> jwkSet = new RemoteJWKSet<>(jwksUrl,
                new DefaultResourceRetriever(JWKS_CONNECT_TIMEOUT_MS, JWKS_READ_TIMEOUT_MS), jwkSetCache);
            jwsKeySelector = JWSAlgorithmFamilyJWSKeySelector.fromJWKSource(jwkSet);
        } catch (KeySourceException | MalformedURLException e) {
            throw new InvalidConfigurationException(e.getMessage());
        }

        DefaultJWTProcessor<SecurityContext> jwtProcessor = new DefaultJWTProcessor<>();
        jwtProcessor.setJWSKeySelector(jwsKeySelector);
        NimbusJwtDecoder jwtDecoder = new NimbusJwtDecoder(jwtProcessor);

        OAuth2TokenValidator<Jwt> audienceValidator = new AudienceValidator(authConfiguration.getAudience());
        OAuth2TokenValidator<Jwt> withIssuer = JwtValidators.createDefaultWithIssuer(authConfiguration.getIssuer());
        OAuth2TokenValidator<Jwt> withAudience = new DelegatingOAuth2TokenValidator<>(withIssuer, audienceValidator);

        jwtDecoder.setJwtValidator(withAudience);

        return jwtDecoder;
    }
}
