feat: auto-detect language for AI writing help
Use ML Kit language identification to detect the post's language before transforming. Maps detected BCP-47 tags to the 7 supported languages (English, Japanese, Korean, German, French, Italian, Spanish). Falls back to English for unsupported languages. Rewriter and proofreader clients are cached per language+outputType combination. https://claude.ai/code/session_01RbCYGrbbapRMike8WQy41F
This commit is contained in:
+74
-61
@@ -21,7 +21,9 @@
|
|||||||
package com.vitorpamplona.amethyst.service.ai
|
package com.vitorpamplona.amethyst.service.ai
|
||||||
|
|
||||||
import android.content.Context
|
import android.content.Context
|
||||||
|
import com.google.android.gms.tasks.Tasks
|
||||||
import com.google.mlkit.genai.common.FeatureStatus
|
import com.google.mlkit.genai.common.FeatureStatus
|
||||||
|
import com.google.mlkit.genai.proofreading.Proofreader
|
||||||
import com.google.mlkit.genai.proofreading.ProofreaderOptions
|
import com.google.mlkit.genai.proofreading.ProofreaderOptions
|
||||||
import com.google.mlkit.genai.proofreading.Proofreading
|
import com.google.mlkit.genai.proofreading.Proofreading
|
||||||
import com.google.mlkit.genai.proofreading.ProofreadingRequest
|
import com.google.mlkit.genai.proofreading.ProofreadingRequest
|
||||||
@@ -29,31 +31,44 @@ import com.google.mlkit.genai.rewriting.Rewriter
|
|||||||
import com.google.mlkit.genai.rewriting.RewriterOptions
|
import com.google.mlkit.genai.rewriting.RewriterOptions
|
||||||
import com.google.mlkit.genai.rewriting.Rewriting
|
import com.google.mlkit.genai.rewriting.Rewriting
|
||||||
import com.google.mlkit.genai.rewriting.RewritingRequest
|
import com.google.mlkit.genai.rewriting.RewritingRequest
|
||||||
|
import com.vitorpamplona.amethyst.service.lang.LanguageTranslatorService
|
||||||
import kotlinx.coroutines.Dispatchers
|
import kotlinx.coroutines.Dispatchers
|
||||||
import kotlinx.coroutines.withContext
|
import kotlinx.coroutines.withContext
|
||||||
|
|
||||||
class MLKitWritingAssistant(
|
class MLKitWritingAssistant(
|
||||||
private val context: Context,
|
private val context: Context,
|
||||||
) : WritingAssistant {
|
) : WritingAssistant {
|
||||||
private var rewriters = mutableMapOf<Int, Rewriter>()
|
private var rewriters = mutableMapOf<Long, Rewriter>()
|
||||||
private var proofreader =
|
private var proofreaders = mutableMapOf<Int, Proofreader>()
|
||||||
Proofreading.getClient(
|
|
||||||
ProofreaderOptions
|
private fun rewriterCacheKey(
|
||||||
.builder(context)
|
@RewriterOptions.OutputType outputType: Int,
|
||||||
.setInputType(ProofreaderOptions.InputType.KEYBOARD)
|
@RewriterOptions.Language language: Int,
|
||||||
.setLanguage(ProofreaderOptions.Language.ENGLISH)
|
): Long = (outputType.toLong() shl 32) or language.toLong()
|
||||||
.build(),
|
|
||||||
)
|
|
||||||
|
|
||||||
private fun getRewriter(
|
private fun getRewriter(
|
||||||
@RewriterOptions.OutputType outputType: Int,
|
@RewriterOptions.OutputType outputType: Int,
|
||||||
|
@RewriterOptions.Language language: Int,
|
||||||
): Rewriter =
|
): Rewriter =
|
||||||
rewriters.getOrPut(outputType) {
|
rewriters.getOrPut(rewriterCacheKey(outputType, language)) {
|
||||||
Rewriting.getClient(
|
Rewriting.getClient(
|
||||||
RewriterOptions
|
RewriterOptions
|
||||||
.builder(context)
|
.builder(context)
|
||||||
.setOutputType(outputType)
|
.setOutputType(outputType)
|
||||||
.setLanguage(RewriterOptions.Language.ENGLISH)
|
.setLanguage(language)
|
||||||
|
.build(),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun getProofreader(
|
||||||
|
@ProofreaderOptions.Language language: Int,
|
||||||
|
): Proofreader =
|
||||||
|
proofreaders.getOrPut(language) {
|
||||||
|
Proofreading.getClient(
|
||||||
|
ProofreaderOptions
|
||||||
|
.builder(context)
|
||||||
|
.setInputType(ProofreaderOptions.InputType.KEYBOARD)
|
||||||
|
.setLanguage(language)
|
||||||
.build(),
|
.build(),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -61,19 +76,12 @@ class MLKitWritingAssistant(
|
|||||||
override suspend fun checkAvailability(): WritingAssistantStatus =
|
override suspend fun checkAvailability(): WritingAssistantStatus =
|
||||||
withContext(Dispatchers.IO) {
|
withContext(Dispatchers.IO) {
|
||||||
try {
|
try {
|
||||||
val status = getRewriter(RewriterOptions.OutputType.REPHRASE).checkFeatureStatus().get()
|
val rewriter = getRewriter(RewriterOptions.OutputType.REPHRASE, RewriterOptions.Language.ENGLISH)
|
||||||
|
val status = rewriter.checkFeatureStatus().get()
|
||||||
when (status) {
|
when (status) {
|
||||||
FeatureStatus.AVAILABLE -> {
|
FeatureStatus.AVAILABLE -> WritingAssistantStatus.Available
|
||||||
WritingAssistantStatus.Available
|
FeatureStatus.DOWNLOADING -> WritingAssistantStatus.Downloading
|
||||||
}
|
else -> WritingAssistantStatus.Unavailable
|
||||||
|
|
||||||
FeatureStatus.DOWNLOADING -> {
|
|
||||||
WritingAssistantStatus.Downloading
|
|
||||||
}
|
|
||||||
|
|
||||||
else -> {
|
|
||||||
WritingAssistantStatus.Unavailable
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
} catch (e: Exception) {
|
} catch (e: Exception) {
|
||||||
WritingAssistantStatus.Unavailable
|
WritingAssistantStatus.Unavailable
|
||||||
@@ -84,43 +92,18 @@ class MLKitWritingAssistant(
|
|||||||
text: String,
|
text: String,
|
||||||
tone: WritingTone,
|
tone: WritingTone,
|
||||||
): WritingResult {
|
): WritingResult {
|
||||||
|
val language = detectLanguage(text)
|
||||||
val transformedText =
|
val transformedText =
|
||||||
when (tone) {
|
when (tone) {
|
||||||
WritingTone.CORRECT -> {
|
WritingTone.CORRECT -> proofread(text, language)
|
||||||
proofread(text)
|
WritingTone.REPHRASE -> rewrite(text, RewriterOptions.OutputType.REPHRASE, language)
|
||||||
}
|
WritingTone.SHORTER -> rewrite(text, RewriterOptions.OutputType.SHORTEN, language)
|
||||||
|
WritingTone.ELABORATE -> rewrite(text, RewriterOptions.OutputType.ELABORATE, language)
|
||||||
WritingTone.REPHRASE -> {
|
WritingTone.FRIENDLY -> rewrite(text, RewriterOptions.OutputType.FRIENDLY, language)
|
||||||
rewrite(text, RewriterOptions.OutputType.REPHRASE)
|
WritingTone.PROFESSIONAL -> rewrite(text, RewriterOptions.OutputType.PROFESSIONAL, language)
|
||||||
}
|
WritingTone.EMOJIFY -> rewrite(text, RewriterOptions.OutputType.EMOJIFY, language)
|
||||||
|
WritingTone.MORE_DIRECT -> rewrite(text, RewriterOptions.OutputType.PROFESSIONAL, language)
|
||||||
WritingTone.SHORTER -> {
|
WritingTone.PUNCHY -> rewrite(text, RewriterOptions.OutputType.SHORTEN, language)
|
||||||
rewrite(text, RewriterOptions.OutputType.SHORTEN)
|
|
||||||
}
|
|
||||||
|
|
||||||
WritingTone.ELABORATE -> {
|
|
||||||
rewrite(text, RewriterOptions.OutputType.ELABORATE)
|
|
||||||
}
|
|
||||||
|
|
||||||
WritingTone.FRIENDLY -> {
|
|
||||||
rewrite(text, RewriterOptions.OutputType.FRIENDLY)
|
|
||||||
}
|
|
||||||
|
|
||||||
WritingTone.PROFESSIONAL -> {
|
|
||||||
rewrite(text, RewriterOptions.OutputType.PROFESSIONAL)
|
|
||||||
}
|
|
||||||
|
|
||||||
WritingTone.EMOJIFY -> {
|
|
||||||
rewrite(text, RewriterOptions.OutputType.EMOJIFY)
|
|
||||||
}
|
|
||||||
|
|
||||||
WritingTone.MORE_DIRECT -> {
|
|
||||||
rewrite(text, RewriterOptions.OutputType.PROFESSIONAL)
|
|
||||||
}
|
|
||||||
|
|
||||||
WritingTone.PUNCHY -> {
|
|
||||||
rewrite(text, RewriterOptions.OutputType.SHORTEN)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return WritingResult(
|
return WritingResult(
|
||||||
@@ -130,19 +113,34 @@ class MLKitWritingAssistant(
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
private suspend fun detectLanguage(text: String): Int =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
try {
|
||||||
|
val langTag = Tasks.await(LanguageTranslatorService.identifyLanguage(text))
|
||||||
|
mapLanguageTag(langTag)
|
||||||
|
} catch (e: Exception) {
|
||||||
|
RewriterOptions.Language.ENGLISH
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
private suspend fun rewrite(
|
private suspend fun rewrite(
|
||||||
text: String,
|
text: String,
|
||||||
@RewriterOptions.OutputType outputType: Int,
|
@RewriterOptions.OutputType outputType: Int,
|
||||||
|
@RewriterOptions.Language language: Int,
|
||||||
): String =
|
): String =
|
||||||
withContext(Dispatchers.IO) {
|
withContext(Dispatchers.IO) {
|
||||||
val rewriter = getRewriter(outputType)
|
val rewriter = getRewriter(outputType, language)
|
||||||
val request = RewritingRequest.builder(text).build()
|
val request = RewritingRequest.builder(text).build()
|
||||||
val result = rewriter.runInference(request).get()
|
val result = rewriter.runInference(request).get()
|
||||||
result.results.firstOrNull()?.text ?: text
|
result.results.firstOrNull()?.text ?: text
|
||||||
}
|
}
|
||||||
|
|
||||||
private suspend fun proofread(text: String): String =
|
private suspend fun proofread(
|
||||||
|
text: String,
|
||||||
|
@ProofreaderOptions.Language language: Int,
|
||||||
|
): String =
|
||||||
withContext(Dispatchers.IO) {
|
withContext(Dispatchers.IO) {
|
||||||
|
val proofreader = getProofreader(language)
|
||||||
val request = ProofreadingRequest.builder(text).build()
|
val request = ProofreadingRequest.builder(text).build()
|
||||||
val result = proofreader.runInference(request).get()
|
val result = proofreader.runInference(request).get()
|
||||||
result.results.firstOrNull()?.text ?: text
|
result.results.firstOrNull()?.text ?: text
|
||||||
@@ -151,6 +149,21 @@ class MLKitWritingAssistant(
|
|||||||
override fun close() {
|
override fun close() {
|
||||||
rewriters.values.forEach { it.close() }
|
rewriters.values.forEach { it.close() }
|
||||||
rewriters.clear()
|
rewriters.clear()
|
||||||
proofreader.close()
|
proofreaders.values.forEach { it.close() }
|
||||||
|
proofreaders.clear()
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
fun mapLanguageTag(tag: String?): Int =
|
||||||
|
when (tag?.lowercase()?.take(2)) {
|
||||||
|
"en" -> RewriterOptions.Language.ENGLISH
|
||||||
|
"ja" -> RewriterOptions.Language.JAPANESE
|
||||||
|
"ko" -> RewriterOptions.Language.KOREAN
|
||||||
|
"de" -> RewriterOptions.Language.GERMAN
|
||||||
|
"fr" -> RewriterOptions.Language.FRENCH
|
||||||
|
"it" -> RewriterOptions.Language.ITALIAN
|
||||||
|
"es" -> RewriterOptions.Language.SPANISH
|
||||||
|
else -> RewriterOptions.Language.ENGLISH
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user