/*
 * Copyright (c) Huawei Technologies Co., Ltd. 2022-2024. All rights reserved.
 */
package jwt4cj

public class BaseVerification <: Verification {
    private var algorithm: Algorithm
    private var expectedChecks: ArrayList<ExpectedCheckHolder>
    private var defaultLeeway: Int64
    private var customLeeways: HashMap<String, Int64>
    private var ignoreIssuedAt1: Bool = false
    private var time: DateTime

    init(algorithm: Algorithm) {
        this.algorithm = algorithm
        this.expectedChecks = ArrayList()
        this.customLeeways = HashMap()
        this.defaultLeeway = 0
        this.time = DateTime.ofEpoch(second: 0, nanosecond: 0)
    }

    public func withAnyOfAudience(audience: Array<String>): Verification {
        let value: ArrayList<String> = ArrayList<String>(audience)
        return withAnyOfAudience(value)
    }

    public func withArrayClaim(name: String, items: Array<String>): Verification {
        let value: ArrayList<String> = ArrayList<String>(items)
        return withArrayClaim(name, value)
    }

    public func withAudience(audience: Array<String>): Verification {
        let value: ArrayList<String> = ArrayList<String>(audience)
        return withAudience(value)
    }

    public func withArrayClaim(name: String, items: Array<Int64>): Verification {
        let value: ArrayList<Int64> = ArrayList<Int64>(items)
        return withArrayClaim(name, value)
    }

    public func withIssuer(issuer: Array<String>): Verification {
        let value: ArrayList<String> = ArrayList<String>(issuer)

        let fun = {
            claim: Claim, _: DecodedJWT =>
            if (claim.isMissing()) {
                return false
            }
            try {
                let _: String = claim.asString()
            } catch (e: Exception) {
                return false
            }

            if (value.size == 0 || !value.contains(claim.asString())) {
                throw IncorrectClaimException(
                    "The Claim 'iss' value doesn't match the required issuer.",
                    RegisteredClaims.ISSUER,
                    claim
                )
            }
            return true
        }
        addCheck(RegisteredClaims.ISSUER, fun)
        return this
    }

    public func withIssuer(issuer: String): Verification {
        return withIssuer([issuer])
    }

    public func withSubject(subject: String): Verification {
        let fun = extractFun({claim: Claim => subject == claim.asString()})

        addCheck(RegisteredClaims.SUBJECT, fun)
        return this
    }

    public func withAudience(audience: ArrayList<String>): Verification {
        let value: ArrayList<String> = audience
        let fun = {
            claim: Claim, decodedJWT: DecodedJWT =>
            if (claim.isMissing()) {
                return false
            }

            if (!assertValidAudienceClaim(decodedJWT.getAudience(), value, true)) {
                throw IncorrectClaimException(
                    "The Claim 'aud' value doesn't contain the required audience.",
                    RegisteredClaims.AUDIENCE,
                    claim
                )
            }
            return true
        }

        addCheck(RegisteredClaims.AUDIENCE, fun)
        return this
    }

    public func withAnyOfAudience(audience: ArrayList<String>): Verification {
        let value: ArrayList<String> = audience

        let fun = {
            claim: Claim, decodedJWT: DecodedJWT =>
            if (claim.isMissing()) {
                return false
            }

            if (!assertValidAudienceClaim(decodedJWT.getAudience(), value, false)) {
                throw IncorrectClaimException(
                    "The Claim 'aud' value doesn't contain the required audience.",
                    RegisteredClaims.AUDIENCE,
                    claim
                )
            }
            return true
        }
        addCheck(RegisteredClaims.AUDIENCE, fun)
        return this
    }

    public func acceptLeeway(leeway: Int64): Verification {
        assertPositive(leeway)
        this.defaultLeeway = leeway
        return this
    }

    public func acceptExpiresAt(leeway: Int64): Verification {
        assertPositive(leeway)
        customLeeways.add(RegisteredClaims.EXPIRES_AT, leeway)
        return this
    }

    public func acceptNotBefore(leeway: Int64): Verification {
        assertPositive(leeway)
        customLeeways.add(RegisteredClaims.NOT_BEFORE, leeway)
        return this
    }

