倍率问题和图片问题调整
This commit is contained in:
1 parent
f1993eb388
commit
3f29ba716a
20 files changed
+1210
-24
No files matched your search
@@ -245,6 +245,7 @@ func (s *Service) Upload(ctx context.Context, scope Scope, cmd UploadCommand) (A
|
||||
if s.blobs == nil {
|
||||
return Asset{}, errors.New("blob store is unavailable")
|
||||
}
|
||||
cmd.FileName, cmd.ContentType = normalizeImageFile(cmd.FileName, cmd.ContentType, cmd.Bytes)
|
||||
name := sanitizeFileName(cmd.FileName)
|
||||
key := path.Join("uploads", s.now().UTC().Format("2006-01-02"), s.id("file")+"-"+name)
|
||||
stored, err := s.blobs.Put(ctx, key, bytes.NewReader(cmd.Bytes), int64(len(cmd.Bytes)), cmd.ContentType)
|
||||
@@ -296,6 +297,7 @@ func (s *Service) ImportGenerated(ctx context.Context, scope Scope, cmd ImportGe
|
||||
if strings.TrimSpace(name) == "" {
|
||||
name = path.Base(strings.SplitN(cmd.URL, "?", 2)[0])
|
||||
}
|
||||
name, blob.ContentType = normalizeImageFile(name, blob.ContentType, content)
|
||||
if s.blobs == nil {
|
||||
return Asset{}, errors.New("blob store is unavailable")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
package assets
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"path"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// normalizeImageFile uses the bytes, not a provider URL suffix or an untrusted
|
||||
// Content-Type, as evidence of a recognized raster format. It does not decode
|
||||
// or transcode the image. Unrecognized data retains its original metadata.
|
||||
func normalizeImageFile(name, contentType string, prefix []byte) (string, string) {
|
||||
if len(prefix) > 512 {
|
||||
prefix = prefix[:512]
|
||||
}
|
||||
detected := http.DetectContentType(prefix)
|
||||
var extension string
|
||||
switch detected {
|
||||
case "image/png":
|
||||
extension = ".png"
|
||||
case "image/jpeg":
|
||||
extension = ".jpg"
|
||||
case "image/webp":
|
||||
extension = ".webp"
|
||||
case "image/gif":
|
||||
extension = ".gif"
|
||||
case "image/bmp":
|
||||
extension = ".bmp"
|
||||
case "image/x-icon":
|
||||
extension = ".ico"
|
||||
default:
|
||||
return name, contentType
|
||||
}
|
||||
name = strings.TrimSpace(name)
|
||||
existing := path.Ext(name)
|
||||
lower := strings.ToLower(existing)
|
||||
if lower == extension || detected == "image/jpeg" && (lower == ".jpeg" || lower == ".jpe") {
|
||||
return name, detected
|
||||
}
|
||||
base := strings.TrimSuffix(name, existing)
|
||||
if base == "" || base == "." || base == "/" {
|
||||
base = "image"
|
||||
}
|
||||
return base + extension, detected
|
||||
}
|
||||
|
||||
// NormalizeImageDownload fixes legacy filenames and response MIME types from
|
||||
// a bounded prefix. The returned body streams every original byte, and Close
|
||||
// still closes the underlying OSS/local/remote stream. No catalog/object write
|
||||
// is performed, including for historical .image objects.
|
||||
func NormalizeImageDownload(name string, blob Blob) (string, Blob, error) {
|
||||
if blob.Body == nil {
|
||||
return name, blob, ErrBlobNotFound
|
||||
}
|
||||
reader := bufio.NewReaderSize(blob.Body, 512)
|
||||
prefix, err := reader.Peek(512)
|
||||
if err != nil && !errors.Is(err, io.EOF) {
|
||||
return name, blob, err
|
||||
}
|
||||
name, blob.ContentType = normalizeImageFile(name, blob.ContentType, prefix)
|
||||
blob.Body = &imageReadCloser{Reader: reader, Closer: blob.Body}
|
||||
return name, blob, nil
|
||||
}
|
||||
|
||||
type imageReadCloser struct {
|
||||
io.Reader
|
||||
io.Closer
|
||||
}
|
||||
@@ -0,0 +1,170 @@
|
||||
package assets
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"image"
|
||||
"image/jpeg"
|
||||
"image/png"
|
||||
"io"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func imageFormatFixtures(t *testing.T) (pngBytes, jpegBytes, webpBytes []byte) {
|
||||
t.Helper()
|
||||
img := image.NewRGBA(image.Rect(0, 0, 2, 2))
|
||||
var p, j bytes.Buffer
|
||||
if err := png.Encode(&p, img); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := jpeg.Encode(&j, img, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
w, err := base64.StdEncoding.DecodeString("UklGRiIAAABXRUJQVlA4IBYAAAAwAQCdASoBAAEADsD+JaQAA3AAAAAA")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return p.Bytes(), j.Bytes(), w
|
||||
}
|
||||
|
||||
func TestNormalizeImageFileUsesBytesAndPreservesCompatibleNames(t *testing.T) {
|
||||
pngData, jpegData, webpData := imageFormatFixtures(t)
|
||||
for _, test := range []struct {
|
||||
label, name, mime, wantName, wantMIME string
|
||||
data []byte
|
||||
}{
|
||||
{"opaque PNG", "result.image", "application/octet-stream", "result.png", "image/png", pngData},
|
||||
{"JPEG", "result.image", "", "result.jpg", "image/jpeg", jpegData},
|
||||
{"WebP despite bad MIME", "result.image", "image/png", "result.webp", "image/webp", webpData},
|
||||
{"wrong extension", "result.png", "image/png", "result.jpg", "image/jpeg", jpegData},
|
||||
{"no suffix", "result", "", "result.png", "image/png", pngData},
|
||||
{"uppercase", "result.PNG", "image/png; charset=binary", "result.PNG", "image/png", pngData},
|
||||
{"JPEG alias", "result.JPEG", "application/octet-stream", "result.JPEG", "image/jpeg", jpegData},
|
||||
{"Chinese filename", "葫芦娃和爷爷.image", "", "葫芦娃和爷爷.png", "image/png", pngData},
|
||||
{"multi-dot filename", "result.v1.image", "", "result.v1.png", "image/png", pngData},
|
||||
{"empty filename", "", "", "image.png", "image/png", pngData},
|
||||
{"video", "result.mp4", "video/mp4", "result.mp4", "video/mp4", []byte("video content")},
|
||||
{"unknown does not trust MIME", "result.image", "image/png", "result.image", "image/png", []byte(`{"error":"not an image"}`)},
|
||||
{"empty content", "result.image", "", "result.image", "", nil},
|
||||
} {
|
||||
t.Run(test.label, func(t *testing.T) {
|
||||
name, mime := normalizeImageFile(test.name, test.mime, test.data)
|
||||
if name != test.wantName || mime != test.wantMIME {
|
||||
t.Fatalf("got %q / %q, want %q / %q", name, mime, test.wantName, test.wantMIME)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportGeneratedNormalizesImageNameStorageAndContentType(t *testing.T) {
|
||||
pngData, jpegData, webpData := imageFormatFixtures(t)
|
||||
for _, test := range []struct {
|
||||
extension, contentType string
|
||||
data []byte
|
||||
}{{"png", "image/png", pngData}, {"jpg", "image/jpeg", jpegData}, {"webp", "image/webp", webpData}} {
|
||||
t.Run(test.extension, func(t *testing.T) {
|
||||
cat, blobs := &memoryCatalog{}, &memoryBlobs{}
|
||||
remote := &memoryRemote{blob: Blob{Body: io.NopCloser(bytes.NewReader(test.data)), ContentType: "application/octet-stream", Size: int64(len(test.data))}}
|
||||
svc := NewService(cat, blobs, remote, nil, func(prefix string) string { return prefix + "-format" })
|
||||
url := "https://provider.test/result.image?signature=fixture"
|
||||
asset, err := svc.ImportGenerated(context.Background(), PlatformScope("owner-a"), ImportGeneratedCommand{URL: url, Capability: "image.generate"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if asset.Name != "result."+test.extension || !strings.HasSuffix(asset.StoragePath, "-result."+test.extension) || asset.Kind != KindImage {
|
||||
t.Fatalf("asset=%#v", asset)
|
||||
}
|
||||
if asset.Metadata["contentType"] != test.contentType || asset.Metadata["importedFrom"] != url || blobs.putContentType != test.contentType {
|
||||
t.Fatalf("metadata=%#v stored MIME=%q", asset.Metadata, blobs.putContentType)
|
||||
}
|
||||
if !bytes.Equal(blobs.putBody, test.data) {
|
||||
t.Fatal("renaming changed the original image bytes")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUploadRecognizesReuploadedLegacyImage(t *testing.T) {
|
||||
pngData, _, _ := imageFormatFixtures(t)
|
||||
cat, blobs := &memoryCatalog{}, &memoryBlobs{}
|
||||
svc := NewService(cat, blobs, nil, nil, nil)
|
||||
asset, err := svc.Upload(context.Background(), PlatformScope("owner-a"), UploadCommand{FileName: "download.image", ContentType: "application/octet-stream", Bytes: pngData})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if asset.Name != "download.png" || asset.Kind != KindImage || asset.Metadata["contentType"] != "image/png" || !strings.HasSuffix(asset.StoragePath, "-download.png") || !bytes.Equal(blobs.putBody, pngData) {
|
||||
t.Fatalf("upload=%#v", asset)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeImageDownloadStreamsAllBytesAndClosesOriginal(t *testing.T) {
|
||||
pngData, _, _ := imageFormatFixtures(t)
|
||||
for _, data := range [][]byte{pngData, append(append([]byte(nil), pngData...), bytes.Repeat([]byte{7}, 4096)...)} {
|
||||
body := &imageTrackedBody{Reader: bytes.NewReader(data)}
|
||||
name, blob, err := NormalizeImageDownload("legacy.image", Blob{Body: body, ContentType: "application/octet-stream", Size: int64(len(data))})
|
||||
if err != nil || name != "legacy.png" || blob.ContentType != "image/png" || blob.Size != int64(len(data)) {
|
||||
t.Fatalf("name=%q blob=%#v err=%v", name, blob, err)
|
||||
}
|
||||
if body.readBytes > 512 || body.closed {
|
||||
t.Fatalf("prefix read consumed %d bytes, closed=%v", body.readBytes, body.closed)
|
||||
}
|
||||
got, err := io.ReadAll(blob.Body)
|
||||
if err != nil || !bytes.Equal(got, data) {
|
||||
t.Fatalf("stream mismatch; error=%v", err)
|
||||
}
|
||||
if err := blob.Body.Close(); err != nil || !body.closed {
|
||||
t.Fatalf("original body not closed; error=%v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeImageDownloadUnknownDataIsNotRenamed(t *testing.T) {
|
||||
for _, data := range [][]byte{[]byte("video bytes"), nil} {
|
||||
name, blob, err := NormalizeImageDownload("legacy.image", Blob{Body: io.NopCloser(bytes.NewReader(data)), ContentType: "application/octet-stream", Size: int64(len(data))})
|
||||
if err != nil || name != "legacy.image" || blob.ContentType != "application/octet-stream" {
|
||||
t.Fatalf("name=%q blob=%#v err=%v", name, blob, err)
|
||||
}
|
||||
got, _ := io.ReadAll(blob.Body)
|
||||
_ = blob.Body.Close()
|
||||
if !bytes.Equal(got, data) {
|
||||
t.Fatal("unknown bytes changed")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeImageDownloadReadError(t *testing.T) {
|
||||
failure := errors.New("read failed")
|
||||
body := &imageTrackedBody{Reader: imageErrorReader{err: failure}}
|
||||
_, _, err := NormalizeImageDownload("legacy.image", Blob{Body: body})
|
||||
if !errors.Is(err, failure) || body.closed {
|
||||
t.Fatalf("error=%v closed=%v", err, body.closed)
|
||||
}
|
||||
_ = body.Close()
|
||||
if _, _, err := NormalizeImageDownload("legacy.image", Blob{}); !errors.Is(err, ErrBlobNotFound) {
|
||||
t.Fatalf("nil body error=%v", err)
|
||||
}
|
||||
}
|
||||
|
||||
type imageTrackedBody struct {
|
||||
io.Reader
|
||||
readBytes int
|
||||
closed bool
|
||||
}
|
||||
|
||||
func (b *imageTrackedBody) Read(p []byte) (int, error) {
|
||||
n, err := b.Reader.Read(p)
|
||||
b.readBytes += n
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (b *imageTrackedBody) Close() error {
|
||||
b.closed = true
|
||||
return nil
|
||||
}
|
||||
|
||||
type imageErrorReader struct{ err error }
|
||||
|
||||
func (r imageErrorReader) Read([]byte) (int, error) { return 0, r.err }
|
||||
@@ -295,9 +295,14 @@ func (h *assetsHandler) download(w http.ResponseWriter, r *http.Request, public
|
||||
return
|
||||
}
|
||||
defer blob.Body.Close()
|
||||
disposition := contentDisposition(a.Name)
|
||||
name, blob, err := assets.NormalizeImageDownload(a.Name, blob)
|
||||
if err != nil {
|
||||
writeAssetError(w, err, public, "")
|
||||
return
|
||||
}
|
||||
disposition := contentDisposition(name)
|
||||
if inline {
|
||||
disposition = inlineContentDisposition(a.Name)
|
||||
disposition = inlineContentDisposition(name)
|
||||
}
|
||||
writeBlob(w, blob, "private, no-store", disposition)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
package httpapi_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/assets"
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/httpapi"
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/publicapi"
|
||||
)
|
||||
|
||||
type formatDownloadBlobs struct {
|
||||
*assetBlobs
|
||||
contentTypes map[string]string
|
||||
reads int
|
||||
}
|
||||
|
||||
type failingImageBody struct {
|
||||
prefix []byte
|
||||
closed bool
|
||||
}
|
||||
|
||||
func (b *failingImageBody) Read(p []byte) (int, error) {
|
||||
if len(b.prefix) > 0 {
|
||||
n := copy(p, b.prefix)
|
||||
b.prefix = b.prefix[n:]
|
||||
return n, nil
|
||||
}
|
||||
return 0, errors.New("storage stream failed")
|
||||
}
|
||||
|
||||
func (b *failingImageBody) Close() error {
|
||||
b.closed = true
|
||||
return nil
|
||||
}
|
||||
|
||||
type failingDownloadBlobs struct {
|
||||
*assetBlobs
|
||||
body *failingImageBody
|
||||
}
|
||||
|
||||
func (b *failingDownloadBlobs) Read(context.Context, string) (assets.Blob, error) {
|
||||
return assets.Blob{Body: b.body, ContentType: "application/octet-stream", Size: 1000}, nil
|
||||
}
|
||||
|
||||
func (b *formatDownloadBlobs) Read(ctx context.Context, key string) (assets.Blob, error) {
|
||||
blob, err := b.assetBlobs.Read(ctx, key)
|
||||
if err == nil {
|
||||
b.reads++
|
||||
blob.ContentType = b.contentTypes[key]
|
||||
}
|
||||
return blob, err
|
||||
}
|
||||
|
||||
func TestLegacyImageDownloadUsesBytesForResponseNameAndType(t *testing.T) {
|
||||
tests := []struct {
|
||||
label, id, name, mediaType, suffix string
|
||||
data []byte
|
||||
}{
|
||||
{"png", "legacy-png", "生成图片.image", "image/png", ".png", append([]byte("\x89PNG\r\n\x1a\n"), bytes.Repeat([]byte("p"), 650)...)},
|
||||
{"jpeg", "legacy-jpeg", "生成图片.image", "image/jpeg", ".jpg", append([]byte("\xff\xd8\xff\xe0"), bytes.Repeat([]byte("j"), 650)...)},
|
||||
{"webp", "legacy-webp", "生成图片.image", "image/webp", ".webp", append([]byte("RIFF\x12\x00\x00\x00WEBPVP8 "), bytes.Repeat([]byte("w"), 650)...)},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.label, func(t *testing.T) {
|
||||
key := "uploads/" + tc.id + ".image"
|
||||
original := assets.Asset{ID: tc.id, OwnerID: "demo-merchant", Name: tc.name, Kind: assets.KindImage, StoragePath: key, URL: "https://cdn.test/" + key}
|
||||
catalog := &assetCatalog{values: []assets.Asset{original}}
|
||||
blobs := &formatDownloadBlobs{
|
||||
assetBlobs: &assetBlobs{values: map[string][]byte{key: bytes.Clone(tc.data)}},
|
||||
contentTypes: map[string]string{key: "application/octet-stream"},
|
||||
}
|
||||
h := newImageDownloadHandler(t, catalog, blobs)
|
||||
for _, inline := range []bool{false, true} {
|
||||
path := "/api/assets/" + tc.id + "/download"
|
||||
kind := "attachment"
|
||||
if inline {
|
||||
path += "?inline=1"
|
||||
kind = "inline"
|
||||
}
|
||||
response := request(t, h, http.MethodGet, path, nil, nil)
|
||||
wantName := strings.TrimSuffix(tc.name, ".image") + tc.suffix
|
||||
if response.Code != http.StatusOK || !bytes.Equal(response.Body.Bytes(), tc.data) {
|
||||
t.Fatalf("inline=%t status=%d body=%q", inline, response.Code, response.Body.Bytes())
|
||||
}
|
||||
if got := response.Header().Get("Content-Type"); got != tc.mediaType {
|
||||
t.Errorf("Content-Type=%q, want %q", got, tc.mediaType)
|
||||
}
|
||||
if got := response.Header().Get("Content-Length"); got != strconv.Itoa(len(tc.data)) {
|
||||
t.Errorf("Content-Length=%q, want %d", got, len(tc.data))
|
||||
}
|
||||
if got := response.Header().Get("Content-Disposition"); !strings.HasPrefix(got, kind+"; ") || !strings.Contains(got, "filename*=UTF-8''"+url.PathEscape(wantName)) {
|
||||
t.Errorf("Content-Disposition=%q, want %s and UTF-8 filename %q", got, kind, wantName)
|
||||
}
|
||||
}
|
||||
if !reflect.DeepEqual(catalog.values, []assets.Asset{original}) || !bytes.Equal(blobs.values[key], tc.data) || len(blobs.values) != 1 {
|
||||
t.Fatal("download changed catalog metadata or stored bytes")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublicLegacyImageDownloadRequiresClientReadAccess(t *testing.T) {
|
||||
data := append([]byte("\x89PNG\r\n\x1a\n"), bytes.Repeat([]byte("x"), 600)...)
|
||||
key := "uploads/public.image"
|
||||
asset := assets.Asset{ID: "public-image", OwnerID: publicapi.OwnerID("agent-a"), Name: "公开图.image", Kind: assets.KindImage, StoragePath: key, Tags: []string{assets.ClientTag("agent-a")}}
|
||||
catalog := &assetCatalog{values: []assets.Asset{asset}}
|
||||
blobs := &formatDownloadBlobs{assetBlobs: &assetBlobs{values: map[string][]byte{key: data}}, contentTypes: map[string]string{key: "application/octet-stream"}}
|
||||
h := newImageDownloadHandler(t, catalog, blobs)
|
||||
path := "/api/v1/assets/public-image/download"
|
||||
|
||||
for _, headers := range []map[string]string{nil, {"Authorization": "Bearer key-b"}} {
|
||||
response := request(t, h, http.MethodGet, path, nil, headers)
|
||||
if response.Code != http.StatusUnauthorized && response.Code != http.StatusNotFound {
|
||||
t.Fatalf("unreadable public download status=%d body=%s", response.Code, response.Body.String())
|
||||
}
|
||||
}
|
||||
if blobs.reads != 0 {
|
||||
t.Fatalf("blob was read without access: %d reads", blobs.reads)
|
||||
}
|
||||
response := request(t, h, http.MethodGet, path, nil, map[string]string{"Authorization": "Bearer key-a"})
|
||||
if response.Code != http.StatusOK || !bytes.Equal(response.Body.Bytes(), data) || response.Header().Get("Content-Type") != "image/png" || !strings.Contains(response.Header().Get("Content-Disposition"), "filename*=UTF-8''"+url.PathEscape("公开图.png")) {
|
||||
t.Fatalf("public download status=%d headers=%v body=%q", response.Code, response.Header(), response.Body.Bytes())
|
||||
}
|
||||
if blobs.reads != 1 {
|
||||
t.Fatalf("authorized blob reads=%d, want 1", blobs.reads)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVideoDownloadKeepsOriginalNameAndType(t *testing.T) {
|
||||
key := "uploads/video.mp4"
|
||||
data := []byte("\x00\x00\x00\x18ftypmp42-video-content")
|
||||
catalog := &assetCatalog{values: []assets.Asset{{ID: "video", OwnerID: "demo-merchant", Kind: assets.KindVideo, Name: "视频.image", StoragePath: key}}}
|
||||
blobs := &formatDownloadBlobs{assetBlobs: &assetBlobs{values: map[string][]byte{key: data}}, contentTypes: map[string]string{key: "video/mp4"}}
|
||||
h := newImageDownloadHandler(t, catalog, blobs)
|
||||
response := request(t, h, http.MethodGet, "/api/assets/video/download", nil, nil)
|
||||
if response.Code != http.StatusOK || !bytes.Equal(response.Body.Bytes(), data) || response.Header().Get("Content-Type") != "video/mp4" || !strings.Contains(response.Header().Get("Content-Disposition"), "filename*=UTF-8''"+url.PathEscape("视频.image")) {
|
||||
t.Fatalf("video download status=%d headers=%v body=%q", response.Code, response.Header(), response.Body.Bytes())
|
||||
}
|
||||
}
|
||||
|
||||
func TestImageDownloadPrefixFailureDoesNotSendSuccessHeaders(t *testing.T) {
|
||||
key := "uploads/failing.image"
|
||||
catalog := &assetCatalog{values: []assets.Asset{{ID: "failing", OwnerID: "demo-merchant", Kind: assets.KindImage, Name: "失败.image", StoragePath: key}}}
|
||||
body := &failingImageBody{prefix: []byte("\x89PNG\r\n\x1a\n")}
|
||||
blobs := &failingDownloadBlobs{assetBlobs: &assetBlobs{values: map[string][]byte{}}, body: body}
|
||||
h := newImageDownloadHandler(t, catalog, blobs)
|
||||
response := request(t, h, http.MethodGet, "/api/assets/failing/download", nil, nil)
|
||||
if response.Code != http.StatusInternalServerError || response.Header().Get("Content-Disposition") != "" || response.Header().Get("Content-Type") == "image/png" || !body.closed {
|
||||
t.Fatalf("failed download status=%d headers=%v closed=%t body=%q", response.Code, response.Header(), body.closed, response.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func newImageDownloadHandler(t *testing.T, catalog *assetCatalog, blobs assets.BlobStore) http.Handler {
|
||||
t.Helper()
|
||||
service := assets.NewService(catalog, blobs, nil, time.Now, nil)
|
||||
platform, err := httpapi.NewPlatformAuthorizer(httpapi.AuthState{}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
h, err := httpapi.NewAssetsHandler(service, platform, publicapi.NewAuthenticator(publicapi.Config{APIKeys: "agent-a:key-a,agent-b:key-b"}), httpapi.AssetsConfig{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
var _ io.ReadCloser = (*failingImageBody)(nil)
|
||||
@@ -109,6 +109,7 @@ func (s *Store) SeedBillingPriceRules(_ context.Context, rules []billing.PriceRu
|
||||
if ok {
|
||||
r.MarkupMultiplier = existing.MarkupMultiplier
|
||||
r.CreatedAt = existing.CreatedAt
|
||||
r.Dimensions = preserveTierMarkups(r.Dimensions, existing.Dimensions)
|
||||
}
|
||||
if r.CreatedAt == "" {
|
||||
r.CreatedAt = now
|
||||
@@ -119,6 +120,30 @@ func (s *Store) SeedBillingPriceRules(_ context.Context, rules []billing.PriceRu
|
||||
return nil
|
||||
}
|
||||
|
||||
// Defaults own the catalog shape and official costs; administrators own the
|
||||
// markup for each matching dimension/tier. Copy before editing so a reused
|
||||
// default catalog is not changed by one store's customizations.
|
||||
func preserveTierMarkups(defaults, existing []billing.ParameterDimension) []billing.ParameterDimension {
|
||||
markups := make(map[[2]string]float64)
|
||||
for _, dimension := range existing {
|
||||
for _, tier := range dimension.Tiers {
|
||||
if tier.MarkupMultiplier >= 1 && tier.MarkupMultiplier <= 1000 {
|
||||
markups[[2]string{dimension.Key, fmt.Sprint(tier.Value)}] = tier.MarkupMultiplier
|
||||
}
|
||||
}
|
||||
}
|
||||
merged := append([]billing.ParameterDimension(nil), defaults...)
|
||||
for di := range merged {
|
||||
merged[di].Tiers = append([]billing.ParameterTier(nil), merged[di].Tiers...)
|
||||
for ti := range merged[di].Tiers {
|
||||
if markup, ok := markups[[2]string{merged[di].Key, fmt.Sprint(merged[di].Tiers[ti].Value)}]; ok {
|
||||
merged[di].Tiers[ti].MarkupMultiplier = markup
|
||||
}
|
||||
}
|
||||
}
|
||||
return merged
|
||||
}
|
||||
|
||||
func (s *Store) PostBillingWalletEntry(ctx context.Context, p billing.WalletPostParams) (billing.WalletPosting, error) {
|
||||
return s.PostWalletEntry(ctx, p)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,188 @@
|
||||
package localstore_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/localstore"
|
||||
)
|
||||
|
||||
func TestPriceReseedingPreservesTierMarkupsAndRefreshesDefinitions(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := localstore.New()
|
||||
old := billing.PriceRule{
|
||||
ID: "price", Provider: "test", Capability: "image.generate", Unit: billing.UnitImage,
|
||||
StandardUnitPriceFen: 30, MarkupMultiplier: 1.8, Enabled: true,
|
||||
Dimensions: []billing.ParameterDimension{
|
||||
{Key: "size", Label: "Old size", BaselineValue: "1K", Tiers: []billing.ParameterTier{
|
||||
{Value: "1K", Label: "Old tier", StandardFactor: 1, MarkupMultiplier: 1.3, Enabled: true},
|
||||
{Value: "removed", StandardFactor: 1, MarkupMultiplier: 1.7, Enabled: true},
|
||||
}},
|
||||
{Key: "referenceImageCount", BaselineValue: 0, Tiers: []billing.ParameterTier{
|
||||
{Value: 0, StandardFactor: 1, MarkupMultiplier: 1.9, Enabled: true},
|
||||
}},
|
||||
{Key: "removed", BaselineValue: "1K", Tiers: []billing.ParameterTier{
|
||||
{Value: "1K", StandardFactor: 1, MarkupMultiplier: 5, Enabled: true},
|
||||
}},
|
||||
},
|
||||
}
|
||||
if err := store.SeedBillingPriceRules(ctx, []billing.PriceRule{old}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
next := billing.PriceRule{
|
||||
ID: old.ID, Provider: old.Provider, Capability: old.Capability, Unit: old.Unit,
|
||||
StandardUnitPriceFen: 40, MarkupMultiplier: 1.2, Enabled: true, Note: "New official cost",
|
||||
Dimensions: []billing.ParameterDimension{
|
||||
{Key: "referenceImageCount", BaselineValue: "0", Tiers: []billing.ParameterTier{
|
||||
{Value: "0", StandardFactor: 1, MarkupMultiplier: 1.2, Enabled: true},
|
||||
}},
|
||||
{Key: "size", Label: "New size", BaselineValue: "2K", DefaultValue: "2K", Tiers: []billing.ParameterTier{
|
||||
{Value: "2K", Label: "New tier", StandardFactor: 4, MarkupMultiplier: 1.2, Enabled: true},
|
||||
{Value: "1K", Label: "New label", StandardFactor: 2, MarkupMultiplier: 1.2, Enabled: false},
|
||||
}},
|
||||
{Key: "new", BaselineValue: "1K", Tiers: []billing.ParameterTier{
|
||||
{Value: "1K", StandardFactor: 1, MarkupMultiplier: 1.2, Enabled: true},
|
||||
}},
|
||||
},
|
||||
}
|
||||
before, _ := json.Marshal(next)
|
||||
for i := 0; i < 3; i++ {
|
||||
if err := store.SeedBillingPriceRules(ctx, []billing.PriceRule{next}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, err := store.GetBillingPriceRule(ctx, old.ID)
|
||||
if err != nil || got == nil {
|
||||
t.Fatalf("rule = %#v, %v", got, err)
|
||||
}
|
||||
want := next
|
||||
want.MarkupMultiplier = old.MarkupMultiplier
|
||||
want.CreatedAt, want.UpdatedAt = got.CreatedAt, got.UpdatedAt
|
||||
// Decode into a fresh value so editing the expectation cannot mutate defaults.
|
||||
raw, _ := json.Marshal(want)
|
||||
want = billing.PriceRule{}
|
||||
if err := json.Unmarshal(raw, &want); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want.Dimensions[0].Tiers[0].MarkupMultiplier = 1.9
|
||||
want.Dimensions[1].Tiers[1].MarkupMultiplier = 1.3
|
||||
if !reflect.DeepEqual(*got, want) {
|
||||
t.Fatalf("reseed %d: got %#v, want %#v", i, got, want)
|
||||
}
|
||||
}
|
||||
after, _ := json.Marshal(next)
|
||||
if string(before) != string(after) {
|
||||
t.Fatal("seeding mutated the supplied default catalog")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSeedreamTierUpdateSurvivesQuotesAndCharges(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := localstore.New()
|
||||
service := billing.NewService(store, nil)
|
||||
if err := store.SeedBillingPriceRules(ctx, billing.DefaultBillingPriceRules()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, update := range []struct {
|
||||
id string
|
||||
markup float64
|
||||
}{
|
||||
{"base-seedream-5-0-pro", 1.3},
|
||||
{"base-seedream-5-0-pro-layers", 2},
|
||||
} {
|
||||
rule, err := service.GetPrice(ctx, update.id)
|
||||
if err != nil || rule == nil {
|
||||
t.Fatalf("get rule = %#v, %v", rule, err)
|
||||
}
|
||||
patch := billing.PricePatch{DimensionKey: "size", TierValue: "1.5K", MarkupMultiplier: update.markup}
|
||||
if err := billing.ValidatePricePatch(rule, patch); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := service.UpdatePrice(ctx, rule.ID, patch); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if _, err := store.PostWalletEntry(ctx, billing.WalletPostParams{
|
||||
LedgerID: "top-up", OrganizationID: "org", Kind: "recharge", DeltaFen: 10000,
|
||||
Currency: billing.CurrencyCNY, IdempotencyKey: "top-up",
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ledger := billing.Ledger{Poster: store, NewID: func() string { return "charge" }}
|
||||
for i := 0; i < 3; i++ {
|
||||
basic, err := service.Quote(ctx, billing.QuoteCommand{
|
||||
Provider: "seedream", Capability: "image.generate", ReqKey: billing.Seedream50ProModel,
|
||||
Parameters: billing.Parameters{"size": "1.5K"},
|
||||
})
|
||||
if err != nil || basic == nil || basic.MarkupMultiplier != 1.3 || basic.AmountFen != 39 {
|
||||
t.Fatalf("basic quote %d = %#v, %v", i, basic, err)
|
||||
}
|
||||
layers, err := service.Quote(ctx, billing.QuoteCommand{
|
||||
Provider: "seedream", Capability: "image.generate", ReqKey: billing.Seedream50ProModel,
|
||||
Parameters: billing.Parameters{"size": "1.5K", "layerDecomposition": true},
|
||||
})
|
||||
if err != nil || layers == nil || layers.MarkupMultiplier != 2 || layers.ReservedAmountFen != 510 {
|
||||
t.Fatalf("layers quote %d = %#v, %v", i, layers, err)
|
||||
}
|
||||
actual, err := billing.CalculateSeedreamLayerAmountFen([]billing.SeedreamLayerImage{
|
||||
{Width: 1024, Height: 1024}, {Width: 2048, Height: 2048},
|
||||
}, layers.MarkupMultiplier)
|
||||
if err != nil || actual != 90 {
|
||||
t.Fatalf("layer settlement amount = %d, %v", actual, err)
|
||||
}
|
||||
other, err := service.Quote(ctx, billing.QuoteCommand{
|
||||
Provider: "seedream", Capability: "image.generate", ReqKey: billing.Seedream50ProModel,
|
||||
Parameters: billing.Parameters{"size": "2K"},
|
||||
})
|
||||
if err != nil || other == nil || other.MarkupMultiplier != 1.2 || other.AmountFen != 72 {
|
||||
t.Fatalf("unmodified tier = %#v, %v", other, err)
|
||||
}
|
||||
posting, err := ledger.Charge(ctx, billing.ChargeRequest{
|
||||
OrganizationID: "org", AccountID: "member", JobID: fmt.Sprintf("job-%d", i), AmountFen: basic.AmountFen,
|
||||
})
|
||||
if err != nil || posting.DeltaFen != -39 || posting.BalanceFen != 10000-int64(i+1)*39 {
|
||||
t.Fatalf("charge = %#v, %v", posting, err)
|
||||
}
|
||||
// A catalog list refresh must still expose the same saved multiplier.
|
||||
rules, err := service.ListPrices(ctx)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, rule := range rules {
|
||||
if rule.ID == "base-seedream-5-0-pro" && rule.Dimensions[0].Tiers[2].MarkupMultiplier != 1.3 {
|
||||
t.Fatal("catalog refresh lost the saved markup")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPriceReseedingUsesDefaultsForInvalidSavedTierMarkup(t *testing.T) {
|
||||
for _, markup := range []float64{0, -1, 1001} {
|
||||
t.Run(fmt.Sprint(markup), func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := localstore.New()
|
||||
rules := billing.DefaultBillingPriceRules()
|
||||
if err := store.SeedBillingPriceRules(ctx, rules); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := store.UpdateBillingPriceRule(ctx, "base-seedream-5-0-pro", billing.PricePatch{
|
||||
DimensionKey: "size", TierValue: "1.5K", MarkupMultiplier: markup,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := store.SeedBillingPriceRules(ctx, rules); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rule, err := store.GetBillingPriceRule(ctx, "base-seedream-5-0-pro")
|
||||
if err != nil || rule == nil {
|
||||
t.Fatalf("get price = %#v, %v", rule, err)
|
||||
}
|
||||
if got := rule.Dimensions[0].Tiers[2].MarkupMultiplier; got != 1.2 {
|
||||
t.Fatalf("invalid saved markup %v survived reseed as %v", markup, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -136,6 +136,9 @@ func cloneRule(r billing.PriceRule) billing.PriceRule {
|
||||
r.Conditions = billing.Conditions(cloneMap(map[string]any(r.Conditions)))
|
||||
r.Source = cloneMap(r.Source)
|
||||
raw, _ := json.Marshal(r.Dimensions)
|
||||
// Unmarshal otherwise reuses the original slice and shares tier storage
|
||||
// with callers, allowing an admin edit to mutate the default catalog too.
|
||||
r.Dimensions = nil
|
||||
_ = json.Unmarshal(raw, &r.Dimensions)
|
||||
return r
|
||||
}
|
||||
|
||||
@@ -8,11 +8,44 @@ import (
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
|
||||
)
|
||||
|
||||
// Merge each default tier with only its saved administrator markup. The default
|
||||
// tier order, factors and metadata remain authoritative. One UPSERT keeps the
|
||||
// merge atomic with concurrent administrator price updates.
|
||||
const SeedBillingPriceRulesSQL = `INSERT INTO public.billing_price_rules (id, provider, capability, req_key, variant_key, unit, standard_unit_price_fen, markup_multiplier, enabled, conditions, quantity_source, priority, note, source, parameter_dimensions)
|
||||
SELECT rule->>'id', rule->>'provider', rule->>'capability', NULLIF(rule->>'reqKey', ''), NULLIF(rule->>'variantKey', ''), rule->>'unit', (rule->>'standardUnitPriceFen')::bigint, (rule->>'markupMultiplier')::numeric, COALESCE((rule->>'enabled')::boolean, true), COALESCE(rule->'conditions', '{}'::jsonb), NULLIF(rule->>'quantitySource', ''), COALESCE((rule->>'priority')::integer, 0), NULLIF(rule->>'note', ''), rule->'source', COALESCE(rule->'parameterDimensions', '[]'::jsonb)
|
||||
FROM jsonb_array_elements($1::jsonb) AS rule
|
||||
ON CONFLICT (provider, capability, COALESCE(req_key, ''), COALESCE(variant_key, ''), COALESCE(conditions, '{}'::jsonb)) DO UPDATE SET
|
||||
req_key = EXCLUDED.req_key, variant_key = EXCLUDED.variant_key, unit = EXCLUDED.unit, standard_unit_price_fen = EXCLUDED.standard_unit_price_fen, enabled = EXCLUDED.enabled, conditions = EXCLUDED.conditions, quantity_source = EXCLUDED.quantity_source, priority = EXCLUDED.priority, note = EXCLUDED.note, source = EXCLUDED.source, parameter_dimensions = EXCLUDED.parameter_dimensions, updated_at = now()`
|
||||
req_key = EXCLUDED.req_key, variant_key = EXCLUDED.variant_key, unit = EXCLUDED.unit, standard_unit_price_fen = EXCLUDED.standard_unit_price_fen, enabled = EXCLUDED.enabled, conditions = EXCLUDED.conditions, quantity_source = EXCLUDED.quantity_source, priority = EXCLUDED.priority, note = EXCLUDED.note, source = EXCLUDED.source,
|
||||
parameter_dimensions = (
|
||||
SELECT COALESCE(jsonb_agg(
|
||||
jsonb_set(default_dimension.dimension, '{tiers}', (
|
||||
SELECT COALESCE(jsonb_agg(
|
||||
CASE
|
||||
WHEN jsonb_typeof(existing_tier.tier->'markupMultiplier') = 'number'
|
||||
THEN CASE
|
||||
WHEN (existing_tier.tier->>'markupMultiplier')::numeric BETWEEN 1 AND 1000
|
||||
THEN jsonb_set(default_tier.tier, '{markupMultiplier}', existing_tier.tier->'markupMultiplier', true)
|
||||
ELSE default_tier.tier
|
||||
END
|
||||
ELSE default_tier.tier
|
||||
END ORDER BY default_tier.position
|
||||
), '[]'::jsonb)
|
||||
FROM jsonb_array_elements(COALESCE(default_dimension.dimension->'tiers', '[]'::jsonb)) WITH ORDINALITY AS default_tier(tier, position)
|
||||
LEFT JOIN LATERAL (
|
||||
SELECT old_tier.tier
|
||||
FROM jsonb_array_elements(COALESCE((
|
||||
SELECT old_dimension.dimension->'tiers'
|
||||
FROM jsonb_array_elements(COALESCE(billing_price_rules.parameter_dimensions, '[]'::jsonb)) AS old_dimension(dimension)
|
||||
WHERE old_dimension.dimension->>'key' = default_dimension.dimension->>'key'
|
||||
LIMIT 1
|
||||
), '[]'::jsonb)) AS old_tier(tier)
|
||||
WHERE old_tier.tier->>'value' = default_tier.tier->>'value'
|
||||
LIMIT 1
|
||||
) AS existing_tier ON true
|
||||
), true) ORDER BY default_dimension.position
|
||||
), '[]'::jsonb)
|
||||
FROM jsonb_array_elements(COALESCE(EXCLUDED.parameter_dimensions, '[]'::jsonb)) WITH ORDINALITY AS default_dimension(dimension, position)
|
||||
), updated_at = now()`
|
||||
|
||||
func (db *Database) SeedBillingPriceRules(ctx context.Context, rules []billing.PriceRule) error {
|
||||
if len(rules) == 0 {
|
||||
|
||||
@@ -3,11 +3,16 @@ package postgres
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
|
||||
"github.com/jackc/pgx/v5"
|
||||
)
|
||||
|
||||
func TestSeedBillingPriceRulesUsesOneIdempotentUpsert(t *testing.T) {
|
||||
@@ -31,3 +36,123 @@ func TestSeedBillingPriceRulesUsesOneIdempotentUpsert(t *testing.T) {
|
||||
t.Fatalf("payload = %#v", decoded)
|
||||
}
|
||||
}
|
||||
|
||||
// Runs against a private temporary PostgreSQL cluster when the local binaries
|
||||
// are supplied. No application database or credentials are used.
|
||||
func TestSeedBillingPriceRulesPreservesTierMarkupInPostgres(t *testing.T) {
|
||||
binDir := os.Getenv("NIAN_TEST_POSTGRES_BIN_DIR")
|
||||
if binDir == "" {
|
||||
t.Skip("set NIAN_TEST_POSTGRES_BIN_DIR to run the PostgreSQL integration test")
|
||||
}
|
||||
base, err := os.MkdirTemp("/tmp", "nian-billing-pg-")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = os.RemoveAll(base) })
|
||||
dataDir := filepath.Join(base, "data")
|
||||
runPG := func(name string, args ...string) {
|
||||
t.Helper()
|
||||
command := exec.Command(filepath.Join(binDir, name), args...)
|
||||
if output, err := command.CombinedOutput(); err != nil {
|
||||
t.Fatalf("%s: %v: %s", name, err, output)
|
||||
}
|
||||
}
|
||||
runPG("initdb", "-D", dataDir, "-A", "trust", "-U", "postgres")
|
||||
options := fmt.Sprintf("-F -c unix_socket_directories=%s -c listen_addresses='' -p 55432", base)
|
||||
runPG("pg_ctl", "-D", dataDir, "-l", filepath.Join(base, "server.log"), "-o", options, "-w", "start")
|
||||
t.Cleanup(func() {
|
||||
command := exec.Command(filepath.Join(binDir, "pg_ctl"), "-D", dataDir, "-m", "immediate", "-w", "stop")
|
||||
if output, err := command.CombinedOutput(); err != nil {
|
||||
t.Errorf("stop temporary PostgreSQL: %v: %s", err, output)
|
||||
}
|
||||
})
|
||||
ctx := context.Background()
|
||||
conn, err := pgx.Connect(ctx, fmt.Sprintf("host=%s port=55432 user=postgres dbname=postgres sslmode=disable", base))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer conn.Close(ctx)
|
||||
_, err = conn.Exec(ctx, `CREATE TABLE public.billing_price_rules (
|
||||
id text PRIMARY KEY, provider text NOT NULL, capability text NOT NULL, req_key text,
|
||||
variant_key text, unit text NOT NULL, standard_unit_price_fen bigint NOT NULL,
|
||||
markup_multiplier numeric NOT NULL, enabled boolean NOT NULL, conditions jsonb NOT NULL,
|
||||
quantity_source text, priority integer NOT NULL, note text, source jsonb,
|
||||
parameter_dimensions jsonb NOT NULL, updated_at timestamptz DEFAULT now());
|
||||
CREATE UNIQUE INDEX billing_price_rules_match_idx ON public.billing_price_rules
|
||||
(provider, capability, coalesce(req_key, ''), coalesce(variant_key, ''), coalesce(conditions, '{}'::jsonb));`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rules := make([]billing.PriceRule, 0, 2)
|
||||
for _, rule := range billing.DefaultBillingPriceRules() {
|
||||
if rule.ID == "base-seedream-5-0-pro" || rule.ID == "base-seedream-5-0-pro-layers" {
|
||||
rules = append(rules, rule)
|
||||
}
|
||||
}
|
||||
if len(rules) != 2 {
|
||||
t.Fatalf("Seedream defaults: %d rules", len(rules))
|
||||
}
|
||||
db := NewDatabase(Config{Backend: BackendPostgres}, seedIntegrationQuerier{conn})
|
||||
if err := db.SeedBillingPriceRules(ctx, rules); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = conn.Exec(ctx, `UPDATE public.billing_price_rules
|
||||
SET parameter_dimensions = jsonb_set(jsonb_set(parameter_dimensions,
|
||||
'{0,tiers,1,markupMultiplier}', '"bad"'::jsonb), '{0,tiers,2,markupMultiplier}', '1.3'::jsonb),
|
||||
markup_multiplier = 1.5
|
||||
WHERE id = 'base-seedream-5-0-pro'`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = conn.Exec(ctx, `UPDATE public.billing_price_rules
|
||||
SET parameter_dimensions = jsonb_set(parameter_dimensions, '{0,tiers,2,markupMultiplier}', '1.4'::jsonb)
|
||||
WHERE id = 'base-seedream-5-0-pro-layers'`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range rules {
|
||||
rules[i].StandardUnitPriceFen++
|
||||
rules[i].Dimensions[0].Tiers = append(rules[i].Dimensions[0].Tiers[1:], billing.ParameterTier{
|
||||
Value: "3K", Label: "3K", StandardFactor: 3, MarkupMultiplier: 1.2, Enabled: true,
|
||||
})
|
||||
rules[i].Dimensions[0].Tiers[1].StandardFactor = 1.5
|
||||
}
|
||||
for attempt := 0; attempt < 2; attempt++ {
|
||||
if err := db.SeedBillingPriceRules(ctx, rules); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
for _, rule := range rules {
|
||||
var cost int64
|
||||
var topLevelMarkup float64
|
||||
var dimensionsJSON []byte
|
||||
if err := conn.QueryRow(ctx, `SELECT standard_unit_price_fen, markup_multiplier, parameter_dimensions FROM public.billing_price_rules WHERE id = $1`, rule.ID).Scan(&cost, &topLevelMarkup, &dimensionsJSON); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var dimensions []billing.ParameterDimension
|
||||
if err := json.Unmarshal(dimensionsJSON, &dimensions); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if cost != rule.StandardUnitPriceFen || len(dimensions) != 1 || len(dimensions[0].Tiers) != 4 {
|
||||
t.Fatalf("%s: cost=%d dimensions=%+v", rule.ID, cost, dimensions)
|
||||
}
|
||||
if dimensions[0].Tiers[1].StandardFactor != 1.5 || dimensions[0].Tiers[3].Value != "3K" || dimensions[0].Tiers[3].MarkupMultiplier != 1.2 {
|
||||
t.Fatalf("%s: default tier definitions not refreshed: %+v", rule.ID, dimensions[0].Tiers)
|
||||
}
|
||||
wantMarkup := 1.3
|
||||
if rule.ID == "base-seedream-5-0-pro-layers" {
|
||||
wantMarkup = 1.4
|
||||
} else if topLevelMarkup != 1.5 || dimensions[0].Tiers[0].MarkupMultiplier != 1.2 {
|
||||
t.Fatalf("basic rule lost top-level markup or malformed tier did not fall back: %+v", dimensions[0].Tiers)
|
||||
}
|
||||
if dimensions[0].Tiers[1].MarkupMultiplier != wantMarkup {
|
||||
t.Fatalf("%s: tier markup = %g, want %g", rule.ID, dimensions[0].Tiers[1].MarkupMultiplier, wantMarkup)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type seedIntegrationQuerier struct{ conn *pgx.Conn }
|
||||
|
||||
func (q seedIntegrationQuerier) Query(ctx context.Context, sql string, args ...any) (Rows, error) {
|
||||
return q.conn.Query(ctx, sql, args...)
|
||||
}
|
||||
Reference in new issue
Block a user