Skip to content
65 changes: 35 additions & 30 deletions Sources/Fluid/Services/LocalAPI/InferenceAPIController.swift
Original file line number Diff line number Diff line change
Expand Up @@ -25,15 +25,16 @@ final class InferenceAPIController: LocalAPIRouteHandler {
let model: String
}

func handle(_ request: LocalAPI.Request) async -> LocalAPI.Response {
guard request.method == "POST" else {
return LocalAPI.error("Method not allowed.", status: 405)
}
private struct AudioUpload {
let data: Data
let suggestedExtension: String
}

switch request.path {
case "/v1/transcribe":
func handle(_ request: LocalAPI.Request) async -> LocalAPI.Response {
switch (request.method, request.path) {
case ("POST", "/v1/transcribe"):
return await self.transcribe(request)
case "/v1/postprocess":
case ("POST", "/v1/postprocess"):
return await self.postprocess(request)
default:
return LocalAPI.error("Route not found.", status: 404)
Expand All @@ -46,27 +47,25 @@ final class InferenceAPIController: LocalAPIRouteHandler {
return try await self.transcribeFile(fileURL)
}

let temporaryFileURL = try await self.decodeUploadedAudioFile(from: request)
do {
let response = try await self.transcribeFile(temporaryFileURL)
await LocalAPIAudioDecoder.removeTemporaryFile(at: temporaryFileURL)
return response
} catch {
await LocalAPIAudioDecoder.removeTemporaryFile(at: temporaryFileURL)
throw error
let upload = try self.decodeUploadedAudio(from: request)
return try await LocalAPIAudioDecoder.withTemporaryAudioFile(
fromAudioData: upload.data,
suggestedExtension: upload.suggestedExtension
) { fileURL in
try await self.transcribeFile(fileURL)
}
} catch {
return LocalAPI.error(error.localizedDescription, status: 400)
}
}

private func transcribeFile(_ fileURL: URL) async throws -> LocalAPI.Response {
let apiResult = try await AppServices.shared.asr.transcribeFileForAPI(fileURL)
let payload = try await LocalAPITranscriptionService.transcribe(fileURL)
return LocalAPI.json(
TranscribeResponse(
text: apiResult.result.text,
confidence: apiResult.result.confidence,
sampleCount: apiResult.sampleCount,
text: payload.text,
confidence: payload.confidence,
sampleCount: payload.sampleCount,
provider: SettingsStore.shared.selectedSpeechModel.displayName
)
)
Expand All @@ -78,7 +77,7 @@ final class InferenceAPIController: LocalAPIRouteHandler {
do {
payload = try LocalAPI.decoder.decode(TranscribeJSONRequest.self, from: request.body)
} catch {
throw NSError(domain: "InferenceAPIController", code: -3, userInfo: [NSLocalizedDescriptionKey: "Invalid JSON audio payload."])
throw self.makeError("Invalid JSON audio payload.", code: -3)
}

guard let path = payload.path, !path.isEmpty else { return nil }
Expand All @@ -101,15 +100,16 @@ final class InferenceAPIController: LocalAPIRouteHandler {
}
}

private func decodeUploadedAudioFile(from request: LocalAPI.Request) async throws -> URL {
private func decodeUploadedAudio(from request: LocalAPI.Request) throws -> AudioUpload {
let data: Data
let suggestedExtension: String

if self.isJSON(request) {
let payload: TranscribeJSONRequest
do {
payload = try LocalAPI.decoder.decode(TranscribeJSONRequest.self, from: request.body)
} catch {
throw NSError(domain: "InferenceAPIController", code: -3, userInfo: [NSLocalizedDescriptionKey: "Invalid JSON audio payload."])
throw self.makeError("Invalid JSON audio payload.", code: -3)
}

if let audioBase64 = payload.audioBase64,
Expand All @@ -118,22 +118,19 @@ final class InferenceAPIController: LocalAPIRouteHandler {
data = decodedData
suggestedExtension = payload.filename.flatMap { URL(fileURLWithPath: $0).pathExtension } ?? "wav"
} else {
throw NSError(domain: "InferenceAPIController", code: -1, userInfo: [NSLocalizedDescriptionKey: "Missing audio path or audioBase64."])
throw self.makeError("Missing audio path or audioBase64.", code: -1)
}
} else {
guard !request.body.isEmpty else {
throw NSError(domain: "InferenceAPIController", code: -1, userInfo: [NSLocalizedDescriptionKey: "Missing audio body."])
throw self.makeError("Missing audio body.", code: -1)
}

data = request.body
let filename = request.headers["x-filename"] ?? "audio.wav"
suggestedExtension = URL(fileURLWithPath: filename).pathExtension
}

return try await LocalAPIAudioDecoder.temporaryFile(
fromAudioData: data,
suggestedExtension: suggestedExtension
)
return AudioUpload(data: data, suggestedExtension: suggestedExtension)
}

private func decodeText(from request: LocalAPI.Request) throws -> String {
Expand All @@ -142,18 +139,26 @@ final class InferenceAPIController: LocalAPIRouteHandler {
do {
payload = try LocalAPI.decoder.decode(TextRequest.self, from: request.body)
} catch {
throw NSError(domain: "InferenceAPIController", code: -4, userInfo: [NSLocalizedDescriptionKey: "Invalid JSON text payload."])
throw self.makeError("Invalid JSON text payload.", code: -4)
}
return payload.text
}

guard let text = String(data: request.body, encoding: .utf8) else {
throw NSError(domain: "InferenceAPIController", code: -2, userInfo: [NSLocalizedDescriptionKey: "Text body must be UTF-8."])
throw self.makeError("Text body must be UTF-8.", code: -2)
}
return text
}

private func isJSON(_ request: LocalAPI.Request) -> Bool {
request.headers["content-type"]?.lowercased().contains("application/json") == true
}

private func makeError(_ message: String, code: Int) -> NSError {
NSError(
domain: "InferenceAPIController",
code: code,
userInfo: [NSLocalizedDescriptionKey: message]
)
}
}
70 changes: 62 additions & 8 deletions Sources/Fluid/Services/LocalAPI/LocalAPIAudioDecoder.swift
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,13 @@ enum LocalAPIAudioDecoder {
fileURL: URL,
chunkDurationSeconds: Double = LocalAPIAudioDecoder.maxChunkDurationSeconds
) throws {
let audioFile = try AVAudioFile(forReading: fileURL)
let audioFile: AVAudioFile
do {
audioFile = try AVAudioFile(forReading: fileURL)
} catch {
throw LocalAPIAudioDecoder.audioDecodeError(underlying: error)
}

let sourceSampleRate = audioFile.processingFormat.sampleRate
guard sourceSampleRate > 0, chunkDurationSeconds > 0 else {
throw NSError(
Expand Down Expand Up @@ -46,11 +52,19 @@ enum LocalAPIAudioDecoder {
)
}

try self.audioFile.read(into: sourceBuffer, frameCount: framesToRead)
return try AudioBufferConverter.monoSamples(
from: sourceBuffer,
targetSampleRate: LocalAPIAudioDecoder.sampleRate
)
do {
try self.audioFile.read(into: sourceBuffer, frameCount: framesToRead)
return try AudioBufferConverter.monoSamples(
from: sourceBuffer,
targetSampleRate: LocalAPIAudioDecoder.sampleRate
)
} catch {
let nsError = error as NSError
if nsError.domain == "LocalAPIAudioDecoder" {
throw error
}
throw LocalAPIAudioDecoder.audioDecodeError(underlying: error)
}
}
}

Expand All @@ -66,19 +80,59 @@ enum LocalAPIAudioDecoder {
}.value
}

static func withTemporaryAudioFile<T>(
fromAudioData data: Data,
suggestedExtension: String,
operation: (URL) async throws -> T
) async throws -> T {
let fileURL = try await self.temporaryFile(
fromAudioData: data,
suggestedExtension: suggestedExtension
)
do {
let result = try await operation(fileURL)
await self.removeTemporaryFile(at: fileURL)
return result
} catch {
await self.removeTemporaryFile(at: fileURL)
throw error
}
}

static func removeTemporaryFile(at fileURL: URL) async {
await Task.detached(priority: .utility) {
try? FileManager.default.removeItem(at: fileURL)
}.value
}

static func estimatedSampleCount(for fileURL: URL) throws -> Int {
let file = try AVAudioFile(forReading: fileURL)
let file: AVAudioFile
do {
file = try AVAudioFile(forReading: fileURL)
} catch {
throw self.audioDecodeError(underlying: error)
}

let sourceFormat = file.processingFormat
guard sourceFormat.sampleRate > 0 else {
throw NSError(domain: "LocalAPIAudioDecoder", code: -6, userInfo: [NSLocalizedDescriptionKey: "Audio file has an invalid sample rate."])
throw NSError(
domain: "LocalAPIAudioDecoder",
code: -6,
userInfo: [NSLocalizedDescriptionKey: "Audio file has an invalid sample rate."]
)
}

return Int((Double(file.length) * self.sampleRate / sourceFormat.sampleRate).rounded())
}

private static func audioDecodeError(underlying error: Error) -> NSError {
NSError(
domain: "LocalAPIAudioDecoder",
code: -7,
userInfo: [
NSLocalizedDescriptionKey: "The uploaded file could not be decoded as audio.",
NSUnderlyingErrorKey: error,
]
)
}
}
Loading
Loading