use crate::{AiConfig, AiResult};
use serde::{Deserialize, Serialize};
pub use wae_types::{BillingDimensions, Decimal};
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", content = "value")]
pub enum AiImageInput {
#[serde(rename = "text")]
Text(String),
#[serde(rename = "image")]
Image(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LoRAConfig {
pub model_id: String,
pub weight: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AiImageTask {
pub inputs: Vec<AiImageInput>,
pub loras: Vec<LoRAConfig>,
pub width: Option<u32>,
pub height: Option<u32>,
pub num_images: Option<u32>,
pub negative_prompt: Option<String>,
pub seed: Option<i64>,
pub extra: Option<serde_json::Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AiImageOutput {
pub images: Vec<String>,
pub text: Option<String>,
pub usage: BillingDimensions,
}
pub trait AiImageCapability: Send + Sync {
#[allow(async_fn_in_trait)]
async fn generate_image(&self, task: &AiImageTask, config: &AiConfig) -> AiResult<AiImageOutput>;
fn estimate_usage(&self, task: &AiImageTask) -> AiResult<BillingDimensions> {
let output_pixels =
(task.width.unwrap_or(1024) as u64) * (task.height.unwrap_or(1024) as u64) * (task.num_images.unwrap_or(1) as u64);
let mut input_text = 0;
let mut input_pixels = 0;
for input in &task.inputs {
match input {
AiImageInput::Text(t) => input_text += t.len() as u64,
AiImageInput::Image(_) => input_pixels += 1024 * 1024,
}
}
Ok(BillingDimensions { input_text, output_text: 0, input_pixels, output_pixels })
}
}