diff --git a/ios/CodexMeterApp/CodexMeterApp/ViewModels/MeterViewModel.swift b/ios/CodexMeterApp/CodexMeterApp/ViewModels/MeterViewModel.swift index a885e82..beed63d 100644 --- a/ios/CodexMeterApp/CodexMeterApp/ViewModels/MeterViewModel.swift +++ b/ios/CodexMeterApp/CodexMeterApp/ViewModels/MeterViewModel.swift @@ -40,6 +40,7 @@ final class MeterViewModel: ObservableObject { @Published var isOK: Bool = false private let ble = BLEManager() + private let urlSession: URLSession private var fetchTimer: Timer? private var bleTimer: Timer? private var mdnsTask: Task? @@ -103,10 +104,12 @@ final class MeterViewModel: ObservableObject { } } - init() { + init(urlSession: URLSession = .shared, startsDiscovery: Bool = true) { + self.urlSession = urlSession ble.onRefreshRequested = { [weak self] in Task { @MainActor [weak self] in await self?.fetchUsage() } } + guard startsDiscovery else { return } // Start an async task to consume async discoveries stream mdnsTask = Task { [weak self] in guard let self else { return } @@ -255,37 +258,55 @@ final class MeterViewModel: ObservableObject { isFetching = true defer { isFetching = false } - do { - var selectedJSON: String? - var selectedStatusData: Data? - for provider in UsageProviderKind.allCases { + var selectedJSON: String? + var selectedStatusData: Data? + var selectedError: Error? + + for provider in UsageProviderKind.allCases { + do { guard let result = try await fetchUsagePayload(provider) else { continue } persist(json: result.json, statusData: result.statusData, provider: provider) if provider == selectedProvider { selectedJSON = result.json selectedStatusData = result.statusData } + } catch { + if provider == selectedProvider { + selectedError = error + } } + } - guard let json = selectedJSON else { - errorMessage = "Cannot reach \(selectedProvider.title) server" + guard let json = selectedJSON else { + if hasCurrentPayload { + errorMessage = nil + await fetchDaemonStatus() return } - usageJSON = json - lastUpdate = Date() - errorMessage = nil - - parseJSON(json) - if let selectedStatusData { - parseDaemonStatus(selectedStatusData) - } else { + if let selectedError { + errorMessage = selectedError.localizedDescription await fetchDaemonStatus() + } else { + errorMessage = "Cannot reach \(selectedProvider.title) server" } - sendToBLE() - } catch { - errorMessage = error.localizedDescription + return + } + + usageJSON = json + lastUpdate = Date() + errorMessage = nil + + parseJSON(json) + if let selectedStatusData { + parseDaemonStatus(selectedStatusData) + } else { await fetchDaemonStatus() } + sendToBLE() + } + + private var hasCurrentPayload: Bool { + usageJSON.trimmingCharacters(in: .whitespacesAndNewlines) != "{}" || lastUpdate != nil } private func fetchUsagePayload(_ provider: UsageProviderKind) async throws -> (json: String, statusData: Data?)? { @@ -294,7 +315,7 @@ final class MeterViewModel: ObservableObject { var request = URLRequest(url: usageURL) request.cachePolicy = .reloadIgnoringLocalAndRemoteCacheData - let (data, response) = try await URLSession.shared.data(for: request) + let (data, response) = try await urlSession.data(for: request) guard let httpResponse = response as? HTTPURLResponse, httpResponse.statusCode == 200 else { return nil } @@ -303,7 +324,7 @@ final class MeterViewModel: ObservableObject { if let statusURL = endpointURL("status", base: base) { var statusRequest = URLRequest(url: statusURL) statusRequest.cachePolicy = .reloadIgnoringLocalAndRemoteCacheData - let (data, response) = try await URLSession.shared.data(for: statusRequest) + let (data, response) = try await urlSession.data(for: statusRequest) statusData = (response as? HTTPURLResponse)?.statusCode == 200 ? data : nil } else { statusData = nil @@ -359,7 +380,7 @@ final class MeterViewModel: ObservableObject { do { var request = URLRequest(url: requestURL) request.cachePolicy = .reloadIgnoringLocalAndRemoteCacheData - let (data, response) = try await URLSession.shared.data(for: request) + let (data, response) = try await urlSession.data(for: request) guard let httpResponse = response as? HTTPURLResponse, httpResponse.statusCode == 200 else { return } parseDaemonStatus(data) } catch { diff --git a/ios/CodexMeterApp/CodexMeterAppTests/MeterDiscoveryTests.swift b/ios/CodexMeterApp/CodexMeterAppTests/MeterDiscoveryTests.swift index 4f77ea0..a2509cd 100644 --- a/ios/CodexMeterApp/CodexMeterAppTests/MeterDiscoveryTests.swift +++ b/ios/CodexMeterApp/CodexMeterAppTests/MeterDiscoveryTests.swift @@ -24,4 +24,110 @@ final class MeterDiscoveryTests: XCTestCase { XCTAssertTrue(received) } + + func testFetchUsageKeepsSelectedProviderWhenOtherProviderFails() async throws { + let configuration = URLSessionConfiguration.ephemeral + configuration.protocolClasses = [MeterMockURLProtocol.self] + let session = URLSession(configuration: configuration) + + MeterMockURLProtocol.handler = { request in + let url = try XCTUnwrap(request.url) + if url.host == "codex.local", url.path == "/usage" { + return ( + HTTPURLResponse(url: url, statusCode: 200, httpVersion: nil, headerFields: nil)!, + #"{"s":45,"sr":120,"w":28,"wr":7200,"st":"$2.31 today","ok":true}"#.data(using: .utf8)! + ) + } + if url.host == "codex.local", url.path == "/status" { + return ( + HTTPURLResponse(url: url, statusCode: 200, httpVersion: nil, headerFields: nil)!, + #"{"source":"codex_oauth","last_success_at":"now","payload_age_seconds":1,"uptime_seconds":60}"#.data(using: .utf8)! + ) + } + throw URLError(.cannotConnectToHost) + } + + let vm: MeterViewModel = await MainActor.run { + let vm = MeterViewModel(urlSession: session, startsDiscovery: false) + vm.codexServerURL = "http://codex.local" + vm.claudeServerURL = "http://claude.local" + vm.selectedProvider = .codex + return vm + } + + await vm.fetchUsage() + + await MainActor.run { + XCTAssertNil(vm.errorMessage) + XCTAssertEqual(vm.statusText, "$2.31 today") + XCTAssertEqual(vm.sessionPct, 45) + } + } + + func testFetchUsageKeepsCurrentPayloadWhenSelectedRefreshIsCancelled() async throws { + let configuration = URLSessionConfiguration.ephemeral + configuration.protocolClasses = [MeterMockURLProtocol.self] + let session = URLSession(configuration: configuration) + + MeterMockURLProtocol.handler = { request in + let url = try XCTUnwrap(request.url) + if url.host == "codex.local", url.path == "/status" { + return ( + HTTPURLResponse(url: url, statusCode: 200, httpVersion: nil, headerFields: nil)!, + #"{"source":"codex_oauth","last_success_at":"now","payload_age_seconds":29,"uptime_seconds":60,"last_error":""}"#.data(using: .utf8)! + ) + } + throw URLError(.cancelled) + } + + let vm: MeterViewModel = await MainActor.run { + let vm = MeterViewModel(urlSession: session, startsDiscovery: false) + vm.codexServerURL = "http://codex.local" + vm.selectedProvider = .codex + vm.usageJSON = #"{"s":64,"sr":289,"w":94,"wr":8895,"st":"0 credits","ok":true}"# + vm.lastUpdate = Date() + vm.statusText = "0 credits" + vm.sessionPct = 64 + return vm + } + + await vm.fetchUsage() + + await MainActor.run { + XCTAssertNil(vm.errorMessage) + XCTAssertEqual(vm.statusText, "0 credits") + XCTAssertEqual(vm.sessionPct, 64) + XCTAssertEqual(vm.daemonLastError, "") + } + } +} + +private final class MeterMockURLProtocol: URLProtocol { + static var handler: ((URLRequest) throws -> (HTTPURLResponse, Data))? + + override class func canInit(with request: URLRequest) -> Bool { + true + } + + override class func canonicalRequest(for request: URLRequest) -> URLRequest { + request + } + + override func startLoading() { + guard let handler = Self.handler else { + client?.urlProtocol(self, didFailWithError: URLError(.badServerResponse)) + return + } + + do { + let (response, data) = try handler(request) + client?.urlProtocol(self, didReceive: response, cacheStoragePolicy: .notAllowed) + client?.urlProtocol(self, didLoad: data) + client?.urlProtocolDidFinishLoading(self) + } catch { + client?.urlProtocol(self, didFailWithError: error) + } + } + + override func stopLoading() {} }