feat(account): add managed credits and cloud gateway

Introduce optional Apple account-backed credits with scoped gateway access while preserving local and BYOK paths. Refresh assistant behavior, tests, privacy disclosures, docs, and the website for the 2.0 experience.
This commit is contained in:
Rocky
2026-08-20 11:43:21 +08:00
parent 0f9280bd00
commit 498f407585
301 changed files with 19221 additions and 10891 deletions
+556 -3
View File
@@ -3,9 +3,10 @@
//
// Hermetic HTTP batch Cloud ASR clients (URLProtocol stub; no live network).
import XCTest
@testable import OSGKeyboardShared
import os
@testable import OSGKeyboardHostSupport
@testable import OSGKeyboardShared
import XCTest
final class CloudASRHTTPClientTests: XCTestCase {
@@ -26,7 +27,7 @@ final class CloudASRHTTPClientTests: XCTestCase {
session: session
)
let dict = PersonalDictionary(entries: [
PersonalDictionary.Entry(term: "OSGKeyboard", category: .productName, source: .manual),
PersonalDictionary.Entry(term: "OSGKeyboard", category: .productName, source: .manual)
])
// 0.1 s @ 16 kHz
let samples = [Float](repeating: 0.01, count: 1_600)
@@ -172,6 +173,18 @@ final class CloudASRHTTPClientTests: XCTestCase {
CloudASRModelCatalog.strategy(for: moonshotStore.asrProviderId),
.localFallback
)
config.credentialSource = .managed
config.save(to: defaults)
let managedStore = AppGroupStore(defaults: defaults)
XCTAssertTrue(
CloudASRClientFactory.make(
store: managedStore,
managedGrants: GatewayGrantCoordinator(
baseURL: URL(string: "https://account.example.test")!
)
) is ManagedVolcengineASRClient
)
}
func testVolcengineProbeConnectionRejectsEmptyAPIKey() async {
@@ -190,4 +203,544 @@ final class CloudASRHTTPClientTests: XCTestCase {
XCTFail("unexpected \(error)")
}
}
func testManagedSessionCreatesAuthorizedRequestAndStreamsBoundedPCM() async throws {
let http = FakeManagedASRHTTPClient(responses: [
.init(status: 201, data: managedSessionDescriptor(maxFrameBytes: 4))
])
let socket = FakeManagedASRWebSocket(receives: [
.data(Data(#"{"result":{"utterances":[{"text":"","definite":true}]}}"#.utf8)),
.closed(
code: URLSessionWebSocketTask.CloseCode.normalClosure.rawValue,
reason: "Complete"
)
])
let sockets = FakeManagedASRWebSocketFactory(socket: socket)
let partials = LockedStrings()
let client = managedClient(http: http, sockets: sockets)
let live = try await client.openStreamingSession(
locale: Locale(identifier: "zh-Hans"),
dictionary: .empty,
onPartial: { partials.append($0) }
)
// Three samples become six PCM16LE bytes and must be split 4 + 2.
try await live.append(samples: [0, 0.5, -0.5])
let final = try await live.finish()
XCTAssertEqual(final, "你好")
XCTAssertEqual(partials.values, ["你好"])
let requests = await http.requests
let create = try XCTUnwrap(requests.first)
XCTAssertEqual(create.httpMethod, "POST")
XCTAssertEqual(create.url?.path, "/v1/gateway/asr/sessions")
XCTAssertEqual(create.value(forHTTPHeaderField: "Authorization"), "Bearer grant-token")
XCTAssertEqual(create.value(forHTTPHeaderField: "X-Request-ID"), "request_12345678")
let createBody = try XCTUnwrap(create.httpBody)
let createJSON = try XCTUnwrap(
JSONSerialization.jsonObject(with: createBody) as? [String: Any]
)
XCTAssertEqual(createJSON["format"] as? String, "pcm")
XCTAssertEqual(createJSON["codec"] as? String, "raw")
XCTAssertEqual(createJSON["sampleRate"] as? Int, 16_000)
XCTAssertEqual(createJSON["estimatedDurationMillis"] as? Int, 210_000)
let webSocketRequest = try XCTUnwrap(sockets.request)
XCTAssertEqual(webSocketRequest.url?.scheme, "wss")
XCTAssertEqual(
webSocketRequest.url?.path,
"/v1/gateway/asr/sessions/11111111-2222-3333-4444-555555555555/stream"
)
XCTAssertEqual(
webSocketRequest.value(forHTTPHeaderField: "Authorization"),
"Bearer grant-token"
)
XCTAssertEqual(
webSocketRequest.value(forHTTPHeaderField: "X-Request-ID"),
"request_12345678"
)
let sent = socket.sentMessages
XCTAssertEqual(sent.count, 3)
guard case .data(let first) = sent[0],
case .data(let second) = sent[1],
case .string(let end) = sent[2] else {
return XCTFail("expected two binary frames followed by the end control frame")
}
XCTAssertEqual(first.count, 4)
XCTAssertEqual(second.count, 2)
XCTAssertEqual(end, #"{"type":"end"}"#)
XCTAssertTrue(socket.wasClosed)
}
func testManagedSessionRefreshesGrantOnceAndReusesRequestIDAfterUnauthorized() async throws {
let unauthorized = Data(
#"{"code":"unauthorized","message":"expired","requestId":"request_12345678"}"#.utf8
)
let http = FakeManagedASRHTTPClient(responses: [
.init(status: 401, data: unauthorized),
.init(status: 201, data: managedSessionDescriptor())
])
let grants = RecordingManagedASRGrantProvider(tokens: ["expired-grant", "fresh-grant"])
let socket = FakeManagedASRWebSocket()
let client = managedClient(
http: http,
sockets: FakeManagedASRWebSocketFactory(socket: socket),
grantProvider: grants
)
let live = try await client.openStreamingSession(
locale: Locale(identifier: "zh-Hans"),
dictionary: .empty,
onPartial: { _ in }
)
live.cancel()
let requests = await http.requests
XCTAssertEqual(requests.count, 2)
XCTAssertEqual(
requests.map { $0.value(forHTTPHeaderField: "Authorization") },
["Bearer expired-grant", "Bearer fresh-grant"]
)
XCTAssertEqual(
Set(requests.compactMap { $0.value(forHTTPHeaderField: "X-Request-ID") }),
["request_12345678"]
)
let refreshFlags = await grants.forceRefreshValues
XCTAssertEqual(refreshFlags, [false, true])
}
func testManagedSessionMapsStableCreationFailures() async {
let cases = [
ManagedSessionFailureCase(
status: 402,
code: "insufficient_credits",
expected: .insufficientCredits
),
ManagedSessionFailureCase(
status: 429,
code: "asr_concurrency_limit",
expected: .concurrencyLimit
),
ManagedSessionFailureCase(
status: 503,
code: "provider_unavailable",
expected: .sessionCreationFailed(
status: 503,
code: "provider_unavailable"
)
)
]
for testCase in cases {
let body = Data(
#"{"code":"\#(testCase.code)","message":"ignored","requestId":"r"}"#.utf8
)
let http = FakeManagedASRHTTPClient(responses: [
.init(status: testCase.status, data: body)
])
let client = managedClient(
http: http,
sockets: FakeManagedASRWebSocketFactory(socket: FakeManagedASRWebSocket())
)
do {
_ = try await client.openStreamingSession(
locale: Locale(identifier: "en-US"),
dictionary: .empty,
onPartial: { _ in }
)
XCTFail("expected \(testCase.expected)")
} catch let error as ManagedCloudASRError {
XCTAssertEqual(error, testCase.expected)
} catch {
XCTFail("unexpected \(error)")
}
}
}
func testManagedSessionConnectTimeoutIsStableAndClosesSocket() async throws {
let http = FakeManagedASRHTTPClient(responses: [
.init(status: 201, data: managedSessionDescriptor())
])
let socket = FakeManagedASRWebSocket(pingDelay: .seconds(10))
let client = managedClient(
http: http,
sockets: FakeManagedASRWebSocketFactory(socket: socket),
connectTimeout: 0.01
)
do {
_ = try await client.openStreamingSession(
locale: Locale(identifier: "en-US"),
dictionary: .empty,
onPartial: { _ in }
)
XCTFail("expected connect timeout")
} catch let error as ManagedCloudASRError {
XCTAssertEqual(error, .connectTimeout)
} catch {
XCTFail("unexpected \(error)")
}
XCTAssertTrue(socket.wasClosed)
}
func testManagedSessionIdleTimeoutIsStable() async throws {
let http = FakeManagedASRHTTPClient(responses: [
.init(status: 201, data: managedSessionDescriptor(idleTimeoutMillis: 10))
])
let socket = FakeManagedASRWebSocket(receiveDelay: .seconds(10))
let client = managedClient(
http: http,
sockets: FakeManagedASRWebSocketFactory(socket: socket)
)
let live = try await client.openStreamingSession(
locale: Locale(identifier: "en-US"),
dictionary: .empty,
onPartial: { _ in }
)
try await live.append(samples: [0.1])
do {
_ = try await live.finish()
XCTFail("expected idle timeout")
} catch let error as ManagedCloudASRError {
XCTAssertEqual(error, .idleTimeout)
} catch {
XCTFail("unexpected \(error)")
}
}
func testManagedSessionMidstreamCancellationClosesAndThrowsCancellation() async throws {
let http = FakeManagedASRHTTPClient(responses: [
.init(status: 201, data: managedSessionDescriptor())
])
let socket = FakeManagedASRWebSocket(receiveDelay: .seconds(10))
let client = managedClient(
http: http,
sockets: FakeManagedASRWebSocketFactory(socket: socket)
)
let live = try await client.openStreamingSession(
locale: Locale(identifier: "en-US"),
dictionary: .empty,
onPartial: { _ in }
)
live.cancel()
do {
try await live.append(samples: [0.1])
XCTFail("expected cancellation")
} catch is CancellationError {
// Expected: StreamingUtterancePipeline maps this to `.cancelled`.
} catch {
XCTFail("unexpected \(error)")
}
XCTAssertTrue(socket.wasClosed)
}
func testManagedSessionEmptyProviderResultIsStable() async throws {
let http = FakeManagedASRHTTPClient(responses: [
.init(status: 201, data: managedSessionDescriptor())
])
let socket = FakeManagedASRWebSocket(receives: [
.data(Data(#"{"result":{"text":""}}"#.utf8)),
.string(#"{"type":"gateway_error","code":"asr_failed"}"#)
])
let client = managedClient(
http: http,
sockets: FakeManagedASRWebSocketFactory(socket: socket)
)
let live = try await client.openStreamingSession(
locale: Locale(identifier: "en-US"),
dictionary: .empty,
onPartial: { _ in }
)
try await live.append(samples: [0.1])
do {
_ = try await live.finish()
XCTFail("expected empty result")
} catch let error as ManagedCloudASRError {
XCTAssertEqual(error, .emptyResult)
} catch {
XCTFail("unexpected \(error)")
}
}
func testManagedBatchFallbackUsesGrantRequestIDAndParsesNDJSON() async throws {
let payload = Data(#"""
{"result":{"utterances":[{"text":"","definite":false}]}}
{"result":{"utterances":[{"text":"","definite":true}]}}
"""#.utf8)
let http = FakeManagedASRHTTPClient(responses: [.init(status: 200, data: payload)])
let client = managedClient(
http: http,
sockets: FakeManagedASRWebSocketFactory(socket: FakeManagedASRWebSocket())
)
let text = try await client.transcribe(
samples: [Float](repeating: 0.1, count: 1_600),
sampleRate: 16_000,
locale: Locale(identifier: "zh-Hans"),
dictionary: .empty
)
XCTAssertEqual(text, "批量结果")
let requests = await http.requests
let request = try XCTUnwrap(requests.first)
XCTAssertEqual(request.url?.path, "/v1/gateway/asr")
XCTAssertEqual(request.value(forHTTPHeaderField: "Authorization"), "Bearer grant-token")
XCTAssertEqual(request.value(forHTTPHeaderField: "X-Request-ID"), "request_12345678")
XCTAssertEqual(request.value(forHTTPHeaderField: "X-Audio-Duration-Ms"), "100")
XCTAssertEqual(request.httpBody?.count, 3_200)
}
func testManagedBatchRefreshesGrantOnceAndKeepsLogicalRequestID() async throws {
let unauthorized = Data(
#"{"code":"unauthorized","message":"expired","requestId":"request_12345678"}"#.utf8
)
let result = Data(
#"{"result":{"utterances":[{"text":"","definite":true}]}}"#.utf8
)
let http = FakeManagedASRHTTPClient(responses: [
.init(status: 401, data: unauthorized),
.init(status: 200, data: result)
])
let grants = RecordingManagedASRGrantProvider(tokens: ["expired-grant", "fresh-grant"])
let client = managedClient(
http: http,
sockets: FakeManagedASRWebSocketFactory(socket: FakeManagedASRWebSocket()),
grantProvider: grants
)
let text = try await client.transcribe(
samples: [Float](repeating: 0.1, count: 1_600),
sampleRate: 16_000,
locale: Locale(identifier: "zh-Hans"),
dictionary: .empty
)
XCTAssertEqual(text, "刷新后结果")
let requests = await http.requests
XCTAssertEqual(requests.count, 2)
XCTAssertEqual(
requests.map { $0.value(forHTTPHeaderField: "Authorization") },
["Bearer expired-grant", "Bearer fresh-grant"]
)
XCTAssertEqual(
Set(requests.compactMap { $0.value(forHTTPHeaderField: "X-Request-ID") }),
["request_12345678"]
)
XCTAssertEqual(requests[0].httpBody, requests[1].httpBody)
let refreshFlags = await grants.forceRefreshValues
XCTAssertEqual(refreshFlags, [false, true])
}
func testManagedBatchFallbackFailureIsStable() async {
let body = Data(
#"{"code":"provider_unavailable","message":"ignored","requestId":"r"}"#.utf8
)
let http = FakeManagedASRHTTPClient(responses: [.init(status: 503, data: body)])
let client = managedClient(
http: http,
sockets: FakeManagedASRWebSocketFactory(socket: FakeManagedASRWebSocket())
)
do {
_ = try await client.transcribe(
samples: [0.1],
sampleRate: 16_000,
locale: Locale(identifier: "en-US"),
dictionary: .empty
)
XCTFail("expected batch failure")
} catch let error as ManagedCloudASRError {
XCTAssertEqual(
error,
.batchFailed(status: 503, code: "provider_unavailable")
)
} catch {
XCTFail("unexpected \(error)")
}
}
private func managedClient(
http: FakeManagedASRHTTPClient,
sockets: FakeManagedASRWebSocketFactory,
grantProvider: any ManagedASRGrantProviding =
StaticManagedASRGrantProvider(token: "grant-token"),
connectTimeout: TimeInterval = 1
) -> ManagedVolcengineASRClient {
ManagedVolcengineASRClient(
baseURL: URL(string: "https://account.example.test")!,
grantProvider: grantProvider,
httpClient: http,
webSocketFactory: sockets,
connectTimeout: connectTimeout,
requestID: { "request_12345678" }
)
}
private func managedSessionDescriptor(
maxFrameBytes: Int = 64 * 1_024,
idleTimeoutMillis: Int = 5_000
) -> Data {
Data(
#"""
{
"sessionId":"11111111-2222-3333-4444-555555555555",
"websocketPath":"/v1/gateway/asr/sessions/11111111-2222-3333-4444-555555555555/stream",
"maxFrameBytes":\#(maxFrameBytes),
"idleTimeoutMillis":\#(idleTimeoutMillis)
}
"""#.utf8
)
}
}
private struct ManagedSessionFailureCase {
let status: Int
let code: String
let expected: ManagedCloudASRError
}
private actor RecordingManagedASRGrantProvider: ManagedASRGrantProviding {
private var tokens: [String]
private(set) var forceRefreshValues: [Bool] = []
init(tokens: [String]) {
self.tokens = tokens
}
func accessToken(forceRefresh: Bool) async throws -> String {
forceRefreshValues.append(forceRefresh)
guard !tokens.isEmpty else {
throw ManagedGatewayError.missingGrant
}
return tokens.removeFirst()
}
}
private actor FakeManagedASRHTTPClient: ManagedASRHTTPClient {
struct Response: Sendable {
let status: Int
let data: Data
}
private var queued: [Response]
private(set) var requests: [URLRequest] = []
init(responses: [Response]) {
queued = responses
}
func data(for request: URLRequest) async throws -> (Data, HTTPURLResponse) {
requests.append(request)
guard !queued.isEmpty else {
throw URLError(.badServerResponse)
}
let response = queued.removeFirst()
let http = HTTPURLResponse(
url: request.url!,
statusCode: response.status,
httpVersion: "HTTP/1.1",
headerFields: ["Content-Type": "application/json"]
)!
return (response.data, http)
}
}
private final class FakeManagedASRWebSocketFactory:
ManagedASRWebSocketFactory, @unchecked Sendable {
private let socket: FakeManagedASRWebSocket
private let lock = OSAllocatedUnfairLock(initialState: Optional<URLRequest>.none)
var request: URLRequest? {
lock.withLock { $0 }
}
init(socket: FakeManagedASRWebSocket) {
self.socket = socket
}
func makeWebSocket(for request: URLRequest) -> any ManagedASRWebSocket {
lock.withLock { $0 = request }
return socket
}
}
private final class FakeManagedASRWebSocket: ManagedASRWebSocket, @unchecked Sendable {
private struct State {
var receives: [ManagedASRWebSocketMessage]
var sent: [ManagedASRWebSocketMessage] = []
var resumed = false
var closed = false
}
private let lock: OSAllocatedUnfairLock<State>
private let pingDelay: Duration?
private let receiveDelay: Duration?
var sentMessages: [ManagedASRWebSocketMessage] {
lock.withLock { $0.sent }
}
var wasClosed: Bool {
lock.withLock { $0.closed }
}
init(
receives: [ManagedASRWebSocketMessage] = [],
pingDelay: Duration? = nil,
receiveDelay: Duration? = nil
) {
lock = OSAllocatedUnfairLock(initialState: State(receives: receives))
self.pingDelay = pingDelay
self.receiveDelay = receiveDelay
}
func resume() {
lock.withLock { $0.resumed = true }
}
func ping() async throws {
if let pingDelay {
try await Task.sleep(for: pingDelay)
}
}
func send(_ message: ManagedASRWebSocketMessage) async throws {
if lock.withLock({ $0.closed }) {
throw URLError(.cancelled)
}
lock.withLock { $0.sent.append(message) }
}
func receive() async throws -> ManagedASRWebSocketMessage {
if let receiveDelay {
try await Task.sleep(for: receiveDelay)
}
if let next = lock.withLock({ state -> ManagedASRWebSocketMessage? in
guard !state.receives.isEmpty else { return nil }
return state.receives.removeFirst()
}) {
return next
}
try await Task.sleep(for: .seconds(10))
throw URLError(.timedOut)
}
func close() {
lock.withLock { $0.closed = true }
}
}
private final class LockedStrings: @unchecked Sendable {
private let lock = OSAllocatedUnfairLock(initialState: [String]())
var values: [String] {
lock.withLock { $0 }
}
func append(_ value: String) {
lock.withLock { $0.append(value) }
}
}