diff --git a/packages/firebase_ai/firebase_ai/lib/src/imagen/imagen_api.dart b/packages/firebase_ai/firebase_ai/lib/src/imagen/imagen_api.dart index b1df92e15423..84c759cf444d 100644 --- a/packages/firebase_ai/firebase_ai/lib/src/imagen/imagen_api.dart +++ b/packages/firebase_ai/firebase_ai/lib/src/imagen/imagen_api.dart @@ -108,7 +108,7 @@ final class ImagenSafetySettings { final ImagenPersonFilterLevel? personFilterLevel; // ignore: public_member_api_docs - Object toJson() => { + Map toJson() => { if (safetyFilterLevel != null) 'safetySetting': safetyFilterLevel!.toJson(), if (personFilterLevel != null) @@ -194,7 +194,7 @@ final class ImagenGenerationConfig { // ignore: public_member_api_docs Map toJson() => { if (negativePrompt != null) 'negativePrompt': negativePrompt, - if (numberOfImages != null) 'numberOfImages': numberOfImages, + 'sampleCount': numberOfImages ?? 1, if (aspectRatio != null) 'aspectRatio': aspectRatio!.toJson(), if (addWatermark != null) 'addWatermark': addWatermark, if (imageFormat != null) 'outputOptions': imageFormat!.toJson(), diff --git a/packages/firebase_ai/firebase_ai/lib/src/imagen/imagen_model.dart b/packages/firebase_ai/firebase_ai/lib/src/imagen/imagen_model.dart index 2957c056522c..4fc6e84d2626 100644 --- a/packages/firebase_ai/firebase_ai/lib/src/imagen/imagen_model.dart +++ b/packages/firebase_ai/firebase_ai/lib/src/imagen/imagen_model.dart @@ -61,17 +61,17 @@ final class ImagenModel extends BaseApiClientModel { if (gcsUri != null) 'storageUri': gcsUri, 'sampleCount': _generationConfig?.numberOfImages ?? 1, if (_generationConfig?.aspectRatio case final aspectRatio?) - 'aspectRatio': aspectRatio, + 'aspectRatio': aspectRatio.toJson(), if (_generationConfig?.negativePrompt case final negativePrompt?) 'negativePrompt': negativePrompt, if (_generationConfig?.addWatermark case final addWatermark?) 'addWatermark': addWatermark, if (_generationConfig?.imageFormat case final imageFormat?) 'outputOption': imageFormat.toJson(), - if (_safetySettings?.personFilterLevel case final personFilterLevel?) - 'personGeneration': personFilterLevel.toJson(), - if (_safetySettings?.safetyFilterLevel case final safetyFilterLevel?) - 'safetySetting': safetyFilterLevel.toJson(), + if (_safetySettings case final safetySettings?) + ...safetySettings.toJson(), + 'includeRaiReason': true, + 'includeSafetyAttributes': true, }; return { @@ -170,10 +170,10 @@ final class ImagenModel extends BaseApiClientModel { 'addWatermark': addWatermark, if (_generationConfig?.imageFormat case final imageFormat?) 'outputOption': imageFormat.toJson(), - if (_safetySettings?.personFilterLevel case final personFilterLevel?) - 'personGeneration': personFilterLevel.toJson(), - if (_safetySettings?.safetyFilterLevel case final safetyFilterLevel?) - 'safetySetting': safetyFilterLevel.toJson(), + if (_safetySettings case final safetySettings?) + ...safetySettings.toJson(), + 'includeRaiReason': true, + 'includeSafetyAttributes': true, }; return { diff --git a/packages/firebase_ai/firebase_ai/test/imagen_model_test.dart b/packages/firebase_ai/firebase_ai/test/imagen_model_test.dart new file mode 100644 index 000000000000..c121ccaddb2b --- /dev/null +++ b/packages/firebase_ai/firebase_ai/test/imagen_model_test.dart @@ -0,0 +1,281 @@ +// Copyright 2024 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. + +import 'dart:typed_data'; + +import 'package:firebase_ai/firebase_ai.dart'; +import 'package:flutter_test/flutter_test.dart'; + +// Copied from imagen_model.dart for testing purposes as it is a private method. +Map generateImagenRequest( + String prompt, { + String? gcsUri, + ImagenGenerationConfig? generationConfig, + ImagenSafetySettings? safetySettings, +}) { + final parameters = { + if (gcsUri != null) 'storageUri': gcsUri, + 'sampleCount': generationConfig?.numberOfImages ?? 1, + if (generationConfig?.aspectRatio case final aspectRatio?) + 'aspectRatio': aspectRatio.toJson(), + if (generationConfig?.negativePrompt case final negativePrompt?) + 'negativePrompt': negativePrompt, + if (generationConfig?.addWatermark case final addWatermark?) + 'addWatermark': addWatermark, + if (generationConfig?.imageFormat case final imageFormat?) + 'outputOption': imageFormat.toJson(), + if (safetySettings case final safetySettings?) ...safetySettings.toJson(), + 'includeRaiReason': true, + 'includeSafetyAttributes': true, + }; + + return { + 'instances': [ + {'prompt': prompt} + ], + 'parameters': parameters, + }; +} + +// Copied from imagen_model.dart for testing +Map generateImagenEditRequest( + List images, + String prompt, { + bool useVertexBackend = true, // Added for testing the throw + ImagenEditingConfig? config, + ImagenGenerationConfig? generationConfig, + ImagenSafetySettings? safetySettings, +}) { + if (!useVertexBackend) { + throw FirebaseAIException( + 'Image editing for Imagen is only supported on Vertex AI backend.'); + } + final parameters = { + 'sampleCount': generationConfig?.numberOfImages ?? 1, + if (config?.editMode case final editMode?) 'editMode': editMode.toJson(), + if (config?.editSteps case final editSteps?) + 'editConfig': {'baseSteps': editSteps}, + if (generationConfig?.negativePrompt case final negativePrompt?) + 'negativePrompt': negativePrompt, + if (generationConfig?.addWatermark case final addWatermark?) + 'addWatermark': addWatermark, + if (generationConfig?.imageFormat case final imageFormat?) + 'outputOption': imageFormat.toJson(), + if (safetySettings case final safetySettings?) ...safetySettings.toJson(), + 'includeRaiReason': true, + 'includeSafetyAttributes': true, + }; + + return { + 'parameters': parameters, + 'instances': [ + { + 'prompt': prompt, + 'referenceImages': images.asMap().entries.map((entry) { + int index = entry.key; + var image = entry.value; + return image.toJson(referenceIdOverrideIfNull: index + images.length); + }).toList(), + } + ], + }; +} + +void main() { + group('ImagenModel request generation', () { + group('generateImagenRequest', () { + test('creates a basic request with default parameters', () { + final request = generateImagenRequest('a beautiful landscape'); + expect(request['instances'], [ + {'prompt': 'a beautiful landscape'} + ]); + final params = request['parameters']! as Map; + expect(params['sampleCount'], 1); + expect(params['includeRaiReason'], true); + expect(params['includeSafetyAttributes'], true); + expect(params.containsKey('storageUri'), isFalse); + expect(params.containsKey('aspectRatio'), isFalse); + expect(params.containsKey('negativePrompt'), isFalse); + expect(params.containsKey('addWatermark'), isFalse); + expect(params.containsKey('outputOption'), isFalse); + expect(params.containsKey('personGeneration'), isFalse); + expect(params.containsKey('safetySetting'), isFalse); + }); + + test('includes all generation config parameters', () { + final config = ImagenGenerationConfig( + numberOfImages: 4, + aspectRatio: ImagenAspectRatio.landscape16x9, + negativePrompt: 'text, watermark', + addWatermark: false, + imageFormat: ImagenFormat.png(), + ); + final request = generateImagenRequest('a futuristic city', + generationConfig: config); + final params = request['parameters']! as Map; + expect(params['sampleCount'], 4); + expect(params['aspectRatio'], '16:9'); + expect(params['negativePrompt'], 'text, watermark'); + expect(params['addWatermark'], false); + expect(params['outputOption'], {'mimeType': 'image/png'}); + expect(params['includeRaiReason'], true); + expect(params['includeSafetyAttributes'], true); + }); + + test('includes all safety settings parameters', () { + final settings = ImagenSafetySettings( + ImagenSafetyFilterLevel.blockNone, + ImagenPersonFilterLevel.allowAdult, + ); + final request = + generateImagenRequest('a robot army', safetySettings: settings); + final params = request['parameters']! as Map; + expect(params['personGeneration'], 'allow_adult'); + expect(params['safetySetting'], 'block_none'); + expect(params['includeRaiReason'], true); + expect(params['includeSafetyAttributes'], true); + }); + + test('includes gcsUri when provided', () { + const uri = 'gs://my-test-bucket/image.png'; + final request = generateImagenRequest('a photo of a cat', gcsUri: uri); + final params = request['parameters']! as Map; + expect(params['storageUri'], uri); + expect(params['includeRaiReason'], true); + expect(params['includeSafetyAttributes'], true); + }); + + test('combines all parameters correctly', () { + final config = ImagenGenerationConfig( + numberOfImages: 2, + negativePrompt: 'dark', + ); + final settings = ImagenSafetySettings( + ImagenSafetyFilterLevel.blockLowAndAbove, + ImagenPersonFilterLevel.blockAll, + ); + const uri = 'gs://my-test-bucket/output/'; + final request = generateImagenRequest( + 'a sunny beach', + gcsUri: uri, + generationConfig: config, + safetySettings: settings, + ); + + final params = request['parameters']! as Map; + expect(params['storageUri'], uri); + expect(params['sampleCount'], 2); + expect(params['negativePrompt'], 'dark'); + expect(params['safetySetting'], 'block_low_and_above'); + expect(params['includeRaiReason'], true); + expect(params['includeSafetyAttributes'], true); + expect(request['instances'], [ + {'prompt': 'a sunny beach'} + ]); + }); + }); + + group('generateImagenEditRequest', () { + late List referenceImages; + + setUp(() { + final dummyBytes = Uint8List.fromList([1, 2, 3]); + final dummyInlineImage = ImagenInlineImage( + bytesBase64Encoded: dummyBytes, mimeType: 'image/jpeg'); + referenceImages = [ImagenRawImage(image: dummyInlineImage)]; + }); + + test('creates a basic edit request', () { + final request = + generateImagenEditRequest(referenceImages, 'make it sunny'); + final params = request['parameters']! as Map; + expect(params['sampleCount'], 1); + expect(params.containsKey('editMode'), isFalse); + expect(params['includeRaiReason'], true); + expect(params['includeSafetyAttributes'], true); + + final instances = request['instances']! as List; + expect(instances, hasLength(1)); + final instance = instances.first as Map; + expect(instance['prompt'], 'make it sunny'); + expect(instance['referenceImages'], isNotNull); + }); + + test('does not include aspectRatio from generation config', () { + final config = ImagenGenerationConfig( + numberOfImages: 2, // This should be included as sampleCount + aspectRatio: ImagenAspectRatio.square1x1, // This should be ignored + ); + final request = generateImagenEditRequest( + referenceImages, + 'add a rainbow', + generationConfig: config, + ); + final params = request['parameters']! as Map; + expect(params['sampleCount'], 2); + expect(params.containsKey('aspectRatio'), isFalse, + reason: 'aspectRatio is not a valid parameter for edit requests.'); + expect(params['includeRaiReason'], true); + expect(params['includeSafetyAttributes'], true); + }); + + test('includes other valid generation config values', () { + final config = ImagenGenerationConfig( + negativePrompt: 'rain', + addWatermark: true, + imageFormat: ImagenFormat.jpeg(), + ); + final request = generateImagenEditRequest( + referenceImages, + 'make it brighter', + generationConfig: config, + ); + final params = request['parameters']! as Map; + expect(params['negativePrompt'], 'rain'); + expect(params['addWatermark'], true); + expect(params['outputOption'], {'mimeType': 'image/jpeg'}); + expect(params['includeRaiReason'], true); + expect(params['includeSafetyAttributes'], true); + }); + + test('includes editing config', () { + final editConfig = ImagenEditingConfig( + editMode: ImagenEditMode.inpaintInsertion, + editSteps: 10, + ); + final request = generateImagenEditRequest( + referenceImages, + 'remove the background', + config: editConfig, + ); + final params = request['parameters']! as Map; + expect(params['editMode'], 'EDIT_MODE_INPAINT_INSERTION'); + expect(params['editConfig'], {'baseSteps': 10}); + expect(params['includeRaiReason'], true); + expect(params['includeSafetyAttributes'], true); + }); + + test('throws exception if not using Vertex backend', () { + expect( + () => generateImagenEditRequest( + referenceImages, + 'a prompt', + useVertexBackend: false, + ), + throwsA(isA()), + ); + }); + }); + }); +} diff --git a/packages/firebase_ai/firebase_ai/test/imagen_test.dart b/packages/firebase_ai/firebase_ai/test/imagen_test.dart index 6f26baffd333..5cdef9734dff 100644 --- a/packages/firebase_ai/firebase_ai/test/imagen_test.dart +++ b/packages/firebase_ai/firebase_ai/test/imagen_test.dart @@ -195,7 +195,7 @@ void main() { final json = config.toJson(); expect(json, { 'negativePrompt': 'blurry, low quality', - 'numberOfImages': 4, + 'sampleCount': 4, 'aspectRatio': '16:9', 'addWatermark': true, 'outputOptions': { @@ -210,9 +210,7 @@ void main() { negativePrompt: 'blurry, low quality', ); final json = config.toJson(); - expect(json, { - 'negativePrompt': 'blurry, low quality', - }); + expect(json, {'negativePrompt': 'blurry, low quality', 'sampleCount': 1}); }); test('toJson with only numberOfImages', () { @@ -221,7 +219,7 @@ void main() { ); final json = config.toJson(); expect(json, { - 'numberOfImages': 2, + 'sampleCount': 2, }); }); @@ -230,9 +228,7 @@ void main() { aspectRatio: ImagenAspectRatio.portrait9x16, ); final json = config.toJson(); - expect(json, { - 'aspectRatio': '9:16', - }); + expect(json, {'aspectRatio': '9:16', 'sampleCount': 1}); }); test('toJson with only imageFormat', () { @@ -244,6 +240,7 @@ void main() { 'outputOptions': { 'mimeType': 'image/png', }, + 'sampleCount': 1 }); }); @@ -252,15 +249,13 @@ void main() { addWatermark: false, ); final json = config.toJson(); - expect(json, { - 'addWatermark': false, - }); + expect(json, {'addWatermark': false, 'sampleCount': 1}); }); test('toJson with empty config', () { final config = ImagenGenerationConfig(); final json = config.toJson(); - expect(json, {}); + expect(json, {'sampleCount': 1}); }); test('toJson with imageFormat uses correct key name "outputOptions"', () {