-
Notifications
You must be signed in to change notification settings - Fork 707
[AI] support on-device structured output #8395
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from 18 commits
a0ffcd5
64ef2ef
35e128a
60b3c12
8d670e8
565be05
b12b96d
e4fc45a
f98bd5c
e3dbf21
7e1f92b
cda79ef
1944371
2981a63
1c8e412
ecd0d71
bc0b6ff
a665954
6ea68a1
ca714f7
6c7e01b
69a8c59
396f724
6794859
17953d2
ef07aea
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -25,12 +25,18 @@ import com.google.devtools.ksp.symbol.KSClassDeclaration | |
| import com.google.devtools.ksp.symbol.KSType | ||
| import com.google.devtools.ksp.symbol.KSVisitorVoid | ||
| import com.google.devtools.ksp.symbol.Modifier | ||
| import com.squareup.kotlinpoet.AnnotationSpec | ||
| import com.squareup.kotlinpoet.ClassName | ||
| import com.squareup.kotlinpoet.CodeBlock | ||
| import com.squareup.kotlinpoet.FileSpec | ||
| import com.squareup.kotlinpoet.FunSpec | ||
| import com.squareup.kotlinpoet.KModifier | ||
| import com.squareup.kotlinpoet.ParameterSpec | ||
| import com.squareup.kotlinpoet.ParameterizedTypeName | ||
| import com.squareup.kotlinpoet.ParameterizedTypeName.Companion.parameterizedBy | ||
| import com.squareup.kotlinpoet.PropertySpec | ||
| import com.squareup.kotlinpoet.TypeName | ||
| import com.squareup.kotlinpoet.TypeSpec | ||
| import com.squareup.kotlinpoet.ksp.toClassName | ||
| import com.squareup.kotlinpoet.ksp.toTypeName | ||
| import com.squareup.kotlinpoet.ksp.writeTo | ||
|
|
@@ -57,30 +63,26 @@ internal class SchemaSymbolProcessorVisitor( | |
| codeGenerator, | ||
| Dependencies(true, containingFile), | ||
| ) | ||
| val companionFile = generateMlKitCompanionFileSpec(classDeclaration) | ||
| companionFile.writeTo( | ||
| codeGenerator, | ||
| Dependencies(true, containingFile), | ||
| ) | ||
| } | ||
|
|
||
| fun generateFileSpec(classDeclaration: KSClassDeclaration): FileSpec { | ||
| val className = classDeclaration.toClassName() | ||
| val simpleNamesJoined = className.simpleNames.joinToString("_") | ||
| return FileSpec.builder( | ||
| classDeclaration.packageName.asString(), | ||
| "${classDeclaration.simpleName.asString()}GeneratedSchema", | ||
| className.packageName, | ||
| "${simpleNamesJoined}GeneratedSchema", | ||
| ) | ||
| .addImport("com.google.firebase.ai.type", "JsonSchema") | ||
| .addFunction( | ||
| FunSpec.builder("firebaseAISchema") | ||
| .receiver( | ||
| ClassName( | ||
| classDeclaration.packageName.asString(), | ||
| classDeclaration.simpleName.asString() + ".Companion" | ||
| ) | ||
| ) | ||
| .receiver(className.nestedClass("Companion")) | ||
| .returns( | ||
| ClassName("com.google.firebase.ai.type", "JsonSchema") | ||
| .parameterizedBy( | ||
| ClassName( | ||
| classDeclaration.packageName.asString(), | ||
| classDeclaration.simpleName.asString() | ||
| ) | ||
| ) | ||
| ClassName("com.google.firebase.ai.type", "JsonSchema").parameterizedBy(className) | ||
| ) | ||
| .addAnnotation(Generated::class) | ||
| .addCode( | ||
|
|
@@ -142,7 +144,17 @@ internal class SchemaSymbolProcessorVisitor( | |
| builder.addStatement("JsonSchema.double(").indent() | ||
| } | ||
| "kotlin.String" -> { | ||
| builder.addStatement("JsonSchema.string(").indent() | ||
| if (!guideValues.enumValues.isNullOrEmpty()) { | ||
| val enumElements = guideValues.enumValues.joinToString { "%S" } | ||
| builder | ||
| .addStatement("JsonSchema.enumeration(") | ||
| .indent() | ||
| .add("values = listOf(") | ||
| .addStatement(enumElements, *guideValues.enumValues.toTypedArray()) | ||
| .addStatement("),") | ||
| } else { | ||
| builder.addStatement("JsonSchema.string(").indent() | ||
| } | ||
| } | ||
| "kotlin.collections.List" -> { | ||
|
|
||
|
|
@@ -155,7 +167,12 @@ internal class SchemaSymbolProcessorVisitor( | |
| throw RuntimeException() | ||
| } | ||
| val listParamCodeBlock = | ||
| generateCodeBlockForSchema(type = listTypeParam.resolve(), parentType = type) | ||
| generateCodeBlockForSchema( | ||
| type = listTypeParam.resolve(), | ||
| parentType = type, | ||
| guideAnnotation = | ||
| if (!guideValues.enumValues.isNullOrEmpty()) guideAnnotation else null, | ||
| ) | ||
| builder | ||
| .addStatement("JsonSchema.array(") | ||
| .indent() | ||
|
|
@@ -171,14 +188,13 @@ internal class SchemaSymbolProcessorVisitor( | |
| .filterIsInstance(KSClassDeclaration::class.java) | ||
| .map { it.simpleName.asString() } | ||
| .toList() | ||
| val enumElements = enumValues.joinToString { "%S" } | ||
| builder | ||
| .addStatement("JsonSchema.enumeration(") | ||
| .indent() | ||
| .addStatement("clazz = ${qualifiedName.asString()}::class,") | ||
| .addStatement("values = listOf(") | ||
| .indent() | ||
| .addStatement(enumValues.joinToString { "\"$it\"" }) | ||
| .unindent() | ||
| .add("values = listOf(") | ||
| .addStatement(enumElements, *enumValues.toTypedArray()) | ||
| .addStatement("),") | ||
| } else { | ||
| builder | ||
|
|
@@ -248,4 +264,125 @@ internal class SchemaSymbolProcessorVisitor( | |
| builder.addStatement("nullable = %L)", className.isNullable).unindent() | ||
| return builder.build() | ||
| } | ||
|
|
||
| private fun isGenerableClass(type: KSType): Boolean { | ||
| return type.declaration.annotations.any { it.shortName.getShortName() == "Generable" } | ||
| } | ||
|
|
||
| private fun isListOfGenerableClass(type: KSType): Boolean { | ||
| val qualifiedName = type.declaration.qualifiedName?.asString() | ||
| if (qualifiedName == "kotlin.collections.List" || qualifiedName == "java.util.List") { | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Is it prudent to analyze hierarchy here? I can imagine weird issues if you wanted to use an immutable list or ArrayList or otherwise.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. We can't safely rely on full hierarchy analysis here because we need to know exactly how to reconstruct the collection when generating the toSdk() map the MLkit response back to developer code. If a developer provided a custom list implementation and we just allowed it through, our KSP processor would generate invalid mapping code that fails to compile on their end. However, I completely agree that we should support the common variants. I've just updated the KSP processor to explicitly support MutableList and ArrayList properties in @generable classes, along with the correct construction logic to convert them safely in the generated toSdk() method. |
||
| val argType = type.arguments.firstOrNull()?.type?.resolve() | ||
| if (argType != null) { | ||
| return isGenerableClass(argType) | ||
| } | ||
| } | ||
| return false | ||
| } | ||
|
|
||
| private fun mapToMlKitCompanionType(type: KSType): TypeName { | ||
| if (isListOfGenerableClass(type)) { | ||
| type.arguments.firstOrNull()?.type?.resolve()?.let { argType -> | ||
| val ksClass = argType.declaration as KSClassDeclaration | ||
| val argClassName = | ||
| ClassName( | ||
| ksClass.packageName.asString(), | ||
| "${ksClass.toClassName().simpleNames.joinToString("_")}_MlKitCompanion", | ||
| ) | ||
| return ClassName("kotlin.collections", "List").parameterizedBy(argClassName) | ||
| } | ||
| } else if (isGenerableClass(type)) { | ||
| val ksClass = type.declaration as KSClassDeclaration | ||
| return ClassName( | ||
| ksClass.packageName.asString(), | ||
| "${ksClass.toClassName().simpleNames.joinToString("_")}_MlKitCompanion" | ||
| ) | ||
| } | ||
| return type.toTypeName() | ||
| } | ||
|
|
||
| fun generateMlKitCompanionFileSpec(classDeclaration: KSClassDeclaration): FileSpec { | ||
| val packageName = classDeclaration.packageName.asString() | ||
| val simpleNamesJoined = classDeclaration.toClassName().simpleNames.joinToString("_") | ||
| val companionClassName = "${simpleNamesJoined}_MlKitCompanion" | ||
| val fileBuilder = | ||
| FileSpec.builder(packageName, companionClassName).addAnnotation(Generated::class) | ||
|
|
||
| val classBuilder = TypeSpec.classBuilder(companionClassName).addModifiers(KModifier.DATA) | ||
| val keepAnnotation = AnnotationSpec.builder(ClassName("androidx.annotation", "Keep")).build() | ||
| classBuilder.addAnnotation(keepAnnotation) | ||
|
|
||
| val generableAnn = | ||
| classDeclaration.annotations.firstOrNull { it.shortName.getShortName() == "Generable" } | ||
| val classDesc = getStringFromAnnotation(generableAnn, "description") | ||
|
milaGGL marked this conversation as resolved.
|
||
| val mlkitGenerableBuilder = | ||
| AnnotationSpec.builder(ClassName("com.google.mlkit.genai.schema.annotations", "Generable")) | ||
| if (!classDesc.isNullOrEmpty()) { | ||
| mlkitGenerableBuilder.addMember("description = %S", classDesc) | ||
| } | ||
| classBuilder.addAnnotation(mlkitGenerableBuilder.build()) | ||
|
|
||
| val primaryConstructor = FunSpec.constructorBuilder() | ||
| val toSdkBuilder = | ||
| FunSpec.builder("toSdk").addAnnotation(keepAnnotation).returns(classDeclaration.toClassName()) | ||
| val toSdkArgs = mutableListOf<String>() | ||
|
|
||
| classDeclaration.getAllProperties().forEach { property -> | ||
| val propName = property.simpleName.asString() | ||
| val propType = property.type.resolve() | ||
| val typeName = mapToMlKitCompanionType(propType) | ||
|
|
||
| val paramBuilder = ParameterSpec.builder(propName, typeName) | ||
| val propBuilder = PropertySpec.builder(propName, typeName).initializer(propName) | ||
|
|
||
| val guideAnn = property.annotations.firstOrNull { it.shortName.getShortName() == "Guide" } | ||
| if (guideAnn != null) { | ||
| val guideValues = | ||
| getGuideValuesFromAnnotation(guideAnn, getStringFromAnnotation(guideAnn, "description")) | ||
|
milaGGL marked this conversation as resolved.
|
||
| val mlkitGuideBuilder = | ||
| AnnotationSpec.builder(ClassName("com.google.mlkit.genai.schema.annotations", "Guide")) | ||
| if (!guideValues.description.isNullOrEmpty()) | ||
| mlkitGuideBuilder.addMember("description = %S", guideValues.description) | ||
| if (guideValues.minimum != null) | ||
| mlkitGuideBuilder.addMember("minimum = %L", guideValues.minimum) | ||
| if (guideValues.maximum != null) | ||
| mlkitGuideBuilder.addMember("maximum = %L", guideValues.maximum) | ||
| if (guideValues.minItems != null) | ||
| mlkitGuideBuilder.addMember("minItems = %L", guideValues.minItems) | ||
| if (guideValues.maxItems != null) | ||
| mlkitGuideBuilder.addMember("maxItems = %L", guideValues.maxItems) | ||
| if (!guideValues.format.isNullOrEmpty()) | ||
| mlkitGuideBuilder.addMember("format = %S", guideValues.format) | ||
| if (!guideValues.enumValues.isNullOrEmpty()) { | ||
| val enumElements = guideValues.enumValues.joinToString { "%S" } | ||
| mlkitGuideBuilder.addMember( | ||
| "enumValues = arrayOf($enumElements)", | ||
| *guideValues.enumValues.toTypedArray() | ||
| ) | ||
| } | ||
| paramBuilder.addAnnotation(mlkitGuideBuilder.build()) | ||
| } | ||
|
|
||
| primaryConstructor.addParameter(paramBuilder.build()) | ||
| classBuilder.addProperty(propBuilder.build()) | ||
|
|
||
| if (isListOfGenerableClass(propType)) { | ||
| toSdkArgs.add("$propName = this.$propName.map { it.toSdk() }") | ||
| } else if (isGenerableClass(propType)) { | ||
| toSdkArgs.add("$propName = this.$propName.toSdk()") | ||
| } else { | ||
| toSdkArgs.add("$propName = this.$propName") | ||
| } | ||
| } | ||
|
|
||
| classBuilder.primaryConstructor(primaryConstructor.build()) | ||
| toSdkBuilder.addStatement( | ||
| "return %T(\n ${toSdkArgs.joinToString(",\n ")}\n)", | ||
| classDeclaration.toClassName() | ||
| ) | ||
| classBuilder.addFunction(toSdkBuilder.build()) | ||
|
|
||
| fileBuilder.addType(classBuilder.build()) | ||
| return fileBuilder.build() | ||
| } | ||
| } | ||
Uh oh!
There was an error while loading. Please reload this page.