From 49e1af241bb0f981de07465f6c62874bc12eb527 Mon Sep 17 00:00:00 2001 From: idevlab Date: Fri, 26 Jun 2026 23:03:33 +0800 Subject: [PATCH] Add screen context mode routing --- Package.resolved | 8 +- Package.swift | 1 + Sources/App/VoicePipeline+Processing.swift | 10 +- Sources/App/VoicePipeline+Replacement.swift | 11 +- Sources/App/VoicePipeline+ScreenContext.swift | 19 ++- Sources/App/VoicePipeline.swift | 2 +- Sources/Config/AppSettings.swift | 5 +- Sources/Config/ScreenContextMode.swift | 41 ++++++ .../InputSessionCoordinator+Output.swift | 14 +- .../Integration/InputSessionCoordinator.swift | 15 ++- Sources/LLM/LLMEngine.swift | 2 +- Sources/LLM/VLMEngine.swift | 69 ++++++++++ .../Processing/TextProcessingOptions.swift | 2 + .../Processing/TextProcessor+Generation.swift | 67 ++++++++++ Sources/Processing/TextProcessor.swift | 125 ++++++++++++------ Sources/Prompts/PromptBuilder.swift | 10 +- Sources/Prompts/PromptCatalog.swift | 22 +++ .../Resources/en.lproj/Localizable.strings | 6 +- .../zh-Hans.lproj/Localizable.strings | 6 +- Sources/Screen/ScreenContextSnapshot.swift | 9 ++ Sources/Screen/ScreenOCR.swift | 31 +++-- Sources/UI/SettingsView.swift | 7 + .../PromptAndProcessingTests.swift | 49 +++++++ .../ScreenContextModeTests.swift | 46 +++++++ 24 files changed, 483 insertions(+), 94 deletions(-) create mode 100644 Sources/Config/ScreenContextMode.swift create mode 100644 Sources/LLM/VLMEngine.swift create mode 100644 Sources/Processing/TextProcessor+Generation.swift create mode 100644 Sources/Screen/ScreenContextSnapshot.swift create mode 100644 Tests/OpenTypeTests/ScreenContextModeTests.swift diff --git a/Package.resolved b/Package.resolved index b366d0f0..23274af3 100644 --- a/Package.resolved +++ b/Package.resolved @@ -1,5 +1,5 @@ { - "originHash" : "653645748f515b7fd79f1d727ac7edfbcf5e6b77ce92e890b05ce8b027bf1082", + "originHash" : "9530b211b9c06ef80f43abb5c395727661ed97308630b252b72fd176620944e6", "pins" : [ { "identity" : "eventsource", @@ -15,8 +15,8 @@ "kind" : "remoteSourceControl", "location" : "https://github.com/ml-explore/mlx-swift", "state" : { - "revision" : "61b9e011e09a62b489f6bd647958f1555bdf2896", - "version" : "0.31.3" + "revision" : "dc43e62d7055353c7f99fa071a4e71d29dfddc44", + "version" : "0.31.4" } }, { @@ -25,7 +25,7 @@ "location" : "https://github.com/ml-explore/mlx-swift-lm", "state" : { "branch" : "main", - "revision" : "82d9cd619a78714dba9d474dbc1c75c65622e710" + "revision" : "40c2ff061bff5f7c5b55e4ff2067df3b7aa74fd7" } }, { diff --git a/Package.swift b/Package.swift index 891cb064..176f277f 100644 --- a/Package.swift +++ b/Package.swift @@ -24,6 +24,7 @@ let package = Package( .product(name: "Hub", package: "swift-transformers"), .product(name: "Tokenizers", package: "swift-transformers"), .product(name: "MLXLLM", package: "mlx-swift-lm"), + .product(name: "MLXVLM", package: "mlx-swift-lm"), .product(name: "MLXLMCommon", package: "mlx-swift-lm"), ], path: "Sources", diff --git a/Sources/App/VoicePipeline+Processing.swift b/Sources/App/VoicePipeline+Processing.swift index f4369eab..60729f70 100644 --- a/Sources/App/VoicePipeline+Processing.swift +++ b/Sources/App/VoicePipeline+Processing.swift @@ -122,7 +122,7 @@ extension VoicePipeline { let screenContext = await finishScreenContextCapture() let inputContext = InputContext.capture( targetApp: targetApp, - screenContext: screenContext, + screenContext: screenContext.text, outputMode: .processed, inputLanguage: settings.inputLanguage, source: .menuBar @@ -138,7 +138,8 @@ extension VoicePipeline { text: raw, stylePrompt: settings.customStylePrompt, model: settings.llmModel, - screenContext: screenContext, + screenContext: screenContext.text, + screenImage: screenContext.image, memoryContext: memoryContext ) recordFormattingDuration(started, label: "Smart Format") @@ -157,7 +158,7 @@ extension VoicePipeline { let screenContext = await finishScreenContextCapture() let inputContext = InputContext.capture( targetApp: targetApp, - screenContext: screenContext, + screenContext: screenContext.text, outputMode: .command, inputLanguage: settings.inputLanguage, source: .menuBar @@ -172,7 +173,8 @@ extension VoicePipeline { let text = await textProcessor.processCommand( text: raw, model: settings.llmModel, - screenContext: screenContext, + screenContext: screenContext.text, + screenImage: screenContext.image, memoryContext: memoryContext ) recordFormattingDuration(started, label: "Voice Command formatting") diff --git a/Sources/App/VoicePipeline+Replacement.swift b/Sources/App/VoicePipeline+Replacement.swift index 48f16464..2e74bf27 100644 --- a/Sources/App/VoicePipeline+Replacement.swift +++ b/Sources/App/VoicePipeline+Replacement.swift @@ -147,14 +147,14 @@ extension VoicePipeline { replacementID: UUID, raw: String, settings: AppSettings, - ocrTask: Task?, + ocrTask: Task?, ocrStartedAt: CFAbsoluteTime? ) async { let started = CFAbsoluteTimeGetCurrent() - let screenContext = await ocrTask?.value ?? "" + let screenContext = await ocrTask?.value ?? .empty if let ocrStartedAt { let elapsed = CFAbsoluteTimeGetCurrent() - ocrStartedAt - Log.info("[VoicePipeline] OCR stage finished in \(String(format: "%.2f", elapsed))s") + Log.info("[VoicePipeline] screen context stage finished in \(String(format: "%.2f", elapsed))s") } guard !Task.isCancelled else { return } @@ -164,7 +164,7 @@ extension VoicePipeline { appName: currentReplacement.targetAppName, bundleIdentifier: currentReplacement.targetBundleIdentifier, windowTitle: currentReplacement.context?.windowTitle, - screenContext: screenContext, + screenContext: screenContext.text, outputMode: .processed, inputLanguage: settings.inputLanguage, source: .menuBar @@ -179,7 +179,8 @@ extension VoicePipeline { text: raw, stylePrompt: settings.customStylePrompt, model: settings.llmModel, - screenContext: screenContext, + screenContext: screenContext.text, + screenImage: screenContext.image, memoryContext: memoryContext ) let elapsed = CFAbsoluteTimeGetCurrent() - started diff --git a/Sources/App/VoicePipeline+ScreenContext.swift b/Sources/App/VoicePipeline+ScreenContext.swift index d3eb5738..55f9cbe5 100644 --- a/Sources/App/VoicePipeline+ScreenContext.swift +++ b/Sources/App/VoicePipeline+ScreenContext.swift @@ -12,23 +12,22 @@ extension VoicePipeline { return } - guard ScreenOCR.hasScreenCapturePermission else { - Log.info("[VoicePipeline] screen capture permission not granted, skipping OCR") - cancelScreenContextCapture() - return - } - screenOCRStartedAt = CFAbsoluteTimeGetCurrent() + let mode = ScreenContextMode.effectiveCaptureMode( + preference: appState.settings.screenContextMode, + useRemoteLLM: appState.settings.useRemoteLLM, + modelID: appState.settings.llmModel + ) screenOCRTask = Task.detached(priority: .utility) { - await ScreenOCR.captureAndRecognize() + await ScreenOCR.capture(mode: mode) } } - func finishScreenContextCapture() async -> String { - let context = await screenOCRTask?.value ?? "" + func finishScreenContextCapture() async -> ScreenContextSnapshot { + let context = await screenOCRTask?.value ?? .empty if let screenOCRStartedAt { let elapsed = CFAbsoluteTimeGetCurrent() - screenOCRStartedAt - Log.info("[VoicePipeline] OCR stage finished in \(String(format: "%.2f", elapsed))s") + Log.info("[VoicePipeline] screen context stage finished in \(String(format: "%.2f", elapsed))s") } screenOCRTask = nil screenOCRStartedAt = nil diff --git a/Sources/App/VoicePipeline.swift b/Sources/App/VoicePipeline.swift index 9511a877..5424c1fc 100644 --- a/Sources/App/VoicePipeline.swift +++ b/Sources/App/VoicePipeline.swift @@ -14,7 +14,7 @@ final class VoicePipeline { var volcSpeechEngine: VolcSpeechEngine? var qwenSpeechEngine: LocalASREngine? var mimoSpeechEngine: LocalASREngine? - var screenOCRTask: Task? + var screenOCRTask: Task? var screenOCRStartedAt: CFAbsoluteTime? var processingTask: Task? var replacementTask: Task? diff --git a/Sources/Config/AppSettings.swift b/Sources/Config/AppSettings.swift index 4777e7a9..54b409b5 100644 --- a/Sources/Config/AppSettings.swift +++ b/Sources/Config/AppSettings.swift @@ -250,6 +250,7 @@ final class AppSettings: ObservableObject { @Published var enableStreamingRecognitionBeta: Bool @Published var inputLanguage: InputLanguage @Published var useScreenContext: Bool + @Published var screenContextMode: ScreenContextMode @Published var enableInstantInsert: Bool @Published var hasCompletedOnboarding: Bool @Published var uiLanguage: UILanguage { @@ -290,7 +291,7 @@ final class AppSettings: ObservableObject { case hotkeyType, activationMode, tapInterval, speechEngine, whisperModel, llmModel case microphoneID, outputMode, languageStyle, customStylePrompt, playSounds case enableStreamingRecognitionBeta - case inputLanguage, useScreenContext, enableInstantInsert, hasCompletedOnboarding, uiLanguage, historyRetention + case inputLanguage, useScreenContext, screenContextMode, enableInstantInsert, hasCompletedOnboarding, uiLanguage, historyRetention case enableMemory, memoryWindowMinutes case useCustomSystemPrompt, customSystemPrompt case useRemoteLLM, remoteProvider, remoteAPIKey, remoteBaseURL, remoteModel @@ -338,6 +339,7 @@ final class AppSettings: ObservableObject { enableStreamingRecognitionBeta = ud.object(forKey: Key.enableStreamingRecognitionBeta.rawValue) as? Bool ?? true inputLanguage = InputLanguage(rawValue: ud.string(forKey: Key.inputLanguage.rawValue) ?? "") ?? .chinese useScreenContext = ud.object(forKey: Key.useScreenContext.rawValue) as? Bool ?? false + screenContextMode = ScreenContextMode(rawValue: ud.string(forKey: Key.screenContextMode.rawValue) ?? "") ?? .ocr enableInstantInsert = ud.object(forKey: Key.enableInstantInsert.rawValue) as? Bool ?? false hasCompletedOnboarding = ud.bool(forKey: Key.hasCompletedOnboarding.rawValue) uiLanguage = loadedUILanguage @@ -397,6 +399,7 @@ final class AppSettings: ObservableObject { $enableStreamingRecognitionBeta.dropFirst().sink { [defaults] in defaults.set($0, forKey: Key.enableStreamingRecognitionBeta.rawValue) }.store(in: &cancellables) $inputLanguage.dropFirst().sink { [defaults] in defaults.set($0.rawValue, forKey: Key.inputLanguage.rawValue) }.store(in: &cancellables) $useScreenContext.dropFirst().sink { [defaults] in defaults.set($0, forKey: Key.useScreenContext.rawValue) }.store(in: &cancellables) + $screenContextMode.dropFirst().sink { [defaults] in defaults.set($0.rawValue, forKey: Key.screenContextMode.rawValue) }.store(in: &cancellables) $enableInstantInsert.dropFirst().sink { [defaults] in defaults.set($0, forKey: Key.enableInstantInsert.rawValue) }.store(in: &cancellables) $hasCompletedOnboarding.dropFirst().sink { [defaults] in defaults.set($0, forKey: Key.hasCompletedOnboarding.rawValue) }.store(in: &cancellables) $uiLanguage.dropFirst().sink { [defaults] in defaults.set($0.rawValue, forKey: Key.uiLanguage.rawValue) }.store(in: &cancellables) diff --git a/Sources/Config/ScreenContextMode.swift b/Sources/Config/ScreenContextMode.swift new file mode 100644 index 00000000..e83704a3 --- /dev/null +++ b/Sources/Config/ScreenContextMode.swift @@ -0,0 +1,41 @@ +import Foundation + +enum ScreenContextMode: String, Codable, CaseIterable { + case ocr = "ocr" + case multimodal = "multimodal" + + static func effectiveCaptureMode( + preference: ScreenContextMode, + useRemoteLLM: Bool, + modelID: String + ) -> ScreenContextMode { + guard preference == .multimodal, + !useRemoteLLM, + supportsScreenImageContext(modelID: modelID) + else { + return .ocr + } + return .multimodal + } + + static func supportsScreenImageContext(modelID: String) -> Bool { + let id = modelID.lowercased() + return id.contains("gemma-4") + || id.contains("gemma4") + || id.contains("gemma_4") + || id.contains("-vl") + || id.contains("_vl") + || id.contains("paligemma") + || id.contains("smolvlm") + || id.contains("fastvlm") + || id.contains("pixtral") + || id.contains("idefics") + } + + var label: String { + switch self { + case .ocr: return L("screen_context_mode.ocr") + case .multimodal: return L("screen_context_mode.multimodal") + } + } +} diff --git a/Sources/Integration/InputSessionCoordinator+Output.swift b/Sources/Integration/InputSessionCoordinator+Output.swift index c3eaccc5..e9a01bac 100644 --- a/Sources/Integration/InputSessionCoordinator+Output.swift +++ b/Sources/Integration/InputSessionCoordinator+Output.swift @@ -14,7 +14,7 @@ extension InputSessionCoordinator { text = textProcessor.basicClean(text: raw) case .processed: let screenContext = await screenContext(from: active) - context = inputContext(for: active, screenContext: screenContext, mode: .processed) + context = inputContext(for: active, screenContext: screenContext.text, mode: .processed) let memoryContext = VoicePipelinePolicy.memoryContext( for: .processed, settings: settings, @@ -23,12 +23,13 @@ extension InputSessionCoordinator { text = await textProcessor.process( text: raw, options: options, - screenContext: screenContext, + screenContext: screenContext.text, + screenImage: screenContext.image, memoryContext: memoryContext ) case .command: let screenContext = await screenContext(from: active) - context = inputContext(for: active, screenContext: screenContext, mode: .command) + context = inputContext(for: active, screenContext: screenContext.text, mode: .command) let memoryContext = VoicePipelinePolicy.memoryContext( for: .command, settings: settings, @@ -37,7 +38,8 @@ extension InputSessionCoordinator { text = await textProcessor.processCommand( text: raw, options: options, - screenContext: screenContext, + screenContext: screenContext.text, + screenImage: screenContext.image, memoryContext: memoryContext ) } @@ -66,7 +68,7 @@ extension InputSessionCoordinator { ) } - private func screenContext(from active: ActiveSession) async -> String { - await active.screenContextTask?.value ?? "" + private func screenContext(from active: ActiveSession) async -> ScreenContextSnapshot { + await active.screenContextTask?.value ?? .empty } } diff --git a/Sources/Integration/InputSessionCoordinator.swift b/Sources/Integration/InputSessionCoordinator.swift index d5ac12e9..a7a45c0f 100644 --- a/Sources/Integration/InputSessionCoordinator.swift +++ b/Sources/Integration/InputSessionCoordinator.swift @@ -12,7 +12,7 @@ final class InputSessionCoordinator { let inputLanguage: InputLanguage let useScreenContext: Bool let streamingEnabled: Bool - let screenContextTask: Task? + let screenContextTask: Task? let client: IntegrationClient? } @@ -213,19 +213,20 @@ final class InputSessionCoordinator { func startScreenContextCaptureIfNeeded( mode: OutputMode, useScreenContext: Bool - ) -> Task? { + ) -> Task? { guard VoicePipelinePolicy.shouldCaptureScreenContext( outputMode: mode, useScreenContext: useScreenContext ) else { return nil } - guard ScreenOCR.hasScreenCapturePermission else { - Log.info("[InputSessionCoordinator] screen capture permission not granted, skipping OCR") - return nil - } + let contextMode = ScreenContextMode.effectiveCaptureMode( + preference: settings.screenContextMode, + useRemoteLLM: settings.useRemoteLLM, + modelID: settings.llmModel + ) return Task.detached(priority: .utility) { - await ScreenOCR.captureAndRecognize() + await ScreenOCR.capture(mode: contextMode) } } diff --git a/Sources/LLM/LLMEngine.swift b/Sources/LLM/LLMEngine.swift index d1e403c8..cfa8dd39 100644 --- a/Sources/LLM/LLMEngine.swift +++ b/Sources/LLM/LLMEngine.swift @@ -136,7 +136,7 @@ actor LLMEngine { return "/no_think\n\(prompt)" } - private static func modelConfiguration(for id: String) -> ModelConfiguration { + static func modelConfiguration(for id: String) -> ModelConfiguration { let extraEOSTokens: Set = id.lowercased().contains("gemma-4") ? [""] : [] return ModelConfiguration(id: id, extraEOSTokens: extraEOSTokens) } diff --git a/Sources/LLM/VLMEngine.swift b/Sources/LLM/VLMEngine.swift new file mode 100644 index 00000000..04f10548 --- /dev/null +++ b/Sources/LLM/VLMEngine.swift @@ -0,0 +1,69 @@ +import CoreGraphics +import CoreImage +import Foundation +import MLXLMCommon +import MLXVLM + +actor VLMEngine { + private var container: ModelContainer? + private var currentModelID: String? + + func loadModel(id: String) async throws { + if currentModelID == id, container != nil { return } + + Log.info("[VLMEngine] loading model: \(id)") + let started = CFAbsoluteTimeGetCurrent() + + if let localURL = ModelStorage.localLLMURL(id) { + container = try await VLMModelFactory.shared.loadContainer( + from: localURL, + using: MLXModelLoading.tokenizerLoader + ) + } else { + let config = LLMEngine.modelConfiguration(for: id) + container = try await VLMModelFactory.shared.loadContainer( + from: MLXModelLoading.downloader, + using: MLXModelLoading.tokenizerLoader, + configuration: config + ) + } + + currentModelID = id + let elapsed = CFAbsoluteTimeGetCurrent() - started + Log.info("[VLMEngine] model loaded in \(String(format: "%.1f", elapsed))s") + } + + func generate( + prompt: String, + systemPrompt: String? = nil, + image: CGImage, + maxTokens: Int = 2048, + temperature: Double = 0.3 + ) async throws -> String { + guard let container else { + throw LLMError.modelNotLoaded + } + + let started = CFAbsoluteTimeGetCurrent() + let params = GenerateParameters(maxTokens: maxTokens, temperature: Float(temperature)) + let session = ChatSession( + container, + instructions: systemPrompt, + generateParameters: params, + processing: .init() + ) + let ciImage = CIImage(cgImage: image) + let result = try await session.respond(to: prompt, image: .ciImage(ciImage)) + + let elapsed = CFAbsoluteTimeGetCurrent() - started + Log.info("[VLMEngine] generated \(result.count) chars in \(String(format: "%.1f", elapsed))s") + return result + } + + var isLoaded: Bool { container != nil } + + func unload() { + container = nil + currentModelID = nil + } +} diff --git a/Sources/Processing/TextProcessingOptions.swift b/Sources/Processing/TextProcessingOptions.swift index 8da1add8..258853cf 100644 --- a/Sources/Processing/TextProcessingOptions.swift +++ b/Sources/Processing/TextProcessingOptions.swift @@ -10,6 +10,7 @@ struct TextProcessingOptions { var remoteAPIKey: String var remoteModel: String var remoteProvider: RemoteProvider + var screenContextMode: ScreenContextMode init(settings: AppSettings, inputLanguage: InputLanguage? = nil) { self.inputLanguage = inputLanguage ?? settings.inputLanguage @@ -21,5 +22,6 @@ struct TextProcessingOptions { self.remoteAPIKey = settings.remoteAPIKey self.remoteModel = settings.remoteModel self.remoteProvider = settings.remoteProvider + self.screenContextMode = settings.screenContextMode } } diff --git a/Sources/Processing/TextProcessor+Generation.swift b/Sources/Processing/TextProcessor+Generation.swift new file mode 100644 index 00000000..a6d00542 --- /dev/null +++ b/Sources/Processing/TextProcessor+Generation.swift @@ -0,0 +1,67 @@ +import CoreGraphics +import Foundation + +extension TextProcessor { + func generateText( + prompt: String, + systemPrompt: String, + options: TextProcessingOptions, + maxTokens: Int, + temperature: Double + ) async throws -> String { + if options.useRemoteLLM { + return try await remoteLLMClient.generate( + prompt: prompt, + systemPrompt: systemPrompt, + baseURL: options.remoteBaseURL, + apiKey: options.remoteAPIKey, + model: options.remoteModel, + provider: options.remoteProvider, + maxTokens: maxTokens, + temperature: temperature + ) + } + + await ensureModelLoaded(options.llmModel) + return try await llm.generate( + prompt: prompt, + systemPrompt: systemPrompt, + maxTokens: maxTokens, + temperature: temperature + ) + } + + func generateWithScreenImage( + prompt: String, + systemPrompt: String, + model: String, + image: CGImage, + maxTokens: Int, + temperature: Double + ) async throws -> String { + try await vlm.loadModel(id: model) + return try await vlm.generate( + prompt: prompt, + systemPrompt: systemPrompt, + image: image, + maxTokens: maxTokens, + temperature: temperature + ) + } + + func shouldUseScreenImage(options: TextProcessingOptions, image: CGImage?) -> Bool { + guard image != nil else { return false } + guard options.screenContextMode == .multimodal, !options.useRemoteLLM else { return false } + return ScreenContextMode.supportsScreenImageContext(modelID: options.llmModel) + } + + private func ensureModelLoaded(_ model: String) async { + guard !(await llm.isLoaded) else { return } + do { + try await llm.loadModel(id: model) + } catch { + Log.error("[TextProcessor] on-demand model load failed: \(error.localizedDescription)") + } + } + +} diff --git a/Sources/Processing/TextProcessor.swift b/Sources/Processing/TextProcessor.swift index c58c3182..8a185d56 100644 --- a/Sources/Processing/TextProcessor.swift +++ b/Sources/Processing/TextProcessor.swift @@ -1,8 +1,10 @@ +import CoreGraphics import Foundation final class TextProcessor { - private let llm = LLMEngine() - private let remoteLLMClient = RemoteLLMClient() + let llm = LLMEngine() + let vlm = VLMEngine() + let remoteLLMClient = RemoteLLMClient() private let dictionary = PersonalDictionary.shared struct GenerationOptions { let maxTokens: Int @@ -18,6 +20,7 @@ final class TextProcessor { func unloadLLM() async { await llm.unload() + await vlm.unload() } @discardableResult @@ -54,18 +57,32 @@ final class TextProcessor { return result.trimmingCharacters(in: .whitespacesAndNewlines) } - func process(text: String, stylePrompt: String, model: String, screenContext: String = "", memoryContext: String = "") async -> String { + func process( + text: String, + stylePrompt: String, + model: String, + screenContext: String = "", + screenImage: CGImage? = nil, + memoryContext: String = "" + ) async -> String { let settings = AppSettings.shared var options = TextProcessingOptions(settings: settings) options.customStylePrompt = stylePrompt options.llmModel = model - return await process(text: text, options: options, screenContext: screenContext, memoryContext: memoryContext) + return await process( + text: text, + options: options, + screenContext: screenContext, + screenImage: screenImage, + memoryContext: memoryContext + ) } func process( text: String, options: TextProcessingOptions, screenContext: String = "", + screenImage: CGImage? = nil, memoryContext: String = "" ) async -> String { let preCleanStarted = CFAbsoluteTimeGetCurrent() @@ -77,6 +94,7 @@ final class TextProcessor { style: options.languageStyle, stylePrompt: options.customStylePrompt, screenContext: screenContext, + screenImageAvailable: shouldUseScreenImage(options: options, image: screenImage), memoryContext: memoryContext, inputLanguage: options.inputLanguage ) @@ -93,22 +111,31 @@ final class TextProcessor { do { var result: String let llmStarted = CFAbsoluteTimeGetCurrent() - if options.useRemoteLLM { - result = try await remoteLLMClient.generate( - prompt: userPrompt, - systemPrompt: systemPrompt, - baseURL: options.remoteBaseURL, - apiKey: options.remoteAPIKey, - model: options.remoteModel, - provider: options.remoteProvider, - maxTokens: generationOptions.maxTokens, - temperature: generationOptions.temperature - ) + if let screenImage, shouldUseScreenImage(options: options, image: screenImage) { + do { + result = try await generateWithScreenImage( + prompt: userPrompt, + systemPrompt: systemPrompt, + model: options.llmModel, + image: screenImage, + maxTokens: generationOptions.maxTokens, + temperature: generationOptions.temperature + ) + } catch { + Log.error("[TextProcessor] VLM failed, falling back to text LLM: \(error.localizedDescription)") + result = try await generateText( + prompt: userPrompt, + systemPrompt: systemPrompt, + options: options, + maxTokens: generationOptions.maxTokens, + temperature: generationOptions.temperature + ) + } } else { - await ensureModelLoaded(options.llmModel) - result = try await llm.generate( + result = try await generateText( prompt: userPrompt, systemPrompt: systemPrompt, + options: options, maxTokens: generationOptions.maxTokens, temperature: generationOptions.temperature ) @@ -126,21 +153,35 @@ final class TextProcessor { } /// Command mode: uses voice command system prompt, higher max tokens. - func processCommand(text: String, model: String, screenContext: String, memoryContext: String = "") async -> String { + func processCommand( + text: String, + model: String, + screenContext: String, + screenImage: CGImage? = nil, + memoryContext: String = "" + ) async -> String { let settings = AppSettings.shared var options = TextProcessingOptions(settings: settings) options.llmModel = model - return await processCommand(text: text, options: options, screenContext: screenContext, memoryContext: memoryContext) + return await processCommand( + text: text, + options: options, + screenContext: screenContext, + screenImage: screenImage, + memoryContext: memoryContext + ) } func processCommand( text: String, options: TextProcessingOptions, screenContext: String, + screenImage: CGImage? = nil, memoryContext: String = "" ) async -> String { let systemPrompt = PromptBuilder.buildCommandSystemPrompt( screenContext: screenContext, + screenImageAvailable: shouldUseScreenImage(options: options, image: screenImage), memoryContext: memoryContext, inputLanguage: options.inputLanguage ) @@ -148,22 +189,33 @@ final class TextProcessor { do { var result: String - if options.useRemoteLLM { - result = try await remoteLLMClient.generate( - prompt: userPrompt, - systemPrompt: systemPrompt, - baseURL: options.remoteBaseURL, - apiKey: options.remoteAPIKey, - model: options.remoteModel, - provider: options.remoteProvider, - maxTokens: 4096 - ) + if let screenImage, shouldUseScreenImage(options: options, image: screenImage) { + do { + result = try await generateWithScreenImage( + prompt: userPrompt, + systemPrompt: systemPrompt, + model: options.llmModel, + image: screenImage, + maxTokens: 4096, + temperature: 0.3 + ) + } catch { + Log.error("[TextProcessor] Command VLM failed, falling back to text LLM: \(error.localizedDescription)") + result = try await generateText( + prompt: userPrompt, + systemPrompt: systemPrompt, + options: options, + maxTokens: 4096, + temperature: 0.3 + ) + } } else { - await ensureModelLoaded(options.llmModel) - result = try await llm.generate( + result = try await generateText( prompt: userPrompt, systemPrompt: systemPrompt, - maxTokens: 4096 + options: options, + maxTokens: 4096, + temperature: 0.3 ) } @@ -176,15 +228,6 @@ final class TextProcessor { } } - private func ensureModelLoaded(_ model: String) async { - guard !(await llm.isLoaded) else { return } - do { - try await llm.loadModel(id: model) - } catch { - Log.error("[TextProcessor] on-demand model load failed: \(error.localizedDescription)") - } - } - private static let thinkTagNames = [ "think", "thinking", "thought", "reason", "reasoning", diff --git a/Sources/Prompts/PromptBuilder.swift b/Sources/Prompts/PromptBuilder.swift index db9a825f..a2c37a2d 100644 --- a/Sources/Prompts/PromptBuilder.swift +++ b/Sources/Prompts/PromptBuilder.swift @@ -5,6 +5,7 @@ enum PromptBuilder { style: LanguageStyle, stylePrompt: String, screenContext: String = "", + screenImageAvailable: Bool = false, memoryContext: String = "", inputLanguage: InputLanguage = .chinese ) -> String { @@ -17,6 +18,7 @@ enum PromptBuilder { ) parts.append(contentsOf: PromptCatalog.processingContextSections( screenContext: screenContext, + screenImageAvailable: screenImageAvailable, memoryContext: memoryContext, inputLanguage: inputLanguage )) @@ -28,10 +30,16 @@ enum PromptBuilder { PromptCatalog.userPrompt(text: text, inputLanguage: inputLanguage) } - static func buildCommandSystemPrompt(screenContext: String, memoryContext: String = "", inputLanguage: InputLanguage = .chinese) -> String { + static func buildCommandSystemPrompt( + screenContext: String, + screenImageAvailable: Bool = false, + memoryContext: String = "", + inputLanguage: InputLanguage = .chinese + ) -> String { var parts = [PromptCatalog.commandSystemPrompt(inputLanguage: inputLanguage)] parts.append(contentsOf: PromptCatalog.commandContextSections( screenContext: screenContext, + screenImageAvailable: screenImageAvailable, memoryContext: memoryContext, inputLanguage: inputLanguage )) diff --git a/Sources/Prompts/PromptCatalog.swift b/Sources/Prompts/PromptCatalog.swift index aa58b967..52ba2ea7 100644 --- a/Sources/Prompts/PromptCatalog.swift +++ b/Sources/Prompts/PromptCatalog.swift @@ -19,11 +19,13 @@ enum PromptCatalog { static func processingContextSections( screenContext: String, + screenImageAvailable: Bool, memoryContext: String, inputLanguage: InputLanguage ) -> [String] { compactSections( processingScreenContext(screenContext, inputLanguage: inputLanguage), + processingScreenImageContext(inputLanguage: inputLanguage, isAvailable: screenImageAvailable), processingMemoryContext(memoryContext, inputLanguage: inputLanguage) ) } @@ -60,11 +62,13 @@ enum PromptCatalog { static func commandContextSections( screenContext: String, + screenImageAvailable: Bool, memoryContext: String, inputLanguage: InputLanguage ) -> [String] { compactSections( commandScreenContext(screenContext, inputLanguage: inputLanguage), + commandScreenImageContext(inputLanguage: inputLanguage, isAvailable: screenImageAvailable), commandMemoryContext(memoryContext, inputLanguage: inputLanguage) ) } @@ -94,6 +98,15 @@ private extension PromptCatalog { """ } + static func processingScreenImageContext(inputLanguage: InputLanguage, isAvailable: Bool) -> String? { + guard isAvailable else { return nil } + if inputLanguage == .chinese { + return "屏幕截图已随本次请求提供。请直接观察截图,仅用于纠错、识别专有名词和理解当前上下文,不要把截图内容无关地混入输出。" + } + + return "A screen image is attached to this request. Inspect it directly for corrections, proper nouns, and current context only. Do not copy unrelated screen content into the output." + } + static func processingMemoryContext(_ memoryContext: String, inputLanguage: InputLanguage) -> String? { guard !memoryContext.isEmpty else { return nil } if inputLanguage == .chinese { @@ -126,6 +139,15 @@ private extension PromptCatalog { """ } + static func commandScreenImageContext(inputLanguage: InputLanguage, isAvailable: Bool) -> String? { + guard isAvailable else { return nil } + if inputLanguage == .chinese { + return "用户当前屏幕截图已随本次请求提供。需要回复、总结、翻译或解释屏幕内容时,请直接依据截图。" + } + + return "The user's current screen image is attached. Use it directly when the command asks you to reply, summarize, translate, or explain visible screen content." + } + static func commandMemoryContext(_ memoryContext: String, inputLanguage: InputLanguage) -> String? { guard !memoryContext.isEmpty else { return nil } if inputLanguage == .chinese { diff --git a/Sources/Resources/en.lproj/Localizable.strings b/Sources/Resources/en.lproj/Localizable.strings index bde32ea7..0632892a 100644 --- a/Sources/Resources/en.lproj/Localizable.strings +++ b/Sources/Resources/en.lproj/Localizable.strings @@ -96,7 +96,9 @@ "settings.memory_window" = "Memory window"; "settings.memory_minutes_fmt" = "%d minutes"; "settings.screen_context" = "Screen context assist"; -"settings.screen_context_help" = "Capture on-screen text to help the LLM correct homophones (requires Screen Recording permission, only active in Smart Format mode)"; +"settings.screen_context_help" = "Capture on-screen text to help the LLM correct homophones (requires Screen Recording; toggle applies to Smart Format, Voice Command uses it automatically)"; +"settings.screen_context_mode" = "Screen reading"; +"settings.screen_context_mode_help" = "OCR extracts screen text before formatting. Multimodal sends a screenshot to the local VLM when supported; remote LLM and unsupported local models fall back to OCR."; "settings.voice_language" = "Voice & Language"; "settings.ui_language" = "Interface language"; "settings.app_icon" = "App icon"; @@ -221,6 +223,8 @@ "engine.volc_asr" = "Doubao ASR"; "engine.qwen3_asr" = "Qwen3-ASR (Local)"; "engine.mimo_asr" = "MiMo-V2.5-ASR (Local)"; +"screen_context_mode.ocr" = "OCR"; +"screen_context_mode.multimodal" = "Multimodal"; "style.prompt.concise" = "Minimalist. Keep only core information, remove repetition and filler, break long sentences short."; "style.prompt.formal" = "Formal written style. Use proper wording, reduce colloquial tone, keep the text clear and well-formed."; "style.prompt.professional" = "Professional cleanup. Actively fix typos, homophones, ASR mistakes, and proper nouns, then turn the text into complete natural written sentences. Use numbered lists only when the raw text is clearly a list or action items."; diff --git a/Sources/Resources/zh-Hans.lproj/Localizable.strings b/Sources/Resources/zh-Hans.lproj/Localizable.strings index 6967c72a..a51f1d8f 100644 --- a/Sources/Resources/zh-Hans.lproj/Localizable.strings +++ b/Sources/Resources/zh-Hans.lproj/Localizable.strings @@ -96,7 +96,9 @@ "settings.memory_window" = "记忆时间窗口"; "settings.memory_minutes_fmt" = "%d 分钟"; "settings.screen_context" = "屏幕上下文辅助"; -"settings.screen_context_help" = "录音时截取屏幕文字辅助 LLM 纠正同音字(需屏幕录制权限,仅智能整理模式生效)"; +"settings.screen_context_help" = "录音时截取屏幕文字辅助 LLM 纠正同音字(需屏幕录制权限;开关用于智能整理,语音命令会自动使用)"; +"settings.screen_context_mode" = "读屏幕方式"; +"settings.screen_context_mode_help" = "OCR 会先提取屏幕文字。多模态会在本地模型支持时把截图交给 VLM;远程 LLM 和不支持截图的本地模型会回退到 OCR。"; "settings.voice_language" = "语音与语言"; "settings.ui_language" = "界面语言"; "settings.app_icon" = "应用图标"; @@ -221,6 +223,8 @@ "engine.volc_asr" = "豆包语音识别"; "engine.qwen3_asr" = "Qwen3-ASR(本地)"; "engine.mimo_asr" = "MiMo-V2.5-ASR(本地)"; +"screen_context_mode.ocr" = "OCR"; +"screen_context_mode.multimodal" = "多模态"; "style.prompt.concise" = "极简。只保留核心信息,删掉修饰、重复和过渡,长句拆短。"; "style.prompt.formal" = "正式书面。用规范表达,减少口语感,保持句子完整、逻辑顺畅。"; "style.prompt.professional" = "专业整理。更主动纠正错别字、同音词、识别错误和专有名词,整理成完整自然的书面句子;只有原文明显是步骤或待办时才用编号。"; diff --git a/Sources/Screen/ScreenContextSnapshot.swift b/Sources/Screen/ScreenContextSnapshot.swift new file mode 100644 index 00000000..4ca66973 --- /dev/null +++ b/Sources/Screen/ScreenContextSnapshot.swift @@ -0,0 +1,9 @@ +import CoreGraphics +import Foundation + +struct ScreenContextSnapshot: @unchecked Sendable { + let text: String + let image: CGImage? + + static let empty = ScreenContextSnapshot(text: "", image: nil) +} diff --git a/Sources/Screen/ScreenOCR.swift b/Sources/Screen/ScreenOCR.swift index c55045c4..b5a60fd2 100644 --- a/Sources/Screen/ScreenOCR.swift +++ b/Sources/Screen/ScreenOCR.swift @@ -5,28 +5,37 @@ import ScreenCaptureKit enum ScreenOCR { - /// Captures the main screen and runs OCR, returning extracted text (truncated to `maxLength`). - /// Silently returns empty if screen capture permission has not been granted. - static func captureAndRecognize(maxLength: Int = 2000) async -> String { - guard hasScreenCapturePermission else { return "" } + static func capture(mode: ScreenContextMode, maxLength: Int = 2000) async -> ScreenContextSnapshot { + guard await checkScreenCapturePermission() else { + Log.info("[ScreenOCR] screen capture permission not granted") + return .empty + } guard let image = await captureMainScreen() else { Log.info("[ScreenOCR] screen capture failed") - return "" + return .empty } - let text = await recognizeText(in: image) - Log.info("[ScreenOCR] OCR extracted \(text.count) chars") - return String(text.prefix(maxLength)) + switch mode { + case .ocr: + let text = await recognizeText(in: image) + Log.info("[ScreenOCR] OCR extracted \(text.count) chars") + return ScreenContextSnapshot(text: String(text.prefix(maxLength)), image: nil) + case .multimodal: + Log.info("[ScreenOCR] captured screen image for multimodal context") + return ScreenContextSnapshot(text: "", image: image) + } } - static var hasScreenCapturePermission: Bool { - CGPreflightScreenCaptureAccess() + /// Captures the main screen and runs OCR, returning extracted text (truncated to `maxLength`). + /// Silently returns empty if screen capture permission has not been granted. + static func captureAndRecognize(maxLength: Int = 2000) async -> String { + await capture(mode: .ocr, maxLength: maxLength).text } static func checkScreenCapturePermission() async -> Bool { do { - _ = try await SCShareableContent.excludingDesktopWindows(false, onScreenWindowsOnly: true) + _ = try await SCShareableContent.current return true } catch { return false diff --git a/Sources/UI/SettingsView.swift b/Sources/UI/SettingsView.swift index ce36244a..f97afd86 100644 --- a/Sources/UI/SettingsView.swift +++ b/Sources/UI/SettingsView.swift @@ -86,6 +86,13 @@ struct SettingsView: View { Text(L("settings.screen_context")) } .help(L("settings.screen_context_help")) + Picker(L("settings.screen_context_mode"), selection: $settings.screenContextMode) { + ForEach(ScreenContextMode.allCases, id: \.self) { mode in + Text(mode.label).tag(mode) + } + } + .pickerStyle(.segmented) + .help(L("settings.screen_context_mode_help")) Toggle(isOn: $settings.playSounds) { Text(L("settings.sound_cues")) } diff --git a/Tests/OpenTypeTests/PromptAndProcessingTests.swift b/Tests/OpenTypeTests/PromptAndProcessingTests.swift index 7bbbdea9..e37b7ca8 100644 --- a/Tests/OpenTypeTests/PromptAndProcessingTests.swift +++ b/Tests/OpenTypeTests/PromptAndProcessingTests.swift @@ -106,6 +106,33 @@ final class PromptAndProcessingTests: XCTestCase { } } + func testSystemPromptIncludesScreenImageContextOnlyWhenAvailable() { + withCleanSettings { + let withoutImage = PromptBuilder.buildSystemPrompt( + style: .professional, + stylePrompt: "", + inputLanguage: .chinese + ) + XCTAssertFalse(withoutImage.contains("屏幕截图已随本次请求提供")) + + let chinese = PromptBuilder.buildSystemPrompt( + style: .professional, + stylePrompt: "", + screenImageAvailable: true, + inputLanguage: .chinese + ) + XCTAssertTrue(chinese.contains("屏幕截图已随本次请求提供")) + + let english = PromptBuilder.buildSystemPrompt( + style: .professional, + stylePrompt: "", + screenImageAvailable: true, + inputLanguage: .english + ) + XCTAssertTrue(english.contains("A screen image is attached to this request")) + } + } + func testCasualStylePromptStillRequiresCorrection() { withCleanSettings { let chinese = PromptBuilder.buildSystemPrompt( @@ -169,6 +196,28 @@ final class PromptAndProcessingTests: XCTestCase { XCTAssertFalse(english.contains("Recent input history")) } + func testCommandPromptIncludesScreenImageContextOnlyWhenAvailable() { + let withoutImage = PromptBuilder.buildCommandSystemPrompt( + screenContext: "", + inputLanguage: .chinese + ) + XCTAssertFalse(withoutImage.contains("用户当前屏幕截图已随本次请求提供")) + + let chinese = PromptBuilder.buildCommandSystemPrompt( + screenContext: "", + screenImageAvailable: true, + inputLanguage: .chinese + ) + XCTAssertTrue(chinese.contains("用户当前屏幕截图已随本次请求提供")) + + let english = PromptBuilder.buildCommandSystemPrompt( + screenContext: "", + screenImageAvailable: true, + inputLanguage: .english + ) + XCTAssertTrue(english.contains("current screen image is attached")) + } + func testPersonalDictionaryReplacementsAndRules() { withCleanSettings { let dictionary = PersonalDictionary.shared diff --git a/Tests/OpenTypeTests/ScreenContextModeTests.swift b/Tests/OpenTypeTests/ScreenContextModeTests.swift new file mode 100644 index 00000000..8deeeffa --- /dev/null +++ b/Tests/OpenTypeTests/ScreenContextModeTests.swift @@ -0,0 +1,46 @@ +import Foundation +import XCTest +@testable import OpenType + +final class ScreenContextModeTests: XCTestCase { + func testCasesAreStable() { + XCTAssertEqual(ScreenContextMode.allCases.map(\.rawValue), [ + "ocr", "multimodal", + ]) + } + + func testDefaultsAndPersists() { + let suiteName = "OpenTypeTests.ScreenContextMode.\(UUID().uuidString)" + let defaults = UserDefaults(suiteName: suiteName)! + defaults.removePersistentDomain(forName: suiteName) + defer { defaults.removePersistentDomain(forName: suiteName) } + + XCTAssertEqual(AppSettings(defaults: defaults).screenContextMode, .ocr) + + defaults.set(ScreenContextMode.multimodal.rawValue, forKey: "screenContextMode") + XCTAssertEqual(AppSettings(defaults: defaults).screenContextMode, .multimodal) + } + + func testFallsBackToOCRWhenImageContextIsUnavailable() { + XCTAssertEqual(ScreenContextMode.effectiveCaptureMode( + preference: .multimodal, + useRemoteLLM: false, + modelID: "mlx-community/Qwen3.5-2B-4bit" + ), .ocr) + XCTAssertEqual(ScreenContextMode.effectiveCaptureMode( + preference: .multimodal, + useRemoteLLM: true, + modelID: "mlx-community/gemma-4-e2b-it-4bit" + ), .ocr) + XCTAssertEqual(ScreenContextMode.effectiveCaptureMode( + preference: .multimodal, + useRemoteLLM: false, + modelID: "mlx-community/gemma-4-e2b-it-4bit" + ), .multimodal) + XCTAssertEqual(ScreenContextMode.effectiveCaptureMode( + preference: .multimodal, + useRemoteLLM: false, + modelID: "mlx-community/gemma4_unified" + ), .multimodal) + } +}