Skip to content
Open
Show file tree
Hide file tree
Changes from 12 commits
Commits
Show all changes
26 commits
Select commit Hold shift + click to select a range
a0ffcd5
initial implementation
milaGGL Jul 6, 2026
64ef2ef
remove objectResponse from Candidate
milaGGL Jul 6, 2026
35e128a
update naming, the on-device GenerateObjectResponse, and ObjectSource
milaGGL Jul 8, 2026
60b3c12
update GenerateObjectResponse
milaGGL Jul 9, 2026
8d670e8
update the name of shadowClass property
milaGGL Jul 10, 2026
565be05
update Generable
milaGGL Jul 27, 2026
b12b96d
remove local dependency
milaGGL Jul 27, 2026
e4fc45a
remove debug log
milaGGL Jul 27, 2026
f98bd5c
Merge branch 'main' into mila-support-structured-ouput-locally
milaGGL Jul 27, 2026
e3dbf21
update api txt file
milaGGL Jul 28, 2026
7e1f92b
Fix ConvertersTest by removing brittle ML Kit default assertions
milaGGL Jul 28, 2026
cda79ef
Update SchemaSymbolProcessorVisitor.kt
milaGGL Jul 28, 2026
1944371
Address code review feedback, add dependency on genai-schema
milaGGL Jul 28, 2026
2981a63
Merge branch 'main' into mila-support-structured-ouput-locally
milaGGL Jul 28, 2026
1c8e412
Merge branch 'main' into mila-support-structured-ouput-locally
milaGGL Aug 4, 2026
ecd0d71
resolve some comments
milaGGL Aug 4, 2026
bc0b6ff
test both nested and composite Generable classes
milaGGL Aug 4, 2026
a665954
Merge branch 'main' into mila-support-structured-ouput-locally
milaGGL Aug 4, 2026
6ea68a1
Merge branch 'main' into mila-support-structured-ouput-locally
milaGGL Aug 4, 2026
ca714f7
format changelog
milaGGL Aug 4, 2026
6c7e01b
Merge branch 'main' into mila-support-structured-ouput-locally
milaGGL Aug 7, 2026
69a8c59
resolve comments
milaGGL Aug 7, 2026
396f724
Merge branch 'mila-support-structured-ouput-locally' of https://githu…
milaGGL Aug 7, 2026
6794859
Update OnDeviceGenerativeModelProvider.kt
milaGGL Aug 7, 2026
17953d2
Update SchemaSymbolProcessorVisitor.kt
milaGGL Aug 7, 2026
ef07aea
Merge branch 'main' into mila-support-structured-ouput-locally
milaGGL Aug 7, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -57,6 +63,11 @@ internal class SchemaSymbolProcessorVisitor(
codeGenerator,
Dependencies(true, containingFile),
)
val companionFile = generateMlKitCompanionFileSpec(classDeclaration)
companionFile.writeTo(
codeGenerator,
Dependencies(true, containingFile),
)
}

