diff --git a/Package.swift b/Package.swift index 506cb5e..96ea295 100644 --- a/Package.swift +++ b/Package.swift @@ -20,7 +20,6 @@ let package = Package( ], dependencies: [ .package(url: "https://github.com/ChimeHQ/OAuthenticator.git", branch: "main"), -// .package(path: "../OAuthenticator"), .package(url: "https://github.com/vapor/jwt-kit.git", from: "5.0.0"), .package(url: "https://github.com/SparrowTek/NetworkingKit.git", branch: "main"), ], diff --git a/Sources/CoreATProtocol/APEnvironment.swift b/Sources/CoreATProtocol/APEnvironment.swift index 45de5e1..1b061c2 100644 --- a/Sources/CoreATProtocol/APEnvironment.swift +++ b/Sources/CoreATProtocol/APEnvironment.swift @@ -2,31 +2,64 @@ // APEnvironment.swift // CoreATProtocol // -// Created by Thomas Rademaker on 10/10/25. -// import JWTKit +/// Session-scoped state for a single AT Protocol session. +/// +/// Today this is a process-wide singleton accessed via ``shared``. A future +/// major release will make it instance-based so a single process can host +/// multiple concurrent sessions without the routing layer needing a global +/// handle. Until then, callers should treat ``shared`` as the one-and-only +/// session and avoid reading or mutating its state from outside CoreATProtocol +/// — use the free functions in `CoreATProtocol.swift` (`setup`, `updateTokens`, +/// etc.) as the public API. @APActor -public class APEnvironment { - public static var current: APEnvironment = APEnvironment() - +public final class ATProtoSession { + public static let shared = ATProtoSession() + public var host: String? public var accessToken: String? public var refreshToken: String? - public var atProtocoldelegate: CoreATProtocolDelegate? + public var atProtocolDelegate: CoreATProtocolDelegate? public var tokenRefreshHandler: (@Sendable () async throws -> Bool)? public var dpopPrivateKey: ES256PrivateKey? public var dpopKeys: JWTKeyCollection? public let dpopNonceStore = DPoPNonceStore() public let clockSkewStore = ClockSkewStore() public let routerDelegate = APRouterDelegate() - - private init() {} - -// func setup(apiKey: String, apiSecret: String, userAgent: String) { -// self.apiKey = apiKey -// self.apiSecret = apiSecret -// self.userAgent = userAgent -// } + + internal init() {} + + /// Clears all mutable session state. Intended for test harnesses. + /// + /// DPoP nonce and clock-skew stores are reset to empty; tokens, keys, host + /// and delegates are nilled out. The router delegate and its coordinator + /// are preserved (they hold no user-scoped state). + public func reset() async { + host = nil + accessToken = nil + refreshToken = nil + atProtocolDelegate = nil + tokenRefreshHandler = nil + dpopPrivateKey = nil + dpopKeys = nil + await dpopNonceStore.clear() + await clockSkewStore.update(offset: 0) + } +} + +// MARK: - Deprecation shim +// +// The original name was `APEnvironment` and the singleton was `.current`. +// The renames below preserve source compatibility for downstream callers +// (bskyKit, EffemKit, Atprosphere) while emitting fix-its that point at the +// new names. A future major release will remove these aliases. + +@available(*, deprecated, renamed: "ATProtoSession") +public typealias APEnvironment = ATProtoSession + +extension ATProtoSession { + @available(*, deprecated, renamed: "shared") + public static var current: ATProtoSession { shared } } diff --git a/Sources/CoreATProtocol/CoreATProtocol.swift b/Sources/CoreATProtocol/CoreATProtocol.swift index dba0263..1b1529c 100644 --- a/Sources/CoreATProtocol/CoreATProtocol.swift +++ b/Sources/CoreATProtocol/CoreATProtocol.swift @@ -36,27 +36,27 @@ public extension CoreATProtocolDelegate { @APActor public func setup(hostURL: String?, accessJWT: String?, refreshJWT: String?, delegate: CoreATProtocolDelegate? = nil) { - APEnvironment.current.host = hostURL - APEnvironment.current.accessToken = accessJWT - APEnvironment.current.refreshToken = refreshJWT - APEnvironment.current.atProtocoldelegate = delegate + ATProtoSession.shared.host = hostURL + ATProtoSession.shared.accessToken = accessJWT + ATProtoSession.shared.refreshToken = refreshJWT + ATProtoSession.shared.atProtocolDelegate = delegate } @APActor public func setDelegate(_ delegate: CoreATProtocolDelegate) { - APEnvironment.current.atProtocoldelegate = delegate + ATProtoSession.shared.atProtocolDelegate = delegate } @APActor public func setTokenRefreshHandler(_ handler: (@Sendable () async throws -> Bool)?) { - APEnvironment.current.tokenRefreshHandler = handler + ATProtoSession.shared.tokenRefreshHandler = handler } @APActor public func setDPoPPrivateKey(pem: String?) async throws { guard let pem, !pem.isEmpty else { - APEnvironment.current.dpopPrivateKey = nil - APEnvironment.current.dpopKeys = nil + ATProtoSession.shared.dpopPrivateKey = nil + ATProtoSession.shared.dpopKeys = nil return } @@ -64,17 +64,17 @@ public func setDPoPPrivateKey(pem: String?) async throws { let keys = JWTKeyCollection() await keys.add(ecdsa: privateKey) - APEnvironment.current.dpopPrivateKey = privateKey - APEnvironment.current.dpopKeys = keys + ATProtoSession.shared.dpopPrivateKey = privateKey + ATProtoSession.shared.dpopKeys = keys } @APActor public func updateTokens(access: String?, refresh: String?) { - APEnvironment.current.accessToken = access - APEnvironment.current.refreshToken = refresh + ATProtoSession.shared.accessToken = access + ATProtoSession.shared.refreshToken = refresh } @APActor public func update(hostURL: String?) { - APEnvironment.current.host = hostURL + ATProtoSession.shared.host = hostURL } diff --git a/Sources/CoreATProtocol/Networking.swift b/Sources/CoreATProtocol/Networking.swift index 5f9206d..7670b67 100644 --- a/Sources/CoreATProtocol/Networking.swift +++ b/Sources/CoreATProtocol/Networking.swift @@ -44,8 +44,8 @@ func shouldPerformRequest(lastFetched: Double, timeLimit: Int = 3600) -> Bool { guard lastFetched != 0 else { return true } let currentTime = Date.now let lastFetchTime = Date(timeIntervalSince1970: lastFetched) - guard let differenceInMinutes = Calendar.current.dateComponents([.second], from: lastFetchTime, to: currentTime).second else { return false } - return differenceInMinutes >= timeLimit + guard let differenceInSeconds = Calendar.current.dateComponents([.second], from: lastFetchTime, to: currentTime).second else { return false } + return differenceInSeconds >= timeLimit } @NetworkingKitActor @@ -53,10 +53,10 @@ public class APRouterDelegate: NetworkRouterDelegate { private let refreshCoordinator = TokenRefreshCoordinator() public func intercept(_ request: inout URLRequest) async { - guard let accessToken = await APEnvironment.current.accessToken else { return } + guard let accessToken = await ATProtoSession.shared.accessToken else { return } - if let dpopKey = await APEnvironment.current.dpopPrivateKey, - let keys = await APEnvironment.current.dpopKeys { + if let dpopKey = await ATProtoSession.shared.dpopPrivateKey, + let keys = await ATProtoSession.shared.dpopKeys { // DPoP-bound token: use "DPoP" scheme + DPoP proof header do { let proof = try await generateDPoPProof(for: request, accessToken: accessToken, privateKey: dpopKey, keys: keys) @@ -97,8 +97,8 @@ public class APRouterDelegate: NetworkRouterDelegate { // Read the nonce at proof-generation time so a concurrent update // between intercept() and sign() is observed on the next retry. - let nonce = await APEnvironment.current.dpopNonceStore.get() - let issuedAt = await APEnvironment.current.clockSkewStore.serverAdjustedNow() + let nonce = await ATProtoSession.shared.dpopNonceStore.get() + let issuedAt = await ATProtoSession.shared.clockSkewStore.serverAdjustedNow() // ath: base64url-encoded SHA-256 hash of the access token (RFC 9449 §4.2) let hash = SHA256.hash(data: Data(accessToken.utf8)) @@ -134,13 +134,13 @@ public class APRouterDelegate: NetworkRouterDelegate { // DPoP rejections often correlate with skew, so harvest this eagerly. let dateHeader = response.value(forHTTPHeaderField: "Date") if let dateHeader { - await APEnvironment.current.clockSkewStore.updateFromServerDate(dateHeader) + await ATProtoSession.shared.clockSkewStore.updateFromServerDate(dateHeader) } let headerNonce = response.value(forHTTPHeaderField: "DPoP-Nonce") ?? response.value(forHTTPHeaderField: "dpop-nonce") if let headerNonce { - await APEnvironment.current.dpopNonceStore.update(headerNonce) + await ATProtoSession.shared.dpopNonceStore.update(headerNonce) lastErrorHadNonceHeader = true } else { lastErrorHadNonceHeader = false @@ -181,7 +181,7 @@ public class APRouterDelegate: NetworkRouterDelegate { } private func refreshViaOAuth() async throws -> Bool { - guard let handler = await APEnvironment.current.tokenRefreshHandler else { + guard let handler = await ATProtoSession.shared.tokenRefreshHandler else { return false } return try await refreshCoordinator.refresh(using: handler) diff --git a/Sources/CoreATProtocol/OAuth/ATProtoOAuth.swift b/Sources/CoreATProtocol/OAuth/ATProtoOAuth.swift index 0c9c795..3401543 100644 --- a/Sources/CoreATProtocol/OAuth/ATProtoOAuth.swift +++ b/Sources/CoreATProtocol/OAuth/ATProtoOAuth.swift @@ -507,7 +507,7 @@ public final class ATProtoOAuth: Sendable { // Strip query params and fragments from htu per DPoP spec let htu = stripQueryAndFragment(from: params.requestEndpoint) - let issuedAt = await APEnvironment.current.clockSkewStore.serverAdjustedNow() + let issuedAt = await ATProtoSession.shared.clockSkewStore.serverAdjustedNow() let payload = DPoPPayload( htm: params.httpMethod, diff --git a/Tests/CoreATProtocolTests/Base64URLTests.swift b/Tests/CoreATProtocolTests/Base64URLTests.swift new file mode 100644 index 0000000..801cb1a --- /dev/null +++ b/Tests/CoreATProtocolTests/Base64URLTests.swift @@ -0,0 +1,32 @@ +import Foundation +import Testing +@testable import CoreATProtocol + +@Suite("Base64URL encoding") +struct Base64URLTests { + @Test("Data encoding strips padding and swaps alphabet") + func dataEncoding() { + // Bytes chosen so the standard base64 output contains both + and /. + let input = Data([0xfb, 0xff, 0xbf, 0xfe]) + #expect(input.base64EncodedString() == "+/+//g==") + #expect(input.base64URLEncodedString() == "-_-__g") + } + + @Test("Empty data round-trips to empty string") + func emptyData() { + #expect(Data().base64URLEncodedString() == "") + } + + @Test("String helper rewrites alphabet and strips padding") + func stringRewrite() { + #expect("+/+//g==".base64URLEncoded() == "-_-__g") + #expect("abcd".base64URLEncoded() == "abcd") + } + + @Test("Variable-length inputs produce unpadded output") + func variableLengthInputs() { + #expect(Data([0x01]).base64URLEncodedString() == "AQ") + #expect(Data([0x01, 0x02]).base64URLEncodedString() == "AQI") + #expect(Data([0x01, 0x02, 0x03]).base64URLEncodedString() == "AQID") + } +} diff --git a/Tests/CoreATProtocolTests/DPoPStoreTests.swift b/Tests/CoreATProtocolTests/DPoPStoreTests.swift new file mode 100644 index 0000000..4ba2008 --- /dev/null +++ b/Tests/CoreATProtocolTests/DPoPStoreTests.swift @@ -0,0 +1,83 @@ +import Foundation +import Testing +@testable import CoreATProtocol + +@Suite("ClockSkewStore") +struct ClockSkewStoreTests { + @Test("Parses RFC 1123 Date header and records offset") + func rfc1123() async { + let store = ClockSkewStore() + let local = Date(timeIntervalSince1970: 1_700_000_000) + let serverHeader = "Thu, 14 Nov 2023 22:13:40 GMT" + let serverDate = try! #require(ClockSkewStore.parse(serverHeader)) + + await store.updateFromServerDate(serverHeader, localNow: local) + let offset = await store.offset + #expect(offset == serverDate.timeIntervalSince(local)) + } + + @Test("Ignores nil or unparseable header without mutating state") + func malformedHeader() async { + let store = ClockSkewStore(offset: 123) + await store.updateFromServerDate(nil) + #expect(await store.offset == 123) + + await store.updateFromServerDate("not a date") + #expect(await store.offset == 123) + } + + @Test("serverAdjustedNow applies the stored offset") + func adjustedNow() async { + let store = ClockSkewStore(offset: 42) + let base = Date(timeIntervalSince1970: 1_000_000) + let adjusted = await store.serverAdjustedNow(localNow: base) + #expect(adjusted.timeIntervalSince1970 == 1_000_042) + } + + @Test("Parses the RFC 850 legacy Date format") + func rfc850() { + #expect(ClockSkewStore.parse("Sunday, 06-Nov-94 08:49:37 GMT") != nil) + } +} + +@Suite("DPoPNonceStore") +struct DPoPNonceStoreTests { + @Test("Starts empty") + func initialState() async { + let store = DPoPNonceStore() + #expect(await store.get() == nil) + } + + @Test("Update and read back") + func update() async { + let store = DPoPNonceStore() + await store.update("abc") + #expect(await store.get() == "abc") + await store.update("xyz") + #expect(await store.get() == "xyz") + } + + @Test("Concurrent updates settle on a valid value") + func concurrentUpdates() async { + let store = DPoPNonceStore() + + await withTaskGroup(of: Void.self) { group in + for i in 0..<100 { + group.addTask { await store.update("nonce-\(i)") } + } + } + + let final = await store.get() + let valid = (0..<100).map { "nonce-\($0)" } + #expect(final != nil) + #expect(valid.contains(final ?? "")) + } + + @Test("clear() resets to nil") + func clearResets() async { + let store = DPoPNonceStore(nonce: "start") + #expect(await store.get() == "start") + await store.clear() + #expect(await store.get() == nil) + } +} diff --git a/Tests/CoreATProtocolTests/IdentityResolverValidationTests.swift b/Tests/CoreATProtocolTests/IdentityResolverValidationTests.swift new file mode 100644 index 0000000..338d772 --- /dev/null +++ b/Tests/CoreATProtocolTests/IdentityResolverValidationTests.swift @@ -0,0 +1,86 @@ +import Foundation +import Testing +@testable import CoreATProtocol + +@Suite("IdentityResolver hostname validation") +struct HostnameValidationTests { + @Test("Valid DNS-style handles are accepted") + func valid() throws { + try IdentityResolver.validateHostname("alice.bsky.social") + try IdentityResolver.validateHostname("example.com") + try IdentityResolver.validateHostname("a-b.c-d.example.com") + } + + @Test("Handles with URL authority characters are rejected", + arguments: [ + "alice@bsky.social", + "alice.bsky.social/extra", + "alice.bsky.social?q=1", + "alice.bsky.social#frag", + "alice.bsky.social:8080", + "alice bsky social", + "aliçe.bsky.social", + ]) + func rejectsInjection(_ input: String) { + #expect(throws: IdentityError.self) { + try IdentityResolver.validateHostname(input) + } + } + + @Test("Empty, single-label, and malformed handles are rejected", + arguments: [ + "", + "com", + ".com", + "alice.", + "alice..com", + "-alice.com", + "alice-.com", + ]) + func rejectsMalformed(_ input: String) { + #expect(throws: IdentityError.self) { + try IdentityResolver.validateHostname(input) + } + } + + @Test("Labels over 63 characters are rejected") + func tooLongLabel() { + let longLabel = String(repeating: "a", count: 64) + #expect(throws: IdentityError.self) { + try IdentityResolver.validateHostname("\(longLabel).com") + } + } +} + +@Suite("DID document PDS lookup") +struct DIDDocumentPDSTests { + @Test("Service matched by #atproto_pds suffix") + func matchBySuffix() { + let doc = DIDDocument( + id: "did:plc:abc", + alsoKnownAs: ["at://alice.bsky.social"], + service: [DIDService(id: "#atproto_pds", type: "Whatever", serviceEndpoint: "https://pds.example")] + ) + #expect(doc.pdsEndpoint == "https://pds.example") + } + + @Test("Service matched by AtprotoPersonalDataServer type") + func matchByType() { + let doc = DIDDocument( + id: "did:plc:abc", + alsoKnownAs: ["at://alice.bsky.social"], + service: [DIDService(id: "#other", type: "AtprotoPersonalDataServer", serviceEndpoint: "https://pds.example")] + ) + #expect(doc.pdsEndpoint == "https://pds.example") + } + + @Test("Document with no matching service returns nil") + func noMatch() { + let doc = DIDDocument( + id: "did:plc:abc", + alsoKnownAs: nil, + service: [DIDService(id: "#other", type: "Unrelated", serviceEndpoint: "https://example.com")] + ) + #expect(doc.pdsEndpoint == nil) + } +} diff --git a/Tests/CoreATProtocolTests/RefreshLoginTests.swift b/Tests/CoreATProtocolTests/RefreshLoginTests.swift new file mode 100644 index 0000000..6be2807 --- /dev/null +++ b/Tests/CoreATProtocolTests/RefreshLoginTests.swift @@ -0,0 +1,85 @@ +import Foundation +import Testing +@testable import CoreATProtocol + +@Suite("ATProtoOAuth refresh error disambiguation") +struct RefreshLoginTests { + private func makeConfig() -> ATProtoOAuthConfig { + ATProtoOAuthConfig( + clientMetadataURL: "https://example.com/client-metadata.json", + redirectURI: "example://callback" + ) + } + + @Test("No stored login returns nil (genuine no-op)") + func noStoredLogin() async throws { + let storage = ATProtoAuthStorage( + retrieveLogin: { nil }, + storeLogin: { _ in }, + retrievePrivateKey: { nil }, + storePrivateKey: { _ in } + ) + let client = await ATProtoOAuth(config: makeConfig(), storage: storage) + let result = try await client.refreshLoginIfNeeded(handle: nil) + #expect(result == nil) + } + + @Test("Missing persisted DPoP key throws .refreshKeyUnavailable") + func missingPersistedKey() async throws { + let expiredLogin = Login( + accessToken: Token(value: "access", expiry: .distantPast), + refreshToken: Token(value: "refresh"), + scopes: "atproto", + issuingServer: "https://bsky.social" + ) + let storage = ATProtoAuthStorage( + retrieveLogin: { expiredLogin }, + storeLogin: { _ in }, + retrievePrivateKey: { nil }, + storePrivateKey: { _ in } + ) + let client = await ATProtoOAuth(config: makeConfig(), storage: storage) + + await #expect(throws: ATProtoOAuthError.self) { + _ = try await client.refreshLoginIfNeeded(handle: nil, force: true) + } + } + + @Test("Expired refresh token throws .refreshTokenUnavailable") + func expiredRefreshToken() async throws { + let expiredLogin = Login( + accessToken: Token(value: "access", expiry: .distantPast), + refreshToken: Token(value: "refresh", expiry: .distantPast), + scopes: "atproto", + issuingServer: "https://bsky.social" + ) + + // Mint a working key via a throwaway client, then feed it back into the + // persisted-key constructor so hasPersistedKey is true and we reach the + // refresh-token-validity check. + let bootstrapStorage = ATProtoAuthStorage( + retrieveLogin: { nil }, + storeLogin: { _ in }, + retrievePrivateKey: { nil }, + storePrivateKey: { _ in } + ) + let bootstrap = await ATProtoOAuth(config: makeConfig(), storage: bootstrapStorage) + let pem = await bootstrap.privateKeyPEM + + let storage = ATProtoAuthStorage( + retrieveLogin: { expiredLogin }, + storeLogin: { _ in }, + retrievePrivateKey: { pem.data(using: .utf8) }, + storePrivateKey: { _ in } + ) + let client = try await ATProtoOAuth( + config: makeConfig(), + storage: storage, + privateKeyPEM: pem + ) + + await #expect(throws: ATProtoOAuthError.self) { + _ = try await client.refreshLoginIfNeeded(handle: nil, force: true) + } + } +}