diff --git a/go.mod b/go.mod index e99535e3..f3c33bf8 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,10 @@ go 1.23 require ( 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/caarlos0/env/v6 v6.10.1 github.com/dchest/captcha v1.1.0 @@ -38,6 +42,20 @@ require ( filippo.io/edwards25519 v1.1.0 // indirect github.com/TylerBrock/colorjson v0.0.0-20200706003622-8a50f05110d2 // 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/bytedance/sonic v1.12.6 // indirect github.com/bytedance/sonic/loader v0.2.1 // indirect diff --git a/go.sum b/go.sum index a9b45ed1..a26a9c8b 100644 --- a/go.sum +++ b/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/andybalholm/cascadia v1.3.3 h1:AG2YHrzJIm4BZ19iwJ/DAua6Btl3IwJX+VI4kktS1LM= 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/go.mod h1:kRBLl5iJ+tD4TcOOxsy/0fnwebNt5EWlYSAyrTnjyyk= github.com/blang/semver/v4 v4.0.0 h1:1PFHFE6yCCTv8C1TeyNNarDzntLi7wMI5i/pzqYIsAM= diff --git a/neo/vision/driver/local/storage.go b/neo/vision/driver/local/storage.go new file mode 100644 index 00000000..ed873129 --- /dev/null +++ b/neo/vision/driver/local/storage.go @@ -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) +} diff --git a/neo/vision/driver/local/storage_test.go b/neo/vision/driver/local/storage_test.go new file mode 100644 index 00000000..15baa1a2 --- /dev/null +++ b/neo/vision/driver/local/storage_test.go @@ -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) + }) +} diff --git a/neo/vision/driver/openai/model.go b/neo/vision/driver/openai/model.go new file mode 100644 index 00000000..a766c380 --- /dev/null +++ b/neo/vision/driver/openai/model.go @@ -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 +} diff --git a/neo/vision/driver/openai/model_test.go b/neo/vision/driver/openai/model_test.go new file mode 100644 index 00000000..e93d516e --- /dev/null +++ b/neo/vision/driver/openai/model_test.go @@ -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") + }) +} diff --git a/neo/vision/driver/s3/storage.go b/neo/vision/driver/s3/storage.go new file mode 100644 index 00000000..addc3b10 --- /dev/null +++ b/neo/vision/driver/s3/storage.go @@ -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) +} diff --git a/neo/vision/driver/s3/storage_test.go b/neo/vision/driver/s3/storage_test.go new file mode 100644 index 00000000..0f3d3847 --- /dev/null +++ b/neo/vision/driver/s3/storage_test.go @@ -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) + }) +} diff --git a/neo/vision/driver/types.go b/neo/vision/driver/types.go new file mode 100644 index 00000000..ccc64ee6 --- /dev/null +++ b/neo/vision/driver/types.go @@ -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"` +} diff --git a/neo/vision/prompt.md b/neo/vision/prompt.md new file mode 100644 index 00000000..2d6e4f34 --- /dev/null +++ b/neo/vision/prompt.md @@ -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. 用一个回调函数,外部传入用来格式化返回的数据。 diff --git a/neo/vision/vision.go b/neo/vision/vision.go new file mode 100644 index 00000000..79c5f372 --- /dev/null +++ b/neo/vision/vision.go @@ -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) +} diff --git a/neo/vision/vision_test.go b/neo/vision/vision_test.go new file mode 100644 index 00000000..73379487 --- /dev/null +++ b/neo/vision/vision_test.go @@ -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) +}