rig-core 0.44.0

An opinionated library for building LLM powered applications.
Documentation
//! Image-generation requests, normalized responses, and model interfaces.
//!
//! ```no_run
//! use rig_core::DynModel;
//! use rig_core::image_generation::ImageGenerationRequestBuilder;
//! use rig_core::operation::ImageGeneration;
//!
//! # async fn example(model: DynModel<ImageGeneration>) -> Result<(), Box<dyn std::error::Error>> {
//! let request = ImageGenerationRequestBuilder::new("A mountain lake")
//!     .width(1024)
//!     .height(1024)
//!     .build();
//! let response = model.call(request).await?;
//! # let _ = response;
//! # Ok(())
//! # }
//! ```
use crate::completion::Usage;
use crate::error::ProviderError;
use serde::{Deserialize, Serialize};
use serde_json::Value;

/// Generated image bytes and normalized provider metadata.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageGenerationResponse {
    /// The generated image, decoded to bytes.
    pub image: Vec<u8>,
    /// Usage as the provider reported it; every counter is `None` when the
    /// provider reported none (see [`Usage`]).
    #[serde(default)]
    pub usage: Usage,
    /// Stable descriptor name of the provider that produced this response,
    /// for example `"openai"`. Always populated.
    pub provider: String,
    /// Provider-reported model identifier, when the wire response named one.
    #[serde(default)]
    pub model: Option<String>,
    /// Provider-assigned response-scoped identifier, when reported.
    #[serde(default, skip_serializing_if = "Option::is_none")]
    pub response_id: Option<String>,
    /// Transport request ID from HTTP headers, or `None` when unreported.
    #[serde(default, skip_serializing_if = "Option::is_none")]
    pub provider_request_id: Option<String>,
    /// Provider response metadata. May be null for byte-only responses or
    /// responses constructed without metadata; image bytes remain in [`Self::image`].
    #[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
    pub raw: serde_json::Value,
}

impl ImageGenerationResponse {
    /// A response carrying `image`. The driver writes the provider, the
    /// transport request id and the reply document; decoders set what the
    /// provider reported.
    pub fn new(image: Vec<u8>) -> Self {
        Self {
            image,
            usage: Usage::default(),
            provider: String::new(),
            model: None,
            response_id: None,
            provider_request_id: None,
            raw: serde_json::Value::Null,
        }
    }
}

/// Normalizes provider image payloads. The driver writes the provider, request
/// id and `raw` afterwards.
pub trait NormalizeImageGenerationResponse {
    /// Normalize this payload.
    fn normalize(self) -> Result<ImageGenerationResponse, ProviderError>;
}

pub struct ImageGenerationRequest {
    pub prompt: String,
    pub width: u32,
    pub height: u32,
    pub additional_params: Option<Value>,
}

/// Builds an image request for a prompt. Defaults to 256 by 256 pixels;
/// supported dimensions depend on the provider.
pub struct ImageGenerationRequestBuilder {
    request: ImageGenerationRequest,
}

impl ImageGenerationRequestBuilder {
    /// A request for `prompt`.
    pub fn new(prompt: impl Into<String>) -> Self {
        Self {
            request: ImageGenerationRequest {
                prompt: prompt.into(),
                width: 256,
                height: 256,
                additional_params: None,
            },
        }
    }

    /// The width of the generated image.
    pub fn width(mut self, width: u32) -> Self {
        self.request.width = width;
        self
    }

    /// The height of the generated image.
    pub fn height(mut self, height: u32) -> Self {
        self.request.height = height;
        self
    }

    /// Merges provider-specific parameters over earlier ones, key by key for
    /// JSON objects; `None` clears existing parameters.
    pub fn additional_params(mut self, params: impl Into<Option<Value>>) -> Self {
        self.request.additional_params =
            crate::json_utils::merge_params(self.request.additional_params.take(), params.into());
        self
    }

    /// Builds the image generation request.
    pub fn build(self) -> ImageGenerationRequest {
        self.request
    }
}

#[cfg(test)]
mod builder_tests;
#[cfg(test)]
mod provider_response_tests;