fun generateFileSpec(classDeclaration: KSClassDeclaration): FileSpec {
Expand Down Expand Up @@ -142,7 +153,18 @@ internal class SchemaSymbolProcessorVisitor(
builder.addStatement("JsonSchema.double(").indent()
}
"kotlin.String" -> {
builder.addStatement("JsonSchema.string(").indent()
if (!guideValues.enumValues.isNullOrEmpty()) {
builder
.addStatement("JsonSchema.enumeration(")
.indent()
.addStatement("values = listOf(")
.indent()
.addStatement(guideValues.enumValues.joinToString { "\"$it\"" })
Comment thread
milaGGL marked this conversation as resolved.
Outdated
.unindent()
.addStatement("),")
} else {
builder.addStatement("JsonSchema.string(").indent()
}
}
"kotlin.collections.List" -> {

Expand All @@ -155,7 +177,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()
Expand Down Expand Up @@ -248,4 +275,127 @@ 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") {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The 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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The 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, packageName: String): TypeName {
if (isListOfGenerableClass(type)) {
val argType = type.arguments.first()!!.type!!.resolve()
val argClassName =
ClassName(
argType.declaration.packageName.asString(),
"${argType.declaration.simpleName.asString()}_MlKitCompanion"
)
return ClassName("kotlin.collections", "List").parameterizedBy(argClassName)
Comment thread
milaGGL marked this conversation as resolved.
Outdated
} else if (isGenerableClass(type)) {
return ClassName(
type.declaration.packageName.asString(),
"${type.declaration.simpleName.asString()}_MlKitCompanion"
)
}
return type.toTypeName()
}

fun generateMlKitCompanionFileSpec(classDeclaration: KSClassDeclaration): FileSpec {
val packageName = classDeclaration.packageName.asString()
val companionClassName = "${classDeclaration.simpleName.asString()}_MlKitCompanion"
Comment thread
milaGGL marked this conversation as resolved.
Outdated
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")
Comment thread
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(ClassName(packageName, classDeclaration.simpleName.asString()))
val toSdkArgs = mutableListOf<String>()

classDeclaration.getAllProperties().forEach { property ->
val propName = property.simpleName.asString()
val propType = property.type.resolve()
val typeName = mapToMlKitCompanionType(propType, packageName)

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"))
Comment thread
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)",
ClassName(packageName, classDeclaration.simpleName.asString())
)
classBuilder.addFunction(toSdkBuilder.build())

fileBuilder.addType(classBuilder.build())
return fileBuilder.build()
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,8 @@ internal data class GuideValues(
val minItems: Int?,
val maxItems: Int?,
val format: String?,
val description: String?
val description: String?,
val enumValues: List<String>? = null,
)

internal fun getGuideValuesFromAnnotation(
Expand All @@ -45,7 +46,8 @@ internal fun getGuideValuesFromAnnotation(
minItems = getIntFromAnnotation(guideAnnotation, "minItems"),
maxItems = getIntFromAnnotation(guideAnnotation, "maxItems"),
format = getStringFromAnnotation(guideAnnotation, "format"),
description = description
description = description,
enumValues = getStringListFromAnnotation(guideAnnotation, "enumValues"),
)

internal fun getDescriptionFromAnnotations(
Expand Down Expand Up @@ -103,6 +105,27 @@ internal fun getStringFromAnnotation(
return guidePropertyStringValue
}

internal fun getStringListFromAnnotation(
guideAnnotation: KSAnnotation?,
listName: String,
): List<String>? {
val rawValue =
guideAnnotation
?.arguments
?.firstOrNull { it.name?.getShortName()?.equals(listName) == true }
?.value
val list =
when (rawValue) {
is List<*> -> rawValue.mapNotNull { it as? String }
is Array<*> -> rawValue.mapNotNull { it as? String }
else -> null
}
if (list.isNullOrEmpty()) {
return null
}
return list
}

internal fun extractBaseKdoc(kdoc: String): String? {
return baseKdocRegex.matchEntire(kdoc)?.groups?.get(1)?.value?.trim().let {
if (it.isNullOrEmpty()) null else it
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ data class RootSchemaTestClass(
val stringTest: String,
val objTest: SecondarySchemaTestClass,
val enumTest: EnumTest,
@Guide(enumValues = ["NORTH", "SOUTH", "EAST", "WEST"]) val stringEnumTest: String,
) {
companion object
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -85,5 +85,11 @@ class FirebaseKspProcessorTest {
.isEqualTo("class kdoc should be used if property kdocs aren't present")
assertThat(objSchema.title).isEqualTo("objTest")
assertThat(objSchema.nullable).isEqualTo(false)

assertThat(rootSchema.properties?.get("stringEnumTest")).isNotNull()
val stringEnumSchema = rootSchema.properties?.get("stringEnumTest")!!
assertThat(stringEnumSchema.enum).isEqualTo(listOf("NORTH", "SOUTH", "EAST", "WEST"))
assertThat(stringEnumSchema.title).isEqualTo("stringEnumTest")
assertThat(stringEnumSchema.nullable).isEqualTo(false)
}
}
11 changes: 9 additions & 2 deletions ai-logic/firebase-ai-ksp-processor/test-app/test-app.gradle.kts
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ android {
compileSdk = 36
defaultConfig {
applicationId = "com.google.firebase.testing.processor"
minSdk = 23
minSdk = 26
Comment thread
milaGGL marked this conversation as resolved.
Outdated
targetSdk = 36
versionCode = 1
versionName = "1.0"
Expand All @@ -41,10 +41,17 @@ android {
}
}

kotlin { compilerOptions { jvmTarget = JvmTarget.JVM_1_8 } }
kotlin {
compilerOptions {
jvmTarget = JvmTarget.JVM_1_8
freeCompilerArgs.add("-Xskip-metadata-version-check")
Comment thread
milaGGL marked this conversation as resolved.
}
}

dependencies {
implementation(project(":ai-logic:firebase-ai"))
implementation(project(":ai-logic:firebase-ai-ondevice"))
implementation(libs.genai.prompt)
ksp(project(":ai-logic:firebase-ai-ksp-processor"))

implementation("com.google.firebase:firebase-common:22.0.0")
Expand Down
8 changes: 8 additions & 0 deletions ai-logic/firebase-ai-ondevice-interop/api.txt
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,13 @@ package com.google.firebase.ai.ondevice.interop {
property public final String? modelVersion;
}

public final class GenerateObjectResponse<T> {
ctor public GenerateObjectResponse(java.util.List<? extends T> instances);
ctor public GenerateObjectResponse(T instance);
method public java.util.List<T> getInstances();
property public final java.util.List<T> instances;
}

public final class GenerationConfig {
ctor public GenerationConfig(com.google.firebase.ai.ondevice.interop.ModelConfig modelConfig);
method public com.google.firebase.ai.ondevice.interop.ModelConfig getModelConfig();
Expand All @@ -117,6 +124,7 @@ package com.google.firebase.ai.ondevice.interop {
method public kotlinx.coroutines.flow.Flow<com.google.firebase.ai.ondevice.interop.DownloadStatusInterop> download();
method public suspend Object? generateContent(com.google.firebase.ai.ondevice.interop.GenerateContentRequest request, kotlin.coroutines.Continuation<? super com.google.firebase.ai.ondevice.interop.GenerateContentResponse>);
method public kotlinx.coroutines.flow.Flow<com.google.firebase.ai.ondevice.interop.GenerateContentResponse> generateContentStream(com.google.firebase.ai.ondevice.interop.GenerateContentRequest request);
method public suspend <T> Object? generateObject(com.google.firebase.ai.ondevice.interop.GenerateContentRequest request, kotlin.reflect.KClass<T> schemaClass, kotlin.coroutines.Continuation<? super com.google.firebase.ai.ondevice.interop.GenerateObjectResponse<T>>);
method public suspend Object? getBaseModelName(kotlin.coroutines.Continuation<? super java.lang.String>);
method public suspend Object? getTokenLimit(kotlin.coroutines.Continuation<? super java.lang.Integer>);
method public suspend Object? isAvailable(kotlin.coroutines.Continuation<? super java.lang.Boolean>);
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
/*
* Copyright 2026 Google LLC
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/

package com.google.firebase.ai.ondevice.interop

/**
* Represents a structured object generation response from the on-device model interop layer.
*
* @property instances The list of generated object instances across all candidates returned
* directly by the model engine.
*/
public class GenerateObjectResponse<T>(public val instances: List<T>) {
public constructor(instance: T) : this(listOf(instance))
}
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,21 @@ public interface GenerativeModel {
*/
public suspend fun generateContent(request: GenerateContentRequest): GenerateContentResponse

/**
* Generates a structured object from the input [GenerateContentRequest] given to the model as a
* prompt.
*
* @param request The input given to the model as a prompt.
* @param schemaClass The target SDK data class (`T`, e.g. `MovieReview::class`), from which the
* underlying engine dynamically resolves and maps the KSP shadow companion class at runtime.
* @throws [FirebaseAIOnDeviceNotAvailableException] if model is not available.
* @return The structured object response generated by the model.
*/
public suspend fun <T : Any> generateObject(
request: GenerateContentRequest,
schemaClass: kotlin.reflect.KClass<T>
): GenerateObjectResponse<T>

/**
* Counts the number of tokens in a prompt using the model's tokenizer.
*
Expand Down
7 changes: 5 additions & 2 deletions ai-logic/firebase-ai-ondevice/firebase-ai-ondevice.gradle.kts
Original file line number Diff line number Diff line change
Expand Up @@ -63,13 +63,16 @@ android {
}

kotlin {
compilerOptions { jvmTarget = JvmTarget.JVM_1_8 }
compilerOptions {
jvmTarget = JvmTarget.JVM_1_8
freeCompilerArgs.add("-Xskip-metadata-version-check")
}
explicitApi()
}

dependencies {
implementation(libs.genai.prompt)
implementation("com.google.firebase:firebase-ai-ondevice-interop:16.0.0-beta03")
implementation(project(":ai-logic:firebase-ai-ondevice-interop"))
Comment thread
VinayGuthal marked this conversation as resolved.

implementation(libs.firebase.common)
implementation(libs.firebase.components)
Expand Down
Loading
Loading