Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 52 additions & 1 deletion Sources/JWTKit/JWTKeyCollection.swift
Original file line number Diff line number Diff line change
Expand Up @@ -11,9 +11,17 @@ import Foundation
/// This actor provides methods to manage multiple keys. It can be used to verify and decode JWTs, as well as to sign and encode JWTs.
/// It also facilitates the encoding and decoding of JWTs using custom or default JSON encoders and decoders.
public actor JWTKeyCollection: Sendable {
private enum Signer {
private enum Signer: Equatable {
case jwt(JWTSigner)
case jwk(JWKSigner)

static func == (lhs: JWTKeyCollection.Signer, rhs: JWTKeyCollection.Signer) -> Bool {
Comment thread
0xTim marked this conversation as resolved.
switch (lhs, rhs) {
case (.jwt(let lhsSigner), .jwt(let rhsSigner)): lhsSigner === rhsSigner
case (.jwk(let lhsSigner), .jwk(let rhsSigner)): lhsSigner === rhsSigner
default: false
}
}
}

private var storage: [JWKIdentifier: Signer]
Expand Down Expand Up @@ -69,6 +77,49 @@ public actor JWTKeyCollection: Sendable {
return self
}

/// Removes the key with the selected KID from the collection.
/// If the default matches, that one is also removed.
/// - Parameter kid: The KID identifying the signer to the collection.
/// - Returns: True if the key was found and removed, false otherwise.
@discardableResult
public func remove(kid: JWKIdentifier) -> Bool {
let value = self.storage.removeValue(forKey: kid)
if value == self.default {
self.default = nil
}
return value != nil
}

/// Removes the default signer.
Comment thread
ptoffy marked this conversation as resolved.
/// The signer might still exist in the collection as non-default one.
/// - Returns: True if the default signer was found and removed, false otherwise.
@discardableResult
public func clearDefault() -> Bool {
if self.default != nil {
self.default = nil
return true
}
return false
}

/// Removes all keys from the collection except the ones defined in the `kids` parameter.
/// - Parameter kids: The KIDs that should be kept during removal.
/// - Parameter clearingDefault: If true, clears the default signer regardless of whether
Comment thread
ptoffy marked this conversation as resolved.
/// its KID is in the exception list, however keeping it in the collection if it's
/// KID is defined in the exception list. If false, preserves the default signer.
/// - Returns: The number of keys removed from storage.
@discardableResult
public func removeAll(except kids: [JWKIdentifier] = [], clearingDefault: Bool = true) -> Int {
let kidsToKeep = Set(kids)
self.storage = self.storage.filter { kidsToKeep.contains($0.key) }
let originalCount = self.storage.count
if clearingDefault {
self.default = nil
}

return originalCount - self.storage.count
}

/// Adds a `JWKS` (JSON Web Key Set) to the collection by decoding a JSON string.
///
/// - Parameter json: A JSON string representing a JWKS.
Expand Down
128 changes: 128 additions & 0 deletions Tests/JWTKitTests/JWTKitTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -852,6 +852,134 @@ struct JWTKitTests {
#expect(header.field1 == nil)
#expect(header.field2 == .string("value2"))
}

@Test("Remove single key by kid")
func testRemoveSingleKey() async throws {
let keys = await JWTKeyCollection()
.add(hmac: "key-1", digestAlgorithm: .sha256, kid: "v1")
.add(hmac: "key-2", digestAlgorithm: .sha256, kid: "v2")

let payload = TestPayload(
sub: "vapor",
name: "Test",
admin: false,
exp: .init(value: Date().addingTimeInterval(3600))
)

let token1 = try await keys.sign(payload, kid: "v1")
let token2 = try await keys.sign(payload, kid: "v2")

_ = try await keys.verify(token1, as: TestPayload.self)
_ = try await keys.verify(token2, as: TestPayload.self)

let removed = await keys.remove(kid: "v1")
#expect(removed == true)

_ = try await keys.verify(token2, as: TestPayload.self)

await #expect(throws: JWTError.self) {
_ = try await keys.verify(token1, as: TestPayload.self)
}
}

@Test("Remove non-existent key returns false")
func testRemoveNonExistentKey() async throws {
let keys = await JWTKeyCollection()
.add(hmac: "key-1", digestAlgorithm: .sha256, kid: "v1")

let removed = await keys.remove(kid: "non-existent")
#expect(removed == false)
}

@Test("Remove default key clears default")
func testRemoveDefaultKey() async throws {
let keys = await JWTKeyCollection()
.add(hmac: "key-1", digestAlgorithm: .sha256) // No kid = default

let payload = TestPayload(
sub: "vapor",
name: "Test",
admin: false,
exp: .init(value: Date().addingTimeInterval(3600))
)

let token = try await keys.sign(payload)
_ = try await keys.verify(token, as: TestPayload.self)

let removed = await keys.clearDefault()
#expect(removed == true)

await #expect(throws: JWTError.noKeyProvided) {
_ = try await keys.sign(payload)
}
}

@Test("RemoveAll except specified keys")
func testRemoveAllExcept() async throws {
Comment thread
ptoffy marked this conversation as resolved.
let keys = await JWTKeyCollection()
.add(hmac: "key-1", digestAlgorithm: .sha256, kid: "v1")
.add(hmac: "key-2", digestAlgorithm: .sha256, kid: "v2")
.add(hmac: "key-3", digestAlgorithm: .sha256, kid: "v3")
.add(hmac: "key-4", digestAlgorithm: .sha256, kid: "v4")

let payload = TestPayload(
sub: "vapor",
name: "Test",
admin: false,
exp: .init(value: Date().addingTimeInterval(3600))
)

let token1 = try await keys.sign(payload, kid: "v1")
let token2 = try await keys.sign(payload, kid: "v2")
let token3 = try await keys.sign(payload, kid: "v3")
let token4 = try await keys.sign(payload, kid: "v4")

await keys.removeAll(except: ["v2", "v4"])

// v2 and v4 still work
_ = try await keys.verify(token2, as: TestPayload.self)
_ = try await keys.verify(token4, as: TestPayload.self)

// v1 and v3 no longer work
await #expect(throws: JWTError.noKeyProvided) {
_ = try await keys.verify(token1, as: TestPayload.self)
}
await #expect(throws: JWTError.noKeyProvided) {
_ = try await keys.verify(token3, as: TestPayload.self)
}
}

@Test("RemoveAll including default")
func testRemoveAllWithDefault() async throws {
let keys = await JWTKeyCollection()
.add(hmac: "key-0", digestAlgorithm: .sha256)
.add(hmac: "key-1", digestAlgorithm: .sha256, kid: "v1")
.add(hmac: "key-2", digestAlgorithm: .sha256, kid: "v2")

let payload = TestPayload(
sub: "vapor",
name: "Test",
admin: false,
exp: .init(value: Date().addingTimeInterval(3600))
)

let token0 = try await keys.sign(payload)
let token1 = try await keys.sign(payload, kid: "v1")
let token2 = try await keys.sign(payload, kid: "v2")

await keys.removeAll(except: ["v2"], clearingDefault: true)

// v2 still works
_ = try await keys.verify(token2, as: TestPayload.self)

// v0 and v1 no longer work
await #expect(throws: JWTError.noKeyProvided) {
_ = try await keys.verify(token0, as: TestPayload.self)
}
await #expect(throws: JWTError.noKeyProvided) {
_ = try await keys.verify(token1, as: TestPayload.self)
}
}
}

let microsoftJWKS = """
Expand Down
Loading