Files
chorus/portal/handler/mysql_integration_test.go
T

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(), "&lt;script&gt;") {
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
}