    public func acceptIssuedAt(leeway: Int64): Verification {
        assertPositive(leeway)
        customLeeways.add(RegisteredClaims.ISSUED_AT, leeway)
        return this
    }

    public func ignoreIssuedAt(): Verification {
        this.ignoreIssuedAt1 = true
        return this
    }

    public func withJWTId(jwtId: String): Verification {
        let fun = extractFun({claim: Claim => jwtId == claim.asString()})
        addCheck(RegisteredClaims.JWT_ID, fun)
        return this
    }

    public func withClaimPresence(name: String): Verification {
        //since addCheck already checks presence, we just return true
        let fun = {
            _: Claim, _: DecodedJWT => true
        }
        withClaim(name, fun)
        return this
    }

    public func withNullClaim(name: String): Verification {
        let fun = {claim: Claim, _: DecodedJWT => claim.isNull()}
        withClaim(name, fun)
        return this
    }

    public func withClaim(name: String, value: Bool): Verification {
        let fun = extractFun({claim: Claim => value == claim.asBool()})
        addCheck(name, fun)
        return this
    }

    func extractFun(fun: (Claim) -> Bool): (Claim, DecodedJWT) -> Bool {
        {
            claim: Claim, _: DecodedJWT =>
            if (claim.isMissing()) {
                return false
            }
            try {
                return fun(claim)
            } catch (e: Exception) {
                return false
            }
        }
    }

    public func withClaim(name: String, value: Int64): Verification {
        let fun = extractFun({claim: Claim => value == claim.asInt()})
        addCheck(name, fun)
        return this
    }

    public func withClaim(name: String, value: Float64): Verification {
        let fun = extractFun({claim: Claim => value == claim.asFloat()})
        addCheck(name, fun)
        return this
    }

    public func withClaim(name: String, value: String): Verification {
        // { a: Int64, b: Int64 => print(a-b);false }
        let fun = extractFun({claim: Claim => value == claim.asString()})
        addCheck(name, fun)
        return this
    }

    public func withClaim(name: String, value: DateTime): Verification { // 比较时间
        // Since date-time claims are serialized as epoch seconds,
        // we need to compare them with only seconds-granularity
        let fun = extractFun({claim: Claim => value == claim.asTime()})
        addCheck(name, fun)
        return this
    }

    public func withClaim(name: String, predicate: (Claim, DecodedJWT) -> Bool): Verification {
        // let fun = {claim: Claim, decodedJWT: DecodedJWT => predicate(claim, decodedJWT)}
        // addCheck(name, fun)
        expectedChecks.add(constructExpectedCheck(name, predicate))
        return this
    }

    public func withArrayClaim(name: String, items: ArrayList<String>): Verification {
        let fun = {claim: Claim, _: DecodedJWT => assertValidCollectionClaim(claim, items)}
        addCheck(name, fun)
        return this
    }

    public func withArrayClaim(name: String, items: ArrayList<Int64>): Verification {
        let fun = {claim: Claim, _: DecodedJWT => assertValidCollectionClaim(claim, items)}
        addCheck(name, fun)
        return this
    }

    public func build(): JWTVerifier {
        return this.build(DateTime.now())
    }

    private func build(time: DateTime): JWTVerifier {
        this.time = time
        addMandatoryClaimChecks()
        return BaseJWTVerifier(algorithm, expectedChecks)
    }

    private func getLeewayFor(name: String): Int64 {
        if (customLeeways.contains(name)) {
            return customLeeways.get(name).getOrThrow()
        } else {
            return defaultLeeway
        }
    }

