Packages
langchain
0.8.13
0.9.2
0.9.1
0.9.0
0.8.14
0.8.13
0.8.12
0.8.11
0.8.10
0.8.9
0.8.8
0.8.7
0.8.6
0.8.5
0.8.4
0.8.3
0.8.2
0.8.1
0.8.0
0.7.0
0.6.3
0.6.2
0.6.1
0.6.0
0.5.2
0.5.1
0.5.0
0.4.1
0.4.0
0.4.0-rc.3
0.4.0-rc.2
0.4.0-rc.1
0.4.0-rc.0
0.3.3
0.3.2
0.3.1
0.3.0
0.3.0-rc.2
0.3.0-rc.1
0.3.0-rc.0
0.2.0
0.1.10
0.1.9
0.1.8
0.1.7
0.1.6
0.1.5
0.1.4
0.1.3
0.1.2
0.1.1
0.1.0
Elixir implementation of a LangChain style framework that lets Elixir projects integrate with and leverage LLMs.
Current section
Files
Jump to
Current section
Files
lib/images/open_ai_image.ex
defmodule LangChain.Images.OpenAIImage do
@moduledoc """
Represents the [OpenAI Images API
endpoint](https://platform.openai.com/docs/api-reference/images) for working
with DALL-E-2 and DALL-E-3.
Parses and validates inputs for making a request from the OpenAI Image API.
Converts responses into more specialized `LangChain` data structures and
provide functions for saving generated images.
"""
use Ecto.Schema
require Logger
import Ecto.Changeset
alias __MODULE__
alias LangChain.Images.GeneratedImage
alias LangChain.Config
alias LangChain.LangChainError
alias LangChain.Utils
# allow up to 2 minutes for response.
@receive_timeout 1_200_000
@primary_key false
embedded_schema do
field :endpoint, :string, default: "https://api.openai.com/v1/images/generations"
# API key for OpenAI. If not set, will use global api key. Allows for usage
# of a different API key per-call if desired. For instance, allowing a
# customer to provide their own.
field :api_key, :string, redact: true
# Duration in seconds for the response to be received. When streaming a very
# lengthy response, a longer time limit may be required. However, when it
# goes on too long by itself, it tends to hallucinate more.
field :receive_timeout, :integer, default: @receive_timeout
# Defaults to `dall-e-2`. The other model option is `dall-e-3` or `gpt-image-1`.
# The model to use for image generation.
field :model, :string, default: "dall-e-2"
# A text description of the desired image(s). The maximum length is 1000
# characters for `dall-e-2`, 4000 characters for `dall-e-3`
# and 32000 characters for `gpt-image-1`.
field :prompt, :string
# The number of images to generate. Must be between 1 and 10. For dall-e-3,
# only n=1 is supported.
field :n, :integer, default: 1
# The quality of the image that will be generated. `hd` creates images with
# finer details and greater consistency across the image. This param is
# only supported for `dall-e-3`.
# `high`, `medium` and `low` are supported for `gpt-image-1`.
field :quality, :string, default: "standard"
# The format in which the generated images are returned. Must be one of
# `url` or `b64_json`. URLs are only valid for 60 minutes after the image
# has been generated.
field :response_format, :string
# The format in which the generated images are returned.
# This parameter is only supported for `gpt-image-1`.
# Must be one of `png`, `jpeg`, or `webp`.
field :output_format, :string
# The size of the generated images. Must be one of `256x256`, `512x512`, or
# `1024x1024` for `dall-e-2`. Must be one of `1024x1024`, `1792x1024`, or `1024x1792`
# for `dall-e-3` models. Must be one of `1024x1024`, `1536x1024` (landscape),
# `1024x1536` (portrait), or `auto` (default value) for `gpt-image-1`.
field :size, :string, default: "1024x1024"
# The style of the generated images. Must be one of `vivid` or `natural`.
# Vivid causes the model to lean towards generating hyper-real and dramatic
# images. Natural causes the model to produce more natural, less hyper-real
# looking images. This param is only supported for `dall-e-3`.
field :style, :string
# A unique identifier representing your end-user, which can help OpenAI to
# monitor and detect abuse
field :user, :string
# Number of retries on closed-connection errors (stale pool). The initial
field :retry_count, :integer, default: 2
end
@type t :: %OpenAIImage{}
@create_fields [
:endpoint,
:api_key,
:receive_timeout,
:model,
:prompt,
:n,
:quality,
:response_format,
:output_format,
:size,
:style,
:user,
:retry_count
]
@required_fields [:endpoint, :model, :prompt]
@spec get_api_key(t()) :: String.t()
defp get_api_key(%OpenAIImage{api_key: api_key}) do
# if no API key is set default to `""` which will raise a OpenAI API error
api_key || Config.resolve(:openai_key, "")
end
@spec get_org_id() :: String.t() | nil
defp get_org_id() do
Config.resolve(:openai_org_id)
end
@spec get_proj_id() :: String.t() | nil
defp get_proj_id() do
Config.resolve(:openai_proj_id)
end
@doc """
Setup a OpenAIImage client configuration.
"""
@spec new(attrs :: map()) :: {:ok, t} | {:error, Ecto.Changeset.t()}
def new(%{} = attrs) do
%OpenAIImage{}
|> cast(attrs, @create_fields)
|> common_validation()
|> conditional_validations_for_model()
|> apply_action(:insert)
end
@doc """
Setup a OpenAIImage client configuration and return it or raise an error if
invalid.
"""
@spec new!(attrs :: map()) :: t() | no_return()
def new!(attrs) do
case new(attrs) do
{:ok, chain} ->
chain
{:error, changeset} ->
raise LangChainError, changeset
end
end
defp common_validation(changeset) do
changeset
|> validate_required(@required_fields)
|> validate_inclusion(:model, ["dall-e-2", "dall-e-3", "gpt-image-1"])
|> validate_number(:receive_timeout, greater_than_or_equal_to: 0)
end
defp conditional_validations_for_model(changeset) do
case get_field(changeset, :model) do
"dall-e-3" ->
changeset
|> validate_length(:prompt, max: 4_000)
|> validate_number(:n, equal_to: 1)
|> validate_inclusion(:size, ["1024x1024", "1792x1024", "1024x1792"])
|> validate_inclusion(:quality, ["standard", "hd"])
|> validate_inclusion(:style, ["vivid", "natural"])
|> validate_inclusion(:response_format, ["url", "b64_json"])
"dall-e-2" ->
changeset
|> validate_length(:prompt, max: 1_000)
|> validate_number(:n, greater_than_or_equal_to: 1, less_than_or_equal_to: 10)
|> validate_inclusion(:size, ["256x256", "512x512", "1024x1024"])
|> validate_inclusion(:quality, ["standard", "hd"])
|> validate_inclusion(:response_format, ["url", "b64_json"])
"gpt-image-1" ->
changeset
|> validate_length(:prompt, max: 32_000)
|> validate_number(:n, greater_than_or_equal_to: 1, less_than_or_equal_to: 10)
|> validate_inclusion(:size, ["1024x1024", "1536x1024", "1024x1536"])
|> validate_inclusion(:quality, ["high", "medium", "low"])
|> validate_inclusion(:output_format, ["png", "jpeg", "webp"])
|> validate_inclusion(:response_format, ["b64_json"])
_other ->
changeset
end
end
@doc """
Return the params formatted for an API request.
"""
@spec for_api(t) :: %{atom() => any()}
def for_api(%OpenAIImage{} = openai) do
%{
model: openai.model,
prompt: openai.prompt,
n: openai.n,
quality: openai.quality,
size: openai.size
}
|> Utils.conditionally_add_to_map(:response_format, openai.response_format)
|> Utils.conditionally_add_to_map(:style, openai.style)
|> Utils.conditionally_add_to_map(:output_format, openai.output_format)
|> Utils.conditionally_add_to_map(:user, openai.user)
end
@doc """
Calls the OpenAI API passing the OpenAIImage struct with configuration.
When successful, it returns `{:ok, generated_images}` where that is a list of
`LangChain.Images.GeneratedImage` structs.
"""
@spec call(t()) :: {:ok, [GeneratedImage.t()]} | {:error, String.t()}
def call(openai)
def call(%OpenAIImage{} = openai) do
try do
# make base api request and perform high-level success/failure checks
case do_api_request(openai) do
{:error, reason} ->
{:error, reason}
{:ok, parsed_data} ->
{:ok, parsed_data}
end
rescue
err in LangChainError ->
{:error, err.message}
end
end
# Make the API request from the OpenAI server.
#
# The result of the function is:
#
# - `{:ok, %{images: [images], prompt: "the re-written prompt"}}
# - `{:error, reason}` - Where reason is a string explanation of what went
# wrong.
#
# Retries the request up to 3 times on transient errors with a brief delay
@doc false
@spec do_api_request(t(), retry_count :: integer() | nil) ::
{:ok, list()} | {:error, String.t()}
def do_api_request(openai, retry_count \\ nil)
def do_api_request(_openai, 0) do
raise LangChainError, "Retries exceeded. Connection failed."
end
def do_api_request(%OpenAIImage{} = openai, retry_count) do
retry_count = retry_count || openai.retry_count + 1
req =
Req.new(
url: openai.endpoint,
json: for_api(openai),
# required for OpenAI API
auth: {:bearer, get_api_key(openai)},
# required for Azure OpenAI version
headers: [
{"api-key", get_api_key(openai)}
],
receive_timeout: openai.receive_timeout,
# Disable Req-level retry to prevent compounding with LangChain's own
# :closed retry. See https://github.com/brainlid/langchain/issues/503
retry: false
)
req
|> maybe_add_org_id_header()
|> maybe_add_proj_id_header()
|> Req.post()
# parse the body and return it as parsed structs
|> case do
{:ok, %Req.Response{body: data}} ->
case do_process_response(data, openai) do
{:error, reason} ->
{:error, reason}
result ->
result
end
{:error, %Req.TransportError{reason: :timeout}} ->
{:error, "Request timed out"}
{:error, %Req.TransportError{reason: :closed}} ->
# Force a retry by making a recursive call decrementing the counter
Logger.debug(fn -> "Mint connection closed: retry count = #{inspect(retry_count)}" end)
do_api_request(openai, retry_count - 1)
other ->
Logger.warning(fn -> "Unexpected and unhandled API response! #{inspect(other)}" end)
other
end
end
@doc false
@spec do_process_response(data :: %{String.t() => any()} | {:error, any()}, t()) ::
{:ok, [GeneratedImage.t()]} | {:error, String.t()}
def do_process_response(%{"data" => images} = response, %OpenAIImage{} = request)
when is_list(images) do
created_at = DateTime.from_unix!(response["created"])
image_type = request.output_format || "png"
results =
Enum.map(images, fn
%{"b64_json" => base64_raw_content} = image_info when is_binary(base64_raw_content) ->
GeneratedImage.new!(%{
type: :base64,
image_type: image_type,
content: base64_raw_content,
created_at: created_at,
prompt: Map.get(image_info, "revised_prompt", request.prompt),
metadata: %{"model" => request.model, "quality" => request.quality}
})
%{"url" => url} = image_info ->
GeneratedImage.new!(%{
type: :url,
image_type: image_type,
content: url,
created_at: created_at,
prompt: Map.get(image_info, "revised_prompt", request.prompt),
metadata: %{"model" => request.model, "quality" => request.quality}
})
other ->
message = "Unsupported image data response from OpenAI! #{inspect(other)}"
Logger.warning(message)
nil
end)
{:ok, results |> Enum.reject(&is_nil(&1))}
end
def do_process_response(%{"error" => %{"code" => code} = error}, %OpenAIImage{} = _request) do
Logger.warning("Error from OpenAI: #{error["message"]}")
reason =
case code do
"content_policy_violation" = value ->
value
other ->
Logger.warning(fn -> "Unhandled error code from API: #{inspect(other)}" end)
other
end
{:error, reason}
end
defp maybe_add_org_id_header(%Req.Request{} = req) do
org_id = get_org_id()
if org_id do
Req.Request.put_header(req, "OpenAI-Organization", org_id)
else
req
end
end
defp maybe_add_proj_id_header(%Req.Request{} = req) do
proj_id = get_proj_id()
if proj_id do
Req.Request.put_header(req, "OpenAI-Project", proj_id)
else
req
end
end
end