9f308fadd2
Rotate AI idle suggestions with optional remote packs, move history/dictionary onto self-sizing Home preview cards, harden clipboard capture/prompting, and simplify keyboard chrome by dropping most liquid-glass shadows.
354 lines
12 KiB
Swift
354 lines
12 KiB
Swift
// OpenAIRealtimeASRClient.swift
|
|
// OSGKeyboard · HostSupport
|
|
//
|
|
// OpenAI Realtime transcription (WebSocket). Streams PCM and transcript
|
|
// deltas for utterance-level ASR. Batch `/audio/transcriptions` remains the
|
|
// fallback path when realtime is unavailable.
|
|
|
|
import Foundation
|
|
import os
|
|
#if canImport(OSGKeyboardShared)
|
|
import OSGKeyboardShared
|
|
#endif
|
|
|
|
struct OpenAIRealtimeASRClient: CloudASRTranscribing, CloudASRStreamingCapable {
|
|
let apiKey: String
|
|
let endpoint: String
|
|
let model: String
|
|
let session: URLSession
|
|
/// Used when streaming fails and Flow falls back to chunked batch ASR.
|
|
private let batchClient: PromptCloudASRClient
|
|
|
|
static let appendChunkBytes = 4_800 // 100 ms @ 24 kHz / 16-bit mono.
|
|
static let finalTimeout: TimeInterval = 15
|
|
|
|
init(
|
|
apiKey: String,
|
|
endpoint: String,
|
|
model: String,
|
|
batchBaseURL: String,
|
|
session: URLSession
|
|
) {
|
|
self.apiKey = apiKey
|
|
self.endpoint = endpoint
|
|
self.model = model
|
|
self.session = session
|
|
self.batchClient = PromptCloudASRClient(
|
|
providerId: "openai",
|
|
baseURL: batchBaseURL.isEmpty ? "https://api.openai.com/v1" : batchBaseURL,
|
|
apiKey: apiKey,
|
|
model: Self.batchModel(from: model),
|
|
session: session
|
|
)
|
|
}
|
|
|
|
func prepare(dictionary: PersonalDictionary) async throws {}
|
|
|
|
func openStreamingSession(
|
|
locale: Locale,
|
|
dictionary: PersonalDictionary,
|
|
onPartial: @escaping @Sendable (String) -> Void
|
|
) async throws -> any CloudASRStreamingSession {
|
|
_ = dictionary
|
|
guard !apiKey.isEmpty else { throw CloudASRError.noAPIKey }
|
|
let url = try resolvedEndpointURL()
|
|
var request = URLRequest(url: url)
|
|
request.timeoutInterval = 8
|
|
request.setValue(
|
|
"Bearer \(apiKey.trimmingCharacters(in: .whitespacesAndNewlines))",
|
|
forHTTPHeaderField: "Authorization"
|
|
)
|
|
|
|
let wsTask = session.webSocketTask(with: request)
|
|
wsTask.resume()
|
|
let live = OpenAIRealtimeStreamingSession(
|
|
wsTask: wsTask,
|
|
model: resolvedRealtimeModel,
|
|
locale: locale,
|
|
onPartial: onPartial
|
|
)
|
|
try await live.start()
|
|
return live
|
|
}
|
|
|
|
func transcribe(
|
|
samples: [Float],
|
|
sampleRate: Int,
|
|
locale: Locale,
|
|
dictionary: PersonalDictionary
|
|
) async throws -> String {
|
|
try await batchClient.transcribe(
|
|
samples: samples,
|
|
sampleRate: sampleRate,
|
|
locale: locale,
|
|
dictionary: dictionary
|
|
)
|
|
}
|
|
|
|
func probeConnection() async throws {
|
|
do {
|
|
let session = try await openStreamingSession(
|
|
locale: Locale(identifier: "zh-CN"),
|
|
dictionary: .empty,
|
|
onPartial: { _ in }
|
|
)
|
|
session.cancel()
|
|
} catch {
|
|
guard Self.shouldFallbackToBatch(afterProbeError: error) else {
|
|
throw CancellationError()
|
|
}
|
|
try await batchClient.probeConnection()
|
|
}
|
|
}
|
|
|
|
static func shouldFallbackToBatch(afterProbeError error: Error) -> Bool {
|
|
!ProviderToolCancellation.matches(error)
|
|
}
|
|
|
|
private var resolvedRealtimeModel: String {
|
|
let trimmed = model.trimmingCharacters(in: .whitespacesAndNewlines)
|
|
if trimmed.isEmpty || trimmed.hasPrefix("gpt-4o") || trimmed == "whisper-1" {
|
|
return CloudASRModelCatalog.openAIRealtimeWhisper
|
|
}
|
|
return trimmed
|
|
}
|
|
|
|
private func resolvedEndpointURL() throws -> URL {
|
|
let raw = endpoint.trimmingCharacters(in: .whitespacesAndNewlines)
|
|
if raw.hasPrefix("wss://") || raw.hasPrefix("ws://") {
|
|
guard let url = URL(string: raw) else { throw CloudASRError.invalidURL }
|
|
return url
|
|
}
|
|
guard let url = URL(string: CloudASRModelCatalog.openAIRealtimeEndpoint) else {
|
|
throw CloudASRError.invalidURL
|
|
}
|
|
return url
|
|
}
|
|
|
|
private static func batchModel(from model: String) -> String {
|
|
let trimmed = model.trimmingCharacters(in: .whitespacesAndNewlines)
|
|
if trimmed.isEmpty || trimmed.contains("realtime") {
|
|
return CloudASRModelCatalog.openAITranscribe
|
|
}
|
|
return trimmed
|
|
}
|
|
}
|
|
|
|
// MARK: - Utterance session
|
|
|
|
private final class OpenAIRealtimeStreamingSession: CloudASRStreamingSession, @unchecked Sendable {
|
|
private let wsTask: URLSessionWebSocketTask
|
|
private let model: String
|
|
private let locale: Locale
|
|
private let onPartial: @Sendable (String) -> Void
|
|
private let lock = OSAllocatedUnfairLock()
|
|
private var receiveTask: Task<Void, Never>?
|
|
private var failure: Error?
|
|
private var finished = false
|
|
private var pcmBuffer = Data()
|
|
private var reducer = OpenAIRealtimeTranscriptReducer()
|
|
private var awaitingCommit = false
|
|
|
|
init(
|
|
wsTask: URLSessionWebSocketTask,
|
|
model: String,
|
|
locale: Locale,
|
|
onPartial: @escaping @Sendable (String) -> Void
|
|
) {
|
|
self.wsTask = wsTask
|
|
self.model = model
|
|
self.locale = locale
|
|
self.onPartial = onPartial
|
|
}
|
|
|
|
func start() async throws {
|
|
receiveTask = Task { [weak self] in
|
|
await self?.receiveLoop()
|
|
}
|
|
let language = OpenAIRealtimeTranscriptReducer.languageHint(from: locale)
|
|
var transcription: [String: Any] = [
|
|
"model": model,
|
|
"delay": "low",
|
|
]
|
|
if let language {
|
|
transcription["language"] = language
|
|
}
|
|
var input: [String: Any] = [
|
|
"format": [
|
|
"type": "audio/pcm",
|
|
"rate": 24_000,
|
|
],
|
|
"transcription": transcription,
|
|
]
|
|
input["turn_detection"] = NSNull()
|
|
let update: [String: Any] = [
|
|
"type": "session.update",
|
|
"session": [
|
|
"type": "transcription",
|
|
"audio": [
|
|
"input": input,
|
|
],
|
|
],
|
|
]
|
|
try await sendJSON(update)
|
|
let deadline = Date().addingTimeInterval(8)
|
|
while Date() < deadline {
|
|
try throwIfFailed()
|
|
if lock.withLock({ reducer.sessionReady }) { return }
|
|
try await Task.sleep(nanoseconds: 20_000_000)
|
|
}
|
|
cancel()
|
|
throw CloudASRError.transport("OpenAI realtime session timed out")
|
|
}
|
|
|
|
func append(samples: [Float]) async throws {
|
|
try throwIfFailed()
|
|
let upsampled = CloudASRStreamingPCM.upsample16kTo24k(samples)
|
|
let pcm = CloudASRStreamingPCM.pcm16LE(samples: upsampled)
|
|
let frames: [Data] = lock.withLock {
|
|
pcmBuffer.append(pcm)
|
|
var frames: [Data] = []
|
|
while pcmBuffer.count >= OpenAIRealtimeASRClient.appendChunkBytes {
|
|
let frame = pcmBuffer.prefix(OpenAIRealtimeASRClient.appendChunkBytes)
|
|
frames.append(Data(frame))
|
|
pcmBuffer.removeFirst(OpenAIRealtimeASRClient.appendChunkBytes)
|
|
}
|
|
return frames
|
|
}
|
|
for frame in frames {
|
|
try await sendAppend(frame)
|
|
}
|
|
}
|
|
|
|
func finish() async throws -> String {
|
|
try throwIfFailed()
|
|
let trailing: Data = lock.withLock {
|
|
let data = pcmBuffer
|
|
pcmBuffer.removeAll(keepingCapacity: false)
|
|
awaitingCommit = true
|
|
return data
|
|
}
|
|
if !trailing.isEmpty {
|
|
try await sendAppend(trailing)
|
|
}
|
|
try await sendJSON(["type": "input_audio_buffer.commit"])
|
|
|
|
let deadline = Date().addingTimeInterval(OpenAIRealtimeASRClient.finalTimeout)
|
|
while Date() < deadline {
|
|
try throwIfFailed()
|
|
let snapshot = lock.withLock {
|
|
(awaitingCommit, reducer.composedFinal(), reducer.composedDisplay())
|
|
}
|
|
if !snapshot.0 {
|
|
let text = snapshot.1.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty
|
|
? snapshot.2
|
|
: snapshot.1
|
|
let trimmed = text.trimmingCharacters(in: .whitespacesAndNewlines)
|
|
cancel()
|
|
if trimmed.isEmpty { throw CloudASRError.emptyTranscript }
|
|
return trimmed
|
|
}
|
|
let settled = lock.withLock {
|
|
!reducer.completedByItem.isEmpty && reducer.partialByItem.isEmpty && !awaitingCommit
|
|
}
|
|
if settled {
|
|
let text = lock.withLock { reducer.composedFinal() }
|
|
.trimmingCharacters(in: .whitespacesAndNewlines)
|
|
cancel()
|
|
if text.isEmpty { throw CloudASRError.emptyTranscript }
|
|
return text
|
|
}
|
|
try await Task.sleep(nanoseconds: 20_000_000)
|
|
}
|
|
let fallback = lock.withLock {
|
|
let final = reducer.composedFinal()
|
|
return final.isEmpty ? reducer.composedDisplay() : final
|
|
}
|
|
.trimmingCharacters(in: .whitespacesAndNewlines)
|
|
cancel()
|
|
if fallback.isEmpty {
|
|
throw CloudASRError.transport("OpenAI realtime final timed out")
|
|
}
|
|
return fallback
|
|
}
|
|
|
|
func cancel() {
|
|
receiveTask?.cancel()
|
|
wsTask.cancel(with: .normalClosure, reason: nil)
|
|
lock.withLock { finished = true }
|
|
}
|
|
|
|
private func receiveLoop() async {
|
|
while !Task.isCancelled {
|
|
let message: URLSessionWebSocketTask.Message
|
|
do {
|
|
message = try await wsTask.receive()
|
|
} catch {
|
|
publishFailure(CloudASRError.transport(error.localizedDescription))
|
|
return
|
|
}
|
|
let text: String
|
|
switch message {
|
|
case .string(let value):
|
|
text = value
|
|
case .data(let data):
|
|
text = String(data: data, encoding: .utf8) ?? ""
|
|
@unknown default:
|
|
continue
|
|
}
|
|
guard !text.isEmpty else { continue }
|
|
|
|
let effect = lock.withLock { () -> OpenAIRealtimeEventEffect in
|
|
let effect = reducer.apply(jsonText: text)
|
|
if case .partial = effect, !reducer.completedByItem.isEmpty {
|
|
awaitingCommit = false
|
|
}
|
|
return effect
|
|
}
|
|
switch effect {
|
|
case .none, .sessionReady:
|
|
continue
|
|
case .partial(let display):
|
|
onPartial(display)
|
|
case .failed(let message):
|
|
publishFailure(CloudASRError.transport(message))
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
private func sendAppend(_ pcm: Data) async throws {
|
|
let audio = pcm.base64EncodedString()
|
|
try await sendJSON([
|
|
"type": "input_audio_buffer.append",
|
|
"audio": audio,
|
|
])
|
|
}
|
|
|
|
private func sendJSON(_ body: [String: Any]) async throws {
|
|
guard JSONSerialization.isValidJSONObject(body),
|
|
let data = try? JSONSerialization.data(withJSONObject: body),
|
|
let string = String(data: data, encoding: .utf8) else {
|
|
throw CloudASRError.decoding("invalid realtime payload")
|
|
}
|
|
do {
|
|
try await wsTask.send(.string(string))
|
|
} catch where ProviderToolCancellation.matches(error) {
|
|
throw CancellationError()
|
|
} catch {
|
|
throw CloudASRError.transport(error.localizedDescription)
|
|
}
|
|
}
|
|
|
|
private func throwIfFailed() throws {
|
|
let (error, done) = lock.withLock { (failure, finished) }
|
|
if let error { throw error }
|
|
if done { throw CloudASRError.transport("OpenAI realtime session cancelled") }
|
|
}
|
|
|
|
private func publishFailure(_ error: Error) {
|
|
lock.withLock { failure = error }
|
|
cancel()
|
|
}
|
|
}
|