contaigents
Version:
Modular AI Content Ecosystem with Audio Generation
237 lines (236 loc) • 8.4 kB
JavaScript
/**
* Image Generator Utility
* Provides a unified interface for generating images from text using different AI providers
*/
import { writeFile } from 'fs';
import { promisify } from 'util';
import path from 'path';
// Convert callback-based writeFile to Promise-based
const writeFileAsync = promisify(writeFile);
/**
* ImageGenerator class for text-to-image conversion using various AI providers
*/
export class ImageGenerator {
/**
* Create a new ImageGenerator instance
* @param {Object} config - Configuration options
* @param {string} config.provider - AI provider ('gemini', 'openai', etc.)
* @param {string} config.apiKey - API key for the provider
* @param {string} config.model - Model to use for image generation
*/
constructor(config) {
this.lastMimeType = '';
this.provider = config.provider || 'gemini';
this.apiKey = config.apiKey;
this.model = config.model || this.getDefaultModel();
}
/**
* Get default model for the provider
*/
getDefaultModel() {
switch (this.provider.toLowerCase()) {
case 'gemini':
case 'google':
return 'imagen-3.0-generate-002';
case 'openai':
return 'dall-e-3';
default:
return 'imagen-3.0-generate-002';
}
}
/**
* Generate image from text prompt
* @param {string} prompt - Text prompt for image generation
* @param {Object} options - Generation options
* @returns {Promise<Buffer>} - Generated image as buffer
*/
async generateImage(prompt, options = {}) {
if (!prompt || typeof prompt !== 'string' || !prompt.trim()) {
throw new Error('Prompt cannot be empty');
}
switch (this.provider.toLowerCase()) {
case 'gemini':
case 'google':
return this._generateWithGemini(prompt, options);
case 'openai':
return this._generateWithOpenAI(prompt, options);
default:
throw new Error(`Unsupported provider: ${this.provider}`);
}
}
/**
* Generate image using Gemini/Google provider (Imagen models)
*/
async _generateWithGemini(prompt, options) {
// Use the Imagen API endpoint with the correct format
const response = await fetch(`https://generativelanguage.googleapis.com/v1beta/models/${this.model}:predict?key=${this.apiKey}`, {
method: 'POST',
headers: {
'Content-Type': 'application/json',
},
body: JSON.stringify({
instances: [
{
prompt: prompt
}
],
parameters: {
sampleCount: 1,
...(options.aspectRatio && { aspectRatio: options.aspectRatio }),
...(options.imageSize && { imageSize: options.imageSize })
}
}),
});
if (!response.ok) {
const errorText = await response.text();
let errorMessage = `Imagen API error: ${response.status} ${response.statusText}`;
try {
const errorData = JSON.parse(errorText);
errorMessage = `Imagen API error: ${errorData.error?.message || errorText}`;
}
catch {
errorMessage = `Imagen API error: ${errorText}`;
}
throw new Error(errorMessage);
}
const data = await response.json();
// Extract image data from Imagen API response
if (!data.predictions || !data.predictions[0]) {
console.error('Imagen API response structure:', JSON.stringify(data, null, 2));
throw new Error('No image data received from Imagen API');
}
const prediction = data.predictions[0];
// The image data is in bytesBase64Encoded field
if (prediction.bytesBase64Encoded) {
this.lastMimeType = prediction.mimeType || 'image/png';
return Buffer.from(prediction.bytesBase64Encoded, 'base64');
}
// Log the actual response structure for debugging
console.error('Imagen API response structure:', JSON.stringify(data, null, 2));
throw new Error('Invalid image data format from Imagen API');
}
/**
* Generate image using OpenAI provider
*/
async _generateWithOpenAI(prompt, options) {
// Standard text-to-image generation
const requestBody = {
prompt: prompt,
model: this.model,
n: 1,
response_format: 'b64_json',
...(options.width && options.height && {
size: `${options.width}x${options.height}`
}),
...(options.quality && { quality: options.quality }),
...(options.style && { style: options.style })
};
const response = await fetch('https://api.openai.com/v1/images/generations', {
method: 'POST',
headers: {
'Content-Type': 'application/json',
'Authorization': `Bearer ${this.apiKey}`,
},
body: JSON.stringify(requestBody),
});
if (!response.ok) {
const errorText = await response.text();
let errorMessage = `OpenAI API error: ${response.status} ${response.statusText}`;
try {
const errorData = JSON.parse(errorText);
errorMessage = `OpenAI API error: ${errorData.error?.message || errorText}`;
}
catch {
errorMessage = `OpenAI API error: ${errorText}`;
}
throw new Error(errorMessage);
}
const data = await response.json();
if (!data.data || !data.data[0] || !data.data[0].b64_json) {
throw new Error('No image data received from OpenAI API');
}
this.lastMimeType = 'image/png';
return Buffer.from(data.data[0].b64_json, 'base64');
}
/**
* Save image buffer to file
* @param {string} filePath - Path to save the image
* @param {Buffer} imageBuffer - Image data buffer
* @returns {Promise<string>} - Path to saved file
*/
async saveImageToFile(filePath, imageBuffer) {
try {
// Ensure the file has an appropriate extension
const ext = path.extname(filePath);
if (!ext) {
const mimeExt = this._getMimeTypeExtension(this.lastMimeType);
filePath = `${filePath}.${mimeExt}`;
}
await writeFileAsync(filePath, imageBuffer);
return filePath;
}
catch (error) {
throw new Error(`Failed to save image file: ${error.message}`);
}
}
/**
* Get file extension from MIME type
*/
_getMimeTypeExtension(mimeType) {
switch (mimeType) {
case 'image/png':
return 'png';
case 'image/jpeg':
case 'image/jpg':
return 'jpg';
case 'image/webp':
return 'webp';
case 'image/gif':
return 'gif';
default:
return 'png';
}
}
/**
* Get available models for the current provider
*/
getAvailableModels() {
switch (this.provider.toLowerCase()) {
case 'gemini':
case 'google':
return [
'imagen-3.0-generate-002',
'imagen-4.0-generate-preview-06-06',
'imagen-4.0-ultra-generate-preview-06-06'
];
case 'openai':
return ['dall-e-3', 'dall-e-2'];
default:
return [this.model];
}
}
/**
* Check if a model is available for the current provider
*/
isModelAvailable(model) {
return this.getAvailableModels().includes(model);
}
/**
* Get the current provider
*/
getProvider() {
return this.provider;
}
/**
* Get the current model
*/
getModel() {
return this.model;
}
/**
* Get the last generated image MIME type
*/
getLastMimeType() {
return this.lastMimeType;
}
}