Update Go module dependencies to include AWS SDK v2
- Added several AWS SDK v2 modules to go.mod and go.sum, enhancing cloud service integration capabilities: - github.com/aws/aws-sdk-go-v2 and its related services (S3, STS, SSO, etc.) for improved AWS functionality. - Updated indirect dependencies to ensure compatibility and performance improvements. - This update supports future development involving AWS services, streamlining the integration process.
This commit is contained in:
parent
af30bec4f0
commit
93a6d7a6a0
12 changed files with 1382 additions and 0 deletions
18
go.mod
18
go.mod
|
|
@ -4,6 +4,10 @@ go 1.23
|
||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/PuerkitoBio/goquery v1.10.1
|
github.com/PuerkitoBio/goquery v1.10.1
|
||||||
|
github.com/aws/aws-sdk-go-v2 v1.32.7
|
||||||
|
github.com/aws/aws-sdk-go-v2/config v1.28.7
|
||||||
|
github.com/aws/aws-sdk-go-v2/credentials v1.17.48
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/s3 v1.71.1
|
||||||
github.com/blang/semver v3.5.1+incompatible
|
github.com/blang/semver v3.5.1+incompatible
|
||||||
github.com/caarlos0/env/v6 v6.10.1
|
github.com/caarlos0/env/v6 v6.10.1
|
||||||
github.com/dchest/captcha v1.1.0
|
github.com/dchest/captcha v1.1.0
|
||||||
|
|
@ -38,6 +42,20 @@ require (
|
||||||
filippo.io/edwards25519 v1.1.0 // indirect
|
filippo.io/edwards25519 v1.1.0 // indirect
|
||||||
github.com/TylerBrock/colorjson v0.0.0-20200706003622-8a50f05110d2 // indirect
|
github.com/TylerBrock/colorjson v0.0.0-20200706003622-8a50f05110d2 // indirect
|
||||||
github.com/andybalholm/cascadia v1.3.3 // indirect
|
github.com/andybalholm/cascadia v1.3.3 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.7 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.16.22 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.26 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.26 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.1 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/v4a v1.3.26 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.12.1 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.4.7 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.12.7 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.18.7 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/sso v1.24.8 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.28.7 // indirect
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/sts v1.33.3 // indirect
|
||||||
|
github.com/aws/smithy-go v1.22.1 // indirect
|
||||||
github.com/blang/semver/v4 v4.0.0 // indirect
|
github.com/blang/semver/v4 v4.0.0 // indirect
|
||||||
github.com/bytedance/sonic v1.12.6 // indirect
|
github.com/bytedance/sonic v1.12.6 // indirect
|
||||||
github.com/bytedance/sonic/loader v0.2.1 // indirect
|
github.com/bytedance/sonic/loader v0.2.1 // indirect
|
||||||
|
|
|
||||||
36
go.sum
36
go.sum
|
|
@ -6,6 +6,42 @@ github.com/TylerBrock/colorjson v0.0.0-20200706003622-8a50f05110d2 h1:ZBbLwSJqkH
|
||||||
github.com/TylerBrock/colorjson v0.0.0-20200706003622-8a50f05110d2/go.mod h1:VSw57q4QFiWDbRnjdX8Cb3Ow0SFncRw+bA/ofY6Q83w=
|
github.com/TylerBrock/colorjson v0.0.0-20200706003622-8a50f05110d2/go.mod h1:VSw57q4QFiWDbRnjdX8Cb3Ow0SFncRw+bA/ofY6Q83w=
|
||||||
github.com/andybalholm/cascadia v1.3.3 h1:AG2YHrzJIm4BZ19iwJ/DAua6Btl3IwJX+VI4kktS1LM=
|
github.com/andybalholm/cascadia v1.3.3 h1:AG2YHrzJIm4BZ19iwJ/DAua6Btl3IwJX+VI4kktS1LM=
|
||||||
github.com/andybalholm/cascadia v1.3.3/go.mod h1:xNd9bqTn98Ln4DwST8/nG+H0yuB8Hmgu1YHNnWw0GeA=
|
github.com/andybalholm/cascadia v1.3.3/go.mod h1:xNd9bqTn98Ln4DwST8/nG+H0yuB8Hmgu1YHNnWw0GeA=
|
||||||
|
github.com/aws/aws-sdk-go-v2 v1.32.7 h1:ky5o35oENWi0JYWUZkB7WYvVPP+bcRF5/Iq7JWSb5Rw=
|
||||||
|
github.com/aws/aws-sdk-go-v2 v1.32.7/go.mod h1:P5WJBrYqqbWVaOxgH0X/FYYD47/nooaPOZPlQdmiN2U=
|
||||||
|
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.7 h1:lL7IfaFzngfx0ZwUGOZdsFFnQ5uLvR0hWqqhyE7Q9M8=
|
||||||
|
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.7/go.mod h1:QraP0UcVlQJsmHfioCrveWOC1nbiWUl3ej08h4mXWoc=
|
||||||
|
github.com/aws/aws-sdk-go-v2/config v1.28.7 h1:GduUnoTXlhkgnxTD93g1nv4tVPILbdNQOzav+Wpg7AE=
|
||||||
|
github.com/aws/aws-sdk-go-v2/config v1.28.7/go.mod h1:vZGX6GVkIE8uECSUHB6MWAUsd4ZcG2Yq/dMa4refR3M=
|
||||||
|
github.com/aws/aws-sdk-go-v2/credentials v1.17.48 h1:IYdLD1qTJ0zanRavulofmqut4afs45mOWEI+MzZtTfQ=
|
||||||
|
github.com/aws/aws-sdk-go-v2/credentials v1.17.48/go.mod h1:tOscxHN3CGmuX9idQ3+qbkzrjVIx32lqDSU1/0d/qXs=
|
||||||
|
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.16.22 h1:kqOrpojG71DxJm/KDPO+Z/y1phm1JlC8/iT+5XRmAn8=
|
||||||
|
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.16.22/go.mod h1:NtSFajXVVL8TA2QNngagVZmUtXciyrHOt7xgz4faS/M=
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.26 h1:I/5wmGMffY4happ8NOCuIUEWGUvvFp5NSeQcXl9RHcI=
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.26/go.mod h1:FR8f4turZtNy6baO0KJ5FJUmXH/cSkI9fOngs0yl6mA=
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.26 h1:zXFLuEuMMUOvEARXFUVJdfqZ4bvvSgdGRq/ATcrQxzM=
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.26/go.mod h1:3o2Wpy0bogG1kyOPrgkXA8pgIfEEv0+m19O9D5+W8y8=
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.1 h1:VaRN3TlFdd6KxX1x3ILT5ynH6HvKgqdiXoTxAF4HQcQ=
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.1/go.mod h1:FbtygfRFze9usAadmnGJNc8KsP346kEe+y2/oyhGAGc=
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/v4a v1.3.26 h1:GeNJsIFHB+WW5ap2Tec4K6dzcVTsRbsT1Lra46Hv9ME=
|
||||||
|
github.com/aws/aws-sdk-go-v2/internal/v4a v1.3.26/go.mod h1:zfgMpwHDXX2WGoG84xG2H+ZlPTkJUU4YUvx2svLQYWo=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.12.1 h1:iXtILhvDxB6kPvEXgsDhGaZCSC6LQET5ZHSdJozeI0Y=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.12.1/go.mod h1:9nu0fVANtYiAePIBh2/pFUSwtJ402hLnp854CNoDOeE=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.4.7 h1:tB4tNw83KcajNAzaIMhkhVI2Nt8fAZd5A5ro113FEMY=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.4.7/go.mod h1:lvpyBGkZ3tZ9iSsUIcC2EWp+0ywa7aK3BLT+FwZi+mQ=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.12.7 h1:8eUsivBQzZHqe/3FE+cqwfH+0p5Jo8PFM/QYQSmeZ+M=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.12.7/go.mod h1:kLPQvGUmxn/fqiCrDeohwG33bq2pQpGeY62yRO6Nrh0=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.18.7 h1:Hi0KGbrnr57bEHWM0bJ1QcBzxLrL/k2DHvGYhb8+W1w=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.18.7/go.mod h1:wKNgWgExdjjrm4qvfbTorkvocEstaoDl4WCvGfeCy9c=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/s3 v1.71.1 h1:aOVVZJgWbaH+EJYPvEgkNhCEbXXvH7+oML36oaPK3zE=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/s3 v1.71.1/go.mod h1:r+xl5yzMk9083rMR+sJ5TYj9Tihvf/l1oxzZXDgGj2Q=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/sso v1.24.8 h1:CvuUmnXI7ebaUAhbJcDy9YQx8wHR69eZ9I7q5hszt/g=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/sso v1.24.8/go.mod h1:XDeGv1opzwm8ubxddF0cgqkZWsyOtw4lr6dxwmb6YQg=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.28.7 h1:F2rBfNAL5UyswqoeWv9zs74N/NanhK16ydHW1pahX6E=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.28.7/go.mod h1:JfyQ0g2JG8+Krq0EuZNnRwX0mU0HrwY/tG6JNfcqh4k=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/sts v1.33.3 h1:Xgv/hyNgvLda/M9l9qxXc4UFSgppnRczLxlMs5Ae/QY=
|
||||||
|
github.com/aws/aws-sdk-go-v2/service/sts v1.33.3/go.mod h1:5Gn+d+VaaRgsjewpMvGazt0WfcFO+Md4wLOuBfGR9Bc=
|
||||||
|
github.com/aws/smithy-go v1.22.1 h1:/HPHZQ0g7f4eUeK6HKglFz8uwVfZKgoI25rb/J+dnro=
|
||||||
|
github.com/aws/smithy-go v1.22.1/go.mod h1:irrKGvNn1InZwb2d7fkIRNucdfwR8R+Ts3wxYa/cJHg=
|
||||||
github.com/blang/semver v3.5.1+incompatible h1:cQNTCjp13qL8KC3Nbxr/y2Bqb63oX6wdnnjpJbkM4JQ=
|
github.com/blang/semver v3.5.1+incompatible h1:cQNTCjp13qL8KC3Nbxr/y2Bqb63oX6wdnnjpJbkM4JQ=
|
||||||
github.com/blang/semver v3.5.1+incompatible/go.mod h1:kRBLl5iJ+tD4TcOOxsy/0fnwebNt5EWlYSAyrTnjyyk=
|
github.com/blang/semver v3.5.1+incompatible/go.mod h1:kRBLl5iJ+tD4TcOOxsy/0fnwebNt5EWlYSAyrTnjyyk=
|
||||||
github.com/blang/semver/v4 v4.0.0 h1:1PFHFE6yCCTv8C1TeyNNarDzntLi7wMI5i/pzqYIsAM=
|
github.com/blang/semver/v4 v4.0.0 h1:1PFHFE6yCCTv8C1TeyNNarDzntLi7wMI5i/pzqYIsAM=
|
||||||
|
|
|
||||||
115
neo/vision/driver/local/storage.go
Normal file
115
neo/vision/driver/local/storage.go
Normal file
|
|
@ -0,0 +1,115 @@
|
||||||
|
package local
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/yaoapp/gou/fs"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Storage the local storage driver
|
||||||
|
type Storage struct {
|
||||||
|
Path string `json:"path" yaml:"path"`
|
||||||
|
Compression bool `json:"compression" yaml:"compression"`
|
||||||
|
BaseURL string `json:"base_url" yaml:"base_url"`
|
||||||
|
PreviewURL func(fileID string) string `json:"-" yaml:"-"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// New create a new local storage
|
||||||
|
func New(options map[string]interface{}) (*Storage, error) {
|
||||||
|
storage := &Storage{
|
||||||
|
Compression: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
if path, ok := options["path"].(string); ok {
|
||||||
|
storage.Path = path
|
||||||
|
}
|
||||||
|
|
||||||
|
if compression, ok := options["compression"].(bool); ok {
|
||||||
|
storage.Compression = compression
|
||||||
|
}
|
||||||
|
|
||||||
|
if baseURL, ok := options["base_url"].(string); ok {
|
||||||
|
storage.BaseURL = baseURL
|
||||||
|
}
|
||||||
|
|
||||||
|
if previewURL, ok := options["preview_url"].(func(string) string); ok {
|
||||||
|
storage.PreviewURL = previewURL
|
||||||
|
}
|
||||||
|
|
||||||
|
if storage.Path == "" {
|
||||||
|
return nil, fmt.Errorf("path is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
return storage, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Upload upload file to local storage
|
||||||
|
func (storage *Storage) Upload(ctx context.Context, filename string, reader io.Reader, contentType string) (string, error) {
|
||||||
|
data, err := fs.Get("data")
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
ext := filepath.Ext(filename)
|
||||||
|
id := storage.makeID(filename, ext)
|
||||||
|
path := filepath.Join(storage.Path, id)
|
||||||
|
|
||||||
|
// Create directory if not exists
|
||||||
|
dir := filepath.Dir(path)
|
||||||
|
if err := data.MkdirAll(dir, 0755); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write file
|
||||||
|
_, err = data.Write(path, reader, 0644)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
return id, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Download download file from local storage
|
||||||
|
func (storage *Storage) Download(ctx context.Context, fileID string) (io.ReadCloser, string, error) {
|
||||||
|
data, err := fs.Get("data")
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
path := filepath.Join(storage.Path, fileID)
|
||||||
|
reader, err := data.ReadCloser(path)
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
contentType := "application/octet-stream"
|
||||||
|
if v, err := data.MimeType(path); err == nil {
|
||||||
|
contentType = v
|
||||||
|
}
|
||||||
|
|
||||||
|
return reader, contentType, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// URL get file url
|
||||||
|
func (storage *Storage) URL(ctx context.Context, fileID string) string {
|
||||||
|
if storage.PreviewURL != nil {
|
||||||
|
return storage.PreviewURL(fileID)
|
||||||
|
}
|
||||||
|
if storage.BaseURL != "" {
|
||||||
|
return fmt.Sprintf("%s/%s", strings.TrimRight(storage.BaseURL, "/"), fileID)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%s/%s", storage.Path, fileID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (storage *Storage) makeID(filename string, ext string) string {
|
||||||
|
date := time.Now().Format("20060102")
|
||||||
|
hash := fmt.Sprintf("%x", sha256.Sum256([]byte(filename)))[:8]
|
||||||
|
name := strings.TrimSuffix(filepath.Base(filename), ext)
|
||||||
|
return fmt.Sprintf("%s/%s-%s%s", date, name, hash, ext)
|
||||||
|
}
|
||||||
74
neo/vision/driver/local/storage_test.go
Normal file
74
neo/vision/driver/local/storage_test.go
Normal file
|
|
@ -0,0 +1,74 @@
|
||||||
|
package local
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/test"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestLocalStorage(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
|
||||||
|
t.Run("Create Storage", func(t *testing.T) {
|
||||||
|
storage, err := New(map[string]interface{}{
|
||||||
|
"path": "/__vision_test",
|
||||||
|
"compression": true,
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, storage)
|
||||||
|
assert.Equal(t, "/__vision_test", storage.Path)
|
||||||
|
assert.True(t, storage.Compression)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Upload and Download", func(t *testing.T) {
|
||||||
|
storage, err := New(map[string]interface{}{
|
||||||
|
"path": "/__vision_test",
|
||||||
|
"compression": true,
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
content := []byte("test content")
|
||||||
|
reader := bytes.NewReader(content)
|
||||||
|
fileID, err := storage.Upload(context.Background(), "test.txt", reader, "text/plain")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotEmpty(t, fileID)
|
||||||
|
|
||||||
|
// Download
|
||||||
|
reader2, contentType, err := storage.Download(context.Background(), fileID)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Contains(t, contentType, "text/plain")
|
||||||
|
|
||||||
|
downloaded, err := io.ReadAll(reader2)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, content, downloaded)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("URL Generation", func(t *testing.T) {
|
||||||
|
storage, err := New(map[string]interface{}{
|
||||||
|
"path": "/__vision_test",
|
||||||
|
"compression": true,
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
fileID := "20240101/test-12345678.txt"
|
||||||
|
url := storage.URL(context.Background(), fileID)
|
||||||
|
assert.Equal(t, "/__vision_test/20240101/test-12345678.txt", url)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Download Non-existent File", func(t *testing.T) {
|
||||||
|
storage, err := New(map[string]interface{}{
|
||||||
|
"path": "/__vision_test",
|
||||||
|
"compression": true,
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
_, _, err = storage.Download(context.Background(), "non-existent.txt")
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
}
|
||||||
184
neo/vision/driver/openai/model.go
Normal file
184
neo/vision/driver/openai/model.go
Normal file
|
|
@ -0,0 +1,184 @@
|
||||||
|
package openai
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/yaoapp/gou/fs"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Model the OpenAI vision model
|
||||||
|
type Model struct {
|
||||||
|
APIKey string `json:"api_key" yaml:"api_key"`
|
||||||
|
Model string `json:"model" yaml:"model"`
|
||||||
|
Compression bool `json:"compression" yaml:"compression"`
|
||||||
|
Prompt string `json:"prompt" yaml:"prompt"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// New create a new OpenAI vision model
|
||||||
|
func New(options map[string]interface{}) (*Model, error) {
|
||||||
|
model := &Model{
|
||||||
|
Model: "gpt-4-vision-preview",
|
||||||
|
Compression: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
if apiKey, ok := options["api_key"].(string); ok {
|
||||||
|
model.APIKey = apiKey
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelName, ok := options["model"].(string); ok {
|
||||||
|
model.Model = modelName
|
||||||
|
}
|
||||||
|
|
||||||
|
if compression, ok := options["compression"].(bool); ok {
|
||||||
|
model.Compression = compression
|
||||||
|
}
|
||||||
|
|
||||||
|
if prompt, ok := options["prompt"].(string); ok {
|
||||||
|
model.Prompt = prompt
|
||||||
|
}
|
||||||
|
|
||||||
|
if model.APIKey == "" {
|
||||||
|
return nil, fmt.Errorf("api_key is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
return model, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Analyze analyze image using OpenAI vision model
|
||||||
|
func (model *Model) Analyze(ctx context.Context, fileID string, prompt string) (map[string]interface{}, error) {
|
||||||
|
if model.APIKey == "" {
|
||||||
|
return nil, fmt.Errorf("api_key is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if fileID is a URL or base64 data
|
||||||
|
var imageURL string
|
||||||
|
if strings.HasPrefix(fileID, "data:image/") {
|
||||||
|
// Already a base64 data URL
|
||||||
|
imageURL = fileID
|
||||||
|
} else if strings.HasPrefix(fileID, "http://") || strings.HasPrefix(fileID, "https://") {
|
||||||
|
// Already a URL
|
||||||
|
imageURL = fileID
|
||||||
|
} else {
|
||||||
|
// Try to read the file and convert to base64
|
||||||
|
data, err := fs.Get("data")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get data fs: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
reader, err := data.ReadCloser(fileID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to read file: %w", err)
|
||||||
|
}
|
||||||
|
defer reader.Close()
|
||||||
|
|
||||||
|
content, err := io.ReadAll(reader)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to read content: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get content type
|
||||||
|
contentType := "image/png" // default
|
||||||
|
if v, err := data.MimeType(fileID); err == nil {
|
||||||
|
contentType = v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert to base64
|
||||||
|
base64Data := base64.StdEncoding.EncodeToString(content)
|
||||||
|
imageURL = fmt.Sprintf("data:%s;base64,%s", contentType, base64Data)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prepare the request body
|
||||||
|
reqBody := map[string]interface{}{
|
||||||
|
"model": model.Model,
|
||||||
|
"messages": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"type": "text",
|
||||||
|
"text": prompt,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": map[string]interface{}{
|
||||||
|
"url": imageURL,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"max_tokens": 1000,
|
||||||
|
}
|
||||||
|
|
||||||
|
jsonBody, err := json.Marshal(reqBody)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal request body: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create request
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "POST", "https://api.openai.com/v1/chat/completions", bytes.NewBuffer(jsonBody))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", model.APIKey))
|
||||||
|
|
||||||
|
// Send request
|
||||||
|
client := &http.Client{}
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to send request: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
// Read response
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return nil, fmt.Errorf("OpenAI API error: %s", string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse response
|
||||||
|
var result map[string]interface{}
|
||||||
|
if err := json.Unmarshal(body, &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract content
|
||||||
|
choices, ok := result["choices"].([]interface{})
|
||||||
|
if !ok || len(choices) == 0 {
|
||||||
|
return nil, fmt.Errorf("invalid response format")
|
||||||
|
}
|
||||||
|
|
||||||
|
message, ok := choices[0].(map[string]interface{})["message"].(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("invalid response format")
|
||||||
|
}
|
||||||
|
|
||||||
|
content, ok := message["content"].(string)
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("invalid response format")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try to parse content as JSON
|
||||||
|
var description map[string]interface{}
|
||||||
|
if err := json.Unmarshal([]byte(content), &description); err != nil {
|
||||||
|
// If not JSON, use the content as description
|
||||||
|
description = map[string]interface{}{
|
||||||
|
"description": content,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return description, nil
|
||||||
|
}
|
||||||
149
neo/vision/driver/openai/model_test.go
Normal file
149
neo/vision/driver/openai/model_test.go
Normal file
|
|
@ -0,0 +1,149 @@
|
||||||
|
package openai
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/gou/fs"
|
||||||
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/neo/vision/driver/s3"
|
||||||
|
"github.com/yaoapp/yao/test"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
// 1x1 transparent PNG
|
||||||
|
testImageBase64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestOpenAIModel(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
|
||||||
|
t.Run("Create Model", func(t *testing.T) {
|
||||||
|
model, err := New(map[string]interface{}{
|
||||||
|
"api_key": os.Getenv("OPENAI_API_KEY"),
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, model)
|
||||||
|
if model != nil {
|
||||||
|
assert.Equal(t, os.Getenv("OPENAI_API_KEY"), model.APIKey)
|
||||||
|
assert.Equal(t, os.Getenv("VISION_MODEL"), model.Model)
|
||||||
|
assert.True(t, model.Compression)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Create Model with Invalid API Key", func(t *testing.T) {
|
||||||
|
_, err := New(map[string]interface{}{})
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "api_key is required")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Analyze with Base64 Image", func(t *testing.T) {
|
||||||
|
model, err := New(map[string]interface{}{
|
||||||
|
"api_key": os.Getenv("OPENAI_API_KEY"),
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Use base64 image data
|
||||||
|
result, err := model.Analyze(context.Background(), "data:image/png;base64,"+testImageBase64, "Describe this image in detail")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.NotEmpty(t, result["description"])
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Analyze with URL", func(t *testing.T) {
|
||||||
|
if os.Getenv("S3_API") == "" || os.Getenv("S3_ACCESS_KEY") == "" ||
|
||||||
|
os.Getenv("S3_SECRET_KEY") == "" || os.Getenv("S3_BUCKET") == "" {
|
||||||
|
t.Skip("S3 environment variables not set")
|
||||||
|
}
|
||||||
|
|
||||||
|
model, err := New(map[string]interface{}{
|
||||||
|
"api_key": os.Getenv("OPENAI_API_KEY"),
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Create S3 client and upload test image
|
||||||
|
s3Client, err := s3.New(map[string]interface{}{
|
||||||
|
"endpoint": os.Getenv("S3_API"),
|
||||||
|
"region": "auto",
|
||||||
|
"key": os.Getenv("S3_ACCESS_KEY"),
|
||||||
|
"secret": os.Getenv("S3_SECRET_KEY"),
|
||||||
|
"bucket": os.Getenv("S3_BUCKET"),
|
||||||
|
"prefix": "vision-test",
|
||||||
|
"expiration": "5m",
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Upload test image
|
||||||
|
imgData, err := base64.StdEncoding.DecodeString(testImageBase64)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
reader := bytes.NewReader(imgData)
|
||||||
|
fileID, err := s3Client.Upload(context.Background(), "test.png", reader, "image/png")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Get URL from S3
|
||||||
|
url := s3Client.URL(context.Background(), fileID)
|
||||||
|
assert.NotEmpty(t, url)
|
||||||
|
|
||||||
|
// Use S3 URL for analysis
|
||||||
|
result, err := model.Analyze(context.Background(), url, "Describe this image in detail")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.NotEmpty(t, result["description"])
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Analyze with File ID", func(t *testing.T) {
|
||||||
|
model, err := New(map[string]interface{}{
|
||||||
|
"api_key": os.Getenv("OPENAI_API_KEY"),
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Create test file
|
||||||
|
data, err := fs.Get("data")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Write test image data
|
||||||
|
imgData, err := base64.StdEncoding.DecodeString(testImageBase64)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
_, err = data.WriteFile("/__vision_test/test.png", imgData, 0644)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Analyze using file ID
|
||||||
|
result, err := model.Analyze(context.Background(), "/__vision_test/test.png", "Describe this image in detail")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.NotEmpty(t, result["description"])
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Analyze with Invalid File ID", func(t *testing.T) {
|
||||||
|
model, err := New(map[string]interface{}{
|
||||||
|
"api_key": os.Getenv("OPENAI_API_KEY"),
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
_, err = model.Analyze(context.Background(), "/non-existent.png", "Describe this image in detail")
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "failed to read file")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Analyze with Invalid API Key", func(t *testing.T) {
|
||||||
|
model, err := New(map[string]interface{}{
|
||||||
|
"api_key": "invalid-key",
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
_, err = model.Analyze(context.Background(), "data:image/png;base64,"+testImageBase64, "Describe this image in detail")
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "OpenAI API error")
|
||||||
|
})
|
||||||
|
}
|
||||||
168
neo/vision/driver/s3/storage.go
Normal file
168
neo/vision/driver/s3/storage.go
Normal file
|
|
@ -0,0 +1,168 @@
|
||||||
|
package s3
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/aws/aws-sdk-go-v2/aws"
|
||||||
|
"github.com/aws/aws-sdk-go-v2/credentials"
|
||||||
|
"github.com/aws/aws-sdk-go-v2/service/s3"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DefaultExpiration default expiration time for presigned URLs (5 minutes)
|
||||||
|
const DefaultExpiration = 5 * time.Minute
|
||||||
|
|
||||||
|
// Storage the S3 storage driver
|
||||||
|
type Storage struct {
|
||||||
|
Endpoint string `json:"endpoint" yaml:"endpoint"`
|
||||||
|
Region string `json:"region" yaml:"region"`
|
||||||
|
Key string `json:"key" yaml:"key"`
|
||||||
|
Secret string `json:"secret" yaml:"secret"`
|
||||||
|
Bucket string `json:"bucket" yaml:"bucket"`
|
||||||
|
Expiration time.Duration `json:"expiration" yaml:"expiration"`
|
||||||
|
client *s3.Client
|
||||||
|
prefix string
|
||||||
|
}
|
||||||
|
|
||||||
|
// New create a new S3 storage
|
||||||
|
func New(options map[string]interface{}) (*Storage, error) {
|
||||||
|
storage := &Storage{
|
||||||
|
Region: "auto",
|
||||||
|
Expiration: DefaultExpiration,
|
||||||
|
}
|
||||||
|
|
||||||
|
if endpoint, ok := options["endpoint"].(string); ok {
|
||||||
|
storage.Endpoint = endpoint
|
||||||
|
}
|
||||||
|
|
||||||
|
if region, ok := options["region"].(string); ok {
|
||||||
|
storage.Region = region
|
||||||
|
}
|
||||||
|
|
||||||
|
if key, ok := options["key"].(string); ok {
|
||||||
|
storage.Key = key
|
||||||
|
}
|
||||||
|
|
||||||
|
if secret, ok := options["secret"].(string); ok {
|
||||||
|
storage.Secret = secret
|
||||||
|
}
|
||||||
|
|
||||||
|
if bucket, ok := options["bucket"].(string); ok {
|
||||||
|
storage.Bucket = bucket
|
||||||
|
}
|
||||||
|
|
||||||
|
if prefix, ok := options["prefix"].(string); ok {
|
||||||
|
storage.prefix = prefix
|
||||||
|
}
|
||||||
|
|
||||||
|
if exp, ok := options["expiration"].(time.Duration); ok {
|
||||||
|
storage.Expiration = exp
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate required fields
|
||||||
|
if storage.Key == "" || storage.Secret == "" {
|
||||||
|
return nil, fmt.Errorf("key and secret are required")
|
||||||
|
}
|
||||||
|
|
||||||
|
if storage.Bucket == "" {
|
||||||
|
return nil, fmt.Errorf("bucket is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create S3 client
|
||||||
|
opts := s3.Options{
|
||||||
|
Region: storage.Region,
|
||||||
|
Credentials: credentials.NewStaticCredentialsProvider(storage.Key, storage.Secret, ""),
|
||||||
|
UsePathStyle: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
if storage.Endpoint != "" {
|
||||||
|
// Remove bucket name from endpoint if present
|
||||||
|
endpoint := storage.Endpoint
|
||||||
|
if strings.Contains(endpoint, "/"+storage.Bucket) {
|
||||||
|
endpoint = strings.TrimSuffix(endpoint, "/"+storage.Bucket)
|
||||||
|
}
|
||||||
|
opts.BaseEndpoint = aws.String(endpoint)
|
||||||
|
}
|
||||||
|
|
||||||
|
storage.client = s3.New(opts)
|
||||||
|
return storage, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Upload upload file to S3
|
||||||
|
func (storage *Storage) Upload(ctx context.Context, filename string, reader io.Reader, contentType string) (string, error) {
|
||||||
|
if storage.client == nil {
|
||||||
|
return "", fmt.Errorf("s3 client not initialized")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generate file ID
|
||||||
|
fileID := storage.makeID(filename, filepath.Ext(filename))
|
||||||
|
key := filepath.Join(storage.prefix, fileID)
|
||||||
|
|
||||||
|
// Upload file
|
||||||
|
_, err := storage.client.PutObject(ctx, &s3.PutObjectInput{
|
||||||
|
Bucket: aws.String(storage.Bucket),
|
||||||
|
Key: aws.String(key),
|
||||||
|
Body: reader,
|
||||||
|
ContentType: aws.String(contentType),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to upload file: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return fileID, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Download download file from S3
|
||||||
|
func (storage *Storage) Download(ctx context.Context, fileID string) (io.ReadCloser, string, error) {
|
||||||
|
if storage.client == nil {
|
||||||
|
return nil, "", fmt.Errorf("s3 client not initialized")
|
||||||
|
}
|
||||||
|
|
||||||
|
key := filepath.Join(storage.prefix, fileID)
|
||||||
|
|
||||||
|
// Get object
|
||||||
|
result, err := storage.client.GetObject(ctx, &s3.GetObjectInput{
|
||||||
|
Bucket: aws.String(storage.Bucket),
|
||||||
|
Key: aws.String(key),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", fmt.Errorf("failed to download file: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
contentType := "application/octet-stream"
|
||||||
|
if result.ContentType != nil {
|
||||||
|
contentType = *result.ContentType
|
||||||
|
}
|
||||||
|
|
||||||
|
return result.Body, contentType, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// URL get file url with expiration
|
||||||
|
func (storage *Storage) URL(ctx context.Context, fileID string) string {
|
||||||
|
if storage.client == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
key := filepath.Join(storage.prefix, fileID)
|
||||||
|
presignClient := s3.NewPresignClient(storage.client)
|
||||||
|
request, err := presignClient.PresignGetObject(ctx, &s3.GetObjectInput{
|
||||||
|
Bucket: aws.String(storage.Bucket),
|
||||||
|
Key: aws.String(key),
|
||||||
|
}, s3.WithPresignExpires(storage.Expiration))
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
return request.URL
|
||||||
|
}
|
||||||
|
|
||||||
|
func (storage *Storage) makeID(filename string, ext string) string {
|
||||||
|
date := time.Now().Format("20060102")
|
||||||
|
name := strings.TrimSuffix(filepath.Base(filename), ext)
|
||||||
|
return fmt.Sprintf("%s/%s-%d%s", date, name, time.Now().UnixNano(), ext)
|
||||||
|
}
|
||||||
143
neo/vision/driver/s3/storage_test.go
Normal file
143
neo/vision/driver/s3/storage_test.go
Normal file
|
|
@ -0,0 +1,143 @@
|
||||||
|
package s3
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/test"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestS3Storage(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
|
||||||
|
t.Run("Create Storage", func(t *testing.T) {
|
||||||
|
options := map[string]interface{}{
|
||||||
|
"endpoint": os.Getenv("S3_API"),
|
||||||
|
"region": "auto",
|
||||||
|
"key": os.Getenv("S3_ACCESS_KEY"),
|
||||||
|
"secret": os.Getenv("S3_SECRET_KEY"),
|
||||||
|
"bucket": os.Getenv("S3_BUCKET"),
|
||||||
|
"prefix": "vision-test",
|
||||||
|
"expiration": 10 * time.Minute,
|
||||||
|
}
|
||||||
|
|
||||||
|
storage, err := New(options)
|
||||||
|
if err != nil {
|
||||||
|
t.Logf("Error creating storage: %v", err)
|
||||||
|
}
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, storage)
|
||||||
|
if storage != nil {
|
||||||
|
assert.Equal(t, os.Getenv("S3_API"), storage.Endpoint)
|
||||||
|
assert.Equal(t, "auto", storage.Region)
|
||||||
|
assert.Equal(t, os.Getenv("S3_ACCESS_KEY"), storage.Key)
|
||||||
|
assert.Equal(t, os.Getenv("S3_SECRET_KEY"), storage.Secret)
|
||||||
|
assert.Equal(t, os.Getenv("S3_BUCKET"), storage.Bucket)
|
||||||
|
assert.Equal(t, "vision-test", storage.prefix)
|
||||||
|
assert.Equal(t, 10*time.Minute, storage.Expiration)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Upload and Download", func(t *testing.T) {
|
||||||
|
storage, err := New(map[string]interface{}{
|
||||||
|
"endpoint": os.Getenv("S3_API"),
|
||||||
|
"region": "auto",
|
||||||
|
"key": os.Getenv("S3_ACCESS_KEY"),
|
||||||
|
"secret": os.Getenv("S3_SECRET_KEY"),
|
||||||
|
"bucket": os.Getenv("S3_BUCKET"),
|
||||||
|
"prefix": "vision-test",
|
||||||
|
"expiration": 5 * time.Minute,
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
content := []byte("test content")
|
||||||
|
reader := bytes.NewReader(content)
|
||||||
|
fileID, err := storage.Upload(context.Background(), "test.txt", reader, "text/plain")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotEmpty(t, fileID)
|
||||||
|
|
||||||
|
// Get presigned URL
|
||||||
|
url := storage.URL(context.Background(), fileID)
|
||||||
|
assert.NotEmpty(t, url)
|
||||||
|
assert.Contains(t, url, "X-Amz-Signature")
|
||||||
|
assert.Contains(t, url, "X-Amz-Expires")
|
||||||
|
|
||||||
|
// Test with different expiration
|
||||||
|
storage.Expiration = 1 * time.Hour
|
||||||
|
url2 := storage.URL(context.Background(), fileID)
|
||||||
|
assert.NotEmpty(t, url2)
|
||||||
|
assert.Contains(t, url2, "X-Amz-Signature")
|
||||||
|
assert.Contains(t, url2, "X-Amz-Expires=3600")
|
||||||
|
|
||||||
|
// Download
|
||||||
|
reader2, contentType, err := storage.Download(context.Background(), fileID)
|
||||||
|
if err != nil {
|
||||||
|
t.Logf("Download error: %v", err)
|
||||||
|
t.FailNow()
|
||||||
|
}
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Contains(t, contentType, "text/plain")
|
||||||
|
|
||||||
|
if reader2 != nil {
|
||||||
|
downloaded, err := io.ReadAll(reader2)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, content, downloaded)
|
||||||
|
reader2.Close()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Upload with Custom Expiration", func(t *testing.T) {
|
||||||
|
storage, err := New(map[string]interface{}{
|
||||||
|
"endpoint": os.Getenv("S3_API"),
|
||||||
|
"region": "auto",
|
||||||
|
"key": os.Getenv("S3_ACCESS_KEY"),
|
||||||
|
"secret": os.Getenv("S3_SECRET_KEY"),
|
||||||
|
"bucket": os.Getenv("S3_BUCKET"),
|
||||||
|
"prefix": "vision-test",
|
||||||
|
"expiration": 5 * time.Minute,
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
content := []byte("test content")
|
||||||
|
reader := bytes.NewReader(content)
|
||||||
|
fileID, err := storage.Upload(context.Background(), "test.txt", reader, "text/plain")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotEmpty(t, fileID)
|
||||||
|
|
||||||
|
// Get URL with default expiration (5 minutes)
|
||||||
|
url := storage.URL(context.Background(), fileID)
|
||||||
|
assert.NotEmpty(t, url)
|
||||||
|
assert.Contains(t, url, "X-Amz-Signature")
|
||||||
|
assert.Contains(t, url, "X-Amz-Expires=300") // 5 minutes = 300 seconds
|
||||||
|
|
||||||
|
// Change expiration and get new URL
|
||||||
|
storage.Expiration = 2 * time.Hour
|
||||||
|
url2 := storage.URL(context.Background(), fileID)
|
||||||
|
assert.NotEmpty(t, url2)
|
||||||
|
assert.Contains(t, url2, "X-Amz-Signature")
|
||||||
|
assert.Contains(t, url2, "X-Amz-Expires=7200") // 2 hours = 7200 seconds
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Download Non-existent File", func(t *testing.T) {
|
||||||
|
storage, err := New(map[string]interface{}{
|
||||||
|
"endpoint": os.Getenv("S3_API"),
|
||||||
|
"region": "auto",
|
||||||
|
"key": os.Getenv("S3_ACCESS_KEY"),
|
||||||
|
"secret": os.Getenv("S3_SECRET_KEY"),
|
||||||
|
"bucket": os.Getenv("S3_BUCKET"),
|
||||||
|
"prefix": "vision-test",
|
||||||
|
"expiration": 5 * time.Minute,
|
||||||
|
})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
_, _, err = storage.Download(context.Background(), "non-existent.txt")
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
}
|
||||||
43
neo/vision/driver/types.go
Normal file
43
neo/vision/driver/types.go
Normal file
|
|
@ -0,0 +1,43 @@
|
||||||
|
package driver
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Config the vision configuration
|
||||||
|
type Config struct {
|
||||||
|
Storage StorageConfig `json:"storage" yaml:"storage"`
|
||||||
|
Model ModelConfig `json:"model" yaml:"model"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// StorageConfig the storage configuration
|
||||||
|
type StorageConfig struct {
|
||||||
|
Driver string `json:"driver" yaml:"driver"`
|
||||||
|
Options map[string]interface{} `json:"options" yaml:"options"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ModelConfig the model configuration
|
||||||
|
type ModelConfig struct {
|
||||||
|
Driver string `json:"driver" yaml:"driver"`
|
||||||
|
Options map[string]interface{} `json:"options" yaml:"options"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Storage the storage interface
|
||||||
|
type Storage interface {
|
||||||
|
Upload(ctx context.Context, filename string, reader io.Reader, contentType string) (string, error)
|
||||||
|
Download(ctx context.Context, fileID string) (io.ReadCloser, string, error)
|
||||||
|
URL(ctx context.Context, fileID string) string
|
||||||
|
}
|
||||||
|
|
||||||
|
// Model the vision model interface
|
||||||
|
type Model interface {
|
||||||
|
Analyze(ctx context.Context, fileID string, prompt string) (map[string]interface{}, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Response the vision response
|
||||||
|
type Response struct {
|
||||||
|
FileID string `json:"file_id" yaml:"file_id"`
|
||||||
|
URL string `json:"url" yaml:"url"`
|
||||||
|
Description map[string]interface{} `json:"description" yaml:"description"`
|
||||||
|
}
|
||||||
8
neo/vision/prompt.md
Normal file
8
neo/vision/prompt.md
Normal file
|
|
@ -0,0 +1,8 @@
|
||||||
|
根据这个数据结构,和说明实现一下对应的逻辑。
|
||||||
|
|
||||||
|
1. driver 单独一个目录. 每个 driver 一个目录,model 和 storage 在一级即可。 放在@vision 下
|
||||||
|
2. model driver: 支持 openai
|
||||||
|
3. storage driver 支持 local 和 s3 . local 使用我框架的 fs 实现,参考@file.go
|
||||||
|
4. 在 程序启动时候,设置视觉配置。(作为可选配置) @types.go
|
||||||
|
5. 统一的创建和调用入口,调用时需传入 chat model (用来判断是否支持视觉) 和图片路径,(使用 fs 读取)。 @vision.go
|
||||||
|
6. 用一个回调函数,外部传入用来格式化返回的数据。
|
||||||
110
neo/vision/vision.go
Normal file
110
neo/vision/vision.go
Normal file
|
|
@ -0,0 +1,110 @@
|
||||||
|
package vision
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/yaoapp/yao/neo/vision/driver"
|
||||||
|
"github.com/yaoapp/yao/neo/vision/driver/local"
|
||||||
|
"github.com/yaoapp/yao/neo/vision/driver/openai"
|
||||||
|
"github.com/yaoapp/yao/neo/vision/driver/s3"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Vision the vision service
|
||||||
|
type Vision struct {
|
||||||
|
storage driver.Storage
|
||||||
|
model driver.Model
|
||||||
|
}
|
||||||
|
|
||||||
|
// New create a new vision service
|
||||||
|
func New(cfg *driver.Config) (*Vision, error) {
|
||||||
|
|
||||||
|
// Create storage driver
|
||||||
|
var storage driver.Storage
|
||||||
|
var err error
|
||||||
|
switch cfg.Storage.Driver {
|
||||||
|
case "local":
|
||||||
|
storage, err = local.New(cfg.Storage.Options)
|
||||||
|
case "s3":
|
||||||
|
// Convert expiration string to duration if present
|
||||||
|
if exp, ok := cfg.Storage.Options["expiration"].(string); ok {
|
||||||
|
if duration, err := time.ParseDuration(exp); err == nil {
|
||||||
|
cfg.Storage.Options["expiration"] = duration
|
||||||
|
}
|
||||||
|
}
|
||||||
|
storage, err = s3.New(cfg.Storage.Options)
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf("storage driver %s not supported", cfg.Storage.Driver)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create storage driver error: %s", err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create model driver
|
||||||
|
var model driver.Model
|
||||||
|
switch cfg.Model.Driver {
|
||||||
|
case "openai":
|
||||||
|
model, err = openai.New(cfg.Model.Options)
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf("model driver %s not supported", cfg.Model.Driver)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create model driver error: %s", err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Vision{
|
||||||
|
storage: storage,
|
||||||
|
model: model,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Upload upload file
|
||||||
|
func (v *Vision) Upload(ctx context.Context, filename string, reader io.Reader, contentType string) (*driver.Response, error) {
|
||||||
|
fileID, err := v.storage.Upload(ctx, filename, reader, contentType)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &driver.Response{
|
||||||
|
FileID: fileID,
|
||||||
|
URL: v.storage.URL(ctx, fileID),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Analyze analyze image using vision model
|
||||||
|
func (v *Vision) Analyze(ctx context.Context, fileID string, prompt string) (*driver.Response, error) {
|
||||||
|
if v.model == nil {
|
||||||
|
return nil, fmt.Errorf("model is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
var url string
|
||||||
|
// If the input is already a base64 data URL or a HTTP(S) URL, use it directly
|
||||||
|
if strings.HasPrefix(fileID, "data:image/") || strings.HasPrefix(fileID, "http://") || strings.HasPrefix(fileID, "https://") {
|
||||||
|
url = fileID
|
||||||
|
} else {
|
||||||
|
// Otherwise, try to get the URL from storage
|
||||||
|
url = v.storage.URL(ctx, fileID)
|
||||||
|
if url == "" {
|
||||||
|
return nil, fmt.Errorf("failed to get URL for file %s", fileID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := v.model.Analyze(ctx, url, prompt)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &driver.Response{
|
||||||
|
FileID: fileID,
|
||||||
|
URL: url,
|
||||||
|
Description: result,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Download download file
|
||||||
|
func (v *Vision) Download(ctx context.Context, fileID string) (io.ReadCloser, string, error) {
|
||||||
|
return v.storage.Download(ctx, fileID)
|
||||||
|
}
|
||||||
334
neo/vision/vision_test.go
Normal file
334
neo/vision/vision_test.go
Normal file
|
|
@ -0,0 +1,334 @@
|
||||||
|
package vision
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/gou/fs"
|
||||||
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/neo/vision/driver"
|
||||||
|
"github.com/yaoapp/yao/test"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
// 1x1 transparent PNG
|
||||||
|
testImageBase64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestVision(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
|
||||||
|
// Setup test server for image hosting
|
||||||
|
imgServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Log request for debugging
|
||||||
|
t.Logf("Received request for: %s", r.URL.Path)
|
||||||
|
|
||||||
|
// Always return the test image
|
||||||
|
imgData, _ := base64.StdEncoding.DecodeString(testImageBase64)
|
||||||
|
w.Header().Set("Content-Type", "image/png")
|
||||||
|
w.Write(imgData)
|
||||||
|
}))
|
||||||
|
defer imgServer.Close()
|
||||||
|
|
||||||
|
t.Logf("Test server running at: %s", imgServer.URL)
|
||||||
|
|
||||||
|
t.Run("Create Vision Service", func(t *testing.T) {
|
||||||
|
vision, err := createTestVision(imgServer.URL)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, vision)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Upload and Download with Local Storage", func(t *testing.T) {
|
||||||
|
vision, err := createTestVision(imgServer.URL)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Test with text file
|
||||||
|
content := []byte("test content")
|
||||||
|
reader := bytes.NewReader(content)
|
||||||
|
resp, err := vision.Upload(context.Background(), "test.txt", reader, "text/plain")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotEmpty(t, resp.FileID)
|
||||||
|
assert.NotEmpty(t, resp.URL)
|
||||||
|
|
||||||
|
// Download
|
||||||
|
reader2, contentType, err := vision.Download(context.Background(), resp.FileID)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Contains(t, contentType, "text/plain")
|
||||||
|
|
||||||
|
if reader2 != nil {
|
||||||
|
downloaded, err := io.ReadAll(reader2)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, content, downloaded)
|
||||||
|
reader2.Close()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Upload and Download with S3 Storage", func(t *testing.T) {
|
||||||
|
vision, err := createTestVisionWithS3()
|
||||||
|
if err != nil {
|
||||||
|
t.Skip("S3 configuration not available")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test with text file
|
||||||
|
content := []byte("test content")
|
||||||
|
reader := bytes.NewReader(content)
|
||||||
|
resp, err := vision.Upload(context.Background(), "test.txt", reader, "text/plain")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotEmpty(t, resp.FileID)
|
||||||
|
assert.NotEmpty(t, resp.URL)
|
||||||
|
|
||||||
|
// Download
|
||||||
|
reader2, contentType, err := vision.Download(context.Background(), resp.FileID)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Contains(t, contentType, "text/plain")
|
||||||
|
|
||||||
|
if reader2 != nil {
|
||||||
|
downloaded, err := io.ReadAll(reader2)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, content, downloaded)
|
||||||
|
reader2.Close()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Analyze Image with Base64", func(t *testing.T) {
|
||||||
|
// Create vision service
|
||||||
|
cfg := &driver.Config{
|
||||||
|
Storage: driver.StorageConfig{
|
||||||
|
Driver: "local",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"path": "/__vision_test",
|
||||||
|
"compression": true,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Model: driver.ModelConfig{
|
||||||
|
Driver: "openai",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"api_key": os.Getenv("OPENAI_API_KEY"),
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
vision, err := New(cfg)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Use base64 data directly
|
||||||
|
result, err := vision.Analyze(context.Background(), "data:image/png;base64,"+testImageBase64, "Describe this image in detail")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.NotEmpty(t, result.Description)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Analyze Image with File", func(t *testing.T) {
|
||||||
|
// Create vision service
|
||||||
|
cfg := &driver.Config{
|
||||||
|
Storage: driver.StorageConfig{
|
||||||
|
Driver: "local",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"path": "/__vision_test",
|
||||||
|
"compression": true,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Model: driver.ModelConfig{
|
||||||
|
Driver: "openai",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"api_key": os.Getenv("OPENAI_API_KEY"),
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
vision, err := New(cfg)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Create test file
|
||||||
|
data, err := fs.Get("data")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Write test image data
|
||||||
|
imgData, err := base64.StdEncoding.DecodeString(testImageBase64)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
_, err = data.WriteFile("/test.png", imgData, 0644)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Analyze using file path
|
||||||
|
result, err := vision.Analyze(context.Background(), "/test.png", "Describe this image in detail")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.NotEmpty(t, result.Description)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Analyze Image with S3 URL", func(t *testing.T) {
|
||||||
|
if os.Getenv("S3_API") == "" || os.Getenv("S3_ACCESS_KEY") == "" ||
|
||||||
|
os.Getenv("S3_SECRET_KEY") == "" || os.Getenv("S3_BUCKET") == "" {
|
||||||
|
t.Skip("S3 environment variables not set")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create vision service
|
||||||
|
cfg := &driver.Config{
|
||||||
|
Storage: driver.StorageConfig{
|
||||||
|
Driver: "s3",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"endpoint": os.Getenv("S3_API"),
|
||||||
|
"region": "auto",
|
||||||
|
"key": os.Getenv("S3_ACCESS_KEY"),
|
||||||
|
"secret": os.Getenv("S3_SECRET_KEY"),
|
||||||
|
"bucket": os.Getenv("S3_BUCKET"),
|
||||||
|
"prefix": "vision-test",
|
||||||
|
"expiration": "5m",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Model: driver.ModelConfig{
|
||||||
|
Driver: "openai",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"api_key": os.Getenv("OPENAI_API_KEY"),
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
vision, err := New(cfg)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Upload test image
|
||||||
|
imgData, err := base64.StdEncoding.DecodeString(testImageBase64)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
reader := bytes.NewReader(imgData)
|
||||||
|
resp, err := vision.Upload(context.Background(), "test.png", reader, "image/png")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotEmpty(t, resp.FileID)
|
||||||
|
assert.NotEmpty(t, resp.URL)
|
||||||
|
|
||||||
|
// Analyze using S3 URL
|
||||||
|
result, err := vision.Analyze(context.Background(), resp.URL, "Describe this image in detail")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.NotEmpty(t, result.Description)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Invalid Model", func(t *testing.T) {
|
||||||
|
cfg := &driver.Config{
|
||||||
|
Storage: driver.StorageConfig{
|
||||||
|
Driver: "local",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"path": "/__vision_test",
|
||||||
|
"compression": true,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Model: driver.ModelConfig{
|
||||||
|
Driver: "invalid",
|
||||||
|
Options: map[string]interface{}{},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := New(cfg)
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "model driver invalid not supported")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Invalid Storage", func(t *testing.T) {
|
||||||
|
cfg := &driver.Config{
|
||||||
|
Storage: driver.StorageConfig{
|
||||||
|
Driver: "invalid",
|
||||||
|
Options: map[string]interface{}{},
|
||||||
|
},
|
||||||
|
Model: driver.ModelConfig{
|
||||||
|
Driver: "openai",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"api_key": "test",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := New(cfg)
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "storage driver invalid not supported")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func createTestVision(baseURL string) (*Vision, error) {
|
||||||
|
cfg := &driver.Config{
|
||||||
|
Storage: driver.StorageConfig{
|
||||||
|
Driver: "local",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"path": "/__vision_test",
|
||||||
|
"compression": true,
|
||||||
|
"base_url": baseURL,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Model: driver.ModelConfig{
|
||||||
|
Driver: "openai",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"api_key": os.Getenv("OPENAI_API_KEY"),
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
"prompt": `# Objective
|
||||||
|
You are a vision assistant, you can help the user to understand the image and describe it.
|
||||||
|
|
||||||
|
## Task Execution Steps
|
||||||
|
1. Understand the image/video and describe it.
|
||||||
|
2. Describe the image/video in detail.
|
||||||
|
|
||||||
|
## Result Format
|
||||||
|
{
|
||||||
|
"description": "The description of the image/video",
|
||||||
|
"content": "The content of the image/video"
|
||||||
|
}`,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
return New(cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
func createTestVisionWithS3() (*Vision, error) {
|
||||||
|
// Check required S3 environment variables
|
||||||
|
if os.Getenv("S3_API") == "" || os.Getenv("S3_ACCESS_KEY") == "" ||
|
||||||
|
os.Getenv("S3_SECRET_KEY") == "" || os.Getenv("S3_BUCKET") == "" {
|
||||||
|
return nil, fmt.Errorf("S3 environment variables not set")
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := &driver.Config{
|
||||||
|
Storage: driver.StorageConfig{
|
||||||
|
Driver: "s3",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"endpoint": os.Getenv("S3_API"),
|
||||||
|
"region": "auto",
|
||||||
|
"key": os.Getenv("S3_ACCESS_KEY"),
|
||||||
|
"secret": os.Getenv("S3_SECRET_KEY"),
|
||||||
|
"bucket": os.Getenv("S3_BUCKET"),
|
||||||
|
"prefix": "vision-test",
|
||||||
|
"expiration": "5m",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Model: driver.ModelConfig{
|
||||||
|
Driver: "openai",
|
||||||
|
Options: map[string]interface{}{
|
||||||
|
"api_key": os.Getenv("OPENAI_API_KEY"),
|
||||||
|
"model": os.Getenv("VISION_MODEL"),
|
||||||
|
"prompt": `# Objective
|
||||||
|
You are a vision assistant, you can help the user to understand the image and describe it.
|
||||||
|
|
||||||
|
## Task Execution Steps
|
||||||
|
1. Understand the image/video and describe it.
|
||||||
|
2. Describe the image/video in detail.
|
||||||
|
|
||||||
|
## Result Format
|
||||||
|
{
|
||||||
|
"description": "The description of the image/video",
|
||||||
|
"content": "The content of the image/video"
|
||||||
|
}`,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
return New(cfg)
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue