Files
OSGKeyboard/OSGKeyboardShared/Services/LLMClient.swift
T
Rocky 498f407585 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.
2026-08-20 11:43:21 +08:00

578 lines
20 KiB
Swift

// LLMClient.swift
// OSGKeyboard · Shared
//
// Protocol-based LLM client. Default implementation is the OpenAI-compatible
// chat completion client. Add other impls (Anthropic, Gemini) as needed.
import Foundation
public enum LLMError: Error, LocalizedError, Sendable, Equatable {
case invalidURL
case noAPIKey
case http(status: Int)
case decoding(String)
case transport(String)
case cancelled
case rateLimited
public var errorDescription: String? {
switch self {
case .invalidURL:
return SharedL10n.string("error.llm.invalidURL")
case .noAPIKey:
return SharedL10n.string("error.llm.noAPIKey")
case .http(let status):
return SharedL10n.format("error.llm.http", status)
case .decoding:
return SharedL10n.string("error.llm.decoding")
case .transport:
return SharedL10n.string("error.llm.transport")
case .rateLimited:
return SharedL10n.string("error.llm.rateLimited")
case .cancelled:
return SharedL10n.string("error.llm.cancelled")
}
}
}
enum LLMHTTPDiagnostics {
static func logFailure(
providerId: String,
statusCode: Int,
responseByteCount: Int,
response: HTTPURLResponse
) {
#if DEBUG
let provider = safeToken(providerId) ?? "unknown"
let requestID = [
"x-request-id",
"request-id",
"x-correlation-id",
"cf-ray"
]
.compactMap { response.value(forHTTPHeaderField: $0) }
.compactMap(safeToken)
.first
let requestMetadata = requestID.map { " requestId=\($0)" } ?? ""
print(
"⚠️ LLM HTTP error provider=\(provider) status=\(statusCode) "
+ "responseBytes=\(responseByteCount)\(requestMetadata)"
)
#endif
}
private static func safeToken(_ value: String) -> String? {
let trimmed = value.trimmingCharacters(in: .whitespacesAndNewlines)
guard !trimmed.isEmpty, trimmed.count <= 128 else { return nil }
let allowed = CharacterSet.alphanumerics.union(CharacterSet(charactersIn: "-._:"))
guard trimmed.unicodeScalars.allSatisfy(allowed.contains) else { return nil }
return trimmed
}
}
public struct LLMGenerationOptions: Sendable, Equatable {
public let temperature: Double?
public let topP: Double?
public let maxTokens: Int?
public init(temperature: Double? = 0.1, topP: Double? = 0.9, maxTokens: Int? = nil) {
self.temperature = temperature
self.topP = topP
self.maxTokens = maxTokens
}
public static let polishDefault = LLMGenerationOptions()
public static let funCreative = LLMGenerationOptions(
temperature: 0.65,
topP: 0.9
)
}
public protocol LLMClient: Sendable {
/// Polish `text` with `systemPrompt`. `timeout` overrides the
/// per-request HTTP timeout for this call; when `nil` the client's
/// `requestTimeout` baseline is used. Long transcripts must pass a
/// larger, length-scaled timeout so the HTTP request is not cut off
/// mid-generation (see `PolishingService.effectiveTimeout`).
func polish(_ text: String, systemPrompt: String, timeout: TimeInterval?) async throws -> String
/// Provider clients override this to support per-attempt generation controls.
func polish(
_ text: String,
systemPrompt: String,
timeout: TimeInterval?,
options: LLMGenerationOptions
) async throws -> String
/// Complete an explicit chat transcript. AI question mode uses this path;
/// dictation polish keeps the narrower `polish` API above.
func complete(
messages: [LLMRequest.Message],
timeout: TimeInterval?,
options: LLMGenerationOptions
) async throws -> String
/// Stream visible answer deltas for AI keyboard mode. Default falls back to
/// a single delta from `complete`. Reasoning / tool scaffolding must not be
/// yielded as answer text.
func completeStreaming(
messages: [LLMRequest.Message],
timeout: TimeInterval?,
options: LLMGenerationOptions
) -> AsyncThrowingStream<LLMStreamEvent, Error>
/// Baseline upper bound for a single LLM HTTP round-trip when no
/// per-request `timeout` is supplied.
var requestTimeout: TimeInterval { get }
}
public extension LLMClient {
/// Convenience overload that uses the baseline `requestTimeout`.
func polish(_ text: String, systemPrompt: String) async throws -> String {
try await polish(text, systemPrompt: systemPrompt, timeout: nil)
}
func polish(
_ text: String,
systemPrompt: String,
timeout: TimeInterval?,
options: LLMGenerationOptions
) async throws -> String {
try await polish(text, systemPrompt: systemPrompt, timeout: timeout)
}
/// Compatibility fallback for injected polish-only clients. Production
/// provider clients override this method to preserve all conversation turns.
func complete(
messages: [LLMRequest.Message],
timeout: TimeInterval?,
options: LLMGenerationOptions
) async throws -> String {
let systemPrompt = messages.first(where: { $0.role == "system" })?.content ?? ""
let userText = messages.last(where: { $0.role == "user" })?.content ?? ""
return try await polish(
userText,
systemPrompt: systemPrompt,
timeout: timeout,
options: options
)
}
/// Non-streaming fallback used by test doubles and polish-only clients.
func completeStreaming(
messages: [LLMRequest.Message],
timeout: TimeInterval?,
options: LLMGenerationOptions
) -> AsyncThrowingStream<LLMStreamEvent, Error> {
AsyncThrowingStream { continuation in
let task = Task {
do {
let text = try await complete(
messages: messages,
timeout: timeout,
options: options
)
let trimmed = text.trimmingCharacters(in: .whitespacesAndNewlines)
if !trimmed.isEmpty {
continuation.yield(.delta(trimmed))
}
continuation.finish()
} catch is CancellationError {
continuation.finish(throwing: LLMError.cancelled)
} catch let error as LLMError {
continuation.finish(throwing: error)
} catch {
continuation.finish(throwing: error)
}
}
continuation.onTermination = { _ in task.cancel() }
}
}
}
// MARK: - OpenAI-compatible implementation
public struct OpenAICompatibleClient: LLMClient {
public let baseURL: String
public let apiKey: String
public let model: String
public let providerId: String
public let thinkingEnabled: Bool
public let session: URLSession
/// Canonical request timeout for a single LLM HTTP round-trip. Both
/// the `URLRequest.timeoutInterval` we set below and any external
/// race that wants to bound the total time spent waiting on the LLM
/// (e.g. `PolishingService`) should derive from this constant.
public let requestTimeout: TimeInterval = 15
public init(
baseURL: String,
apiKey: String,
model: String,
providerId: String = "",
thinkingEnabled: Bool = false,
session: URLSession = .shared
) {
self.baseURL = baseURL
self.apiKey = apiKey
self.model = model
self.providerId = providerId
self.thinkingEnabled = thinkingEnabled
self.session = session
}
public func polish(_ text: String, systemPrompt: String, timeout: TimeInterval?) async throws -> String {
try await polish(
text,
systemPrompt: systemPrompt,
timeout: timeout,
options: .polishDefault
)
}
public func polish(
_ text: String,
systemPrompt: String,
timeout: TimeInterval?,
options: LLMGenerationOptions
) async throws -> String {
try await complete(
messages: [
.system(systemPrompt),
.user(text)
],
timeout: timeout,
options: options
)
}
public func complete(
messages: [LLMRequest.Message],
timeout: TimeInterval?,
options: LLMGenerationOptions
) async throws -> String {
let req = try makeChatRequest(
messages: messages,
timeout: timeout,
options: options,
stream: false
)
do {
let (data, response) = try await session.data(for: req)
guard let http = response as? HTTPURLResponse else {
throw LLMError.transport("non-HTTP response")
}
if !(200..<300).contains(http.statusCode) {
LLMHTTPDiagnostics.logFailure(
providerId: providerId,
statusCode: http.statusCode,
responseByteCount: data.count,
response: http
)
if http.statusCode == 429 { throw LLMError.rateLimited }
throw LLMError.http(status: http.statusCode)
}
do {
let decoded = try JSONDecoder().decode(LLMResponse.self, from: data)
LLMCacheMetricsStore.record(
providerId: providerId,
promptTokens: decoded.usage?.promptTokens,
cachedTokens: decoded.usage?.cachedTokens
)
return decoded.content.trimmingCharacters(in: .whitespacesAndNewlines)
} catch {
throw LLMError.decoding(String(describing: error))
}
} catch let err as LLMError {
throw err
} catch is CancellationError {
throw LLMError.cancelled
} catch let urlError as URLError where urlError.code == .cancelled {
throw LLMError.cancelled
} catch {
throw LLMError.transport(String(describing: error))
}
}
public func completeStreaming(
messages: [LLMRequest.Message],
timeout: TimeInterval?,
options: LLMGenerationOptions
) -> AsyncThrowingStream<LLMStreamEvent, Error> {
AsyncThrowingStream { continuation in
let task = Task {
do {
let req = try makeChatRequest(
messages: messages,
timeout: timeout,
options: options,
stream: true
)
for try await event in LLMStreamingSession.mapSSE(
session: session,
request: req,
providerId: providerId,
parse: LLMStreamDeltaParser.chatCompletionsDelta(from:)
) {
continuation.yield(event)
}
continuation.finish()
} catch is CancellationError {
continuation.finish(throwing: LLMError.cancelled)
} catch let error as LLMError {
continuation.finish(throwing: error)
} catch {
continuation.finish(throwing: LLMError.transport(String(describing: error)))
}
}
continuation.onTermination = { _ in task.cancel() }
}
}
private func makeChatRequest(
messages: [LLMRequest.Message],
timeout: TimeInterval?,
options: LLMGenerationOptions,
stream: Bool
) throws -> URLRequest {
guard !apiKey.isEmpty else { throw LLMError.noAPIKey }
let urlString = baseURL.hasSuffix("/")
? "\(baseURL)chat/completions"
: "\(baseURL)/chat/completions"
guard let url = URL(string: urlString) else { throw LLMError.invalidURL }
let omitSampling = LLMThinkingControl.shouldOmitSamplingParameters(
providerId: providerId,
baseURL: baseURL,
model: model,
thinkingEnabled: thinkingEnabled
)
let request = LLMRequest(
model: model,
messages: messages,
temperature: omitSampling ? nil : options.temperature,
maxTokens: options.maxTokens ?? Self.outputTokenLimit(for: messages),
topP: omitSampling ? nil : options.topP
)
var req = URLRequest(url: url)
req.httpMethod = "POST"
req.setValue("application/json", forHTTPHeaderField: "Content-Type")
req.setValue("Bearer \(apiKey)", forHTTPHeaderField: "Authorization")
// Per-request timeout scales with transcript length; fall back to
// the baseline when the caller does not supply one.
req.timeoutInterval = timeout ?? requestTimeout
req.httpBody = try Self.encodedBody(
request,
providerId: providerId,
baseURL: baseURL,
model: model,
thinkingEnabled: thinkingEnabled,
stream: stream
)
return req
}
private static func outputTokenLimit(
for messages: [LLMRequest.Message]
) -> Int {
let combined = messages.map(\.content).joined(separator: "\n")
return LLMRequest.outputTokenLimit(for: combined)
}
private static func encodedBody(
_ request: LLMRequest,
providerId: String,
baseURL: String,
model: String,
thinkingEnabled: Bool,
stream: Bool = false
) throws -> Data {
let encoded = try JSONEncoder().encode(request)
guard var body = try JSONSerialization.jsonObject(with: encoded) as? [String: Any] else {
return encoded
}
LLMThinkingControl.apply(
to: &body,
providerId: providerId,
baseURL: baseURL,
model: model,
enabled: thinkingEnabled
)
if stream {
body["stream"] = true
}
return try JSONSerialization.data(withJSONObject: body)
}
}
// MARK: - Factory
public enum LLMClientFactory {
/// Build a client from the current `ProviderConfig`.
public static func make(from config: ProviderConfig) -> LLMClient {
make(
providerId: config.providerId,
baseURL: config.baseURL,
apiKey: config.apiKey,
model: config.model,
thinkingEnabled: config.llmThinkingEnabled
)
}
/// Provider-aware factory used by `PolishingService`.
public static func make(
providerId: String,
baseURL: String,
apiKey: String,
model: String,
thinkingEnabled: Bool = false,
session: URLSession = .shared
) -> LLMClient {
switch providerId {
case "anthropic":
return AnthropicMessagesClient(apiKey: apiKey, model: model, session: session)
default:
let resolvedBase = resolvedOpenAICompatibleBaseURL(providerId: providerId, baseURL: baseURL)
return OpenAICompatibleClient(
baseURL: resolvedBase,
apiKey: apiKey,
model: model,
providerId: providerId,
thinkingEnabled: thinkingEnabled,
session: session
)
}
}
/// Gemini exposes an OpenAI-compatible shim under `/v1beta/openai`.
private static func resolvedOpenAICompatibleBaseURL(providerId: String, baseURL: String) -> String {
if !baseURL.isEmpty { return baseURL }
switch providerId {
case "gemini":
return "https://generativelanguage.googleapis.com/v1beta/openai"
default:
return baseURL
}
}
/// Single source of truth for the LLM request timeout, shared by
/// `LLMClient.requestTimeout` implementations and any caller that
/// wants to bound total time spent waiting on the LLM (e.g.
/// `PolishingService`'s safety-net `withThrowingTaskGroup`). Use
/// this instead of hard-coding `15` so all timeouts stay aligned.
public static var defaultRequestTimeout: TimeInterval {
OpenAICompatibleClient(baseURL: "", apiKey: "", model: "").requestTimeout
}
}
// MARK: - Provider-specific thinking controls
//
// Cloud polish defaults to thinking OFF (`llmThinkingEnabled == false`).
// DeepSeek V4 thinking defaults to *enabled* server-side, so we must send an
// explicit `thinking: { type: "disabled" }` — merely omitting the field (or
// sending `reasoning_effort: "low"`, which DeepSeek maps to `high`) leaves
// CoT on and makes polish appear stuck.
enum LLMThinkingControl {
static func shouldOmitSamplingParameters(
providerId: String,
baseURL: String,
model: String,
thinkingEnabled: Bool
) -> Bool {
if thinkingEnabled { return true }
return control(providerId: providerId, baseURL: baseURL, model: model) == .openAIReasoning
}
static func apply(
to body: inout [String: Any],
providerId: String,
baseURL: String,
model: String,
enabled: Bool
) {
switch control(providerId: providerId, baseURL: baseURL, model: model) {
case .deepSeek:
// Official toggle; do not send reasoning_effort when disabled —
// DeepSeek maps low/medium → high while thinking stays on.
body["thinking"] = ["type": enabled ? "enabled" : "disabled"]
if enabled {
body["reasoning_effort"] = "high"
} else {
body.removeValue(forKey: "reasoning_effort")
}
case .miniMax:
body["thinking"] = ["type": enabled ? "adaptive" : "disabled"]
case .gemini:
body["thinking_config"] = [
"thinking_budget": enabled ? -1 : 0
]
case .openAIReasoning:
// o-series / gpt-5: only touch the field when the user opts in,
// or when disabling an always-on reasoner with the lowest effort.
if enabled {
body["reasoning_effort"] = "medium"
} else {
body["reasoning_effort"] = "low"
}
case .none:
return
}
}
private enum Control {
/// DeepSeek / Ark: explicit thinking type toggle.
case deepSeek
case miniMax
case gemini
case openAIReasoning
}
private static func control(
providerId: String,
baseURL: String,
model: String
) -> Control? {
switch providerId {
case "deepseek", "ark":
return .deepSeek
case "minimax":
return .miniMax
case "gemini":
return .gemini
case "openai":
return isOpenAIReasoningModel(model) ? .openAIReasoning : nil
default:
return control(baseURL: baseURL, model: model)
}
}
private static func control(baseURL: String, model: String) -> Control? {
let lower = baseURL.lowercased()
if lower.contains("minimax") || lower.contains("minimaxi") {
return .miniMax
}
if lower.contains("generativelanguage.googleapis.com") {
return .gemini
}
// Hosted DeepSeek (SiliconFlow / OpenRouter / custom proxies).
if lower.contains("deepseek") || model.lowercased().contains("deepseek") {
return .deepSeek
}
return nil
}
private static func isOpenAIReasoningModel(_ model: String) -> Bool {
let lower = model.trimmingCharacters(in: .whitespacesAndNewlines).lowercased()
return lower.hasPrefix("o1")
|| lower.hasPrefix("o3")
|| lower.hasPrefix("o4")
|| lower.hasPrefix("gpt-5")
|| lower.contains("reasoning")
}
}