    private func addMandatoryClaimChecks() {
        let expiresAtLeeway: Int64 = getLeewayFor(RegisteredClaims.EXPIRES_AT)
        let notBeforeLeeway: Int64 = getLeewayFor(RegisteredClaims.NOT_BEFORE)
        let issuedAtLeeway: Int64 = getLeewayFor(RegisteredClaims.ISSUED_AT)

        let fun = {
            claim: Claim, _: DecodedJWT => assertValidInstantClaim(
                RegisteredClaims.EXPIRES_AT,
                claim,
                expiresAtLeeway,
                true
            )
        }
        let fun1 = {
            claim: Claim, _: DecodedJWT => assertValidInstantClaim(
                RegisteredClaims.NOT_BEFORE,
                claim,
                notBeforeLeeway,
                false
            )
        }
        let fun2 = {
            claim: Claim, _: DecodedJWT => assertValidInstantClaim(
                RegisteredClaims.ISSUED_AT,
                claim,
                issuedAtLeeway,
                false
            )
        }

        expectedChecks.add(constructExpectedCheck(RegisteredClaims.EXPIRES_AT, fun))
        expectedChecks.add(constructExpectedCheck(RegisteredClaims.NOT_BEFORE, fun1))
        if (!ignoreIssuedAt1) {
            expectedChecks.add(constructExpectedCheck(RegisteredClaims.ISSUED_AT, fun2))
        }
    }

    private func assertValidCollectionClaim(claim: Claim, expectedClaimValue: ArrayList<String>): Bool {
        if (claim.isMissing()) {
            return false
        }
        let claimArr = asStringList(claim)
        for (pattern in expectedClaimValue) {
            if (!claimArr.contains(pattern)) {
                return false
            }
        }
        return true
    }

    private func assertValidCollectionClaim(claim: Claim, expectedClaimValue: ArrayList<Int64>): Bool {
        if (claim.isMissing()) {
            return false
        }
        let claimArrInt = asIntList(claim)
        for (pattern in expectedClaimValue) {
            if (!claimArrInt.contains(pattern)) {
                return false
            }
        }
        return true
    }

    private func assertValidInstantClaim(claimName: String, claim: Claim, leeway: Int64, shouldBeFuture: Bool): Bool {
        if (claim.isMissing()) {
            return true
        }

        var claimVal: DateTime = claim.asTime()
        var now: DateTime = time
        var isValid: Bool
        if (shouldBeFuture) {
            isValid = assertInstantIsFuture(claimVal, leeway, now)
            if (!isValid) {
                throw TokenExpiredException("The Token has expired on ${claimVal}.", claimVal)
            }
        } else {
            isValid = assertInstantIsPast(claimVal, leeway, now)
            print(isValid)
            if (!isValid) {
                throw IncorrectClaimException("The Token can't be used before ${claimVal}.", claimName, claim)
            }
        }
        return true
    }

    private func assertInstantIsFuture(claimVal: DateTime, leeway: Int64, now: DateTime): Bool {
        var between = Duration.second * leeway
        return !((now - between) > claimVal)
    }

    private func assertInstantIsPast(claimVal: DateTime, leeway: Int64, now: DateTime): Bool {
        var between = Duration.second * leeway
        return !((now + between) < claimVal)
    }

    private func assertValidAudienceClaim(
        audience: ArrayList<String>,
        values: ArrayList<String>,
        shouldContainAll: Bool
    ): Bool {
        var containsAll: Bool = true
        for (value in values) {
            if (!audience.contains(value)) {
                containsAll = false
                break
            }
        }

        var disjoint: Bool = true
        for (value in values) {
            if (audience.contains(value)) {
                disjoint = false
                break
            }
        }

        return !((shouldContainAll && !containsAll) || (!shouldContainAll && disjoint))
    }

    private func assertPositive(leeway: Int64) {
        if (leeway < 0) {
            throw IllegalArgumentException("Leeway value can't be negative.")
        }
    }

    private func addCheck(name: String, predicate: (Claim, DecodedJWT) -> Bool) {
        let fun = {
            claim: Claim, decodedJWT: DecodedJWT =>
            if (claim.isMissing()) {
                throw MissingClaimException(name)
            }
            return predicate(claim, decodedJWT)
        }
        expectedChecks.add(constructExpectedCheck(name, fun))
    }

    private func constructExpectedCheck(claimName: String, check: (Claim, DecodedJWT) -> Bool): ExpectedCheckHolder {
        return ExpectedCheckHolderImpl(claimName, check)
    }
}