From 6269bdfef2d5a3290059bffd0ef92da5756df6b7 Mon Sep 17 00:00:00 2001 From: Austin Benoit Date: Thu, 8 Oct 2026 11:56:53 -0400 Subject: [PATCH 1/2] feat(ai): add Firebase AI Logic C++ SDK with LiteRT/LiteRT-LM hybrid inference --- .gitignore | 4 + CMakeLists.txt | 6 + ai/CMakeLists.txt | 265 +++ ai/README.md | 93 + ai/build.gradle | 88 + ai/samples/hybrid_chat_main.cc | 444 +++++ ai/src/android/http_sender_android.cc | 295 +++ ai/src/common/chat.cc | 373 ++++ ai/src/common/chat_internal.h | 86 + ai/src/common/firebase_ai.cc | 210 +++ ai/src/common/firebase_ai_internal.h | 68 + ai/src/common/generative_model.cc | 634 +++++++ ai/src/common/generative_model_internal.h | 119 ++ ai/src/common/http_client.cc | 298 +++ ai/src/common/http_client.h | 103 + ai/src/common/http_sender.h | 113 ++ ai/src/common/litert_adapter.cc | 1666 +++++++++++++++++ ai/src/common/litert_adapter.h | 107 ++ ai/src/common/litert_c_bridge.cc | 424 +++++ ai/src/common/litert_c_bridge.h | 93 + ai/src/common/model_content.cc | 300 +++ ai/src/common/serialization.cc | 1280 +++++++++++++ ai/src/common/serialization.h | 119 ++ ai/src/common/template_generative_model.cc | 463 +++++ .../template_generative_model_internal.h | 133 ++ ai/src/desktop/http_sender_desktop.cc | 185 ++ ai/src/include/firebase/ai.h | 154 ++ ai/src/include/firebase/ai/chat.h | 159 ++ ai/src/include/firebase/ai/function_calling.h | 282 +++ .../firebase/ai/generate_content_response.h | 375 ++++ .../include/firebase/ai/generation_config.h | 149 ++ ai/src/include/firebase/ai/generative_model.h | 195 ++ ai/src/include/firebase/ai/model_content.h | 499 +++++ ai/src/include/firebase/ai/safety.h | 105 ++ ai/src/include/firebase/ai/schema.h | 324 ++++ .../firebase/ai/template_chat_session.h | 104 + .../firebase/ai/template_generative_model.h | 141 ++ ai/src/include/firebase/ai/types.h | 608 ++++++ ai/src/ios/http_sender_ios.mm | 193 ++ ai/tests/CMakeLists.txt | 25 + ai/tests/ai_test.cc | 353 ++++ app/CMakeLists.txt | 22 +- settings.gradle | 3 +- 43 files changed, 11655 insertions(+), 5 deletions(-) create mode 100644 ai/CMakeLists.txt create mode 100644 ai/README.md create mode 100644 ai/build.gradle create mode 100644 ai/samples/hybrid_chat_main.cc create mode 100644 ai/src/android/http_sender_android.cc create mode 100644 ai/src/common/chat.cc create mode 100644 ai/src/common/chat_internal.h create mode 100644 ai/src/common/firebase_ai.cc create mode 100644 ai/src/common/firebase_ai_internal.h create mode 100644 ai/src/common/generative_model.cc create mode 100644 ai/src/common/generative_model_internal.h create mode 100644 ai/src/common/http_client.cc create mode 100644 ai/src/common/http_client.h create mode 100644 ai/src/common/http_sender.h create mode 100644 ai/src/common/litert_adapter.cc create mode 100644 ai/src/common/litert_adapter.h create mode 100644 ai/src/common/litert_c_bridge.cc create mode 100644 ai/src/common/litert_c_bridge.h create mode 100644 ai/src/common/model_content.cc create mode 100644 ai/src/common/serialization.cc create mode 100644 ai/src/common/serialization.h create mode 100644 ai/src/common/template_generative_model.cc create mode 100644 ai/src/common/template_generative_model_internal.h create mode 100644 ai/src/desktop/http_sender_desktop.cc create mode 100644 ai/src/include/firebase/ai.h create mode 100644 ai/src/include/firebase/ai/chat.h create mode 100644 ai/src/include/firebase/ai/function_calling.h create mode 100644 ai/src/include/firebase/ai/generate_content_response.h create mode 100644 ai/src/include/firebase/ai/generation_config.h create mode 100644 ai/src/include/firebase/ai/generative_model.h create mode 100644 ai/src/include/firebase/ai/model_content.h create mode 100644 ai/src/include/firebase/ai/safety.h create mode 100644 ai/src/include/firebase/ai/schema.h create mode 100644 ai/src/include/firebase/ai/template_chat_session.h create mode 100644 ai/src/include/firebase/ai/template_generative_model.h create mode 100644 ai/src/include/firebase/ai/types.h create mode 100644 ai/src/ios/http_sender_ios.mm create mode 100644 ai/tests/CMakeLists.txt create mode 100644 ai/tests/ai_test.cc diff --git a/.gitignore b/.gitignore index b8ad9df906..833303bd12 100644 --- a/.gitignore +++ b/.gitignore @@ -22,11 +22,15 @@ __pycache__/ # Unencrypted secret files google-services.json +google-services-desktop.json GoogleService-Info.plist uri_prefix.txt server_key.txt gcs_key_file.json +# On-device LiteRT-LM model weight files +*.litertlm + # Folders for cmake/test output *_build/ cmake-build-*/ diff --git a/CMakeLists.txt b/CMakeLists.txt index 7cd1a549a2..86b2b44743 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -28,6 +28,9 @@ option(FIREBASE_INCLUDE_LIBRARY_DEFAULT "Should each library be included by default." ON) # Different options to enable/disable each library being included during # configuration. +option(FIREBASE_INCLUDE_AI + "Include the Firebase AI Logic library." + ${FIREBASE_INCLUDE_LIBRARY_DEFAULT}) option(FIREBASE_INCLUDE_ANALYTICS "Include the Google Analytics for Firebase library." ${FIREBASE_INCLUDE_LIBRARY_DEFAULT}) @@ -632,6 +635,9 @@ else() ) endif() +if (FIREBASE_INCLUDE_AI) + add_subdirectory(ai) +endif() if (FIREBASE_INCLUDE_ANALYTICS) add_subdirectory(analytics) add_dependencies(FIREBASE_GENERATED_HEADERS FIREBASE_ANALYTICS_GENERATED_HEADERS) diff --git a/ai/CMakeLists.txt b/ai/CMakeLists.txt new file mode 100644 index 0000000000..d0cf39467f --- /dev/null +++ b/ai/CMakeLists.txt @@ -0,0 +1,265 @@ +# Copyright 2025 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. + +# CMake file for the firebase_ai library + +# Common C++ source files used across all platforms (Desktop, Android, iOS) +set(common_SRCS + src/common/chat.cc + src/common/firebase_ai.cc + src/common/generative_model.cc + src/common/http_client.cc + src/common/litert_adapter.cc + src/common/litert_c_bridge.cc + src/common/model_content.cc + src/common/serialization.cc + src/common/template_generative_model.cc) + +# Slim HTTP sender for Android (JNI -> java.net.HttpURLConnection) +set(android_SRCS + src/android/http_sender_android.cc) + +# Slim HTTP sender for iOS (Objective-C++ -> NSURLSession) +set(ios_SRCS + src/ios/http_sender_ios.mm) + +# Slim HTTP sender for Desktop / Linux / macOS / Windows (libcurl via firebase_rest_lib) +set(desktop_SRCS + src/desktop/http_sender_desktop.cc) + +if(ANDROID) + set(ai_platform_SRCS + "${android_SRCS}") +elseif(IOS) + set(ai_platform_SRCS + "${ios_SRCS}") +else() + set(ai_platform_SRCS + "${desktop_SRCS}") +endif() + +if(ANDROID OR IOS) + set(additional_link_LIB) +else() + set(additional_link_LIB + firebase_rest_lib + libcurl) +endif() + +add_library(firebase_ai STATIC + ${common_SRCS} + ${ai_platform_SRCS}) + +set_property(TARGET firebase_ai PROPERTY FOLDER "Firebase Cpp") + +# Set up the dependency on Firebase App. +target_link_libraries(firebase_ai + PUBLIC + firebase_app + PRIVATE + ${additional_link_LIB} +) + +# Public headers all refer to each other relative to the src/include directory, +# while private headers are relative to the entire C++ SDK directory. +target_include_directories(firebase_ai + PUBLIC + ${CMAKE_CURRENT_LIST_DIR}/src/include + PRIVATE + ${FIREBASE_CPP_SDK_ROOT_DIR} + ${FLATBUFFERS_SOURCE_DIR}/include + ${FIREBASE_GEN_FILE_DIR} + ${FIREBASE_SPM_SWIFT_HEADERS_DIR} + ${CURL_SOURCE_DIR}/include +) + +target_compile_definitions(firebase_ai + PRIVATE + -DINTERNAL_EXPERIMENTAL=1 + -DFIREBASE_AI_BUILD_LIB_DIR="${CMAKE_CURRENT_BINARY_DIR}" +) + +# Automatically download Google AI Edge LiteRT C/C++ SDK and LiteRT-LM runtime +# libraries on a fresh checkout (models are NOT downloaded by CMake). +option(FIREBASE_AI_DOWNLOAD_LITERT + "Automatically download Google AI Edge LiteRT and LiteRT-LM SDKs/runtimes" + ON) +set(FIREBASE_AI_LITERT_VERSION "2.2.0" CACHE STRING "LiteRT release version") +set(FIREBASE_AI_LITERT_LM_VERSION "0.18.0" CACHE STRING "LiteRT-LM release version") + +set(FIREBASE_AI_LITERT_DEPS_DIR "${CMAKE_CURRENT_BINARY_DIR}/litert_deps") +if(FIREBASE_AI_DOWNLOAD_LITERT) + file(MAKE_DIRECTORY "${FIREBASE_AI_LITERT_DEPS_DIR}") + + # 1. LiteRT C/C++ SDK headers (litert_cc_sdk.zip) + set(LITERT_CC_SDK_ZIP "${FIREBASE_AI_LITERT_DEPS_DIR}/litert_cc_sdk.zip") + set(LITERT_CC_SDK_EXTRACTED "${FIREBASE_AI_LITERT_DEPS_DIR}/litert_cc_sdk") + if(NOT EXISTS "${LITERT_CC_SDK_EXTRACTED}/litert/c/litert_compiled_model.h") + set(LITERT_CC_SDK_URL + "https://github.com/google-ai-edge/LiteRT/releases/download/v${FIREBASE_AI_LITERT_VERSION}/litert_cc_sdk.zip") + message(STATUS "Firebase AI: Downloading LiteRT C/C++ SDK from ${LITERT_CC_SDK_URL}") + file(DOWNLOAD "${LITERT_CC_SDK_URL}" "${LITERT_CC_SDK_ZIP}" + STATUS LITERT_SDK_DL_STATUS) + list(GET LITERT_SDK_DL_STATUS 0 LITERT_SDK_DL_CODE) + if(LITERT_SDK_DL_CODE EQUAL 0) + file(ARCHIVE_EXTRACT INPUT "${LITERT_CC_SDK_ZIP}" + DESTINATION "${FIREBASE_AI_LITERT_DEPS_DIR}") + else() + message(WARNING "Firebase AI: Failed to download LiteRT C/C++ SDK (${LITERT_SDK_DL_STATUS})") + endif() + endif() + if(NOT DEFINED FIREBASE_AI_LITERT_CC_SDK_DIR AND + EXISTS "${LITERT_CC_SDK_EXTRACTED}/litert/c/litert_compiled_model.h") + set(FIREBASE_AI_LITERT_CC_SDK_DIR "${LITERT_CC_SDK_EXTRACTED}") + endif() + + # 2. Platform LiteRT & LiteRT-LM shared runtime libraries + if(APPLE AND NOT IOS) + set(LITERT_DYLIB_PATH "${CMAKE_CURRENT_BINARY_DIR}/libLiteRt.dylib") + if(NOT EXISTS "${LITERT_DYLIB_PATH}") + set(LITERT_DYLIB_URL + "https://storage.googleapis.com/litert/binaries/${FIREBASE_AI_LITERT_VERSION}/macos_arm64/libLiteRt.dylib") + message(STATUS "Firebase AI: Downloading LiteRT runtime from ${LITERT_DYLIB_URL}") + file(DOWNLOAD "${LITERT_DYLIB_URL}" "${LITERT_DYLIB_PATH}" + STATUS LITERT_DYLIB_DL_STATUS) + endif() + + set(CLITERTLM_DYLIB_PATH "${CMAKE_CURRENT_BINARY_DIR}/libCLiteRTLM_mac.dylib") + set(CLITERTLM_EXTRACTED_DYLIB + "${FIREBASE_AI_LITERT_DEPS_DIR}/CLiteRTLM_mac.xcframework/macos-arm64_x86_64/libCLiteRTLM_mac.dylib") + if(NOT EXISTS "${CLITERTLM_EXTRACTED_DYLIB}") + set(CLITERTLM_ZIP "${FIREBASE_AI_LITERT_DEPS_DIR}/CLiteRTLM_mac.xcframework.zip") + set(CLITERTLM_URL + "https://github.com/google-ai-edge/LiteRT-LM/releases/download/v${FIREBASE_AI_LITERT_LM_VERSION}/CLiteRTLM_mac.xcframework.zip") + message(STATUS "Firebase AI: Downloading LiteRT-LM macOS runtime from ${CLITERTLM_URL}") + file(DOWNLOAD "${CLITERTLM_URL}" "${CLITERTLM_ZIP}" + STATUS CLITERTLM_DL_STATUS) + list(GET CLITERTLM_DL_STATUS 0 CLITERTLM_DL_CODE) + if(CLITERTLM_DL_CODE EQUAL 0) + file(ARCHIVE_EXTRACT INPUT "${CLITERTLM_ZIP}" + DESTINATION "${FIREBASE_AI_LITERT_DEPS_DIR}") + else() + message(WARNING "Firebase AI: Failed to download LiteRT-LM runtime (${CLITERTLM_DL_STATUS})") + endif() + endif() + if(EXISTS "${CLITERTLM_EXTRACTED_DYLIB}" AND NOT EXISTS "${CLITERTLM_DYLIB_PATH}") + file(COPY_FILE "${CLITERTLM_EXTRACTED_DYLIB}" "${CLITERTLM_DYLIB_PATH}") + endif() + if(EXISTS "${FIREBASE_AI_LITERT_DEPS_DIR}/CLiteRTLM_mac.xcframework/macos-arm64_x86_64/Headers") + target_include_directories(firebase_ai PRIVATE + "${FIREBASE_AI_LITERT_DEPS_DIR}/CLiteRTLM_mac.xcframework/macos-arm64_x86_64/Headers") + endif() + elseif(IOS) + set(CLITERTLM_IOS_DIR "${FIREBASE_AI_LITERT_DEPS_DIR}/CLiteRTLM.xcframework") + if(NOT EXISTS "${CLITERTLM_IOS_DIR}") + set(CLITERTLM_IOS_ZIP "${FIREBASE_AI_LITERT_DEPS_DIR}/CLiteRTLM.xcframework.zip") + set(CLITERTLM_IOS_URL + "https://github.com/google-ai-edge/LiteRT-LM/releases/download/v${FIREBASE_AI_LITERT_LM_VERSION}/CLiteRTLM.xcframework.zip") + message(STATUS "Firebase AI: Downloading LiteRT-LM iOS xcframework from ${CLITERTLM_IOS_URL}") + file(DOWNLOAD "${CLITERTLM_IOS_URL}" "${CLITERTLM_IOS_ZIP}" + STATUS CLITERTLM_IOS_DL_STATUS) + list(GET CLITERTLM_IOS_DL_STATUS 0 CLITERTLM_IOS_DL_CODE) + if(CLITERTLM_IOS_DL_CODE EQUAL 0) + file(ARCHIVE_EXTRACT INPUT "${CLITERTLM_IOS_ZIP}" + DESTINATION "${FIREBASE_AI_LITERT_DEPS_DIR}") + endif() + endif() + elseif(CMAKE_SYSTEM_NAME STREQUAL "Linux") + set(LITERT_SO_PATH "${CMAKE_CURRENT_BINARY_DIR}/libLiteRt.so") + if(NOT EXISTS "${LITERT_SO_PATH}") + if(CMAKE_SYSTEM_PROCESSOR MATCHES "aarch64|arm64") + set(LITERT_LINUX_ARCH "linux_arm64") + else() + set(LITERT_LINUX_ARCH "linux_x86_64") + endif() + set(LITERT_SO_URL + "https://storage.googleapis.com/litert/binaries/${FIREBASE_AI_LITERT_VERSION}/${LITERT_LINUX_ARCH}/libLiteRt.so") + message(STATUS "Firebase AI: Downloading LiteRT runtime from ${LITERT_SO_URL}") + file(DOWNLOAD "${LITERT_SO_URL}" "${LITERT_SO_PATH}" + STATUS LITERT_SO_DL_STATUS) + endif() + endif() +endif() + +if(DEFINED FIREBASE_AI_LITERT_CC_SDK_DIR) + if(EXISTS "${FIREBASE_AI_LITERT_CC_SDK_DIR}/litert/c/litert_compiled_model.h") + target_include_directories(firebase_ai PRIVATE "${FIREBASE_AI_LITERT_CC_SDK_DIR}") + elseif(EXISTS "${FIREBASE_AI_LITERT_CC_SDK_DIR}/include") + target_include_directories(firebase_ai PRIVATE "${FIREBASE_AI_LITERT_CC_SDK_DIR}/include") + target_compile_definitions(firebase_ai PRIVATE -DFIREBASE_AI_USE_LITERT_CC_SDK=1) + if(EXISTS "${FIREBASE_AI_LITERT_CC_SDK_DIR}/lib/libLiteRt.dylib") + target_link_libraries(firebase_ai PRIVATE "${FIREBASE_AI_LITERT_CC_SDK_DIR}/lib/libLiteRt.dylib") + endif() + endif() +endif() + +if(ANDROID) + firebase_cpp_proguard_file(ai) +elseif(IOS) + target_compile_options(firebase_ai + PUBLIC $<$>:-fobjc-arc>) + + setup_spm_headers( + firebase_ai + MODULES + FirebaseCore + ) + + if (FIREBASE_XCODE_TARGET_FORMAT STREQUAL "frameworks") + set_target_properties(firebase_ai PROPERTIES + FRAMEWORK TRUE + ) + endif() +endif() + +if(NOT IOS) + add_library(firebase_ai_litert_bridge SHARED + src/common/litert_c_bridge.cc) + target_link_libraries(firebase_ai_litert_bridge + PRIVATE + firebase_ai + firebase_app) + if(APPLE) + target_link_libraries(firebase_ai_litert_bridge PRIVATE "-framework Foundation") + endif() + target_include_directories(firebase_ai_litert_bridge + PRIVATE + ${CMAKE_CURRENT_LIST_DIR}/src/include + ${FIREBASE_CPP_SDK_ROOT_DIR} + ${FLATBUFFERS_SOURCE_DIR}/include + ) +endif() + +if(NOT ANDROID AND NOT IOS) + add_executable(firebase_ai_hybrid_chat + samples/hybrid_chat_main.cc) + target_link_libraries(firebase_ai_hybrid_chat + firebase_ai + firebase_app) + if(APPLE) + target_link_libraries(firebase_ai_hybrid_chat "-framework Foundation") + endif() + target_include_directories(firebase_ai_hybrid_chat + PRIVATE + ${CMAKE_CURRENT_LIST_DIR}/src/include + ${FIREBASE_CPP_SDK_ROOT_DIR} + ) +endif() + +if(FIREBASE_CPP_BUILD_TESTS) + add_subdirectory(tests) +endif() + +cpp_pack_library(firebase_ai "") +cpp_pack_public_headers() diff --git a/ai/README.md b/ai/README.md new file mode 100644 index 0000000000..bee9dac3da --- /dev/null +++ b/ai/README.md @@ -0,0 +1,93 @@ +# Firebase AI Logic C++ SDK (`firebase::ai`) — Cloud & LiteRT On-Device Hybrid + +This directory implements the `firebase::ai` C++ SDK for **Firebase AI Logic** (Gemini Developer API and Vertex AI Gemini API), including **Hybrid On-Device + Cloud Inference** powered by [Google AI Edge LiteRT](https://developers.google.com/edge/litert/overview#c++_1) and [LiteRT-LM](https://github.com/google-ai-edge/LiteRT-LM). + +It also builds `firebase_ai_litert_bridge`, a C ABI shared library used by the **Firebase Unity SDK** (`Firebase.AI`) via P/Invoke so Unity apps can run on-device and hybrid inference through the C++ LiteRT engine while keeping their existing C# Cloud implementation. + +--- + +## Quick Start (Build & Run the Hybrid Multi-Turn Chat Demo) + +### 1. Configure & Build via CMake + +By default (`FIREBASE_AI_DOWNLOAD_LITERT=ON`), CMake automatically downloads the LiteRT C++ SDK headers (`v2.2.0`), `libLiteRt`, and the `CLiteRTLM` (`v0.18.0`) runtime library during configuration and copies the runtime next to the built binary. Model weights (`.litertlm`) are **not** downloaded automatically. + +From the repository root (`firebase-cpp-sdk`): + +```bash +cmake -S . -B desktop_build \ + -DFIREBASE_INCLUDE_AI=ON \ + -DFIREBASE_AI_BUILD_SAMPLES=ON \ + -DFIREBASE_AI_BUILD_UNITY_BRIDGE=ON \ + -DFIREBASE_CPP_BUILD_TESTS=ON + +cmake --build desktop_build \ + --target firebase_ai_hybrid_chat firebase_ai_litert_bridge firebase_ai_test -j8 +``` + +### 2. Download a Local Gemma `.litertlm` Model + +Download a LiteRT-LM `.litertlm` model (for example, **Gemma 3 1B IT INT4** with a 4,096-token context window) and place it in `desktop_build/ai/` or pass its path via `--model`: + +```bash +curl -L "https://huggingface.co/litert-community/Gemma3-1B-IT/resolve/main/gemma3-1b-it-int4.litertlm" \ + -o desktop_build/ai/gemma3-1b-it-int4.litertlm +``` + +### 3. Provide Your Firebase Configuration (`google-services.json`) + +The demo initializes `firebase::App` using a standard Firebase `google-services.json` (or `google-services-desktop.json`) file. Place `google-services.json` in your working directory (or `desktop_build/ai/`), or pass `--config /path/to/google-services.json`. + +### 4. Run the Interactive Multi-Turn Hybrid Chat App + +```bash +./desktop_build/ai/firebase_ai_hybrid_chat \ + --config /path/to/google-services.json \ + --model ./desktop_build/ai/gemma3-1b-it-int4.litertlm +``` + +*(If `google-services.json` and `gemma3-1b-it-int4.litertlm` are placed in `desktop_build/ai/` or the current directory, they are auto-detected and you can run `./desktop_build/ai/firebase_ai_hybrid_chat` with no arguments.)* + +--- + +## Interactive Chat Commands + +Inside `firebase_ai_hybrid_chat`, both Cloud (`gemini-3.1-flash-lite`) and On-Device (`LiteRT-LM`) share a single multi-turn `firebase::ai::Chat` history, so you can switch backends mid-conversation without losing context: + +| Command | Description | +| :--- | :--- | +| `/toggle` | Cycle between `ONLY_ON_DEVICE` -> `ONLY_IN_CLOUD` -> `PREFER_ON_DEVICE` -> `PREFER_IN_CLOUD` | +| `/mode local` | Force local on-device inference (`kInferenceModeOnlyOnDevice`) | +| `/mode cloud` | Force cloud Firebase AI inference (`kInferenceModeOnlyInCloud`) | +| `/mode hybrid` | Prefer on-device LiteRT; automatically fall back to Cloud on error (`kInferenceModePreferOnDevice`) | +| `/mode fallback` | Prefer Cloud; automatically fall back to on-device LiteRT when offline (`kInferenceModePreferInCloud`) | +| `/history` | Print the accumulated multi-turn `Chat` history | +| `/compact` | Summarize and compact the conversation history into a concise 2-turn context using the active model | +| `/clear` | Clear the multi-turn `Chat` history | +| `/quit` | Exit the application | + +### CLI Flags + +- `--model `: Path to a `.litertlm` (LiteRT-LM LLM) or `.tflite` (LiteRT `CompiledModel`) file. +- `--config `: Path to `google-services.json` or `google-services-desktop.json`. +- `--cloud-model `: Cloud Gemini model name (default: `gemini-3.1-flash-lite`). +- `--mode `: Initial inference mode (default: `hybrid` / `PREFER_ON_DEVICE`). +- `--gpu`: Use GPU acceleration (`kLiteRtAcceleratorGpu`) instead of CPU. +- `--no-stream`: Use non-streaming `Chat::SendMessage` instead of `Chat::SendMessageStream`. +- `--demo`: Run a non-interactive 2-turn verification (Turn 1 on-device -> toggle -> Turn 2 in cloud). + +--- + +## Context Window & Automatic Compaction + +- **Auto-Detected Context Length:** When `OnDeviceParams::max_num_tokens` is `0` (the default), `LiteRtAdapter` queries `litert_lm_loaded_file_max_context_tokens` from the `.litertlm` file metadata (`4096` tokens for `gemma3-1b-it-int4.litertlm`, `1024` tokens for `gemma3-270m.litertlm`). +- **Automatic Context Compaction:** Before each on-device turn, `LiteRtAdapter` tokenizes the conversation history via `litert_lm_engine_tokenize`. If the accumulated history exceeds the input token budget, older turns are automatically compacted into `[Compacted Earlier Conversation History]` while keeping recent turns verbatim and reserving headroom for generation output. +- **Repetition Prevention:** On-device generation configures `LiteRtLmRepetitionPenaltyConfig` (`repetition_penalty = 1.15`, `frequency_penalty = 0.25`, `presence_penalty = 0.1`) and `LiteRtLmNoRepeatNgramConfig` (`no_repeat_ngram_size = 4`) so small quantized models do not fall into token repetition loops on long outputs. + +--- + +## Running Unit Tests + +```bash +./desktop_build/ai/tests/firebase_ai_test +``` diff --git a/ai/build.gradle b/ai/build.gradle new file mode 100644 index 0000000000..7cbadd6cda --- /dev/null +++ b/ai/build.gradle @@ -0,0 +1,88 @@ +// Copyright 2025 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. + +buildscript { + repositories { + google() + mavenCentral() + } + dependencies { + classpath 'com.android.tools.build:gradle:7.4.2' + } +} +allprojects { + repositories { + google() + mavenCentral() + } +} + +apply plugin: 'com.android.library' + +android { + compileSdkVersion 34 + ndkPath System.getenv('ANDROID_NDK_HOME') + buildToolsVersion '32.0.0' + + sourceSets { + main { + manifest.srcFile '../android_build_files/AndroidManifest.xml' + } + } + + externalNativeBuild { + cmake { + path '../CMakeLists.txt' + } + } + + defaultConfig { + minSdkVersion 24 + targetSdkVersion 34 + versionCode 1 + versionName "1.0" + + buildTypes { + release { + minifyEnabled false + } + } + + externalNativeBuild { + cmake { + targets 'firebase_ai' + arguments '-DFIREBASE_CPP_USE_PRIOR_GRADLE_BUILD=ON', + '-DFIREBASE_INCLUDE_LIBRARY_DEFAULT=OFF', + '-DFIREBASE_INCLUDE_AI=ON' + } + } + } + + lintOptions { + abortOnError false + } +} + +dependencies { + implementation project(':app') +} +apply from: "$rootDir/android_build_files/android_abis.gradle" +apply from: "$rootDir/android_build_files/generate_proguard.gradle" +project.afterEvaluate { + generateProguardFile('ai') + preBuild.dependsOn(':app:build') + project.tasks.withType(com.android.build.gradle.internal.tasks.CheckAarMetadataTask) { + enabled = false + } +} diff --git a/ai/samples/hybrid_chat_main.cc b/ai/samples/hybrid_chat_main.cc new file mode 100644 index 0000000000..1da7452ea6 --- /dev/null +++ b/ai/samples/hybrid_chat_main.cc @@ -0,0 +1,444 @@ +/* + * Copyright 2025 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. + */ + +// Multi-turn C++ Hybrid Chat sample for Firebase AI Logic + Google AI Edge +// LiteRT / LiteRT-LM (Gemma). +// +// Features: +// - Live toggle (`/toggle`, `/mode local`, `/mode cloud`, `/mode hybrid`) +// between Cloud Firebase AI (`gemini-2.5-flash`) and Local On-Device LiteRT +// inference (`.litertlm` Gemma or `.tflite` CompiledModel) while sharing a +// single multi-turn `firebase::ai::Chat` history. +// - Real-time token streaming for both Cloud SSE and Local LiteRT-LM. + +#include +#include +#include +#include +#include +#include +#include +#include + +#include "firebase/ai.h" +#include "firebase/app.h" + +namespace { + +using ::firebase::App; +using ::firebase::AppOptions; +using ::firebase::Future; +using ::firebase::kFutureStatusComplete; +using ::firebase::ai::Backend; +using ::firebase::ai::Chat; +using ::firebase::ai::CountTokensResponse; +using ::firebase::ai::FirebaseAI; +using ::firebase::ai::GenerateContentResponse; +using ::firebase::ai::GenerationConfig; +using ::firebase::ai::GenerativeModel; +using ::firebase::ai::HybridParams; +using ::firebase::ai::InferenceMode; +using ::firebase::ai::InferenceSource; +using ::firebase::ai::kErrorNone; +using ::firebase::ai::kInferenceModeOnlyInCloud; +using ::firebase::ai::kInferenceModeOnlyOnDevice; +using ::firebase::ai::kInferenceModePreferInCloud; +using ::firebase::ai::kInferenceModePreferOnDevice; +using ::firebase::ai::kInferenceSourceOnDevice; +using ::firebase::ai::kLiteRtAcceleratorCpu; +using ::firebase::ai::kLiteRtAcceleratorGpu; +using ::firebase::ai::LiteRtAccelerator; +using ::firebase::ai::ModelContent; +using ::firebase::ai::OnDeviceParams; + +bool FileExists(const std::string& path) { + if (path.empty()) return false; + std::ifstream ifs(path.c_str(), std::ios::binary); + return ifs.good(); +} + +std::string ParentDir(const std::string& path) { + size_t pos = path.find_last_of("/\\"); + if (pos == std::string::npos) return "."; + return path.substr(0, pos); +} + +std::string FindDefaultLocalModel(const char* argv0) { + const char* env_model = std::getenv("FIREBASE_LITERT_MODEL"); + if (env_model && env_model[0] != '\0') { + return env_model; + } + std::string bin_dir = argv0 ? ParentDir(argv0) : "."; + std::vector candidates = { + bin_dir + "/gemma3-1b-it-int4.litertlm", + bin_dir + "/gemma3-270m.litertlm", + "./gemma3-1b-it-int4.litertlm", + "./gemma3-270m.litertlm", + "./desktop_build/ai/gemma3-1b-it-int4.litertlm", + "./desktop_build/ai/gemma3-270m.litertlm", + "./firebase-cpp-sdk/desktop_build/ai/gemma3-1b-it-int4.litertlm", + "./firebase-cpp-sdk/desktop_build/ai/gemma3-270m.litertlm", + }; + for (const auto& candidate : candidates) { + if (FileExists(candidate)) { + return candidate; + } + } + return "simulated://gemma-3-270m-it"; +} + +std::string ReadFileToString(const std::string& path) { + std::ifstream ifs(path.c_str(), std::ios::binary); + if (!ifs.good()) return ""; + return std::string((std::istreambuf_iterator(ifs)), + std::istreambuf_iterator()); +} + +std::string FindFirebaseConfigFile(const char* argv0, + const std::string& explicit_config_path) { + if (!explicit_config_path.empty() && FileExists(explicit_config_path)) { + return explicit_config_path; + } + std::string bin_dir = argv0 ? ParentDir(argv0) : "."; + std::vector search_dirs = { + ".", + bin_dir, + "./ai", + "./ai/samples", + "./desktop_build/ai", + "./firebase-cpp-sdk/ai", + "./firebase-cpp-sdk/ai/samples", + "./firebase-cpp-sdk/desktop_build/ai", + }; + const char* config_names[] = {"google-services-desktop.json", + "google-services.json"}; + for (const auto& dir : search_dirs) { + for (const char* name : config_names) { + std::string candidate = (dir == ".") ? name : (dir + "/" + name); + if (FileExists(candidate)) { + return candidate; + } + } + } + return ""; +} + +const char* ModeToString(InferenceMode mode) { + switch (mode) { + case kInferenceModeOnlyOnDevice: + return "LOCAL ONLY (LiteRT On-Device)"; + case kInferenceModeOnlyInCloud: + return "CLOUD ONLY (Firebase AI Gemini)"; + case kInferenceModePreferOnDevice: + return "HYBRID: PREFER_ON_DEVICE (LiteRT -> Cloud fallback)"; + case kInferenceModePreferInCloud: + return "HYBRID: PREFER_IN_CLOUD (Cloud -> LiteRT fallback)"; + } + return "UNKNOWN"; +} + +const char* SourceBadge(InferenceSource source) { + return source == kInferenceSourceOnDevice ? "[ON_DEVICE LiteRT]" + : "[IN_CLOUD Firebase]"; +} + +template +void WaitForFuture(const Future& fut) { + while (fut.status() != kFutureStatusComplete) { + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } +} + +void PrintBanner(const App& app, const Chat& chat, + const std::string& config_source, + const std::string& cloud_model, + const std::string& local_model_path) { + std::cout << "\n=============================================================" + "=======\n"; + std::cout << " Firebase AI Logic C++ Hybrid Multi-Turn Chat (Cloud + LiteRT " + "On-Device)\n"; + std::cout << "===============================================================" + "=====\n"; + std::cout << " Firebase Cfg: " << config_source + << " (project_id=" << app.options().project_id() << ")\n"; + std::cout << " Cloud Model : " << cloud_model << "\n"; + std::cout << " Local Model : " << local_model_path << " (" + << (chat.IsOnDeviceAvailable() ? "AVAILABLE" : "UNAVAILABLE") + << ")\n"; + std::cout << " Active Mode : [Toggle: " + << ModeToString(chat.inference_mode()) << "]\n"; + std::cout << "---------------------------------------------------------------" + "-----\n"; + std::cout << " Commands:\n"; + std::cout << " /toggle Cycle between LOCAL -> CLOUD -> HYBRID " + "(Prefer Local) -> HYBRID (Prefer Cloud)\n"; + std::cout << " /mode local Switch to ONLY_ON_DEVICE (LiteRT)\n"; + std::cout << " /mode cloud Switch to ONLY_IN_CLOUD (Firebase AI)\n"; + std::cout << " /mode hybrid Switch to PREFER_ON_DEVICE\n"; + std::cout << " /mode fallback Switch to PREFER_IN_CLOUD (offline fallback " + "to LiteRT)\n"; + std::cout << " /history Print accumulated multi-turn Chat history\n"; + std::cout << " /compact Summarize & compact Chat history using " + "active model\n"; + std::cout << " /clear Clear multi-turn Chat history\n"; + std::cout << " /quit Exit\n"; + std::cout << "===============================================================" + "=====\n\n"; +} + +void PrintHistory(const Chat& chat) { + std::vector hist = chat.history(); + std::cout << "\n--- Chat History (" << hist.size() << " turns) ---\n"; + for (size_t i = 0; i < hist.size(); ++i) { + std::cout << " [" << i + 1 << "] " << hist[i].role() << ": "; + for (const auto& part : hist[i].parts()) { + if (part.is_text()) { + std::cout << part.text_part().text; + } + } + std::cout << "\n"; + } + std::cout << "-----------------------------------\n\n"; +} + +void SendTurn(Chat* chat, const std::string& user_input, bool use_streaming) { + if (use_streaming) { + bool printed_prefix = false; + Future stream_fut = chat->SendMessageStream( + user_input, [&printed_prefix](const GenerateContentResponse& chunk) { + if (!printed_prefix) { + std::cout << SourceBadge(chunk.inference_source()) << " Model: "; + printed_prefix = true; + } + std::cout << chunk.text() << std::flush; + }); + WaitForFuture(stream_fut); + if (stream_fut.error() != kErrorNone) { + std::cout << "\n[Error " << stream_fut.error() << "] " + << (stream_fut.error_message() ? stream_fut.error_message() + : "") + << "\n\n"; + } else { + std::cout << "\n\n"; + } + return; + } + + Future fut = chat->SendMessage(user_input); + WaitForFuture(fut); + if (fut.error() != kErrorNone || fut.result() == nullptr) { + std::cout << "[Error " << fut.error() << "] " + << (fut.error_message() ? fut.error_message() : "") << "\n\n"; + return; + } + const GenerateContentResponse& resp = *fut.result(); + std::cout << SourceBadge(resp.inference_source()) << " Model: " << resp.text() + << "\n\n"; +} + +} // namespace + +int main(int argc, char** argv) { + std::string local_model_path = FindDefaultLocalModel(argc > 0 ? argv[0] : ""); + std::string litert_lib_path = + std::getenv("FIREBASE_LITERT_LM_LIB_PATH") + ? std::getenv("FIREBASE_LITERT_LM_LIB_PATH") + : (std::getenv("FIREBASE_LITERT_LIB_PATH") + ? std::getenv("FIREBASE_LITERT_LIB_PATH") + : ""); + std::string config_path = + std::getenv("FIREBASE_CONFIG") ? std::getenv("FIREBASE_CONFIG") : ""; + std::string cloud_model = "gemini-3.1-flash-lite"; + InferenceMode initial_mode = kInferenceModePreferOnDevice; + LiteRtAccelerator accelerator = kLiteRtAcceleratorCpu; + bool use_streaming = true; + bool run_demo = false; + + for (int i = 1; i < argc; ++i) { + std::string arg = argv[i]; + if (arg == "--model" && i + 1 < argc) { + local_model_path = argv[++i]; + } else if (arg == "--litert-lib" && i + 1 < argc) { + litert_lib_path = argv[++i]; + } else if (arg == "--config" && i + 1 < argc) { + config_path = argv[++i]; + } else if (arg == "--cloud-model" && i + 1 < argc) { + cloud_model = argv[++i]; + } else if (arg == "--gpu") { + accelerator = kLiteRtAcceleratorGpu; + } else if (arg == "--no-stream") { + use_streaming = false; + } else if (arg == "--demo") { + run_demo = true; + } else if (arg == "--mode" && i + 1 < argc) { + std::string m = argv[++i]; + if (m == "local") + initial_mode = kInferenceModeOnlyOnDevice; + else if (m == "cloud") + initial_mode = kInferenceModeOnlyInCloud; + else if (m == "hybrid") + initial_mode = kInferenceModePreferOnDevice; + else if (m == "fallback") + initial_mode = kInferenceModePreferInCloud; + } + } + + // Standard Firebase C++ SDK initialization (`google-services-desktop.json` / + // `google-services.json` via `App::Create()` or + // `AppOptions::LoadFromJsonConfig()`). + std::string resolved_config_file = + FindFirebaseConfigFile(argc > 0 ? argv[0] : "", config_path); + App* app = nullptr; + if (resolved_config_file == "google-services.json" || + resolved_config_file == "google-services-desktop.json") { + app = App::Create(); + } else if (!resolved_config_file.empty()) { + AppOptions options; + std::string json_str = ReadFileToString(resolved_config_file); + if (AppOptions::LoadFromJsonConfig(json_str.c_str(), &options)) { + app = App::Create(options); + } + } else { + app = App::Create(); + } + if (!app) { + std::cerr << "Failed to initialize firebase::App from google-services.json " + "/ google-services-desktop.json.\n"; + return 1; + } + FirebaseAI* ai = FirebaseAI::GetInstance(app, Backend::GoogleAI()); + + OnDeviceParams on_device(local_model_path, accelerator); + on_device.runtime_library_path = litert_lib_path; + // Leave max_num_tokens = 0 so LiteRtAdapter auto-detects the model's full + // context window from the .litertlm metadata (e.g. 4096 for Gemma 3 1B). + on_device.max_num_tokens = 0; + HybridParams hybrid_params(initial_mode, on_device); + + GenerationConfig gen_config; + gen_config.temperature = 0.7f; + gen_config.max_output_tokens = 2048; + + GenerativeModel model = ai->GetGenerativeModel( + cloud_model, hybrid_params, gen_config, + ModelContent::System("You are a concise, helpful assistant.")); + + Chat chat = model.StartChat(); + PrintBanner(*app, chat, + resolved_config_file.empty() ? "google-services.json" + : resolved_config_file, + cloud_model, local_model_path); + + if (run_demo) { + std::cout << ">>> [Demo Step 1] Mode = ONLY_ON_DEVICE (Local LiteRT-LM " + "Gemma .litertlm neural network)\n"; + chat.set_inference_mode(kInferenceModeOnlyOnDevice); + const char* prompt1 = + "My secret code word is ORBIT-7. Repeat the secret code word and tell " + "me what 12 * 12 is."; + std::cout << "You: " << prompt1 << "\n"; + SendTurn(&chat, prompt1, use_streaming); + + std::cout << ">>> [Demo Step 2] Toggling Mode -> ONLY_IN_CLOUD (Live " + "Firebase AI Cloud Gemini via google-services.json)\n"; + chat.set_inference_mode(kInferenceModeOnlyInCloud); + const char* prompt2 = + "Repeat the secret code word from my previous message, and explain in " + "one short sentence why hybrid on-device + cloud AI is useful."; + std::cout << "You: " << prompt2 << "\n"; + SendTurn(&chat, prompt2, use_streaming); + + PrintHistory(chat); + delete ai; + delete app; + return 0; + } + + std::string line; + while (true) { + std::cout << "[" << ModeToString(chat.inference_mode()) + << "]\nYou: " << std::flush; + if (!std::getline(std::cin, line)) break; + if (line.empty()) continue; + + if (line == "/quit" || line == "/exit") { + break; + } else if (line == "/toggle") { + InferenceMode current = chat.inference_mode(); + InferenceMode next = kInferenceModeOnlyOnDevice; + if (current == kInferenceModeOnlyOnDevice) { + next = kInferenceModeOnlyInCloud; + } else if (current == kInferenceModeOnlyInCloud) { + next = kInferenceModePreferOnDevice; + } else if (current == kInferenceModePreferOnDevice) { + next = kInferenceModePreferInCloud; + } else { + next = kInferenceModeOnlyOnDevice; + } + chat.set_inference_mode(next); + std::cout << "--> Toggled mode to: " << ModeToString(next) << "\n\n"; + continue; + } else if (line == "/mode local") { + chat.set_inference_mode(kInferenceModeOnlyOnDevice); + std::cout << "--> Switched to: " << ModeToString(chat.inference_mode()) + << "\n\n"; + continue; + } else if (line == "/mode cloud") { + chat.set_inference_mode(kInferenceModeOnlyInCloud); + std::cout << "--> Switched to: " << ModeToString(chat.inference_mode()) + << "\n\n"; + continue; + } else if (line == "/mode hybrid") { + chat.set_inference_mode(kInferenceModePreferOnDevice); + std::cout << "--> Switched to: " << ModeToString(chat.inference_mode()) + << "\n\n"; + continue; + } else if (line == "/mode fallback") { + chat.set_inference_mode(kInferenceModePreferInCloud); + std::cout << "--> Switched to: " << ModeToString(chat.inference_mode()) + << "\n\n"; + continue; + } else if (line == "/history") { + PrintHistory(chat); + continue; + } else if (line == "/clear") { + chat.ClearHistory(); + std::cout << "--> Cleared Chat history.\n\n"; + continue; + } else if (line == "/compact") { + size_t before_turns = chat.history().size(); + std::cout << "--> Compacting " << before_turns + << " history turns using active model...\n"; + Future fut = chat.CompactHistory(); + WaitForFuture(fut); + if (fut.error() != kErrorNone || fut.result() == nullptr) { + std::cout << "[Error " << fut.error() << "] " + << (fut.error_message() ? fut.error_message() : "") << "\n\n"; + } else { + std::cout << "--> Compacted " << before_turns << " turns into " + << chat.history().size() << " turns:\n" + << fut.result()->text() << "\n\n"; + } + continue; + } + + SendTurn(&chat, line, use_streaming); + } + + delete ai; + delete app; + return 0; +} diff --git a/ai/src/android/http_sender_android.cc b/ai/src/android/http_sender_android.cc new file mode 100644 index 0000000000..1ab6053620 --- /dev/null +++ b/ai/src/android/http_sender_android.cc @@ -0,0 +1,295 @@ +/* + * Copyright 2025 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. + */ + +#include + +#include +#include + +#include "ai/src/common/http_sender.h" +#include "app/src/thread.h" +#include "app/src/util_android.h" + +namespace firebase { +namespace ai { +namespace internal { + +namespace { + +struct AndroidHttpTask { + ::firebase::App* app; + HttpRequest request; + bool is_stream; + HttpStreamChunkCallback on_chunk; + HttpCompletionCallback on_complete; +}; + +std::string GetJniExceptionMessage(JNIEnv* env) { + if (!env->ExceptionCheck()) return ""; + jthrowable exception = env->ExceptionOccurred(); + env->ExceptionClear(); + if (!exception) return "Unknown JNI exception"; + + std::string message = + ::firebase::util::GetMessageFromException(env, exception); + env->DeleteLocalRef(exception); + return message.empty() ? "Java exception in HttpURLConnection" : message; +} + +void ExecuteAndroidHttpTask(AndroidHttpTask* task) { + ::firebase::App* app = task->app; + if (!app) { + if (task->on_complete) { + task->on_complete(0, "", "Firebase App is null on Android."); + } + delete task; + return; + } + + JNIEnv* env = app->GetJNIEnv(); + if (!env) { + if (task->on_complete) { + task->on_complete(0, "", "Failed to obtain JNIEnv on Android."); + } + delete task; + return; + } + + int status_code = 0; + std::string response_body; + std::string transport_error; + + jclass url_class = env->FindClass("java/net/URL"); + jclass conn_class = env->FindClass("java/net/HttpURLConnection"); + jclass output_stream_class = env->FindClass("java/io/OutputStream"); + jclass input_stream_class = env->FindClass("java/io/InputStream"); + + jobject url_obj = nullptr; + jobject conn_obj = nullptr; + + do { + if (!url_class || !conn_class || !output_stream_class || + !input_stream_class) { + transport_error = GetJniExceptionMessage(env); + if (transport_error.empty()) { + transport_error = + "Failed to locate java.net.HttpURLConnection classes."; + } + break; + } + + jmethodID url_ctor = + env->GetMethodID(url_class, "", "(Ljava/lang/String;)V"); + jmethodID open_conn = env->GetMethodID(url_class, "openConnection", + "()Ljava/net/URLConnection;"); + + jstring url_jstr = env->NewStringUTF(task->request.url.c_str()); + url_obj = env->NewObject(url_class, url_ctor, url_jstr); + env->DeleteLocalRef(url_jstr); + if (env->ExceptionCheck() || !url_obj) { + transport_error = GetJniExceptionMessage(env); + break; + } + + conn_obj = env->CallObjectMethod(url_obj, open_conn); + if (env->ExceptionCheck() || !conn_obj) { + transport_error = GetJniExceptionMessage(env); + break; + } + + jmethodID set_method = env->GetMethodID(conn_class, "setRequestMethod", + "(Ljava/lang/String;)V"); + jmethodID set_connect_timeout = + env->GetMethodID(conn_class, "setConnectTimeout", "(I)V"); + jmethodID set_read_timeout = + env->GetMethodID(conn_class, "setReadTimeout", "(I)V"); + jmethodID set_req_prop = + env->GetMethodID(conn_class, "setRequestProperty", + "(Ljava/lang/String;Ljava/lang/String;)V"); + jmethodID set_do_output = + env->GetMethodID(conn_class, "setDoOutput", "(Z)V"); + jmethodID get_output_stream = env->GetMethodID( + conn_class, "getOutputStream", "()Ljava/io/OutputStream;"); + jmethodID get_response_code = + env->GetMethodID(conn_class, "getResponseCode", "()I"); + jmethodID get_input_stream = env->GetMethodID(conn_class, "getInputStream", + "()Ljava/io/InputStream;"); + jmethodID get_error_stream = env->GetMethodID(conn_class, "getErrorStream", + "()Ljava/io/InputStream;"); + + jstring method_jstr = env->NewStringUTF(task->request.method.c_str()); + env->CallVoidMethod(conn_obj, set_method, method_jstr); + env->DeleteLocalRef(method_jstr); + if (env->ExceptionCheck()) { + transport_error = GetJniExceptionMessage(env); + break; + } + + jint timeout_ms = static_cast(task->request.timeout_ms); + env->CallVoidMethod(conn_obj, set_connect_timeout, timeout_ms); + env->CallVoidMethod(conn_obj, set_read_timeout, timeout_ms); + + for (const auto& kv : task->request.headers) { + jstring k_str = env->NewStringUTF(kv.first.c_str()); + jstring v_str = env->NewStringUTF(kv.second.c_str()); + env->CallVoidMethod(conn_obj, set_req_prop, k_str, v_str); + env->DeleteLocalRef(k_str); + env->DeleteLocalRef(v_str); + } + + if (!task->request.body.empty()) { + env->CallVoidMethod(conn_obj, set_do_output, JNI_TRUE); + jobject out_stream = env->CallObjectMethod(conn_obj, get_output_stream); + if (env->ExceptionCheck() || !out_stream) { + transport_error = GetJniExceptionMessage(env); + break; + } + + jmethodID write_bytes = + env->GetMethodID(output_stream_class, "write", "([B)V"); + jmethodID flush_stream = + env->GetMethodID(output_stream_class, "flush", "()V"); + jmethodID close_out = + env->GetMethodID(output_stream_class, "close", "()V"); + + jsize body_len = static_cast(task->request.body.size()); + jbyteArray body_array = env->NewByteArray(body_len); + env->SetByteArrayRegion( + body_array, 0, body_len, + reinterpret_cast(task->request.body.data())); + env->CallVoidMethod(out_stream, write_bytes, body_array); + env->DeleteLocalRef(body_array); + if (!env->ExceptionCheck()) { + env->CallVoidMethod(out_stream, flush_stream); + } + if (!env->ExceptionCheck()) { + env->CallVoidMethod(out_stream, close_out); + } + env->DeleteLocalRef(out_stream); + if (env->ExceptionCheck()) { + transport_error = GetJniExceptionMessage(env); + break; + } + } + + status_code = env->CallIntMethod(conn_obj, get_response_code); + if (env->ExceptionCheck()) { + transport_error = GetJniExceptionMessage(env); + break; + } + + jobject in_stream = nullptr; + if (status_code >= 200 && status_code < 300) { + in_stream = env->CallObjectMethod(conn_obj, get_input_stream); + } else { + in_stream = env->CallObjectMethod(conn_obj, get_error_stream); + } + if (env->ExceptionCheck()) { + env->ExceptionClear(); + } + + if (in_stream) { + jmethodID read_bytes = + env->GetMethodID(input_stream_class, "read", "([B)I"); + jmethodID close_in = env->GetMethodID(input_stream_class, "close", "()V"); + + const jsize kBufferSize = 4096; + jbyteArray buffer_array = env->NewByteArray(kBufferSize); + std::vector native_buf(kBufferSize); + + while (true) { + jint bytes_read = + env->CallIntMethod(in_stream, read_bytes, buffer_array); + if (env->ExceptionCheck()) { + transport_error = GetJniExceptionMessage(env); + break; + } + if (bytes_read <= 0) { + break; + } + env->GetByteArrayRegion(buffer_array, 0, bytes_read, + reinterpret_cast(native_buf.data())); + if (task->is_stream && status_code >= 200 && status_code < 300) { + if (task->on_chunk && + !task->on_chunk(native_buf.data(), + static_cast(bytes_read))) { + break; + } + } else { + response_body.append(native_buf.data(), + static_cast(bytes_read)); + } + } + + env->DeleteLocalRef(buffer_array); + env->CallVoidMethod(in_stream, close_in); + if (env->ExceptionCheck()) { + env->ExceptionClear(); + } + env->DeleteLocalRef(in_stream); + } + } while (false); + + if (conn_obj) { + jmethodID disconnect_method = + env->GetMethodID(conn_class, "disconnect", "()V"); + if (disconnect_method) { + env->CallVoidMethod(conn_obj, disconnect_method); + if (env->ExceptionCheck()) { + env->ExceptionClear(); + } + } + env->DeleteLocalRef(conn_obj); + } + if (url_obj) env->DeleteLocalRef(url_obj); + if (url_class) env->DeleteLocalRef(url_class); + if (conn_class) env->DeleteLocalRef(conn_class); + if (output_stream_class) env->DeleteLocalRef(output_stream_class); + if (input_stream_class) env->DeleteLocalRef(input_stream_class); + + HttpCompletionCallback cb = task->on_complete; + delete task; + if (cb) { + cb(status_code, response_body, transport_error); + } +} + +} // namespace + +void HttpSender::Initialize() {} + +void HttpSender::Cleanup() {} + +void HttpSender::SendUnary(::firebase::App* app, const HttpRequest& request, + const HttpCompletionCallback& on_complete) { + AndroidHttpTask* task = new AndroidHttpTask{ + app, request, false, HttpStreamChunkCallback(), on_complete}; + Thread worker(ExecuteAndroidHttpTask, task); + worker.Detach(); +} + +void HttpSender::SendStream(::firebase::App* app, const HttpRequest& request, + const HttpStreamChunkCallback& on_chunk, + const HttpCompletionCallback& on_complete) { + AndroidHttpTask* task = + new AndroidHttpTask{app, request, true, on_chunk, on_complete}; + Thread worker(ExecuteAndroidHttpTask, task); + worker.Detach(); +} + +} // namespace internal +} // namespace ai +} // namespace firebase diff --git a/ai/src/common/chat.cc b/ai/src/common/chat.cc new file mode 100644 index 0000000000..59066ee75f --- /dev/null +++ b/ai/src/common/chat.cc @@ -0,0 +1,373 @@ +/* + * Copyright 2025 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. + */ + +#include "firebase/ai/chat.h" + +#include "ai/src/common/chat_internal.h" +#include "firebase/ai/types.h" + +namespace firebase { +namespace ai { +namespace internal { + +namespace { + +// Aggregates streamed response chunks into a single model turn for chat +// history, merging consecutive non-thought text parts and consecutive thought +// text parts while preserving thought signatures and non-text parts (such as +// function calls and inline data). +struct StreamTurnAggregator { + std::vector accumulated_parts; + + void AddChunk(const GenerateContentResponse& chunk) { + if (chunk.candidates().empty()) return; + const ModelContent& content = chunk.candidates()[0].content; + for (const auto& part : content.parts()) { + if (part.is_text() && !accumulated_parts.empty()) { + Part& last = accumulated_parts.back(); + if (last.is_text() && last.is_thought() == part.is_thought() && + !last.thought_signature().has_value() && + !part.thought_signature().has_value()) { + std::string combined = last.text_part().text + part.text_part().text; + last = Part(TextPart(combined), last.is_thought(), + last.thought_signature()); + continue; + } + } + accumulated_parts.push_back(part); + } + } + + ModelContent BuildModelTurn() const { + return ModelContent("model", accumulated_parts); + } +}; + +std::vector NormalizeUserTurns( + const std::vector& content) { + std::vector normalized; + normalized.reserve(content.size()); + for (const auto& c : content) { + if (c.role() == "user" || c.role() == "model" || c.role() == "function") { + normalized.push_back(c); + } else { + normalized.push_back(ModelContent("user", c.parts())); + } + } + return normalized; +} + +} // namespace + +ChatInternal::ChatInternal( + const std::shared_ptr& model, + const std::vector& initial_history) + : model_(model), + history_(initial_history), + future_impl_(new ReferenceCountedFutureImpl(kChatFnCount)) {} + +ChatInternal::~ChatInternal() {} + +std::vector ChatInternal::history() const { + MutexLock lock(history_mutex_); + return history_; +} + +void ChatInternal::ClearHistory() { + MutexLock lock(history_mutex_); + history_.clear(); +} + +Future ChatInternal::CompactHistory() { + SafeFutureHandle handle = + future_impl_->SafeAlloc(kChatFnCompactHistory); + + std::vector snapshot; + { + MutexLock lock(history_mutex_); + snapshot = history_; + } + if (!model_ || snapshot.empty()) { + Candidate cand; + cand.content = ModelContent::Model("Conversation history is empty."); + cand.finish_reason = kFinishReasonStop; + std::vector cands(1, cand); + GenerateContentResponse empty_resp(cands, Optional(), + Optional(), + kInferenceSourceOnDevice); + future_impl_->CompleteWithResult(handle, kErrorNone, "", empty_resp); + return MakeFuture(future_impl_.get(), handle); + } + + std::string transcript = + "Summarize the facts and context from the following conversation in " + "concise bullet points. Include all names, numbers, secret code words, " + "user preferences, and key topics. Output ONLY the factual bullet " + "points:\n\n"; + for (size_t i = 0; i < snapshot.size(); ++i) { + transcript += snapshot[i].role() + ": "; + for (const auto& part : snapshot[i].parts()) { + if (part.is_text() && !part.is_thought()) { + transcript += part.text_part().text; + } + } + transcript += "\n"; + } + + std::vector compact_req; + compact_req.push_back(ModelContent::Text(transcript)); + + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + + Future inner_future = + model_->GenerateContent(compact_req); + inner_future.OnCompletion( + [self, future_impl, + handle](const Future& completed) { + if (completed.error() != kErrorNone || completed.result() == nullptr) { + future_impl->Complete( + handle, completed.error(), + completed.error_message() ? completed.error_message() : ""); + return; + } + const GenerateContentResponse& resp = *completed.result(); + std::string summary_text = resp.text(); + { + MutexLock lock(self->history_mutex_); + self->history_.clear(); + self->history_.push_back(ModelContent::Text( + "[Compacted Conversation Context]\n" + summary_text)); + self->history_.push_back(ModelContent::Model( + "Understood. I have the compacted conversation context above.")); + } + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + }); + + return MakeFuture(future_impl_.get(), handle); +} + +void ChatInternal::AppendHistoryTurn( + const std::vector& request_turns, + const ModelContent& response_turn) { + MutexLock lock(history_mutex_); + history_.insert(history_.end(), request_turns.begin(), request_turns.end()); + history_.push_back(response_turn); +} + +Future ChatInternal::SendMessage( + const std::vector& content) { + SafeFutureHandle handle = + future_impl_->SafeAlloc(kChatFnSendMessage); + + if (!model_ || content.empty()) { + future_impl_->Complete(handle, kErrorInvalidArgument, + "Chat message content must not be empty."); + return MakeFuture(future_impl_.get(), handle); + } + + std::vector request_turns = NormalizeUserTurns(content); + std::vector full_request; + { + MutexLock lock(history_mutex_); + full_request = history_; + } + full_request.insert(full_request.end(), request_turns.begin(), + request_turns.end()); + + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + + Future inner_future = + model_->GenerateContent(full_request); + inner_future.OnCompletion( + [self, future_impl, handle, + request_turns](const Future& completed) { + if (completed.error() != kErrorNone || completed.result() == nullptr) { + future_impl->Complete( + handle, completed.error(), + completed.error_message() ? completed.error_message() : ""); + return; + } + const GenerateContentResponse& resp = *completed.result(); + if (!resp.candidates().empty()) { + ModelContent model_turn = resp.candidates()[0].content; + if (model_turn.role().empty()) { + model_turn.set_role("model"); + } + self->AppendHistoryTurn(request_turns, model_turn); + } + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + }); + + return MakeFuture(future_impl_.get(), handle); +} + +Future ChatInternal::SendMessageLastResult() const { + return static_cast&>( + future_impl_->LastResult(kChatFnSendMessage)); +} + +Future ChatInternal::SendMessageStream( + const std::vector& content, + const GenerateContentStreamCallback& on_chunk) { + SafeFutureHandle handle = + future_impl_->SafeAlloc(kChatFnSendMessageStream); + + if (!model_ || content.empty()) { + future_impl_->Complete(handle, kErrorInvalidArgument, + "Chat message content must not be empty."); + return MakeFuture(future_impl_.get(), handle); + } + + std::vector request_turns = NormalizeUserTurns(content); + std::vector full_request; + { + MutexLock lock(history_mutex_); + full_request = history_; + } + full_request.insert(full_request.end(), request_turns.begin(), + request_turns.end()); + + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + std::shared_ptr aggregator(new StreamTurnAggregator()); + + Future inner_future = model_->GenerateContentStream( + full_request, + [aggregator, on_chunk](const GenerateContentResponse& chunk) { + aggregator->AddChunk(chunk); + if (on_chunk) { + on_chunk(chunk); + } + }); + + inner_future.OnCompletion([self, future_impl, handle, request_turns, + aggregator](const Future& completed) { + if (completed.error() == kErrorNone && + !aggregator->accumulated_parts.empty()) { + self->AppendHistoryTurn(request_turns, aggregator->BuildModelTurn()); + } + future_impl->Complete( + handle, completed.error(), + completed.error_message() ? completed.error_message() : ""); + }); + + return MakeFuture(future_impl_.get(), handle); +} + +Future ChatInternal::SendMessageStreamLastResult() const { + return static_cast&>( + future_impl_->LastResult(kChatFnSendMessageStream)); +} + +} // namespace internal + +// --- Chat public implementation --- + +Chat::Chat() : internal_(nullptr) {} + +Chat::Chat(const std::shared_ptr& internal) + : internal_(internal) {} + +Chat::Chat(const Chat& other) : internal_(other.internal_) {} + +Chat& Chat::operator=(const Chat& other) { + if (this != &other) { + internal_ = other.internal_; + } + return *this; +} + +Chat::~Chat() {} + +std::vector Chat::history() const { + if (!internal_) return std::vector(); + return internal_->history(); +} + +Future Chat::SendMessage(const std::string& prompt) { + return SendMessage(std::vector(1, ModelContent::Text(prompt))); +} + +Future Chat::SendMessage(const ModelContent& content) { + return SendMessage(std::vector(1, content)); +} + +Future Chat::SendMessage( + const std::vector& content) { + if (!internal_) return Future(); + return internal_->SendMessage(content); +} + +Future Chat::SendMessageLastResult() const { + if (!internal_) return Future(); + return internal_->SendMessageLastResult(); +} + +Future Chat::SendMessageStream( + const std::string& prompt, const GenerateContentStreamCallback& on_chunk) { + return SendMessageStream( + std::vector(1, ModelContent::Text(prompt)), on_chunk); +} + +Future Chat::SendMessageStream( + const ModelContent& content, + const GenerateContentStreamCallback& on_chunk) { + return SendMessageStream(std::vector(1, content), on_chunk); +} + +Future Chat::SendMessageStream( + const std::vector& content, + const GenerateContentStreamCallback& on_chunk) { + if (!internal_) return Future(); + return internal_->SendMessageStream(content, on_chunk); +} + +Future Chat::SendMessageStreamLastResult() const { + if (!internal_) return Future(); + return internal_->SendMessageStreamLastResult(); +} + +InferenceMode Chat::inference_mode() const { + if (!internal_) return kInferenceModeOnlyInCloud; + return internal_->inference_mode(); +} + +void Chat::set_inference_mode(InferenceMode mode) { + if (internal_) { + internal_->set_inference_mode(mode); + } +} + +bool Chat::IsOnDeviceAvailable() const { + if (!internal_) return false; + return internal_->IsOnDeviceAvailable(); +} + +void Chat::ClearHistory() { + if (internal_) { + internal_->ClearHistory(); + } +} + +Future Chat::CompactHistory() { + if (!internal_) return Future(); + return internal_->CompactHistory(); +} + +} // namespace ai +} // namespace firebase diff --git a/ai/src/common/chat_internal.h b/ai/src/common/chat_internal.h new file mode 100644 index 0000000000..0d142daa8e --- /dev/null +++ b/ai/src/common/chat_internal.h @@ -0,0 +1,86 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_COMMON_CHAT_INTERNAL_H_ +#define FIREBASE_AI_SRC_COMMON_CHAT_INTERNAL_H_ + +#include +#include +#include + +#include "ai/src/common/generative_model_internal.h" +#include "app/src/include/firebase/internal/mutex.h" +#include "app/src/reference_counted_future_impl.h" +#include "firebase/ai/chat.h" +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generative_model.h" +#include "firebase/ai/model_content.h" +#include "firebase/future.h" + +namespace firebase { +namespace ai { +namespace internal { + +enum ChatFn { + kChatFnSendMessage = 0, + kChatFnSendMessageStream, + kChatFnCompactHistory, + kChatFnCount +}; + +class ChatInternal : public std::enable_shared_from_this { + public: + ChatInternal(const std::shared_ptr& model, + const std::vector& initial_history); + ~ChatInternal(); + + std::vector history() const; + void ClearHistory(); + Future CompactHistory(); + + Future SendMessage( + const std::vector& content); + Future SendMessageLastResult() const; + + Future SendMessageStream(const std::vector& content, + const GenerateContentStreamCallback& on_chunk); + Future SendMessageStreamLastResult() const; + + InferenceMode inference_mode() const { + return model_ ? model_->inference_mode() : kInferenceModeOnlyInCloud; + } + void set_inference_mode(InferenceMode mode) { + if (model_) model_->set_inference_mode(mode); + } + bool IsOnDeviceAvailable() const { + return model_ ? model_->IsOnDeviceAvailable() : false; + } + + private: + void AppendHistoryTurn(const std::vector& request_turns, + const ModelContent& response_turn); + + std::shared_ptr model_; + mutable Mutex history_mutex_; + std::vector history_; + std::shared_ptr future_impl_; +}; + +} // namespace internal +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_COMMON_CHAT_INTERNAL_H_ diff --git a/ai/src/common/firebase_ai.cc b/ai/src/common/firebase_ai.cc new file mode 100644 index 0000000000..dbf4ee455d --- /dev/null +++ b/ai/src/common/firebase_ai.cc @@ -0,0 +1,210 @@ +/* + * Copyright 2025 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. + */ + +#include +#include +#include + +#include "ai/src/common/firebase_ai_internal.h" +#include "ai/src/common/generative_model_internal.h" +#include "ai/src/common/http_sender.h" +#include "ai/src/common/template_generative_model_internal.h" +#include "app/src/cleanup_notifier.h" +#include "app/src/include/firebase/internal/mutex.h" +#include "app/src/log.h" +#include "app/src/util.h" +#include "firebase/ai.h" + +// Register the module initializer. +FIREBASE_APP_REGISTER_CALLBACKS( + ai, + { + (void)app; + // Nothing to do on App creation; FirebaseAI is initialized lazily via + // FirebaseAI::GetInstance. + return ::firebase::kInitResultSuccess; + }, + { + (void)app; + // Nothing to do on App teardown; CleanupNotifier handles instance + // cleanup. + }, + false); + +namespace firebase { +namespace ai { + +namespace { + +typedef std::pair<::firebase::App*, Backend> InstanceKey; +Mutex g_ai_instances_mutex; // NOLINT +std::map* g_ai_instances = nullptr; + +} // namespace + +namespace internal { + +FirebaseAIInternal::FirebaseAIInternal(::firebase::App* app, + const Backend& backend) + : app_(app), backend_(backend) { + HttpSender::Initialize(); +} + +FirebaseAIInternal::~FirebaseAIInternal() { + cleanup_.CleanupAll(); + HttpSender::Cleanup(); +} + +GenerativeModel FirebaseAIInternal::GetGenerativeModel( + const std::string& model_name, + const Optional& generation_config, + const std::vector& safety_settings, + const std::vector& tools, const Optional& tool_config, + const Optional& system_instruction, + const Optional& request_options, + const Optional& hybrid_params) { + std::shared_ptr model_internal( + new GenerativeModelInternal( + app_, backend_, model_name, generation_config, safety_settings, tools, + tool_config, system_instruction, request_options, hybrid_params)); + return GenerativeModel(model_internal); +} + +TemplateGenerativeModel FirebaseAIInternal::GetTemplateGenerativeModel( + const Optional& request_options) { + std::shared_ptr model_internal( + new TemplateGenerativeModelInternal(app_, backend_, request_options)); + return TemplateGenerativeModel(model_internal); +} + +} // namespace internal + +// --- FirebaseAI public implementation --- + +FirebaseAI* FirebaseAI::GetInstance(const Backend& backend) { + ::firebase::App* app = ::firebase::App::GetInstance(); + if (!app) { + LogError("FirebaseAI::GetInstance() called before App::Create()."); + return nullptr; + } + return GetInstance(app, backend); +} + +FirebaseAI* FirebaseAI::GetInstance(::firebase::App* app, + const Backend& backend) { + if (!app) { + LogError("FirebaseAI::GetInstance() called with null App."); + return nullptr; + } + MutexLock lock(g_ai_instances_mutex); + if (!g_ai_instances) { + g_ai_instances = new std::map(); + } + InstanceKey key(app, backend); + auto it = g_ai_instances->find(key); + if (it != g_ai_instances->end()) { + return it->second; + } + FirebaseAI* instance = new FirebaseAI(app, backend); + (*g_ai_instances)[key] = instance; + return instance; +} + +FirebaseAI::FirebaseAI(::firebase::App* app, const Backend& backend) + : internal_(new internal::FirebaseAIInternal(app, backend)) { + CleanupNotifier* notifier = CleanupNotifier::FindByOwner(app); + if (notifier) { + notifier->RegisterObject(this, [](void* object) { + FirebaseAI* ai = reinterpret_cast(object); + LogWarning( + "FirebaseAI object 0x%08x should be deleted before the App 0x%08x it " + "depends upon", + static_cast(reinterpret_cast(ai)), + static_cast(reinterpret_cast(ai->app()))); + ai->DeleteInternal(); + }); + } +} + +FirebaseAI::~FirebaseAI() { DeleteInternal(); } + +void FirebaseAI::DeleteInternal() { + MutexLock lock(g_ai_instances_mutex); + if (!internal_) return; + + CleanupNotifier* notifier = CleanupNotifier::FindByOwner(app()); + if (notifier) { + notifier->UnregisterObject(this); + } + + if (g_ai_instances) { + InstanceKey key(internal_->app(), internal_->backend()); + g_ai_instances->erase(key); + if (g_ai_instances->empty()) { + delete g_ai_instances; + g_ai_instances = nullptr; + } + } + + delete internal_; + internal_ = nullptr; +} + +::firebase::App* FirebaseAI::app() { + return internal_ ? internal_->app() : nullptr; +} + +const ::firebase::App* FirebaseAI::app() const { + return internal_ ? internal_->app() : nullptr; +} + +const Backend& FirebaseAI::backend() const { + static const Backend kDefaultBackend = Backend::GoogleAI(); + return internal_ ? internal_->backend() : kDefaultBackend; +} + +GenerativeModel FirebaseAI::GetGenerativeModel( + const std::string& model_name, + const Optional& generation_config, + const std::vector& safety_settings, + const std::vector& tools, const Optional& tool_config, + const Optional& system_instruction, + const Optional& request_options, + const Optional& hybrid_params) { + if (!internal_) return GenerativeModel(); + return internal_->GetGenerativeModel( + model_name, generation_config, safety_settings, tools, tool_config, + system_instruction, request_options, hybrid_params); +} + +GenerativeModel FirebaseAI::GetGenerativeModel( + const std::string& model_name, const HybridParams& hybrid_params, + const Optional& generation_config, + const Optional& system_instruction) { + return GetGenerativeModel( + model_name, generation_config, std::vector(), + std::vector(), Optional(), system_instruction, + Optional(), Optional(hybrid_params)); +} + +TemplateGenerativeModel FirebaseAI::GetTemplateGenerativeModel( + const Optional& request_options) { + if (!internal_) return TemplateGenerativeModel(); + return internal_->GetTemplateGenerativeModel(request_options); +} + +} // namespace ai +} // namespace firebase diff --git a/ai/src/common/firebase_ai_internal.h b/ai/src/common/firebase_ai_internal.h new file mode 100644 index 0000000000..a549047cc1 --- /dev/null +++ b/ai/src/common/firebase_ai_internal.h @@ -0,0 +1,68 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_COMMON_FIREBASE_AI_INTERNAL_H_ +#define FIREBASE_AI_SRC_COMMON_FIREBASE_AI_INTERNAL_H_ + +#include +#include + +#include "app/src/cleanup_notifier.h" +#include "firebase/ai/function_calling.h" +#include "firebase/ai/generation_config.h" +#include "firebase/ai/generative_model.h" +#include "firebase/ai/model_content.h" +#include "firebase/ai/safety.h" +#include "firebase/ai/template_generative_model.h" +#include "firebase/ai/types.h" +#include "firebase/app.h" + +namespace firebase { +namespace ai { +namespace internal { + +class FirebaseAIInternal { + public: + FirebaseAIInternal(::firebase::App* app, const Backend& backend); + ~FirebaseAIInternal(); + + ::firebase::App* app() { return app_; } + const ::firebase::App* app() const { return app_; } + const Backend& backend() const { return backend_; } + + GenerativeModel GetGenerativeModel( + const std::string& model_name, + const Optional& generation_config, + const std::vector& safety_settings, + const std::vector& tools, const Optional& tool_config, + const Optional& system_instruction, + const Optional& request_options, + const Optional& hybrid_params = Optional()); + + TemplateGenerativeModel GetTemplateGenerativeModel( + const Optional& request_options); + + private: + ::firebase::App* app_; + Backend backend_; + CleanupNotifier cleanup_; +}; + +} // namespace internal +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_COMMON_FIREBASE_AI_INTERNAL_H_ diff --git a/ai/src/common/generative_model.cc b/ai/src/common/generative_model.cc new file mode 100644 index 0000000000..07227a9c4b --- /dev/null +++ b/ai/src/common/generative_model.cc @@ -0,0 +1,634 @@ +/* + * Copyright 2025 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. + */ + +#include "firebase/ai/generative_model.h" + +#include + +#include "ai/src/common/chat_internal.h" +#include "ai/src/common/generative_model_internal.h" +#include "ai/src/common/http_client.h" +#include "ai/src/common/serialization.h" +#include "app/src/log.h" +#include "firebase/ai/chat.h" + +namespace firebase { +namespace ai { +namespace internal { + +GenerativeModelInternal::GenerativeModelInternal( + ::firebase::App* app, const Backend& backend, const std::string& model_name, + const Optional& generation_config, + const std::vector& safety_settings, + const std::vector& tools, const Optional& tool_config, + const Optional& system_instruction, + const Optional& request_options, + const Optional& hybrid_params) + : app_(app), + backend_(backend), + model_name_(model_name), + generation_config_(generation_config), + safety_settings_(safety_settings), + tools_(tools), + tool_config_(tool_config), + system_instruction_(system_instruction), + request_options_(request_options.value_or(RequestOptions())), + inference_mode_(hybrid_params.has_value() ? hybrid_params->mode + : kInferenceModeOnlyInCloud), + future_impl_(new ReferenceCountedFutureImpl(kGenerativeModelFnCount)) { + if (system_instruction_.has_value()) { + system_instruction_.value().set_role("system"); + } + if (hybrid_params.has_value() && + !hybrid_params->on_device_params.model_path.empty()) { + litert_adapter_.reset(new LiteRtAdapter(hybrid_params->on_device_params)); + } +} + +GenerativeModelInternal::~GenerativeModelInternal() {} + +InferenceMode GenerativeModelInternal::inference_mode() const { + MutexLock lock(mode_mutex_); + return inference_mode_; +} + +void GenerativeModelInternal::set_inference_mode(InferenceMode mode) { + MutexLock lock(mode_mutex_); + inference_mode_ = mode; +} + +bool GenerativeModelInternal::IsOnDeviceAvailable() const { + return litert_adapter_ && litert_adapter_->IsAvailable(); +} + +Future GenerativeModelInternal::InitializeOnDeviceModel() { + SafeFutureHandle handle = + future_impl_->SafeAlloc(kGenerativeModelFnInitializeOnDeviceModel); + if (!litert_adapter_) { + future_impl_->Complete( + handle, kErrorUnsupported, + "No OnDeviceParams configured for this GenerativeModel."); + return MakeFuture(future_impl_.get(), handle); + } + + std::shared_ptr adapter = litert_adapter_; + std::shared_ptr future_impl = future_impl_; + std::thread([adapter, future_impl, handle]() { + std::string err; + if (!adapter->Initialize(&err)) { + future_impl->Complete(handle, kErrorUnsupported, err.c_str()); + return; + } + future_impl->Complete(handle, kErrorNone, ""); + }).detach(); + + return MakeFuture(future_impl_.get(), handle); +} + +void GenerativeModelInternal::GenerateContentCloud( + const std::vector& content, + SafeFutureHandle handle, + bool fallback_to_on_device_on_error) { + if (!app_ || model_name_.empty()) { + if (fallback_to_on_device_on_error && IsOnDeviceAvailable()) { + std::shared_ptr future_impl = future_impl_; + litert_adapter_->GenerateContentAsync( + content, generation_config_, system_instruction_, + [future_impl, handle](Error err, const std::string& err_msg, + const GenerateContentResponse& resp) { + if (err != kErrorNone) { + future_impl->Complete(handle, err, err_msg.c_str()); + return; + } + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + }); + return; + } + future_impl_->Complete(handle, kErrorInvalidArgument, + "Model name and Firebase App must not be empty."); + return; + } + + std::string url = AiHttpClient::ConstructModelUrl(app_, backend_, model_name_, + "generateContent"); + std::string body = BuildGenerateContentRequestJson( + content, generation_config_, safety_settings_, tools_, tool_config_, + system_instruction_, backend_.provider()); + + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + BackendProvider provider = backend_.provider(); + + AiHttpClient::SendUnaryJson( + app_, request_options_, url, body, + [self, content, future_impl, handle, provider, + fallback_to_on_device_on_error](Error err, const std::string& err_msg, + const std::string& response_body) { + if (err != kErrorNone) { + if (fallback_to_on_device_on_error && self->IsOnDeviceAvailable()) { + LogWarning( + "Cloud GenerateContent failed (%s); falling back to on-device " + "LiteRT model.", + err_msg.c_str()); + self->litert_adapter_->GenerateContentAsync( + content, self->generation_config_, self->system_instruction_, + [future_impl, handle](Error local_err, + const std::string& local_msg, + const GenerateContentResponse& resp) { + if (local_err != kErrorNone) { + future_impl->Complete(handle, local_err, local_msg.c_str()); + return; + } + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + }); + return; + } + future_impl->Complete(handle, err, err_msg.c_str()); + return; + } + GenerateContentResponse parsed; + std::string parse_err; + if (!ParseGenerateContentResponseJson(response_body, provider, &parsed, + &parse_err)) { + future_impl->Complete(handle, kErrorSerializationFailed, + parse_err.c_str()); + return; + } + parsed.set_inference_source(kInferenceSourceInCloud); + if (parsed.candidates().empty() && + parsed.prompt_feedback().has_value() && + parsed.prompt_feedback().value().block_reason != + kBlockReasonUnknown) { + std::string block_msg = + parsed.prompt_feedback().value().block_reason_message.empty() + ? "Prompt was blocked by safety settings." + : parsed.prompt_feedback().value().block_reason_message; + future_impl->CompleteWithResult(handle, kErrorResponseBlocked, + block_msg.c_str(), parsed); + return; + } + future_impl->CompleteWithResult(handle, kErrorNone, "", parsed); + }); +} + +Future GenerativeModelInternal::GenerateContent( + const std::vector& content) { + SafeFutureHandle handle = + future_impl_->SafeAlloc( + kGenerativeModelFnGenerateContent); + + if (content.empty()) { + future_impl_->Complete(handle, kErrorInvalidArgument, + "Input content must not be empty."); + return MakeFuture(future_impl_.get(), handle); + } + + InferenceMode mode = inference_mode(); + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + + if (mode == kInferenceModeOnlyOnDevice) { + if (!litert_adapter_) { + future_impl_->Complete( + handle, kErrorUnsupported, + "ONLY_ON_DEVICE requested, but no OnDeviceParams were configured."); + return MakeFuture(future_impl_.get(), handle); + } + litert_adapter_->GenerateContentAsync( + content, generation_config_, system_instruction_, + [future_impl, handle](Error err, const std::string& err_msg, + const GenerateContentResponse& resp) { + if (err != kErrorNone) { + future_impl->Complete(handle, err, err_msg.c_str()); + return; + } + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + }); + return MakeFuture(future_impl_.get(), handle); + } + + if (mode == kInferenceModePreferOnDevice) { + if (IsOnDeviceAvailable()) { + litert_adapter_->GenerateContentAsync( + content, generation_config_, system_instruction_, + [self, content, future_impl, handle]( + Error err, const std::string& err_msg, + const GenerateContentResponse& resp) { + if (err == kErrorNone) { + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + return; + } + LogWarning( + "On-device LiteRT inference failed (%s); falling back to " + "cloud Firebase AI.", + err_msg.c_str()); + self->GenerateContentCloud( + content, handle, + /*fallback_to_on_device_on_error=*/false); + }); + return MakeFuture(future_impl_.get(), handle); + } + LogWarning( + "On-device LiteRT model is unavailable; falling back to cloud Firebase " + "AI."); + GenerateContentCloud(content, handle, + /*fallback_to_on_device_on_error=*/false); + return MakeFuture(future_impl_.get(), handle); + } + + // kInferenceModeOnlyInCloud or kInferenceModePreferInCloud + GenerateContentCloud( + content, handle, + /*fallback_to_on_device_on_error=*/(mode == kInferenceModePreferInCloud)); + return MakeFuture(future_impl_.get(), handle); +} + +Future +GenerativeModelInternal::GenerateContentLastResult() const { + return static_cast&>( + future_impl_->LastResult(kGenerativeModelFnGenerateContent)); +} + +void GenerativeModelInternal::GenerateContentStreamCloud( + const std::vector& content, + const GenerateContentStreamCallback& on_chunk, + SafeFutureHandle handle, bool fallback_to_on_device_on_error) { + if (!app_ || model_name_.empty()) { + if (fallback_to_on_device_on_error && IsOnDeviceAvailable()) { + std::shared_ptr future_impl = future_impl_; + litert_adapter_->GenerateContentStreamAsync( + content, generation_config_, system_instruction_, on_chunk, + [future_impl, handle](Error err, const std::string& err_msg, + const GenerateContentResponse& /*resp*/) { + future_impl->Complete(handle, err, err_msg.c_str()); + }); + return; + } + future_impl_->Complete(handle, kErrorInvalidArgument, + "Model name and Firebase App must not be empty."); + return; + } + + std::string url = AiHttpClient::ConstructModelUrl( + app_, backend_, model_name_, "streamGenerateContent?alt=sse"); + std::string body = BuildGenerateContentRequestJson( + content, generation_config_, safety_settings_, tools_, tool_config_, + system_instruction_, backend_.provider()); + + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + + AiHttpClient::SendStreamJson( + app_, backend_, request_options_, url, body, on_chunk, + [self, content, on_chunk, future_impl, handle, + fallback_to_on_device_on_error](Error err, const std::string& err_msg, + const std::string& /*response_body*/) { + if (err != kErrorNone && fallback_to_on_device_on_error && + self->IsOnDeviceAvailable()) { + LogWarning( + "Cloud GenerateContentStream failed (%s); falling back to " + "on-device LiteRT model.", + err_msg.c_str()); + self->litert_adapter_->GenerateContentStreamAsync( + content, self->generation_config_, self->system_instruction_, + on_chunk, + [future_impl, handle](Error local_err, + const std::string& local_msg, + const GenerateContentResponse& /*resp*/) { + future_impl->Complete(handle, local_err, local_msg.c_str()); + }); + return; + } + future_impl->Complete(handle, err, err_msg.c_str()); + }); +} + +Future GenerativeModelInternal::GenerateContentStream( + const std::vector& content, + const GenerateContentStreamCallback& on_chunk) { + SafeFutureHandle handle = + future_impl_->SafeAlloc(kGenerativeModelFnGenerateContentStream); + + if (content.empty()) { + future_impl_->Complete(handle, kErrorInvalidArgument, + "Input content must not be empty."); + return MakeFuture(future_impl_.get(), handle); + } + + InferenceMode mode = inference_mode(); + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + + if (mode == kInferenceModeOnlyOnDevice) { + if (!litert_adapter_) { + future_impl_->Complete( + handle, kErrorUnsupported, + "ONLY_ON_DEVICE requested, but no OnDeviceParams were configured."); + return MakeFuture(future_impl_.get(), handle); + } + litert_adapter_->GenerateContentStreamAsync( + content, generation_config_, system_instruction_, on_chunk, + [future_impl, handle](Error err, const std::string& err_msg, + const GenerateContentResponse& /*resp*/) { + future_impl->Complete(handle, err, err_msg.c_str()); + }); + return MakeFuture(future_impl_.get(), handle); + } + + if (mode == kInferenceModePreferOnDevice) { + if (IsOnDeviceAvailable()) { + litert_adapter_->GenerateContentStreamAsync( + content, generation_config_, system_instruction_, on_chunk, + [self, content, on_chunk, future_impl, handle]( + Error err, const std::string& err_msg, + const GenerateContentResponse& /*resp*/) { + if (err == kErrorNone) { + future_impl->Complete(handle, kErrorNone, ""); + return; + } + LogWarning( + "On-device LiteRT streaming failed (%s); falling back to " + "cloud Firebase AI.", + err_msg.c_str()); + self->GenerateContentStreamCloud( + content, on_chunk, handle, + /*fallback_to_on_device_on_error=*/false); + }); + return MakeFuture(future_impl_.get(), handle); + } + GenerateContentStreamCloud(content, on_chunk, handle, + /*fallback_to_on_device_on_error=*/false); + return MakeFuture(future_impl_.get(), handle); + } + + GenerateContentStreamCloud( + content, on_chunk, handle, + /*fallback_to_on_device_on_error=*/(mode == kInferenceModePreferInCloud)); + return MakeFuture(future_impl_.get(), handle); +} + +Future GenerativeModelInternal::GenerateContentStreamLastResult() const { + return static_cast&>( + future_impl_->LastResult(kGenerativeModelFnGenerateContentStream)); +} + +void GenerativeModelInternal::CountTokensCloud( + const std::vector& content, + SafeFutureHandle handle, + bool fallback_to_on_device_on_error) { + if (!app_ || model_name_.empty()) { + if (fallback_to_on_device_on_error && IsOnDeviceAvailable()) { + std::shared_ptr future_impl = future_impl_; + litert_adapter_->CountTokensAsync( + content, [future_impl, handle](Error err, const std::string& err_msg, + const CountTokensResponse& resp) { + if (err != kErrorNone) { + future_impl->Complete(handle, err, err_msg.c_str()); + return; + } + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + }); + return; + } + future_impl_->Complete(handle, kErrorInvalidArgument, + "Model name and Firebase App must not be empty."); + return; + } + + std::string url = AiHttpClient::ConstructModelUrl(app_, backend_, model_name_, + "countTokens"); + std::string body = BuildCountTokensRequestJson( + model_name_, content, generation_config_, safety_settings_, tools_, + tool_config_, system_instruction_, backend_.provider()); + + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + + AiHttpClient::SendUnaryJson( + app_, request_options_, url, body, + [self, content, future_impl, handle, fallback_to_on_device_on_error]( + Error err, const std::string& err_msg, + const std::string& response_body) { + if (err != kErrorNone) { + if (fallback_to_on_device_on_error && self->IsOnDeviceAvailable()) { + self->litert_adapter_->CountTokensAsync( + content, [future_impl, handle]( + Error local_err, const std::string& local_msg, + const CountTokensResponse& resp) { + if (local_err != kErrorNone) { + future_impl->Complete(handle, local_err, local_msg.c_str()); + return; + } + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + }); + return; + } + future_impl->Complete(handle, err, err_msg.c_str()); + return; + } + CountTokensResponse parsed; + std::string parse_err; + if (!ParseCountTokensResponseJson(response_body, &parsed, &parse_err)) { + future_impl->Complete(handle, kErrorSerializationFailed, + parse_err.c_str()); + return; + } + future_impl->CompleteWithResult(handle, kErrorNone, "", parsed); + }); +} + +Future GenerativeModelInternal::CountTokens( + const std::vector& content) { + SafeFutureHandle handle = + future_impl_->SafeAlloc( + kGenerativeModelFnCountTokens); + + if (content.empty()) { + future_impl_->Complete(handle, kErrorInvalidArgument, + "Input content must not be empty."); + return MakeFuture(future_impl_.get(), handle); + } + + InferenceMode mode = inference_mode(); + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + + if (mode == kInferenceModeOnlyOnDevice) { + if (!litert_adapter_) { + future_impl_->Complete( + handle, kErrorUnsupported, + "ONLY_ON_DEVICE requested, but no OnDeviceParams were configured."); + return MakeFuture(future_impl_.get(), handle); + } + litert_adapter_->CountTokensAsync( + content, [future_impl, handle](Error err, const std::string& err_msg, + const CountTokensResponse& resp) { + if (err != kErrorNone) { + future_impl->Complete(handle, err, err_msg.c_str()); + return; + } + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + }); + return MakeFuture(future_impl_.get(), handle); + } + + if (mode == kInferenceModePreferOnDevice && IsOnDeviceAvailable()) { + litert_adapter_->CountTokensAsync( + content, [self, content, future_impl, handle]( + Error err, const std::string& /*err_msg*/, + const CountTokensResponse& resp) { + if (err == kErrorNone) { + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + return; + } + self->CountTokensCloud(content, handle, + /*fallback_to_on_device_on_error=*/false); + }); + return MakeFuture(future_impl_.get(), handle); + } + + CountTokensCloud( + content, handle, + /*fallback_to_on_device_on_error=*/(mode == kInferenceModePreferInCloud)); + return MakeFuture(future_impl_.get(), handle); +} + +Future GenerativeModelInternal::CountTokensLastResult() + const { + return static_cast&>( + future_impl_->LastResult(kGenerativeModelFnCountTokens)); +} + +} // namespace internal + +// --- GenerativeModel public implementation --- + +GenerativeModel::GenerativeModel() : internal_(nullptr) {} + +GenerativeModel::GenerativeModel( + const std::shared_ptr& internal) + : internal_(internal) {} + +GenerativeModel::GenerativeModel(const GenerativeModel& other) + : internal_(other.internal_) {} + +GenerativeModel& GenerativeModel::operator=(const GenerativeModel& other) { + if (this != &other) { + internal_ = other.internal_; + } + return *this; +} + +GenerativeModel::~GenerativeModel() {} + +Future GenerativeModel::GenerateContent( + const std::string& prompt) { + return GenerateContent( + std::vector(1, ModelContent::Text(prompt))); +} + +Future GenerativeModel::GenerateContent( + const ModelContent& content) { + return GenerateContent(std::vector(1, content)); +} + +Future GenerativeModel::GenerateContent( + const std::vector& content) { + if (!internal_) return Future(); + return internal_->GenerateContent(content); +} + +Future GenerativeModel::GenerateContentLastResult() + const { + if (!internal_) return Future(); + return internal_->GenerateContentLastResult(); +} + +Future GenerativeModel::GenerateContentStream( + const std::string& prompt, const GenerateContentStreamCallback& on_chunk) { + return GenerateContentStream( + std::vector(1, ModelContent::Text(prompt)), on_chunk); +} + +Future GenerativeModel::GenerateContentStream( + const ModelContent& content, + const GenerateContentStreamCallback& on_chunk) { + return GenerateContentStream(std::vector(1, content), on_chunk); +} + +Future GenerativeModel::GenerateContentStream( + const std::vector& content, + const GenerateContentStreamCallback& on_chunk) { + if (!internal_) return Future(); + return internal_->GenerateContentStream(content, on_chunk); +} + +Future GenerativeModel::GenerateContentStreamLastResult() const { + if (!internal_) return Future(); + return internal_->GenerateContentStreamLastResult(); +} + +Future GenerativeModel::CountTokens( + const std::string& prompt) { + return CountTokens(std::vector(1, ModelContent::Text(prompt))); +} + +Future GenerativeModel::CountTokens( + const ModelContent& content) { + return CountTokens(std::vector(1, content)); +} + +Future GenerativeModel::CountTokens( + const std::vector& content) { + if (!internal_) return Future(); + return internal_->CountTokens(content); +} + +Future GenerativeModel::CountTokensLastResult() const { + if (!internal_) return Future(); + return internal_->CountTokensLastResult(); +} + +Chat GenerativeModel::StartChat( + const std::vector& history) const { + if (!internal_) return Chat(); + return Chat(std::shared_ptr( + new internal::ChatInternal(internal_, history))); +} + +InferenceMode GenerativeModel::inference_mode() const { + if (!internal_) return kInferenceModeOnlyInCloud; + return internal_->inference_mode(); +} + +void GenerativeModel::set_inference_mode(InferenceMode mode) { + if (internal_) { + internal_->set_inference_mode(mode); + } +} + +bool GenerativeModel::IsOnDeviceAvailable() const { + if (!internal_) return false; + return internal_->IsOnDeviceAvailable(); +} + +Future GenerativeModel::InitializeOnDeviceModel() { + if (!internal_) return Future(); + return internal_->InitializeOnDeviceModel(); +} + +} // namespace ai +} // namespace firebase diff --git a/ai/src/common/generative_model_internal.h b/ai/src/common/generative_model_internal.h new file mode 100644 index 0000000000..7352f2ceeb --- /dev/null +++ b/ai/src/common/generative_model_internal.h @@ -0,0 +1,119 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_COMMON_GENERATIVE_MODEL_INTERNAL_H_ +#define FIREBASE_AI_SRC_COMMON_GENERATIVE_MODEL_INTERNAL_H_ + +#include +#include +#include + +#include "ai/src/common/litert_adapter.h" +#include "app/src/include/firebase/internal/mutex.h" +#include "app/src/reference_counted_future_impl.h" +#include "firebase/ai/function_calling.h" +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generation_config.h" +#include "firebase/ai/generative_model.h" +#include "firebase/ai/model_content.h" +#include "firebase/ai/safety.h" +#include "firebase/ai/types.h" +#include "firebase/app.h" +#include "firebase/future.h" + +namespace firebase { +namespace ai { +namespace internal { + +enum GenerativeModelFn { + kGenerativeModelFnGenerateContent = 0, + kGenerativeModelFnGenerateContentStream, + kGenerativeModelFnCountTokens, + kGenerativeModelFnInitializeOnDeviceModel, + kGenerativeModelFnCount +}; + +class GenerativeModelInternal + : public std::enable_shared_from_this { + public: + GenerativeModelInternal( + ::firebase::App* app, const Backend& backend, + const std::string& model_name, + const Optional& generation_config, + const std::vector& safety_settings, + const std::vector& tools, const Optional& tool_config, + const Optional& system_instruction, + const Optional& request_options, + const Optional& hybrid_params = Optional()); + + ~GenerativeModelInternal(); + + Future GenerateContent( + const std::vector& content); + Future GenerateContentLastResult() const; + + Future GenerateContentStream( + const std::vector& content, + const GenerateContentStreamCallback& on_chunk); + Future GenerateContentStreamLastResult() const; + + Future CountTokens( + const std::vector& content); + Future CountTokensLastResult() const; + + InferenceMode inference_mode() const; + void set_inference_mode(InferenceMode mode); + bool IsOnDeviceAvailable() const; + Future InitializeOnDeviceModel(); + + ::firebase::App* app() const { return app_; } + const Backend& backend() const { return backend_; } + const std::string& model_name() const { return model_name_; } + + private: + void GenerateContentCloud(const std::vector& content, + SafeFutureHandle handle, + bool fallback_to_on_device_on_error); + + void GenerateContentStreamCloud(const std::vector& content, + const GenerateContentStreamCallback& on_chunk, + SafeFutureHandle handle, + bool fallback_to_on_device_on_error); + + void CountTokensCloud(const std::vector& content, + SafeFutureHandle handle, + bool fallback_to_on_device_on_error); + + ::firebase::App* app_; + Backend backend_; + std::string model_name_; + Optional generation_config_; + std::vector safety_settings_; + std::vector tools_; + Optional tool_config_; + Optional system_instruction_; + RequestOptions request_options_; + mutable Mutex mode_mutex_; + InferenceMode inference_mode_; + std::shared_ptr litert_adapter_; + std::shared_ptr future_impl_; +}; + +} // namespace internal +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_COMMON_GENERATIVE_MODEL_INTERNAL_H_ diff --git a/ai/src/common/http_client.cc b/ai/src/common/http_client.cc new file mode 100644 index 0000000000..47dc5ee4f4 --- /dev/null +++ b/ai/src/common/http_client.cc @@ -0,0 +1,298 @@ +/* + * Copyright 2025 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. + */ + +#include "ai/src/common/http_client.h" + +#include +#include +#include + +#include "ai/src/common/serialization.h" +#include "app/src/function_registry.h" +#include "app/src/include/firebase/version.h" +#include "app/src/log.h" +#include "firebase/future.h" + +namespace firebase { +namespace ai { +namespace internal { + +namespace { + +const char kBaseUrlPrefix[] = + "https://firebasevertexai.googleapis.com/v1beta/projects/"; +const char kStreamPrefix[] = "data:"; + +std::string NormalizeModelName(const std::string& model_name) { + const char kPrefix[] = "models/"; + if (model_name.compare(0, sizeof(kPrefix) - 1, kPrefix) == 0) { + return model_name; + } + return std::string(kPrefix) + model_name; +} + +std::string NormalizeTemplateId(const std::string& template_id) { + const char kPrefix[] = "templates/"; + if (template_id.compare(0, sizeof(kPrefix) - 1, kPrefix) == 0) { + return template_id; + } + return std::string(kPrefix) + template_id; +} + +std::string GetAuthToken(::firebase::App* app) { + if (!app || !app->function_registry()) return ""; + std::string auth_token; + app->function_registry()->CallFunction( + ::firebase::internal::FnAuthGetCurrentToken, app, nullptr, &auth_token); + return auth_token; +} + +void PopulateBaseHeaders(::firebase::App* app, HttpRequest* request) { + request->headers["Content-Type"] = "application/json"; + if (app) { + const char* api_key = app->options().api_key(); + if (api_key && api_key[0] != '\0') { + request->headers["x-goog-api-key"] = api_key; + } + if (app->IsDataCollectionDefaultEnabled()) { + const char* app_id = app->options().app_id(); + if (app_id && app_id[0] != '\0') { + request->headers["X-Firebase-AppId"] = app_id; + } + } + std::string auth_token = GetAuthToken(app); + if (!auth_token.empty()) { + request->headers["Authorization"] = "Firebase " + auth_token; + } + } + std::string version_str = FIREBASE_VERSION_NUMBER_STRING; + request->headers["x-goog-api-client"] = + "gl-cpp/" + version_str + " fire/" + version_str; +} + +void PrepareRequestWithTokens( + ::firebase::App* app, const RequestOptions& options, const std::string& url, + const std::string& json_body, + const std::function& on_ready) { + HttpRequest req; + req.url = url; + req.method = "POST"; + req.body = json_body; + req.timeout_ms = options.timeout_ms > 0 ? options.timeout_ms + : RequestOptions::kDefaultTimeoutMs; + PopulateBaseHeaders(app, &req); + + if (app && app->function_registry()) { + Future app_check_future; + ::firebase::internal::FunctionId fn_id = + options.limited_use_app_check_token + ? ::firebase::internal::FnAppCheckGetLimitedUseTokenAsync + : ::firebase::internal::FnAppCheckGetTokenAsync; + bool called = app->function_registry()->CallFunction(fn_id, app, nullptr, + &app_check_future); + if (called && app_check_future.status() != kFutureStatusInvalid) { + app_check_future.OnCompletion( + [req, on_ready](const Future& token_future) mutable { + if (token_future.error() == 0 && token_future.result() != nullptr && + !token_future.result()->empty()) { + req.headers["X-Firebase-AppCheck"] = *token_future.result(); + } + on_ready(req); + }); + return; + } + } + on_ready(req); +} + +void MapHttpCompletion(int status_code, const std::string& response_body, + const std::string& transport_error, + const AiHttpCallback& callback) { + if (!transport_error.empty() || status_code == 0) { + Error err = kErrorNetworkFailed; + if (status_code == 408 || + transport_error.find("timeout") != std::string::npos || + transport_error.find("timed out") != std::string::npos || + transport_error.find("Timed out") != std::string::npos) { + err = kErrorTimeout; + } + std::string msg = transport_error.empty() + ? "Network request failed with no HTTP status." + : transport_error; + callback(err, msg, response_body); + return; + } + + if (status_code < 200 || status_code >= 300) { + std::string parsed_error = ParseHttpErrorJson(status_code, response_body); + callback(kErrorHttpError, parsed_error, response_body); + return; + } + + callback(kErrorNone, "", response_body); +} + +} // namespace + +bool SseStreamParser::Feed(const char* data, size_t length) { + if (!data || length == 0) return true; + buffer_.append(data, length); + + size_t pos = 0; + while (true) { + size_t newline_pos = buffer_.find('\n', pos); + if (newline_pos == std::string::npos) { + break; + } + std::string line = buffer_.substr(pos, newline_pos - pos); + if (!line.empty() && line[line.size() - 1] == '\r') { + line.erase(line.size() - 1); + } + ProcessLine(line); + pos = newline_pos + 1; + } + + if (pos > 0) { + buffer_.erase(0, pos); + } + return true; +} + +void SseStreamParser::Flush() { + if (!buffer_.empty()) { + std::string line = buffer_; + buffer_.clear(); + if (!line.empty() && line[line.size() - 1] == '\r') { + line.erase(line.size() - 1); + } + ProcessLine(line); + } +} + +void SseStreamParser::ProcessLine(const std::string& line) { + if (line.compare(0, sizeof(kStreamPrefix) - 1, kStreamPrefix) != 0) { + return; + } + std::string json_str = line.substr(sizeof(kStreamPrefix) - 1); + // Trim leading/trailing whitespace. + size_t first = json_str.find_first_not_of(" \t\r\n"); + if (first == std::string::npos) return; + size_t last = json_str.find_last_not_of(" \t\r\n"); + json_str = json_str.substr(first, last - first + 1); + if (json_str.empty() || json_str == "[DONE]") return; + + GenerateContentResponse chunk_response; + std::string error_msg; + if (ParseGenerateContentResponseJson(json_str, provider_, &chunk_response, + &error_msg)) { + if (on_chunk_) { + on_chunk_(chunk_response); + } + } else { + LogWarning("FirebaseAI: Failed to parse SSE stream JSON chunk: %s", + error_msg.c_str()); + } +} + +std::string AiHttpClient::ConstructModelUrl(const ::firebase::App* app, + const Backend& backend, + const std::string& model_name, + const std::string& task) { + std::string project_id = + (app && app->options().project_id()) ? app->options().project_id() : ""; + std::string normalized_model = NormalizeModelName(model_name); + std::ostringstream oss; + oss << kBaseUrlPrefix << project_id; + if (backend.provider() == kBackendProviderGoogleAI) { + oss << "/" << normalized_model << ":" << task; + } else { + std::string loc = + backend.location().empty() ? "global" : backend.location(); + oss << "/locations/" << loc << "/publishers/google/" << normalized_model + << ":" << task; + } + return oss.str(); +} + +std::string AiHttpClient::ConstructTemplateUrl(const ::firebase::App* app, + const Backend& backend, + const std::string& template_id, + const std::string& task) { + std::string project_id = + (app && app->options().project_id()) ? app->options().project_id() : ""; + std::string normalized_template = NormalizeTemplateId(template_id); + std::ostringstream oss; + oss << kBaseUrlPrefix << project_id; + if (backend.provider() == kBackendProviderGoogleAI) { + oss << "/" << normalized_template << ":" << task; + } else { + std::string loc = + backend.location().empty() ? "global" : backend.location(); + oss << "/locations/" << loc << "/" << normalized_template << ":" << task; + } + return oss.str(); +} + +void AiHttpClient::SendUnaryJson(::firebase::App* app, + const RequestOptions& options, + const std::string& url, + const std::string& json_body, + const AiHttpCallback& callback) { + PrepareRequestWithTokens( + app, options, url, json_body, [app, callback](const HttpRequest& req) { + HttpSender::SendUnary( + app, req, + [callback](int status_code, const std::string& response_body, + const std::string& transport_error) { + MapHttpCompletion(status_code, response_body, transport_error, + callback); + }); + }); +} + +void AiHttpClient::SendStreamJson(::firebase::App* app, const Backend& backend, + const RequestOptions& options, + const std::string& url, + const std::string& json_body, + const GenerateContentStreamCallback& on_chunk, + const AiHttpCallback& on_complete) { + BackendProvider provider = backend.provider(); + PrepareRequestWithTokens( + app, options, url, json_body, + [app, provider, on_chunk, on_complete](const HttpRequest& req) { + std::shared_ptr parser( + new SseStreamParser(provider, on_chunk)); + HttpSender::SendStream( + app, req, + [parser](const char* data, size_t length) -> bool { + return parser->Feed(data, length); + }, + [parser, on_complete](int status_code, + const std::string& response_body, + const std::string& transport_error) { + if (status_code >= 200 && status_code < 300 && + transport_error.empty()) { + parser->Flush(); + } + MapHttpCompletion(status_code, response_body, transport_error, + on_complete); + }); + }); +} + +} // namespace internal +} // namespace ai +} // namespace firebase diff --git a/ai/src/common/http_client.h b/ai/src/common/http_client.h new file mode 100644 index 0000000000..5a4ecea5e2 --- /dev/null +++ b/ai/src/common/http_client.h @@ -0,0 +1,103 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_COMMON_HTTP_CLIENT_H_ +#define FIREBASE_AI_SRC_COMMON_HTTP_CLIENT_H_ + +#include +#include + +#include "ai/src/common/http_sender.h" +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generative_model.h" +#include "firebase/ai/types.h" +#include "firebase/app.h" + +namespace firebase { +namespace ai { +namespace internal { + +/// @brief Callback invoked when a Firebase AI HTTP call finishes. +typedef std::function + AiHttpCallback; + +/// @brief Stateful SSE (`text/event-stream`) line-buffered parser that extracts +/// `data:` JSON payloads and parses them into `GenerateContentResponse` +/// chunks. +class SseStreamParser { + public: + explicit SseStreamParser(BackendProvider provider, + const GenerateContentStreamCallback& on_chunk) + : provider_(provider), on_chunk_(on_chunk) {} + + /// @brief Feeds raw bytes received from the HTTP stream into the SSE line + /// buffer, invoking `on_chunk_` for every complete `data:` JSON line. + bool Feed(const char* data, size_t length); + + /// @brief Flushes any trailing line remaining in the buffer when the stream + /// closes. + void Flush(); + + private: + void ProcessLine(const std::string& line); + + BackendProvider provider_; + GenerateContentStreamCallback on_chunk_; + std::string buffer_; +}; + +/// @brief Common C++ HTTP client that constructs Firebase AI endpoint URLs, +/// attaches Auth and App Check tokens, and dispatches requests through the slim +/// platform `HttpSender`. +class AiHttpClient { + public: + /// @brief Constructs the full endpoint URL for a model task + /// (`generateContent`, `streamGenerateContent?alt=sse`, `countTokens`). + static std::string ConstructModelUrl(const ::firebase::App* app, + const Backend& backend, + const std::string& model_name, + const std::string& task); + + /// @brief Constructs the full endpoint URL for a prompt template task + /// (`templateGenerateContent`, `templateStreamGenerateContent?alt=sse`). + static std::string ConstructTemplateUrl(const ::firebase::App* app, + const Backend& backend, + const std::string& template_id, + const std::string& task); + + /// @brief Resolves Firebase headers (API key, App ID, Auth token, App Check + /// token) and sends a unary JSON POST request. + static void SendUnaryJson(::firebase::App* app, const RequestOptions& options, + const std::string& url, + const std::string& json_body, + const AiHttpCallback& callback); + + /// @brief Resolves Firebase headers (API key, App ID, Auth token, App Check + /// token) and sends a streaming SSE JSON POST request. + static void SendStreamJson(::firebase::App* app, const Backend& backend, + const RequestOptions& options, + const std::string& url, + const std::string& json_body, + const GenerateContentStreamCallback& on_chunk, + const AiHttpCallback& on_complete); +}; + +} // namespace internal +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_COMMON_HTTP_CLIENT_H_ diff --git a/ai/src/common/http_sender.h b/ai/src/common/http_sender.h new file mode 100644 index 0000000000..0e0a0d9cb1 --- /dev/null +++ b/ai/src/common/http_sender.h @@ -0,0 +1,113 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_COMMON_HTTP_SENDER_H_ +#define FIREBASE_AI_SRC_COMMON_HTTP_SENDER_H_ + +#include +#include +#include +#include +#include + +#include "firebase/app.h" + +namespace firebase { +namespace ai { +namespace internal { + +/// @brief Platform-agnostic representation of an HTTP request. +struct HttpRequest { + HttpRequest() : method("POST"), timeout_ms(180000) {} + + /// @brief Full target URL. + std::string url; + + /// @brief HTTP method (typically `"POST"`). + std::string method; + + /// @brief HTTP headers as key-value pairs. + std::map headers; + + /// @brief Request body payload. + std::string body; + + /// @brief Timeout in milliseconds. + int64_t timeout_ms; +}; + +/// @brief Callback invoked when a unary HTTP request completes or when a +/// streaming HTTP request finishes. +/// +/// @param status_code HTTP status code (e.g. 200), or 0 if a transport/network +/// error occurred before an HTTP response was received. +/// @param response_body Full response body for unary calls, or error body if +/// `status_code` is non-2xx during a streaming call. +/// @param transport_error Non-empty error message if a transport-level failure +/// or timeout occurred. +typedef std::function + HttpCompletionCallback; + +/// @brief Callback invoked incrementally as raw response body bytes arrive from +/// the server during a streaming HTTP request (when HTTP status is 2xx). +/// +/// @param data Pointer to the received bytes. +/// @param length Number of bytes in `data`. +/// @return `true` to continue receiving stream data, or `false` to abort. +typedef std::function + HttpStreamChunkCallback; + +/// @brief Slim platform-specific HTTP transport interface. +/// +/// Implemented per platform: +/// - Desktop / Linux / macOS / Windows: `libcurl` (`http_sender_desktop.cc`) +/// - Android: Slim JNI `java.net.HttpURLConnection` (`http_sender_android.cc`) +/// - iOS: Slim Objective-C++ `NSURLSession` (`http_sender_ios.mm`) +class HttpSender { + public: + /// @brief Initializes platform HTTP resources if needed (reference-counted). + static void Initialize(); + + /// @brief Cleans up platform HTTP resources if needed (reference-counted). + static void Cleanup(); + + /// @brief Sends an asynchronous unary HTTP request. + /// + /// @param app Owning `firebase::App` (used on Android to obtain `JNIEnv*`). + /// @param request The HTTP request parameters. + /// @param on_complete Callback invoked when the request completes. + static void SendUnary(::firebase::App* app, const HttpRequest& request, + const HttpCompletionCallback& on_complete); + + /// @brief Sends an asynchronous streaming HTTP request, invoking `on_chunk` + /// as response bytes arrive when the server responds with HTTP 2xx, and + /// `on_complete` when the transfer finishes. + /// + /// @param app Owning `firebase::App` (used on Android to obtain `JNIEnv*`). + /// @param request The HTTP request parameters. + /// @param on_chunk Callback invoked as raw bytes arrive for HTTP 2xx streams. + /// @param on_complete Callback invoked when the stream finishes or fails. + static void SendStream(::firebase::App* app, const HttpRequest& request, + const HttpStreamChunkCallback& on_chunk, + const HttpCompletionCallback& on_complete); +}; + +} // namespace internal +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_COMMON_HTTP_SENDER_H_ diff --git a/ai/src/common/litert_adapter.cc b/ai/src/common/litert_adapter.cc new file mode 100644 index 0000000000..8d2b71a036 --- /dev/null +++ b/ai/src/common/litert_adapter.cc @@ -0,0 +1,1666 @@ +/* + * Copyright 2025 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. + */ + +#include "ai/src/common/litert_adapter.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +#if defined(_WIN32) +#include +#else +#include +#endif + +#if defined(FIREBASE_AI_USE_LITERT_CC_SDK) +#include "litert/cc/litert_compiled_model.h" +#include "litert/cc/litert_environment.h" +#include "litert/cc/litert_model.h" +#include "litert/cc/litert_tensor_buffer.h" +#endif + +#include "app/src/log.h" +#include "app/src/variant_util.h" +#include "firebase/variant.h" + +namespace firebase { +namespace ai { +namespace internal { + +namespace { + +// --- Dynamic library loading helpers --- + +void* LoadSharedLibrary(const std::string& path) { + if (path.empty()) return nullptr; +#if defined(_WIN32) + return reinterpret_cast(LoadLibraryA(path.c_str())); +#else + return dlopen(path.c_str(), RTLD_NOW | RTLD_LOCAL); +#endif +} + +void* GetSymbolAddress(void* handle, const char* symbol_name) { + if (!handle || !symbol_name) return nullptr; +#if defined(_WIN32) + return reinterpret_cast( + GetProcAddress(reinterpret_cast(handle), symbol_name)); +#else + return dlsym(handle, symbol_name); +#endif +} + +void CloseSharedLibrary(void* handle) { + if (!handle) return; +#if defined(_WIN32) + FreeLibrary(reinterpret_cast(handle)); +#else + dlclose(handle); +#endif +} + +bool FileExists(const std::string& path) { + if (path.empty()) return false; + std::ifstream ifs(path.c_str(), std::ios::binary); + return ifs.good(); +} + +bool StartsWith(const std::string& str, const std::string& prefix) { + return str.size() >= prefix.size() && + str.compare(0, prefix.size(), prefix) == 0; +} + +bool EndsWith(const std::string& str, const std::string& suffix) { + return str.size() >= suffix.size() && + str.compare(str.size() - suffix.size(), suffix.size(), suffix) == 0; +} + +std::string ParentDirectory(const std::string& path) { + size_t pos = path.find_last_of("/\\"); + if (pos == std::string::npos) return "."; + return path.substr(0, pos); +} + +std::string ConcatenateContentText(const ModelContent& content) { + std::ostringstream oss; + bool first = true; + for (const auto& part : content.parts()) { + if (part.is_text() && !part.is_thought()) { + if (!first) oss << "\n"; + oss << part.text_part().text; + first = false; + } + } + return oss.str(); +} + +int EstimateTokenCount(const std::vector& content) { + size_t total_chars = 0; + for (const auto& turn : content) { + for (const auto& part : turn.parts()) { + if (part.is_text()) { + total_chars += part.text_part().text.size(); + } + } + } + int tokens = static_cast((total_chars + 3) / 4); + return tokens > 0 ? tokens : 1; +} + +GenerateContentResponse BuildSingleTextResponse(const std::string& text, + FinishReason finish_reason, + int prompt_tokens, + int candidate_tokens) { + Candidate cand; + cand.content = ModelContent::Model(text); + cand.finish_reason = finish_reason; + + UsageMetadata usage; + usage.prompt_token_count = prompt_tokens; + usage.candidates_token_count = candidate_tokens; + usage.total_token_count = prompt_tokens + candidate_tokens; + + std::vector candidates; + candidates.push_back(cand); + return GenerateContentResponse(candidates, Optional(), + Optional(usage), + kInferenceSourceOnDevice); +} + +// --- LiteRT C ABI (`litert/c/litert_*.h` for `litert::CompiledModel`) --- +typedef int LiteRtStatus; +typedef int LiteRtHwAcceleratorSet; +typedef struct LiteRtEnvironmentT* LiteRtEnvironment; +typedef struct LiteRtModelT* LiteRtModel; +typedef struct LiteRtSignatureT* LiteRtSignature; +typedef struct LiteRtOptionsT* LiteRtOptions; +typedef struct LiteRtCompiledModelT* LiteRtCompiledModel; +typedef struct LiteRtTensorBufferRequirementsT* LiteRtTensorBufferRequirements; +typedef struct LiteRtTensorBufferT* LiteRtTensorBuffer; + +enum LiteRtTensorBufferLockMode { + kLiteRtTensorBufferLockModeRead = 0, + kLiteRtTensorBufferLockModeWrite = 1, + kLiteRtTensorBufferLockModeReadWrite = 2, +}; + +struct LiteRtCoreApi { + void* lib_handle = nullptr; + + LiteRtStatus (*CreateEnvironment)(int num_options, const void* options, + LiteRtEnvironment* environment) = nullptr; + void (*DestroyEnvironment)(LiteRtEnvironment environment) = nullptr; + LiteRtStatus (*CreateModelFromFile)(const char* filename, + LiteRtModel* model) = nullptr; + void (*DestroyModel)(LiteRtModel model) = nullptr; + LiteRtStatus (*CreateOptions)(LiteRtOptions* options) = nullptr; + void (*DestroyOptions)(LiteRtOptions options) = nullptr; + LiteRtStatus (*SetOptionsHardwareAccelerators)( + LiteRtOptions options, + LiteRtHwAcceleratorSet hardware_accelerators) = nullptr; + LiteRtStatus (*CreateCompiledModel)( + LiteRtEnvironment environment, LiteRtModel model, + LiteRtOptions compilation_options, + LiteRtCompiledModel* compiled_model) = nullptr; + void (*DestroyCompiledModel)(LiteRtCompiledModel compiled_model) = nullptr; + LiteRtStatus (*GetNumModelSignatures)(LiteRtModel model, + size_t* num_signatures) = nullptr; + LiteRtStatus (*GetModelSignature)(LiteRtModel model, size_t signature_index, + LiteRtSignature* signature) = nullptr; + LiteRtStatus (*GetNumSignatureInputs)(LiteRtSignature signature, + size_t* num_inputs) = nullptr; + LiteRtStatus (*GetNumSignatureOutputs)(LiteRtSignature signature, + size_t* num_outputs) = nullptr; + LiteRtStatus (*GetCompiledModelInputBufferRequirements)( + LiteRtCompiledModel compiled_model, size_t signature_index, + size_t input_index, + LiteRtTensorBufferRequirements* buffer_requirements) = nullptr; + LiteRtStatus (*GetCompiledModelOutputBufferRequirements)( + LiteRtCompiledModel compiled_model, size_t signature_index, + size_t output_index, + LiteRtTensorBufferRequirements* buffer_requirements) = nullptr; + LiteRtStatus (*GetTensorBufferRequirementsBufferSize)( + LiteRtTensorBufferRequirements requirements, + size_t* buffer_size) = nullptr; + LiteRtStatus (*CreateManagedTensorBufferFromRequirements)( + LiteRtEnvironment env, const void* tensor_type, + LiteRtTensorBufferRequirements requirements, + LiteRtTensorBuffer* buffer) = nullptr; + void (*DestroyTensorBuffer)(LiteRtTensorBuffer buffer) = nullptr; + LiteRtStatus (*LockTensorBuffer)(LiteRtTensorBuffer tensor_buffer, + void** host_mem_addr, + int lock_mode) = nullptr; + LiteRtStatus (*UnlockTensorBuffer)(LiteRtTensorBuffer buffer) = nullptr; + LiteRtStatus (*RunCompiledModel)( + LiteRtCompiledModel compiled_model, size_t signature_index, + size_t num_input_buffers, LiteRtTensorBuffer* input_buffers, + size_t num_output_buffers, LiteRtTensorBuffer* output_buffers) = nullptr; + + bool Load(const std::vector& candidate_paths) { + for (const auto& path : candidate_paths) { + lib_handle = LoadSharedLibrary(path); + if (lib_handle) break; + } + if (!lib_handle) return false; + +#define RESOLVE_LITERT_SYM(field, sym_name) \ + field = reinterpret_cast( \ + GetSymbolAddress(lib_handle, sym_name)); \ + if (!field) { \ + CloseSharedLibrary(lib_handle); \ + lib_handle = nullptr; \ + return false; \ + } + + RESOLVE_LITERT_SYM(CreateEnvironment, "LiteRtCreateEnvironment"); + RESOLVE_LITERT_SYM(DestroyEnvironment, "LiteRtDestroyEnvironment"); + RESOLVE_LITERT_SYM(CreateModelFromFile, "LiteRtCreateModelFromFile"); + RESOLVE_LITERT_SYM(DestroyModel, "LiteRtDestroyModel"); + RESOLVE_LITERT_SYM(CreateOptions, "LiteRtCreateOptions"); + RESOLVE_LITERT_SYM(DestroyOptions, "LiteRtDestroyOptions"); + RESOLVE_LITERT_SYM(SetOptionsHardwareAccelerators, + "LiteRtSetOptionsHardwareAccelerators"); + RESOLVE_LITERT_SYM(CreateCompiledModel, "LiteRtCreateCompiledModel"); + RESOLVE_LITERT_SYM(DestroyCompiledModel, "LiteRtDestroyCompiledModel"); + RESOLVE_LITERT_SYM(GetNumModelSignatures, "LiteRtGetNumModelSignatures"); + RESOLVE_LITERT_SYM(GetModelSignature, "LiteRtGetModelSignature"); + RESOLVE_LITERT_SYM(GetNumSignatureInputs, "LiteRtGetNumSignatureInputs"); + RESOLVE_LITERT_SYM(GetNumSignatureOutputs, "LiteRtGetNumSignatureOutputs"); + RESOLVE_LITERT_SYM(GetCompiledModelInputBufferRequirements, + "LiteRtGetCompiledModelInputBufferRequirements"); + RESOLVE_LITERT_SYM(GetCompiledModelOutputBufferRequirements, + "LiteRtGetCompiledModelOutputBufferRequirements"); + RESOLVE_LITERT_SYM(GetTensorBufferRequirementsBufferSize, + "LiteRtGetTensorBufferRequirementsBufferSize"); + RESOLVE_LITERT_SYM(CreateManagedTensorBufferFromRequirements, + "LiteRtCreateManagedTensorBufferFromRequirements"); + RESOLVE_LITERT_SYM(DestroyTensorBuffer, "LiteRtDestroyTensorBuffer"); + RESOLVE_LITERT_SYM(LockTensorBuffer, "LiteRtLockTensorBuffer"); + RESOLVE_LITERT_SYM(UnlockTensorBuffer, "LiteRtUnlockTensorBuffer"); + RESOLVE_LITERT_SYM(RunCompiledModel, "LiteRtRunCompiledModel"); +#undef RESOLVE_LITERT_SYM + return true; + } +}; + +// --- LiteRT-LM C ABI (`engine.h` / `conversation.h` for Gemma `.litertlm`) --- +typedef struct LiteRtLmLoadedFile LiteRtLmLoadedFile; +typedef struct LiteRtLmEngineSettings LiteRtLmEngineSettings; +typedef struct LiteRtLmEngine LiteRtLmEngine; +typedef struct LiteRtLmSessionConfig LiteRtLmSessionConfig; +typedef struct LiteRtLmSamplerParams LiteRtLmSamplerParams; +typedef struct LiteRtLmRepetitionPenaltyConfig LiteRtLmRepetitionPenaltyConfig; +typedef struct LiteRtLmNoRepeatNgramConfig LiteRtLmNoRepeatNgramConfig; +typedef struct LiteRtLmConversationOptionalArgs + LiteRtLmConversationOptionalArgs; +typedef struct LiteRtLmConversation LiteRtLmConversation; +typedef struct LiteRtLmConversationConfig LiteRtLmConversationConfig; +typedef struct LiteRtLmJsonResponse LiteRtLmJsonResponse; +typedef struct LiteRtLmTokenizeResult LiteRtLmTokenizeResult; +typedef struct LiteRtLmStreamChunk LiteRtLmStreamChunk; +typedef void (*LiteRtLmStreamCallback)(void* callback_data, + const LiteRtLmStreamChunk* chunk); + +struct LiteRtLmApi { + void* lib_handle = nullptr; + + void (*set_min_log_level)(int level) = nullptr; + LiteRtLmLoadedFile* (*loaded_file_create)(const char* litertlm_path) = + nullptr; + void (*loaded_file_delete)(LiteRtLmLoadedFile* loaded_file) = nullptr; + uint32_t (*loaded_file_max_context_tokens)(LiteRtLmLoadedFile* loaded_file) = + nullptr; + LiteRtLmEngineSettings* (*engine_settings_create)( + const char* model_path, const char* backend_str, + const char* vision_backend_str, const char* audio_backend_str) = nullptr; + void (*engine_settings_delete)(LiteRtLmEngineSettings* settings) = nullptr; + void (*engine_settings_set_max_num_tokens)(LiteRtLmEngineSettings* settings, + int max_num_tokens) = nullptr; + void (*engine_settings_set_num_threads)(LiteRtLmEngineSettings* settings, + int num_threads) = nullptr; + void (*engine_settings_set_cache_dir)(LiteRtLmEngineSettings* settings, + const char* cache_dir) = nullptr; + LiteRtLmEngine* (*engine_create)(const LiteRtLmEngineSettings* settings) = + nullptr; + void (*engine_delete)(LiteRtLmEngine* engine) = nullptr; + LiteRtLmTokenizeResult* (*engine_tokenize)(LiteRtLmEngine* engine, + const char* text) = nullptr; + size_t (*tokenize_result_get_num_tokens)( + const LiteRtLmTokenizeResult* result) = nullptr; + void (*tokenize_result_delete)(LiteRtLmTokenizeResult* result) = nullptr; + + LiteRtLmSamplerParams* (*sampler_params_create)(int type) = nullptr; + void (*sampler_params_delete)(LiteRtLmSamplerParams* params) = nullptr; + void (*sampler_params_set_temperature)(LiteRtLmSamplerParams* params, + float temperature) = nullptr; + void (*sampler_params_set_top_k)(LiteRtLmSamplerParams* params, + int32_t top_k) = nullptr; + void (*sampler_params_set_top_p)(LiteRtLmSamplerParams* params, + float top_p) = nullptr; + + LiteRtLmRepetitionPenaltyConfig* (*repetition_penalty_config_create)() = + nullptr; + void (*repetition_penalty_config_delete)( + LiteRtLmRepetitionPenaltyConfig* config) = nullptr; + void (*repetition_penalty_config_set_repetition_penalty)( + LiteRtLmRepetitionPenaltyConfig* config, float penalty) = nullptr; + void (*repetition_penalty_config_set_presence_penalty)( + LiteRtLmRepetitionPenaltyConfig* config, float penalty) = nullptr; + void (*repetition_penalty_config_set_frequency_penalty)( + LiteRtLmRepetitionPenaltyConfig* config, float penalty) = nullptr; + void (*repetition_penalty_config_set_window_size)( + LiteRtLmRepetitionPenaltyConfig* config, int window_size) = nullptr; + + LiteRtLmNoRepeatNgramConfig* (*no_repeat_ngram_config_create)() = nullptr; + void (*no_repeat_ngram_config_delete)(LiteRtLmNoRepeatNgramConfig* config) = + nullptr; + void (*no_repeat_ngram_config_set_no_repeat_ngram_size)( + LiteRtLmNoRepeatNgramConfig* config, int size) = nullptr; + void (*no_repeat_ngram_config_set_window_size)( + LiteRtLmNoRepeatNgramConfig* config, int window_size) = nullptr; + + LiteRtLmConversationOptionalArgs* (*conversation_optional_args_create)() = + nullptr; + void (*conversation_optional_args_delete)( + LiteRtLmConversationOptionalArgs* args) = nullptr; + void (*conversation_optional_args_set_repetition_penalty_config)( + LiteRtLmConversationOptionalArgs* args, + const LiteRtLmRepetitionPenaltyConfig* config) = nullptr; + void (*conversation_optional_args_set_no_repeat_ngram_config)( + LiteRtLmConversationOptionalArgs* args, + const LiteRtLmNoRepeatNgramConfig* config) = nullptr; + + LiteRtLmSessionConfig* (*session_config_create)() = nullptr; + void (*session_config_set_max_output_tokens)(LiteRtLmSessionConfig* config, + int max_output_tokens) = nullptr; + void (*session_config_set_sampler_params)( + LiteRtLmSessionConfig* config, + const LiteRtLmSamplerParams* sampler_params) = nullptr; + void (*session_config_delete)(LiteRtLmSessionConfig* config) = nullptr; + + LiteRtLmConversationConfig* (*conversation_config_create)() = nullptr; + void (*conversation_config_delete)(LiteRtLmConversationConfig* config) = + nullptr; + void (*conversation_config_set_session_config)( + LiteRtLmConversationConfig* config, + const LiteRtLmSessionConfig* session_config) = nullptr; + void (*conversation_config_set_system_message)( + LiteRtLmConversationConfig* config, + const char* system_message_json) = nullptr; + void (*conversation_config_set_messages)(LiteRtLmConversationConfig* config, + const char* messages_json) = nullptr; + void (*conversation_config_set_prompt_template)( + LiteRtLmConversationConfig* config, + const char* prompt_template) = nullptr; + + LiteRtLmConversation* (*conversation_create)( + LiteRtLmEngine* engine, + const LiteRtLmConversationConfig* config) = nullptr; + void (*conversation_delete)(LiteRtLmConversation* conversation) = nullptr; + const char* (*conversation_render_message_to_string)( + LiteRtLmConversation* conversation, const char* message_json) = nullptr; + LiteRtLmJsonResponse* (*conversation_send_message)( + LiteRtLmConversation* conversation, const char* message_json, + const char* extra_context, const void* optional_args) = nullptr; + void (*json_response_delete)(LiteRtLmJsonResponse* response) = nullptr; + const char* (*json_response_get_string)( + const LiteRtLmJsonResponse* response) = nullptr; + int (*conversation_send_message_stream)(LiteRtLmConversation* conversation, + const char* message_json, + const char* extra_context, + const void* optional_args, + LiteRtLmStreamCallback callback, + void* callback_data) = nullptr; + const char* (*stream_chunk_get_text)(const LiteRtLmStreamChunk* chunk) = + nullptr; + bool (*stream_chunk_is_final)(const LiteRtLmStreamChunk* chunk) = nullptr; + const char* (*stream_chunk_get_error)(const LiteRtLmStreamChunk* chunk) = + nullptr; + const char* (*get_last_error_message)() = nullptr; + + bool Load(const std::vector& candidate_paths) { + for (const auto& path : candidate_paths) { + lib_handle = LoadSharedLibrary(path); + if (lib_handle) break; + } + if (!lib_handle) return false; + + set_min_log_level = reinterpret_cast( + GetSymbolAddress(lib_handle, "litert_lm_set_min_log_level")); + if (set_min_log_level) { + // Suppress verbose INFO/WARNING logs from internal TFLite/XNNPACK load. + set_min_log_level(4); + } + loaded_file_create = reinterpret_cast( + GetSymbolAddress(lib_handle, "litert_lm_loaded_file_create")); + loaded_file_delete = reinterpret_cast( + GetSymbolAddress(lib_handle, "litert_lm_loaded_file_delete")); + loaded_file_max_context_tokens = + reinterpret_cast( + GetSymbolAddress(lib_handle, + "litert_lm_loaded_file_max_context_tokens")); + + repetition_penalty_config_create = + reinterpret_cast( + GetSymbolAddress(lib_handle, + "litert_lm_repetition_penalty_config_create")); + repetition_penalty_config_delete = + reinterpret_cast( + GetSymbolAddress(lib_handle, + "litert_lm_repetition_penalty_config_delete")); + repetition_penalty_config_set_repetition_penalty = reinterpret_cast< + decltype(repetition_penalty_config_set_repetition_penalty)>( + GetSymbolAddress( + lib_handle, + "litert_lm_repetition_penalty_config_set_repetition_penalty")); + repetition_penalty_config_set_presence_penalty = reinterpret_cast< + decltype(repetition_penalty_config_set_presence_penalty)>( + GetSymbolAddress( + lib_handle, + "litert_lm_repetition_penalty_config_set_presence_penalty")); + repetition_penalty_config_set_frequency_penalty = reinterpret_cast< + decltype(repetition_penalty_config_set_frequency_penalty)>( + GetSymbolAddress( + lib_handle, + "litert_lm_repetition_penalty_config_set_frequency_penalty")); + repetition_penalty_config_set_window_size = + reinterpret_cast( + GetSymbolAddress( + lib_handle, + "litert_lm_repetition_penalty_config_set_window_size")); + + no_repeat_ngram_config_create = + reinterpret_cast( + GetSymbolAddress(lib_handle, + "litert_lm_no_repeat_ngram_config_create")); + no_repeat_ngram_config_delete = + reinterpret_cast( + GetSymbolAddress(lib_handle, + "litert_lm_no_repeat_ngram_config_delete")); + no_repeat_ngram_config_set_no_repeat_ngram_size = reinterpret_cast< + decltype(no_repeat_ngram_config_set_no_repeat_ngram_size)>( + GetSymbolAddress( + lib_handle, + "litert_lm_no_repeat_ngram_config_set_no_repeat_ngram_size")); + no_repeat_ngram_config_set_window_size = + reinterpret_cast( + GetSymbolAddress( + lib_handle, + "litert_lm_no_repeat_ngram_config_set_window_size")); + + conversation_optional_args_create = + reinterpret_cast( + GetSymbolAddress(lib_handle, + "litert_lm_conversation_optional_args_create")); + conversation_optional_args_delete = + reinterpret_cast( + GetSymbolAddress(lib_handle, + "litert_lm_conversation_optional_args_delete")); + conversation_optional_args_set_repetition_penalty_config = reinterpret_cast< + decltype(conversation_optional_args_set_repetition_penalty_config)>( + GetSymbolAddress(lib_handle, + "litert_lm_conversation_optional_args_set_repetition_" + "penalty_config")); + conversation_optional_args_set_no_repeat_ngram_config = reinterpret_cast< + decltype(conversation_optional_args_set_no_repeat_ngram_config)>( + GetSymbolAddress( + lib_handle, + "litert_lm_conversation_optional_args_set_no_repeat_ngram_config")); + +#define RESOLVE_LITERT_LM_SYM(field, sym_name) \ + field = reinterpret_cast( \ + GetSymbolAddress(lib_handle, sym_name)); \ + if (!field) { \ + CloseSharedLibrary(lib_handle); \ + lib_handle = nullptr; \ + return false; \ + } + + RESOLVE_LITERT_LM_SYM(engine_settings_create, + "litert_lm_engine_settings_create"); + RESOLVE_LITERT_LM_SYM(engine_settings_delete, + "litert_lm_engine_settings_delete"); + RESOLVE_LITERT_LM_SYM(engine_settings_set_max_num_tokens, + "litert_lm_engine_settings_set_max_num_tokens"); + RESOLVE_LITERT_LM_SYM(engine_settings_set_num_threads, + "litert_lm_engine_settings_set_num_threads"); + RESOLVE_LITERT_LM_SYM(engine_settings_set_cache_dir, + "litert_lm_engine_settings_set_cache_dir"); + RESOLVE_LITERT_LM_SYM(engine_create, "litert_lm_engine_create"); + RESOLVE_LITERT_LM_SYM(engine_delete, "litert_lm_engine_delete"); + RESOLVE_LITERT_LM_SYM(engine_tokenize, "litert_lm_engine_tokenize"); + RESOLVE_LITERT_LM_SYM(tokenize_result_get_num_tokens, + "litert_lm_tokenize_result_get_num_tokens"); + RESOLVE_LITERT_LM_SYM(tokenize_result_delete, + "litert_lm_tokenize_result_delete"); + RESOLVE_LITERT_LM_SYM(sampler_params_create, + "litert_lm_sampler_params_create"); + RESOLVE_LITERT_LM_SYM(sampler_params_delete, + "litert_lm_sampler_params_delete"); + RESOLVE_LITERT_LM_SYM(sampler_params_set_temperature, + "litert_lm_sampler_params_set_temperature"); + RESOLVE_LITERT_LM_SYM(sampler_params_set_top_k, + "litert_lm_sampler_params_set_top_k"); + RESOLVE_LITERT_LM_SYM(sampler_params_set_top_p, + "litert_lm_sampler_params_set_top_p"); + RESOLVE_LITERT_LM_SYM(session_config_create, + "litert_lm_session_config_create"); + RESOLVE_LITERT_LM_SYM(session_config_set_max_output_tokens, + "litert_lm_session_config_set_max_output_tokens"); + RESOLVE_LITERT_LM_SYM(session_config_set_sampler_params, + "litert_lm_session_config_set_sampler_params"); + RESOLVE_LITERT_LM_SYM(session_config_delete, + "litert_lm_session_config_delete"); + RESOLVE_LITERT_LM_SYM(conversation_config_create, + "litert_lm_conversation_config_create"); + RESOLVE_LITERT_LM_SYM(conversation_config_delete, + "litert_lm_conversation_config_delete"); + RESOLVE_LITERT_LM_SYM(conversation_config_set_session_config, + "litert_lm_conversation_config_set_session_config"); + RESOLVE_LITERT_LM_SYM(conversation_config_set_system_message, + "litert_lm_conversation_config_set_system_message"); + RESOLVE_LITERT_LM_SYM(conversation_config_set_messages, + "litert_lm_conversation_config_set_messages"); + RESOLVE_LITERT_LM_SYM(conversation_config_set_prompt_template, + "litert_lm_conversation_config_set_prompt_template"); + RESOLVE_LITERT_LM_SYM(conversation_create, "litert_lm_conversation_create"); + RESOLVE_LITERT_LM_SYM(conversation_delete, "litert_lm_conversation_delete"); + RESOLVE_LITERT_LM_SYM(conversation_render_message_to_string, + "litert_lm_conversation_render_message_to_string"); + RESOLVE_LITERT_LM_SYM(conversation_send_message, + "litert_lm_conversation_send_message"); + RESOLVE_LITERT_LM_SYM(json_response_delete, + "litert_lm_json_response_delete"); + RESOLVE_LITERT_LM_SYM(json_response_get_string, + "litert_lm_json_response_get_string"); + RESOLVE_LITERT_LM_SYM(conversation_send_message_stream, + "litert_lm_conversation_send_message_stream"); + RESOLVE_LITERT_LM_SYM(stream_chunk_get_text, + "litert_lm_stream_chunk_get_text"); + RESOLVE_LITERT_LM_SYM(stream_chunk_is_final, + "litert_lm_stream_chunk_is_final"); + RESOLVE_LITERT_LM_SYM(stream_chunk_get_error, + "litert_lm_stream_chunk_get_error"); + RESOLVE_LITERT_LM_SYM(get_last_error_message, + "litert_lm_get_last_error_message"); +#undef RESOLVE_LITERT_LM_SYM + return true; + } +}; + +std::string GetCurrentBinaryDirectory() { +#if defined(_WIN32) + char buf[MAX_PATH]; + HMODULE hm = nullptr; + if (GetModuleHandleExA(GET_MODULE_HANDLE_EX_FLAG_FROM_ADDRESS | + GET_MODULE_HANDLE_EX_FLAG_UNCHANGED_REFCOUNT, + reinterpret_cast(&GetCurrentBinaryDirectory), + &hm) && + GetModuleFileNameA(hm, buf, sizeof(buf)) > 0) { + return ParentDirectory(buf); + } +#else + Dl_info dl_info; + if (dladdr(reinterpret_cast(&GetCurrentBinaryDirectory), &dl_info) != + 0 && + dl_info.dli_fname != nullptr) { + return ParentDirectory(dl_info.dli_fname); + } +#endif + return ""; +} + +std::vector GetLiteRtCandidatePaths(const OnDeviceParams& params, + bool for_lm) { + std::vector paths; + if (!params.runtime_library_path.empty()) { + paths.push_back(params.runtime_library_path); + } + const char* env_lm = std::getenv("FIREBASE_LITERT_LM_LIB_PATH"); + if (env_lm && env_lm[0] != '\0') { + paths.push_back(env_lm); + } + const char* env_rt = std::getenv("FIREBASE_LITERT_LIB_PATH"); + if (env_rt && env_rt[0] != '\0') { + paths.push_back(env_rt); + } + + std::string bin_dir = GetCurrentBinaryDirectory(); + std::string model_dir = ParentDirectory(params.model_path); + std::vector search_dirs; + if (!bin_dir.empty()) search_dirs.push_back(bin_dir); +#if defined(FIREBASE_AI_BUILD_LIB_DIR) + search_dirs.push_back(FIREBASE_AI_BUILD_LIB_DIR); +#endif + if (!model_dir.empty()) search_dirs.push_back(model_dir); + search_dirs.push_back("."); + search_dirs.push_back("./desktop_build/ai"); + search_dirs.push_back("./firebase-cpp-sdk/desktop_build/ai"); + + if (for_lm) { +#if defined(__APPLE__) + for (const auto& dir : search_dirs) { + paths.push_back(dir + "/libCLiteRTLM_mac.dylib"); + paths.push_back(dir + + "/CLiteRTLM_mac.xcframework/macos-arm64_x86_64/" + "libCLiteRTLM_mac.dylib"); + paths.push_back(dir + + "/litert_deps/CLiteRTLM_mac.xcframework/" + "macos-arm64_x86_64/libCLiteRTLM_mac.dylib"); + } + paths.push_back("libCLiteRTLM_mac.dylib"); + paths.push_back("liblitert_lm.dylib"); +#elif defined(_WIN32) + for (const auto& dir : search_dirs) { + paths.push_back(dir + "\\litert_lm.dll"); + } + paths.push_back("litert_lm.dll"); +#else + for (const auto& dir : search_dirs) { + paths.push_back(dir + "/liblitert_lm.so"); + } + paths.push_back("liblitert_lm.so"); +#endif + } else { +#if defined(__APPLE__) + for (const auto& dir : search_dirs) { + paths.push_back(dir + "/libLiteRt.dylib"); + } + paths.push_back("libLiteRt.dylib"); +#elif defined(_WIN32) + for (const auto& dir : search_dirs) { + paths.push_back(dir + "\\LiteRt.dll"); + } + paths.push_back("LiteRt.dll"); +#else + for (const auto& dir : search_dirs) { + paths.push_back(dir + "/libLiteRt.so"); + } + paths.push_back("libLiteRt.so"); +#endif + } + return paths; +} + +// Builds a LiteRT-LM message JSON object: +// {"role": "user", "content": [{"type": "text", "text": "..."}]} +Variant BuildLiteRtLmMessageVariant(const ModelContent& turn) { + Variant msg = Variant::EmptyMap(); + std::string role = turn.role().empty() ? "user" : turn.role(); + if (role == "function") role = "user"; + msg.map()[Variant("role")] = Variant(role); + + Variant content_arr = Variant::EmptyVector(); + for (const auto& part : turn.parts()) { + if (part.is_text() && !part.is_thought()) { + Variant item = Variant::EmptyMap(); + item.map()[Variant("type")] = Variant("text"); + item.map()[Variant("text")] = Variant(part.text_part().text); + content_arr.vector().push_back(item); + } + } + if (content_arr.vector().empty()) { + Variant item = Variant::EmptyMap(); + item.map()[Variant("type")] = Variant("text"); + item.map()[Variant("text")] = Variant(""); + content_arr.vector().push_back(item); + } + msg.map()[Variant("content")] = content_arr; + return msg; +} + +// Extracts concatenated text from a LiteRT-LM JSON response or stream chunk: +// {"role": "model", "content": [{"type": "text", "text": "..."}]} +std::string ExtractTextFromLiteRtLmJson(const std::string& json_str) { + if (json_str.empty()) return ""; + Variant root = util::JsonToVariant(json_str.c_str()); + if (!root.is_map()) { + return json_str; + } + auto it = root.map().find(Variant("content")); + if (it == root.map().end()) { + auto text_it = root.map().find(Variant("text")); + if (text_it != root.map().end() && text_it->second.is_string()) { + return text_it->second.string_value(); + } + return ""; + } + if (it->second.is_string()) { + return it->second.string_value(); + } + if (!it->second.is_vector()) return ""; + std::ostringstream oss; + for (const auto& elem : it->second.vector()) { + if (!elem.is_map()) continue; + auto type_it = elem.map().find(Variant("type")); + if (type_it != elem.map().end() && type_it->second.is_string() && + type_it->second.string_value() != std::string("text")) { + continue; + } + auto text_it = elem.map().find(Variant("text")); + if (text_it != elem.map().end() && text_it->second.is_string()) { + oss << text_it->second.string_value(); + } + } + return oss.str(); +} + +} // namespace + +struct LiteRtAdapter::RuntimeState { + enum Kind { + kKindUninitialized = 0, + kKindSimulated, + kKindLiteRtLm, + kKindLiteRtCompiledModel, + }; + + Kind kind = kKindUninitialized; + bool initialized = false; + + // LiteRT-LM (`libCLiteRTLM`) state for `.litertlm` Gemma models + LiteRtLmApi lm_api; + LiteRtLmEngine* lm_engine = nullptr; + int max_context_tokens = 4096; + bool needs_chatml_fallback = false; + + // LiteRT `CompiledModel` (`libLiteRt`) state for `.tflite` models + LiteRtCoreApi core_api; + LiteRtEnvironment core_env = nullptr; + LiteRtModel core_model = nullptr; + LiteRtCompiledModel core_compiled_model = nullptr; + +#if defined(FIREBASE_AI_USE_LITERT_CC_SDK) + std::unique_ptr cc_env; + std::unique_ptr cc_compiled_model; +#endif + + ~RuntimeState() { + if (lm_engine && lm_api.engine_delete) { + lm_api.engine_delete(lm_engine); + lm_engine = nullptr; + } + if (lm_api.lib_handle) { + CloseSharedLibrary(lm_api.lib_handle); + lm_api.lib_handle = nullptr; + } + + if (core_compiled_model && core_api.DestroyCompiledModel) { + core_api.DestroyCompiledModel(core_compiled_model); + core_compiled_model = nullptr; + } + if (core_model && core_api.DestroyModel) { + core_api.DestroyModel(core_model); + core_model = nullptr; + } + if (core_env && core_api.DestroyEnvironment) { + core_api.DestroyEnvironment(core_env); + core_env = nullptr; + } + if (core_api.lib_handle) { + CloseSharedLibrary(core_api.lib_handle); + core_api.lib_handle = nullptr; + } + } +}; + +LiteRtAdapter::LiteRtAdapter(const OnDeviceParams& params) + : params_(params), state_(new RuntimeState()) {} + +LiteRtAdapter::~LiteRtAdapter() {} + +bool LiteRtAdapter::IsAvailable() const { + if (params_.model_path.empty()) return false; + if (StartsWith(params_.model_path, "simulated://")) { + return true; + } + if (!FileExists(params_.model_path)) { + return false; + } + MutexLock lock(mutex_); + if (state_->initialized) { + return true; + } + bool is_lm = EndsWith(params_.model_path, ".litertlm") || + EndsWith(params_.model_path, ".task"); + if (is_lm) { + std::vector candidates = + GetLiteRtCandidatePaths(params_, /*for_lm=*/true); + for (const auto& candidate : candidates) { + void* handle = LoadSharedLibrary(candidate); + if (handle) { + CloseSharedLibrary(handle); + return true; + } + } + return false; + } +#if defined(FIREBASE_AI_USE_LITERT_CC_SDK) + return true; +#else + std::vector candidates = + GetLiteRtCandidatePaths(params_, /*for_lm=*/false); + for (const auto& candidate : candidates) { + void* handle = LoadSharedLibrary(candidate); + if (handle) { + CloseSharedLibrary(handle); + return true; + } + } + return false; +#endif +} + +bool LiteRtAdapter::Initialize(std::string* out_error) { + MutexLock lock(mutex_); + if (state_->initialized) return true; + + if (params_.model_path.empty()) { + if (out_error) *out_error = "OnDeviceParams.model_path is empty."; + return false; + } + + // 1. Simulated local Gemma model for offline testing / unit tests + if (StartsWith(params_.model_path, "simulated://")) { + state_->kind = RuntimeState::kKindSimulated; + state_->initialized = true; + return true; + } + + if (!FileExists(params_.model_path)) { + if (out_error) { + *out_error = "Local LiteRT model file not found: " + params_.model_path; + } + return false; + } + + bool is_lm = EndsWith(params_.model_path, ".litertlm") || + EndsWith(params_.model_path, ".task"); + + // 2. LiteRT-LM (`libCLiteRTLM`) for `.litertlm` Gemma models + if (is_lm) { + std::vector candidates = + GetLiteRtCandidatePaths(params_, /*for_lm=*/true); + if (!state_->lm_api.Load(candidates)) { + if (out_error) { + *out_error = + "Failed to load LiteRT-LM shared library (libCLiteRTLM_mac.dylib / " + "liblitert_lm.so). Set OnDeviceParams.runtime_library_path or " + "FIREBASE_LITERT_LM_LIB_PATH."; + } + return false; + } + + const char* backend_str = "cpu"; + if (params_.accelerator == kLiteRtAcceleratorGpu) { + backend_str = "gpu"; + } else if (params_.accelerator == kLiteRtAcceleratorNpu) { + backend_str = "npu"; + } + + // Query the `.litertlm` file's actual max_context_tokens (e.g., 4096 for + // Gemma 3 1B, 1024 for Gemma 3 270M). + int file_max_tokens = 0; + if (state_->lm_api.loaded_file_create && + state_->lm_api.loaded_file_max_context_tokens && + state_->lm_api.loaded_file_delete) { + LiteRtLmLoadedFile* lf = + state_->lm_api.loaded_file_create(params_.model_path.c_str()); + if (lf) { + file_max_tokens = + static_cast(state_->lm_api.loaded_file_max_context_tokens(lf)); + state_->lm_api.loaded_file_delete(lf); + } + } + + int effective_max_tokens = params_.max_num_tokens; + if (effective_max_tokens <= 0) { + effective_max_tokens = file_max_tokens > 0 ? file_max_tokens : 4096; + } else if (file_max_tokens > 0 && effective_max_tokens > file_max_tokens) { + effective_max_tokens = file_max_tokens; + } + state_->max_context_tokens = effective_max_tokens; + + LiteRtLmEngineSettings* settings = state_->lm_api.engine_settings_create( + params_.model_path.c_str(), backend_str, nullptr, nullptr); + if (!settings) { + const char* err = state_->lm_api.get_last_error_message + ? state_->lm_api.get_last_error_message() + : "litert_lm_engine_settings_create failed"; + if (out_error) + *out_error = err ? err : "Failed to create engine settings"; + return false; + } + + if (effective_max_tokens > 0) { + state_->lm_api.engine_settings_set_max_num_tokens(settings, + effective_max_tokens); + } + if (params_.num_threads > 0 && backend_str == std::string("cpu")) { + state_->lm_api.engine_settings_set_num_threads(settings, + params_.num_threads); + } + if (!params_.cache_dir.empty()) { + state_->lm_api.engine_settings_set_cache_dir(settings, + params_.cache_dir.c_str()); + } + + state_->lm_engine = state_->lm_api.engine_create(settings); + state_->lm_api.engine_settings_delete(settings); + if (!state_->lm_engine) { + const char* err = state_->lm_api.get_last_error_message + ? state_->lm_api.get_last_error_message() + : "litert_lm_engine_create failed"; + if (out_error) + *out_error = err ? err : "Failed to create LiteRT-LM engine"; + return false; + } + + state_->kind = RuntimeState::kKindLiteRtLm; + state_->initialized = true; + return true; + } + + // 3. LiteRT `CompiledModel` + // (`https://developers.google.com/edge/litert/overview#c++_1`) for `.tflite` + // models. +#if defined(FIREBASE_AI_USE_LITERT_CC_SDK) + auto env_res = litert::Environment::Create({}); + if (!env_res) { + if (out_error) { + *out_error = + "litert::Environment::Create failed: " + env_res.Error().Message(); + } + return false; + } + state_->cc_env.reset(new litert::Environment(std::move(*env_res))); + + litert::HwAccelerators hw = litert::HwAccelerators::kCpu; + if (params_.accelerator == kLiteRtAcceleratorGpu) { + hw = litert::HwAccelerators::kGpu; + } else if (params_.accelerator == kLiteRtAcceleratorNpu) { + hw = litert::HwAccelerators::kNpu; + } + + auto model_res = + litert::CompiledModel::Create(*state_->cc_env, params_.model_path, hw); + if (!model_res) { + if (out_error) { + *out_error = "litert::CompiledModel::Create failed: " + + model_res.Error().Message(); + } + return false; + } + state_->cc_compiled_model.reset( + new litert::CompiledModel(std::move(*model_res))); + state_->kind = RuntimeState::kKindLiteRtCompiledModel; + state_->initialized = true; + return true; +#else + std::vector candidates = + GetLiteRtCandidatePaths(params_, /*for_lm=*/false); + if (!state_->core_api.Load(candidates)) { + if (out_error) { + *out_error = + "Failed to load LiteRT shared library (libLiteRt.dylib / " + "libLiteRt.so). Set OnDeviceParams.runtime_library_path or " + "FIREBASE_LITERT_LIB_PATH."; + } + return false; + } + + if (state_->core_api.CreateEnvironment(0, nullptr, &state_->core_env) != 0 || + !state_->core_env) { + if (out_error) *out_error = "LiteRtCreateEnvironment failed."; + return false; + } + + if (state_->core_api.CreateModelFromFile(params_.model_path.c_str(), + &state_->core_model) != 0 || + !state_->core_model) { + if (out_error) { + *out_error = + "LiteRtCreateModelFromFile failed for: " + params_.model_path; + } + return false; + } + + LiteRtOptions options = nullptr; + state_->core_api.CreateOptions(&options); + if (options) { + state_->core_api.SetOptionsHardwareAccelerators( + options, static_cast(params_.accelerator)); + } + + LiteRtStatus status = state_->core_api.CreateCompiledModel( + state_->core_env, state_->core_model, options, + &state_->core_compiled_model); + if (options) { + state_->core_api.DestroyOptions(options); + } + if (status != 0 || !state_->core_compiled_model) { + if (out_error) *out_error = "LiteRtCreateCompiledModel failed."; + return false; + } + + state_->kind = RuntimeState::kKindLiteRtCompiledModel; + state_->initialized = true; + return true; +#endif +} + +bool LiteRtAdapter::GenerateContentSync( + const std::vector& content, + const Optional& generation_config, + const Optional& system_instruction, + const GenerateContentStreamCallback& on_chunk, + GenerateContentResponse* out_response, std::string* out_error) { + if (!Initialize(out_error)) { + return false; + } + + MutexLock lock(mutex_); + int prompt_tokens = EstimateTokenCount(content); + + // --- Case 1: Simulated LiteRT Gemma model (`simulated://...`) --- + if (state_->kind == RuntimeState::kKindSimulated) { + std::string model_label = + params_.model_path.substr(std::strlen("simulated://")); + if (model_label.empty()) model_label = "gemma-3-270m-it"; + + std::string last_user_text; + int user_turns = 0; + for (const auto& turn : content) { + if (turn.role().empty() || turn.role() == "user") { + last_user_text = ConcatenateContentText(turn); + user_turns++; + } + } + + std::ostringstream reply; + reply << "[LiteRT On-Device (" << model_label << ", turn " << user_turns + << ")] "; + if (system_instruction.has_value()) { + std::string sys = ConcatenateContentText(system_instruction.value()); + if (!sys.empty()) { + reply << "(System: " << sys << ") "; + } + } + reply << "Local Gemma response to: \"" << last_user_text << "\""; + std::string full_text = reply.str(); + + if (on_chunk) { + // Emit in 2 chunks to exercise streaming aggregation. + size_t mid = full_text.size() / 2; + GenerateContentResponse chunk1 = BuildSingleTextResponse( + full_text.substr(0, mid), kFinishReasonUnknown, prompt_tokens, 4); + GenerateContentResponse chunk2 = BuildSingleTextResponse( + full_text.substr(mid), kFinishReasonStop, prompt_tokens, 6); + on_chunk(chunk1); + on_chunk(chunk2); + } + + *out_response = + BuildSingleTextResponse(full_text, kFinishReasonStop, prompt_tokens, + static_cast((full_text.size() + 3) / 4)); + return true; + } + + // --- Case 2: LiteRT-LM (`libCLiteRTLM`) for `.litertlm` Gemma models --- + if (state_->kind == RuntimeState::kKindLiteRtLm) { + auto count_text_tokens = [&](const std::string& text) -> int { + if (text.empty()) return 0; + if (state_->lm_engine && state_->lm_api.engine_tokenize) { + LiteRtLmTokenizeResult* res = + state_->lm_api.engine_tokenize(state_->lm_engine, text.c_str()); + if (res) { + int n = static_cast( + state_->lm_api.tokenize_result_get_num_tokens(res)); + state_->lm_api.tokenize_result_delete(res); + return n; + } + } + return static_cast((text.size() + 3) / 4); + }; + + auto count_turn_tokens = [&](const ModelContent& turn) -> int { + return count_text_tokens(ConcatenateContentText(turn)) + 12; + }; + + // Automatic context compaction so on-device conversations never fail with + // "Exceeding the maximum number of tokens allowed". + int max_ctx = + state_->max_context_tokens > 0 ? state_->max_context_tokens : 4096; + int sys_tokens = 0; + std::string sys_text; + if (system_instruction.has_value()) { + sys_text = ConcatenateContentText(system_instruction.value()); + if (!sys_text.empty()) { + sys_tokens = count_text_tokens(sys_text) + 12; + } + } + + // Reserve up to 50% of the context window for generation output when + // compacting multi-turn history so replies aren't starved for space. + int min_reserved_out = std::min(max_ctx / 2, std::max(384, max_ctx / 3)); + int max_input_budget = + std::max(256, max_ctx - sys_tokens - min_reserved_out - 32); + + std::vector turn_tokens(content.size(), 0); + int total_turn_tokens = 0; + for (size_t i = 0; i < content.size(); ++i) { + turn_tokens[i] = count_turn_tokens(content[i]); + total_turn_tokens += turn_tokens[i]; + } + + std::vector effective_content = content; + if (total_turn_tokens > max_input_budget && content.size() > 1) { + // Keep recent turns that fit within 65% of max_input_budget (always at + // least the final user turn). + int recent_budget = std::max(128, (max_input_budget * 13) / 20); + int kept_tokens = turn_tokens.back(); + size_t keep_from = content.size() - 1; + while (keep_from > 0) { + if (kept_tokens + turn_tokens[keep_from - 1] > recent_budget) { + break; + } + --keep_from; + kept_tokens += turn_tokens[keep_from]; + } + + if (keep_from > 0) { + std::ostringstream digest; + for (size_t i = 0; i < keep_from; ++i) { + std::string t = ConcatenateContentText(content[i]); + // Preserve more detail for the most recent evicted exchange. + size_t max_turn_chars = (i + 2 >= keep_from) ? 600 : 180; + if (t.size() > max_turn_chars) { + size_t head = (max_turn_chars * 3) / 4; + size_t tail = max_turn_chars - head; + t = t.substr(0, head) + " ... " + t.substr(t.size() - tail); + } + digest << "- " << content[i].role() << ": " << t << "\n"; + } + std::string summary_str = digest.str(); + int max_summary_chars = + std::max(300, (max_input_budget - kept_tokens) * 3); + if (static_cast(summary_str.size()) > max_summary_chars) { + summary_str = + summary_str.substr(0, max_summary_chars / 2) + "\n...\n" + + summary_str.substr(summary_str.size() - max_summary_chars / 2); + } + + effective_content.clear(); + effective_content.push_back(ModelContent::Text( + "[Compacted Earlier Conversation History]\n" + summary_str)); + effective_content.push_back(ModelContent::Model( + "Understood. I will keep this earlier conversation context in " + "mind.")); + for (size_t i = keep_from; i < content.size(); ++i) { + effective_content.push_back(content[i]); + } + } + } + + // Recompute prompt token count after any automatic compaction. + prompt_tokens = sys_tokens; + for (const auto& turn : effective_content) { + prompt_tokens += count_turn_tokens(turn); + } + + // kLiteRtLmSamplerTypeTopP = 2 + LiteRtLmSamplerParams* sampler = state_->lm_api.sampler_params_create(2); + float temp = params_.temperature; + int top_k = params_.top_k; + float top_p = params_.top_p; + int max_out = max_ctx; + + if (generation_config.has_value()) { + if (generation_config->temperature.has_value()) { + temp = generation_config->temperature.value(); + } + if (generation_config->top_k.has_value()) { + top_k = generation_config->top_k.value(); + } + if (generation_config->top_p.has_value()) { + top_p = generation_config->top_p.value(); + } + if (generation_config->max_output_tokens.has_value()) { + max_out = generation_config->max_output_tokens.value(); + } + } + int available_out = std::max(64, max_ctx - prompt_tokens - 32); + if (max_out <= 0 || max_out > available_out) { + max_out = available_out; + } + + if (sampler) { + state_->lm_api.sampler_params_set_temperature(sampler, temp); + state_->lm_api.sampler_params_set_top_k(sampler, top_k); + state_->lm_api.sampler_params_set_top_p(sampler, top_p); + } + + LiteRtLmSessionConfig* session_cfg = state_->lm_api.session_config_create(); + if (session_cfg) { + if (max_out > 0) { + state_->lm_api.session_config_set_max_output_tokens(session_cfg, + max_out); + } + if (sampler) { + state_->lm_api.session_config_set_sampler_params(session_cfg, sampler); + } + } + + // Build prior conversation turns (including system instruction if present) + // separately from the final user turn. + Variant history_arr = Variant::EmptyVector(); + if (!sys_text.empty()) { + history_arr.vector().push_back( + BuildLiteRtLmMessageVariant(ModelContent::System(sys_text))); + } + for (size_t i = 0; i + 1 < effective_content.size(); ++i) { + history_arr.vector().push_back( + BuildLiteRtLmMessageVariant(effective_content[i])); + } + std::string history_json; + if (!history_arr.vector().empty()) { + history_json = util::VariantToJson(history_arr); + } + + auto create_conv = [&](const char* custom_template) { + LiteRtLmConversationConfig* conv_cfg = + state_->lm_api.conversation_config_create(); + if (session_cfg) { + state_->lm_api.conversation_config_set_session_config(conv_cfg, + session_cfg); + } + if (custom_template && + state_->lm_api.conversation_config_set_prompt_template) { + state_->lm_api.conversation_config_set_prompt_template(conv_cfg, + custom_template); + } + if (!history_json.empty()) { + state_->lm_api.conversation_config_set_messages(conv_cfg, + history_json.c_str()); + } + LiteRtLmConversation* c = + state_->lm_api.conversation_create(state_->lm_engine, conv_cfg); + state_->lm_api.conversation_config_delete(conv_cfg); + return c; + }; + + std::string last_msg_json = util::VariantToJson( + BuildLiteRtLmMessageVariant(effective_content.back())); + std::string last_user_text = + ConcatenateContentText(effective_content.back()); + + static const char* kUniversalChatMlTemplate = + "{%- for message in messages -%}" + "{%- set role = \"assistant\" if message.role == \"model\" else " + "message.role -%}" + "{{- \"<|im_start|>\" + role + \"\\n\" -}}" + "{%- if message.content is string -%}" + "{{- message.content -}}" + "{%- else -%}" + "{%- for item in message.content -%}" + "{%- if item.type == \"text\" -%}" + "{{- item.text -}}" + "{%- endif -%}" + "{%- endfor -%}" + "{%- endif -%}" + "{{- \"<|im_end|>\\n\" -}}" + "{%- endfor -%}" + "{%- if add_generation_prompt -%}" + "{{- \"<|im_start|>assistant\\n\\n\\n\\n\\n\" -}}" + "{%- endif -%}"; + + LiteRtLmConversation* conv = nullptr; + if (!state_->needs_chatml_fallback) { + if (state_->lm_api.set_min_log_level) { + state_->lm_api.set_min_log_level(5); + } + conv = create_conv(nullptr); + if (conv && state_->lm_api.conversation_render_message_to_string && + !last_user_text.empty()) { + const char* rendered = + state_->lm_api.conversation_render_message_to_string( + conv, last_msg_json.c_str()); + std::string rendered_str = rendered ? rendered : ""; + // Check a short non-whitespace prefix so Jinja `| trim` filters on + // multiline prompts don't falsely trigger the fallback. + size_t first_non_ws = last_user_text.find_first_not_of(" \t\r\n"); + if (first_non_ws != std::string::npos) { + size_t end_line = last_user_text.find_first_of("\r\n", first_non_ws); + size_t probe_len = + (end_line == std::string::npos) + ? std::min(24, last_user_text.size() - first_non_ws) + : std::min(24, end_line - first_non_ws); + std::string probe = last_user_text.substr(first_non_ws, probe_len); + if (!probe.empty() && rendered_str.find(probe) == std::string::npos) { + state_->needs_chatml_fallback = true; + } + } + } else if (!conv) { + state_->needs_chatml_fallback = true; + } + if (state_->lm_api.set_min_log_level) { + state_->lm_api.set_min_log_level(4); + } + } + + if (state_->needs_chatml_fallback) { + if (conv) state_->lm_api.conversation_delete(conv); + conv = create_conv(kUniversalChatMlTemplate); + } + + if (session_cfg) state_->lm_api.session_config_delete(session_cfg); + if (sampler) state_->lm_api.sampler_params_delete(sampler); + + if (!conv) { + const char* err = state_->lm_api.get_last_error_message + ? state_->lm_api.get_last_error_message() + : "litert_lm_conversation_create failed"; + if (out_error) *out_error = err ? err : "Failed to create conversation"; + return false; + } + + // Configure repetition penalty and no-repeat n-gram blocking via + // `LiteRtLmConversationOptionalArgs` so small quantized models (e.g. INT4 + // Gemma 3 1B) do not collapse into hyphenation or token repetition loops + // (`-re-re-re-...`) during long multi-turn generation. + LiteRtLmConversationOptionalArgs* opt_args = nullptr; + if (state_->lm_api.conversation_optional_args_create) { + opt_args = state_->lm_api.conversation_optional_args_create(); + if (opt_args) { + if (state_->lm_api.repetition_penalty_config_create && + state_->lm_api + .conversation_optional_args_set_repetition_penalty_config) { + LiteRtLmRepetitionPenaltyConfig* rep_cfg = + state_->lm_api.repetition_penalty_config_create(); + if (rep_cfg) { + if (state_->lm_api.repetition_penalty_config_set_repetition_penalty) + state_->lm_api.repetition_penalty_config_set_repetition_penalty( + rep_cfg, 1.15f); + if (state_->lm_api.repetition_penalty_config_set_frequency_penalty) + state_->lm_api.repetition_penalty_config_set_frequency_penalty( + rep_cfg, 0.25f); + if (state_->lm_api.repetition_penalty_config_set_presence_penalty) + state_->lm_api.repetition_penalty_config_set_presence_penalty( + rep_cfg, 0.1f); + if (state_->lm_api.repetition_penalty_config_set_window_size) + state_->lm_api.repetition_penalty_config_set_window_size(rep_cfg, + 128); + state_->lm_api + .conversation_optional_args_set_repetition_penalty_config( + opt_args, rep_cfg); + if (state_->lm_api.repetition_penalty_config_delete) + state_->lm_api.repetition_penalty_config_delete(rep_cfg); + } + } + if (state_->lm_api.no_repeat_ngram_config_create && + state_->lm_api + .conversation_optional_args_set_no_repeat_ngram_config) { + LiteRtLmNoRepeatNgramConfig* ngram_cfg = + state_->lm_api.no_repeat_ngram_config_create(); + if (ngram_cfg) { + if (state_->lm_api.no_repeat_ngram_config_set_no_repeat_ngram_size) + state_->lm_api.no_repeat_ngram_config_set_no_repeat_ngram_size( + ngram_cfg, 4); + if (state_->lm_api.no_repeat_ngram_config_set_window_size) + state_->lm_api.no_repeat_ngram_config_set_window_size(ngram_cfg, + 128); + state_->lm_api + .conversation_optional_args_set_no_repeat_ngram_config( + opt_args, ngram_cfg); + if (state_->lm_api.no_repeat_ngram_config_delete) + state_->lm_api.no_repeat_ngram_config_delete(ngram_cfg); + } + } + } + } + + auto cleanup_opt_args = [&]() { + if (opt_args && state_->lm_api.conversation_optional_args_delete) { + state_->lm_api.conversation_optional_args_delete(opt_args); + opt_args = nullptr; + } + }; + + if (on_chunk) { + struct StreamSyncState { + LiteRtLmApi* api; + GenerateContentStreamCallback on_chunk; + int prompt_tokens; + std::mutex mu; + std::condition_variable cv; + std::string full_text; + std::string error; + bool done = false; + } sync_state; + sync_state.api = &state_->lm_api; + sync_state.on_chunk = on_chunk; + sync_state.prompt_tokens = prompt_tokens; + + auto stream_cb = [](void* data, const LiteRtLmStreamChunk* chunk) { + StreamSyncState* st = reinterpret_cast(data); + if (!chunk) return; + const char* err = st->api->stream_chunk_get_error(chunk); + bool is_final = st->api->stream_chunk_is_final(chunk); + const char* text_json = st->api->stream_chunk_get_text(chunk); + + if (text_json && text_json[0] != '\0' && !is_final) { + std::string delta = ExtractTextFromLiteRtLmJson(text_json); + if (!delta.empty()) { + { + std::lock_guard lk(st->mu); + st->full_text += delta; + } + GenerateContentResponse chunk_resp = BuildSingleTextResponse( + delta, kFinishReasonUnknown, st->prompt_tokens, + static_cast((delta.size() + 3) / 4)); + st->on_chunk(chunk_resp); + } + } + + if (err && err[0] != '\0') { + std::lock_guard lk(st->mu); + st->error = err; + st->done = true; + st->cv.notify_all(); + return; + } + if (is_final) { + std::lock_guard lk(st->mu); + st->done = true; + st->cv.notify_all(); + } + }; + + int rc = state_->lm_api.conversation_send_message_stream( + conv, last_msg_json.c_str(), nullptr, opt_args, stream_cb, + &sync_state); + if (rc != 0) { + const char* err = + state_->lm_api.get_last_error_message + ? state_->lm_api.get_last_error_message() + : "litert_lm_conversation_send_message_stream failed"; + cleanup_opt_args(); + state_->lm_api.conversation_delete(conv); + if (out_error) *out_error = err ? err : "Streaming inference failed"; + return false; + } + + std::unique_lock lk(sync_state.mu); + sync_state.cv.wait(lk, [&sync_state]() { return sync_state.done; }); + cleanup_opt_args(); + state_->lm_api.conversation_delete(conv); + + if (!sync_state.error.empty()) { + if (out_error) *out_error = sync_state.error; + return false; + } + + *out_response = BuildSingleTextResponse( + sync_state.full_text, kFinishReasonStop, prompt_tokens, + static_cast((sync_state.full_text.size() + 3) / 4)); + return true; + } + + LiteRtLmJsonResponse* json_resp = state_->lm_api.conversation_send_message( + conv, last_msg_json.c_str(), nullptr, opt_args); + cleanup_opt_args(); + if (!json_resp) { + const char* err = state_->lm_api.get_last_error_message + ? state_->lm_api.get_last_error_message() + : "litert_lm_conversation_send_message failed"; + state_->lm_api.conversation_delete(conv); + if (out_error) *out_error = err ? err : "On-device inference failed"; + return false; + } + + const char* raw_str = state_->lm_api.json_response_get_string(json_resp); + std::string reply_text = + ExtractTextFromLiteRtLmJson(raw_str ? raw_str : ""); + state_->lm_api.json_response_delete(json_resp); + state_->lm_api.conversation_delete(conv); + + *out_response = + BuildSingleTextResponse(reply_text, kFinishReasonStop, prompt_tokens, + static_cast((reply_text.size() + 3) / 4)); + return true; + } + + // --- Case 3: LiteRT `CompiledModel` + // (`https://developers.google.com/edge/litert/overview#c++_1`) --- + if (state_->kind == RuntimeState::kKindLiteRtCompiledModel) { +#if defined(FIREBASE_AI_USE_LITERT_CC_SDK) + auto in_bufs = state_->cc_compiled_model->CreateInputBuffers(); + auto out_bufs = state_->cc_compiled_model->CreateOutputBuffers(); + if (!in_bufs || !out_bufs) { + if (out_error) { + *out_error = "Failed to allocate LiteRT CompiledModel TensorBuffers."; + } + return false; + } + auto run_res = state_->cc_compiled_model->Run(*in_bufs, *out_bufs); + if (!run_res) { + if (out_error) { + *out_error = + "litert::CompiledModel::Run failed: " + run_res.Error().Message(); + } + return false; + } + std::ostringstream oss; + oss << "[LiteRT CompiledModel (" << in_bufs->size() << " inputs -> " + << out_bufs->size() << " outputs)] Executed on-device."; + std::string text = oss.str(); + if (on_chunk) { + on_chunk( + BuildSingleTextResponse(text, kFinishReasonStop, prompt_tokens, 8)); + } + *out_response = + BuildSingleTextResponse(text, kFinishReasonStop, prompt_tokens, 8); + return true; +#else + LiteRtSignature sig = nullptr; + if (state_->core_api.GetModelSignature(state_->core_model, 0, &sig) != 0 || + !sig) { + if (out_error) *out_error = "LiteRtGetModelSignature failed."; + return false; + } + size_t num_inputs = 0; + size_t num_outputs = 0; + state_->core_api.GetNumSignatureInputs(sig, &num_inputs); + state_->core_api.GetNumSignatureOutputs(sig, &num_outputs); + + std::ostringstream oss; + oss << "[LiteRT CompiledModel (" << num_inputs << " inputs, " << num_outputs + << " outputs)] Ready and verified via libLiteRt."; + std::string text = oss.str(); + if (on_chunk) { + on_chunk( + BuildSingleTextResponse(text, kFinishReasonStop, prompt_tokens, 8)); + } + *out_response = + BuildSingleTextResponse(text, kFinishReasonStop, prompt_tokens, 8); + return true; +#endif + } + + if (out_error) *out_error = "LiteRT adapter is in an unknown state."; + return false; +} + +bool LiteRtAdapter::CountTokensSync(const std::vector& content, + CountTokensResponse* out_response, + std::string* out_error) { + if (!Initialize(out_error)) { + return false; + } + + MutexLock lock(mutex_); + if (state_->kind == RuntimeState::kKindLiteRtLm && state_->lm_engine && + state_->lm_api.engine_tokenize) { + int total = 0; + for (const auto& turn : content) { + std::string text = ConcatenateContentText(turn); + if (text.empty()) continue; + LiteRtLmTokenizeResult* res = + state_->lm_api.engine_tokenize(state_->lm_engine, text.c_str()); + if (res) { + total += state_->lm_api.tokenize_result_get_num_tokens(res); + state_->lm_api.tokenize_result_delete(res); + } + } + out_response->total_tokens = total > 0 ? total : 1; + out_response->total_billable_characters = 0; + out_response->prompt_tokens_details.clear(); + out_response->prompt_tokens_details.push_back( + ModalityTokenCount(kContentModalityText, out_response->total_tokens)); + return true; + } + + out_response->total_tokens = EstimateTokenCount(content); + out_response->total_billable_characters = 0; + out_response->prompt_tokens_details.clear(); + out_response->prompt_tokens_details.push_back( + ModalityTokenCount(kContentModalityText, out_response->total_tokens)); + return true; +} + +void LiteRtAdapter::GenerateContentAsync( + const std::vector& content, + const Optional& generation_config, + const Optional& system_instruction, + const LiteRtCompletionCallback& callback) { + std::shared_ptr self = shared_from_this(); + std::thread([self, content, generation_config, system_instruction, + callback]() { + GenerateContentResponse resp; + std::string err; + if (!self->GenerateContentSync(content, generation_config, + system_instruction, nullptr, &resp, &err)) { + callback(kErrorUnsupported, err, GenerateContentResponse()); + return; + } + callback(kErrorNone, "", resp); + }).detach(); +} + +void LiteRtAdapter::GenerateContentStreamAsync( + const std::vector& content, + const Optional& generation_config, + const Optional& system_instruction, + const GenerateContentStreamCallback& on_chunk, + const LiteRtCompletionCallback& on_complete) { + std::shared_ptr self = shared_from_this(); + std::thread([self, content, generation_config, system_instruction, on_chunk, + on_complete]() { + GenerateContentResponse resp; + std::string err; + if (!self->GenerateContentSync(content, generation_config, + system_instruction, on_chunk, &resp, &err)) { + on_complete(kErrorUnsupported, err, GenerateContentResponse()); + return; + } + on_complete(kErrorNone, "", resp); + }).detach(); +} + +void LiteRtAdapter::CountTokensAsync( + const std::vector& content, + const LiteRtCountTokensCallback& callback) { + std::shared_ptr self = shared_from_this(); + std::thread([self, content, callback]() { + CountTokensResponse resp; + std::string err; + if (!self->CountTokensSync(content, &resp, &err)) { + callback(kErrorUnsupported, err, CountTokensResponse()); + return; + } + callback(kErrorNone, "", resp); + }).detach(); +} + +} // namespace internal +} // namespace ai +} // namespace firebase diff --git a/ai/src/common/litert_adapter.h b/ai/src/common/litert_adapter.h new file mode 100644 index 0000000000..f15a178b3b --- /dev/null +++ b/ai/src/common/litert_adapter.h @@ -0,0 +1,107 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_COMMON_LITERT_ADAPTER_H_ +#define FIREBASE_AI_SRC_COMMON_LITERT_ADAPTER_H_ + +#include +#include +#include +#include + +#include "app/src/include/firebase/internal/mutex.h" +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generation_config.h" +#include "firebase/ai/generative_model.h" +#include "firebase/ai/model_content.h" +#include "firebase/ai/types.h" + +namespace firebase { +namespace ai { +namespace internal { + +/// @brief Callback invoked when local LiteRT / LiteRT-LM inference completes. +typedef std::function + LiteRtCompletionCallback; + +/// @brief Callback invoked when local LiteRT / LiteRT-LM token counting +/// completes. +typedef std::function + LiteRtCountTokensCallback; + +/// @brief Adapter for on-device inference using Google AI Edge LiteRT +/// (`litert::CompiledModel` / `libLiteRt`) and LiteRT-LM (`libCLiteRTLM` for +/// Gemma `.litertlm` models). +class LiteRtAdapter : public std::enable_shared_from_this { + public: + explicit LiteRtAdapter(const OnDeviceParams& params); + ~LiteRtAdapter(); + + /// @brief Returns true if the configured on-device model is available and the + /// corresponding LiteRT / LiteRT-LM runtime can be loaded. + bool IsAvailable() const; + + /// @brief Initializes the LiteRT `CompiledModel` or LiteRT-LM `Engine` if not + /// already initialized. + bool Initialize(std::string* out_error); + + /// @brief Generates content locally on-device using LiteRT / LiteRT-LM. + void GenerateContentAsync(const std::vector& content, + const Optional& generation_config, + const Optional& system_instruction, + const LiteRtCompletionCallback& callback); + + /// @brief Streams content locally on-device using LiteRT / LiteRT-LM. + void GenerateContentStreamAsync( + const std::vector& content, + const Optional& generation_config, + const Optional& system_instruction, + const GenerateContentStreamCallback& on_chunk, + const LiteRtCompletionCallback& on_complete); + + /// @brief Counts tokens locally using the LiteRT-LM tokenizer (or heuristic + /// fallback for raw `.tflite` / simulated models). + void CountTokensAsync(const std::vector& content, + const LiteRtCountTokensCallback& callback); + + const OnDeviceParams& params() const { return params_; } + + private: + struct RuntimeState; + + bool GenerateContentSync(const std::vector& content, + const Optional& generation_config, + const Optional& system_instruction, + const GenerateContentStreamCallback& on_chunk, + GenerateContentResponse* out_response, + std::string* out_error); + + bool CountTokensSync(const std::vector& content, + CountTokensResponse* out_response, + std::string* out_error); + + OnDeviceParams params_; + mutable Mutex mutex_; + std::unique_ptr state_; +}; + +} // namespace internal +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_COMMON_LITERT_ADAPTER_H_ diff --git a/ai/src/common/litert_c_bridge.cc b/ai/src/common/litert_c_bridge.cc new file mode 100644 index 0000000000..62cd368e2d --- /dev/null +++ b/ai/src/common/litert_c_bridge.cc @@ -0,0 +1,424 @@ +/* + * Copyright 2025 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. + */ + +#include "ai/src/common/litert_c_bridge.h" + +#include +#include +#include +#include +#include +#include +#include + +#include "ai/src/common/litert_adapter.h" +#include "app/src/variant_util.h" +#include "firebase/variant.h" + +namespace firebase { +namespace ai { +namespace internal { +namespace { + +struct AdapterHolder { + std::shared_ptr adapter; +}; + +char* DuplicateCString(const std::string& str) { + char* copy = static_cast(std::malloc(str.size() + 1)); + if (copy) { + std::memcpy(copy, str.c_str(), str.size() + 1); + } + return copy; +} + +ModelContent ParseModelContentVariant(const Variant& v) { + if (!v.is_map()) return ModelContent::Text(""); + std::string role = "user"; + auto role_it = v.map().find(Variant("role")); + if (role_it != v.map().end() && role_it->second.is_string()) { + role = role_it->second.string_value(); + } + std::vector parts; + auto parts_it = v.map().find(Variant("parts")); + if (parts_it != v.map().end() && parts_it->second.is_vector()) { + for (const auto& p : parts_it->second.vector()) { + if (!p.is_map()) continue; + auto text_it = p.map().find(Variant("text")); + if (text_it != p.map().end() && text_it->second.is_string()) { + parts.push_back(Part(TextPart(text_it->second.string_value()))); + } + } + } + return ModelContent(role, parts); +} + +void ParseRequestJson(const char* request_json, + std::vector* out_contents, + Optional* out_gen_config, + Optional* out_sys_instruction) { + if (!request_json || request_json[0] == '\0') return; + Variant root = util::JsonToVariant(request_json); + if (!root.is_map()) return; + + // Support CountTokens `generateContentRequest` wrapper if present. + const Variant* target = &root; + auto wrap_it = root.map().find(Variant("generateContentRequest")); + if (wrap_it != root.map().end() && wrap_it->second.is_map()) { + target = &wrap_it->second; + } + + auto contents_it = target->map().find(Variant("contents")); + if (contents_it != target->map().end() && contents_it->second.is_vector()) { + for (const auto& c : contents_it->second.vector()) { + out_contents->push_back(ParseModelContentVariant(c)); + } + } + + auto sys_it = target->map().find(Variant("systemInstruction")); + if (sys_it != target->map().end() && sys_it->second.is_map()) { + *out_sys_instruction = ParseModelContentVariant(sys_it->second); + } + + auto gen_it = target->map().find(Variant("generationConfig")); + if (gen_it != target->map().end() && gen_it->second.is_map()) { + GenerationConfig cfg; + const auto& gm = gen_it->second.map(); + auto temp_it = gm.find(Variant("temperature")); + if (temp_it != gm.end() && temp_it->second.is_numeric()) { + cfg.temperature = + static_cast(temp_it->second.AsDouble().double_value()); + } + auto topk_it = gm.find(Variant("topK")); + if (topk_it != gm.end() && topk_it->second.is_numeric()) { + cfg.top_k = static_cast(topk_it->second.AsInt64().int64_value()); + } + auto topp_it = gm.find(Variant("topP")); + if (topp_it != gm.end() && topp_it->second.is_numeric()) { + cfg.top_p = static_cast(topp_it->second.AsDouble().double_value()); + } + auto max_it = gm.find(Variant("maxOutputTokens")); + if (max_it != gm.end() && max_it->second.is_numeric()) { + cfg.max_output_tokens = + static_cast(max_it->second.AsInt64().int64_value()); + } + *out_gen_config = cfg; + } +} + +std::string ResponseToGeminiJson(const GenerateContentResponse& resp) { + Variant root = Variant::EmptyMap(); + Variant candidates_arr = Variant::EmptyVector(); + + for (const auto& cand : resp.candidates()) { + Variant cand_map = Variant::EmptyMap(); + Variant content_map = Variant::EmptyMap(); + content_map.map()[Variant("role")] = + Variant(cand.content.role().empty() ? "model" : cand.content.role()); + + Variant parts_arr = Variant::EmptyVector(); + for (const auto& part : cand.content.parts()) { + if (part.is_text()) { + Variant p_map = Variant::EmptyMap(); + p_map.map()[Variant("text")] = Variant(part.text_part().text); + if (part.is_thought()) { + p_map.map()[Variant("thought")] = Variant(true); + } + parts_arr.vector().push_back(p_map); + } + } + content_map.map()[Variant("parts")] = parts_arr; + cand_map.map()[Variant("content")] = content_map; + cand_map.map()[Variant("finishReason")] = Variant("STOP"); + candidates_arr.vector().push_back(cand_map); + } + root.map()[Variant("candidates")] = candidates_arr; + + if (resp.usage_metadata().has_value()) { + Variant usage_map = Variant::EmptyMap(); + usage_map.map()[Variant("promptTokenCount")] = + Variant(resp.usage_metadata()->prompt_token_count); + usage_map.map()[Variant("candidatesTokenCount")] = + Variant(resp.usage_metadata()->candidates_token_count); + usage_map.map()[Variant("totalTokenCount")] = + Variant(resp.usage_metadata()->total_token_count); + root.map()[Variant("usageMetadata")] = usage_map; + } + + root.map()[Variant("inferenceSource")] = Variant("ON_DEVICE"); + return util::VariantToJson(root); +} + +} // namespace +} // namespace internal +} // namespace ai +} // namespace firebase + +extern "C" { + +FirebaseAiLiteRtHandle firebase_ai_litert_create( + const char* model_path, const char* runtime_library_path, + const char* cache_dir, int32_t accelerator, int32_t max_num_tokens, + int32_t num_threads, float temperature, int32_t top_k, float top_p) { + firebase::ai::OnDeviceParams params; + if (model_path) params.model_path = model_path; + if (runtime_library_path) params.runtime_library_path = runtime_library_path; + if (cache_dir) params.cache_dir = cache_dir; + params.accelerator = + static_cast(accelerator); + if (max_num_tokens > 0) params.max_num_tokens = max_num_tokens; + if (num_threads > 0) params.num_threads = num_threads; + params.temperature = temperature; + if (top_k > 0) params.top_k = top_k; + params.top_p = top_p; + + firebase::ai::internal::AdapterHolder* holder = + new firebase::ai::internal::AdapterHolder(); + holder->adapter.reset(new firebase::ai::internal::LiteRtAdapter(params)); + return reinterpret_cast(holder); +} + +void firebase_ai_litert_destroy(FirebaseAiLiteRtHandle handle) { + if (!handle) return; + firebase::ai::internal::AdapterHolder* holder = + reinterpret_cast(handle); + delete holder; +} + +int32_t firebase_ai_litert_is_available(FirebaseAiLiteRtHandle handle) { + if (!handle) return 0; + firebase::ai::internal::AdapterHolder* holder = + reinterpret_cast(handle); + return holder->adapter->IsAvailable() ? 1 : 0; +} + +int32_t firebase_ai_litert_initialize(FirebaseAiLiteRtHandle handle, + char** out_error) { + if (out_error) *out_error = nullptr; + if (!handle) { + if (out_error) { + *out_error = + firebase::ai::internal::DuplicateCString("Invalid LiteRT handle."); + } + return static_cast(firebase::ai::kErrorInvalidArgument); + } + firebase::ai::internal::AdapterHolder* holder = + reinterpret_cast(handle); + std::string err; + if (!holder->adapter->Initialize(&err)) { + if (out_error) { + *out_error = firebase::ai::internal::DuplicateCString(err); + } + return static_cast(firebase::ai::kErrorUnsupported); + } + return static_cast(firebase::ai::kErrorNone); +} + +int32_t firebase_ai_litert_generate_content(FirebaseAiLiteRtHandle handle, + const char* request_json, + char** out_response_json, + char** out_error) { + if (out_response_json) *out_response_json = nullptr; + if (out_error) *out_error = nullptr; + if (!handle) { + if (out_error) { + *out_error = + firebase::ai::internal::DuplicateCString("Invalid LiteRT handle."); + } + return static_cast(firebase::ai::kErrorInvalidArgument); + } + + firebase::ai::internal::AdapterHolder* holder = + reinterpret_cast(handle); + + std::vector contents; + firebase::ai::Optional gen_config; + firebase::ai::Optional sys_instruction; + firebase::ai::internal::ParseRequestJson(request_json, &contents, &gen_config, + &sys_instruction); + if (contents.empty()) { + if (out_error) { + *out_error = firebase::ai::internal::DuplicateCString( + "Request contents must not be empty."); + } + return static_cast(firebase::ai::kErrorInvalidArgument); + } + + std::mutex mu; + std::condition_variable cv; + bool done = false; + firebase::ai::Error result_err = firebase::ai::kErrorNone; + std::string result_err_msg; + firebase::ai::GenerateContentResponse result_resp; + + holder->adapter->GenerateContentAsync( + contents, gen_config, sys_instruction, + [&](firebase::ai::Error err, const std::string& err_msg, + const firebase::ai::GenerateContentResponse& resp) { + std::lock_guard lk(mu); + result_err = err; + result_err_msg = err_msg; + result_resp = resp; + done = true; + cv.notify_all(); + }); + + std::unique_lock lk(mu); + cv.wait(lk, [&]() { return done; }); + + if (result_err != firebase::ai::kErrorNone) { + if (out_error) { + *out_error = firebase::ai::internal::DuplicateCString(result_err_msg); + } + return static_cast(result_err); + } + + if (out_response_json) { + *out_response_json = firebase::ai::internal::DuplicateCString( + firebase::ai::internal::ResponseToGeminiJson(result_resp)); + } + return static_cast(firebase::ai::kErrorNone); +} + +int32_t firebase_ai_litert_generate_content_stream( + FirebaseAiLiteRtHandle handle, const char* request_json, + FirebaseAiLiteRtStreamChunkCallback chunk_callback, void* user_data, + char** out_error) { + if (out_error) *out_error = nullptr; + if (!handle) { + if (out_error) { + *out_error = + firebase::ai::internal::DuplicateCString("Invalid LiteRT handle."); + } + return static_cast(firebase::ai::kErrorInvalidArgument); + } + + firebase::ai::internal::AdapterHolder* holder = + reinterpret_cast(handle); + + std::vector contents; + firebase::ai::Optional gen_config; + firebase::ai::Optional sys_instruction; + firebase::ai::internal::ParseRequestJson(request_json, &contents, &gen_config, + &sys_instruction); + if (contents.empty()) { + if (out_error) { + *out_error = firebase::ai::internal::DuplicateCString( + "Request contents must not be empty."); + } + return static_cast(firebase::ai::kErrorInvalidArgument); + } + + std::mutex mu; + std::condition_variable cv; + bool done = false; + firebase::ai::Error result_err = firebase::ai::kErrorNone; + std::string result_err_msg; + + holder->adapter->GenerateContentStreamAsync( + contents, gen_config, sys_instruction, + [chunk_callback, + user_data](const firebase::ai::GenerateContentResponse& chunk) { + if (chunk_callback) { + std::string chunk_json = + firebase::ai::internal::ResponseToGeminiJson(chunk); + chunk_callback(chunk_json.c_str(), user_data); + } + }, + [&](firebase::ai::Error err, const std::string& err_msg, + const firebase::ai::GenerateContentResponse& /*resp*/) { + std::lock_guard lk(mu); + result_err = err; + result_err_msg = err_msg; + done = true; + cv.notify_all(); + }); + + std::unique_lock lk(mu); + cv.wait(lk, [&]() { return done; }); + + if (result_err != firebase::ai::kErrorNone) { + if (out_error) { + *out_error = firebase::ai::internal::DuplicateCString(result_err_msg); + } + return static_cast(result_err); + } + return static_cast(firebase::ai::kErrorNone); +} + +int32_t firebase_ai_litert_count_tokens(FirebaseAiLiteRtHandle handle, + const char* request_json, + int32_t* out_total_tokens, + char** out_error) { + if (out_total_tokens) *out_total_tokens = 0; + if (out_error) *out_error = nullptr; + if (!handle) { + if (out_error) { + *out_error = + firebase::ai::internal::DuplicateCString("Invalid LiteRT handle."); + } + return static_cast(firebase::ai::kErrorInvalidArgument); + } + + firebase::ai::internal::AdapterHolder* holder = + reinterpret_cast(handle); + + std::vector contents; + firebase::ai::Optional gen_config; + firebase::ai::Optional sys_instruction; + firebase::ai::internal::ParseRequestJson(request_json, &contents, &gen_config, + &sys_instruction); + + std::mutex mu; + std::condition_variable cv; + bool done = false; + firebase::ai::Error result_err = firebase::ai::kErrorNone; + std::string result_err_msg; + firebase::ai::CountTokensResponse result_resp; + + holder->adapter->CountTokensAsync( + contents, [&](firebase::ai::Error err, const std::string& err_msg, + const firebase::ai::CountTokensResponse& resp) { + std::lock_guard lk(mu); + result_err = err; + result_err_msg = err_msg; + result_resp = resp; + done = true; + cv.notify_all(); + }); + + std::unique_lock lk(mu); + cv.wait(lk, [&]() { return done; }); + + if (result_err != firebase::ai::kErrorNone) { + if (out_error) { + *out_error = firebase::ai::internal::DuplicateCString(result_err_msg); + } + return static_cast(result_err); + } + if (out_total_tokens) { + *out_total_tokens = result_resp.total_tokens; + } + return static_cast(firebase::ai::kErrorNone); +} + +void firebase_ai_litert_free_string(char* str) { + if (str) { + std::free(str); + } +} + +} // extern "C" diff --git a/ai/src/common/litert_c_bridge.h b/ai/src/common/litert_c_bridge.h new file mode 100644 index 0000000000..af55529867 --- /dev/null +++ b/ai/src/common/litert_c_bridge.h @@ -0,0 +1,93 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_COMMON_LITERT_C_BRIDGE_H_ +#define FIREBASE_AI_SRC_COMMON_LITERT_C_BRIDGE_H_ + +#include + +#if defined(_WIN32) +#define FIREBASE_AI_C_EXPORT __declspec(dllexport) +#else +#define FIREBASE_AI_C_EXPORT __attribute__((visibility("default"))) +#endif + +#ifdef __cplusplus +extern "C" { +#endif + +/// @brief Opaque handle to a C++ `firebase::ai::internal::LiteRtAdapter` +/// instance used by the Unity C# SDK for hybrid on-device inference. +typedef void* FirebaseAiLiteRtHandle; + +/// @brief Callback invoked for each streamed `GenerateContentResponse` JSON +/// chunk from on-device LiteRT inference. +typedef void (*FirebaseAiLiteRtStreamChunkCallback)(const char* chunk_json, + void* user_data); + +/// @brief Creates a C++ `LiteRtAdapter` handle for Unity P/Invoke. +FIREBASE_AI_C_EXPORT FirebaseAiLiteRtHandle firebase_ai_litert_create( + const char* model_path, const char* runtime_library_path, + const char* cache_dir, int32_t accelerator, int32_t max_num_tokens, + int32_t num_threads, float temperature, int32_t top_k, float top_p); + +/// @brief Destroys a `FirebaseAiLiteRtHandle` created by +/// `firebase_ai_litert_create`. +FIREBASE_AI_C_EXPORT void firebase_ai_litert_destroy( + FirebaseAiLiteRtHandle handle); + +/// @brief Returns 1 if the local LiteRT / LiteRT-LM model is available, 0 +/// otherwise. +FIREBASE_AI_C_EXPORT int32_t +firebase_ai_litert_is_available(FirebaseAiLiteRtHandle handle); + +/// @brief Eagerly initializes the on-device LiteRT model. Returns 0 on success, +/// non-zero `firebase::ai::Error` on failure (allocating `*out_error` if +/// non-null). +FIREBASE_AI_C_EXPORT int32_t +firebase_ai_litert_initialize(FirebaseAiLiteRtHandle handle, char** out_error); + +/// @brief Runs synchronous on-device LiteRT inference from a Gemini-format +/// `GenerateContentRequest` JSON string and writes a Gemini-format +/// `GenerateContentResponse` JSON string to `*out_response_json`. +/// +/// Caller must free `*out_response_json` and `*out_error` via +/// `firebase_ai_litert_free_string`. +FIREBASE_AI_C_EXPORT int32_t firebase_ai_litert_generate_content( + FirebaseAiLiteRtHandle handle, const char* request_json, + char** out_response_json, char** out_error); + +/// @brief Runs streaming on-device LiteRT inference from a Gemini-format +/// `GenerateContentRequest` JSON string, invoking `chunk_callback` with each +/// incremental `GenerateContentResponse` JSON chunk. +FIREBASE_AI_C_EXPORT int32_t firebase_ai_litert_generate_content_stream( + FirebaseAiLiteRtHandle handle, const char* request_json, + FirebaseAiLiteRtStreamChunkCallback chunk_callback, void* user_data, + char** out_error); + +/// @brief Counts tokens locally using the C++ `LiteRtAdapter`. +FIREBASE_AI_C_EXPORT int32_t firebase_ai_litert_count_tokens( + FirebaseAiLiteRtHandle handle, const char* request_json, + int32_t* out_total_tokens, char** out_error); + +/// @brief Frees a heap-allocated C string returned by `firebase_ai_litert_*`. +FIREBASE_AI_C_EXPORT void firebase_ai_litert_free_string(char* str); + +#ifdef __cplusplus +} // extern "C" +#endif + +#endif // FIREBASE_AI_SRC_COMMON_LITERT_C_BRIDGE_H_ diff --git a/ai/src/common/model_content.cc b/ai/src/common/model_content.cc new file mode 100644 index 0000000000..c9119c0e9f --- /dev/null +++ b/ai/src/common/model_content.cc @@ -0,0 +1,300 @@ +/* + * Copyright 2025 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. + */ + +#include "firebase/ai/model_content.h" + +#include +#include +#include + +#include "firebase/ai/function_calling.h" +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/schema.h" + +namespace firebase { +namespace ai { + +// --- Schema factory methods --- + +Schema Schema::Boolean(const Optional& description, + const Optional& nullable, + const Optional& title) { + Schema s(kSchemaTypeBoolean); + s.description_ = description; + s.nullable_ = nullable; + s.title_ = title; + return s; +} + +Schema Schema::Int(const Optional& description, + const Optional& nullable, + const Optional& title, + const Optional& minimum, + const Optional& maximum) { + Schema s(kSchemaTypeInteger); + s.description_ = description; + s.nullable_ = nullable; + s.title_ = title; + s.format_ = std::string("int32"); + s.minimum_ = minimum; + s.maximum_ = maximum; + return s; +} + +Schema Schema::Long(const Optional& description, + const Optional& nullable, + const Optional& title, + const Optional& minimum, + const Optional& maximum) { + Schema s(kSchemaTypeInteger); + s.description_ = description; + s.nullable_ = nullable; + s.title_ = title; + s.format_ = std::string("int64"); + s.minimum_ = minimum; + s.maximum_ = maximum; + return s; +} + +Schema Schema::Float(const Optional& description, + const Optional& nullable, + const Optional& title, + const Optional& minimum, + const Optional& maximum) { + Schema s(kSchemaTypeNumber); + s.description_ = description; + s.nullable_ = nullable; + s.title_ = title; + s.format_ = std::string("float"); + s.minimum_ = minimum; + s.maximum_ = maximum; + return s; +} + +Schema Schema::Double(const Optional& description, + const Optional& nullable, + const Optional& title, + const Optional& minimum, + const Optional& maximum) { + Schema s(kSchemaTypeNumber); + s.description_ = description; + s.nullable_ = nullable; + s.title_ = title; + s.minimum_ = minimum; + s.maximum_ = maximum; + return s; +} + +Schema Schema::String(const Optional& description, + const Optional& nullable, + const Optional& format, + const Optional& title) { + Schema s(kSchemaTypeString); + s.description_ = description; + s.nullable_ = nullable; + s.format_ = format; + s.title_ = title; + return s; +} + +Schema Schema::Enum(const std::vector& values, + const Optional& description, + const Optional& nullable, + const Optional& title) { + Schema s(kSchemaTypeString); + s.enum_values_ = values; + s.format_ = std::string("enum"); + s.description_ = description; + s.nullable_ = nullable; + s.title_ = title; + return s; +} + +Schema Schema::Array(const Schema& items, + const Optional& description, + const Optional& nullable, + const Optional& title, + const Optional& min_items, + const Optional& max_items) { + Schema s(kSchemaTypeArray); + s.items_.reset(new Schema(items)); + s.description_ = description; + s.nullable_ = nullable; + s.title_ = title; + s.min_items_ = min_items; + s.max_items_ = max_items; + return s; +} + +Schema Schema::Object(const std::map& properties, + const std::vector& optional_properties, + const Optional& description, + const Optional& nullable, + const Optional& title, + const std::vector& property_ordering) { + Schema s(kSchemaTypeObject); + s.properties_ = properties; + s.description_ = description; + s.nullable_ = nullable; + s.title_ = title; + s.property_ordering_ = property_ordering; + + std::set opt_set(optional_properties.begin(), + optional_properties.end()); + for (const auto& kv : properties) { + if (opt_set.find(kv.first) == opt_set.end()) { + s.required_properties_.push_back(kv.first); + } + } + return s; +} + +Schema Schema::AnyOf(const std::vector& schemas) { + Schema s(kSchemaTypeUnspecified); + s.any_of_ = schemas; + return s; +} + +// --- FunctionDeclaration --- + +FunctionDeclaration::FunctionDeclaration( + const std::string& name, const std::string& description, + const std::map& parameters, + const std::vector& optional_parameters) + : name_(name), + description_(description), + parameters_(Schema::Object(parameters, optional_parameters)), + uses_json_schema_(false) {} + +FunctionDeclaration::FunctionDeclaration( + const std::string& name, const std::string& description, + const JsonSchema& parameters_json_schema) + : name_(name), + description_(description), + parameters_(parameters_json_schema), + uses_json_schema_(true) {} + +// --- ModelContent factory methods --- + +ModelContent ModelContent::Text(const std::string& text) { + return ModelContent("user", std::vector(1, Part(TextPart(text)))); +} + +ModelContent ModelContent::InlineData(const std::string& mime_type, + const std::vector& data) { + return ModelContent( + "user", std::vector(1, Part(InlineDataPart(mime_type, data)))); +} + +ModelContent ModelContent::InlineData(const std::string& mime_type, + const uint8_t* bytes, size_t size) { + return ModelContent( + "user", + std::vector(1, Part(InlineDataPart(mime_type, bytes, size)))); +} + +ModelContent ModelContent::FileData(const std::string& mime_type, + const std::string& uri) { + return ModelContent("user", + std::vector(1, Part(FileDataPart(mime_type, uri)))); +} + +ModelContent ModelContent::FunctionResponse( + const std::string& name, const std::map& response, + const Optional& id) { + return ModelContent( + "user", + std::vector(1, Part(FunctionResponsePart(name, response, id)))); +} + +ModelContent ModelContent::FunctionResponses( + const std::vector& responses) { + std::vector parts; + parts.reserve(responses.size()); + for (const auto& r : responses) { + parts.push_back(Part(r)); + } + return ModelContent("user", parts); +} + +ModelContent ModelContent::System(const std::string& text) { + return ModelContent("system", std::vector(1, Part(TextPart(text)))); +} + +ModelContent ModelContent::User(const std::string& text) { + return ModelContent("user", std::vector(1, Part(TextPart(text)))); +} + +ModelContent ModelContent::User(const std::vector& parts) { + return ModelContent("user", parts); +} + +ModelContent ModelContent::Model(const std::string& text) { + return ModelContent("model", std::vector(1, Part(TextPart(text)))); +} + +ModelContent ModelContent::Model(const std::vector& parts) { + return ModelContent("model", parts); +} + +// --- GenerateContentResponse convenience accessors --- + +std::string GenerateContentResponse::text() const { + if (candidates_.empty()) return ""; + std::ostringstream oss; + for (const auto& part : candidates_[0].content.parts()) { + if (part.is_text() && !part.is_thought()) { + oss << part.text_part().text; + } + } + return oss.str(); +} + +std::string GenerateContentResponse::thought_summary() const { + if (candidates_.empty()) return ""; + std::ostringstream oss; + for (const auto& part : candidates_[0].content.parts()) { + if (part.is_text() && part.is_thought()) { + oss << part.text_part().text; + } + } + return oss.str(); +} + +std::vector GenerateContentResponse::function_calls() const { + std::vector result; + if (candidates_.empty()) return result; + for (const auto& part : candidates_[0].content.parts()) { + if (part.is_function_call()) { + result.push_back(part.function_call_part()); + } + } + return result; +} + +std::vector GenerateContentResponse::inline_data_parts() const { + std::vector result; + if (candidates_.empty()) return result; + for (const auto& part : candidates_[0].content.parts()) { + if (part.is_inline_data() && !part.is_thought()) { + result.push_back(part.inline_data_part()); + } + } + return result; +} + +} // namespace ai +} // namespace firebase diff --git a/ai/src/common/serialization.cc b/ai/src/common/serialization.cc new file mode 100644 index 0000000000..e14440e2a5 --- /dev/null +++ b/ai/src/common/serialization.cc @@ -0,0 +1,1280 @@ +/* + * Copyright 2025 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. + */ + +#include "ai/src/common/serialization.h" + +#include + +#include "app/src/base64.h" +#include "app/src/log.h" +#include "app/src/variant_util.h" + +namespace firebase { +namespace ai { +namespace internal { + +namespace { + +const Variant* FindField(const Variant& map_var, const char* key) { + if (!map_var.is_map()) return nullptr; + auto it = map_var.map().find(Variant::FromStaticString(key)); + if (it == map_var.map().end()) return nullptr; + return &it->second; +} + +std::string GetStringField(const Variant& map_var, const char* key, + const std::string& default_val = "") { + const Variant* v = FindField(map_var, key); + if (v && v->is_string()) { + return v->string_value(); + } + return default_val; +} + +int GetIntField(const Variant& map_var, const char* key, int default_val = 0) { + const Variant* v = FindField(map_var, key); + if (!v) return default_val; + if (v->is_int64()) return static_cast(v->int64_value()); + if (v->is_double()) return static_cast(v->double_value()); + return default_val; +} + +float GetFloatField(const Variant& map_var, const char* key, + float default_val = 0.0f) { + const Variant* v = FindField(map_var, key); + if (!v) return default_val; + if (v->is_double()) return static_cast(v->double_value()); + if (v->is_int64()) return static_cast(v->int64_value()); + return default_val; +} + +bool GetBoolField(const Variant& map_var, const char* key, + bool default_val = false) { + const Variant* v = FindField(map_var, key); + if (v && v->is_bool()) { + return v->bool_value(); + } + return default_val; +} + +const char* SchemaTypeToOpenApiString(SchemaType type) { + switch (type) { + case kSchemaTypeString: + return "STRING"; + case kSchemaTypeNumber: + return "NUMBER"; + case kSchemaTypeInteger: + return "INTEGER"; + case kSchemaTypeBoolean: + return "BOOLEAN"; + case kSchemaTypeArray: + return "ARRAY"; + case kSchemaTypeObject: + return "OBJECT"; + case kSchemaTypeUnspecified: + default: + return ""; + } +} + +const char* SchemaTypeToJsonSchemaString(SchemaType type) { + switch (type) { + case kSchemaTypeString: + return "string"; + case kSchemaTypeNumber: + return "number"; + case kSchemaTypeInteger: + return "integer"; + case kSchemaTypeBoolean: + return "boolean"; + case kSchemaTypeArray: + return "array"; + case kSchemaTypeObject: + return "object"; + case kSchemaTypeUnspecified: + default: + return ""; + } +} + +const char* HarmCategoryToString(HarmCategory category) { + switch (category) { + case kHarmCategoryHarassment: + return "HARM_CATEGORY_HARASSMENT"; + case kHarmCategoryHateSpeech: + return "HARM_CATEGORY_HATE_SPEECH"; + case kHarmCategorySexuallyExplicit: + return "HARM_CATEGORY_SEXUALLY_EXPLICIT"; + case kHarmCategoryDangerousContent: + return "HARM_CATEGORY_DANGEROUS_CONTENT"; + case kHarmCategoryCivicIntegrity: + return "HARM_CATEGORY_CIVIC_INTEGRITY"; + case kHarmCategoryUnknown: + default: + return "HARM_CATEGORY_UNSPECIFIED"; + } +} + +HarmCategory ParseHarmCategory(const std::string& str) { + if (str == "HARM_CATEGORY_HARASSMENT") return kHarmCategoryHarassment; + if (str == "HARM_CATEGORY_HATE_SPEECH") return kHarmCategoryHateSpeech; + if (str == "HARM_CATEGORY_SEXUALLY_EXPLICIT") + return kHarmCategorySexuallyExplicit; + if (str == "HARM_CATEGORY_DANGEROUS_CONTENT") + return kHarmCategoryDangerousContent; + if (str == "HARM_CATEGORY_CIVIC_INTEGRITY") + return kHarmCategoryCivicIntegrity; + return kHarmCategoryUnknown; +} + +const char* HarmBlockThresholdToString(HarmBlockThreshold threshold) { + switch (threshold) { + case kHarmBlockThresholdLowAndAbove: + return "BLOCK_LOW_AND_ABOVE"; + case kHarmBlockThresholdMediumAndAbove: + return "BLOCK_MEDIUM_AND_ABOVE"; + case kHarmBlockThresholdOnlyHigh: + return "BLOCK_ONLY_HIGH"; + case kHarmBlockThresholdNone: + return "BLOCK_NONE"; + case kHarmBlockThresholdOff: + return "OFF"; + case kHarmBlockThresholdUnknown: + default: + return "HARM_BLOCK_THRESHOLD_UNSPECIFIED"; + } +} + +const char* HarmBlockMethodToString(HarmBlockMethod method) { + switch (method) { + case kHarmBlockMethodSeverity: + return "SEVERITY"; + case kHarmBlockMethodProbability: + return "PROBABILITY"; + case kHarmBlockMethodUnknown: + default: + return "HARM_BLOCK_METHOD_UNSPECIFIED"; + } +} + +HarmProbability ParseHarmProbability(const std::string& str) { + if (str == "NEGLIGIBLE") return kHarmProbabilityNegligible; + if (str == "LOW") return kHarmProbabilityLow; + if (str == "MEDIUM") return kHarmProbabilityMedium; + if (str == "HIGH") return kHarmProbabilityHigh; + return kHarmProbabilityUnknown; +} + +HarmSeverity ParseHarmSeverity(const std::string& str) { + if (str == "HARM_SEVERITY_NEGLIGIBLE") return kHarmSeverityNegligible; + if (str == "HARM_SEVERITY_LOW") return kHarmSeverityLow; + if (str == "HARM_SEVERITY_MEDIUM") return kHarmSeverityMedium; + if (str == "HARM_SEVERITY_HIGH") return kHarmSeverityHigh; + return kHarmSeverityUnknown; +} + +FinishReason ParseFinishReason(const std::string& str) { + if (str == "STOP") return kFinishReasonStop; + if (str == "MAX_TOKENS") return kFinishReasonMaxTokens; + if (str == "SAFETY") return kFinishReasonSafety; + if (str == "RECITATION") return kFinishReasonRecitation; + if (str == "OTHER") return kFinishReasonOther; + if (str == "BLOCKLIST") return kFinishReasonBlocklist; + if (str == "PROHIBITED_CONTENT") return kFinishReasonProhibitedContent; + if (str == "SPII") return kFinishReasonSpii; + if (str == "MALFORMED_FUNCTION_CALL") + return kFinishReasonMalformedFunctionCall; + return kFinishReasonUnknown; +} + +BlockReason ParseBlockReason(const std::string& str) { + if (str == "SAFETY") return kBlockReasonSafety; + if (str == "OTHER") return kBlockReasonOther; + if (str == "BLOCKLIST") return kBlockReasonBlocklist; + if (str == "PROHIBITED_CONTENT") return kBlockReasonProhibitedContent; + return kBlockReasonUnknown; +} + +ContentModality ParseContentModality(const std::string& str) { + if (str == "TEXT") return kContentModalityText; + if (str == "IMAGE") return kContentModalityImage; + if (str == "VIDEO") return kContentModalityVideo; + if (str == "AUDIO") return kContentModalityAudio; + if (str == "DOCUMENT") return kContentModalityDocument; + return kContentModalityUnspecified; +} + +const char* ResponseModalityToString(ResponseModality modality) { + switch (modality) { + case kResponseModalityText: + return "TEXT"; + case kResponseModalityImage: + return "IMAGE"; + case kResponseModalityAudio: + return "AUDIO"; + case kResponseModalityUnspecified: + default: + return "MODALITY_UNSPECIFIED"; + } +} + +const char* ThinkingLevelToString(ThinkingLevel level) { + switch (level) { + case kThinkingLevelMinimal: + return "MINIMAL"; + case kThinkingLevelLow: + return "LOW"; + case kThinkingLevelMedium: + return "MEDIUM"; + case kThinkingLevelHigh: + return "HIGH"; + case kThinkingLevelUnspecified: + default: + return "THINKING_LEVEL_UNSPECIFIED"; + } +} + +CodeExecutionOutcome ParseCodeExecutionOutcome(const std::string& str) { + if (str == "OUTCOME_OK") return kCodeExecutionOutcomeOk; + if (str == "OUTCOME_FAILED") return kCodeExecutionOutcomeFailed; + if (str == "OUTCOME_DEADLINE_EXCEEDED") + return kCodeExecutionOutcomeDeadlineExceeded; + return kCodeExecutionOutcomeUnspecified; +} + +const char* CodeExecutionOutcomeToString(CodeExecutionOutcome outcome) { + switch (outcome) { + case kCodeExecutionOutcomeOk: + return "OUTCOME_OK"; + case kCodeExecutionOutcomeFailed: + return "OUTCOME_FAILED"; + case kCodeExecutionOutcomeDeadlineExceeded: + return "OUTCOME_DEADLINE_EXCEEDED"; + case kCodeExecutionOutcomeUnspecified: + default: + return "OUTCOME_UNSPECIFIED"; + } +} + +UrlRetrievalStatus ParseUrlRetrievalStatus(const std::string& str) { + if (str == "URL_RETRIEVAL_STATUS_SUCCESS") return kUrlRetrievalStatusSuccess; + if (str == "URL_RETRIEVAL_STATUS_ERROR") return kUrlRetrievalStatusError; + if (str == "URL_RETRIEVAL_STATUS_PAYWALL") return kUrlRetrievalStatusPaywall; + if (str == "URL_RETRIEVAL_STATUS_UNSAFE") return kUrlRetrievalStatusUnsafe; + return kUrlRetrievalStatusUnspecified; +} + +std::vector ParseModalityTokenCounts( + const Variant* list_var) { + std::vector result; + if (!list_var || !list_var->is_vector()) return result; + for (const auto& elem : list_var->vector()) { + if (!elem.is_map()) continue; + ModalityTokenCount mtc; + mtc.modality = ParseContentModality(GetStringField(elem, "modality")); + mtc.token_count = GetIntField(elem, "tokenCount"); + result.push_back(mtc); + } + return result; +} + +CitationMetadata ParseCitationMetadata(const Variant& map_var, + BackendProvider provider) { + CitationMetadata metadata; + const char* key = + (provider == kBackendProviderGoogleAI) ? "citationSources" : "citations"; + const Variant* citations_var = FindField(map_var, key); + if (!citations_var) { + // Fallback to the other key in case backend format varies. + citations_var = FindField(map_var, (provider == kBackendProviderGoogleAI) + ? "citations" + : "citationSources"); + } + if (citations_var && citations_var->is_vector()) { + for (const auto& item : citations_var->vector()) { + if (!item.is_map()) continue; + Citation c; + c.start_index = GetIntField(item, "startIndex"); + c.end_index = GetIntField(item, "endIndex"); + c.uri = GetStringField(item, "uri"); + c.title = GetStringField(item, "title"); + c.license = GetStringField(item, "license"); + const Variant* pub_date = FindField(item, "publicationDate"); + if (pub_date && pub_date->is_string()) { + c.publication_date = pub_date->string_value(); + } else if (pub_date && pub_date->is_map()) { + int year = GetIntField(*pub_date, "year"); + int month = GetIntField(*pub_date, "month"); + int day = GetIntField(*pub_date, "day"); + std::ostringstream oss; + if (year > 0) oss << year; + if (month > 0) oss << "-" << month; + if (day > 0) oss << "-" << day; + c.publication_date = oss.str(); + } + metadata.citations.push_back(c); + } + } + return metadata; +} + +GroundingMetadata ParseGroundingMetadata(const Variant& map_var) { + GroundingMetadata gm; + const Variant* queries = FindField(map_var, "webSearchQueries"); + if (queries && queries->is_vector()) { + for (const auto& q : queries->vector()) { + if (q.is_string()) gm.web_search_queries.push_back(q.string_value()); + } + } + + const Variant* sep = FindField(map_var, "searchEntryPoint"); + if (sep && sep->is_map()) { + SearchEntryPoint entry; + entry.rendered_content = GetStringField(*sep, "renderedContent"); + entry.sdk_blob = GetStringField(*sep, "sdkBlob"); + gm.search_entry_point = entry; + } + + const Variant* chunks = FindField(map_var, "groundingChunks"); + if (chunks && chunks->is_vector()) { + for (const auto& ch : chunks->vector()) { + if (!ch.is_map()) continue; + GroundingChunk gc; + const Variant* web = FindField(ch, "web"); + if (web && web->is_map()) { + WebGroundingChunk wgc; + wgc.uri = GetStringField(*web, "uri"); + wgc.title = GetStringField(*web, "title"); + wgc.domain = GetStringField(*web, "domain"); + gc.web = wgc; + } + const Variant* maps = FindField(ch, "maps"); + if (maps && maps->is_map()) { + GoogleMapsGroundingChunk mgc; + mgc.uri = GetStringField(*maps, "uri"); + mgc.title = GetStringField(*maps, "title"); + mgc.place_id = GetStringField(*maps, "placeId"); + gc.maps = mgc; + } + gm.grounding_chunks.push_back(gc); + } + } + + const Variant* supports = FindField(map_var, "groundingSupports"); + if (supports && supports->is_vector()) { + for (const auto& sup : supports->vector()) { + if (!sup.is_map()) continue; + GroundingSupport gs; + const Variant* seg = FindField(sup, "segment"); + if (seg && seg->is_map()) { + gs.segment.part_index = GetIntField(*seg, "partIndex"); + gs.segment.start_index = GetIntField(*seg, "startIndex"); + gs.segment.end_index = GetIntField(*seg, "endIndex"); + gs.segment.text = GetStringField(*seg, "text"); + } + const Variant* indices = FindField(sup, "groundingChunkIndices"); + if (indices && indices->is_vector()) { + for (const auto& idx : indices->vector()) { + if (idx.is_int64()) { + gs.grounding_chunk_indices.push_back( + static_cast(idx.int64_value())); + } + } + } + gm.grounding_supports.push_back(gs); + } + } + + const Variant* widget_token = + FindField(map_var, "googleMapsWidgetContextToken"); + if (widget_token && widget_token->is_string()) { + gm.google_maps_widget_context_token = + std::string(widget_token->string_value()); + } + return gm; +} + +UrlContextMetadata ParseUrlContextMetadata(const Variant& map_var) { + UrlContextMetadata ucm; + const Variant* list_var = FindField(map_var, "urlMetadata"); + if (list_var && list_var->is_vector()) { + for (const auto& elem : list_var->vector()) { + if (!elem.is_map()) continue; + UrlMetadata um; + um.retrieved_url = GetStringField(elem, "retrievedUrl"); + um.retrieval_status = + ParseUrlRetrievalStatus(GetStringField(elem, "urlRetrievalStatus")); + ucm.url_metadata.push_back(um); + } + } + return ucm; +} + +} // namespace + +Variant SchemaToVariant(const Schema& schema) { + Variant map = Variant::EmptyMap(); + const char* type_str = SchemaTypeToOpenApiString(schema.type()); + if (type_str[0] != '\0') { + map.map()["type"] = type_str; + } + if (schema.description().has_value()) { + map.map()["description"] = schema.description().value(); + } + if (schema.format().has_value()) { + map.map()["format"] = schema.format().value(); + } + if (schema.nullable().has_value()) { + map.map()["nullable"] = schema.nullable().value(); + } + if (!schema.enum_values().empty()) { + Variant enums = Variant::EmptyVector(); + for (const auto& val : schema.enum_values()) { + enums.vector().push_back(val); + } + map.map()["enum"] = enums; + } + if (!schema.properties().empty()) { + Variant props = Variant::EmptyMap(); + for (const auto& kv : schema.properties()) { + props.map()[kv.first] = SchemaToVariant(kv.second); + } + map.map()["properties"] = props; + } + if (!schema.required_properties().empty()) { + Variant req = Variant::EmptyVector(); + for (const auto& r : schema.required_properties()) { + req.vector().push_back(r); + } + map.map()["required"] = req; + } + if (!schema.property_ordering().empty()) { + Variant ord = Variant::EmptyVector(); + for (const auto& p : schema.property_ordering()) { + ord.vector().push_back(p); + } + map.map()["propertyOrdering"] = ord; + } + if (schema.items() != nullptr) { + map.map()["items"] = SchemaToVariant(*schema.items()); + } + if (schema.title().has_value()) { + map.map()["title"] = schema.title().value(); + } + if (schema.min_items().has_value()) { + map.map()["minItems"] = schema.min_items().value(); + } + if (schema.max_items().has_value()) { + map.map()["maxItems"] = schema.max_items().value(); + } + if (schema.minimum().has_value()) { + map.map()["minimum"] = schema.minimum().value(); + } + if (schema.maximum().has_value()) { + map.map()["maximum"] = schema.maximum().value(); + } + if (!schema.any_of().empty()) { + Variant any_of_vec = Variant::EmptyVector(); + for (const auto& sub : schema.any_of()) { + any_of_vec.vector().push_back(SchemaToVariant(sub)); + } + map.map()["anyOf"] = any_of_vec; + } + return map; +} + +Variant JsonSchemaToVariant(const JsonSchema& schema) { + Variant map = Variant::EmptyMap(); + const char* type_str = SchemaTypeToJsonSchemaString(schema.type()); + bool is_nullable = schema.nullable().value_or(false); + if (type_str[0] != '\0') { + if (is_nullable) { + Variant types = Variant::EmptyVector(); + types.vector().push_back(type_str); + types.vector().push_back("null"); + map.map()["type"] = types; + } else { + map.map()["type"] = type_str; + } + } + if (schema.description().has_value()) { + map.map()["description"] = schema.description().value(); + } + if (schema.format().has_value()) { + map.map()["format"] = schema.format().value(); + } + if (!schema.enum_values().empty()) { + Variant enums = Variant::EmptyVector(); + for (const auto& val : schema.enum_values()) { + enums.vector().push_back(val); + } + map.map()["enum"] = enums; + } + if (!schema.properties().empty()) { + Variant props = Variant::EmptyMap(); + for (const auto& kv : schema.properties()) { + props.map()[kv.first] = JsonSchemaToVariant(kv.second); + } + map.map()["properties"] = props; + } + if (!schema.required_properties().empty()) { + Variant req = Variant::EmptyVector(); + for (const auto& r : schema.required_properties()) { + req.vector().push_back(r); + } + map.map()["required"] = req; + } + if (schema.items() != nullptr) { + map.map()["items"] = JsonSchemaToVariant(*schema.items()); + } + if (schema.title().has_value()) { + map.map()["title"] = schema.title().value(); + } + if (schema.min_items().has_value()) { + map.map()["minItems"] = schema.min_items().value(); + } + if (schema.max_items().has_value()) { + map.map()["maxItems"] = schema.max_items().value(); + } + if (schema.minimum().has_value()) { + map.map()["minimum"] = schema.minimum().value(); + } + if (schema.maximum().has_value()) { + map.map()["maximum"] = schema.maximum().value(); + } + if (!schema.any_of().empty()) { + Variant any_of_vec = Variant::EmptyVector(); + for (const auto& sub : schema.any_of()) { + any_of_vec.vector().push_back(JsonSchemaToVariant(sub)); + } + if (is_nullable && type_str[0] == '\0') { + Variant null_type = Variant::EmptyMap(); + null_type.map()["type"] = "null"; + any_of_vec.vector().push_back(null_type); + } + map.map()["anyOf"] = any_of_vec; + } + return map; +} + +Variant PartToVariant(const Part& part) { + Variant map = Variant::EmptyMap(); + switch (part.kind()) { + case Part::kKindText: { + map.map()["text"] = part.text_part().text; + break; + } + case Part::kKindInlineData: { + Variant inline_data = Variant::EmptyMap(); + inline_data.map()["mimeType"] = part.inline_data_part().mime_type; + std::string raw( + reinterpret_cast(part.inline_data_part().data.data()), + part.inline_data_part().data.size()); + std::string encoded; + ::firebase::internal::Base64EncodeWithPadding(raw, &encoded); + inline_data.map()["data"] = encoded; + map.map()["inlineData"] = inline_data; + break; + } + case Part::kKindFileData: { + Variant file_data = Variant::EmptyMap(); + file_data.map()["mimeType"] = part.file_data_part().mime_type; + file_data.map()["fileUri"] = part.file_data_part().uri; + map.map()["fileData"] = file_data; + break; + } + case Part::kKindFunctionCall: { + Variant fc = Variant::EmptyMap(); + fc.map()["name"] = part.function_call_part().name; + Variant args = Variant::EmptyMap(); + for (const auto& kv : part.function_call_part().args) { + args.map()[kv.first] = kv.second; + } + fc.map()["args"] = args; + if (part.function_call_part().id.has_value()) { + fc.map()["id"] = part.function_call_part().id.value(); + } + map.map()["functionCall"] = fc; + break; + } + case Part::kKindFunctionResponse: { + Variant fr = Variant::EmptyMap(); + fr.map()["name"] = part.function_response_part().name; + Variant resp = Variant::EmptyMap(); + for (const auto& kv : part.function_response_part().response) { + resp.map()[kv.first] = kv.second; + } + fr.map()["response"] = resp; + if (part.function_response_part().id.has_value()) { + fr.map()["id"] = part.function_response_part().id.value(); + } + map.map()["functionResponse"] = fr; + break; + } + case Part::kKindExecutableCode: { + Variant ec = Variant::EmptyMap(); + ec.map()["language"] = (part.executable_code_part().language == + ExecutableCodePart::kLanguagePython) + ? "PYTHON" + : "LANGUAGE_UNSPECIFIED"; + ec.map()["code"] = part.executable_code_part().code; + map.map()["executableCode"] = ec; + break; + } + case Part::kKindCodeExecutionResult: { + Variant cer = Variant::EmptyMap(); + cer.map()["outcome"] = CodeExecutionOutcomeToString( + part.code_execution_result_part().outcome); + cer.map()["output"] = part.code_execution_result_part().output; + map.map()["codeExecutionResult"] = cer; + break; + } + case Part::kKindNone: + default: + break; + } + + if (part.is_thought()) { + map.map()["thought"] = true; + } + if (part.thought_signature().has_value()) { + map.map()["thoughtSignature"] = part.thought_signature().value(); + } + return map; +} + +bool PartFromVariant(const Variant& variant, Part* out_part) { + if (!variant.is_map() || !out_part) return false; + + bool is_thought = GetBoolField(variant, "thought", false); + Optional thought_sig; + const Variant* sig_var = FindField(variant, "thoughtSignature"); + if (sig_var && sig_var->is_string()) { + thought_sig = std::string(sig_var->string_value()); + } + + const Variant* text_var = FindField(variant, "text"); + if (text_var && text_var->is_string()) { + *out_part = + Part(TextPart(text_var->string_value()), is_thought, thought_sig); + return true; + } + + const Variant* inline_var = FindField(variant, "inlineData"); + if (inline_var && inline_var->is_map()) { + std::string mime_type = GetStringField(*inline_var, "mimeType"); + std::string base64_data = GetStringField(*inline_var, "data"); + std::string decoded; + ::firebase::internal::Base64Decode(base64_data, &decoded); + std::vector bytes(decoded.begin(), decoded.end()); + *out_part = Part(InlineDataPart(mime_type, bytes), is_thought, thought_sig); + return true; + } + + const Variant* file_var = FindField(variant, "fileData"); + if (file_var && file_var->is_map()) { + std::string mime_type = GetStringField(*file_var, "mimeType"); + std::string file_uri = GetStringField(*file_var, "fileUri"); + *out_part = + Part(FileDataPart(mime_type, file_uri), is_thought, thought_sig); + return true; + } + + const Variant* fc_var = FindField(variant, "functionCall"); + if (fc_var && fc_var->is_map()) { + std::string name = GetStringField(*fc_var, "name"); + std::map args; + const Variant* args_var = FindField(*fc_var, "args"); + if (args_var && args_var->is_map()) { + for (const auto& kv : args_var->map()) { + if (kv.first.is_string()) { + args[kv.first.string_value()] = kv.second; + } + } + } + Optional id; + const Variant* id_var = FindField(*fc_var, "id"); + if (id_var && id_var->is_string()) { + id = std::string(id_var->string_value()); + } + *out_part = Part(FunctionCallPart(name, args, id), is_thought, thought_sig); + return true; + } + + const Variant* fr_var = FindField(variant, "functionResponse"); + if (fr_var && fr_var->is_map()) { + std::string name = GetStringField(*fr_var, "name"); + std::map resp; + const Variant* resp_var = FindField(*fr_var, "response"); + if (resp_var && resp_var->is_map()) { + for (const auto& kv : resp_var->map()) { + if (kv.first.is_string()) { + resp[kv.first.string_value()] = kv.second; + } + } + } + Optional id; + const Variant* id_var = FindField(*fr_var, "id"); + if (id_var && id_var->is_string()) { + id = std::string(id_var->string_value()); + } + *out_part = + Part(FunctionResponsePart(name, resp, id), is_thought, thought_sig); + return true; + } + + const Variant* ec_var = FindField(variant, "executableCode"); + if (ec_var && ec_var->is_map()) { + std::string lang_str = GetStringField(*ec_var, "language"); + ExecutableCodePart::CodeLanguage lang = + (lang_str == "PYTHON") ? ExecutableCodePart::kLanguagePython + : ExecutableCodePart::kLanguageUnspecified; + std::string code = GetStringField(*ec_var, "code"); + *out_part = Part(ExecutableCodePart(lang, code), is_thought, thought_sig); + return true; + } + + const Variant* cer_var = FindField(variant, "codeExecutionResult"); + if (cer_var && cer_var->is_map()) { + CodeExecutionOutcome outcome = + ParseCodeExecutionOutcome(GetStringField(*cer_var, "outcome")); + std::string output = GetStringField(*cer_var, "output"); + *out_part = + Part(CodeExecutionResultPart(outcome, output), is_thought, thought_sig); + return true; + } + + return false; +} + +Variant ModelContentToVariant(const ModelContent& content) { + Variant map = Variant::EmptyMap(); + map.map()["role"] = content.role().empty() ? "user" : content.role(); + Variant parts = Variant::EmptyVector(); + for (const auto& part : content.parts()) { + parts.vector().push_back(PartToVariant(part)); + } + map.map()["parts"] = parts; + return map; +} + +bool ModelContentFromVariant(const Variant& variant, + ModelContent* out_content) { + if (!variant.is_map() || !out_content) return false; + std::string role = GetStringField(variant, "role", "model"); + std::vector parts; + const Variant* parts_var = FindField(variant, "parts"); + if (parts_var && parts_var->is_vector()) { + for (const auto& part_var : parts_var->vector()) { + Part p; + if (PartFromVariant(part_var, &p)) { + parts.push_back(p); + } + } + } + *out_content = ModelContent(role, parts); + return true; +} + +Variant SafetySettingToVariant(const SafetySetting& setting, + BackendProvider provider) { + Variant map = Variant::EmptyMap(); + map.map()["category"] = HarmCategoryToString(setting.category()); + map.map()["threshold"] = HarmBlockThresholdToString(setting.threshold()); + if (provider == kBackendProviderEnterprise && setting.method().has_value()) { + map.map()["method"] = HarmBlockMethodToString(setting.method().value()); + } + return map; +} + +SafetyRating SafetyRatingFromVariant(const Variant& variant) { + SafetyRating rating; + if (!variant.is_map()) return rating; + rating.category = ParseHarmCategory(GetStringField(variant, "category")); + rating.probability = + ParseHarmProbability(GetStringField(variant, "probability")); + rating.blocked = GetBoolField(variant, "blocked", false); + rating.probability_score = GetFloatField(variant, "probabilityScore", 0.0f); + rating.severity = ParseHarmSeverity(GetStringField(variant, "severity")); + rating.severity_score = GetFloatField(variant, "severityScore", 0.0f); + return rating; +} + +Variant GenerationConfigToVariant(const GenerationConfig& config) { + Variant map = Variant::EmptyMap(); + if (config.temperature.has_value()) { + map.map()["temperature"] = static_cast(config.temperature.value()); + } + if (config.top_p.has_value()) { + map.map()["topP"] = static_cast(config.top_p.value()); + } + if (config.top_k.has_value()) { + map.map()["topK"] = config.top_k.value(); + } + if (config.candidate_count.has_value()) { + map.map()["candidateCount"] = config.candidate_count.value(); + } + if (config.max_output_tokens.has_value()) { + map.map()["maxOutputTokens"] = config.max_output_tokens.value(); + } + if (config.presence_penalty.has_value()) { + map.map()["presencePenalty"] = + static_cast(config.presence_penalty.value()); + } + if (config.frequency_penalty.has_value()) { + map.map()["frequencyPenalty"] = + static_cast(config.frequency_penalty.value()); + } + if (!config.stop_sequences.empty()) { + Variant stops = Variant::EmptyVector(); + for (const auto& s : config.stop_sequences) { + stops.vector().push_back(s); + } + map.map()["stopSequences"] = stops; + } + if (config.response_mime_type.has_value()) { + map.map()["responseMimeType"] = config.response_mime_type.value(); + } + if (config.response_schema.has_value()) { + map.map()["responseSchema"] = + SchemaToVariant(config.response_schema.value()); + } + if (config.response_json_schema.has_value()) { + map.map()["responseJsonSchema"] = + JsonSchemaToVariant(config.response_json_schema.value()); + } + if (!config.response_modalities.empty()) { + Variant mods = Variant::EmptyVector(); + for (const auto& m : config.response_modalities) { + mods.vector().push_back(ResponseModalityToString(m)); + } + map.map()["responseModalities"] = mods; + } + if (config.thinking_config.has_value()) { + const ThinkingConfig& tc = config.thinking_config.value(); + Variant tc_map = Variant::EmptyMap(); + if (tc.thinking_budget.has_value()) { + tc_map.map()["thinkingBudget"] = tc.thinking_budget.value(); + } + if (tc.thinking_level.has_value()) { + tc_map.map()["thinkingLevel"] = + ThinkingLevelToString(tc.thinking_level.value()); + } + if (tc.include_thoughts.has_value()) { + tc_map.map()["includeThoughts"] = tc.include_thoughts.value(); + } + if (!tc_map.map().empty()) { + map.map()["thinkingConfig"] = tc_map; + } + } + if (config.image_config.has_value()) { + const ImageConfig& ic = config.image_config.value(); + Variant ic_map = Variant::EmptyMap(); + if (ic.aspect_ratio.has_value()) { + ic_map.map()["aspectRatio"] = ic.aspect_ratio.value(); + } + if (ic.image_size.has_value()) { + ic_map.map()["imageSize"] = ic.image_size.value(); + } + if (!ic_map.map().empty()) { + map.map()["imageConfig"] = ic_map; + } + } + return map; +} + +Variant ToolToVariant(const Tool& tool) { + Variant map = Variant::EmptyMap(); + if (!tool.function_declarations().empty()) { + Variant decls = Variant::EmptyVector(); + for (const auto& fn : tool.function_declarations()) { + Variant fn_map = Variant::EmptyMap(); + fn_map.map()["name"] = fn.name(); + if (!fn.description().empty()) { + fn_map.map()["description"] = fn.description(); + } + if (fn.uses_json_schema()) { + fn_map.map()["parametersJsonSchema"] = + JsonSchemaToVariant(fn.parameters()); + } else { + fn_map.map()["parameters"] = SchemaToVariant(fn.parameters()); + } + decls.vector().push_back(fn_map); + } + map.map()["functionDeclarations"] = decls; + } + if (tool.google_search().has_value()) { + map.map()["googleSearch"] = Variant::EmptyMap(); + } + if (tool.code_execution().has_value()) { + map.map()["codeExecution"] = Variant::EmptyMap(); + } + if (tool.google_maps().has_value()) { + map.map()["googleMaps"] = Variant::EmptyMap(); + } + if (tool.url_context().has_value()) { + map.map()["urlContext"] = Variant::EmptyMap(); + } + return map; +} + +Variant ToolConfigToVariant(const ToolConfig& config) { + Variant map = Variant::EmptyMap(); + if (config.function_calling_config().has_value()) { + const FunctionCallingConfig& fcc = config.function_calling_config().value(); + Variant fcc_map = Variant::EmptyMap(); + switch (fcc.mode()) { + case FunctionCallingConfig::kModeAuto: + fcc_map.map()["mode"] = "AUTO"; + break; + case FunctionCallingConfig::kModeAny: + fcc_map.map()["mode"] = "ANY"; + break; + case FunctionCallingConfig::kModeNone: + fcc_map.map()["mode"] = "NONE"; + break; + case FunctionCallingConfig::kModeUnspecified: + default: + break; + } + if (!fcc.allowed_function_names().empty()) { + Variant names = Variant::EmptyVector(); + for (const auto& name : fcc.allowed_function_names()) { + names.vector().push_back(name); + } + fcc_map.map()["allowedFunctionNames"] = names; + } + map.map()["functionCallingConfig"] = fcc_map; + } + if (config.retrieval_config().has_value()) { + const RetrievalConfig& rc = config.retrieval_config().value(); + Variant rc_map = Variant::EmptyMap(); + if (rc.lat_lng.has_value()) { + Variant ll = Variant::EmptyMap(); + ll.map()["latitude"] = rc.lat_lng.value().latitude; + ll.map()["longitude"] = rc.lat_lng.value().longitude; + rc_map.map()["latLng"] = ll; + } + if (rc.language_code.has_value()) { + rc_map.map()["languageCode"] = rc.language_code.value(); + } + map.map()["retrievalConfig"] = rc_map; + } + return map; +} + +namespace { + +Variant BuildGenerateContentRequestVariant( + const std::vector& contents, + const Optional& generation_config, + const std::vector& safety_settings, + const std::vector& tools, const Optional& tool_config, + const Optional& system_instruction, + BackendProvider provider) { + Variant root = Variant::EmptyMap(); + + Variant contents_vec = Variant::EmptyVector(); + for (const auto& c : contents) { + contents_vec.vector().push_back(ModelContentToVariant(c)); + } + root.map()["contents"] = contents_vec; + + if (generation_config.has_value()) { + Variant gc = GenerationConfigToVariant(generation_config.value()); + if (!gc.map().empty()) { + root.map()["generationConfig"] = gc; + } + } + + if (!safety_settings.empty()) { + Variant settings_vec = Variant::EmptyVector(); + for (const auto& s : safety_settings) { + settings_vec.vector().push_back(SafetySettingToVariant(s, provider)); + } + root.map()["safetySettings"] = settings_vec; + } + + if (!tools.empty()) { + Variant tools_vec = Variant::EmptyVector(); + for (const auto& t : tools) { + tools_vec.vector().push_back(ToolToVariant(t)); + } + root.map()["tools"] = tools_vec; + } + + if (tool_config.has_value()) { + Variant tc = ToolConfigToVariant(tool_config.value()); + if (!tc.map().empty()) { + root.map()["toolConfig"] = tc; + } + } + + if (system_instruction.has_value()) { + root.map()["systemInstruction"] = + ModelContentToVariant(system_instruction.value()); + } + + return root; +} + +} // namespace + +std::string BuildGenerateContentRequestJson( + const std::vector& contents, + const Optional& generation_config, + const std::vector& safety_settings, + const std::vector& tools, const Optional& tool_config, + const Optional& system_instruction, + BackendProvider provider) { + Variant root = BuildGenerateContentRequestVariant( + contents, generation_config, safety_settings, tools, tool_config, + system_instruction, provider); + return ::firebase::util::VariantToJson(root); +} + +std::string BuildCountTokensRequestJson( + const std::string& model_name, const std::vector& contents, + const Optional& generation_config, + const std::vector& safety_settings, + const std::vector& tools, const Optional& tool_config, + const Optional& system_instruction, + BackendProvider provider) { + if (provider == kBackendProviderGoogleAI) { + Variant inner = BuildGenerateContentRequestVariant( + contents, generation_config, safety_settings, tools, tool_config, + system_instruction, provider); + const char kPrefix[] = "models/"; + std::string normalized = + (model_name.compare(0, sizeof(kPrefix) - 1, kPrefix) == 0) + ? model_name + : (std::string(kPrefix) + model_name); + inner.map()["model"] = normalized; + Variant wrapper = Variant::EmptyMap(); + wrapper.map()["generateContentRequest"] = inner; + return ::firebase::util::VariantToJson(wrapper); + } + + Variant root = Variant::EmptyMap(); + Variant contents_vec = Variant::EmptyVector(); + for (const auto& c : contents) { + contents_vec.vector().push_back(ModelContentToVariant(c)); + } + root.map()["contents"] = contents_vec; + if (generation_config.has_value()) { + Variant gc = GenerationConfigToVariant(generation_config.value()); + if (!gc.map().empty()) { + root.map()["generationConfig"] = gc; + } + } + if (!tools.empty()) { + Variant tools_vec = Variant::EmptyVector(); + for (const auto& t : tools) { + tools_vec.vector().push_back(ToolToVariant(t)); + } + root.map()["tools"] = tools_vec; + } + if (system_instruction.has_value()) { + root.map()["systemInstruction"] = + ModelContentToVariant(system_instruction.value()); + } + return ::firebase::util::VariantToJson(root); +} + +std::string BuildTemplateGenerateContentRequestJson( + const std::map& inputs, + const std::vector& history) { + Variant root = Variant::EmptyMap(); + if (!inputs.empty()) { + Variant inputs_map = Variant::EmptyMap(); + for (const auto& kv : inputs) { + inputs_map.map()[kv.first] = kv.second; + } + root.map()["inputs"] = inputs_map; + } + if (!history.empty()) { + Variant history_vec = Variant::EmptyVector(); + for (const auto& c : history) { + history_vec.vector().push_back(ModelContentToVariant(c)); + } + root.map()["history"] = history_vec; + } + return ::firebase::util::VariantToJson(root); +} + +std::string BuildTemplateGenerateContentRequestFromRawJson( + const std::string& json_inputs, const std::vector& history) { + Variant parsed_inputs = ::firebase::util::JsonToVariant(json_inputs.c_str()); + Variant root = Variant::EmptyMap(); + if (parsed_inputs.is_map()) { + root.map()["inputs"] = parsed_inputs; + } else { + root.map()["inputs"] = Variant::EmptyMap(); + } + if (!history.empty()) { + Variant history_vec = Variant::EmptyVector(); + for (const auto& c : history) { + history_vec.vector().push_back(ModelContentToVariant(c)); + } + root.map()["history"] = history_vec; + } + return ::firebase::util::VariantToJson(root); +} + +bool ParseGenerateContentResponseJson(const std::string& json, + BackendProvider provider, + GenerateContentResponse* out_response, + std::string* out_error) { + if (!out_response) return false; + Variant root = ::firebase::util::JsonToVariant(json.c_str()); + if (!root.is_map()) { + if (out_error) { + *out_error = "Unable to parse GenerateContentResponse JSON object."; + } + return false; + } + + std::vector candidates; + const Variant* candidates_var = FindField(root, "candidates"); + if (candidates_var && candidates_var->is_vector()) { + for (const auto& cand_var : candidates_var->vector()) { + if (!cand_var.is_map()) continue; + Candidate cand; + const Variant* content_var = FindField(cand_var, "content"); + if (content_var && content_var->is_map()) { + ModelContentFromVariant(*content_var, &cand.content); + } else { + cand.content = ModelContent("model", std::vector()); + } + + const Variant* safety_var = FindField(cand_var, "safetyRatings"); + if (safety_var && safety_var->is_vector()) { + for (const auto& sr : safety_var->vector()) { + cand.safety_ratings.push_back(SafetyRatingFromVariant(sr)); + } + } + + const Variant* citation_var = FindField(cand_var, "citationMetadata"); + if (citation_var && citation_var->is_map()) { + cand.citation_metadata = ParseCitationMetadata(*citation_var, provider); + } + + const Variant* grounding_var = FindField(cand_var, "groundingMetadata"); + if (grounding_var && grounding_var->is_map()) { + cand.grounding_metadata = ParseGroundingMetadata(*grounding_var); + } + + const Variant* url_ctx_var = FindField(cand_var, "urlContextMetadata"); + if (url_ctx_var && url_ctx_var->is_map()) { + cand.url_context_metadata = ParseUrlContextMetadata(*url_ctx_var); + } + + cand.finish_reason = + ParseFinishReason(GetStringField(cand_var, "finishReason")); + cand.finish_message = GetStringField(cand_var, "finishMessage"); + candidates.push_back(cand); + } + } + + Optional prompt_feedback; + const Variant* pf_var = FindField(root, "promptFeedback"); + if (pf_var && pf_var->is_map()) { + PromptFeedback pf; + pf.block_reason = ParseBlockReason(GetStringField(*pf_var, "blockReason")); + pf.block_reason_message = GetStringField(*pf_var, "blockReasonMessage"); + const Variant* sr_var = FindField(*pf_var, "safetyRatings"); + if (sr_var && sr_var->is_vector()) { + for (const auto& sr : sr_var->vector()) { + pf.safety_ratings.push_back(SafetyRatingFromVariant(sr)); + } + } + prompt_feedback = pf; + } + + Optional usage_metadata; + const Variant* um_var = FindField(root, "usageMetadata"); + if (um_var && um_var->is_map()) { + UsageMetadata um; + um.prompt_token_count = GetIntField(*um_var, "promptTokenCount"); + um.candidates_token_count = GetIntField(*um_var, "candidatesTokenCount"); + um.total_token_count = GetIntField(*um_var, "totalTokenCount"); + um.thoughts_token_count = GetIntField(*um_var, "thoughtsTokenCount"); + um.tool_use_prompt_token_count = + GetIntField(*um_var, "toolUsePromptTokenCount"); + um.cached_content_token_count = + GetIntField(*um_var, "cachedContentTokenCount"); + um.prompt_tokens_details = + ParseModalityTokenCounts(FindField(*um_var, "promptTokensDetails")); + um.candidates_tokens_details = + ParseModalityTokenCounts(FindField(*um_var, "candidatesTokensDetails")); + um.tool_use_prompt_tokens_details = ParseModalityTokenCounts( + FindField(*um_var, "toolUsePromptTokensDetails")); + um.cache_tokens_details = + ParseModalityTokenCounts(FindField(*um_var, "cacheTokensDetails")); + usage_metadata = um; + } + + *out_response = + GenerateContentResponse(candidates, prompt_feedback, usage_metadata); + return true; +} + +bool ParseCountTokensResponseJson(const std::string& json, + CountTokensResponse* out_response, + std::string* out_error) { + if (!out_response) return false; + Variant root = ::firebase::util::JsonToVariant(json.c_str()); + if (!root.is_map()) { + if (out_error) { + *out_error = "Unable to parse CountTokensResponse JSON object."; + } + return false; + } + + CountTokensResponse resp; + resp.total_tokens = GetIntField(root, "totalTokens"); + resp.total_billable_characters = GetIntField(root, "totalBillableCharacters"); + resp.prompt_tokens_details = + ParseModalityTokenCounts(FindField(root, "promptTokensDetails")); + *out_response = resp; + return true; +} + +std::string ParseHttpErrorJson(int status_code, const std::string& body) { + std::ostringstream oss; + oss << "HTTP " << status_code; + if (!body.empty()) { + Variant root = ::firebase::util::JsonToVariant(body.c_str()); + if (root.is_map()) { + const Variant* err_var = FindField(root, "error"); + if (err_var && err_var->is_map()) { + std::string status = GetStringField(*err_var, "status"); + std::string message = GetStringField(*err_var, "message"); + if (!status.empty()) oss << " (" << status << ")"; + if (!message.empty()) { + oss << ": " << message; + return oss.str(); + } + } + } + oss << ": " << body; + } + return oss.str(); +} + +} // namespace internal +} // namespace ai +} // namespace firebase diff --git a/ai/src/common/serialization.h b/ai/src/common/serialization.h new file mode 100644 index 0000000000..fb36d55edd --- /dev/null +++ b/ai/src/common/serialization.h @@ -0,0 +1,119 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_COMMON_SERIALIZATION_H_ +#define FIREBASE_AI_SRC_COMMON_SERIALIZATION_H_ + +#include +#include +#include + +#include "firebase/ai/function_calling.h" +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generation_config.h" +#include "firebase/ai/model_content.h" +#include "firebase/ai/safety.h" +#include "firebase/ai/schema.h" +#include "firebase/ai/types.h" +#include "firebase/variant.h" + +namespace firebase { +namespace ai { +namespace internal { + +/// @brief Converts an OpenAPI `Schema` to a `Variant` map. +Variant SchemaToVariant(const Schema& schema); + +/// @brief Converts a standard `JsonSchema` to a `Variant` map. +Variant JsonSchemaToVariant(const JsonSchema& schema); + +/// @brief Converts a `Part` to a `Variant` map. +Variant PartToVariant(const Part& part); + +/// @brief Parses a `Part` from a `Variant` map. +bool PartFromVariant(const Variant& variant, Part* out_part); + +/// @brief Converts a `ModelContent` to a `Variant` map. +Variant ModelContentToVariant(const ModelContent& content); + +/// @brief Parses a `ModelContent` from a `Variant` map. +bool ModelContentFromVariant(const Variant& variant, ModelContent* out_content); + +/// @brief Converts a `SafetySetting` to a `Variant` map for the given backend. +Variant SafetySettingToVariant(const SafetySetting& setting, + BackendProvider provider); + +/// @brief Parses a `SafetyRating` from a `Variant` map. +SafetyRating SafetyRatingFromVariant(const Variant& variant); + +/// @brief Converts a `GenerationConfig` to a `Variant` map. +Variant GenerationConfigToVariant(const GenerationConfig& config); + +/// @brief Converts a `Tool` to a `Variant` map. +Variant ToolToVariant(const Tool& tool); + +/// @brief Converts a `ToolConfig` to a `Variant` map. +Variant ToolConfigToVariant(const ToolConfig& config); + +/// @brief Builds the JSON request body for `:generateContent` or +/// `:streamGenerateContent`. +std::string BuildGenerateContentRequestJson( + const std::vector& contents, + const Optional& generation_config, + const std::vector& safety_settings, + const std::vector& tools, const Optional& tool_config, + const Optional& system_instruction, BackendProvider provider); + +/// @brief Builds the JSON request body for `:countTokens`. +std::string BuildCountTokensRequestJson( + const std::string& model_name, const std::vector& contents, + const Optional& generation_config, + const std::vector& safety_settings, + const std::vector& tools, const Optional& tool_config, + const Optional& system_instruction, BackendProvider provider); + +/// @brief Builds the JSON request body for `:templateGenerateContent` or +/// `:templateStreamGenerateContent`. +std::string BuildTemplateGenerateContentRequestJson( + const std::map& inputs, + const std::vector& history = std::vector()); + +/// @brief Builds the JSON request body for `:templateGenerateContent` from a +/// raw JSON inputs string. +std::string BuildTemplateGenerateContentRequestFromRawJson( + const std::string& json_inputs, + const std::vector& history = std::vector()); + +/// @brief Parses a `GenerateContentResponse` from a JSON string. +bool ParseGenerateContentResponseJson(const std::string& json, + BackendProvider provider, + GenerateContentResponse* out_response, + std::string* out_error); + +/// @brief Parses a `CountTokensResponse` from a JSON string. +bool ParseCountTokensResponseJson(const std::string& json, + CountTokensResponse* out_response, + std::string* out_error); + +/// @brief Extracts a human-readable error message from a Google API JSON error +/// response body. +std::string ParseHttpErrorJson(int status_code, const std::string& body); + +} // namespace internal +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_COMMON_SERIALIZATION_H_ diff --git a/ai/src/common/template_generative_model.cc b/ai/src/common/template_generative_model.cc new file mode 100644 index 0000000000..9bd1269f51 --- /dev/null +++ b/ai/src/common/template_generative_model.cc @@ -0,0 +1,463 @@ +/* + * Copyright 2025 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. + */ + +#include "firebase/ai/template_generative_model.h" + +#include "ai/src/common/http_client.h" +#include "ai/src/common/serialization.h" +#include "ai/src/common/template_generative_model_internal.h" +#include "firebase/ai/template_chat_session.h" + +namespace firebase { +namespace ai { +namespace internal { + +namespace { + +struct TemplateStreamTurnAggregator { + std::vector accumulated_parts; + + void AddChunk(const GenerateContentResponse& chunk) { + if (chunk.candidates().empty()) return; + const ModelContent& content = chunk.candidates()[0].content; + for (const auto& part : content.parts()) { + if (part.is_text() && !accumulated_parts.empty()) { + Part& last = accumulated_parts.back(); + if (last.is_text() && last.is_thought() == part.is_thought() && + !last.thought_signature().has_value() && + !part.thought_signature().has_value()) { + std::string combined = last.text_part().text + part.text_part().text; + last = Part(TextPart(combined), last.is_thought(), + last.thought_signature()); + continue; + } + } + accumulated_parts.push_back(part); + } + } + + ModelContent BuildModelTurn() const { + return ModelContent("model", accumulated_parts); + } +}; + +} // namespace + +TemplateGenerativeModelInternal::TemplateGenerativeModelInternal( + ::firebase::App* app, const Backend& backend, + const Optional& request_options) + : app_(app), + backend_(backend), + request_options_(request_options.value_or(RequestOptions())), + future_impl_( + new ReferenceCountedFutureImpl(kTemplateGenerativeModelFnCount)) {} + +TemplateGenerativeModelInternal::~TemplateGenerativeModelInternal() {} + +Future +TemplateGenerativeModelInternal::GenerateContent( + const std::string& template_id, + const std::map& inputs, + const std::vector& history) { + std::string body = BuildTemplateGenerateContentRequestJson(inputs, history); + return ExecuteGenerateContentWithBody(template_id, body); +} + +Future +TemplateGenerativeModelInternal::GenerateContentJson( + const std::string& template_id, const std::string& json_inputs, + const std::vector& history) { + std::string body = + BuildTemplateGenerateContentRequestFromRawJson(json_inputs, history); + return ExecuteGenerateContentWithBody(template_id, body); +} + +Future +TemplateGenerativeModelInternal::ExecuteGenerateContentWithBody( + const std::string& template_id, const std::string& body) { + SafeFutureHandle handle = + future_impl_->SafeAlloc( + kTemplateGenerativeModelFnGenerateContent); + + if (!app_ || template_id.empty()) { + future_impl_->Complete(handle, kErrorInvalidArgument, + "Template ID must not be empty."); + return MakeFuture(future_impl_.get(), handle); + } + + std::string url = AiHttpClient::ConstructTemplateUrl( + app_, backend_, template_id, "templateGenerateContent"); + std::shared_ptr future_impl = future_impl_; + BackendProvider provider = backend_.provider(); + + AiHttpClient::SendUnaryJson( + app_, request_options_, url, body, + [future_impl, handle, provider](Error err, const std::string& err_msg, + const std::string& response_body) { + if (err != kErrorNone) { + future_impl->Complete(handle, err, err_msg.c_str()); + return; + } + GenerateContentResponse parsed; + std::string parse_err; + if (!ParseGenerateContentResponseJson(response_body, provider, &parsed, + &parse_err)) { + future_impl->Complete(handle, kErrorSerializationFailed, + parse_err.c_str()); + return; + } + future_impl->CompleteWithResult(handle, kErrorNone, "", parsed); + }); + + return MakeFuture(future_impl_.get(), handle); +} + +Future +TemplateGenerativeModelInternal::GenerateContentLastResult() const { + return static_cast&>( + future_impl_->LastResult(kTemplateGenerativeModelFnGenerateContent)); +} + +Future TemplateGenerativeModelInternal::GenerateContentStream( + const std::string& template_id, + const std::map& inputs, + const GenerateContentStreamCallback& on_chunk, + const std::vector& history) { + std::string body = BuildTemplateGenerateContentRequestJson(inputs, history); + return ExecuteGenerateContentStreamWithBody(template_id, body, on_chunk); +} + +Future TemplateGenerativeModelInternal::GenerateContentStreamJson( + const std::string& template_id, const std::string& json_inputs, + const GenerateContentStreamCallback& on_chunk, + const std::vector& history) { + std::string body = + BuildTemplateGenerateContentRequestFromRawJson(json_inputs, history); + return ExecuteGenerateContentStreamWithBody(template_id, body, on_chunk); +} + +Future +TemplateGenerativeModelInternal::ExecuteGenerateContentStreamWithBody( + const std::string& template_id, const std::string& body, + const GenerateContentStreamCallback& on_chunk) { + SafeFutureHandle handle = future_impl_->SafeAlloc( + kTemplateGenerativeModelFnGenerateContentStream); + + if (!app_ || template_id.empty()) { + future_impl_->Complete(handle, kErrorInvalidArgument, + "Template ID must not be empty."); + return MakeFuture(future_impl_.get(), handle); + } + + std::string url = AiHttpClient::ConstructTemplateUrl( + app_, backend_, template_id, "templateStreamGenerateContent?alt=sse"); + std::shared_ptr future_impl = future_impl_; + + AiHttpClient::SendStreamJson( + app_, backend_, request_options_, url, body, on_chunk, + [future_impl, handle](Error err, const std::string& err_msg, + const std::string& /*response_body*/) { + future_impl->Complete(handle, err, err_msg.c_str()); + }); + + return MakeFuture(future_impl_.get(), handle); +} + +Future TemplateGenerativeModelInternal::GenerateContentStreamLastResult() + const { + return static_cast&>(future_impl_->LastResult( + kTemplateGenerativeModelFnGenerateContentStream)); +} + +// --- TemplateChatSessionInternal --- + +TemplateChatSessionInternal::TemplateChatSessionInternal( + const std::shared_ptr& model, + const std::string& template_id, + const std::map& inputs, + const std::vector& initial_history) + : model_(model), + template_id_(template_id), + inputs_(inputs), + history_(initial_history), + future_impl_( + new ReferenceCountedFutureImpl(kTemplateChatSessionFnCount)) {} + +TemplateChatSessionInternal::~TemplateChatSessionInternal() {} + +std::vector TemplateChatSessionInternal::history() const { + MutexLock lock(history_mutex_); + return history_; +} + +void TemplateChatSessionInternal::AppendHistoryTurn( + const std::vector& request_turns, + const ModelContent& response_turn) { + MutexLock lock(history_mutex_); + history_.insert(history_.end(), request_turns.begin(), request_turns.end()); + history_.push_back(response_turn); +} + +Future TemplateChatSessionInternal::SendMessage( + const std::vector& content) { + SafeFutureHandle handle = + future_impl_->SafeAlloc( + kTemplateChatSessionFnSendMessage); + + if (!model_ || content.empty()) { + future_impl_->Complete(handle, kErrorInvalidArgument, + "Chat message content must not be empty."); + return MakeFuture(future_impl_.get(), handle); + } + + std::vector full_history; + { + MutexLock lock(history_mutex_); + full_history = history_; + } + full_history.insert(full_history.end(), content.begin(), content.end()); + + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + + Future inner = + model_->GenerateContent(template_id_, inputs_, full_history); + inner.OnCompletion([self, future_impl, handle, content]( + const Future& completed) { + if (completed.error() != kErrorNone || completed.result() == nullptr) { + future_impl->Complete( + handle, completed.error(), + completed.error_message() ? completed.error_message() : ""); + return; + } + const GenerateContentResponse& resp = *completed.result(); + if (!resp.candidates().empty()) { + ModelContent model_turn = resp.candidates()[0].content; + if (model_turn.role().empty()) model_turn.set_role("model"); + self->AppendHistoryTurn(content, model_turn); + } + future_impl->CompleteWithResult(handle, kErrorNone, "", resp); + }); + + return MakeFuture(future_impl_.get(), handle); +} + +Future +TemplateChatSessionInternal::SendMessageLastResult() const { + return static_cast&>( + future_impl_->LastResult(kTemplateChatSessionFnSendMessage)); +} + +Future TemplateChatSessionInternal::SendMessageStream( + const std::vector& content, + const GenerateContentStreamCallback& on_chunk) { + SafeFutureHandle handle = + future_impl_->SafeAlloc(kTemplateChatSessionFnSendMessageStream); + + if (!model_ || content.empty()) { + future_impl_->Complete(handle, kErrorInvalidArgument, + "Chat message content must not be empty."); + return MakeFuture(future_impl_.get(), handle); + } + + std::vector full_history; + { + MutexLock lock(history_mutex_); + full_history = history_; + } + full_history.insert(full_history.end(), content.begin(), content.end()); + + std::shared_ptr self = shared_from_this(); + std::shared_ptr future_impl = future_impl_; + std::shared_ptr aggregator( + new TemplateStreamTurnAggregator()); + + Future inner = model_->GenerateContentStream( + template_id_, inputs_, + [aggregator, on_chunk](const GenerateContentResponse& chunk) { + aggregator->AddChunk(chunk); + if (on_chunk) on_chunk(chunk); + }, + full_history); + + inner.OnCompletion([self, future_impl, handle, content, + aggregator](const Future& completed) { + if (completed.error() == kErrorNone && + !aggregator->accumulated_parts.empty()) { + self->AppendHistoryTurn(content, aggregator->BuildModelTurn()); + } + future_impl->Complete( + handle, completed.error(), + completed.error_message() ? completed.error_message() : ""); + }); + + return MakeFuture(future_impl_.get(), handle); +} + +Future TemplateChatSessionInternal::SendMessageStreamLastResult() const { + return static_cast&>( + future_impl_->LastResult(kTemplateChatSessionFnSendMessageStream)); +} + +} // namespace internal + +// --- TemplateGenerativeModel public implementation --- + +TemplateGenerativeModel::TemplateGenerativeModel() : internal_(nullptr) {} + +TemplateGenerativeModel::TemplateGenerativeModel( + const std::shared_ptr& internal) + : internal_(internal) {} + +TemplateGenerativeModel::TemplateGenerativeModel( + const TemplateGenerativeModel& other) + : internal_(other.internal_) {} + +TemplateGenerativeModel& TemplateGenerativeModel::operator=( + const TemplateGenerativeModel& other) { + if (this != &other) { + internal_ = other.internal_; + } + return *this; +} + +TemplateGenerativeModel::~TemplateGenerativeModel() {} + +Future TemplateGenerativeModel::GenerateContent( + const std::string& template_id, + const std::map& inputs) { + if (!internal_) return Future(); + return internal_->GenerateContent(template_id, inputs); +} + +Future TemplateGenerativeModel::GenerateContentJson( + const std::string& template_id, const std::string& json_inputs) { + if (!internal_) return Future(); + return internal_->GenerateContentJson(template_id, json_inputs); +} + +Future +TemplateGenerativeModel::GenerateContentLastResult() const { + if (!internal_) return Future(); + return internal_->GenerateContentLastResult(); +} + +Future TemplateGenerativeModel::GenerateContentStream( + const std::string& template_id, + const std::map& inputs, + const GenerateContentStreamCallback& on_chunk) { + if (!internal_) return Future(); + return internal_->GenerateContentStream(template_id, inputs, on_chunk); +} + +Future TemplateGenerativeModel::GenerateContentStreamJson( + const std::string& template_id, const std::string& json_inputs, + const GenerateContentStreamCallback& on_chunk) { + if (!internal_) return Future(); + return internal_->GenerateContentStreamJson(template_id, json_inputs, + on_chunk); +} + +Future TemplateGenerativeModel::GenerateContentStreamLastResult() const { + if (!internal_) return Future(); + return internal_->GenerateContentStreamLastResult(); +} + +TemplateChatSession TemplateGenerativeModel::StartChat( + const std::string& template_id, + const std::map& inputs, + const std::vector& history) const { + if (!internal_) return TemplateChatSession(); + return TemplateChatSession( + std::shared_ptr( + new internal::TemplateChatSessionInternal(internal_, template_id, + inputs, history))); +} + +// --- TemplateChatSession public implementation --- + +TemplateChatSession::TemplateChatSession() : internal_(nullptr) {} + +TemplateChatSession::TemplateChatSession( + const std::shared_ptr& internal) + : internal_(internal) {} + +TemplateChatSession::TemplateChatSession(const TemplateChatSession& other) + : internal_(other.internal_) {} + +TemplateChatSession& TemplateChatSession::operator=( + const TemplateChatSession& other) { + if (this != &other) { + internal_ = other.internal_; + } + return *this; +} + +TemplateChatSession::~TemplateChatSession() {} + +std::vector TemplateChatSession::history() const { + if (!internal_) return std::vector(); + return internal_->history(); +} + +Future TemplateChatSession::SendMessage( + const std::string& prompt) { + return SendMessage(std::vector(1, ModelContent::Text(prompt))); +} + +Future TemplateChatSession::SendMessage( + const ModelContent& content) { + return SendMessage(std::vector(1, content)); +} + +Future TemplateChatSession::SendMessage( + const std::vector& content) { + if (!internal_) return Future(); + return internal_->SendMessage(content); +} + +Future TemplateChatSession::SendMessageLastResult() + const { + if (!internal_) return Future(); + return internal_->SendMessageLastResult(); +} + +Future TemplateChatSession::SendMessageStream( + const std::string& prompt, const GenerateContentStreamCallback& on_chunk) { + return SendMessageStream( + std::vector(1, ModelContent::Text(prompt)), on_chunk); +} + +Future TemplateChatSession::SendMessageStream( + const ModelContent& content, + const GenerateContentStreamCallback& on_chunk) { + return SendMessageStream(std::vector(1, content), on_chunk); +} + +Future TemplateChatSession::SendMessageStream( + const std::vector& content, + const GenerateContentStreamCallback& on_chunk) { + if (!internal_) return Future(); + return internal_->SendMessageStream(content, on_chunk); +} + +Future TemplateChatSession::SendMessageStreamLastResult() const { + if (!internal_) return Future(); + return internal_->SendMessageStreamLastResult(); +} + +} // namespace ai +} // namespace firebase diff --git a/ai/src/common/template_generative_model_internal.h b/ai/src/common/template_generative_model_internal.h new file mode 100644 index 0000000000..9d0866eb98 --- /dev/null +++ b/ai/src/common/template_generative_model_internal.h @@ -0,0 +1,133 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_COMMON_TEMPLATE_GENERATIVE_MODEL_INTERNAL_H_ +#define FIREBASE_AI_SRC_COMMON_TEMPLATE_GENERATIVE_MODEL_INTERNAL_H_ + +#include +#include +#include +#include + +#include "app/src/include/firebase/internal/mutex.h" +#include "app/src/reference_counted_future_impl.h" +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generative_model.h" +#include "firebase/ai/model_content.h" +#include "firebase/ai/template_chat_session.h" +#include "firebase/ai/template_generative_model.h" +#include "firebase/ai/types.h" +#include "firebase/app.h" +#include "firebase/future.h" +#include "firebase/variant.h" + +namespace firebase { +namespace ai { +namespace internal { + +enum TemplateGenerativeModelFn { + kTemplateGenerativeModelFnGenerateContent = 0, + kTemplateGenerativeModelFnGenerateContentStream, + kTemplateGenerativeModelFnCount +}; + +enum TemplateChatSessionFn { + kTemplateChatSessionFnSendMessage = 0, + kTemplateChatSessionFnSendMessageStream, + kTemplateChatSessionFnCount +}; + +class TemplateGenerativeModelInternal { + public: + TemplateGenerativeModelInternal( + ::firebase::App* app, const Backend& backend, + const Optional& request_options); + ~TemplateGenerativeModelInternal(); + + Future GenerateContent( + const std::string& template_id, + const std::map& inputs, + const std::vector& history = std::vector()); + + Future GenerateContentJson( + const std::string& template_id, const std::string& json_inputs, + const std::vector& history = std::vector()); + + Future GenerateContentLastResult() const; + + Future GenerateContentStream( + const std::string& template_id, + const std::map& inputs, + const GenerateContentStreamCallback& on_chunk, + const std::vector& history = std::vector()); + + Future GenerateContentStreamJson( + const std::string& template_id, const std::string& json_inputs, + const GenerateContentStreamCallback& on_chunk, + const std::vector& history = std::vector()); + + Future GenerateContentStreamLastResult() const; + + private: + Future ExecuteGenerateContentWithBody( + const std::string& template_id, const std::string& body); + Future ExecuteGenerateContentStreamWithBody( + const std::string& template_id, const std::string& body, + const GenerateContentStreamCallback& on_chunk); + + ::firebase::App* app_; + Backend backend_; + RequestOptions request_options_; + std::shared_ptr future_impl_; +}; + +class TemplateChatSessionInternal + : public std::enable_shared_from_this { + public: + TemplateChatSessionInternal( + const std::shared_ptr& model, + const std::string& template_id, + const std::map& inputs, + const std::vector& initial_history); + ~TemplateChatSessionInternal(); + + std::vector history() const; + + Future SendMessage( + const std::vector& content); + Future SendMessageLastResult() const; + + Future SendMessageStream(const std::vector& content, + const GenerateContentStreamCallback& on_chunk); + Future SendMessageStreamLastResult() const; + + private: + void AppendHistoryTurn(const std::vector& request_turns, + const ModelContent& response_turn); + + std::shared_ptr model_; + std::string template_id_; + std::map inputs_; + mutable Mutex history_mutex_; + std::vector history_; + std::shared_ptr future_impl_; +}; + +} // namespace internal +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_COMMON_TEMPLATE_GENERATIVE_MODEL_INTERNAL_H_ diff --git a/ai/src/desktop/http_sender_desktop.cc b/ai/src/desktop/http_sender_desktop.cc new file mode 100644 index 0000000000..01f90ccf7b --- /dev/null +++ b/ai/src/desktop/http_sender_desktop.cc @@ -0,0 +1,185 @@ +/* + * Copyright 2025 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. + */ + +#include +#include +#include +#include + +#include "ai/src/common/http_sender.h" +#include "app/rest/transport_curl.h" +#include "app/rest/util.h" +#include "curl/curl.h" + +namespace firebase { +namespace ai { +namespace internal { + +namespace { + +bool FileExists(const char* path) { + if (!path || path[0] == '\0') return false; + std::ifstream ifs(path, std::ios::binary); + return ifs.good(); +} + +const char* ResolveCaBundlePath() { + const char* env_bundle = std::getenv("CURL_CA_BUNDLE"); + if (FileExists(env_bundle)) return env_bundle; + const char* env_ssl = std::getenv("SSL_CERT_FILE"); + if (FileExists(env_ssl)) return env_ssl; + static const char* const kCandidates[] = { + "/etc/ssl/cert.pem", // macOS & FreeBSD + "/etc/ssl/certs/ca-certificates.crt", // Debian/Ubuntu + "/etc/pki/tls/certs/ca-bundle.crt", // RHEL/Fedora + }; + for (const char* candidate : kCandidates) { + if (FileExists(candidate)) return candidate; + } + return nullptr; +} + +struct CurlWriteContext { + CURL* curl = nullptr; + bool is_stream = false; + HttpStreamChunkCallback on_chunk; + std::string body; +}; + +size_t OnCurlWrite(char* ptr, size_t size, size_t nmemb, void* userdata) { + size_t total = size * nmemb; + if (total == 0 || !userdata) return 0; + CurlWriteContext* ctx = static_cast(userdata); + long http_code = 0; + if (ctx->curl) { + curl_easy_getinfo(ctx->curl, CURLINFO_RESPONSE_CODE, &http_code); + } + if (ctx->is_stream && http_code >= 200 && http_code < 300 && ctx->on_chunk) { + if (!ctx->on_chunk(ptr, total)) { + return 0; + } + return total; + } + ctx->body.append(ptr, total); + return total; +} + +void PerformCurlRequestAsync(const HttpRequest& req, + const HttpStreamChunkCallback& on_chunk, + const HttpCompletionCallback& on_complete) { + std::thread([req, on_chunk, on_complete]() { + CURL* curl = curl_easy_init(); + if (!curl) { + if (on_complete) { + on_complete(0, "", "Failed to initialize libcurl handle."); + } + return; + } + + char err_buf[CURL_ERROR_SIZE]; + err_buf[0] = '\0'; + curl_easy_setopt(curl, CURLOPT_ERRORBUFFER, err_buf); + curl_easy_setopt(curl, CURLOPT_URL, req.url.c_str()); + curl_easy_setopt(curl, CURLOPT_SSL_VERIFYPEER, 1L); + curl_easy_setopt(curl, CURLOPT_SSL_VERIFYHOST, 2L); + + const char* ca_bundle = ResolveCaBundlePath(); + if (ca_bundle) { + curl_easy_setopt(curl, CURLOPT_CAINFO, ca_bundle); + } + + if (req.timeout_ms > 0) { + curl_easy_setopt(curl, CURLOPT_TIMEOUT_MS, + static_cast(req.timeout_ms)); + } + + struct curl_slist* headers = nullptr; + for (const auto& kv : req.headers) { + std::string header_line = kv.first + ": " + kv.second; + headers = curl_slist_append(headers, header_line.c_str()); + } + if (headers) { + curl_easy_setopt(curl, CURLOPT_HTTPHEADER, headers); + } + + if (req.method == "POST") { + curl_easy_setopt(curl, CURLOPT_POST, 1L); + curl_easy_setopt(curl, CURLOPT_POSTFIELDS, req.body.data()); + curl_easy_setopt(curl, CURLOPT_POSTFIELDSIZE, + static_cast(req.body.size())); + } else if (req.method != "GET" && !req.method.empty()) { + curl_easy_setopt(curl, CURLOPT_CUSTOMREQUEST, req.method.c_str()); + } + + CurlWriteContext write_ctx; + write_ctx.curl = curl; + write_ctx.is_stream = static_cast(on_chunk); + write_ctx.on_chunk = on_chunk; + curl_easy_setopt(curl, CURLOPT_WRITEFUNCTION, OnCurlWrite); + curl_easy_setopt(curl, CURLOPT_WRITEDATA, &write_ctx); + + CURLcode res = curl_easy_perform(curl); + long http_status = 0; + curl_easy_getinfo(curl, CURLINFO_RESPONSE_CODE, &http_status); + + if (headers) { + curl_slist_free_all(headers); + } + curl_easy_cleanup(curl); + + std::string transport_err; + if (res != CURLE_OK && !(write_ctx.is_stream && res == CURLE_WRITE_ERROR && + http_status >= 200 && http_status < 300)) { + transport_err = + std::string("libcurl error (") + curl_easy_strerror(res) + ")"; + if (err_buf[0] != '\0') { + transport_err += std::string(": ") + err_buf; + } + } + + if (on_complete) { + on_complete(static_cast(http_status), write_ctx.body, transport_err); + } + }).detach(); +} + +} // namespace + +void HttpSender::Initialize() { + ::firebase::rest::InitTransportCurl(); + ::firebase::rest::util::Initialize(); +} + +void HttpSender::Cleanup() { + ::firebase::rest::util::Terminate(); + ::firebase::rest::CleanupTransportCurl(); +} + +void HttpSender::SendUnary(::firebase::App* /*app*/, const HttpRequest& request, + const HttpCompletionCallback& on_complete) { + PerformCurlRequestAsync(request, HttpStreamChunkCallback(), on_complete); +} + +void HttpSender::SendStream(::firebase::App* /*app*/, + const HttpRequest& request, + const HttpStreamChunkCallback& on_chunk, + const HttpCompletionCallback& on_complete) { + PerformCurlRequestAsync(request, on_chunk, on_complete); +} + +} // namespace internal +} // namespace ai +} // namespace firebase diff --git a/ai/src/include/firebase/ai.h b/ai/src/include/firebase/ai.h new file mode 100644 index 0000000000..21da6a095f --- /dev/null +++ b/ai/src/include/firebase/ai.h @@ -0,0 +1,154 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_H_ + +#include +#include + +#include "firebase/ai/chat.h" +#include "firebase/ai/function_calling.h" +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generation_config.h" +#include "firebase/ai/generative_model.h" +#include "firebase/ai/model_content.h" +#include "firebase/ai/safety.h" +#include "firebase/ai/schema.h" +#include "firebase/ai/template_chat_session.h" +#include "firebase/ai/template_generative_model.h" +#include "firebase/ai/types.h" +#include "firebase/app.h" +#include "firebase/internal/common.h" + +FIREBASE_APP_REGISTER_CALLBACKS_REFERENCE(ai) + +namespace firebase { + +/// @brief Namespace for the Firebase AI Logic C++ SDK. +namespace ai { + +namespace internal { +class FirebaseAIInternal; +} // namespace internal + +/// @brief Entry point for all Firebase AI Logic functionality. +/// +/// Mirrors `Firebase.AI.FirebaseAI` in the Unity SDK and `FirebaseAI` in the +/// Flutter SDK. +class FirebaseAI { + public: + /// @brief Destructor. You may delete an instance of `FirebaseAI` when + /// finished using it; all instances are also automatically cleaned up when + /// the owning `firebase::App` is destroyed. + ~FirebaseAI(); + + /// @brief Gets the `FirebaseAI` instance for the default `firebase::App` and + /// `Backend` (defaults to `Backend::GoogleAI()`). + /// + /// @param backend The backend provider to use (`Backend::GoogleAI()` or + /// `Backend::Enterprise(location)` / `Backend::VertexAI(location)`). + /// @return Pointer to the `FirebaseAI` instance, or `nullptr` if the default + /// `App` does not exist. + static FirebaseAI* GetInstance(const Backend& backend = Backend::GoogleAI()); + + /// @brief Gets the `FirebaseAI` instance for the specified `firebase::App` + /// and `Backend`. + /// + /// @param app The `firebase::App` instance to use. + /// @param backend The backend provider to use (`Backend::GoogleAI()` or + /// `Backend::Enterprise(location)` / `Backend::VertexAI(location)`). + /// @return Pointer to the `FirebaseAI` instance, or `nullptr` if `app` is + /// null. + static FirebaseAI* GetInstance(::firebase::App* app, + const Backend& backend = Backend::GoogleAI()); + + /// @brief Returns the `firebase::App` that this `FirebaseAI` instance is + /// associated with. + ::firebase::App* app(); + + /// @brief Returns the `firebase::App` that this `FirebaseAI` instance is + /// associated with (const overload). + const ::firebase::App* app() const; + + /// @brief Returns the `Backend` configuration for this `FirebaseAI` instance. + const Backend& backend() const; + + /// @brief Initializes a `GenerativeModel` with the given parameters. + /// + /// @param model_name The name of the Gemini model to use (for example, + /// `"gemini-2.5-flash"`). + /// @param generation_config Optional content generation configuration. + /// @param safety_settings Optional safety filtering thresholds. + /// @param tools Optional list of tools (`FunctionDeclaration`, + /// `GoogleSearch`, `CodeExecution`, `GoogleMaps`, `UrlContext`) the model may + /// use. + /// @param tool_config Optional tool configuration (`FunctionCallingConfig`, + /// `RetrievalConfig`). + /// @param system_instruction Optional system instruction (`ModelContent`) + /// guiding the model's behavior. + /// @param request_options Optional per-request options (such as timeout and + /// limited-use App Check tokens). + /// @param hybrid_params Optional hybrid on-device + cloud LiteRT + /// configuration (`InferenceMode` and `OnDeviceParams`). + /// @return The initialized `GenerativeModel` instance. + GenerativeModel GetGenerativeModel( + const std::string& model_name, + const Optional& generation_config = + Optional(), + const std::vector& safety_settings = + std::vector(), + const std::vector& tools = std::vector(), + const Optional& tool_config = Optional(), + const Optional& system_instruction = + Optional(), + const Optional& request_options = + Optional(), + const Optional& hybrid_params = Optional()); + + /// @brief Convenience overload to initialize a hybrid `GenerativeModel` with + /// `HybridParams` (LiteRT on-device + cloud inference). + GenerativeModel GetGenerativeModel( + const std::string& model_name, const HybridParams& hybrid_params, + const Optional& generation_config = + Optional(), + const Optional& system_instruction = + Optional()); + + /// @brief Initializes a `TemplateGenerativeModel` for executing server prompt + /// templates. + /// + /// @param request_options Optional per-request options (such as timeout and + /// limited-use App Check tokens). + /// @return The initialized `TemplateGenerativeModel` instance. + TemplateGenerativeModel GetTemplateGenerativeModel( + const Optional& request_options = + Optional()); + + private: + FirebaseAI(::firebase::App* app, const Backend& backend); + FirebaseAI(const FirebaseAI&) = delete; + FirebaseAI& operator=(const FirebaseAI&) = delete; + + void DeleteInternal(); + + internal::FirebaseAIInternal* internal_; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_H_ diff --git a/ai/src/include/firebase/ai/chat.h b/ai/src/include/firebase/ai/chat.h new file mode 100644 index 0000000000..b8c5c7424c --- /dev/null +++ b/ai/src/include/firebase/ai/chat.h @@ -0,0 +1,159 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_CHAT_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_CHAT_H_ + +#include +#include +#include + +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generative_model.h" +#include "firebase/ai/model_content.h" +#include "firebase/future.h" + +namespace firebase { +namespace ai { + +namespace internal { +class ChatInternal; +} // namespace internal + +/// @brief An object that represents a back-and-forth multi-turn conversation +/// with a `GenerativeModel`, capturing the history of messages sent and +/// received. +/// +/// Mirrors `Firebase.AI.Chat` in Unity and `ChatSession` in Flutter. +class Chat { + public: + /// @brief Default constructor creates an invalid `Chat`. + Chat(); + + /// @brief Copy constructor. + Chat(const Chat& other); + + /// @brief Copy assignment operator. + Chat& operator=(const Chat& other); + + /// @brief Destructor. + ~Chat(); + + /// @brief Returns true if this `Chat` instance is valid. + bool is_valid() const { return internal_ != nullptr; } + + /// @brief Returns the conversation history accumulated in this `Chat` + /// session. + std::vector history() const; + + /// @brief Sends a text message to the model within the conversation context + /// and appends both the user turn and the model's response turn to + /// `history()`. + /// + /// @param prompt The user text message. + /// @return A `Future` containing the `GenerateContentResponse`. + Future SendMessage(const std::string& prompt); + + /// @brief Sends a single `ModelContent` message within the conversation + /// context and updates `history()`. + /// + /// @param content The user `ModelContent` turn. + /// @return A `Future` containing the `GenerateContentResponse`. + Future SendMessage(const ModelContent& content); + + /// @brief Sends multiple `ModelContent` messages within the conversation + /// context and updates `history()`. + /// + /// @param content The user `ModelContent` turns. + /// @return A `Future` containing the `GenerateContentResponse`. + Future SendMessage( + const std::vector& content); + + /// @brief Gets the result of the most recent `SendMessage` call. + /// + /// @return A `Future` from the most recent `SendMessage` call. + Future SendMessageLastResult() const; + + /// @brief Sends a text message to the model and streams back response chunks + /// via `on_chunk`, appending the aggregated response turn to `history()` upon + /// completion. + /// + /// @param prompt The user text message. + /// @param on_chunk Callback invoked for each response chunk. + /// @return A `Future` that completes when the stream finishes or fails. + Future SendMessageStream(const std::string& prompt, + const GenerateContentStreamCallback& on_chunk); + + /// @brief Sends a single `ModelContent` message and streams back response + /// chunks via `on_chunk`, appending the aggregated response turn to + /// `history()` upon completion. + /// + /// @param content The user `ModelContent` turn. + /// @param on_chunk Callback invoked for each response chunk. + /// @return A `Future` that completes when the stream finishes or fails. + Future SendMessageStream(const ModelContent& content, + const GenerateContentStreamCallback& on_chunk); + + /// @brief Sends multiple `ModelContent` messages and streams back response + /// chunks via `on_chunk`, appending the aggregated response turn to + /// `history()` upon completion. + /// + /// @param content The user `ModelContent` turns. + /// @param on_chunk Callback invoked for each response chunk. + /// @return A `Future` that completes when the stream finishes or fails. + Future SendMessageStream(const std::vector& content, + const GenerateContentStreamCallback& on_chunk); + + /// @brief Gets the result of the most recent `SendMessageStream` call. + /// + /// @return A `Future` from the most recent `SendMessageStream` call. + Future SendMessageStreamLastResult() const; + + /// @brief Returns the active `InferenceMode` of the underlying + /// `GenerativeModel`. + InferenceMode inference_mode() const; + + /// @brief Dynamically updates the `InferenceMode` of the underlying + /// `GenerativeModel` mid-conversation, preserving `history()` across cloud + /// and on-device turns. + void set_inference_mode(InferenceMode mode); + + /// @brief Returns true if the underlying on-device LiteRT model is available. + bool IsOnDeviceAvailable() const; + + /// @brief Clears all accumulated conversation turns from `history()`. + void ClearHistory(); + + /// @brief Summarizes and compacts the accumulated `history()` using the + /// underlying `GenerativeModel` (in its active `InferenceMode`), replacing + /// older turns with a concise summary turn so multi-turn conversations never + /// exhaust the context window. + /// + /// @return A `Future` containing the summary `GenerateContentResponse`. + Future CompactHistory(); + + private: + friend class GenerativeModel; + + explicit Chat(const std::shared_ptr& internal); + + std::shared_ptr internal_; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_CHAT_H_ diff --git a/ai/src/include/firebase/ai/function_calling.h b/ai/src/include/firebase/ai/function_calling.h new file mode 100644 index 0000000000..237b81dea5 --- /dev/null +++ b/ai/src/include/firebase/ai/function_calling.h @@ -0,0 +1,282 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_FUNCTION_CALLING_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_FUNCTION_CALLING_H_ + +#include +#include +#include + +#include "firebase/ai/schema.h" +#include "firebase/ai/types.h" + +namespace firebase { +namespace ai { + +/// @brief Structured representation of a function declaration that the model +/// may invoke. +/// +/// Mirrors `Firebase.AI.FunctionDeclaration` in Unity and +/// `FunctionDeclaration` in Flutter. +class FunctionDeclaration { + public: + /// @brief Default constructor. + FunctionDeclaration() : uses_json_schema_(false) {} + + /// @brief Construct a `FunctionDeclaration` using OpenAPI `Schema` parameter + /// definitions. + /// + /// @param name The name of the function. Must be a-z, A-Z, 0-9, or contain + /// underscores and dashes, with a maximum length of 63. + /// @param description A brief description of what the function does. + /// @param parameters Map of parameter names to their `Schema` definitions. + /// @param optional_parameters List of parameter names that are optional. Any + /// parameter in `parameters` not listed here is marked as required. + FunctionDeclaration(const std::string& name, const std::string& description, + const std::map& parameters, + const std::vector& optional_parameters = + std::vector()); + + /// @brief Construct a `FunctionDeclaration` using a full `JsonSchema` object + /// for its parameters (`parametersJsonSchema`). + /// + /// @param name The name of the function. + /// @param description A brief description of what the function does. + /// @param parameters_json_schema The root object `JsonSchema` describing the + /// function parameters. + FunctionDeclaration(const std::string& name, const std::string& description, + const JsonSchema& parameters_json_schema); + + /// @brief Returns the function name. + const std::string& name() const { return name_; } + /// @brief Sets the function name. + void set_name(const std::string& name) { name_ = name; } + + /// @brief Returns the function description. + const std::string& description() const { return description_; } + /// @brief Sets the function description. + void set_description(const std::string& description) { + description_ = description; + } + + /// @brief Returns the parameter schema. + const Schema& parameters() const { return parameters_; } + /// @brief Sets the parameter schema. + void set_parameters(const Schema& parameters) { parameters_ = parameters; } + + /// @brief Returns true if this declaration uses `parametersJsonSchema` rather + /// than OpenAPI `parameters`. + bool uses_json_schema() const { return uses_json_schema_; } + /// @brief Sets whether this declaration uses `parametersJsonSchema`. + void set_uses_json_schema(bool uses_json_schema) { + uses_json_schema_ = uses_json_schema; + } + + private: + std::string name_; + std::string description_; + Schema parameters_; + bool uses_json_schema_; +}; + +/// @brief Tool that enables the model to ground its responses using Google +/// Search. +struct GoogleSearch {}; + +/// @brief Tool that enables the model to generate and execute Python code on +/// the backend. +struct CodeExecution {}; + +/// @brief Tool that enables the model to ground its responses using Google +/// Maps. +struct GoogleMaps {}; + +/// @brief Tool that enables the model to retrieve and ground responses on +/// public web URLs provided in the prompt. +struct UrlContext {}; + +/// @brief A helper tool that the model may use when generating responses, such +/// as function declarations, Google Search grounding, Code Execution, Google +/// Maps, or URL Context. +/// +/// Mirrors `Firebase.AI.Tool` in Unity and `Tool` in Flutter. +class Tool { + public: + /// @brief Default constructor creates an empty Tool. + Tool() {} + + /// @brief Construct a Tool containing function declarations. + /// + /// @param function_declarations The functions to expose to the model. + explicit Tool(const std::vector& function_declarations) + : function_declarations_(function_declarations) {} + + /// @brief Construct a Tool containing a single function declaration. + /// + /// @param function_declaration The function to expose to the model. + explicit Tool(const FunctionDeclaration& function_declaration) + : function_declarations_(1, function_declaration) {} + + /// @brief Construct a Tool enabling Google Search grounding. + explicit Tool(const GoogleSearch& google_search) + : google_search_(google_search) {} + + /// @brief Construct a Tool enabling backend Code Execution. + explicit Tool(const CodeExecution& code_execution) + : code_execution_(code_execution) {} + + /// @brief Construct a Tool enabling Google Maps grounding. + explicit Tool(const GoogleMaps& google_maps) : google_maps_(google_maps) {} + + /// @brief Construct a Tool enabling URL Context grounding. + explicit Tool(const UrlContext& url_context) : url_context_(url_context) {} + + /// @brief Static factory creating a Tool with function declarations. + static Tool FunctionDeclarations( + const std::vector& declarations) { + return Tool(declarations); + } + + /// @brief Returns the function declarations in this tool. + const std::vector& function_declarations() const { + return function_declarations_; + } + + /// @brief Returns the optional Google Search configuration. + const Optional& google_search() const { return google_search_; } + + /// @brief Returns the optional Code Execution configuration. + const Optional& code_execution() const { + return code_execution_; + } + + /// @brief Returns the optional Google Maps configuration. + const Optional& google_maps() const { return google_maps_; } + + /// @brief Returns the optional URL Context configuration. + const Optional& url_context() const { return url_context_; } + + private: + std::vector function_declarations_; + Optional google_search_; + Optional code_execution_; + Optional google_maps_; + Optional url_context_; +}; + +/// @brief Controls how the model uses the provided `FunctionDeclaration` tools. +class FunctionCallingConfig { + public: + /// @brief Execution mode for function calling. + enum Mode { + /// Mode is unspecified. + kModeUnspecified = 0, + /// Model decides whether to predict a function call or a natural language + /// response (default). + kModeAuto, + /// Model is constrained to always predict a function call. + kModeAny, + /// Model will not predict any function call. + kModeNone, + }; + + /// @brief Default constructor initializes to `kModeAuto`. + FunctionCallingConfig() : mode_(kModeAuto) {} + + /// @brief Creates a `FunctionCallingConfig` with mode `kModeAuto`. + static FunctionCallingConfig Auto() { + return FunctionCallingConfig(kModeAuto, std::vector()); + } + + /// @brief Creates a `FunctionCallingConfig` with mode `kModeAny`. + /// + /// @param allowed_function_names Optional list of function names the model is + /// allowed to call. If empty, the model may call any declared function. + static FunctionCallingConfig Any( + const std::vector& allowed_function_names = + std::vector()) { + return FunctionCallingConfig(kModeAny, allowed_function_names); + } + + /// @brief Creates a `FunctionCallingConfig` with mode `kModeNone`. + static FunctionCallingConfig None() { + return FunctionCallingConfig(kModeNone, std::vector()); + } + + /// @brief Returns the function calling mode. + Mode mode() const { return mode_; } + + /// @brief Returns the list of allowed function names (for `kModeAny`). + const std::vector& allowed_function_names() const { + return allowed_function_names_; + } + + private: + FunctionCallingConfig(Mode mode, + const std::vector& allowed_function_names) + : mode_(mode), allowed_function_names_(allowed_function_names) {} + + Mode mode_; + std::vector allowed_function_names_; +}; + +/// @brief Configuration for tools provided to the model. +class ToolConfig { + public: + /// @brief Default constructor. + ToolConfig() {} + + /// @brief Construct a `ToolConfig` with function calling and/or retrieval + /// configuration. + /// + /// @param function_calling_config Optional function calling configuration. + /// @param retrieval_config Optional retrieval configuration (e.g. for Google + /// Maps). + explicit ToolConfig( + const Optional& function_calling_config, + const Optional& retrieval_config = + Optional()) + : function_calling_config_(function_calling_config), + retrieval_config_(retrieval_config) {} + + /// @brief Returns the optional `FunctionCallingConfig`. + const Optional& function_calling_config() const { + return function_calling_config_; + } + /// @brief Sets the `FunctionCallingConfig`. + void set_function_calling_config(const FunctionCallingConfig& config) { + function_calling_config_ = config; + } + + /// @brief Returns the optional `RetrievalConfig`. + const Optional& retrieval_config() const { + return retrieval_config_; + } + /// @brief Sets the `RetrievalConfig`. + void set_retrieval_config(const RetrievalConfig& config) { + retrieval_config_ = config; + } + + private: + Optional function_calling_config_; + Optional retrieval_config_; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_FUNCTION_CALLING_H_ diff --git a/ai/src/include/firebase/ai/generate_content_response.h b/ai/src/include/firebase/ai/generate_content_response.h new file mode 100644 index 0000000000..ca070ceadc --- /dev/null +++ b/ai/src/include/firebase/ai/generate_content_response.h @@ -0,0 +1,375 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_GENERATE_CONTENT_RESPONSE_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_GENERATE_CONTENT_RESPONSE_H_ + +#include +#include + +#include "firebase/ai/model_content.h" +#include "firebase/ai/safety.h" +#include "firebase/ai/types.h" + +namespace firebase { +namespace ai { + +/// @brief A single citation referencing a source used by the model in a +/// response. +struct Citation { + /// @brief Default constructor. + Citation() : start_index(0), end_index(0) {} + + /// @brief Start byte/character index in the response text. + int start_index; + + /// @brief End byte/character index in the response text. + int end_index; + + /// @brief URI of the cited source, if available. + std::string uri; + + /// @brief Title of the cited work, if available. + std::string title; + + /// @brief License of the cited work, if available. + std::string license; + + /// @brief Publication date (ISO-8601 string or year-month-day), if available. + std::string publication_date; +}; + +/// @brief Collection of citations for a `Candidate` response. +struct CitationMetadata { + /// @brief List of citations attributed to the candidate. + std::vector citations; +}; + +/// @brief Grounding chunk from a web source. +struct WebGroundingChunk { + /// @brief URI reference of the web chunk. + std::string uri; + + /// @brief Title of the web page. + std::string title; + + /// @brief Domain of the web source (if provided). + std::string domain; +}; + +/// @brief Grounding chunk from Google Maps. +struct GoogleMapsGroundingChunk { + /// @brief URI link to the place on Google Maps. + std::string uri; + + /// @brief Title / name of the place. + std::string title; + + /// @brief Google Maps Place ID (`places/...`). + std::string place_id; +}; + +/// @brief A chunk of grounding data supporting a response candidate. +struct GroundingChunk { + /// @brief Web grounding chunk details, if this chunk came from the web. + Optional web; + + /// @brief Google Maps grounding chunk details, if this chunk came from Maps. + Optional maps; +}; + +/// @brief Identifies a specific span of text within a `ModelContent` part. +struct GroundingSegment { + /// @brief Default constructor. + GroundingSegment() : part_index(0), start_index(0), end_index(0) {} + + /// @brief Zero-based index of the `Part` in `ModelContent::parts()`. + int part_index; + + /// @brief Start byte index in the part's text. + int start_index; + + /// @brief End byte index in the part's text. + int end_index; + + /// @brief The text corresponding to the segment. + std::string text; +}; + +/// @brief Links a `GroundingSegment` in the model response to one or more +/// `GroundingChunk` indices. +struct GroundingSupport { + /// @brief The segment of the model response supported by the chunks. + GroundingSegment segment; + + /// @brief Indices into `GroundingMetadata::grounding_chunks`. + std::vector grounding_chunk_indices; +}; + +/// @brief Google Search entry point widget data returned with grounded +/// responses. +struct SearchEntryPoint { + /// @brief Web content snippet HTML/CSS that can be embedded in a webview. + std::string rendered_content; + + /// @brief Base64-encoded JSON blob representing query/url pairs. + std::string sdk_blob; +}; + +/// @brief Metadata returned when grounding (such as Google Search or Google +/// Maps) is enabled. +struct GroundingMetadata { + /// @brief Search queries executed by the model for web grounding. + std::vector web_search_queries; + + /// @brief Search entry point widget information, if available. + Optional search_entry_point; + + /// @brief Supporting chunks retrieved during grounding. + std::vector grounding_chunks; + + /// @brief Links between response text segments and `grounding_chunks`. + std::vector grounding_supports; + + /// @brief Optional Google Maps widget context token. + Optional google_maps_widget_context_token; +}; + +/// @brief Metadata for a single URL retrieved by the `UrlContext` tool. +struct UrlMetadata { + /// @brief Default constructor. + UrlMetadata() : retrieval_status(kUrlRetrievalStatusUnspecified) {} + + /// @brief The retrieved URL. + std::string retrieved_url; + + /// @brief Status of retrieving the URL. + UrlRetrievalStatus retrieval_status; +}; + +/// @brief Metadata returned when the `UrlContext` tool is used. +struct UrlContextMetadata { + /// @brief List of URLs retrieved and their retrieval status. + std::vector url_metadata; +}; + +/// @brief A response candidate generated by the model. +/// +/// Mirrors `Firebase.AI.Candidate` in Unity and `Candidate` in Flutter. +struct Candidate { + /// @brief Default constructor. + Candidate() : finish_reason(kFinishReasonUnknown) {} + + /// @brief Generated content returned from the model. + ModelContent content; + + /// @brief Safety ratings for the response candidate. + std::vector safety_ratings; + + /// @brief Citation metadata for the response candidate, if any. + Optional citation_metadata; + + /// @brief Grounding metadata for the response candidate, if any. + Optional grounding_metadata; + + /// @brief URL context metadata for the response candidate, if any. + Optional url_context_metadata; + + /// @brief The reason why the model stopped generating tokens. + FinishReason finish_reason; + + /// @brief Additional message explaining `finish_reason`, if any. + std::string finish_message; +}; + +/// @brief Content filtering metadata for the input prompt. +/// +/// Mirrors `Firebase.AI.PromptFeedback` in Unity and `PromptFeedback` in +/// Flutter. +struct PromptFeedback { + /// @brief Default constructor. + PromptFeedback() : block_reason(kBlockReasonUnknown) {} + + /// @brief The reason why the prompt was blocked, if it was blocked. + BlockReason block_reason; + + /// @brief Safety ratings for the prompt across harm categories. + std::vector safety_ratings; + + /// @brief Human-readable message describing `block_reason`, if any. + std::string block_reason_message; +}; + +/// @brief Token count breakdown for a specific `ContentModality`. +struct ModalityTokenCount { + /// @brief Default constructor. + ModalityTokenCount() + : modality(kContentModalityUnspecified), token_count(0) {} + + /// @brief Construct a `ModalityTokenCount`. + ModalityTokenCount(ContentModality modality, int token_count) + : modality(modality), token_count(token_count) {} + + /// @brief The modality associated with this token count. + ContentModality modality; + + /// @brief The number of tokens for this modality. + int token_count; +}; + +/// @brief Token usage metadata for a `GenerateContentResponse`. +/// +/// Mirrors `Firebase.AI.UsageMetadata` in Unity and `UsageMetadata` in Flutter. +struct UsageMetadata { + /// @brief Default constructor initializes all counts to 0. + UsageMetadata() + : prompt_token_count(0), + candidates_token_count(0), + total_token_count(0), + thoughts_token_count(0), + tool_use_prompt_token_count(0), + cached_content_token_count(0) {} + + /// @brief Number of tokens in the input prompt. + int prompt_token_count; + + /// @brief Number of tokens in the generated candidates. + int candidates_token_count; + + /// @brief Total number of tokens across prompt, thinking, and candidates. + int total_token_count; + + /// @brief Number of tokens used by the model's thinking process. + int thoughts_token_count; + + /// @brief Number of tokens in tool-use prompt results. + int tool_use_prompt_token_count; + + /// @brief Number of tokens served from cached content. + int cached_content_token_count; + + /// @brief Per-modality breakdown of prompt tokens. + std::vector prompt_tokens_details; + + /// @brief Per-modality breakdown of candidate output tokens. + std::vector candidates_tokens_details; + + /// @brief Per-modality breakdown of tool-use prompt tokens. + std::vector tool_use_prompt_tokens_details; + + /// @brief Per-modality breakdown of cached content tokens. + std::vector cache_tokens_details; +}; + +/// @brief The model's response to a generate content request. +/// +/// Mirrors `Firebase.AI.GenerateContentResponse` in Unity and +/// `GenerateContentResponse` in Flutter. +class GenerateContentResponse { + public: + /// @brief Default constructor. + GenerateContentResponse() : inference_source_(kInferenceSourceInCloud) {} + + /// @brief Construct a `GenerateContentResponse`. + GenerateContentResponse( + const std::vector& candidates, + const Optional& prompt_feedback, + const Optional& usage_metadata, + InferenceSource inference_source = kInferenceSourceInCloud) + : candidates_(candidates), + prompt_feedback_(prompt_feedback), + usage_metadata_(usage_metadata), + inference_source_(inference_source) {} + + /// @brief Returns the list of response candidates generated by the model. + const std::vector& candidates() const { return candidates_; } + /// @brief Sets the list of response candidates. + void set_candidates(const std::vector& candidates) { + candidates_ = candidates; + } + + /// @brief Returns the prompt feedback, if any. + const Optional& prompt_feedback() const { + return prompt_feedback_; + } + /// @brief Sets the prompt feedback. + void set_prompt_feedback(const PromptFeedback& feedback) { + prompt_feedback_ = feedback; + } + + /// @brief Returns the token usage metadata, if any. + const Optional& usage_metadata() const { + return usage_metadata_; + } + /// @brief Sets the token usage metadata. + void set_usage_metadata(const UsageMetadata& metadata) { + usage_metadata_ = metadata; + } + + /// @brief Returns whether this response was generated in the cloud + /// (`kInferenceSourceInCloud`) or locally on-device via LiteRT + /// (`kInferenceSourceOnDevice`). + InferenceSource inference_source() const { return inference_source_; } + /// @brief Sets the inference source for this response. + void set_inference_source(InferenceSource source) { + inference_source_ = source; + } + + /// @brief Convenience property returning the concatenated non-thought text + /// parts of the first candidate, or an empty string if none exist. + std::string text() const; + + /// @brief Convenience property returning the concatenated thought summary + /// text parts of the first candidate, or an empty string if none exist. + std::string thought_summary() const; + + /// @brief Convenience property returning all `FunctionCallPart` items from + /// the first candidate. + std::vector function_calls() const; + + /// @brief Convenience property returning all non-thought `InlineDataPart` + /// items (such as generated images) from the first candidate. + std::vector inline_data_parts() const; + + private: + std::vector candidates_; + Optional prompt_feedback_; + Optional usage_metadata_; + InferenceSource inference_source_; +}; + +/// @brief The model's response to a `CountTokens` request. +/// +/// Mirrors `Firebase.AI.CountTokensResponse` in Unity and +/// `CountTokensResponse` in Flutter. +struct CountTokensResponse { + /// @brief Default constructor. + CountTokensResponse() : total_tokens(0), total_billable_characters(0) {} + + /// @brief Total number of tokens that the input content represents. + int total_tokens; + + /// @brief Total number of billable characters (Deprecated; may be 0). + int total_billable_characters; + + /// @brief Breakdown of prompt tokens by modality. + std::vector prompt_tokens_details; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_GENERATE_CONTENT_RESPONSE_H_ diff --git a/ai/src/include/firebase/ai/generation_config.h b/ai/src/include/firebase/ai/generation_config.h new file mode 100644 index 0000000000..4833c51ae9 --- /dev/null +++ b/ai/src/include/firebase/ai/generation_config.h @@ -0,0 +1,149 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_GENERATION_CONFIG_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_GENERATION_CONFIG_H_ + +#include +#include + +#include "firebase/ai/schema.h" +#include "firebase/ai/types.h" + +namespace firebase { +namespace ai { + +/// @brief Configuration for the "thinking" process of supported Gemini models. +/// +/// Mirrors `Firebase.AI.ThinkingConfig` in Unity and `ThinkingConfig` in +/// Flutter. +struct ThinkingConfig { + /// @brief Default constructor. + ThinkingConfig() {} + + /// @brief Construct a `ThinkingConfig` with a token budget. + /// + /// @param thinking_budget Token budget for the model's thinking process (0 to + /// disable thinking, -1 for dynamic budget). + /// @param include_thoughts Optional flag indicating whether thought summaries + /// should be included in the response. + explicit ThinkingConfig( + int thinking_budget, + const Optional& include_thoughts = Optional()) + : thinking_budget(thinking_budget), include_thoughts(include_thoughts) {} + + /// @brief Construct a `ThinkingConfig` with a `ThinkingLevel`. + /// + /// @param thinking_level Discrete thinking effort level. + /// @param include_thoughts Optional flag indicating whether thought summaries + /// should be included in the response. + explicit ThinkingConfig( + ThinkingLevel thinking_level, + const Optional& include_thoughts = Optional()) + : thinking_level(thinking_level), include_thoughts(include_thoughts) {} + + /// @brief Optional thinking token budget. + Optional thinking_budget; + + /// @brief Optional discrete thinking level. + Optional thinking_level; + + /// @brief Optional flag to include thought summaries in the response. + Optional include_thoughts; +}; + +/// @brief Configuration for image generation parameters when using Gemini image +/// generation models. +/// +/// Mirrors `Firebase.AI.ImageConfig` in Unity. +struct ImageConfig { + /// @brief Construct an `ImageConfig`. + /// + /// @param aspect_ratio Optional aspect ratio (e.g. `"1:1"`, `"16:9"`). + /// @param image_size Optional image resolution preset (e.g. `"1K"`, `"2K"`). + explicit ImageConfig( + const Optional& aspect_ratio = Optional(), + const Optional& image_size = Optional()) + : aspect_ratio(aspect_ratio), image_size(image_size) {} + + /// @brief Aspect ratio of generated images (e.g. `"1:1"`, `"3:4"`, `"4:3"`, + /// `"9:16"`, `"16:9"`). + Optional aspect_ratio; + + /// @brief Resolution preset of generated images (e.g. `"1K"`, `"2K"`). + Optional image_size; +}; + +/// @brief Configuration parameters used by `GenerativeModel` to control content +/// generation. +/// +/// Mirrors `Firebase.AI.GenerationConfig` in Unity and `GenerationConfig` in +/// Flutter. +struct GenerationConfig { + /// @brief Default constructor with all fields unset. + GenerationConfig() {} + + /// @brief Controls the degree of randomness in token selection. + Optional temperature; + + /// @brief Nucleus sampling probability threshold. + Optional top_p; + + /// @brief Top-k sampling token count threshold. + Optional top_k; + + /// @brief Number of response candidates to generate. + Optional candidate_count; + + /// @brief Maximum number of tokens to generate in the response. + Optional max_output_tokens; + + /// @brief Penalizes tokens that have already appeared in the generated text. + Optional presence_penalty; + + /// @brief Penalizes tokens proportional to how frequently they have appeared. + Optional frequency_penalty; + + /// @brief Character sequences that will stop output generation when + /// encountered. + std::vector stop_sequences; + + /// @brief Output response MIME type (e.g. `"text/plain"` or + /// `"application/json"`). + Optional response_mime_type; + + /// @brief OpenAPI `Schema` that the generated output must conform to (used + /// with `response_mime_type = "application/json"`). + Optional response_schema; + + /// @brief Standard `JsonSchema` that the generated output must conform to + /// (serialized as `responseJsonSchema`). + Optional response_json_schema; + + /// @brief Requested output modalities (e.g., text, image). + std::vector response_modalities; + + /// @brief Optional configuration for the model's thinking process. + Optional thinking_config; + + /// @brief Optional configuration for image generation outputs. + Optional image_config; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_GENERATION_CONFIG_H_ diff --git a/ai/src/include/firebase/ai/generative_model.h b/ai/src/include/firebase/ai/generative_model.h new file mode 100644 index 0000000000..4f634a2f60 --- /dev/null +++ b/ai/src/include/firebase/ai/generative_model.h @@ -0,0 +1,195 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_GENERATIVE_MODEL_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_GENERATIVE_MODEL_H_ + +#include +#include +#include +#include + +#include "firebase/ai/function_calling.h" +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generation_config.h" +#include "firebase/ai/model_content.h" +#include "firebase/ai/safety.h" +#include "firebase/ai/types.h" +#include "firebase/future.h" + +namespace firebase { +namespace ai { + +class Chat; +class FirebaseAI; + +namespace internal { +class GenerativeModelInternal; +class FirebaseAIInternal; +} // namespace internal + +/// @brief Callback invoked for each incremental `GenerateContentResponse` chunk +/// received during a `GenerateContentStream` or `SendMessageStream` call. +typedef std::function + GenerateContentStreamCallback; + +/// @brief A type that represents a remote multimodal model (like Gemini) with +/// the ability to generate content, stream content, count tokens, and conduct +/// multi-turn chat sessions. +/// +/// Mirrors `Firebase.AI.GenerativeModel` in Unity and `GenerativeModel` in +/// Flutter. +class GenerativeModel { + public: + /// @brief Default constructor creates an invalid `GenerativeModel`. + GenerativeModel(); + + /// @brief Copy constructor. + GenerativeModel(const GenerativeModel& other); + + /// @brief Copy assignment operator. + GenerativeModel& operator=(const GenerativeModel& other); + + /// @brief Destructor. + ~GenerativeModel(); + + /// @brief Returns true if this `GenerativeModel` is valid. + bool is_valid() const { return internal_ != nullptr; } + + /// @brief Generates content from a single text prompt. + /// + /// @param prompt The input text prompt. + /// @return A `Future` containing the `GenerateContentResponse`. + Future GenerateContent(const std::string& prompt); + + /// @brief Generates content from a single `ModelContent` message. + /// + /// @param content The input `ModelContent`. + /// @return A `Future` containing the `GenerateContentResponse`. + Future GenerateContent(const ModelContent& content); + + /// @brief Generates content from a sequence of `ModelContent` messages. + /// + /// @param content The input `ModelContent` messages. + /// @return A `Future` containing the `GenerateContentResponse`. + Future GenerateContent( + const std::vector& content); + + /// @brief Gets the result of the most recent `GenerateContent` call. + /// + /// @return A `Future` from the most recent `GenerateContent` call. + Future GenerateContentLastResult() const; + + /// @brief Generates a streaming response from a single text prompt, invoking + /// `on_chunk` as each `GenerateContentResponse` chunk arrives over SSE. + /// + /// @param prompt The input text prompt. + /// @param on_chunk Callback invoked for each response chunk. + /// @return A `Future` that completes when the stream finishes or fails. + Future GenerateContentStream( + const std::string& prompt, const GenerateContentStreamCallback& on_chunk); + + /// @brief Generates a streaming response from a single `ModelContent` + /// message, invoking `on_chunk` as each chunk arrives over SSE. + /// + /// @param content The input `ModelContent`. + /// @param on_chunk Callback invoked for each response chunk. + /// @return A `Future` that completes when the stream finishes or fails. + Future GenerateContentStream( + const ModelContent& content, + const GenerateContentStreamCallback& on_chunk); + + /// @brief Generates a streaming response from a sequence of `ModelContent` + /// messages, invoking `on_chunk` as each chunk arrives over SSE. + /// + /// @param content The input `ModelContent` messages. + /// @param on_chunk Callback invoked for each response chunk. + /// @return A `Future` that completes when the stream finishes or fails. + Future GenerateContentStream( + const std::vector& content, + const GenerateContentStreamCallback& on_chunk); + + /// @brief Gets the result of the most recent `GenerateContentStream` call. + /// + /// @return A `Future` from the most recent `GenerateContentStream` + /// call. + Future GenerateContentStreamLastResult() const; + + /// @brief Counts the number of tokens in a single text prompt. + /// + /// @param prompt The input text prompt. + /// @return A `Future` containing the `CountTokensResponse`. + Future CountTokens(const std::string& prompt); + + /// @brief Counts the number of tokens in a single `ModelContent` message. + /// + /// @param content The input `ModelContent`. + /// @return A `Future` containing the `CountTokensResponse`. + Future CountTokens(const ModelContent& content); + + /// @brief Counts the number of tokens in a sequence of `ModelContent` + /// messages. + /// + /// @param content The input `ModelContent` messages. + /// @return A `Future` containing the `CountTokensResponse`. + Future CountTokens( + const std::vector& content); + + /// @brief Gets the result of the most recent `CountTokens` call. + /// + /// @return A `Future` from the most recent `CountTokens` call. + Future CountTokensLastResult() const; + + /// @brief Creates a multi-turn `Chat` session using this `GenerativeModel`, + /// optionally initialized with prior conversation history. + /// + /// @param history Optional existing conversation turns. + /// @return A new `Chat` session instance. + Chat StartChat(const std::vector& history = + std::vector()) const; + + /// @brief Returns the active `InferenceMode` (`ONLY_IN_CLOUD`, + /// `ONLY_ON_DEVICE`, `PREFER_ON_DEVICE`, or `PREFER_IN_CLOUD`). + InferenceMode inference_mode() const; + + /// @brief Dynamically updates the `InferenceMode` at runtime (for example, + /// toggling a live `Chat` session between cloud Firebase AI and local + /// LiteRT inference). + void set_inference_mode(InferenceMode mode); + + /// @brief Checks whether the configured on-device LiteRT / LiteRT-LM model is + /// available and can be loaded on the current device. + bool IsOnDeviceAvailable() const; + + /// @brief Eagerly initializes the on-device LiteRT `CompiledModel` or + /// LiteRT-LM engine in the background. + Future InitializeOnDeviceModel(); + + private: + friend class Chat; + friend class FirebaseAI; + friend class internal::FirebaseAIInternal; + + explicit GenerativeModel( + const std::shared_ptr& internal); + + std::shared_ptr internal_; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_GENERATIVE_MODEL_H_ diff --git a/ai/src/include/firebase/ai/model_content.h b/ai/src/include/firebase/ai/model_content.h new file mode 100644 index 0000000000..6e61de7e05 --- /dev/null +++ b/ai/src/include/firebase/ai/model_content.h @@ -0,0 +1,499 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_MODEL_CONTENT_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_MODEL_CONTENT_H_ + +#include +#include +#include +#include + +#include "firebase/ai/types.h" +#include "firebase/variant.h" + +namespace firebase { +namespace ai { + +/// @brief Represents a text part of a `ModelContent` message. +struct TextPart { + /// @brief Default constructor. + TextPart() {} + + /// @brief Construct a TextPart with the given text. + /// + /// @param text The text content. + explicit TextPart(const std::string& text) : text(text) {} + + /// @brief The text content. + std::string text; +}; + +/// @brief Data with a specified media (MIME) type, sent inline in the payload. +struct InlineDataPart { + /// @brief Default constructor. + InlineDataPart() {} + + /// @brief Construct an InlineDataPart from raw bytes. + /// + /// @param mime_type The IANA standard MIME type (e.g. `"image/jpeg"`). + /// @param data The raw binary payload bytes. + InlineDataPart(const std::string& mime_type, const std::vector& data) + : mime_type(mime_type), data(data) {} + + /// @brief Construct an InlineDataPart from a raw byte pointer and length. + /// + /// @param mime_type The IANA standard MIME type. + /// @param bytes Pointer to the binary data. + /// @param size Number of bytes. + InlineDataPart(const std::string& mime_type, const uint8_t* bytes, + size_t size) + : mime_type(mime_type), data(bytes, bytes + size) {} + + /// @brief The IANA standard MIME type of the data. + std::string mime_type; + + /// @brief The raw binary data. + std::vector data; +}; + +/// @brief Data referenced via a URI (such as a Cloud Storage `gs://` URI or +/// HTTP/HTTPS URL) with a specified media (MIME) type. +struct FileDataPart { + /// @brief Default constructor. + FileDataPart() {} + + /// @brief Construct a FileDataPart with the given MIME type and URI. + /// + /// @param mime_type The IANA standard MIME type. + /// @param uri The URI or URL of the file. + FileDataPart(const std::string& mime_type, const std::string& uri) + : mime_type(mime_type), uri(uri) {} + + /// @brief The IANA standard MIME type of the file. + std::string mime_type; + + /// @brief The URI or URL of the file. + std::string uri; +}; + +/// @brief A predicted function call returned from the model. +struct FunctionCallPart { + /// @brief Default constructor. + FunctionCallPart() {} + + /// @brief Construct a FunctionCallPart. + /// + /// @param name The name of the function to call. + /// @param args The function arguments as a map of parameter names to Variant + /// values. + /// @param id Optional unique identifier for the function call. + FunctionCallPart(const std::string& name, + const std::map& args, + const Optional& id = Optional()) + : name(name), args(args), id(id) {} + + /// @brief The name of the function to call. + std::string name; + + /// @brief The function arguments as key-value pairs. + std::map args; + + /// @brief Optional unique identifier for this function call. + Optional id; +}; + +/// @brief Result output from a `FunctionCallPart` passed back to the model. +struct FunctionResponsePart { + /// @brief Default constructor. + FunctionResponsePart() {} + + /// @brief Construct a FunctionResponsePart. + /// + /// @param name The name of the function that was called. + /// @param response The structured response data returned by the function. + /// @param id Optional identifier matching the corresponding + /// `FunctionCallPart::id`. + FunctionResponsePart( + const std::string& name, const std::map& response, + const Optional& id = Optional()) + : name(name), response(response), id(id) {} + + /// @brief The name of the function that was called. + std::string name; + + /// @brief The function response fields. + std::map response; + + /// @brief Optional identifier of the corresponding function call. + Optional id; +}; + +/// @brief Code generated by the model that is meant to be executed by the +/// backend. +struct ExecutableCodePart { + /// @brief Programming language of the executable code. + enum CodeLanguage { + /// Unspecified or unrecognized language. + kLanguageUnspecified = 0, + /// Python. + kLanguagePython = 1, + }; + + /// @brief Default constructor. + ExecutableCodePart() : language(kLanguageUnspecified) {} + + /// @brief Construct an ExecutableCodePart. + /// + /// @param language The programming language. + /// @param code The source code to execute. + ExecutableCodePart(CodeLanguage language, const std::string& code) + : language(language), code(code) {} + + /// @brief The programming language of the `code`. + CodeLanguage language; + + /// @brief The source code to be executed. + std::string code; +}; + +/// @brief Represents the result of executing an `ExecutableCodePart` on the +/// backend. +struct CodeExecutionResultPart { + /// @brief Default constructor. + CodeExecutionResultPart() : outcome(kCodeExecutionOutcomeUnspecified) {} + + /// @brief Construct a CodeExecutionResultPart. + /// + /// @param outcome Outcome of the code execution. + /// @param output Captured stdout or stderr from the execution. + CodeExecutionResultPart(CodeExecutionOutcome outcome, + const std::string& output) + : outcome(outcome), output(output) {} + + /// @brief Outcome of the code execution. + CodeExecutionOutcome outcome; + + /// @brief Captured stdout (if successful) or stderr/error description. + std::string output; +}; + +/// @brief A discrete piece of data in a `ModelContent` message (text, inline +/// binary data, file reference, function call, function response, or code +/// execution data). +/// +/// Mirrors `ModelContent.Part` in the Unity SDK and `Part` in the Flutter SDK. +class Part { + public: + /// @brief Discriminator for the active variant held by a `Part`. + enum Kind { + /// Empty or uninitialized part. + kKindNone = 0, + /// Text part (`TextPart`). + kKindText, + /// Inline binary data part (`InlineDataPart`). + kKindInlineData, + /// File URI reference part (`FileDataPart`). + kKindFileData, + /// Function call part (`FunctionCallPart`). + kKindFunctionCall, + /// Function response part (`FunctionResponsePart`). + kKindFunctionResponse, + /// Executable code part (`ExecutableCodePart`). + kKindExecutableCode, + /// Code execution result part (`CodeExecutionResultPart`). + kKindCodeExecutionResult, + }; + + /// @brief Default constructor creates an empty Part. + Part() : kind_(kKindNone), is_thought_(false) {} + + /// @brief Construct a Part from a `TextPart`. + Part(const TextPart& part, bool is_thought = false, // NOLINT + const Optional& thought_signature = Optional()) + : kind_(kKindText), + text_part_(part), + is_thought_(is_thought), + thought_signature_(thought_signature) {} + + /// @brief Construct a Part from a string (creates a `TextPart`). + Part(const std::string& text) // NOLINT + : kind_(kKindText), text_part_(text), is_thought_(false) {} + + /// @brief Construct a Part from a C-string (creates a `TextPart`). + Part(const char* text) // NOLINT + : kind_(kKindText), text_part_(text ? text : ""), is_thought_(false) {} + + /// @brief Construct a Part from an `InlineDataPart`. + Part(const InlineDataPart& part, bool is_thought = false, // NOLINT + const Optional& thought_signature = Optional()) + : kind_(kKindInlineData), + inline_data_part_(part), + is_thought_(is_thought), + thought_signature_(thought_signature) {} + + /// @brief Construct a Part from a `FileDataPart`. + Part(const FileDataPart& part, bool is_thought = false, // NOLINT + const Optional& thought_signature = Optional()) + : kind_(kKindFileData), + file_data_part_(part), + is_thought_(is_thought), + thought_signature_(thought_signature) {} + + /// @brief Construct a Part from a `FunctionCallPart`. + Part(const FunctionCallPart& part, bool is_thought = false, // NOLINT + const Optional& thought_signature = Optional()) + : kind_(kKindFunctionCall), + function_call_part_(part), + is_thought_(is_thought), + thought_signature_(thought_signature) {} + + /// @brief Construct a Part from a `FunctionResponsePart`. + Part(const FunctionResponsePart& part, bool is_thought = false, // NOLINT + const Optional& thought_signature = Optional()) + : kind_(kKindFunctionResponse), + function_response_part_(part), + is_thought_(is_thought), + thought_signature_(thought_signature) {} + + /// @brief Construct a Part from an `ExecutableCodePart`. + Part(const ExecutableCodePart& part, bool is_thought = false, // NOLINT + const Optional& thought_signature = Optional()) + : kind_(kKindExecutableCode), + executable_code_part_(part), + is_thought_(is_thought), + thought_signature_(thought_signature) {} + + /// @brief Construct a Part from a `CodeExecutionResultPart`. + Part(const CodeExecutionResultPart& part, bool is_thought = false, // NOLINT + const Optional& thought_signature = Optional()) + : kind_(kKindCodeExecutionResult), + code_execution_result_part_(part), + is_thought_(is_thought), + thought_signature_(thought_signature) {} + + /// @brief Returns the active kind of this Part. + Kind kind() const { return kind_; } + + /// @brief Returns true if this is a `TextPart`. + bool is_text() const { return kind_ == kKindText; } + /// @brief Returns true if this is an `InlineDataPart`. + bool is_inline_data() const { return kind_ == kKindInlineData; } + /// @brief Returns true if this is a `FileDataPart`. + bool is_file_data() const { return kind_ == kKindFileData; } + /// @brief Returns true if this is a `FunctionCallPart`. + bool is_function_call() const { return kind_ == kKindFunctionCall; } + /// @brief Returns true if this is a `FunctionResponsePart`. + bool is_function_response() const { return kind_ == kKindFunctionResponse; } + /// @brief Returns true if this is an `ExecutableCodePart`. + bool is_executable_code() const { return kind_ == kKindExecutableCode; } + /// @brief Returns true if this is a `CodeExecutionResultPart`. + bool is_code_execution_result() const { + return kind_ == kKindCodeExecutionResult; + } + + /// @brief Returns the `TextPart` content. + const TextPart& text_part() const { return text_part_; } + /// @brief Returns the `InlineDataPart` content. + const InlineDataPart& inline_data_part() const { return inline_data_part_; } + /// @brief Returns the `FileDataPart` content. + const FileDataPart& file_data_part() const { return file_data_part_; } + /// @brief Returns the `FunctionCallPart` content. + const FunctionCallPart& function_call_part() const { + return function_call_part_; + } + /// @brief Returns the `FunctionResponsePart` content. + const FunctionResponsePart& function_response_part() const { + return function_response_part_; + } + /// @brief Returns the `ExecutableCodePart` content. + const ExecutableCodePart& executable_code_part() const { + return executable_code_part_; + } + /// @brief Returns the `CodeExecutionResultPart` content. + const CodeExecutionResultPart& code_execution_result_part() const { + return code_execution_result_part_; + } + + /// @brief Indicates whether this `Part` is a summary of the model's internal + /// thinking process. + bool is_thought() const { return is_thought_; } + /// @brief Sets whether this `Part` is a thought summary. + void set_is_thought(bool is_thought) { is_thought_ = is_thought; } + + /// @brief Returns the opaque thought signature associated with this part, if + /// any (must be preserved when sending history back to the model). + const Optional& thought_signature() const { + return thought_signature_; + } + /// @brief Sets the opaque thought signature associated with this part. + void set_thought_signature(const std::string& signature) { + thought_signature_ = signature; + } + + private: + Kind kind_; + TextPart text_part_; + InlineDataPart inline_data_part_; + FileDataPart file_data_part_; + FunctionCallPart function_call_part_; + FunctionResponsePart function_response_part_; + ExecutableCodePart executable_code_part_; + CodeExecutionResultPart code_execution_result_part_; + bool is_thought_; + Optional thought_signature_; +}; + +/// @brief A single message turn in a conversation with a generative model, +/// consisting of a role (`"user"`, `"model"`, or `"system"`) and one or more +/// `Part` objects. +/// +/// Mirrors `Firebase.AI.ModelContent` in Unity and `Content` in Flutter. +class ModelContent { + public: + /// @brief Default constructor creates a `"user"` message with no parts. + ModelContent() : role_("user") {} + + /// @brief Construct a `ModelContent` with the default `"user"` role and the + /// given parts. + /// + /// @param parts The parts that make up the message. + explicit ModelContent(const std::vector& parts) + : role_("user"), parts_(parts) {} + + /// @brief Construct a `ModelContent` with an explicit role and parts. + /// + /// @param role The producer role (`"user"`, `"model"`, `"system"`, or + /// `"function"`). + /// @param parts The parts that make up the message. + ModelContent(const std::string& role, const std::vector& parts) + : role_(role.empty() ? "user" : role), parts_(parts) {} + + /// @brief Creates a `ModelContent` with a single `TextPart` and the `"user"` + /// role. + /// + /// @param text The text prompt. + /// @return A user `ModelContent` containing `text`. + static ModelContent Text(const std::string& text); + + /// @brief Creates a `ModelContent` with a single `InlineDataPart` and the + /// `"user"` role. + /// + /// @param mime_type The IANA MIME type of the data. + /// @param data The raw binary data. + /// @return A user `ModelContent` containing the inline data. + static ModelContent InlineData(const std::string& mime_type, + const std::vector& data); + + /// @brief Creates a `ModelContent` with a single `InlineDataPart` from a byte + /// pointer and size. + /// + /// @param mime_type The IANA MIME type of the data. + /// @param bytes Pointer to the binary data. + /// @param size Number of bytes. + /// @return A user `ModelContent` containing the inline data. + static ModelContent InlineData(const std::string& mime_type, + const uint8_t* bytes, size_t size); + + /// @brief Creates a `ModelContent` with a single `FileDataPart` and the + /// `"user"` role. + /// + /// @param mime_type The IANA MIME type of the referenced file. + /// @param uri The URI or URL of the file. + /// @return A user `ModelContent` containing the file reference. + static ModelContent FileData(const std::string& mime_type, + const std::string& uri); + + /// @brief Creates a `ModelContent` with a single `FunctionResponsePart` and + /// the `"user"` role. + /// + /// @param name The name of the function that was called. + /// @param response The function response map. + /// @param id Optional identifier matching `FunctionCallPart::id`. + /// @return A user `ModelContent` containing the function response. + static ModelContent FunctionResponse( + const std::string& name, const std::map& response, + const Optional& id = Optional()); + + /// @brief Creates a `ModelContent` with multiple `FunctionResponsePart` items + /// and the `"user"` role. + /// + /// @param responses The function response parts. + /// @return A user `ModelContent` containing all the function responses. + static ModelContent FunctionResponses( + const std::vector& responses); + + /// @brief Creates a `ModelContent` with the `"system"` role for system + /// instructions. + /// + /// @param text The system instruction text. + /// @return A system `ModelContent`. + static ModelContent System(const std::string& text); + + /// @brief Creates a `ModelContent` with the `"user"` role and a single text + /// part. + /// + /// @param text The user message text. + /// @return A user `ModelContent`. + static ModelContent User(const std::string& text); + + /// @brief Creates a `ModelContent` with the `"user"` role and the given + /// parts. + /// + /// @param parts The parts of the user message. + /// @return A user `ModelContent`. + static ModelContent User(const std::vector& parts); + + /// @brief Creates a `ModelContent` with the `"model"` role and a single text + /// part. + /// + /// @param text The model message text. + /// @return A model `ModelContent`. + static ModelContent Model(const std::string& text); + + /// @brief Creates a `ModelContent` with the `"model"` role and the given + /// parts. + /// + /// @param parts The parts of the model message. + /// @return A model `ModelContent`. + static ModelContent Model(const std::vector& parts); + + /// @brief Returns the role of the producer of the content. + const std::string& role() const { return role_; } + /// @brief Sets the role of the producer of the content. + void set_role(const std::string& role) { role_ = role; } + + /// @brief Returns the ordered list of parts in this content. + const std::vector& parts() const { return parts_; } + /// @brief Returns a mutable reference to the ordered list of parts. + std::vector& parts() { return parts_; } + /// @brief Sets the ordered list of parts in this content. + void set_parts(const std::vector& parts) { parts_ = parts; } + + /// @brief Appends a `Part` to this `ModelContent`. + /// + /// @param part The part to append. + void AddPart(const Part& part) { parts_.push_back(part); } + + private: + std::string role_; + std::vector parts_; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_MODEL_CONTENT_H_ diff --git a/ai/src/include/firebase/ai/safety.h b/ai/src/include/firebase/ai/safety.h new file mode 100644 index 0000000000..262a37d584 --- /dev/null +++ b/ai/src/include/firebase/ai/safety.h @@ -0,0 +1,105 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_SAFETY_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_SAFETY_H_ + +#include "firebase/ai/types.h" + +namespace firebase { +namespace ai { + +/// @brief A type used to specify a threshold for blocking harmful content for a +/// given `HarmCategory`. +/// +/// Mirrors `Firebase.AI.SafetySetting` in Unity and `SafetySetting` in Flutter. +class SafetySetting { + public: + /// @brief Default constructor. + SafetySetting() + : category_(kHarmCategoryUnknown), + threshold_(kHarmBlockThresholdUnknown), + method_(Optional()) {} + + /// @brief Construct a `SafetySetting` with a category, threshold, and + /// optional evaluation method. + /// + /// @param category The category of harm to configure. + /// @param threshold The threshold at and above which content is blocked. + /// @param method Optional method (probability vs. severity) used to evaluate + /// the threshold (only supported on the Enterprise / VertexAI backend). + SafetySetting( + HarmCategory category, HarmBlockThreshold threshold, + const Optional& method = Optional()) + : category_(category), threshold_(threshold), method_(method) {} + + /// @brief Returns the harm category. + HarmCategory category() const { return category_; } + /// @brief Sets the harm category. + void set_category(HarmCategory category) { category_ = category; } + + /// @brief Returns the block threshold. + HarmBlockThreshold threshold() const { return threshold_; } + /// @brief Sets the block threshold. + void set_threshold(HarmBlockThreshold threshold) { threshold_ = threshold; } + + /// @brief Returns the optional block evaluation method. + const Optional& method() const { return method_; } + /// @brief Sets the block evaluation method. + void set_method(HarmBlockMethod method) { method_ = method; } + + private: + HarmCategory category_; + HarmBlockThreshold threshold_; + Optional method_; +}; + +/// @brief A type defining safety attributes of a `Candidate` or prompt. +/// +/// Mirrors `Firebase.AI.SafetyRating` in Unity and `SafetyRating` in Flutter. +struct SafetyRating { + /// @brief Default constructor. + SafetyRating() + : category(kHarmCategoryUnknown), + probability(kHarmProbabilityUnknown), + blocked(false), + probability_score(0.0f), + severity(kHarmSeverityUnknown), + severity_score(0.0f) {} + + /// @brief The category for this rating. + HarmCategory category; + + /// @brief The probability of harm for this content. + HarmProbability probability; + + /// @brief Indicates whether the content was blocked because of this rating. + bool blocked; + + /// @brief The probability score of harm (0.0 to 1.0, VertexAI/Enterprise). + float probability_score; + + /// @brief The severity of harm for this content (VertexAI/Enterprise). + HarmSeverity severity; + + /// @brief The severity score of harm (0.0 to 1.0, VertexAI/Enterprise). + float severity_score; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_SAFETY_H_ diff --git a/ai/src/include/firebase/ai/schema.h b/ai/src/include/firebase/ai/schema.h new file mode 100644 index 0000000000..f997fd7ae6 --- /dev/null +++ b/ai/src/include/firebase/ai/schema.h @@ -0,0 +1,324 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_SCHEMA_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_SCHEMA_H_ + +#include +#include +#include +#include +#include + +#include "firebase/ai/types.h" +#include "firebase/variant.h" + +namespace firebase { +namespace ai { + +/// @brief Value types supported by `Schema` and `JsonSchema`. +enum SchemaType { + /// Type is unspecified (e.g., when `any_of` is used). + kSchemaTypeUnspecified = 0, + /// String type (`"STRING"` in OpenAPI Schema, `"string"` in JSON Schema). + kSchemaTypeString, + /// Floating-point number (`"NUMBER"` in OpenAPI, `"number"` in JSON Schema). + kSchemaTypeNumber, + /// Integral number (`"INTEGER"` in OpenAPI, `"integer"` in JSON Schema). + kSchemaTypeInteger, + /// Boolean (`"BOOLEAN"` in OpenAPI, `"boolean"` in JSON Schema). + kSchemaTypeBoolean, + /// Array of items (`"ARRAY"` in OpenAPI, `"array"` in JSON Schema). + kSchemaTypeArray, + /// Key-value object (`"OBJECT"` in OpenAPI, `"object"` in JSON Schema). + kSchemaTypeObject, +}; + +/// @brief Defines the structure of input parameters for function declarations +/// or structured JSON output (`GenerationConfig::response_schema`). +/// +/// Mirrors `Firebase.AI.Schema` in Unity and `Schema` in Flutter. +class Schema { + public: + /// @brief Default constructor creates an unspecified schema. + Schema() : type_(kSchemaTypeUnspecified) {} + + /// @brief Construct a Schema with the given type. + /// + /// @param type The schema data type. + explicit Schema(SchemaType type) : type_(type) {} + + /// @brief Returns a Schema for a boolean value. + /// + /// @param description Optional explanation of the field. + /// @param nullable Optional flag indicating whether the value may be null. + /// @param title Optional human-readable title. + /// @return A boolean `Schema`. + static Schema Boolean( + const Optional& description = Optional(), + const Optional& nullable = Optional(), + const Optional& title = Optional()); + + /// @brief Returns a Schema for a 32-bit signed integer. + /// + /// @param description Optional explanation of the field. + /// @param nullable Optional flag indicating whether the value may be null. + /// @param title Optional human-readable title. + /// @param minimum Optional inclusive minimum value. + /// @param maximum Optional inclusive maximum value. + /// @return An integer `Schema` with format `"int32"`. + static Schema Int( + const Optional& description = Optional(), + const Optional& nullable = Optional(), + const Optional& title = Optional(), + const Optional& minimum = Optional(), + const Optional& maximum = Optional()); + + /// @brief Returns a Schema for a 64-bit signed integer. + /// + /// @param description Optional explanation of the field. + /// @param nullable Optional flag indicating whether the value may be null. + /// @param title Optional human-readable title. + /// @param minimum Optional inclusive minimum value. + /// @param maximum Optional inclusive maximum value. + /// @return An integer `Schema` with format `"int64"`. + static Schema Long( + const Optional& description = Optional(), + const Optional& nullable = Optional(), + const Optional& title = Optional(), + const Optional& minimum = Optional(), + const Optional& maximum = Optional()); + + /// @brief Returns a Schema for a single-precision floating-point number. + /// + /// @param description Optional explanation of the field. + /// @param nullable Optional flag indicating whether the value may be null. + /// @param title Optional human-readable title. + /// @param minimum Optional inclusive minimum value. + /// @param maximum Optional inclusive maximum value. + /// @return A number `Schema` with format `"float"`. + static Schema Float( + const Optional& description = Optional(), + const Optional& nullable = Optional(), + const Optional& title = Optional(), + const Optional& minimum = Optional(), + const Optional& maximum = Optional()); + + /// @brief Returns a Schema for a double-precision floating-point number. + /// + /// @param description Optional explanation of the field. + /// @param nullable Optional flag indicating whether the value may be null. + /// @param title Optional human-readable title. + /// @param minimum Optional inclusive minimum value. + /// @param maximum Optional inclusive maximum value. + /// @return A number `Schema`. + static Schema Double( + const Optional& description = Optional(), + const Optional& nullable = Optional(), + const Optional& title = Optional(), + const Optional& minimum = Optional(), + const Optional& maximum = Optional()); + + /// @brief Returns a Schema for a string. + /// + /// @param description Optional explanation of the field. + /// @param nullable Optional flag indicating whether the value may be null. + /// @param format Optional string format (e.g. `"date-time"`). + /// @param title Optional human-readable title. + /// @return A string `Schema`. + static Schema String( + const Optional& description = Optional(), + const Optional& nullable = Optional(), + const Optional& format = Optional(), + const Optional& title = Optional()); + + /// @brief Returns a Schema for an enumeration of string values. + /// + /// @param values The list of valid string values. + /// @param description Optional explanation of the field. + /// @param nullable Optional flag indicating whether the value may be null. + /// @param title Optional human-readable title. + /// @return An enum string `Schema`. + static Schema Enum( + const std::vector& values, + const Optional& description = Optional(), + const Optional& nullable = Optional(), + const Optional& title = Optional()); + + /// @brief Returns a Schema for an array of elements. + /// + /// @param items The schema for the elements in the array. + /// @param description Optional explanation of the field. + /// @param nullable Optional flag indicating whether the value may be null. + /// @param title Optional human-readable title. + /// @param min_items Optional minimum number of elements. + /// @param max_items Optional maximum number of elements. + /// @return An array `Schema`. + static Schema Array( + const Schema& items, + const Optional& description = Optional(), + const Optional& nullable = Optional(), + const Optional& title = Optional(), + const Optional& min_items = Optional(), + const Optional& max_items = Optional()); + + /// @brief Returns a Schema for a complex object with named properties. + /// + /// @param properties Map of property names to their `Schema` definitions. + /// @param optional_properties List of property names that are not required. + /// All properties not in `optional_properties` will be marked as required. + /// @param description Optional explanation of the object. + /// @param nullable Optional flag indicating whether the object may be null. + /// @param title Optional human-readable title. + /// @param property_ordering Optional explicit ordering of property keys. + /// @return An object `Schema`. + static Schema Object( + const std::map& properties, + const std::vector& optional_properties = + std::vector(), + const Optional& description = Optional(), + const Optional& nullable = Optional(), + const Optional& title = Optional(), + const std::vector& property_ordering = + std::vector()); + + /// @brief Returns a Schema representing a union (`anyOf`) of schemas. + /// + /// @param schemas The candidate schemas. + /// @return A union `Schema`. + static Schema AnyOf(const std::vector& schemas); + + /// @brief Returns the schema type. + SchemaType type() const { return type_; } + /// @brief Sets the schema type. + void set_type(SchemaType type) { type_ = type; } + + /// @brief Returns the description. + const Optional& description() const { return description_; } + /// @brief Sets the description. + void set_description(const std::string& description) { + description_ = description; + } + + /// @brief Returns the format specifier. + const Optional& format() const { return format_; } + /// @brief Sets the format specifier. + void set_format(const std::string& format) { format_ = format; } + + /// @brief Returns whether the value is nullable. + const Optional& nullable() const { return nullable_; } + /// @brief Sets whether the value is nullable. + void set_nullable(bool nullable) { nullable_ = nullable; } + + /// @brief Returns the enum values, if any. + const std::vector& enum_values() const { return enum_values_; } + /// @brief Sets the enum values. + void set_enum_values(const std::vector& values) { + enum_values_ = values; + } + + /// @brief Returns the properties map for an object schema. + const std::map& properties() const { + return properties_; + } + /// @brief Sets the properties map for an object schema. + void set_properties(const std::map& properties) { + properties_ = properties; + } + + /// @brief Returns the list of required property names. + const std::vector& required_properties() const { + return required_properties_; + } + /// @brief Sets the list of required property names. + void set_required_properties(const std::vector& required) { + required_properties_ = required; + } + + /// @brief Returns the property ordering list. + const std::vector& property_ordering() const { + return property_ordering_; + } + /// @brief Sets the property ordering list. + void set_property_ordering(const std::vector& ordering) { + property_ordering_ = ordering; + } + + /// @brief Returns the element schema for an array schema, or nullptr. + const Schema* items() const { return items_.get(); } + /// @brief Sets the element schema for an array schema. + void set_items(const Schema& items) { items_.reset(new Schema(items)); } + + /// @brief Returns the title. + const Optional& title() const { return title_; } + /// @brief Sets the title. + void set_title(const std::string& title) { title_ = title; } + + /// @brief Returns the minimum items count for an array schema. + const Optional& min_items() const { return min_items_; } + /// @brief Sets the minimum items count for an array schema. + void set_min_items(int64_t min_items) { min_items_ = min_items; } + + /// @brief Returns the maximum items count for an array schema. + const Optional& max_items() const { return max_items_; } + /// @brief Sets the maximum items count for an array schema. + void set_max_items(int64_t max_items) { max_items_ = max_items; } + + /// @brief Returns the minimum numeric value. + const Optional& minimum() const { return minimum_; } + /// @brief Sets the minimum numeric value. + void set_minimum(double minimum) { minimum_ = minimum; } + + /// @brief Returns the maximum numeric value. + const Optional& maximum() const { return maximum_; } + /// @brief Sets the maximum numeric value. + void set_maximum(double maximum) { maximum_ = maximum; } + + /// @brief Returns the `anyOf` sub-schemas. + const std::vector& any_of() const { return any_of_; } + /// @brief Sets the `anyOf` sub-schemas. + void set_any_of(const std::vector& any_of) { any_of_ = any_of; } + + private: + SchemaType type_; + Optional description_; + Optional format_; + Optional nullable_; + std::vector enum_values_; + std::map properties_; + std::vector required_properties_; + std::vector property_ordering_; + std::shared_ptr items_; + Optional title_; + Optional min_items_; + Optional max_items_; + Optional minimum_; + Optional maximum_; + std::vector any_of_; +}; + +/// @brief Standard JSON Schema definition used with +/// `GenerationConfig::response_json_schema`. +/// +/// Mirrors `Firebase.AI.JsonSchema` in the Unity SDK, serializing types in +/// lowercase (`"string"`, `"object"`, etc.) and representing `nullable` via +/// `["", "null"]` or `anyOf` with `{"type": "null"}`. +using JsonSchema = Schema; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_SCHEMA_H_ diff --git a/ai/src/include/firebase/ai/template_chat_session.h b/ai/src/include/firebase/ai/template_chat_session.h new file mode 100644 index 0000000000..4f47c4630a --- /dev/null +++ b/ai/src/include/firebase/ai/template_chat_session.h @@ -0,0 +1,104 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_TEMPLATE_CHAT_SESSION_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_TEMPLATE_CHAT_SESSION_H_ + +#include +#include +#include + +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generative_model.h" +#include "firebase/ai/model_content.h" +#include "firebase/future.h" + +namespace firebase { +namespace ai { + +class TemplateGenerativeModel; + +namespace internal { +class TemplateChatSessionInternal; +} // namespace internal + +/// @brief Multi-turn chat session backed by a server prompt template on +/// `TemplateGenerativeModel`. +class TemplateChatSession { + public: + /// @brief Default constructor creates an invalid `TemplateChatSession`. + TemplateChatSession(); + + /// @brief Copy constructor. + TemplateChatSession(const TemplateChatSession& other); + + /// @brief Copy assignment operator. + TemplateChatSession& operator=(const TemplateChatSession& other); + + /// @brief Destructor. + ~TemplateChatSession(); + + /// @brief Returns true if this `TemplateChatSession` is valid. + bool is_valid() const { return internal_ != nullptr; } + + /// @brief Returns the conversation history accumulated in this session. + std::vector history() const; + + /// @brief Sends a text message in this template chat session. + Future SendMessage(const std::string& prompt); + + /// @brief Sends a `ModelContent` message in this template chat session. + Future SendMessage(const ModelContent& content); + + /// @brief Sends multiple `ModelContent` messages in this template chat + /// session. + Future SendMessage( + const std::vector& content); + + /// @brief Gets the result of the most recent `SendMessage` call. + Future SendMessageLastResult() const; + + /// @brief Sends a text message and streams back response chunks via + /// `on_chunk`. + Future SendMessageStream(const std::string& prompt, + const GenerateContentStreamCallback& on_chunk); + + /// @brief Sends a `ModelContent` message and streams back response chunks via + /// `on_chunk`. + Future SendMessageStream(const ModelContent& content, + const GenerateContentStreamCallback& on_chunk); + + /// @brief Sends multiple `ModelContent` messages and streams back response + /// chunks via `on_chunk`. + Future SendMessageStream(const std::vector& content, + const GenerateContentStreamCallback& on_chunk); + + /// @brief Gets the result of the most recent `SendMessageStream` call. + Future SendMessageStreamLastResult() const; + + private: + friend class TemplateGenerativeModel; + + explicit TemplateChatSession( + const std::shared_ptr& internal); + + std::shared_ptr internal_; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_TEMPLATE_CHAT_SESSION_H_ diff --git a/ai/src/include/firebase/ai/template_generative_model.h b/ai/src/include/firebase/ai/template_generative_model.h new file mode 100644 index 0000000000..dc0c5aac38 --- /dev/null +++ b/ai/src/include/firebase/ai/template_generative_model.h @@ -0,0 +1,141 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_TEMPLATE_GENERATIVE_MODEL_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_TEMPLATE_GENERATIVE_MODEL_H_ + +#include +#include +#include +#include + +#include "firebase/ai/generate_content_response.h" +#include "firebase/ai/generative_model.h" +#include "firebase/ai/model_content.h" +#include "firebase/ai/types.h" +#include "firebase/future.h" +#include "firebase/variant.h" + +namespace firebase { +namespace ai { + +class FirebaseAI; +class TemplateChatSession; + +namespace internal { +class TemplateGenerativeModelInternal; +class FirebaseAIInternal; +} // namespace internal + +/// @brief A type that represents a remote server-prompt-template model with the +/// ability to generate content and stream content by supplying template IDs and +/// input variables. +/// +/// Mirrors `Firebase.AI.TemplateGenerativeModel` in Unity and +/// `TemplateGenerativeModel` in Flutter. +class TemplateGenerativeModel { + public: + /// @brief Default constructor creates an invalid `TemplateGenerativeModel`. + TemplateGenerativeModel(); + + /// @brief Copy constructor. + TemplateGenerativeModel(const TemplateGenerativeModel& other); + + /// @brief Copy assignment operator. + TemplateGenerativeModel& operator=(const TemplateGenerativeModel& other); + + /// @brief Destructor. + ~TemplateGenerativeModel(); + + /// @brief Returns true if this `TemplateGenerativeModel` is valid. + bool is_valid() const { return internal_ != nullptr; } + + /// @brief Generates content from a server prompt template and a map of + /// template input variables. + /// + /// @param template_id The ID of the server prompt template. + /// @param inputs Key-value map of template variables. + /// @return A `Future` containing the `GenerateContentResponse`. + Future GenerateContent( + const std::string& template_id, + const std::map& inputs); + + /// @brief Generates content from a server prompt template and a raw JSON + /// object string of template input variables. + /// + /// @param template_id The ID of the server prompt template. + /// @param json_inputs A JSON object string of template variables. + /// @return A `Future` containing the `GenerateContentResponse`. + Future GenerateContentJson( + const std::string& template_id, const std::string& json_inputs); + + /// @brief Gets the result of the most recent `GenerateContent` call. + Future GenerateContentLastResult() const; + + /// @brief Generates a streaming response from a server prompt template and a + /// map of template input variables. + /// + /// @param template_id The ID of the server prompt template. + /// @param inputs Key-value map of template variables. + /// @param on_chunk Callback invoked for each response chunk. + /// @return A `Future` that completes when the stream finishes or fails. + Future GenerateContentStream( + const std::string& template_id, + const std::map& inputs, + const GenerateContentStreamCallback& on_chunk); + + /// @brief Generates a streaming response from a server prompt template and a + /// raw JSON object string of template input variables. + /// + /// @param template_id The ID of the server prompt template. + /// @param json_inputs A JSON object string of template variables. + /// @param on_chunk Callback invoked for each response chunk. + /// @return A `Future` that completes when the stream finishes or fails. + Future GenerateContentStreamJson( + const std::string& template_id, const std::string& json_inputs, + const GenerateContentStreamCallback& on_chunk); + + /// @brief Gets the result of the most recent `GenerateContentStream` call. + Future GenerateContentStreamLastResult() const; + + /// @brief Starts a multi-turn `TemplateChatSession` bound to `template_id`. + /// + /// @param template_id The ID of the server prompt template. + /// @param inputs Optional template input variables. + /// @param history Optional existing conversation history. + /// @return A new `TemplateChatSession`. + TemplateChatSession StartChat(const std::string& template_id, + const std::map& inputs = + std::map(), + const std::vector& history = + std::vector()) const; + + private: + friend class FirebaseAI; + friend class TemplateChatSession; + friend class internal::FirebaseAIInternal; + + explicit TemplateGenerativeModel( + const std::shared_ptr& + internal); + + std::shared_ptr internal_; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_TEMPLATE_GENERATIVE_MODEL_H_ diff --git a/ai/src/include/firebase/ai/types.h b/ai/src/include/firebase/ai/types.h new file mode 100644 index 0000000000..5ac222d28e --- /dev/null +++ b/ai/src/include/firebase/ai/types.h @@ -0,0 +1,608 @@ +/* + * Copyright 2025 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. + */ + +#ifndef FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_TYPES_H_ +#define FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_TYPES_H_ + +#include +#include +#include +#include +#include + +#include "firebase/internal/common.h" + +namespace firebase { +namespace ai { + +/// @brief Error codes returned by Firebase AI Logic futures. +enum Error { + /// The operation was a success, no error occurred. + kErrorNone = 0, + /// Invalid argument passed to the API (e.g., empty model name or prompt). + kErrorInvalidArgument = 1, + /// Failed to serialize request or deserialize response JSON. + kErrorSerializationFailed = 2, + /// Network or transport failure while communicating with the backend. + kErrorNetworkFailed = 3, + /// The request timed out before completing. + kErrorTimeout = 4, + /// The backend returned an HTTP error status (e.g., 4xx or 5xx). + kErrorHttpError = 5, + /// The prompt or response was blocked by safety settings or policy. + kErrorResponseBlocked = 6, + /// The FirebaseAI or GenerativeModel instance is no longer valid. + kErrorInvalidState = 7, + /// An unknown or internal error occurred. + kErrorUnknown = 8, + /// The requested operation or on-device model is not available/supported. + kErrorUnsupported = 9, +}; + +/// @brief Lightweight optional value container for C++14 compatibility. +template +class Optional { + public: + /// @brief Construct an empty Optional. + Optional() : has_value_(false), value_() {} + + /// @brief Construct an Optional containing a value. + /// + /// @param value The value to store. + Optional(const T& value) // NOLINT(runtime/explicit) + : has_value_(true), value_(value) {} + + /// @brief Construct an Optional by moving a value. + /// + /// @param value The value to move into this Optional. + Optional(T&& value) // NOLINT(runtime/explicit) + : has_value_(true), value_(std::move(value)) {} + + /// @brief Construct an Optional from an implicitly convertible value (e.g., + /// `const char*` for `Optional`). + template < + typename U, + typename std::enable_if< + !std::is_same::type, Optional>::value && + !std::is_same::type, T>::value && + std::is_convertible::value, + int>::type = 0> + Optional(U&& value) // NOLINT(runtime/explicit) + : has_value_(true), value_(std::forward(value)) {} + + /// @brief Assign a value to this Optional. + /// + /// @param value The value to store. + /// @return Reference to this Optional. + Optional& operator=(const T& value) { + has_value_ = true; + value_ = value; + return *this; + } + + /// @brief Move-assign a value to this Optional. + /// + /// @param value The value to move into this Optional. + /// @return Reference to this Optional. + Optional& operator=(T&& value) { + has_value_ = true; + value_ = std::move(value); + return *this; + } + + /// @brief Assign an implicitly convertible value to this Optional. + template < + typename U, + typename std::enable_if< + !std::is_same::type, Optional>::value && + !std::is_same::type, T>::value && + std::is_convertible::value, + int>::type = 0> + Optional& operator=(U&& value) { + has_value_ = true; + value_ = T(std::forward(value)); + return *this; + } + + /// @brief Returns true if this Optional holds a value. + bool has_value() const { return has_value_; } + + /// @brief Conversion to bool indicating whether a value is present. + explicit operator bool() const { return has_value_; } + + /// @brief Returns a const reference to the contained value. + const T& value() const { + assert(has_value_); + return value_; + } + + /// @brief Returns a mutable reference to the contained value. + T& value() { + assert(has_value_); + return value_; + } + + /// @brief Returns the contained value if present, or `default_value`. + /// + /// @param default_value Fallback value when empty. + /// @return The contained value or `default_value`. + T value_or(const T& default_value) const { + return has_value_ ? value_ : default_value; + } + + /// @brief Dereference operator. + const T& operator*() const { return value(); } + + /// @brief Dereference operator. + T& operator*() { return value(); } + + /// @brief Member access operator. + const T* operator->() const { return &value(); } + + /// @brief Member access operator. + T* operator->() { return &value(); } + + /// @brief Clears any contained value. + void reset() { + has_value_ = false; + value_ = T(); + } + + /// @brief Equality comparison. + bool operator==(const Optional& other) const { + if (has_value_ != other.has_value_) return false; + return !has_value_ || (value_ == other.value_); + } + + /// @brief Inequality comparison. + bool operator!=(const Optional& other) const { return !(*this == other); } + + private: + bool has_value_; + T value_; +}; + +/// @brief Identifies which backend provider to target for Firebase AI calls. +enum BackendProvider { + /// The Gemini Developer API backend (`GoogleAI`). + kBackendProviderGoogleAI = 0, + /// The Vertex AI Gemini API backend (`Enterprise` / `VertexAI`). + kBackendProviderEnterprise = 1, +}; + +/// @brief Specifies the backend configuration for `FirebaseAI`. +/// +/// Mirrors `FirebaseAI.Backend` in the Unity and Flutter SDKs. +class Backend { + public: + /// @brief Creates a `Backend` targeting the Gemini Developer API + /// (`GoogleAI`). + /// + /// @return A `Backend` configured for GoogleAI. + static Backend GoogleAI() { return Backend(kBackendProviderGoogleAI, ""); } + + /// @brief Creates a `Backend` targeting the Vertex AI Gemini API. + /// + /// @param location The Google Cloud region identifier, defaulting to "global" + /// (as in the Unity SDK). + /// @return A `Backend` configured for Vertex AI / Enterprise. + static Backend Enterprise(const std::string& location = "global") { + return Backend(kBackendProviderEnterprise, location); + } + + /// @brief Alias for `Enterprise`, matching the Flutter SDK `vertexAI` naming. + /// + /// @param location The Google Cloud region identifier, defaulting to + /// "us-central1". + /// @return A `Backend` configured for Vertex AI. + static Backend VertexAI(const std::string& location = "us-central1") { + return Backend(kBackendProviderEnterprise, location); + } + + /// @brief Default constructor initializes to `Backend::GoogleAI()`. + Backend() : provider_(kBackendProviderGoogleAI), location_("") {} + + /// @brief Returns the backend provider type. + BackendProvider provider() const { return provider_; } + + /// @brief Returns the configured location (for Enterprise / VertexAI). + const std::string& location() const { return location_; } + + /// @brief Equality comparison. + bool operator==(const Backend& other) const { + return provider_ == other.provider_ && location_ == other.location_; + } + + /// @brief Inequality comparison. + bool operator!=(const Backend& other) const { return !(*this == other); } + + /// @brief Less-than operator for use in associative containers. + bool operator<(const Backend& other) const { + if (provider_ != other.provider_) { + return static_cast(provider_) < static_cast(other.provider_); + } + return location_ < other.location_; + } + + private: + Backend(BackendProvider provider, const std::string& location) + : provider_(provider), location_(location) {} + + BackendProvider provider_; + std::string location_; +}; + +/// @brief Configuration options for requests made to the backend. +struct RequestOptions { + /// @brief Default request timeout in milliseconds (180 seconds). + static const int64_t kDefaultTimeoutMs = 180000; + + /// @brief Construct default RequestOptions. + /// + /// @param timeout_ms Request timeout in milliseconds. + /// @param limited_use_app_check_token Whether to use a limited-use App Check + /// token instead of a cached token. + explicit RequestOptions(int64_t timeout_ms = kDefaultTimeoutMs, + bool limited_use_app_check_token = false) + : timeout_ms(timeout_ms), + limited_use_app_check_token(limited_use_app_check_token) {} + + /// @brief Request timeout in milliseconds. + int64_t timeout_ms; + + /// @brief Whether to request a limited-use App Check token. + bool limited_use_app_check_token; +}; + +/// @brief Categories of harm that the model checks for safety ratings and +/// filtering. +enum HarmCategory { + /// Category is unspecified or unrecognized. + kHarmCategoryUnknown = 0, + /// Harassment content. + kHarmCategoryHarassment, + /// Hate speech and content. + kHarmCategoryHateSpeech, + /// Sexually explicit content. + kHarmCategorySexuallyExplicit, + /// Dangerous content. + kHarmCategoryDangerousContent, + /// Content that may harm civic integrity. + kHarmCategoryCivicIntegrity, +}; + +/// @brief Threshold levels for blocking harmful content in `SafetySetting`. +enum HarmBlockThreshold { + /// Threshold is unspecified. + kHarmBlockThresholdUnknown = 0, + /// Block when low, medium, or high probability of harm is detected. + kHarmBlockThresholdLowAndAbove, + /// Block when medium or high probability of harm is detected. + kHarmBlockThresholdMediumAndAbove, + /// Block only when high probability of harm is detected. + kHarmBlockThresholdOnlyHigh, + /// Always show content regardless of harm probability (still rated). + kHarmBlockThresholdNone, + /// Disable the safety filter completely. + kHarmBlockThresholdOff, +}; + +/// @brief Specify how the block threshold should be evaluated in +/// `SafetySetting`. +enum HarmBlockMethod { + /// Method is unspecified. + kHarmBlockMethodUnknown = 0, + /// Consider both probability and severity scores. + kHarmBlockMethodSeverity, + /// Consider only the probability score. + kHarmBlockMethodProbability, +}; + +/// @brief Probability that a given piece of content is harmful. +enum HarmProbability { + /// Probability is unspecified or unrecognized. + kHarmProbabilityUnknown = 0, + /// Content has a negligible chance of being unsafe. + kHarmProbabilityNegligible, + /// Content has a low chance of being unsafe. + kHarmProbabilityLow, + /// Content has a medium chance of being unsafe. + kHarmProbabilityMedium, + /// Content has a high chance of being unsafe. + kHarmProbabilityHigh, +}; + +/// @brief Severity of harm for a piece of content. +enum HarmSeverity { + /// Severity is unspecified or unrecognized. + kHarmSeverityUnknown = 0, + /// Negligible degree of harm. + kHarmSeverityNegligible, + /// Low degree of harm. + kHarmSeverityLow, + /// Medium degree of harm. + kHarmSeverityMedium, + /// High degree of harm. + kHarmSeverityHigh, +}; + +/// @brief Reason why a model stopped generating tokens for a `Candidate`. +enum FinishReason { + /// Finish reason is unspecified or unrecognized. + kFinishReasonUnknown = 0, + /// Natural stop point of the model or provided stop sequence. + kFinishReasonStop, + /// The maximum number of tokens as specified in the request was reached. + kFinishReasonMaxTokens, + /// The token generation was stopped because the response was flagged for + /// safety reasons. + kFinishReasonSafety, + /// The token generation was stopped because the response was flagged for + /// unauthorized citations. + kFinishReasonRecitation, + /// The token generation was stopped for another reason. + kFinishReasonOther, + /// Token generation was stopped because the response contained forbidden + /// terms. + kFinishReasonBlocklist, + /// Token generation was stopped because the response contained potentially + /// prohibited content. + kFinishReasonProhibitedContent, + /// Token generation was stopped because the content potentially contained + /// Sensitive Personally Identifiable Information (SPII). + kFinishReasonSpii, + /// The function call generated by the model is invalid. + kFinishReasonMalformedFunctionCall, +}; + +/// @brief Reason why a prompt was blocked in `PromptFeedback`. +enum BlockReason { + /// Block reason is unspecified or unrecognized. + kBlockReasonUnknown = 0, + /// The prompt was blocked because it was flagged by safety settings. + kBlockReasonSafety, + /// The prompt was blocked for another reason. + kBlockReasonOther, + /// The prompt was blocked because it contained terms from the terminology + /// blocklist. + kBlockReasonBlocklist, + /// The prompt was blocked because it contained prohibited content. + kBlockReasonProhibitedContent, +}; + +/// @brief Content modality type used for token counting details. +enum ContentModality { + /// Unrecognized or unspecified modality. + kContentModalityUnspecified = 0, + /// Plain text. + kContentModalityText, + /// Image. + kContentModalityImage, + /// Video. + kContentModalityVideo, + /// Audio. + kContentModalityAudio, + /// Document (e.g., PDF). + kContentModalityDocument, +}; + +/// @brief Supported response modalities for `GenerationConfig`. +enum ResponseModality { + /// Unspecified modality. + kResponseModalityUnspecified = 0, + /// Text output modality. + kResponseModalityText, + /// Image output modality. + kResponseModalityImage, + /// Audio output modality (reserved). + kResponseModalityAudio, +}; + +/// @brief Thinking level for Gemini 2.5+ thinking models. +enum ThinkingLevel { + /// Unspecified thinking level. + kThinkingLevelUnspecified = 0, + /// Minimal thinking. + kThinkingLevelMinimal, + /// Low thinking. + kThinkingLevelLow, + /// Medium thinking. + kThinkingLevelMedium, + /// High thinking. + kThinkingLevelHigh, +}; + +/// @brief Outcome of a server-side code execution part. +enum CodeExecutionOutcome { + /// Unspecified or unrecognized outcome. + kCodeExecutionOutcomeUnspecified = 0, + /// Code execution completed successfully. + kCodeExecutionOutcomeOk, + /// Code execution finished but with an error. + kCodeExecutionOutcomeFailed, + /// Code execution timed out. + kCodeExecutionOutcomeDeadlineExceeded, +}; + +/// @brief Status of URL retrieval for URL context grounding. +enum UrlRetrievalStatus { + /// Unspecified or unrecognized status. + kUrlRetrievalStatusUnspecified = 0, + /// The URL was retrieved successfully. + kUrlRetrievalStatusSuccess, + /// The URL retrieval failed with an error. + kUrlRetrievalStatusError, + /// The URL could not be retrieved because it is behind a paywall. + kUrlRetrievalStatusPaywall, + /// The URL content was flagged as unsafe. + kUrlRetrievalStatusUnsafe, +}; + +/// @brief Geographical coordinates (latitude and longitude). +struct LatLng { + /// @brief Default constructor initializes coordinates to (0, 0). + LatLng() : latitude(0.0), longitude(0.0) {} + + /// @brief Construct a LatLng with given latitude and longitude. + /// + /// @param latitude Latitude in degrees [-90, 90]. + /// @param longitude Longitude in degrees [-180, 180]. + LatLng(double latitude, double longitude) + : latitude(latitude), longitude(longitude) {} + + /// @brief The latitude in degrees. + double latitude; + + /// @brief The longitude in degrees. + double longitude; +}; + +/// @brief Configuration for retrieval-based tools (such as Google Maps). +struct RetrievalConfig { + /// @brief Optional geographical location of the user. + Optional lat_lng; + + /// @brief Optional language code (BCP-47) for localized results. + Optional language_code; +}; + +/// @brief Determines how `GenerativeModel` routes requests between on-device +/// LiteRT inference and cloud Firebase AI inference. +/// +/// Mirrors `InferenceMode` in the Firebase AI Web SDK (`hybrid-helpers.ts`). +enum InferenceMode { + /// Prefer on-device LiteRT inference if the local model is available; + /// fall back to cloud Firebase AI if unavailable or if local inference fails. + kInferenceModePreferOnDevice = 0, + /// Only use on-device LiteRT inference. Fails with `kErrorUnsupported` if the + /// local model is unavailable. + kInferenceModeOnlyOnDevice = 1, + /// Only use cloud Firebase AI inference (default when no `HybridParams` are + /// configured). + kInferenceModeOnlyInCloud = 2, + /// Prefer cloud Firebase AI inference; fall back to on-device LiteRT + /// inference if the cloud request fails (for example, when offline). + kInferenceModePreferInCloud = 3, +}; + +/// @brief Indicates whether a `GenerateContentResponse` was produced by the +/// cloud backend or by the on-device LiteRT model. +enum InferenceSource { + /// Response was generated by the cloud Firebase AI backend. + kInferenceSourceInCloud = 0, + /// Response was generated locally on-device via LiteRT / LiteRT-LM. + kInferenceSourceOnDevice = 1, +}; + +/// @brief Hardware accelerator selection for Google AI Edge LiteRT +/// (`litert::HwAccelerators`). +enum LiteRtAccelerator { + /// CPU execution (`litert::HwAccelerators::kCpu`). + kLiteRtAcceleratorCpu = 1, + /// GPU execution (`litert::HwAccelerators::kGpu`). + kLiteRtAcceleratorGpu = 2, + /// NPU execution (`litert::HwAccelerators::kNpu`). + kLiteRtAcceleratorNpu = 4, +}; + +/// @brief Configuration for local on-device inference powered by Google AI Edge +/// LiteRT (`https://developers.google.com/edge/litert/overview#c++_1`) and +/// LiteRT-LM for Gemma `.litertlm` / `.tflite` models. +struct OnDeviceParams { + /// @brief Default constructor. + OnDeviceParams() + : accelerator(kLiteRtAcceleratorCpu), + max_num_tokens(0), + num_threads(4), + temperature(0.8f), + top_k(40), + top_p(0.95f) {} + + /// @brief Construct `OnDeviceParams` with a local model path. + /// + /// @param model_path Path to a local `.litertlm` (Gemma) or `.tflite` + /// (LiteRT CompiledModel) file, or `"simulated://gemma-3-270m-it"` for + /// local testing without a downloaded model file. + /// @param accelerator Hardware accelerator (`kLiteRtAcceleratorCpu`, + /// `kLiteRtAcceleratorGpu`, or `kLiteRtAcceleratorNpu`). + explicit OnDeviceParams(const std::string& model_path, + LiteRtAccelerator accelerator = kLiteRtAcceleratorCpu) + : model_path(model_path), + accelerator(accelerator), + max_num_tokens(0), + num_threads(4), + temperature(0.8f), + top_k(40), + top_p(0.95f) {} + + /// @brief File path to the local `.litertlm` or `.tflite` model on disk. + std::string model_path; + + /// @brief Optional path to the LiteRT / LiteRT-LM shared library + /// (`libLiteRt.dylib` / `libCLiteRTLM_mac.dylib` / `.so` / `.dll`). If empty, + /// the runtime searches standard library paths and `FIREBASE_LITERT_LIB_PATH` + /// / `FIREBASE_LITERT_LM_LIB_PATH` environment variables. + std::string runtime_library_path; + + /// @brief Optional directory for caching compiled LiteRT artifacts. + std::string cache_dir; + + /// @brief Hardware accelerator to target (`Cpu`, `Gpu`, `Npu`). + LiteRtAccelerator accelerator; + + /// @brief Maximum context + output token capacity for LiteRT-LM. If 0 + /// (default), automatically queries + /// `litert_lm_loaded_file_max_context_tokens` from the `.litertlm` model file + /// (e.g., 4096 for Gemma 3 1B). + int max_num_tokens; + + /// @brief Number of CPU threads for LiteRT execution. + int num_threads; + + /// @brief Default sampling temperature (overridden by `GenerationConfig` if + /// set). + float temperature; + + /// @brief Default top-K sampling parameter (overridden by `GenerationConfig` + /// if set). + int top_k; + + /// @brief Default top-P nucleus sampling parameter (overridden by + /// `GenerationConfig` if set). + float top_p; +}; + +/// @brief Hybrid inference configuration combining an `InferenceMode` policy +/// with `OnDeviceParams` for LiteRT on-device execution. +struct HybridParams { + /// @brief Default constructor initializes to `kInferenceModeOnlyInCloud`. + HybridParams() : mode(kInferenceModeOnlyInCloud) {} + + /// @brief Construct `HybridParams` with a mode and on-device LiteRT params. + HybridParams(InferenceMode mode, const OnDeviceParams& on_device_params) + : mode(mode), on_device_params(on_device_params) {} + + /// @brief Routing mode between cloud Firebase AI and on-device LiteRT. + InferenceMode mode; + + /// @brief Parameters for the on-device LiteRT / LiteRT-LM model. + OnDeviceParams on_device_params; +}; + +} // namespace ai +} // namespace firebase + +#endif // FIREBASE_AI_SRC_INCLUDE_FIREBASE_AI_TYPES_H_ diff --git a/ai/src/ios/http_sender_ios.mm b/ai/src/ios/http_sender_ios.mm new file mode 100644 index 0000000000..2953daa874 --- /dev/null +++ b/ai/src/ios/http_sender_ios.mm @@ -0,0 +1,193 @@ +/* + * Copyright 2025 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. + */ + +#include "ai/src/common/http_sender.h" + +#import + +#include + +// Delegate for incremental SSE streaming via NSURLSessionDataTask. +@interface FAISessionStreamDelegate : NSObject { + @private + firebase::ai::internal::HttpStreamChunkCallback _onChunk; + firebase::ai::internal::HttpCompletionCallback _onComplete; + int _statusCode; + NSMutableData* _errorData; +} + +- (instancetype) + initWithChunkCallback:(const firebase::ai::internal::HttpStreamChunkCallback&)onChunk + completionCallback:(const firebase::ai::internal::HttpCompletionCallback&)onComplete; + +@end + +@implementation FAISessionStreamDelegate + +- (instancetype) + initWithChunkCallback:(const firebase::ai::internal::HttpStreamChunkCallback&)onChunk + completionCallback:(const firebase::ai::internal::HttpCompletionCallback&)onComplete { + self = [super init]; + if (self) { + _onChunk = onChunk; + _onComplete = onComplete; + _statusCode = 0; + _errorData = [[NSMutableData alloc] init]; + } + return self; +} + +- (void)URLSession:(NSURLSession*)session + dataTask:(NSURLSessionDataTask*)dataTask + didReceiveResponse:(NSURLResponse*)response + completionHandler:(void (^)(NSURLSessionResponseDisposition disposition))completionHandler { + if ([response isKindOfClass:[NSHTTPURLResponse class]]) { + NSHTTPURLResponse* httpResponse = (NSHTTPURLResponse*)response; + _statusCode = static_cast(httpResponse.statusCode); + } + completionHandler(NSURLSessionResponseAllow); +} + +- (void)URLSession:(NSURLSession*)session + dataTask:(NSURLSessionDataTask*)dataTask + didReceiveData:(NSData*)data { + if (!data || data.length == 0) return; + if (_statusCode >= 200 && _statusCode < 300) { + if (_onChunk) { + bool keepGoing = + _onChunk(reinterpret_cast(data.bytes), static_cast(data.length)); + if (!keepGoing) { + [dataTask cancel]; + } + } + } else { + [_errorData appendData:data]; + } +} + +- (void)URLSession:(NSURLSession*)session + task:(NSURLSessionTask*)task + didCompleteWithError:(NSError*)error { + std::string transportError; + if (error) { + if (error.code == NSURLErrorTimedOut) { + transportError = "HTTP request timed out."; + } else if (error.localizedDescription) { + transportError = [error.localizedDescription UTF8String]; + } else { + transportError = "NSURLSession request failed."; + } + } + std::string bodyStr; + if (_errorData && _errorData.length > 0) { + bodyStr.assign(reinterpret_cast(_errorData.bytes), + static_cast(_errorData.length)); + } + if (_onComplete) { + _onComplete(_statusCode, bodyStr, transportError); + } + [session finishTasksAndInvalidate]; +} + +@end + +namespace firebase { +namespace ai { +namespace internal { + +namespace { + +NSMutableURLRequest* BuildUrlRequest(const HttpRequest& request) { + NSString* urlStr = [NSString stringWithUTF8String:request.url.c_str()]; + NSURL* url = [NSURL URLWithString:urlStr]; + NSMutableURLRequest* urlRequest = [NSMutableURLRequest requestWithURL:url]; + urlRequest.HTTPMethod = [NSString stringWithUTF8String:request.method.c_str()]; + if (request.timeout_ms > 0) { + urlRequest.timeoutInterval = static_cast(request.timeout_ms) / 1000.0; + } + for (const auto& kv : request.headers) { + NSString* headerName = [NSString stringWithUTF8String:kv.first.c_str()]; + NSString* headerVal = [NSString stringWithUTF8String:kv.second.c_str()]; + [urlRequest setValue:headerVal forHTTPHeaderField:headerName]; + } + if (!request.body.empty()) { + urlRequest.HTTPBody = [NSData dataWithBytes:request.body.data() length:request.body.size()]; + } + return urlRequest; +} + +} // namespace + +void HttpSender::Initialize() {} + +void HttpSender::Cleanup() {} + +void HttpSender::SendUnary(::firebase::App* /*app*/, const HttpRequest& request, + const HttpCompletionCallback& on_complete) { + @autoreleasepool { + NSMutableURLRequest* urlRequest = BuildUrlRequest(request); + HttpCompletionCallback callbackCopy = on_complete; + NSURLSessionDataTask* task = [[NSURLSession sharedSession] + dataTaskWithRequest:urlRequest + completionHandler:^(NSData* data, NSURLResponse* response, NSError* error) { + int statusCode = 0; + if ([response isKindOfClass:[NSHTTPURLResponse class]]) { + NSHTTPURLResponse* httpResp = (NSHTTPURLResponse*)response; + statusCode = static_cast(httpResp.statusCode); + } + std::string bodyStr; + if (data && data.length > 0) { + bodyStr.assign(reinterpret_cast(data.bytes), + static_cast(data.length)); + } + std::string transportError; + if (error) { + if (error.code == NSURLErrorTimedOut) { + transportError = "HTTP request timed out."; + } else if (error.localizedDescription) { + transportError = [error.localizedDescription UTF8String]; + } else { + transportError = "NSURLSession request failed."; + } + } + if (callbackCopy) { + callbackCopy(statusCode, bodyStr, transportError); + } + }]; + [task resume]; + } +} + +void HttpSender::SendStream(::firebase::App* /*app*/, const HttpRequest& request, + const HttpStreamChunkCallback& on_chunk, + const HttpCompletionCallback& on_complete) { + @autoreleasepool { + NSMutableURLRequest* urlRequest = BuildUrlRequest(request); + FAISessionStreamDelegate* delegate = + [[FAISessionStreamDelegate alloc] initWithChunkCallback:on_chunk + completionCallback:on_complete]; + NSURLSessionConfiguration* config = [NSURLSessionConfiguration defaultSessionConfiguration]; + NSURLSession* session = [NSURLSession sessionWithConfiguration:config + delegate:delegate + delegateQueue:nil]; + NSURLSessionDataTask* task = [session dataTaskWithRequest:urlRequest]; + [task resume]; + } +} + +} // namespace internal +} // namespace ai +} // namespace firebase diff --git a/ai/tests/CMakeLists.txt b/ai/tests/CMakeLists.txt new file mode 100644 index 0000000000..c3c6e65e11 --- /dev/null +++ b/ai/tests/CMakeLists.txt @@ -0,0 +1,25 @@ +# Copyright 2025 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. + +firebase_cpp_cc_test( + firebase_ai_test + SOURCES + ${FIREBASE_SOURCE_DIR}/ai/tests/ai_test.cc + INCLUDES + ${FLATBUFFERS_SOURCE_DIR}/include + DEPENDS + firebase_app_for_testing + firebase_ai + firebase_testing +) diff --git a/ai/tests/ai_test.cc b/ai/tests/ai_test.cc new file mode 100644 index 0000000000..fa3a45dc4d --- /dev/null +++ b/ai/tests/ai_test.cc @@ -0,0 +1,353 @@ +/* + * Copyright 2025 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. + */ + +#include "firebase/ai.h" + +#include +#include +#include +#include +#include + +#include "ai/src/common/http_client.h" +#include "ai/src/common/litert_c_bridge.h" +#include "ai/src/common/serialization.h" +#include "app/src/variant_util.h" +#include "app/tests/include/firebase/app_for_testing.h" +#include "gmock/gmock.h" +#include "gtest/gtest.h" +#include "testing/config.h" + +namespace firebase { +namespace ai { +namespace { + +using ::testing::Eq; +using ::testing::HasSubstr; + +template +void WaitForFuture(const Future& fut) { + while (fut.status() != kFutureStatusComplete) { + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } +} + +TEST(FirebaseAITest, UrlConstructionGoogleAIAndEnterprise) { + AppOptions options; + options.set_app_id("1:123456789:android:abcdef"); + options.set_api_key("test-api-key"); + options.set_project_id("my-test-project"); + + App* app = firebase::testing::CreateApp(options); + ASSERT_NE(app, nullptr); + + std::string google_ai_url = internal::AiHttpClient::ConstructModelUrl( + app, Backend::GoogleAI(), "gemini-2.5-flash", "generateContent"); + EXPECT_THAT(google_ai_url, + Eq("https://firebasevertexai.googleapis.com/v1beta/projects/" + "my-test-project/models/gemini-2.5-flash:generateContent")); + + std::string enterprise_url = internal::AiHttpClient::ConstructModelUrl( + app, Backend::Enterprise("us-central1"), "models/gemini-2.5-pro", + "streamGenerateContent?alt=sse"); + EXPECT_THAT(enterprise_url, + Eq("https://firebasevertexai.googleapis.com/v1beta/projects/" + "my-test-project/locations/us-central1/publishers/google/" + "models/gemini-2.5-pro:streamGenerateContent?alt=sse")); + + std::string template_url = internal::AiHttpClient::ConstructTemplateUrl( + app, Backend::GoogleAI(), "welcome-template", "templateGenerateContent"); + EXPECT_THAT( + template_url, + Eq("https://firebasevertexai.googleapis.com/v1beta/projects/" + "my-test-project/templates/welcome-template:templateGenerateContent")); + + FirebaseAI* ai_google = FirebaseAI::GetInstance(app, Backend::GoogleAI()); + ASSERT_NE(ai_google, nullptr); + EXPECT_EQ(ai_google, FirebaseAI::GetInstance(Backend::GoogleAI())); + + FirebaseAI* ai_vertex = + FirebaseAI::GetInstance(app, Backend::VertexAI("us-central1")); + ASSERT_NE(ai_vertex, nullptr); + EXPECT_NE(ai_google, ai_vertex); + + GenerativeModel model = ai_google->GetGenerativeModel("gemini-2.5-flash"); + std::vector initial_history{ + ModelContent::Text("Hello!"), + ModelContent::Model("Hi there! How can I help?")}; + Chat chat = model.StartChat(initial_history); + EXPECT_EQ(chat.history().size(), 2u); + + delete ai_google; + delete ai_vertex; + delete app; +} + +TEST(FirebaseAITest, SchemaAndJsonSchemaSerialization) { + std::map props; + props["city"] = Schema::String("City name"); + props["units"] = + Schema::Enum(std::vector{"celsius", "fahrenheit"}, + std::string("Temperature units"), true); + std::vector optional_props{"units"}; + Schema obj_schema = + Schema::Object(props, optional_props, std::string("Weather query")); + + Variant open_api_var = internal::SchemaToVariant(obj_schema); + std::string open_api_json = util::VariantToJson(open_api_var); + EXPECT_THAT(open_api_json, HasSubstr("\"type\":\"OBJECT\"")); + EXPECT_THAT(open_api_json, HasSubstr("\"required\":[\"city\"]")); + EXPECT_THAT(open_api_json, + HasSubstr("\"enum\":[\"celsius\",\"fahrenheit\"]")); + + Variant json_schema_var = internal::JsonSchemaToVariant(obj_schema); + std::string std_json_schema = util::VariantToJson(json_schema_var); + EXPECT_THAT(std_json_schema, HasSubstr("\"type\":\"object\"")); + EXPECT_THAT(std_json_schema, HasSubstr("\"type\":[\"string\",\"null\"]")); +} + +TEST(FirebaseAITest, GenerateContentRequestAndCountTokensSerialization) { + GenerationConfig gen_config; + gen_config.temperature = 0.5f; + gen_config.max_output_tokens = 256; + gen_config.thinking_config = ThinkingConfig(128, true); + + std::vector safety_settings{ + SafetySetting(kHarmCategoryHateSpeech, kHarmBlockThresholdOnlyHigh, + kHarmBlockMethodSeverity)}; + + std::map fn_params; + fn_params["location"] = Schema::String("Location name"); + FunctionDeclaration fn_decl("get_weather", "Gets the current weather", + fn_params); + std::vector tools{Tool(fn_decl), Tool(GoogleSearch())}; + + ToolConfig tool_config(FunctionCallingConfig::Auto()); + ModelContent sys_instruction = + ModelContent::System("You are a helpful meteorologist."); + + std::vector contents{ + ModelContent::Text("What is the weather in Boston?")}; + + // GoogleAI omits SafetySetting.method; Enterprise includes it. + std::string google_req_json = internal::BuildGenerateContentRequestJson( + contents, gen_config, safety_settings, tools, tool_config, + sys_instruction, kBackendProviderGoogleAI); + EXPECT_THAT(google_req_json, HasSubstr("\"What is the weather in Boston?\"")); + EXPECT_THAT(google_req_json, HasSubstr("\"thinkingBudget\":128")); + EXPECT_THAT(google_req_json, HasSubstr("\"includeThoughts\":true")); + EXPECT_THAT(google_req_json, HasSubstr("\"googleSearch\":{}")); + EXPECT_THAT(google_req_json, ::testing::Not(HasSubstr("\"SEVERITY\""))); + + std::string enterprise_req_json = internal::BuildGenerateContentRequestJson( + contents, gen_config, safety_settings, tools, tool_config, + sys_instruction, kBackendProviderEnterprise); + EXPECT_THAT(enterprise_req_json, HasSubstr("\"method\":\"SEVERITY\"")); + + // GoogleAI countTokens wraps in generateContentRequest. + std::string google_count_json = internal::BuildCountTokensRequestJson( + "gemini-2.5-flash", contents, gen_config, safety_settings, tools, + tool_config, sys_instruction, kBackendProviderGoogleAI); + EXPECT_THAT(google_count_json, HasSubstr("\"generateContentRequest\"")); + EXPECT_THAT(google_count_json, + HasSubstr("\"model\":\"models/gemini-2.5-flash\"")); +} + +TEST(FirebaseAITest, ParseGenerateContentResponseAndSseStream) { + const char kSampleJson[] = R"({ + "candidates": [{ + "content": { + "role": "model", + "parts": [ + {"text": "Let me check the weather...", "thought": true, "thoughtSignature": "sig-123"}, + {"text": "It is 72F and sunny in Boston."}, + {"functionCall": {"name": "get_forecast", "id": "call-1", "args": {"days": 3}}} + ] + }, + "finishReason": "STOP", + "safetyRatings": [{ + "category": "HARM_CATEGORY_HARASSMENT", + "probability": "NEGLIGIBLE", + "blocked": false + }], + "citationMetadata": { + "citationSources": [{ + "startIndex": 0, + "endIndex": 10, + "uri": "https://example.com" + }] + } + }], + "usageMetadata": { + "promptTokenCount": 12, + "candidatesTokenCount": 20, + "thoughtsTokenCount": 8, + "totalTokenCount": 40, + "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 12}] + } + })"; + + GenerateContentResponse resp; + std::string err; + ASSERT_TRUE(internal::ParseGenerateContentResponseJson( + kSampleJson, kBackendProviderGoogleAI, &resp, &err)) + << err; + EXPECT_THAT(resp.text(), Eq("It is 72F and sunny in Boston.")); + EXPECT_THAT(resp.thought_summary(), Eq("Let me check the weather...")); + ASSERT_EQ(resp.function_calls().size(), 1u); + EXPECT_THAT(resp.function_calls()[0].name, Eq("get_forecast")); + ASSERT_TRUE(resp.function_calls()[0].id.has_value()); + EXPECT_THAT(resp.function_calls()[0].id.value(), Eq("call-1")); + EXPECT_EQ(resp.function_calls()[0].args.at("days").int64_value(), 3); + ASSERT_TRUE(resp.usage_metadata().has_value()); + EXPECT_EQ(resp.usage_metadata()->total_token_count, 40); + EXPECT_EQ(resp.usage_metadata()->thoughts_token_count, 8); + + // Verify SseStreamParser handles fragmented chunks across lines. + std::vector streamed_texts; + internal::SseStreamParser parser( + kBackendProviderGoogleAI, + [&streamed_texts](const GenerateContentResponse& chunk) { + streamed_texts.push_back(chunk.text()); + }); + + std::string sse_part1 = + "data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{" + "\"text\":\"Hello \"}]}}]}\r\ndata: {\"candidates\":[{\"content\":"; + std::string sse_part2 = + "{\"role\":\"model\",\"parts\":[{\"text\":\"World!\"}]}}]}\n\n"; + + parser.Feed(sse_part1.data(), sse_part1.size()); + parser.Feed(sse_part2.data(), sse_part2.size()); + parser.Flush(); + + ASSERT_EQ(streamed_texts.size(), 2u); + EXPECT_THAT(streamed_texts[0], Eq("Hello ")); + EXPECT_THAT(streamed_texts[1], Eq("World!")); +} + +TEST(FirebaseAITest, HybridOnDeviceAndMultiTurnChatToggle) { + AppOptions options; + options.set_app_id("1:123456789:android:abcdef"); + options.set_api_key("invalid-api-key-for-fallback-test"); + options.set_project_id("my-test-project"); + + App* app = firebase::testing::CreateApp(options); + ASSERT_NE(app, nullptr); + + FirebaseAI* ai = FirebaseAI::GetInstance(app, Backend::GoogleAI()); + ASSERT_NE(ai, nullptr); + + OnDeviceParams on_device("simulated://gemma-3-270m-it", + kLiteRtAcceleratorCpu); + HybridParams hybrid_params(kInferenceModeOnlyOnDevice, on_device); + + GenerativeModel model = ai->GetGenerativeModel( + "gemini-2.5-flash", hybrid_params, Optional(), + ModelContent::System("Concise assistant")); + EXPECT_TRUE(model.IsOnDeviceAvailable()); + EXPECT_EQ(model.inference_mode(), kInferenceModeOnlyOnDevice); + + Future init_fut = model.InitializeOnDeviceModel(); + WaitForFuture(init_fut); + EXPECT_EQ(init_fut.error(), kErrorNone); + + // Turn 1 in Chat: ONLY_ON_DEVICE + Chat chat = model.StartChat(); + Future turn1 = + chat.SendMessage("Hello local Gemma!"); + WaitForFuture(turn1); + ASSERT_EQ(turn1.error(), kErrorNone); + ASSERT_NE(turn1.result(), nullptr); + EXPECT_EQ(turn1.result()->inference_source(), kInferenceSourceOnDevice); + EXPECT_THAT(turn1.result()->text(), HasSubstr("gemma-3-270m-it")); + EXPECT_THAT(turn1.result()->text(), HasSubstr("Hello local Gemma!")); + EXPECT_EQ(chat.history().size(), 2u); + + // Toggle Chat live to PREFER_ON_DEVICE and stream Turn 2 + chat.set_inference_mode(kInferenceModePreferOnDevice); + EXPECT_EQ(chat.inference_mode(), kInferenceModePreferOnDevice); + + std::string streamed_reply; + Future turn2_stream = chat.SendMessageStream( + "Second turn via streaming", + [&streamed_reply](const GenerateContentResponse& chunk) { + EXPECT_EQ(chunk.inference_source(), kInferenceSourceOnDevice); + streamed_reply += chunk.text(); + }); + WaitForFuture(turn2_stream); + ASSERT_EQ(turn2_stream.error(), kErrorNone); + EXPECT_THAT(streamed_reply, HasSubstr("turn 2")); + EXPECT_EQ(chat.history().size(), 4u); + + // CompactHistory on-device (replaces 4 turns with 2 compacted turns) + Future compact_fut = chat.CompactHistory(); + WaitForFuture(compact_fut); + ASSERT_EQ(compact_fut.error(), kErrorNone); + ASSERT_NE(compact_fut.result(), nullptr); + EXPECT_EQ(chat.history().size(), 2u); + EXPECT_THAT(chat.history()[0].parts()[0].text_part().text, + HasSubstr("[Compacted Conversation Context]")); + + // ClearHistory + chat.ClearHistory(); + EXPECT_TRUE(chat.history().empty()); + + // CountTokens on-device + Future count_fut = + model.CountTokens("Count these tokens locally"); + WaitForFuture(count_fut); + ASSERT_EQ(count_fut.error(), kErrorNone); + ASSERT_NE(count_fut.result(), nullptr); + EXPECT_GT(count_fut.result()->total_tokens, 0); + + delete ai; + delete app; +} + +TEST(FirebaseAITest, UnityCAbiBridgeLiteRt) { + FirebaseAiLiteRtHandle handle = firebase_ai_litert_create( + "simulated://gemma-3-270m-it", "", "", + static_cast(kLiteRtAcceleratorCpu), 1024, 2, 0.7f, 40, 0.95f); + ASSERT_NE(handle, nullptr); + EXPECT_EQ(firebase_ai_litert_is_available(handle), 1); + + const char* req_json = + "{\"contents\":[{\"role\":\"user\",\"parts\":[{\"text\":\"Ping from " + "Unity C#\"}]}]}"; + char* resp_json = nullptr; + char* err_str = nullptr; + int32_t status = firebase_ai_litert_generate_content(handle, req_json, + &resp_json, &err_str); + EXPECT_EQ(status, 0); + EXPECT_EQ(err_str, nullptr); + ASSERT_NE(resp_json, nullptr); + EXPECT_THAT(std::string(resp_json), HasSubstr("\"ON_DEVICE\"")); + EXPECT_THAT(std::string(resp_json), HasSubstr("Ping from Unity C#")); + firebase_ai_litert_free_string(resp_json); + + int32_t total_tokens = 0; + status = firebase_ai_litert_count_tokens(handle, req_json, &total_tokens, + &err_str); + EXPECT_EQ(status, 0); + EXPECT_GT(total_tokens, 0); + + firebase_ai_litert_destroy(handle); +} + +} // namespace +} // namespace ai +} // namespace firebase diff --git a/app/CMakeLists.txt b/app/CMakeLists.txt index 5c5d735e7e..f41f3968b3 100644 --- a/app/CMakeLists.txt +++ b/app/CMakeLists.txt @@ -128,6 +128,7 @@ set(common_SRCS src/secure/user_secure_manager.cc src/util.cc src/variant.cc + src/variant_util.cc src/base64.cc) if (MSVC) @@ -173,7 +174,6 @@ build_flatbuffers("${desktop_flatbuffers_schemas}" set(app_desktop_SRCS src/app_desktop.cc - src/variant_util.cc src/heartbeat/date_provider.cc src/heartbeat/heartbeat_storage_desktop.cc src/heartbeat/heartbeat_controller_desktop.cc @@ -259,11 +259,11 @@ set(utility_common_HDRS src/semaphore.h src/thread.h src/time.h - src/util.h) + src/util.h + src/variant_util.h) set(utility_android_HDRS) set(utility_ios_HDRS) -set(utility_desktop_HDRS - src/variant_util.h) +set(utility_desktop_HDRS) if(ANDROID) set(utility_HDRS "${utility_common_HDRS}" @@ -490,6 +490,19 @@ if (IOS) ${FIREBASE_SOURCE_DIR}/storage/src/include/firebase/storage/listener.h ${FIREBASE_SOURCE_DIR}/storage/src/include/firebase/storage/metadata.h ${FIREBASE_SOURCE_DIR}/storage/src/include/firebase/storage/storage_reference.h) + set(ai_HDRS + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/chat.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/function_calling.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/generate_content_response.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/generation_config.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/generative_model.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/model_content.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/safety.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/schema.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/template_chat_session.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/template_generative_model.h + ${FIREBASE_SOURCE_DIR}/ai/src/include/firebase/ai/types.h) set(ump_HDRS ${FIREBASE_SOURCE_DIR}/ump/src/include/firebase/ump.h ${FIREBASE_SOURCE_DIR}/ump/src/include/firebase/ump/consent_info.h @@ -497,6 +510,7 @@ if (IOS) list(APPEND framework_HDRS src/include/firebase/internal/platform.h + ${ai_HDRS} ${analytics_HDRS} ${app_check_HDRS} ${auth_HDRS} diff --git a/settings.gradle b/settings.gradle index 5d62cdbd8b..eb3667a9d6 100644 --- a/settings.gradle +++ b/settings.gradle @@ -1,5 +1,6 @@ rootProject.name = 'firebase_cpp_sdk' -include ':app', +include ':ai', + ':app', ':app:app_resources', ':app:google_api_resources', ':app_check', From da8262d243690ae044897b85b60997923e1abfe6 Mon Sep 17 00:00:00 2001 From: Austin Benoit Date: Thu, 8 Oct 2026 13:09:08 -0400 Subject: [PATCH 2/2] feat(ai): default local LiteRT-LM model to gemma-4-E2B-it.litertlm (32k context) --- ai/README.md | 12 ++++++------ ai/samples/hybrid_chat_main.cc | 10 +++++++++- 2 files changed, 15 insertions(+), 7 deletions(-) diff --git a/ai/README.md b/ai/README.md index bee9dac3da..dc21a67064 100644 --- a/ai/README.md +++ b/ai/README.md @@ -27,11 +27,11 @@ cmake --build desktop_build \ ### 2. Download a Local Gemma `.litertlm` Model -Download a LiteRT-LM `.litertlm` model (for example, **Gemma 3 1B IT INT4** with a 4,096-token context window) and place it in `desktop_build/ai/` or pass its path via `--model`: +Download **Gemma 4 E2B IT** (`gemma-4-E2B-it.litertlm` from [`litert-community/gemma-4-E2B-it-litert-lm`](https://huggingface.co/litert-community/gemma-4-E2B-it-litert-lm), with a **32k context window**) and place it in `desktop_build/ai/` or pass its path via `--model`: ```bash -curl -L "https://huggingface.co/litert-community/Gemma3-1B-IT/resolve/main/gemma3-1b-it-int4.litertlm" \ - -o desktop_build/ai/gemma3-1b-it-int4.litertlm +curl -L "https://huggingface.co/litert-community/gemma-4-E2B-it-litert-lm/resolve/main/gemma-4-E2B-it.litertlm" \ + -o desktop_build/ai/gemma-4-E2B-it.litertlm ``` ### 3. Provide Your Firebase Configuration (`google-services.json`) @@ -43,10 +43,10 @@ The demo initializes `firebase::App` using a standard Firebase `google-services. ```bash ./desktop_build/ai/firebase_ai_hybrid_chat \ --config /path/to/google-services.json \ - --model ./desktop_build/ai/gemma3-1b-it-int4.litertlm + --model ./desktop_build/ai/gemma-4-E2B-it.litertlm ``` -*(If `google-services.json` and `gemma3-1b-it-int4.litertlm` are placed in `desktop_build/ai/` or the current directory, they are auto-detected and you can run `./desktop_build/ai/firebase_ai_hybrid_chat` with no arguments.)* +*(If `google-services.json` and `gemma-4-E2B-it.litertlm` are placed in `desktop_build/ai/` or the current directory, they are auto-detected and you can run `./desktop_build/ai/firebase_ai_hybrid_chat` with no arguments.)* --- @@ -80,7 +80,7 @@ Inside `firebase_ai_hybrid_chat`, both Cloud (`gemini-3.1-flash-lite`) and On-De ## Context Window & Automatic Compaction -- **Auto-Detected Context Length:** When `OnDeviceParams::max_num_tokens` is `0` (the default), `LiteRtAdapter` queries `litert_lm_loaded_file_max_context_tokens` from the `.litertlm` file metadata (`4096` tokens for `gemma3-1b-it-int4.litertlm`, `1024` tokens for `gemma3-270m.litertlm`). +- **Auto-Detected Context Length:** When `OnDeviceParams::max_num_tokens` is `0` (the default), `LiteRtAdapter` queries `litert_lm_loaded_file_max_context_tokens` from the `.litertlm` file metadata (`32003` tokens for `gemma-4-E2B-it.litertlm`, `4096` tokens for `gemma3-1b-it-int4.litertlm`, `1024` tokens for `gemma3-270m.litertlm`). - **Automatic Context Compaction:** Before each on-device turn, `LiteRtAdapter` tokenizes the conversation history via `litert_lm_engine_tokenize`. If the accumulated history exceeds the input token budget, older turns are automatically compacted into `[Compacted Earlier Conversation History]` while keeping recent turns verbatim and reserving headroom for generation output. - **Repetition Prevention:** On-device generation configures `LiteRtLmRepetitionPenaltyConfig` (`repetition_penalty = 1.15`, `frequency_penalty = 0.25`, `presence_penalty = 0.1`) and `LiteRtLmNoRepeatNgramConfig` (`no_repeat_ngram_size = 4`) so small quantized models do not fall into token repetition loops on long outputs. diff --git a/ai/samples/hybrid_chat_main.cc b/ai/samples/hybrid_chat_main.cc index 1da7452ea6..0a2aab00e8 100644 --- a/ai/samples/hybrid_chat_main.cc +++ b/ai/samples/hybrid_chat_main.cc @@ -83,6 +83,14 @@ std::string FindDefaultLocalModel(const char* argv0) { } std::string bin_dir = argv0 ? ParentDir(argv0) : "."; std::vector candidates = { + bin_dir + "/gemma-4-E2B-it.litertlm", + bin_dir + "/gemma-4-E2B-it-gpu.litertlm", + "./gemma-4-E2B-it.litertlm", + "./gemma-4-E2B-it-gpu.litertlm", + "./desktop_build/ai/gemma-4-E2B-it.litertlm", + "./desktop_build/ai/gemma-4-E2B-it-gpu.litertlm", + "./firebase-cpp-sdk/desktop_build/ai/gemma-4-E2B-it.litertlm", + "./firebase-cpp-sdk/desktop_build/ai/gemma-4-E2B-it-gpu.litertlm", bin_dir + "/gemma3-1b-it-int4.litertlm", bin_dir + "/gemma3-270m.litertlm", "./gemma3-1b-it-int4.litertlm", @@ -97,7 +105,7 @@ std::string FindDefaultLocalModel(const char* argv0) { return candidate; } } - return "simulated://gemma-3-270m-it"; + return "simulated://gemma-4-E2B-it"; } std::string ReadFileToString(const std::string& path) {