8000
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
1 change: 1 addition & 0 deletions firebase-ai/CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
# Unreleased

- [changed] Added the `hybrid` component to request headers coming from `prefer_in_cloud` configurations (#7857)
- [feature] Added experimental support for on-device inference (#7739)
- [feature] Added automatic function calling support with `AutoFunctionDeclaration`.
- [feature] Added no-argument overloads for `Tool.urlContext()` and `Tool.googleSearch()`.
Expand Down
8000
Original file line numberDiff line number Diff line change
Expand Up @@ -315,10 +315,10 @@ internal constructor(
/**
* Returns a [GenerativeModelProvider] that uses the cloud backend.
*
* @param isFallback Whether this provider is being used as a fallback for another provider.
* @param isHybrid Whether this provider is being used in a hybrid configuration.
* @return A [GenerativeModelProvider] that uses the cloud backend.
*/
internal fun buildCloudModelProvider(isFallback: Boolean = false): GenerativeModelProvider {
internal fun buildCloudModelProvider(isHybrid: Boolean = false): GenerativeModelProvider {
return CloudGenerativeModelProvider(
modelName = modelName,
generationConfig = generationConfig,
Expand All @@ -332,7 +332,7 @@ internal constructor(
apiKey,
modelName,
requestOptions,
if (isFallback) "${apiClient} hybrid" else apiClient,
if (isHybrid) "${apiClient} hybrid" else apiClient,
firebaseApp,
AppCheckHeaderProvider(
TAG,
Expand All @@ -359,13 +359,13 @@ internal constructor(
InferenceMode.PREFER_ON_DEVICE -> {
FallbackGenerativeModelProvider(
defaultModel = buildOnDeviceModelProvider(),
fallbackModel = buildCloudModelProvider(isFallback = true),
fallbackModel = buildCloudModelProvider(isHybrid = true),
shouldFallbackInException = true
)
}
InferenceMode.PREFER_IN_CLOUD ->
FallbackGenerativeModelProvider(
defaultModel = buildCloudModelProvider(),
defaultModel = buildCloudModelProvider(isHybrid = true),
fallbackModel = buildOnDeviceModelProvider(),
precondition =
NetworkStatusChecker(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -113,7 +113,7 @@ internal constructor(
model: String,
private val requestOptions: RequestOptions,
httpEngine: HttpClientEngine,
private val apiClient: String,
internal val apiClient: String,
private val firebaseApp: FirebaseApp,
private val appVersion: Int = 0,
private val googleAppId: String,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,7 @@ internal class CloudGenerativeModelProvider(
private val toolConfig: ToolConfig? = null,
private val systemInstruction: Content? = null,
private val generativeBackend: GenerativeBackend = GenerativeBackend.googleAI(),
private val controller: APIController,
internal val controller: APIController,
) : GenerativeModelProvider {

override suspend fun generateContent(prompt: List<Content>): GenerateContentResponse =
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -47,8 +47,8 @@ import kotlinx.coroutines.flow.onEach
* Defaults to `true`.
*/
internal class FallbackGenerativeModelProvider(
private val defaultModel: GenerativeModelProvider,
private val fallbackModel: GenerativeModelProvider,
internal val defaultModel: GenerativeModelProvider,
internal val fallbackModel: GenerativeModelProvider,
private val precondition: () -> Boolean = { true },
private val shouldFallbackInException: Boolean = true
) : GenerativeModelProvider {
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,116 @@
/*
* 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

import android.content.Context
import android.net.ConnectivityManager
import com.google.firebase.FirebaseApp
import com.google.firebase.ai.generativemodel.CloudGenerativeModelProvider
import com.google.firebase.ai.generativemodel.FallbackGenerativeModelProvider
import com.google.firebase.ai.type.GenerativeBackend
import com.google.firebase.ai.type.PublicPreviewAPI
import io.kotest.matchers.string.shouldContain
import io.kotest.matchers.string.shouldNotContain
import io.kotest.matchers.types.shouldBeInstanceOf
import io.mockk.every
import io.mockk.mockk
import org.junit.Before
import org.junit.Test

@OptIn(PublicPreviewAPI::class)
internal class GenerativeModelBuilderTests {
Comment thread
rlazo marked this conversation as resolved.

val firebaseApp = mockk<FirebaseApp>(relaxed = true)

@Before
fun setUp() {
every { firebaseApp.options.applicationId } returns "1:12345:android:67890"
val context = mockk<Context>(relaxed = true)
val connectivityManager = mockk<ConnectivityManager>(relaxed = true)
every { firebaseApp.applicationContext } returns context
every { context.getSystemService(Context.CONNECTIVITY_SERVICE) } returns connectivityManager
}

@Test
fun `getModelProvider uses hybrid suffix in PREFER_ON_DEVICE mode`() {
val builder =
GenerativeModel.Builder(
modelName = "gemini-1.5-flash",
apiKey = "apiKey",
firebaseApp = firebaseApp,
useLimitedUseAppCheckTokens = false,
generativeBackend = GenerativeBackend.googleAI()
)
.apply { }

val provider = builder.getModelProvider()

provider.shouldBeInstanceOf<FallbackGenerativeModelProvider>()
val fallbackModel = provider.fallbackModel

fallbackModel.shouldBeInstanceOf<CloudGenerativeModelProvider>()
val controller = fallbackModel.controller

val apiClient = controller.apiClient
apiClient shouldContain " hybrid"
}

@Test
fun `getModelProvider uses hybrid suffix in PREFER_IN_CLOUD mode`() {
val builder =
GenerativeModel.Builder(
modelName = "gemini-1.5-flash",
apiKey = "apiKey",
firebaseApp = firebaseApp,
useLimitedUseAppCheckTokens = false,
generativeBackend = GenerativeBackend.googleAI()
)
.apply { }

val provider = builder.getModelProvider()

provider.shouldBeInstanceOf<FallbackGenerativeModelProvider>()
val defaultModel = provider.defaultModel

defaultModel.shouldBeInstanceOf<CloudGenerativeModelProvider>()
val controller = defaultModel.controller

val apiClient = controller.apiClient
apiClient shouldContain " hybrid"
}

@Test
fun `getModelProvider does NOT use hybrid suffix in ONLY_IN_CLOUD mode`() {
val builder =
GenerativeModel.Builder(
modelName = "gemini-1.5-flash",
apiKey = "apiKey",
firebaseApp = firebaseApp,
useLimitedUseAppCheckTokens = false,
generativeBackend = GenerativeBackend.googleAI()
)
.apply { }

val provider = builder.getModelProvider()

provider.shouldBeInstanceOf<CloudGenerativeModelProvider>()
val controller = provider.controller

val apiClient = controller.apiClient
apiClient shouldNotContain " hybrid"
}
}
Loading
0