865 lines
53 KiB
Go
865 lines
53 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"
|
|
platformapikey "git.ilapage.cn/OPC/chorus/internal/platform/apikey"
|
|
passwordpkg "git.ilapage.cn/OPC/chorus/internal/platform/password"
|
|
"git.ilapage.cn/OPC/chorus/internal/platform/ratelimit"
|
|
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, account, password string) *httptest.ResponseRecorder {
|
|
t.Helper()
|
|
body, _ := json.Marshal(map[string]string{"account": account, "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
|
|
expectedDatabase := strings.TrimSpace(os.Getenv("CHORUS_MIGRATION_TEST_DATABASE"))
|
|
if expectedDatabase == "" {
|
|
expectedDatabase = "chorus_test"
|
|
}
|
|
if err := db.Raw("SELECT DATABASE()").Scan(&databaseName).Error; err != nil || databaseName != expectedDatabase {
|
|
t.Fatalf("portal integration requires %s database, got %q: %v", expectedDatabase, 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{{Username: "portal-a-" + suffix, Email: "portal-a-" + suffix + "@example.invalid", PasswordHash: encoded, DisplayName: "User A", Status: "active"}, {Username: "portal-b-" + suffix, 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.APIAuditEvent{})
|
|
db.Where("user_id IN ?", []uint64{users[0].ID, users[1].ID}).Delete(&model.APIKey{})
|
|
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, RateLimitConfig{
|
|
Limiter: ratelimit.New(),
|
|
UserPolicy: ratelimit.Policy{Capacity: 10_000, Window: time.Minute},
|
|
APIKeyPolicy: ratelimit.Policy{Capacity: 10_000, Window: time.Minute},
|
|
})
|
|
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())
|
|
}
|
|
legacy := &apiClient{router: router}
|
|
legacy.start(t)
|
|
legacyBody, _ := json.Marshal(map[string]string{"email": users[0].Email, "password": password})
|
|
if response := legacy.do(http.MethodPost, "/api/session/login", legacyBody, "application/json"); response.Code != http.StatusBadRequest {
|
|
t.Fatalf("legacy email contract=%d %s", response.Code, response.Body.String())
|
|
}
|
|
|
|
unknown := &apiClient{router: router}
|
|
unknown.start(t)
|
|
unknownResponse := unknown.login(t, "missing-"+suffix, "wrong-password-value")
|
|
wrong := &apiClient{router: router}
|
|
wrong.start(t)
|
|
wrongResponse := wrong.login(t, users[0].Username, "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)
|
|
if response := clientA.login(t, users[0].Email, password); response.Code != http.StatusUnauthorized {
|
|
t.Fatalf("email login=%d %s", response.Code, response.Body.String())
|
|
}
|
|
clientA = &apiClient{router: router}
|
|
clientA.start(t)
|
|
oldCookie := *clientA.cookie
|
|
if response := clientA.login(t, users[0].Username, 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(), "文生图") || !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())
|
|
}
|
|
apiKeyPage := clientA.do(http.MethodGet, "/api-keys", nil, "")
|
|
if apiKeyPage.Code != http.StatusOK || !strings.Contains(apiKeyPage.Body.String(), "完整 Key 只在创建成功后显示一次") || !strings.Contains(apiKeyPage.Body.String(), `/static/api-keys.js`) {
|
|
t.Fatalf("API key page=%d %s", apiKeyPage.Code, apiKeyPage.Body.String())
|
|
}
|
|
createKeyBody, _ := json.Marshal(map[string]any{"name": "自动化脚本", "expires_in_days": 90})
|
|
validCSRF := clientA.csrf
|
|
clientA.csrf = ""
|
|
if response := clientA.do(http.MethodPost, "/api/api-keys", createKeyBody, "application/json"); response.Code != http.StatusForbidden || response.Header().Get("Cache-Control") != "no-store" {
|
|
t.Fatalf("API key CSRF response=%d cache=%q %s", response.Code, response.Header().Get("Cache-Control"), response.Body.String())
|
|
}
|
|
clientA.csrf = validCSRF
|
|
createKeyResponse := clientA.do(http.MethodPost, "/api/api-keys", createKeyBody, "application/json")
|
|
var createdKey struct {
|
|
ID uint64 `json:"id"`
|
|
Name string `json:"name"`
|
|
KeyPrefix string `json:"key_prefix"`
|
|
Token string `json:"token"`
|
|
ExpiresAt *time.Time `json:"expires_at"`
|
|
}
|
|
if createKeyResponse.Code != http.StatusCreated || createKeyResponse.Header().Get("Cache-Control") != "no-store" || createKeyResponse.Header().Get("Pragma") != "no-cache" || json.Unmarshal(createKeyResponse.Body.Bytes(), &createdKey) != nil || createdKey.ID == 0 || createdKey.Token == "" || !strings.HasPrefix(createdKey.Token, createdKey.KeyPrefix) || createdKey.ExpiresAt == nil {
|
|
t.Fatalf("create API key=%d cache=%q body=%s", createKeyResponse.Code, createKeyResponse.Header().Get("Cache-Control"), createKeyResponse.Body.String())
|
|
}
|
|
var storedKey model.APIKey
|
|
if err := db.First(&storedKey, createdKey.ID).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var createAudit model.APIAuditEvent
|
|
if err := db.Where("api_key_id = ? AND action = ?", createdKey.ID, "api_key.created").Take(&createAudit).Error; err != nil || strings.Contains(string(createAudit.Summary), createdKey.Token) || strings.Contains(string(createAudit.Summary), storedKey.PublicID) {
|
|
t.Fatalf("created key audit is missing or sensitive: %#v %v", createAudit, err)
|
|
}
|
|
parsedKey, err := platformapikey.Parse(createdKey.Token)
|
|
if err != nil || parsedKey.PublicID != storedKey.PublicID || !bytes.Equal(parsedKey.SecretHash, storedKey.SecretHash) {
|
|
t.Fatalf("stored API key does not match one-time token: %v", err)
|
|
}
|
|
openAPIKeyA := storeTestAPIKey(t, db, users[0].ID, "OpenAPI A", nil, nil)
|
|
openAPIKeyB := storeTestAPIKey(t, db, users[1].ID, "OpenAPI B", nil, nil)
|
|
var openAPIKeyARow model.APIKey
|
|
if err := db.Where("public_id = ?", openAPIKeyA.PublicID).Take(&openAPIKeyARow).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
expiredAt := time.Now().UTC().Add(-time.Minute)
|
|
expiredKey := storeTestAPIKey(t, db, users[0].ID, "Expired OpenAPI", &expiredAt, nil)
|
|
var auditCountBeforeInvalid int64
|
|
if err := db.Model(&model.APIAuditEvent{}).Where("user_id IN ?", []uint64{users[0].ID, users[1].ID}).Count(&auditCountBeforeInvalid).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
missingAuth := openAPIRequest(router, "", http.MethodGet, "/openapi/v1/openapi.json", nil, "", "")
|
|
expiredAuth := openAPIRequest(router, expiredKey.Token, http.MethodGet, "/openapi/v1/openapi.json", nil, "", "")
|
|
invalidAuth := openAPIRequest(router, "administrator-jwt", http.MethodGet, "/openapi/v1/openapi.json", nil, "", "")
|
|
queryAuth := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/openapi.json?Authorization=ignored", nil, "", "")
|
|
for name, response := range map[string]*httptest.ResponseRecorder{"missing": missingAuth, "expired": expiredAuth, "invalid": invalidAuth, "query": queryAuth} {
|
|
if response.Code != http.StatusUnauthorized || response.Body.String() != missingAuth.Body.String() || response.Header().Get("WWW-Authenticate") != "Bearer" || response.Header().Get("Cache-Control") != "no-store" || response.Header().Get("X-Request-ID") == "" || response.Header().Get("Set-Cookie") != "" {
|
|
t.Fatalf("%s OpenAPI auth=%d headers=%v body=%s", name, response.Code, response.Header(), response.Body.String())
|
|
}
|
|
}
|
|
var auditCountAfterInvalid int64
|
|
if err := db.Model(&model.APIAuditEvent{}).Where("user_id IN ?", []uint64{users[0].ID, users[1].ID}).Count(&auditCountAfterInvalid).Error; err != nil || auditCountAfterInvalid != auditCountBeforeInvalid {
|
|
t.Fatalf("unknown or unusable API keys wrote database audit rows: before=%d after=%d err=%v", auditCountBeforeInvalid, auditCountAfterInvalid, err)
|
|
}
|
|
auditHandler := &Handler{
|
|
service: generationService, rateLimiter: ratelimit.New(),
|
|
userRatePolicy: ratelimit.Policy{Capacity: 10, Window: time.Minute},
|
|
apiKeyRatePolicy: ratelimit.Policy{Capacity: 1, Window: time.Minute},
|
|
now: func() time.Time { return time.Now().UTC() },
|
|
}
|
|
auditRouter := gin.New()
|
|
auditRouter.Use(auditHandler.requestID, func(c *gin.Context) {
|
|
c.Set(openAPIPrincipalContextKey, service.APIPrincipal{UserID: users[0].ID, APIKeyID: openAPIKeyARow.ID})
|
|
c.Next()
|
|
})
|
|
auditRouter.POST("/openapi/v1/generations/text", auditHandler.auditOpenAPISubmission, auditHandler.limitOpenAPIRequest, func(c *gin.Context) { c.Status(http.StatusAccepted) })
|
|
if response := rateRequest(auditRouter, http.MethodPost, "/openapi/v1/generations/text"); response.Code != http.StatusAccepted {
|
|
t.Fatalf("audit rate first request = %d", response.Code)
|
|
}
|
|
if response := rateRequest(auditRouter, http.MethodPost, "/openapi/v1/generations/text"); response.Code != http.StatusTooManyRequests {
|
|
t.Fatalf("audit rate limited request = %d", response.Code)
|
|
}
|
|
var rateAudit model.APIAuditEvent
|
|
if err := db.Where("api_key_id = ? AND action = ? AND status_code = ?", openAPIKeyARow.ID, "openapi.generation.submit", http.StatusTooManyRequests).Order("id DESC").Take(&rateAudit).Error; err != nil || rateAudit.Result != "denied" || rateAudit.ErrorCode == nil || *rateAudit.ErrorCode != "rate_limited" || strings.Contains(string(rateAudit.Summary), openAPIKeyA.Token) {
|
|
t.Fatalf("authenticated rate-limit audit is missing or sensitive: %#v %v", rateAudit, err)
|
|
}
|
|
browserCookieRequest := httptest.NewRequest(http.MethodGet, "/openapi/v1/openapi.json", nil)
|
|
browserCookieRequest.Header.Set("Authorization", "Bearer "+openAPIKeyA.Token)
|
|
browserCookieRequest.AddCookie(clientA.cookie)
|
|
browserCookieResponse := httptest.NewRecorder()
|
|
router.ServeHTTP(browserCookieResponse, browserCookieRequest)
|
|
if browserCookieResponse.Code != http.StatusUnauthorized {
|
|
t.Fatalf("OpenAPI accepted browser cookie=%d", browserCookieResponse.Code)
|
|
}
|
|
specResponse := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/openapi.json", nil, "", "")
|
|
var specification struct {
|
|
OpenAPI string `json:"openapi"`
|
|
}
|
|
if specResponse.Code != http.StatusOK || json.Unmarshal(specResponse.Body.Bytes(), &specification) != nil || specification.OpenAPI != "3.1.0" || specResponse.Header().Get("Set-Cookie") != "" || specResponse.Header().Get("X-Request-ID") == "" {
|
|
t.Fatalf("OpenAPI specification=%d headers=%v body=%s", specResponse.Code, specResponse.Header(), specResponse.Body.String())
|
|
}
|
|
var usedKey model.APIKey
|
|
if err := db.Where("public_id = ?", openAPIKeyA.PublicID).First(&usedKey).Error; err != nil || usedKey.LastUsedAt == nil {
|
|
t.Fatalf("OpenAPI key last_used_at was not updated: %#v %v", usedKey, err)
|
|
}
|
|
recentUse := time.Now().UTC().Truncate(time.Microsecond).Add(-10 * time.Second)
|
|
if err := db.Model(&model.APIKey{}).Where("id = ?", usedKey.ID).Update("last_used_at", recentUse).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/openapi.json", nil, "", ""); response.Code != http.StatusOK {
|
|
t.Fatalf("recent API key reuse=%d", response.Code)
|
|
}
|
|
var throttledUse model.APIKey
|
|
if err := db.First(&throttledUse, usedKey.ID).Error; err != nil || throttledUse.LastUsedAt == nil || !throttledUse.LastUsedAt.Equal(recentUse) {
|
|
t.Fatalf("last_used_at was not throttled: %#v %v", throttledUse.LastUsedAt, err)
|
|
}
|
|
staleUse := time.Now().UTC().Truncate(time.Microsecond).Add(-2 * time.Minute)
|
|
if err := db.Model(&model.APIKey{}).Where("id = ?", usedKey.ID).Update("last_used_at", staleUse).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/openapi.json", nil, "", ""); response.Code != http.StatusOK {
|
|
t.Fatalf("stale API key reuse=%d", response.Code)
|
|
}
|
|
var refreshedUse model.APIKey
|
|
if err := db.First(&refreshedUse, usedKey.ID).Error; err != nil || refreshedUse.LastUsedAt == nil || !refreshedUse.LastUsedAt.After(staleUse) {
|
|
t.Fatalf("stale last_used_at was not refreshed: %#v %v", refreshedUse.LastUsedAt, err)
|
|
}
|
|
browserOnly := clientA.do(http.MethodGet, "/openapi/v1/generations", nil, "")
|
|
if browserOnly.Code != http.StatusUnauthorized {
|
|
t.Fatalf("OpenAPI accepted browser session=%d", browserOnly.Code)
|
|
}
|
|
apiKeyOnBrowserRoute := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/api/generations", nil, "", "")
|
|
if apiKeyOnBrowserRoute.Code != http.StatusUnauthorized {
|
|
t.Fatalf("browser route accepted API key=%d", apiKeyOnBrowserRoute.Code)
|
|
}
|
|
|
|
openAPITextKey := "openapi-text-" + suffix
|
|
openAPITextBody, _ := json.Marshal(map[string]string{"prompt": "OpenAPI text prompt"})
|
|
openAPIText := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", openAPITextBody, "application/json", openAPITextKey)
|
|
if openAPIText.Code != http.StatusAccepted {
|
|
t.Fatalf("OpenAPI text submit=%d %s", openAPIText.Code, openAPIText.Body.String())
|
|
}
|
|
openAPITextID := responseID(t, openAPIText)
|
|
var submitAudit model.APIAuditEvent
|
|
if err := db.Where("api_key_id = ? AND generation_id = ? AND action = ?", openAPIKeyARow.ID, openAPITextID, "openapi.generation.submit").Take(&submitAudit).Error; err != nil || submitAudit.Result != "succeeded" || strings.Contains(string(submitAudit.Summary), "OpenAPI text prompt") || strings.Contains(string(submitAudit.Summary), openAPIKeyA.Token) {
|
|
t.Fatalf("OpenAPI submit audit is missing or sensitive: %#v %v", submitAudit, err)
|
|
}
|
|
openAPITextReplay := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", openAPITextBody, "application/json", openAPITextKey)
|
|
if openAPITextReplay.Code != http.StatusOK || responseID(t, openAPITextReplay) != openAPITextID {
|
|
t.Fatalf("OpenAPI text replay=%d %s", openAPITextReplay.Code, openAPITextReplay.Body.String())
|
|
}
|
|
conflictBody, _ := json.Marshal(map[string]string{"prompt": "different OpenAPI prompt"})
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", conflictBody, "application/json", openAPITextKey); response.Code != http.StatusConflict || !strings.Contains(response.Body.String(), `"code":"idempotency_conflict"`) {
|
|
t.Fatalf("OpenAPI text conflict=%d %s", response.Code, response.Body.String())
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", openAPITextBody, "application/json", ""); response.Code != http.StatusBadRequest {
|
|
t.Fatalf("missing Idempotency-Key=%d %s", response.Code, response.Body.String())
|
|
}
|
|
unknownFieldBody, _ := json.Marshal(map[string]string{"prompt": "prompt", "idempotency_key": "body-key"})
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", unknownFieldBody, "application/json", "header-key-123"); response.Code != http.StatusBadRequest {
|
|
t.Fatalf("body idempotency key was accepted=%d %s", response.Code, response.Body.String())
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", openAPITextBody, "application/json", "非ASCII幂等键"); response.Code != http.StatusBadRequest {
|
|
t.Fatalf("non-ASCII Idempotency-Key=%d %s", response.Code, response.Body.String())
|
|
}
|
|
openAPICrossChannelBody, _ := json.Marshal(map[string]string{"idempotency_key": openAPITextKey, "prompt": "OpenAPI text prompt"})
|
|
if response := clientA.do(http.MethodPost, "/api/generations/text", openAPICrossChannelBody, "application/json"); response.Code != http.StatusOK || responseID(t, response) != openAPITextID {
|
|
t.Fatalf("cross-channel idempotent replay=%d %s", response.Code, response.Body.String())
|
|
}
|
|
var openAPIGeneration model.Generation
|
|
if err := db.First(&openAPIGeneration, openAPITextID).Error; err != nil || openAPIGeneration.Status != model.StatusPending || openAPIGeneration.AttemptCount != 0 || openAPIGeneration.ProviderModelID != nil {
|
|
t.Fatalf("OpenAPI synchronous submit touched worker fields: %#v %v", openAPIGeneration, err)
|
|
}
|
|
listKeys := clientA.do(http.MethodGet, "/api/api-keys", nil, "")
|
|
if listKeys.Code != http.StatusOK || !strings.Contains(listKeys.Body.String(), "自动化脚本") || strings.Contains(listKeys.Body.String(), createdKey.Token) || strings.Contains(listKeys.Body.String(), `secret_hash`) {
|
|
t.Fatalf("list API keys=%d %s", listKeys.Code, listKeys.Body.String())
|
|
}
|
|
if refreshedPage := clientA.do(http.MethodGet, "/api-keys", nil, ""); strings.Contains(refreshedPage.Body.String(), createdKey.Token) {
|
|
t.Fatal("API key page recovered the complete token")
|
|
}
|
|
invalidExpiry, _ := json.Marshal(map[string]any{"name": "Invalid expiry", "expires_in_days": 365})
|
|
if response := clientA.do(http.MethodPost, "/api/api-keys", invalidExpiry, "application/json"); response.Code != http.StatusBadRequest || strings.Contains(response.Body.String(), createdKey.Token) {
|
|
t.Fatalf("invalid API key expiry=%d %s", response.Code, response.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
|
|
imageParts := []uploadPart{{clientID: "primary-a", name: "first.png", mime: "image/png", data: testPNG(4, 4), position: 0, role: model.RolePrimary, note: "keep edges"}}
|
|
imageResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", model.CapabilityImageEdit, "user role rule", imageParts), 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, "user role rule", imageParts), lastMultipartType)
|
|
if imageReplay.Code != 200 || responseID(t, imageReplay) != imageID {
|
|
t.Fatalf("image replay=%d %s", imageReplay.Code, imageReplay.Body.String())
|
|
}
|
|
imageConflict := 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 imageConflict.Code != http.StatusConflict || !strings.Contains(imageConflict.Body.String(), `"code":"idempotency_conflict"`) {
|
|
t.Fatalf("image conflict=%d %s", imageConflict.Code, imageConflict.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)
|
|
}
|
|
openAPIImageKey := "openapi-image-" + suffix
|
|
openAPIImageParts := []uploadPart{{clientID: "openapi-primary", name: "openapi.png", mime: "image/png", data: testPNG(3, 3), position: 0, role: model.RolePrimary, note: "preserve subject"}}
|
|
openAPIImageBody, openAPIImageType := openAPIMultipartBody(t, "OpenAPI image edit", model.CapabilityImageEdit, "OpenAPI role rule", openAPIImageParts)
|
|
openAPIImage := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/image", openAPIImageBody, openAPIImageType, openAPIImageKey)
|
|
if openAPIImage.Code != http.StatusAccepted {
|
|
t.Fatalf("OpenAPI image submit=%d %s", openAPIImage.Code, openAPIImage.Body.String())
|
|
}
|
|
openAPIImageID := responseID(t, openAPIImage)
|
|
openAPIImageBody, openAPIImageType = openAPIMultipartBody(t, "OpenAPI image edit", model.CapabilityImageEdit, "OpenAPI role rule", openAPIImageParts)
|
|
openAPIImageReplay := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/image", openAPIImageBody, openAPIImageType, openAPIImageKey)
|
|
if openAPIImageReplay.Code != http.StatusOK || responseID(t, openAPIImageReplay) != openAPIImageID {
|
|
t.Fatalf("OpenAPI image replay=%d %s", openAPIImageReplay.Code, openAPIImageReplay.Body.String())
|
|
}
|
|
changedOpenAPIImageParts := append([]uploadPart(nil), openAPIImageParts...)
|
|
changedOpenAPIImageParts[0].data = testPNG(5, 5)
|
|
openAPIImageBody, openAPIImageType = openAPIMultipartBody(t, "OpenAPI image edit", model.CapabilityImageEdit, "OpenAPI role rule", changedOpenAPIImageParts)
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/image", openAPIImageBody, openAPIImageType, openAPIImageKey); response.Code != http.StatusConflict || !strings.Contains(response.Body.String(), `"code":"idempotency_conflict"`) {
|
|
t.Fatalf("OpenAPI image conflict=%d %s", response.Code, response.Body.String())
|
|
}
|
|
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())
|
|
}
|
|
openAPIHistory := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/generations?limit=2", nil, "", "")
|
|
var openAPIHistoryPage struct {
|
|
Items []struct {
|
|
ID uint64 `json:"id"`
|
|
} `json:"items"`
|
|
NextCursor string `json:"next_cursor"`
|
|
HasMore bool `json:"has_more"`
|
|
}
|
|
if openAPIHistory.Code != http.StatusOK || json.Unmarshal(openAPIHistory.Body.Bytes(), &openAPIHistoryPage) != nil || len(openAPIHistoryPage.Items) != 2 || !openAPIHistoryPage.HasMore || openAPIHistoryPage.NextCursor == "" {
|
|
t.Fatalf("OpenAPI history=%d %s", openAPIHistory.Code, openAPIHistory.Body.String())
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyB.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d", openAPITextID), nil, "", ""); response.Code != http.StatusNotFound {
|
|
t.Fatalf("cross-user OpenAPI detail=%d %s", response.Code, response.Body.String())
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyB.Token, http.MethodGet, "/openapi/v1/generations?limit=2&cursor="+openAPIHistoryPage.NextCursor, nil, "", ""); response.Code != http.StatusBadRequest || !strings.Contains(response.Body.String(), `"code":"invalid_cursor"`) {
|
|
t.Fatalf("cross-user OpenAPI cursor=%d %s", response.Code, response.Body.String())
|
|
}
|
|
|
|
clientB := &apiClient{router: router}
|
|
clientB.start(t)
|
|
if response := clientB.login(t, users[1].Username, 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)
|
|
}
|
|
renameKeyBody, _ := json.Marshal(map[string]string{"name": "其他用户不能改名"})
|
|
if response := clientB.do(http.MethodPatch, fmt.Sprintf("/api/api-keys/%d", createdKey.ID), renameKeyBody, "application/json"); response.Code != http.StatusNotFound {
|
|
t.Fatalf("cross-user API key rename=%d %s", response.Code, response.Body.String())
|
|
}
|
|
if response := clientB.do(http.MethodDelete, fmt.Sprintf("/api/api-keys/%d", createdKey.ID), nil, ""); response.Code != http.StatusNotFound {
|
|
t.Fatalf("cross-user API key revoke=%d %s", response.Code, response.Body.String())
|
|
}
|
|
renameKeyBody, _ = json.Marshal(map[string]string{"name": "已改名脚本"})
|
|
if response := clientA.do(http.MethodPatch, fmt.Sprintf("/api/api-keys/%d", createdKey.ID), renameKeyBody, "application/json"); response.Code != http.StatusOK || !strings.Contains(response.Body.String(), "已改名脚本") || strings.Contains(response.Body.String(), createdKey.Token) {
|
|
t.Fatalf("rename own API key=%d %s", response.Code, response.Body.String())
|
|
}
|
|
if err := db.Model(&model.User{}).Where("id = ?", users[0].ID).Update("status", "disabled").Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if response := clientA.do(http.MethodGet, "/api/api-keys", nil, ""); response.Code != http.StatusForbidden || !strings.Contains(response.Body.String(), `"code":"account_disabled"`) {
|
|
t.Fatalf("disabled user API keys=%d %s", response.Code, response.Body.String())
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/generations", nil, "", ""); response.Code != http.StatusUnauthorized || !strings.Contains(response.Body.String(), `"code":"invalid_api_key"`) {
|
|
t.Fatalf("disabled user OpenAPI=%d %s", response.Code, response.Body.String())
|
|
}
|
|
if err := db.Model(&model.User{}).Where("id = ?", users[0].ID).Update("status", "active").Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
firstRevoke := clientA.do(http.MethodDelete, fmt.Sprintf("/api/api-keys/%d", createdKey.ID), nil, "")
|
|
secondRevoke := clientA.do(http.MethodDelete, fmt.Sprintf("/api/api-keys/%d", createdKey.ID), nil, "")
|
|
if firstRevoke.Code != http.StatusOK || secondRevoke.Code != http.StatusOK || firstRevoke.Body.String() != secondRevoke.Body.String() || !strings.Contains(firstRevoke.Body.String(), `"status":"revoked"`) || strings.Contains(firstRevoke.Body.String(), createdKey.Token) {
|
|
t.Fatalf("idempotent API key revoke=%d/%d %s / %s", firstRevoke.Code, secondRevoke.Code, firstRevoke.Body.String(), secondRevoke.Body.String())
|
|
}
|
|
var lifecycleAudits int64
|
|
if err := db.Model(&model.APIAuditEvent{}).Where("api_key_id = ? AND action IN ?", createdKey.ID, []string{"api_key.created", "api_key.renamed", "api_key.revoked"}).Count(&lifecycleAudits).Error; err != nil || lifecycleAudits != 4 {
|
|
t.Fatalf("API key lifecycle audits = %d, want 4: %v", lifecycleAudits, err)
|
|
}
|
|
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)
|
|
}
|
|
openAPIDetail := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d", imageID), nil, "", "")
|
|
if openAPIDetail.Code != http.StatusOK || !strings.Contains(openAPIDetail.Body.String(), fmt.Sprintf("/openapi/v1/generations/%d/inputs/%d", imageID, inputs[0].ID)) || !strings.Contains(openAPIDetail.Body.String(), fmt.Sprintf("/openapi/v1/generations/%d/outputs/%d", imageID, output.ID)) || strings.Contains(openAPIDetail.Body.String(), storageRoot) || strings.Contains(openAPIDetail.Body.String(), original.Key) {
|
|
t.Fatalf("OpenAPI detail=%d %s", openAPIDetail.Code, openAPIDetail.Body.String())
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d/inputs/%d", imageID, inputs[0].ID), nil, "", ""); response.Code != http.StatusOK || response.Header().Get("Content-Type") != "image/png" {
|
|
t.Fatalf("OpenAPI input=%d %s", response.Code, response.Header().Get("Content-Type"))
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d/outputs/%d", imageID, output.ID), nil, "", ""); response.Code != http.StatusOK || response.Header().Get("Content-Type") != "image/png" {
|
|
t.Fatalf("OpenAPI output=%d %s", response.Code, response.Header().Get("Content-Type"))
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d/outputs/%d/thumbnail", imageID, output.ID), nil, "", ""); response.Code != http.StatusOK || response.Header().Get("Content-Type") != "image/png" {
|
|
t.Fatalf("OpenAPI thumbnail=%d %s", response.Code, response.Header().Get("Content-Type"))
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyB.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d/outputs/%d", imageID, output.ID), nil, "", ""); response.Code != http.StatusNotFound {
|
|
t.Fatalf("cross-user OpenAPI output=%d %s", response.Code, response.Body.String())
|
|
}
|
|
revokedAt := time.Now().UTC()
|
|
if err := db.Model(&model.APIKey{}).Where("public_id = ?", openAPIKeyA.PublicID).Update("revoked_at", revokedAt).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/generations", nil, "", ""); response.Code != http.StatusUnauthorized || response.Body.String() != missingAuth.Body.String() {
|
|
t.Fatalf("revoked OpenAPI key=%d %s", response.Code, response.Body.String())
|
|
}
|
|
|
|
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 openAPIRequest(router http.Handler, token, method, path string, body []byte, contentType, idempotencyKey string) *httptest.ResponseRecorder {
|
|
request := httptest.NewRequest(method, path, bytes.NewReader(body))
|
|
request.RemoteAddr = "192.0.2.20:1234"
|
|
if token != "" {
|
|
request.Header.Set("Authorization", "Bearer "+token)
|
|
}
|
|
if contentType != "" {
|
|
request.Header.Set("Content-Type", contentType)
|
|
}
|
|
if idempotencyKey != "" {
|
|
request.Header.Set("Idempotency-Key", idempotencyKey)
|
|
}
|
|
response := httptest.NewRecorder()
|
|
router.ServeHTTP(response, request)
|
|
return response
|
|
}
|
|
|
|
func storeTestAPIKey(t *testing.T, db *gorm.DB, userID uint64, name string, expiresAt, revokedAt *time.Time) platformapikey.Credential {
|
|
t.Helper()
|
|
credential, err := platformapikey.Generate()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
row := model.APIKey{UserID: userID, Name: name, PublicID: credential.PublicID, KeyPrefix: credential.KeyPrefix, SecretHash: credential.SecretHash, ExpiresAt: expiresAt, RevokedAt: revokedAt}
|
|
if err := db.Create(&row).Error; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return credential
|
|
}
|
|
|
|
func openAPIMultipartBody(t *testing.T, prompt string, capability model.Capability, roleRule string, parts []uploadPart) ([]byte, string) {
|
|
t.Helper()
|
|
var body bytes.Buffer
|
|
writer := multipart.NewWriter(&body)
|
|
_ = 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()
|
|
return body.Bytes(), writer.FormDataContentType()
|
|
}
|
|
|
|
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
|
|
}
|