Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
146 lines
5.1 KiB
Go
146 lines
5.1 KiB
Go
package handler
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
|
"github.com/Wei-Shaw/sub2api/internal/service"
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/stretchr/testify/require"
|
|
)
|
|
|
|
type asyncImageMemoryStore struct {
|
|
mu sync.RWMutex
|
|
tasks map[string]*service.ImageTaskRecord
|
|
}
|
|
|
|
func (s *asyncImageMemoryStore) Save(_ context.Context, task *service.ImageTaskRecord, _ time.Duration) error {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
copy := *task
|
|
copy.Result = append(json.RawMessage(nil), task.Result...)
|
|
copy.Error = append(json.RawMessage(nil), task.Error...)
|
|
s.tasks[task.ID] = ©
|
|
return nil
|
|
}
|
|
|
|
func (s *asyncImageMemoryStore) Get(_ context.Context, id string) (*service.ImageTaskRecord, error) {
|
|
s.mu.RLock()
|
|
defer s.mu.RUnlock()
|
|
task := s.tasks[id]
|
|
if task == nil {
|
|
return nil, service.ErrImageTaskNotFound
|
|
}
|
|
copy := *task
|
|
copy.Result = append(json.RawMessage(nil), task.Result...)
|
|
copy.Error = append(json.RawMessage(nil), task.Error...)
|
|
return ©, nil
|
|
}
|
|
|
|
func TestAsyncImageHandlerSubmitAndPoll(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
store := &asyncImageMemoryStore{tasks: make(map[string]*service.ImageTaskRecord)}
|
|
tasks := service.NewImageTaskServiceWithUploader(store, nil, time.Hour, time.Minute)
|
|
release := make(chan struct{})
|
|
h := &AsyncImageHandler{tasks: tasks}
|
|
h.execute = func(_ string, c *gin.Context) {
|
|
<-release
|
|
c.JSON(http.StatusOK, gin.H{"created": 123, "data": []gin.H{{"url": "https://example.test/image.png"}}})
|
|
}
|
|
|
|
router := gin.New()
|
|
router.Use(func(c *gin.Context) {
|
|
groupID := int64(3)
|
|
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
|
|
ID: 9,
|
|
UserID: 7,
|
|
GroupID: &groupID,
|
|
Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI, AllowImageGeneration: true},
|
|
})
|
|
c.Next()
|
|
})
|
|
router.POST("/v1/images/generations/async", h.Submit)
|
|
router.GET("/v1/images/tasks/:task_id", h.Get)
|
|
|
|
requestCtx, cancelRequest := context.WithCancel(context.Background())
|
|
req := httptest.NewRequest(http.MethodPost, "/v1/images/generations/async", strings.NewReader(`{"model":"gpt-image-1","prompt":"cat"}`)).WithContext(requestCtx)
|
|
req.Header.Set("Content-Type", "application/json")
|
|
w := httptest.NewRecorder()
|
|
router.ServeHTTP(w, req)
|
|
require.Equal(t, http.StatusAccepted, w.Code)
|
|
require.Equal(t, "no-store", w.Header().Get("Cache-Control"))
|
|
require.Equal(t, "3", w.Header().Get("Retry-After"))
|
|
|
|
var accepted struct {
|
|
TaskID string `json:"task_id"`
|
|
Status string `json:"status"`
|
|
PollURL string `json:"poll_url"`
|
|
}
|
|
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &accepted))
|
|
require.Equal(t, service.ImageTaskStatusProcessing, accepted.Status)
|
|
require.Equal(t, "/v1/images/tasks/"+accepted.TaskID, accepted.PollURL)
|
|
require.Equal(t, accepted.PollURL, w.Header().Get("Location"))
|
|
|
|
// The detached background request must survive completion of/cancellation
|
|
// from the short submission request.
|
|
cancelRequest()
|
|
close(release)
|
|
require.Eventually(t, func() bool {
|
|
got, err := tasks.Get(context.Background(), service.ImageTaskOwner{UserID: 7, APIKeyID: 9}, accepted.TaskID)
|
|
return err == nil && got.Status == service.ImageTaskStatusCompleted
|
|
}, time.Second, 10*time.Millisecond)
|
|
|
|
pollReq := httptest.NewRequest(http.MethodGet, accepted.PollURL, nil)
|
|
pollWriter := httptest.NewRecorder()
|
|
router.ServeHTTP(pollWriter, pollReq)
|
|
require.Equal(t, http.StatusOK, pollWriter.Code)
|
|
require.Equal(t, "no-store", pollWriter.Header().Get("Cache-Control"))
|
|
require.Empty(t, pollWriter.Header().Get("Retry-After"))
|
|
require.Contains(t, pollWriter.Body.String(), "https://example.test/image.png")
|
|
}
|
|
|
|
// When object storage is not configured the feature is fully disabled: the
|
|
// endpoints must return 404 without creating a task or writing to Redis.
|
|
func TestAsyncImageHandlerDisabledReturns404(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
store := &asyncImageMemoryStore{tasks: make(map[string]*service.ImageTaskRecord)}
|
|
tasks := service.NewImageTaskServiceWithOptions(store, time.Hour, time.Minute) // enabled == false
|
|
h := &AsyncImageHandler{tasks: tasks}
|
|
|
|
router := gin.New()
|
|
router.Use(func(c *gin.Context) {
|
|
groupID := int64(3)
|
|
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
|
|
ID: 9,
|
|
UserID: 7,
|
|
GroupID: &groupID,
|
|
Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI, AllowImageGeneration: true},
|
|
})
|
|
c.Next()
|
|
})
|
|
router.POST("/v1/images/generations/async", h.Submit)
|
|
router.GET("/v1/images/tasks/:task_id", h.Get)
|
|
|
|
req := httptest.NewRequest(http.MethodPost, "/v1/images/generations/async", strings.NewReader(`{"model":"gpt-image-1","prompt":"cat"}`))
|
|
req.Header.Set("Content-Type", "application/json")
|
|
w := httptest.NewRecorder()
|
|
router.ServeHTTP(w, req)
|
|
require.Equal(t, http.StatusNotFound, w.Code)
|
|
require.Contains(t, w.Body.String(), "not enabled")
|
|
|
|
pollReq := httptest.NewRequest(http.MethodGet, "/v1/images/tasks/imgtask_missing", nil)
|
|
pollWriter := httptest.NewRecorder()
|
|
router.ServeHTTP(pollWriter, pollReq)
|
|
require.Equal(t, http.StatusNotFound, pollWriter.Code)
|
|
|
|
// No task was created / persisted.
|
|
require.Empty(t, store.tasks)
|
|
}
|