Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 21 additions & 0 deletions ai/application.go
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,10 @@ func (r *Application) Agent(agent contractsai.Agent, options ...contractsai.Opti
return NewConversation(r.ctx, agent, provider, model, middlewares), nil
}

func (r *Application) Image(prompt string, options ...contractsai.Option) contractsai.ImageRequest {
return NewImageRequest(r.ctx, r, prompt, options...)
}

func (r *Application) putFile(ctx context.Context, file contractsai.StorableFile, options ...contractsai.Option) (contractsai.StoredFileResponse, error) {
_, providerName, provider, err := r.resolveProvider(options)
if err != nil {
Expand All @@ -50,6 +54,23 @@ func (r *Application) putFile(ctx context.Context, file contractsai.StorableFile
return fileProvider.PutFile(ctx, file)
}

func (r *Application) image(ctx context.Context, prompt contractsai.ImagePrompt, options ...contractsai.Option) (contractsai.ImageResponse, error) {
opts, providerName, provider, err := r.resolveProvider(options)
if err != nil {
return nil, err
}
if prompt.Model == "" {
prompt.Model = opts.Model
}

imageProvider, ok := provider.(contractsai.ImageProvider)
if !ok {
return nil, errors.AIProviderDoesNotSupportImages.Args(providerName)
}

return imageProvider.Image(ctx, prompt)
}

func (r *Application) resolveProvider(options []contractsai.Option) (*contractsai.Options, string, contractsai.Provider, error) {
opts := &contractsai.Options{}
for _, option := range options {
Expand Down
163 changes: 163 additions & 0 deletions ai/application_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package ai
import (
"context"
"testing"
"time"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
Expand Down Expand Up @@ -335,6 +336,129 @@ func TestApplication_putFile(t *testing.T) {
}
}

func TestApplication_Image(t *testing.T) {
ctx := context.Background()
config := contractsai.Config{
Default: "default",
Providers: map[string]contractsai.ProviderConfig{
"default": {Via: mocksai.NewProvider(t)},
},
}

app := NewApplication(ctx, config)
request := app.Image("draw a cat", WithProvider("default"), WithModel("gpt-image-1"))

req, ok := request.(*imageRequest)
assert.True(t, ok)
assert.Equal(t, ctx, req.ctx)
assert.Equal(t, app, req.app)
assert.Equal(t, "draw a cat", req.prompt)
assert.Equal(t, "default", req.provider)
assert.Equal(t, "gpt-image-1", req.model)

assert.Same(t, req, request.Square())
assert.Same(t, req, request.Portrait())
assert.Same(t, req, request.Landscape())
assert.Same(t, req, request.Quality(contractsai.ImageQualityHigh))
assert.Same(t, req, request.Timeout(2*time.Second))

attachment := ImageFromByte([]byte("image"), WithMimeType("image/png"))
assert.Same(t, req, request.Attachments(attachment))
assert.Equal(t, contractsai.ImageSizeLandscape, req.size)
assert.Equal(t, contractsai.ImageQualityHigh, req.quality)
assert.Equal(t, 2*time.Second, req.timeout)
assert.Equal(t, []contractsai.Attachment{attachment}, req.attachments)
}

func TestImageRequest_Generate(t *testing.T) {
ctx := context.Background()
provider := &applicationImageProviderStub{}
config := contractsai.Config{
Default: "default",
Providers: map[string]contractsai.ProviderConfig{
"default": {Via: provider},
},
}

app := NewApplication(context.Background(), config)
attachment := ImageFromByte([]byte("image"), WithMimeType("image/png"))
response := &applicationImageResponseStub{}
provider.response = response

result, err := app.Image("draw a cat").
Landscape().
Quality(contractsai.ImageQualityHigh).
Attachments(attachment).
Timeout(3 * time.Second).
Generate()

require.NoError(t, err)
assert.Equal(t, response, result)
assert.Equal(t, ctx, provider.ctx)
assert.Equal(t, contractsai.ImagePrompt{
Prompt: "draw a cat",
Size: contractsai.ImageSizeLandscape,
Quality: contractsai.ImageQualityHigh,
Attachments: []contractsai.Attachment{attachment},
Timeout: 3 * time.Second,
}, provider.prompt)
}

func TestApplication_image(t *testing.T) {
tests := []struct {
name string
options []contractsai.Option
setup func() contractsai.Config
expectError error
}{
{
name: "success",
options: []contractsai.Option{WithProvider("openai"), WithModel("gpt-image-override")},
setup: func() contractsai.Config {
provider := &applicationImageProviderStub{}
provider.response = &applicationImageResponseStub{}
return contractsai.Config{
Default: "default",
Providers: map[string]contractsai.ProviderConfig{
"default": {Via: mocksai.NewProvider(t)},
"openai": {Via: provider},
},
}
},
},
{
name: "provider does not support images",
setup: func() contractsai.Config {
return contractsai.Config{
Default: "default",
Providers: map[string]contractsai.ProviderConfig{
"default": {Via: mocksai.NewProvider(t)},
},
}
},
expectError: errors.AIProviderDoesNotSupportImages.Args("default"),
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
app := NewApplication(context.Background(), tt.setup())
response, err := app.image(context.Background(), contractsai.ImagePrompt{Prompt: "draw a cat"}, tt.options...)
assert.Equal(t, tt.expectError, err)
if tt.expectError != nil {
assert.Nil(t, response)
return
}

require.NotNil(t, response)
provider, ok := app.config.Providers["openai"].Via.(*applicationImageProviderStub)
if ok {
assert.Equal(t, "gpt-image-override", provider.prompt.Model)
}
})
}
}

type applicationTestMiddleware struct{}

type uploadTestProvider struct {
Expand All @@ -353,6 +477,45 @@ func (p uploadTestProvider) PutFile(ctx context.Context, file contractsai.Storab
return p.fileProvider.PutFile(ctx, file)
}

type applicationImageProviderStub struct {
ctx context.Context
prompt contractsai.ImagePrompt
response contractsai.ImageResponse
err error
}

func (p *applicationImageProviderStub) Prompt(context.Context, contractsai.AgentPrompt) (contractsai.Response, error) {
return nil, nil
}

func (p *applicationImageProviderStub) Stream(context.Context, contractsai.AgentPrompt) (contractsai.StreamableResponse, error) {
return nil, nil
}

func (p *applicationImageProviderStub) Image(ctx context.Context, prompt contractsai.ImagePrompt) (contractsai.ImageResponse, error) {
p.ctx = ctx
p.prompt = prompt
return p.response, p.err
}

type applicationImageResponseStub struct{}

func (r *applicationImageResponseStub) Content(context.Context) ([]byte, error) {
return []byte("image"), nil
}

func (r *applicationImageResponseStub) MimeType() string { return "image/png" }

func (r *applicationImageResponseStub) Usage() contractsai.Usage { return nil }

func (r *applicationImageResponseStub) Then(callback func(contractsai.ImageResponse)) contractsai.ImageResponse {
if callback != nil {
callback(r)
}

return r
}

func (m *applicationTestMiddleware) Handle(ctx context.Context, prompt contractsai.AgentPrompt, next contractsai.Next) (contractsai.Response, error) {
response, err := next(ctx, prompt)
if err != nil {
Expand Down
16 changes: 16 additions & 0 deletions ai/image/image.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,16 @@
package image

import contractsai "github.com/goravel/framework/contracts/ai"

type Quality = contractsai.ImageQuality
type Size = contractsai.ImageSize

const (
QualityLow = contractsai.ImageQualityLow
QualityMedium = contractsai.ImageQualityMedium
QualityHigh = contractsai.ImageQualityHigh

SizeSquare = contractsai.ImageSizeSquare
SizePortrait = contractsai.ImageSizePortrait
SizeLandscape = contractsai.ImageSizeLandscape
)
98 changes: 98 additions & 0 deletions ai/image_request.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,98 @@
package ai

import (
"context"
"time"

contractsai "github.com/goravel/framework/contracts/ai"
)

type imageRequest struct {
ctx context.Context
app *Application
prompt string
provider string
model string
size contractsai.ImageSize
quality contractsai.ImageQuality
attachments []contractsai.Attachment
timeout time.Duration
}

func NewImageRequest(ctx context.Context, app *Application, prompt string, options ...contractsai.Option) contractsai.ImageRequest {
resolvedOptions := &contractsai.Options{}
for _, option := range options {
if option == nil {
continue
}

option(resolvedOptions)
}

return &imageRequest{
ctx: ctx,
app: app,
prompt: prompt,
provider: resolvedOptions.Provider,
model: resolvedOptions.Model,
}
}

func (r *imageRequest) Model(model string) contractsai.ImageRequest {
r.model = model
return r
}

func (r *imageRequest) Provider(provider string) contractsai.ImageRequest {
r.provider = provider
return r
}

func (r *imageRequest) Square() contractsai.ImageRequest {
r.size = contractsai.ImageSizeSquare
return r
}

func (r *imageRequest) Portrait() contractsai.ImageRequest {
r.size = contractsai.ImageSizePortrait
return r
}

func (r *imageRequest) Landscape() contractsai.ImageRequest {
r.size = contractsai.ImageSizeLandscape
return r
}

func (r *imageRequest) Quality(quality contractsai.ImageQuality) contractsai.ImageRequest {
r.quality = quality
return r
}

func (r *imageRequest) Attachments(attachments ...contractsai.Attachment) contractsai.ImageRequest {
r.attachments = append(r.attachments, filterNilAttachments(attachments)...)
return r
}

func (r *imageRequest) Timeout(timeout time.Duration) contractsai.ImageRequest {
r.timeout = timeout
return r
}

func (r *imageRequest) Generate() (contractsai.ImageResponse, error) {
options := make([]contractsai.Option, 0, 2)
if r.provider != "" {
options = append(options, WithProvider(r.provider))
}
if r.model != "" {
options = append(options, WithModel(r.model))
}
Comment thread
hwbrzzl marked this conversation as resolved.

return r.app.image(r.ctx, contractsai.ImagePrompt{
Prompt: r.prompt,
Model: r.model,
Size: r.size,
Quality: r.quality,
Attachments: filterNilAttachments(r.attachments),
Timeout: r.timeout,
}, options...)
}
Loading
Loading