515 lines
28 KiB
Go
515 lines
28 KiB
Go
package handler
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"image"
|
|
"image/color"
|
|
"image/png"
|
|
"mime/multipart"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"net/textproto"
|
|
"os"
|
|
"strconv"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
|
"git.ilapage.cn/OPC/chorus/internal/core/queue"
|
|
corerouter "git.ilapage.cn/OPC/chorus/internal/core/router"
|
|
corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage"
|
|
passwordpkg "git.ilapage.cn/OPC/chorus/internal/platform/password"
|
|
platformstorage "git.ilapage.cn/OPC/chorus/internal/platform/storage"
|
|
"git.ilapage.cn/OPC/chorus/portal/auth"
|
|
"git.ilapage.cn/OPC/chorus/portal/service"
|
|
"git.ilapage.cn/OPC/chorus/portal/session"
|
|
"github.com/gin-gonic/gin"
|
|
"gorm.io/driver/mysql"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
type apiClient struct {
|
|
router http.Handler
|
|
cookie *http.Cookie
|
|
csrf string
|
|
}
|
|
|
|
func (c *apiClient) do(method, path string, body []byte, contentType string) *httptest.ResponseRecorder {
|
|
request := httptest.NewRequest(method, path, bytes.NewReader(body))
|
|
request.RemoteAddr = "192.0.2.10:1234"
|
|
if contentType != "" {
|
|
request.Header.Set("Content-Type", contentType)
|
|
}
|
|
if c.cookie != nil {
|
|
request.AddCookie(c.cookie)
|
|
}
|
|
if c.csrf != "" {
|
|
request.Header.Set("X-CSRF-Token", c.csrf)
|
|
}
|
|
response := httptest.NewRecorder()
|
|
c.router.ServeHTTP(response, request)
|
|
for _, cookie := range response.Result().Cookies() {
|
|
if cookie.Name == "chorus_session" {
|
|
c.cookie = cookie
|
|
}
|
|
}
|
|
return response
|
|
}
|
|
func (c *apiClient) start(t *testing.T) {
|
|
t.Helper()
|
|
response := c.do(http.MethodGet, "/api/session", nil, "")
|
|
var body struct {
|
|
CSRF string `json:"csrf_token"`
|
|
}
|
|
if response.Code != 200 || json.Unmarshal(response.Body.Bytes(), &body) != nil || body.CSRF == "" {
|
|
t.Fatalf("session response=%d %s", response.Code, response.Body.String())
|
|
}
|
|
c.csrf = body.CSRF
|
|
}
|
|
func (c *apiClient) login(t *testing.T, email, password string) *httptest.ResponseRecorder {
|
|
t.Helper()
|
|
body, _ := json.Marshal(map[string]string{"email": email, "password": password})
|
|
response := c.do(http.MethodPost, "/api/session/login", body, "application/json")
|
|
if response.Code == 200 {
|
|
var result struct {
|
|
CSRF string `json:"csrf_token"`
|
|
}
|
|
if json.Unmarshal(response.Body.Bytes(), &result) != nil || result.CSRF == "" {
|
|
t.Fatalf("login body=%s", response.Body.String())
|
|
}
|
|
c.csrf = result.CSRF
|
|
}
|
|
return response
|
|
}
|
|
|
|
func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) {
|
|
dsn := os.Getenv("CHORUS_TEST_DSN")
|
|
if dsn == "" {
|
|
t.Skip("CHORUS_TEST_DSN is not set")
|
|
}
|
|
gin.SetMode(gin.TestMode)
|
|
db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var databaseName string
|
|
if err := db.Raw("SELECT DATABASE()").Scan(&databaseName).Error; err != nil || databaseName != "chorus_test" {
|
|
t.Fatalf("portal integration requires chorus_test database, got %q: %v", databaseName, err)
|
|
}
|
|
sqlDB, _ := db.DB()
|
|
defer sqlDB.Close()
|
|
suffix := strconv.FormatInt(time.Now().UnixNano(), 10)
|
|
password := "test-password-" + suffix
|
|
encoded, err := passwordpkg.Encode(password)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
users := []model.User{{Email: "portal-a-" + suffix + "@example.invalid", PasswordHash: encoded, DisplayName: "User A", Status: "active"}, {Email: "portal-b-" + suffix + "@example.invalid", PasswordHash: encoded, DisplayName: "User B", Status: "active"}}
|
|
if err := db.Create(&users).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
templates := []model.PromptTemplate{
|
|
{TemplateKey: "portal-text-" + suffix, Kind: model.KindText, APIType: model.APIChat, Capability: model.CapabilityText, Name: "Portal Text", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true},
|
|
{TemplateKey: "portal-generate-" + suffix, Kind: model.KindImage, APIType: model.APIImages, Capability: model.CapabilityImageGenerate, Name: "Portal Generate", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true},
|
|
{TemplateKey: "portal-edit-" + suffix, Kind: model.KindImage, APIType: model.APIImagesEdits, Capability: model.CapabilityImageEdit, Name: "Portal Edit", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "default role rule", Enabled: true},
|
|
}
|
|
if err := db.Create(&templates).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
routeIDs, providerID := createPortalRoutes(t, db, suffix, templates)
|
|
defer func() {
|
|
db.Where("user_id IN ?", []uint64{users[0].ID, users[1].ID}).Delete(&model.Generation{})
|
|
db.Exec("DELETE FROM active_routes WHERE route_pool_id IN ?", routeIDs)
|
|
db.Exec("DELETE FROM route_pools WHERE id IN ?", routeIDs)
|
|
db.Exec("DELETE FROM provider_model_capabilities WHERE provider_model_id IN (SELECT id FROM provider_models WHERE provider_id=?)", providerID)
|
|
db.Exec("DELETE FROM provider_models WHERE provider_id=?", providerID)
|
|
db.Exec("DELETE FROM providers WHERE id=?", providerID)
|
|
db.Where("id IN ?", []uint64{templates[0].ID, templates[1].ID, templates[2].ID}).Delete(&model.PromptTemplate{})
|
|
db.Where("id IN ?", []uint64{users[0].ID, users[1].ID}).Delete(&model.User{})
|
|
}()
|
|
queueRepository, _ := queue.NewMySQLRepository(db)
|
|
storageRoot := t.TempDir()
|
|
local, err := platformstorage.NewLocal(platformstorage.Config{Root: storageRoot, MaxObjectBytes: 1 << 20, MaxImagePixels: 10000, ThumbnailMaxSide: 32, AllowedImageMIME: map[string]bool{"image/png": true, "image/jpeg": true}})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
routeRepository, _ := corerouter.NewGORMRepository(db)
|
|
generationService, err := service.New(db, queueRepository, local, routeRepository, service.Config{MaxPromptBytes: 1000, MaxImages: 3, MaxImageBytes: 200 << 10, MaxUploadBytes: 1 << 20, MaxImagePixels: 10000, HistoryLimit: 20, CursorKey: []byte(strings.Repeat("c", 32)), CursorTTL: time.Hour})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
sessions, _ := session.New([]byte(strings.Repeat("s", 32)), time.Hour, false)
|
|
authService, _ := auth.NewService(db, 20, time.Minute)
|
|
router, err := NewRouter(sessions, authService, generationService, 1<<20)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
staticRequest := httptest.NewRequest(http.MethodGet, "/static/app.css", nil)
|
|
staticResponse := httptest.NewRecorder()
|
|
router.ServeHTTP(staticResponse, staticRequest)
|
|
if staticResponse.Code != http.StatusOK || staticResponse.Header().Get("Set-Cookie") != "" || !strings.Contains(staticResponse.Header().Get("Cache-Control"), "max-age") {
|
|
t.Fatalf("static asset response=%d cookie=%q cache=%q", staticResponse.Code, staticResponse.Header().Get("Set-Cookie"), staticResponse.Header().Get("Cache-Control"))
|
|
}
|
|
loginPage := (&apiClient{router: router}).do(http.MethodGet, "/login?return_to=%2F%2Fevil.example", nil, "")
|
|
if loginPage.Code != http.StatusOK || !strings.Contains(loginPage.Body.String(), `data-return-to="/"`) || strings.Contains(loginPage.Body.String(), "cdn.") || strings.Contains(loginPage.Body.String(), "test-password") {
|
|
t.Fatalf("login page=%d %s", loginPage.Code, loginPage.Body.String())
|
|
}
|
|
|
|
unknown := &apiClient{router: router}
|
|
unknown.start(t)
|
|
unknownResponse := unknown.login(t, "missing-"+suffix+"@example.invalid", "wrong-password-value")
|
|
wrong := &apiClient{router: router}
|
|
wrong.start(t)
|
|
wrongResponse := wrong.login(t, users[0].Email, "wrong-password-value")
|
|
if unknownResponse.Code != 401 || wrongResponse.Code != 401 || unknownResponse.Body.String() != wrongResponse.Body.String() {
|
|
t.Fatalf("credential responses differ: %d %s / %d %s", unknownResponse.Code, unknownResponse.Body.String(), wrongResponse.Code, wrongResponse.Body.String())
|
|
}
|
|
|
|
clientA := &apiClient{router: router}
|
|
clientA.start(t)
|
|
oldCookie := *clientA.cookie
|
|
if response := clientA.login(t, users[0].Email, password); response.Code != 200 {
|
|
t.Fatalf("login=%d %s", response.Code, response.Body.String())
|
|
}
|
|
emptyWorkspace := clientA.do(http.MethodGet, "/", nil, "")
|
|
if emptyWorkspace.Code != http.StatusOK || !strings.Contains(emptyWorkspace.Body.String(), "还没有生成记录") || !strings.Contains(emptyWorkspace.Body.String(), "开始第一次生成") {
|
|
t.Fatalf("empty workspace=%d %s", emptyWorkspace.Code, emptyWorkspace.Body.String())
|
|
}
|
|
oldRequest := httptest.NewRequest(http.MethodGet, "/api/generations", nil)
|
|
oldRequest.AddCookie(&oldCookie)
|
|
oldResponse := httptest.NewRecorder()
|
|
router.ServeHTTP(oldResponse, oldRequest)
|
|
if oldResponse.Code != 401 {
|
|
t.Fatalf("old session code=%d", oldResponse.Code)
|
|
}
|
|
|
|
textBody, _ := json.Marshal(map[string]string{"idempotency_key": "text-key-" + suffix, "prompt": "write a title"})
|
|
textResponse := clientA.do(http.MethodPost, "/api/generations/text", textBody, "application/json")
|
|
if textResponse.Code != 202 {
|
|
t.Fatalf("text submit=%d %s", textResponse.Code, textResponse.Body.String())
|
|
}
|
|
textID := responseID(t, textResponse)
|
|
replay := clientA.do(http.MethodPost, "/api/generations/text", textBody, "application/json")
|
|
if replay.Code != 200 || responseID(t, replay) != textID {
|
|
t.Fatalf("text replay=%d %s", replay.Code, replay.Body.String())
|
|
}
|
|
var generation model.Generation
|
|
if err := db.First(&generation, textID).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if generation.Status != model.StatusPending || generation.ProviderModelID != nil || generation.AttemptCount != 0 || generation.RoutePoolID == nil || generation.PromptTemplateID == nil || len(generation.RouteSnapshot) == 0 {
|
|
t.Fatalf("synchronous submission touched worker fields: %#v", generation)
|
|
}
|
|
workspace := clientA.do(http.MethodGet, "/", nil, "")
|
|
if workspace.Code != http.StatusOK || !strings.Contains(workspace.Body.String(), "User A") || !strings.Contains(workspace.Body.String(), `content="1000"`) || !strings.Contains(workspace.Body.String(), `content="3"`) || strings.Contains(workspace.Body.String(), "cdn.") {
|
|
t.Fatalf("workspace=%d %s", workspace.Code, workspace.Body.String())
|
|
}
|
|
pendingPage := clientA.do(http.MethodGet, fmt.Sprintf("/generations/%d", textID), nil, "")
|
|
if pendingPage.Code != http.StatusOK || !strings.Contains(pendingPage.Body.String(), "任务正在排队") || !strings.Contains(pendingPage.Body.String(), "hx-get=") {
|
|
t.Fatalf("pending page=%d %s", pendingPage.Code, pendingPage.Body.String())
|
|
}
|
|
leaseOwner := "ui-test-worker"
|
|
leaseToken := "ui-test-lease"
|
|
if err := db.Model(&model.Generation{}).Where("id=?", textID).Updates(map[string]any{"status": model.StatusRunning, "started_at": time.Now(), "lease_owner": leaseOwner, "lease_token": leaseToken, "lease_until": time.Now().Add(time.Minute)}).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
runningFragment := clientA.do(http.MethodGet, fmt.Sprintf("/ui/generations/%d/result", textID), nil, "")
|
|
if runningFragment.Code != http.StatusOK || !strings.Contains(runningFragment.Body.String(), "正在生成内容") || !strings.Contains(runningFragment.Body.String(), "hx-get=") {
|
|
t.Fatalf("running fragment=%d %s", runningFragment.Code, runningFragment.Body.String())
|
|
}
|
|
errorMessage := `<script>alert("unsafe")</script>`
|
|
if err := db.Model(&model.Generation{}).Where("id=?", textID).Updates(map[string]any{"status": model.StatusFailed, "error_code": "upstream_error", "error_message": errorMessage, "completed_at": time.Now(), "lease_owner": nil, "lease_token": nil, "lease_until": nil}).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
failedFragment := clientA.do(http.MethodGet, fmt.Sprintf("/ui/generations/%d/result", textID), nil, "")
|
|
if failedFragment.Code != http.StatusOK || !strings.Contains(failedFragment.Body.String(), "这次生成没有完成") || strings.Contains(failedFragment.Body.String(), "hx-get=") || strings.Contains(failedFragment.Body.String(), errorMessage) || !strings.Contains(failedFragment.Body.String(), "<script>") {
|
|
t.Fatalf("failed fragment=%d %s", failedFragment.Code, failedFragment.Body.String())
|
|
}
|
|
|
|
successBody, _ := json.Marshal(map[string]string{"idempotency_key": "text-success-" + suffix, "prompt": "write a successful title"})
|
|
successResponse := clientA.do(http.MethodPost, "/api/generations/text", successBody, "application/json")
|
|
if successResponse.Code != http.StatusAccepted {
|
|
t.Fatalf("successful text submit=%d %s", successResponse.Code, successResponse.Body.String())
|
|
}
|
|
successID := responseID(t, successResponse)
|
|
textResult := "generated text result"
|
|
if err := db.Create(&model.GenerationOutput{GenerationID: successID, Kind: model.KindText, TextContent: &textResult}).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := db.Model(&model.Generation{}).Where("id=?", successID).Updates(map[string]any{"status": model.StatusSucceeded, "completed_at": time.Now()}).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
succeededTextFragment := clientA.do(http.MethodGet, fmt.Sprintf("/ui/generations/%d/result", successID), nil, "")
|
|
if succeededTextFragment.Code != http.StatusOK || !strings.Contains(succeededTextFragment.Body.String(), textResult) || !strings.Contains(succeededTextFragment.Body.String(), "复制") || strings.Contains(succeededTextFragment.Body.String(), "hx-get=") {
|
|
t.Fatalf("succeeded text fragment=%d %s", succeededTextFragment.Code, succeededTextFragment.Body.String())
|
|
}
|
|
|
|
generateKey := "generate-key-" + suffix
|
|
generateResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, generateKey, "generate image", model.CapabilityImageGenerate, "", nil), lastMultipartType)
|
|
if generateResponse.Code != http.StatusAccepted {
|
|
t.Fatalf("image generate submit=%d %s", generateResponse.Code, generateResponse.Body.String())
|
|
}
|
|
generateID := responseID(t, generateResponse)
|
|
var generated model.Generation
|
|
if err := db.First(&generated, generateID).Error; err != nil || generated.RenderedPrompt != "generate image" || generated.RoleRule != nil {
|
|
t.Fatalf("image generate=%#v err=%v", generated, err)
|
|
}
|
|
|
|
imageKey := "image-key-" + suffix
|
|
imageResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", model.CapabilityImageEdit, "user role rule", []uploadPart{{clientID: "primary-a", name: "first.png", mime: "image/png", data: testPNG(4, 4), position: 0, role: model.RolePrimary, note: "keep edges"}}), lastMultipartType)
|
|
if imageResponse.Code != 202 {
|
|
t.Fatalf("image submit=%d %s", imageResponse.Code, imageResponse.Body.String())
|
|
}
|
|
imageID := responseID(t, imageResponse)
|
|
var inputs []model.GenerationInput
|
|
if err := db.Where("generation_id=?", imageID).Find(&inputs).Error; err != nil || len(inputs) != 1 || inputs[0].Role != model.RolePrimary || inputs[0].Note == nil || *inputs[0].Note != "keep edges" {
|
|
t.Fatalf("inputs=%#v err=%v", inputs, err)
|
|
}
|
|
var imageGeneration model.Generation
|
|
if err := db.First(&imageGeneration, imageID).Error; err != nil || imageGeneration.RoleRule == nil || *imageGeneration.RoleRule != "user role rule" || !strings.Contains(imageGeneration.RenderedPrompt, "user role rule") || strings.Contains(imageGeneration.RenderedPrompt, "default role rule") || !strings.Contains(imageGeneration.RenderedPrompt, "position=0 role=primary note=keep edges") || imageGeneration.RoutePoolID == nil {
|
|
t.Fatalf("image generation=%#v err=%v", imageGeneration, err)
|
|
}
|
|
if err := db.Table("route_pools").Where("id=?", routeIDs[2]).Update("enabled", false).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
imageReplay := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", model.CapabilityImageEdit, "changed rule", []uploadPart{{clientID: "primary-b", name: "second.png", mime: "image/png", data: testPNG(4, 4), position: 0, role: model.RolePrimary}}), lastMultipartType)
|
|
if imageReplay.Code != 200 || responseID(t, imageReplay) != imageID {
|
|
t.Fatalf("image replay=%d %s", imageReplay.Code, imageReplay.Body.String())
|
|
}
|
|
if err := db.Table("route_pools").Where("id=?", routeIDs[2]).Update("enabled", true).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var inputCount int64
|
|
db.Model(&model.GenerationInput{}).Where("generation_id=?", imageID).Count(&inputCount)
|
|
if inputCount != 1 {
|
|
t.Fatalf("replay wrote %d inputs", inputCount)
|
|
}
|
|
invalidResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, "invalid-key-"+suffix, "edit image", model.CapabilityImageEdit, "", []uploadPart{{clientID: "bad-file", name: "fake.txt", mime: "image/png", data: testPNG(4, 4), position: 0, role: model.RolePrimary}}), lastMultipartType)
|
|
if invalidResponse.Code != 400 || !strings.Contains(invalidResponse.Body.String(), `"code":"invalid_image_extension"`) || !strings.Contains(invalidResponse.Body.String(), `"index":0`) || !strings.Contains(invalidResponse.Body.String(), `"name":"fake.txt"`) {
|
|
t.Fatalf("invalid image response=%d %s", invalidResponse.Code, invalidResponse.Body.String())
|
|
}
|
|
historyResponse := clientA.do(http.MethodGet, "/api/generations?limit=2", nil, "")
|
|
var firstHistory struct {
|
|
Items []struct {
|
|
ID uint64 `json:"id"`
|
|
} `json:"items"`
|
|
NextCursor string `json:"next_cursor"`
|
|
HasMore bool `json:"has_more"`
|
|
}
|
|
if historyResponse.Code != 200 || json.Unmarshal(historyResponse.Body.Bytes(), &firstHistory) != nil || len(firstHistory.Items) != 2 || !firstHistory.HasMore || firstHistory.NextCursor == "" {
|
|
t.Fatalf("first history=%d %s", historyResponse.Code, historyResponse.Body.String())
|
|
}
|
|
nextHistoryResponse := clientA.do(http.MethodGet, "/api/generations?limit=2&cursor="+firstHistory.NextCursor, nil, "")
|
|
var nextHistory struct {
|
|
Items []struct {
|
|
ID uint64 `json:"id"`
|
|
} `json:"items"`
|
|
}
|
|
if nextHistoryResponse.Code != 200 || json.Unmarshal(nextHistoryResponse.Body.Bytes(), &nextHistory) != nil || len(nextHistory.Items) == 0 || nextHistory.Items[0].ID == firstHistory.Items[0].ID || nextHistory.Items[0].ID == firstHistory.Items[1].ID {
|
|
t.Fatalf("next history=%d %s", nextHistoryResponse.Code, nextHistoryResponse.Body.String())
|
|
}
|
|
tamperedBytes := []byte(firstHistory.NextCursor)
|
|
tamperAt := len(tamperedBytes) / 2
|
|
if tamperedBytes[tamperAt] == 'A' {
|
|
tamperedBytes[tamperAt] = 'B'
|
|
} else {
|
|
tamperedBytes[tamperAt] = 'A'
|
|
}
|
|
tamperedCursor := string(tamperedBytes)
|
|
if response := clientA.do(http.MethodGet, "/api/generations?limit=2&cursor="+tamperedCursor, nil, ""); response.Code != 400 || !strings.Contains(response.Body.String(), `"code":"invalid_cursor"`) {
|
|
t.Fatalf("tampered cursor=%d %s", response.Code, response.Body.String())
|
|
}
|
|
if response := clientA.do(http.MethodGet, "/api/generations?limit=21", nil, ""); response.Code != 400 || !strings.Contains(response.Body.String(), `"code":"invalid_limit"`) {
|
|
t.Fatalf("invalid limit=%d %s", response.Code, response.Body.String())
|
|
}
|
|
if response := clientA.do(http.MethodGet, "/api/generations?cursor=", nil, ""); response.Code != 400 || !strings.Contains(response.Body.String(), `"code":"invalid_cursor"`) {
|
|
t.Fatalf("empty cursor=%d %s", response.Code, response.Body.String())
|
|
}
|
|
|
|
clientB := &apiClient{router: router}
|
|
clientB.start(t)
|
|
if response := clientB.login(t, users[1].Email, password); response.Code != 200 {
|
|
t.Fatalf("login B=%d", response.Code)
|
|
}
|
|
if response := clientB.do(http.MethodGet, fmt.Sprintf("/api/generations/%d", imageID), nil, ""); response.Code != 404 {
|
|
t.Fatalf("cross-user detail=%d", response.Code)
|
|
}
|
|
if response := clientB.do(http.MethodGet, "/api/generations?limit=2&cursor="+firstHistory.NextCursor, nil, ""); response.Code != 400 || !strings.Contains(response.Body.String(), `"code":"invalid_cursor"`) {
|
|
t.Fatalf("cross-user cursor=%d %s", response.Code, response.Body.String())
|
|
}
|
|
if response := clientB.do(http.MethodGet, fmt.Sprintf("/api/generations/%d/inputs/%d", imageID, inputs[0].ID), nil, ""); response.Code != 404 {
|
|
t.Fatalf("cross-user input=%d", response.Code)
|
|
}
|
|
if response := clientA.do(http.MethodGet, fmt.Sprintf("/api/generations/%d/inputs/%d", imageID, inputs[0].ID), nil, ""); response.Code != 200 || response.Header().Get("Content-Type") != "image/png" {
|
|
t.Fatalf("own input=%d %s", response.Code, response.Header().Get("Content-Type"))
|
|
}
|
|
|
|
outputData := testPNG(8, 8)
|
|
original, _ := local.Put(context.Background(), corestorage.PutRequest{Key: fmt.Sprintf("test-output/%d/original", imageID), OwnerID: users[0].ID, GenerationID: imageID, ContentType: "image/png", Source: bytes.NewReader(outputData)})
|
|
thumbnail, _ := local.Put(context.Background(), corestorage.PutRequest{Key: fmt.Sprintf("test-output/%d/thumbnail", imageID), OwnerID: users[0].ID, GenerationID: imageID, ContentType: "image/png", Source: bytes.NewReader(testPNG(2, 2))})
|
|
mimeType := "image/png"
|
|
size := uint64(original.Size)
|
|
output := model.GenerationOutput{GenerationID: imageID, Kind: model.KindImage, StorageKey: &original.Key, ThumbnailStorageKey: &thumbnail.Key, MIMEType: &mimeType, SizeBytes: &size}
|
|
if err := db.Create(&output).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := db.Model(&model.Generation{}).Where("id=?", imageID).Updates(map[string]any{"status": model.StatusSucceeded, "completed_at": time.Now()}).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
succeededPage := clientA.do(http.MethodGet, fmt.Sprintf("/generations/%d", imageID), nil, "")
|
|
if succeededPage.Code != http.StatusOK || !strings.Contains(succeededPage.Body.String(), "已完成") || !strings.Contains(succeededPage.Body.String(), fmt.Sprintf("/api/generations/%d/outputs/%d", imageID, output.ID)) || strings.Contains(succeededPage.Body.String(), storageRoot) || strings.Contains(succeededPage.Body.String(), original.Key) || strings.Contains(succeededPage.Body.String(), "hx-get=") {
|
|
t.Fatalf("succeeded page=%d %s", succeededPage.Code, succeededPage.Body.String())
|
|
}
|
|
detail := clientA.do(http.MethodGet, fmt.Sprintf("/api/generations/%d", imageID), nil, "")
|
|
if detail.Code != 200 || !strings.Contains(detail.Body.String(), `"terminal":true`) || strings.Contains(detail.Body.String(), storageRoot) || strings.Contains(detail.Body.String(), inputs[0].StorageKey) {
|
|
t.Fatalf("terminal detail=%d %s", detail.Code, detail.Body.String())
|
|
}
|
|
if response := clientB.do(http.MethodGet, fmt.Sprintf("/api/generations/%d/outputs/%d", imageID, output.ID), nil, ""); response.Code != 404 {
|
|
t.Fatalf("cross-user output=%d", response.Code)
|
|
}
|
|
if response := clientB.do(http.MethodGet, fmt.Sprintf("/generations/%d", imageID), nil, ""); response.Code != 404 {
|
|
t.Fatalf("cross-user page=%d", response.Code)
|
|
}
|
|
if response := clientA.do(http.MethodGet, fmt.Sprintf("/api/generations/%d/outputs/%d/thumbnail", imageID, output.ID), nil, ""); response.Code != 200 {
|
|
t.Fatalf("own thumbnail=%d", response.Code)
|
|
}
|
|
|
|
anonymous := &apiClient{router: router}
|
|
anonymous.start(t)
|
|
request := httptest.NewRequest(http.MethodGet, "/api/generations/1/card", nil)
|
|
request.Header.Set("HX-Request", "true")
|
|
request.AddCookie(anonymous.cookie)
|
|
response := httptest.NewRecorder()
|
|
router.ServeHTTP(response, request)
|
|
if response.Code != 401 || response.Header().Get("HX-Redirect") != "/login" {
|
|
t.Fatalf("expired HTMX=%d redirect=%q", response.Code, response.Header().Get("HX-Redirect"))
|
|
}
|
|
fragmentRequest := httptest.NewRequest(http.MethodGet, fmt.Sprintf("/ui/generations/%d/result", imageID), nil)
|
|
fragmentRequest.Header.Set("HX-Request", "true")
|
|
fragmentRequest.AddCookie(anonymous.cookie)
|
|
fragmentResponse := httptest.NewRecorder()
|
|
router.ServeHTTP(fragmentResponse, fragmentRequest)
|
|
if fragmentResponse.Code != 401 || !strings.Contains(fragmentResponse.Header().Get("HX-Trigger"), "authExpired") || fragmentResponse.Header().Get("HX-Redirect") != "" {
|
|
t.Fatalf("expired result fragment=%d trigger=%q redirect=%q", fragmentResponse.Code, fragmentResponse.Header().Get("HX-Trigger"), fragmentResponse.Header().Get("HX-Redirect"))
|
|
}
|
|
logout := clientA.do(http.MethodPost, "/api/session/logout", nil, "")
|
|
if logout.Code != 200 {
|
|
t.Fatalf("logout=%d %s", logout.Code, logout.Body.String())
|
|
}
|
|
if response := clientA.do(http.MethodGet, "/api/generations", nil, ""); response.Code != 401 {
|
|
t.Fatalf("post-logout access=%d", response.Code)
|
|
}
|
|
}
|
|
|
|
type uploadPart struct {
|
|
clientID, name, mime, note string
|
|
data []byte
|
|
position uint32
|
|
role model.InputRole
|
|
}
|
|
|
|
var lastMultipartType string
|
|
|
|
func multipartBody(t *testing.T, key, prompt string, capability model.Capability, roleRule string, parts []uploadPart) []byte {
|
|
t.Helper()
|
|
var body bytes.Buffer
|
|
writer := multipart.NewWriter(&body)
|
|
_ = writer.WriteField("idempotency_key", key)
|
|
_ = writer.WriteField("prompt", prompt)
|
|
metadata := service.ImageMetadata{Capability: capability, RoleRule: roleRule, Images: make([]service.ImageInput, 0, len(parts))}
|
|
for _, item := range parts {
|
|
metadata.Images = append(metadata.Images, service.ImageInput{ClientID: item.clientID, Position: item.position, Role: item.role, Note: item.note})
|
|
}
|
|
encodedMetadata, _ := json.Marshal(metadata)
|
|
_ = writer.WriteField("metadata", string(encodedMetadata))
|
|
for _, item := range parts {
|
|
header := make(textproto.MIMEHeader)
|
|
header.Set("Content-Disposition", fmt.Sprintf(`form-data; name="files[%s]"; filename="%s"`, item.clientID, item.name))
|
|
header.Set("Content-Type", item.mime)
|
|
part, _ := writer.CreatePart(header)
|
|
_, _ = part.Write(item.data)
|
|
}
|
|
writer.Close()
|
|
lastMultipartType = writer.FormDataContentType()
|
|
return body.Bytes()
|
|
}
|
|
func responseID(t *testing.T, response *httptest.ResponseRecorder) uint64 {
|
|
t.Helper()
|
|
var body struct {
|
|
ID uint64 `json:"id"`
|
|
}
|
|
if json.Unmarshal(response.Body.Bytes(), &body) != nil || body.ID == 0 {
|
|
t.Fatalf("response id missing: %s", response.Body.String())
|
|
}
|
|
return body.ID
|
|
}
|
|
func testPNG(width, height int) []byte {
|
|
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
|
for y := 0; y < height; y++ {
|
|
for x := 0; x < width; x++ {
|
|
img.Set(x, y, color.RGBA{R: 20, G: 100, B: 180, A: 255})
|
|
}
|
|
}
|
|
var body bytes.Buffer
|
|
_ = png.Encode(&body, img)
|
|
return body.Bytes()
|
|
}
|
|
|
|
func createPortalRoutes(t *testing.T, db *gorm.DB, suffix string, templates []model.PromptTemplate) ([]uint64, uint64) {
|
|
t.Helper()
|
|
tx := db.Begin()
|
|
if tx.Error != nil {
|
|
t.Fatal(tx.Error)
|
|
}
|
|
fail := func(err error) {
|
|
t.Helper()
|
|
tx.Rollback()
|
|
t.Fatal(err)
|
|
}
|
|
slug := "portal-provider-" + suffix
|
|
if err := tx.Exec("INSERT INTO providers (slug,name,base_url,auth_type,enabled) VALUES (?,?,?,'none',TRUE)", slug, "Portal Test", "https://example.invalid").Error; err != nil {
|
|
fail(err)
|
|
}
|
|
var providerID uint64
|
|
if err := tx.Raw("SELECT id FROM providers WHERE slug=?", slug).Scan(&providerID).Error; err != nil || providerID == 0 {
|
|
fail(fmt.Errorf("load provider id: %w", err))
|
|
}
|
|
routeIDs := make([]uint64, 0, len(templates))
|
|
for index, promptTemplate := range templates {
|
|
modelIDValue := fmt.Sprintf("portal-model-%d-%s", index, suffix)
|
|
if err := tx.Exec("INSERT INTO provider_models (provider_id,name,model_id,api_type,kind,extra_body,timeout_ms,weight,enabled) VALUES (?,?,?,?,?,JSON_OBJECT(),1000,100,TRUE)", providerID, "Portal Model", modelIDValue, promptTemplate.APIType, promptTemplate.Kind).Error; err != nil {
|
|
fail(err)
|
|
}
|
|
var providerModelID uint64
|
|
if err := tx.Raw("SELECT id FROM provider_models WHERE provider_id=? AND model_id=? AND api_type=?", providerID, modelIDValue, promptTemplate.APIType).Scan(&providerModelID).Error; err != nil || providerModelID == 0 {
|
|
fail(fmt.Errorf("load provider model id: %w", err))
|
|
}
|
|
if err := tx.Exec("INSERT INTO provider_model_capabilities (provider_model_id,capability) VALUES (?,?)", providerModelID, promptTemplate.Capability).Error; err != nil {
|
|
fail(err)
|
|
}
|
|
routeSlug := fmt.Sprintf("portal-route-%d-%s", index, suffix)
|
|
if err := tx.Exec("INSERT INTO route_pools (slug,name,capability,prompt_template_id,max_failover,version,enabled) VALUES (?,?,?,?,0,1,TRUE)", routeSlug, "Portal Route", promptTemplate.Capability, promptTemplate.ID).Error; err != nil {
|
|
fail(err)
|
|
}
|
|
var routeID uint64
|
|
if err := tx.Raw("SELECT id FROM route_pools WHERE slug=?", routeSlug).Scan(&routeID).Error; err != nil || routeID == 0 {
|
|
fail(fmt.Errorf("load route id: %w", err))
|
|
}
|
|
routeIDs = append(routeIDs, routeID)
|
|
if err := tx.Exec("INSERT INTO route_pool_members (route_pool_id,provider_model_id,weight,failure_threshold,open_seconds,half_open_max,enabled,position) VALUES (?,?,100,3,60,1,TRUE,0)", routeID, providerModelID).Error; err != nil {
|
|
fail(err)
|
|
}
|
|
if err := tx.Exec("INSERT INTO active_routes (capability,route_pool_id) VALUES (?,?)", promptTemplate.Capability, routeID).Error; err != nil {
|
|
fail(err)
|
|
}
|
|
}
|
|
if err := tx.Commit().Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return routeIDs, providerID
|
|
}
|