Files
chorus/portal/handler/mysql_integration_test.go
T

379 lines
19 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"
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)
}
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-image-" + suffix, Kind: model.KindImage, APIType: model.APIImagesEdits, Capability: model.CapabilityImageEdit, Name: "Portal Image", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true},
}
if err := db.Create(&templates).Error; err != nil {
t.Fatal(err)
}
defer func() {
db.Where("user_id IN ?", []uint64{users[0].ID, users[1].ID}).Delete(&model.Generation{})
db.Where("id IN ?", []uint64{templates[0].ID, templates[1].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)
}
generationService, err := service.New(db, queueRepository, local, service.Config{MaxPromptBytes: 1000, MaxImages: 3, MaxImageBytes: 200 << 10, MaxUploadBytes: 1 << 20, MaxImagePixels: 10000, HistoryLimit: 20})
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 {
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())
}
imageKey := "image-key-" + suffix
imageResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", []uploadPart{{"first.png", "image/png", testPNG(4, 4)}}), 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 {
t.Fatalf("inputs=%#v err=%v", inputs, err)
}
imageReplay := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", []uploadPart{{"second.png", "image/png", testPNG(4, 4)}}), lastMultipartType)
if imageReplay.Code != 200 || responseID(t, imageReplay) != imageID {
t.Fatalf("image replay=%d %s", imageReplay.Code, imageReplay.Body.String())
}
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", []uploadPart{{"fake.txt", "image/png", testPNG(4, 4)}}), 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())
}
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, 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 {
name, mime string
data []byte
}
var lastMultipartType string
func multipartBody(t *testing.T, key, prompt string, parts []uploadPart) []byte {
t.Helper()
var body bytes.Buffer
writer := multipart.NewWriter(&body)
_ = writer.WriteField("idempotency_key", key)
_ = writer.WriteField("prompt", prompt)
for _, item := range parts {
header := make(textproto.MIMEHeader)
header.Set("Content-Disposition", fmt.Sprintf(`form-data; name="images"; filename="%s"`, 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()
}