diff --git a/backend/internal/assets/assets.go b/backend/internal/assets/assets.go index 76aed66..152767f 100644 --- a/backend/internal/assets/assets.go +++ b/backend/internal/assets/assets.go @@ -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") } diff --git a/backend/internal/assets/image_format.go b/backend/internal/assets/image_format.go new file mode 100644 index 0000000..97e23e3 --- /dev/null +++ b/backend/internal/assets/image_format.go @@ -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 +} diff --git a/backend/internal/assets/image_format_test.go b/backend/internal/assets/image_format_test.go new file mode 100644 index 0000000..abf03cb --- /dev/null +++ b/backend/internal/assets/image_format_test.go @@ -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 } diff --git a/backend/internal/httpapi/assets.go b/backend/internal/httpapi/assets.go index 1c29db1..3d5dd62 100644 --- a/backend/internal/httpapi/assets.go +++ b/backend/internal/httpapi/assets.go @@ -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) } diff --git a/backend/internal/httpapi/assets_image_names_test.go b/backend/internal/httpapi/assets_image_names_test.go new file mode 100644 index 0000000..7636c70 --- /dev/null +++ b/backend/internal/httpapi/assets_image_names_test.go @@ -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) diff --git a/backend/internal/localstore/billing.go b/backend/internal/localstore/billing.go index 613a019..850ac56 100644 --- a/backend/internal/localstore/billing.go +++ b/backend/internal/localstore/billing.go @@ -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) } diff --git a/backend/internal/localstore/billing_defaults_test.go b/backend/internal/localstore/billing_defaults_test.go new file mode 100644 index 0000000..ae6517a --- /dev/null +++ b/backend/internal/localstore/billing_defaults_test.go @@ -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) + } + }) + } +} diff --git a/backend/internal/localstore/store.go b/backend/internal/localstore/store.go index 6d520b1..13d7648 100644 --- a/backend/internal/localstore/store.go +++ b/backend/internal/localstore/store.go @@ -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 } diff --git a/backend/internal/postgres/billing_defaults.go b/backend/internal/postgres/billing_defaults.go index 88e77c2..7894f4d 100644 --- a/backend/internal/postgres/billing_defaults.go +++ b/backend/internal/postgres/billing_defaults.go @@ -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 { diff --git a/backend/internal/postgres/billing_defaults_test.go b/backend/internal/postgres/billing_defaults_test.go index e81d8ce..b4bfe1a 100644 --- a/backend/internal/postgres/billing_defaults_test.go +++ b/backend/internal/postgres/billing_defaults_test.go @@ -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...) +} diff --git a/components/billing-manager.tsx b/components/billing-manager.tsx index 13264b6..eca2cb1 100644 --- a/components/billing-manager.tsx +++ b/components/billing-manager.tsx @@ -21,6 +21,7 @@ import { } from "lucide-react"; import { billingUnitLabel, formatBillingAmount } from "@/lib/billing"; import { parseAdminBillingPayload, parseMemberBillingPayload } from "@/lib/client/billing-api"; +import { groupedPriceRuleSections, priceTierEditTarget, type PriceEditTarget } from "@/lib/client/billing-price-catalog"; import { pulseFeedback, revealChildren, runScopedMotion } from "@/lib/ui/motion"; import type { BillingAccountConfig, BillingParameterDimension, BillingParameterTier, BillingPriceRule, OrganizationWallet } from "@/lib/types"; @@ -89,13 +90,6 @@ type AdjustmentDraft = { note: string; }; -type PriceEditTarget = { - dimensionKey?: string; - tierValue?: string; - label: string; - markupMultiplier: number; -}; - type PriceEditState = { rule: BillingPriceRule; target: PriceEditTarget; @@ -615,7 +609,7 @@ function priceEditStandardUnitPriceFen(rule: BillingPriceRule, target: PriceEdit return Math.ceil(rule.standardUnitPriceFen * Number(tier?.standardFactor || 1)); } -function PriceCatalog({ rules, onEdit, disabled }: { rules: BillingPriceRule[]; onEdit: (rule: BillingPriceRule, target: PriceEditTarget) => void; disabled: boolean }) { +export function PriceCatalog({ rules, onEdit, disabled }: { rules: BillingPriceRule[]; onEdit: (rule: BillingPriceRule, target: PriceEditTarget) => void; disabled: boolean }) { const groups = groupPriceRules(rules); return
{groups.length ? groups.map(([key, group]) => group.length === 1 ? : ) :
暂无平台价格标准。
} @@ -652,40 +646,48 @@ function PriceServiceCard({ rule, onEdit, disabled }: { rule: BillingPriceRule; function PriceServiceGroup({ rules, onEdit, disabled }: { rules: BillingPriceRule[]; onEdit: (rule: BillingPriceRule, target: PriceEditTarget) => void; disabled: boolean }) { const rule = rules[0]; + const sections = groupedPriceRuleSections(rules); + const hasDimensions = sections.some((section) => section.dimensions.length > 0); return
{providerName(rule.provider).slice(0, 2)}
{providerName(rule.provider)} · {rule.capability === "video.generate" ? "视频生成" : "图片生成"}{rule.reqKey || "默认服务"}
-
平台参数档位{rules.length} 个
+
{hasDimensions ? "计费业务" : "平台参数档位"}{rules.length} 个
{rule.source?.url ? 查看价格来源 : 平台标准目录} 平台标准目录按参数档位列出,用户选择后自动匹配。
-
+ {hasDimensions ?
{sections.flatMap((section) => section.dimensions.length + ? section.dimensions.map((dimension) => ) + : [
+
{section.label}不同标准费率独立计费,倍率仅影响用户价
+
参数档位标准成本用户价倍率操作
+
+
])}
:
参数档位不同标准费率独立计费,倍率仅影响用户价
参数档位标准成本用户价倍率操作
{rules.map((item) => )}
-
+
}
; } -function PriceDimension({ rule, dimension, onEdit, disabled }: { rule: BillingPriceRule; dimension: BillingParameterDimension; onEdit: (rule: BillingPriceRule, target: PriceEditTarget) => void; disabled: boolean }) { +function PriceDimension({ rule, dimension, businessLabel, onEdit, disabled }: { rule: BillingPriceRule; dimension: BillingParameterDimension; businessLabel?: string; onEdit: (rule: BillingPriceRule, target: PriceEditTarget) => void; disabled: boolean }) { const baseline = dimension.tiers.find((tier) => String(tier.value).toLowerCase() === String(dimension.baselineValue).toLowerCase()); return
-
{dimension.label}基准档位:{baseline?.label || String(dimension.baselineValue)} · 组合报价按实际选择自动计算
+
{businessLabel ? `${businessLabel} · ${dimension.label || dimension.key}` : dimension.label || dimension.key}基准档位:{baseline?.label || String(dimension.baselineValue)} · 组合报价按实际选择自动计算
参数档位标准成本用户价倍率操作
-
{dimension.tiers.map((tier) => )}
+
{dimension.tiers.map((tier) => )}
; } -function PriceTierRow({ rule, dimension, tier, onEdit, disabled }: { rule: BillingPriceRule; dimension: BillingParameterDimension; tier: BillingParameterTier; onEdit: (rule: BillingPriceRule, target: PriceEditTarget) => void; disabled: boolean }) { +function PriceTierRow({ rule, dimension, tier, businessLabel, onEdit, disabled }: { rule: BillingPriceRule; dimension: BillingParameterDimension; tier: BillingParameterTier; businessLabel?: string; onEdit: (rule: BillingPriceRule, target: PriceEditTarget) => void; disabled: boolean }) { const standardUnitPriceFen = Math.ceil(rule.standardUnitPriceFen * tier.standardFactor); return
-
{tier.label}{String(tier.value)}{tier.note ? ` · ${tier.note}` : ""}
+
{tier.label || String(tier.value)}{String(tier.value)}{tier.note ? ` · ${tier.note}` : ""}
{formatBillingAmount(standardUnitPriceFen)}/{billingUnitLabel(rule.unit)}
{formatBillingAmount(customerUnitPriceFen(standardUnitPriceFen, tier.markupMultiplier))}/{billingUnitLabel(rule.unit)}
{tier.markupMultiplier.toFixed(2)}× -
+
; } diff --git a/lib/client/billing-price-catalog.ts b/lib/client/billing-price-catalog.ts new file mode 100644 index 0000000..e91d7ce --- /dev/null +++ b/lib/client/billing-price-catalog.ts @@ -0,0 +1,27 @@ +import type { BillingParameterDimension, BillingParameterTier, BillingPriceRule } from "@/lib/types"; + +export type PriceEditTarget = { + dimensionKey?: string; + tierValue?: string; + label: string; + markupMultiplier: number; +}; + +export function groupedPriceRuleSections(rules: BillingPriceRule[]) { + return rules.map((rule, index) => ({ + rule, + label: rule.provider === "seedream" && rule.capability === "image.generate" + ? rule.conditions?.layerDecomposition === true ? "图层拆分" : "基础生图" + : rule.variantKey || `规则 ${index + 1}`, + dimensions: rule.parameterDimensions?.filter((dimension) => dimension.tiers.length > 0) || [] + })); +} + +export function priceTierEditTarget(dimension: BillingParameterDimension, tier: BillingParameterTier, businessLabel?: string): PriceEditTarget { + return { + dimensionKey: dimension.key, + tierValue: String(tier.value), + label: [businessLabel, dimension.label || dimension.key, tier.label || String(tier.value)].filter(Boolean).join(" · "), + markupMultiplier: tier.markupMultiplier + }; +} diff --git a/lib/client/zip-download.ts b/lib/client/zip-download.ts index c85335d..2632cc8 100644 --- a/lib/client/zip-download.ts +++ b/lib/client/zip-download.ts @@ -19,7 +19,7 @@ export async function downloadFilesAsZip(files: DownloadableFile[], archiveName: const response = await fetch(file.url, { credentials: "same-origin" }); if (!response.ok) throw new Error(`下载第 ${index + 1} 个图层失败`); const data = new Uint8Array(await response.arrayBuffer()); - const safeName = uniqueFileName(sanitizeFileName(file.name) || `layer-${index}.png`, usedNames); + const safeName = uniqueFileName(imageFileName(sanitizeFileName(file.name) || `layer-${index}`, data), usedNames); entries.push({ name: new TextEncoder().encode(safeName), data, crc: crc32(data), offset: 0 }); } const blob = zipBlob(entries); @@ -93,6 +93,42 @@ function sanitizeFileName(value: string) { return value.trim().replace(/[\\/:*?"<>|\u0000-\u001f]/g, "-").replace(/^\.+/, "").slice(0, 120); } +function imageFileName(name: string, data: Uint8Array) { + let extensions: string[] | undefined; + let preferred: string | undefined; + if (data.length >= 8 && data[0] === 0x89 && data[1] === 0x50 && data[2] === 0x4e && data[3] === 0x47 && data[4] === 0x0d && data[5] === 0x0a && data[6] === 0x1a && data[7] === 0x0a) { + extensions = ["png"]; + preferred = "png"; + } else if (data.length >= 3 && data[0] === 0xff && data[1] === 0xd8 && data[2] === 0xff) { + extensions = ["jpg", "jpeg", "jpe"]; + preferred = "jpg"; + } else if (data.length >= 14 && asciiAt(data, 0, "RIFF") && asciiAt(data, 8, "WEBPVP")) { + extensions = ["webp"]; + preferred = "webp"; + } else if (data.length >= 6 && (asciiAt(data, 0, "GIF87a") || asciiAt(data, 0, "GIF89a"))) { + extensions = ["gif"]; + preferred = "gif"; + } else if (data.length >= 2 && asciiAt(data, 0, "BM")) { + extensions = ["bmp"]; + preferred = "bmp"; + } else if (data.length >= 4 && data[0] === 0 && data[1] === 0 && (data[2] === 1 || data[2] === 2) && data[3] === 0) { + extensions = ["ico"]; + preferred = "ico"; + } + if (!preferred || !extensions) return name; + const dot = name.lastIndexOf("."); + const current = dot > 0 ? name.slice(dot + 1) : ""; + if (extensions.includes(current.toLowerCase())) return name; + return `${dot > 0 ? name.slice(0, dot) : name}.${preferred}`; +} + +function asciiAt(data: Uint8Array, offset: number, value: string) { + for (let index = 0; index < value.length; index += 1) { + if (data[offset + index] !== value.charCodeAt(index)) return false; + } + return true; +} + function uniqueFileName(name: string, used: Set) { if (!used.has(name)) { used.add(name); diff --git a/lib/server/billing-catalog.ts b/lib/server/billing-catalog.ts index 4a995b5..991aa2c 100644 --- a/lib/server/billing-catalog.ts +++ b/lib/server/billing-catalog.ts @@ -4,7 +4,7 @@ import { updateBillingPriceRule, type BillingPriceRuleInput } from "@/lib/server/billing-store"; -import type { BillingPriceRule } from "@/lib/types"; +import type { BillingParameterDimension, BillingPriceRule } from "@/lib/types"; export const BILLING_USD_CNY_RATE = 7.2; export const DEFAULT_BILLING_MARKUP_MULTIPLIER = 1.2; @@ -390,7 +390,7 @@ export async function ensureDefaultBillingPriceRules(): Promise rule.id === updated.id); @@ -412,6 +412,31 @@ export async function ensureDefaultBillingPriceRules(): Promise { + const savedDimension = saved?.find((item) => item.key === dimension.key); + return { + ...dimension, + tiers: dimension.tiers.map((tier) => { + const savedTier = savedDimension?.tiers.find((item) => String(item.value) === String(tier.value)); + const savedMarkup = savedTier?.markupMultiplier; + return { + ...tier, + markupMultiplier: typeof savedMarkup === "number" + && Number.isFinite(savedMarkup) + && savedMarkup >= 1 + && savedMarkup <= 1000 + ? savedMarkup + : tier.markupMultiplier + }; + }) + }; + }); +} + function priceRuleMatchKey(rule: Pick): string { const conditions = Object.fromEntries(Object.entries(rule.conditions || {}).sort(([left], [right]) => left.localeCompare(right))); return [rule.provider, rule.capability, rule.reqKey || "", rule.variantKey || "", JSON.stringify(conditions)].join("\u0000"); diff --git a/next-env.d.ts b/next-env.d.ts index 830fb59..dd052b5 100644 --- a/next-env.d.ts +++ b/next-env.d.ts @@ -1,6 +1,6 @@ /// /// -/// +/// // NOTE: This file should not be edited // see https://nextjs.org/docs/app/api-reference/config/typescript for more information. diff --git a/next.config.ts b/next.config.ts index f13478d..4630879 100644 --- a/next.config.ts +++ b/next.config.ts @@ -6,6 +6,30 @@ const nextConfig: NextConfig = { // while a dev server is alive must not invalidate the dev server's chunks. distDir: process.env.NODE_ENV === "development" ? ".next-dev" : ".next", devIndicators: false, + + ...(process.env.NODE_ENV === "development" + ? { + async rewrites() { + const backend = "http://127.0.0.1:8080"; + + return [ + { + source: "/api/:path*", + destination: `${backend}/api/:path*` + }, + { + source: "/uploads/:path*", + destination: `${backend}/uploads/:path*` + }, + { + source: "/generated-results/:path*", + destination: `${backend}/generated-results/:path*` + } + ]; + } + } + : {}), + images: { unoptimized: true, remotePatterns: [ diff --git a/tests/billing-catalog-seed.test.ts b/tests/billing-catalog-seed.test.ts new file mode 100644 index 0000000..a32e560 --- /dev/null +++ b/tests/billing-catalog-seed.test.ts @@ -0,0 +1,78 @@ +import { mkdtemp, rm } from "node:fs/promises"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { DEFAULT_BILLING_PRICE_RULES, ensureDefaultBillingPriceRules } from "@/lib/server/billing-catalog"; +import { quoteGenerationCharge } from "@/lib/server/billing-service"; +import { createBillingPriceRule, getBillingPriceRule, updateBillingPriceTierMultiplier } from "@/lib/server/billing-store"; + +let runtimeDir = ""; + +describe("default billing catalog refresh", () => { + beforeEach(async () => { + runtimeDir = await mkdtemp(join(tmpdir(), "zhinian-billing-seed-")); + vi.stubEnv("ZHINIAN_RUNTIME_DIR", runtimeDir); + vi.stubEnv("ZHINIAN_DATA_BACKEND", "local"); + }); + + afterEach(async () => { + vi.unstubAllEnvs(); + await rm(runtimeDir, { recursive: true, force: true }); + }); + + it("keeps a saved tier multiplier across catalog refresh and quotation", async () => { + await ensureDefaultBillingPriceRules(); + await updateBillingPriceTierMultiplier({ + ruleId: "base-seedream-5-0-pro", + dimensionKey: "size", + tierValue: "1.5K", + markupMultiplier: 1.3 + }); + + const quote = await quoteGenerationCharge({ + provider: "seedream", + capability: "image.generate", + reqKey: "doubao-seedream-5-0-pro-260628", + requestPayload: { settings: { size: "1.5K" } }, + usageContext: { source: "platform", accountId: "user-1", displayName: "测试用户", organizationId: "org-1" } + }); + expect(quote).toMatchObject({ standardUnitPriceFen: 30, markupMultiplier: 1.3, amountFen: 39 }); + + const saved = await getBillingPriceRule("base-seedream-5-0-pro"); + const tiers = saved?.parameterDimensions?.find((dimension) => dimension.key === "size")?.tiers; + expect(tiers?.find((tier) => tier.value === "1.5K")?.markupMultiplier).toBe(1.3); + expect(tiers?.find((tier) => tier.value === "2K")?.markupMultiplier).toBe(1.2); + }); + + it("takes current default definitions while retaining valid matching tier markups", async () => { + const candidate = DEFAULT_BILLING_PRICE_RULES.find((rule) => rule.id === "base-seedream-5-0-pro")!; + const size = candidate.parameterDimensions!.find((dimension) => dimension.key === "size")!; + await createBillingPriceRule({ + ...candidate, + standardUnitPriceFen: 5, + markupMultiplier: 1.7, + parameterDimensions: [{ + ...size, + label: "旧分辨率", + tiers: [ + { ...size.tiers[0], standardFactor: 99, markupMultiplier: 0 }, + { ...size.tiers[1], standardFactor: 99, markupMultiplier: 1.4 }, + { value: "3K", label: "旧档位", standardFactor: 9, markupMultiplier: 3, enabled: true } + ] + }] + }); + + await ensureDefaultBillingPriceRules(); + const saved = await getBillingPriceRule(candidate.id!); + expect(saved?.standardUnitPriceFen).toBe(30); + expect(saved?.markupMultiplier).toBe(1.7); + const dimension = saved?.parameterDimensions?.find((item) => item.key === "size"); + expect(dimension?.label).toBe("分辨率"); + expect(dimension?.tiers.map((tier) => tier.value)).toEqual(["1K", "1.5K", "2K"]); + expect(dimension?.tiers.map((tier) => [tier.standardFactor, tier.markupMultiplier])).toEqual([ + [1, 1.2], + [1, 1.4], + [2, 1.2] + ]); + }); +}); diff --git a/tests/billing-price-catalog.test.ts b/tests/billing-price-catalog.test.ts new file mode 100644 index 0000000..aa8ea6a --- /dev/null +++ b/tests/billing-price-catalog.test.ts @@ -0,0 +1,113 @@ +import { createElement } from "react"; +import { renderToStaticMarkup } from "react-dom/server"; +import { describe, expect, it } from "vitest"; +import { PriceCatalog } from "@/components/billing-manager"; +import { groupedPriceRuleSections, priceTierEditTarget } from "@/lib/client/billing-price-catalog"; +import type { BillingParameterDimension, BillingPriceRule } from "@/lib/types"; + +const size: BillingParameterDimension = { + key: "size", + label: "分辨率", + baselineValue: "1.5K", + tiers: [ + { value: "1.5K", label: "1.5K", standardFactor: 1, markupMultiplier: 1.2, enabled: true }, + { value: "2K", label: "2K", standardFactor: 2, markupMultiplier: 1.3, enabled: true } + ] +}; + +function rule(id: string, overrides: Partial = {}): BillingPriceRule { + return { + id, + provider: "seedream", + capability: "image.generate", + reqKey: "doubao-seedream-5-0-pro-260628", + unit: "image", + standardUnitPriceFen: 30, + markupMultiplier: 1.2, + enabled: true, + createdAt: "2026-01-01T00:00:00Z", + updatedAt: "2026-01-01T00:00:00Z", + ...overrides + }; +} + +describe("billing price catalog grouped rules", () => { + it("keeps separate Seedream businesses and their resolution edit targets", () => { + const basic = rule("basic", { parameterDimensions: [size] }); + const layers = rule("layers", { + standardUnitPriceFen: 15, + conditions: { layerDecomposition: true }, + parameterDimensions: [size] + }); + const sections = groupedPriceRuleSections([basic, layers]); + expect(sections.map(({ label, dimensions }) => [label, dimensions[0]?.key])).toEqual([ + ["基础生图", "size"], + ["图层拆分", "size"] + ]); + expect(priceTierEditTarget(sections[1].dimensions[0], size.tiers[1], sections[1].label)).toEqual({ + dimensionKey: "size", + tierValue: "2K", + label: "图层拆分 · 分辨率 · 2K", + markupMultiplier: 1.3 + }); + + expect(sections[0].rule.id).toBe("basic"); + expect(sections[1].rule.id).toBe("layers"); + expect(sections.map((section) => section.dimensions[0].tiers.map((tier) => tier.value))).toEqual([ + ["1.5K", "2K"], + ["1.5K", "2K"] + ]); + const html = renderToStaticMarkup(createElement(PriceCatalog, { rules: [basic, layers], onEdit: () => {}, disabled: false })); + expect(html).toContain("基础生图 · 分辨率"); + expect(html).toContain("图层拆分 · 分辨率"); + expect(html.match(/>调整倍率<\/button>/g)).toHaveLength(4); + expect(html).toContain("¥0.15"); + expect(html).toContain("¥0.30"); + }); + + it("keeps grouped legacy video rules as independent flat rows", () => { + const rules = [ + rule("720p", { provider: "seedance", capability: "video.generate", reqKey: "seedance-2.0", variantKey: "resolution=720p", unit: "video_second", standardUnitPriceFen: 99 }), + rule("1080p", { provider: "seedance", capability: "video.generate", reqKey: "seedance-2.0", variantKey: "resolution=1080p", unit: "video_second", standardUnitPriceFen: 248 }) + ]; + expect(groupedPriceRuleSections(rules).every((section) => section.dimensions.length === 0)).toBe(true); + expect(groupedPriceRuleSections(rules).map((section) => section.label)).toEqual(["resolution=720p", "resolution=1080p"]); + const html = renderToStaticMarkup(createElement(PriceCatalog, { rules, onEdit: () => {}, disabled: false })); + expect(html).toContain("resolution=720p"); + expect(html).toContain("resolution=1080p"); + expect(html.match(/>调整倍率<\/button>/g)).toHaveLength(2); + }); + + it("retains unnamed dimension and tier identifiers in edit labels", () => { + const dimension: BillingParameterDimension = { + key: "quality", + label: "", + baselineValue: "medium", + tiers: [{ value: "high", label: "", standardFactor: 1, markupMultiplier: 1.25, enabled: true }] + }; + expect(priceTierEditTarget(dimension, dimension.tiers[0])).toEqual({ + dimensionKey: "quality", + tierValue: "high", + label: "quality · high", + markupMultiplier: 1.25 + }); + const single = rule("evolink", { + provider: "evolink", + reqKey: "gpt-image-2", + parameterDimensions: [dimension] + }); + const singleHtml = renderToStaticMarkup(createElement(PriceCatalog, { rules: [single], onEdit: () => {}, disabled: false })); + expect(singleHtml).toContain(">quality"); + expect(singleHtml).toContain(">high"); + expect(singleHtml.match(/>调整倍率<\/button>/g)).toHaveLength(1); + + const mixedHtml = renderToStaticMarkup(createElement(PriceCatalog, { + rules: [single, rule("legacy", { provider: "evolink", reqKey: "gpt-image-2", variantKey: "legacy" })], + onEdit: () => {}, + disabled: false + })); + expect(mixedHtml).toContain("规则 1 · quality"); + expect(mixedHtml).toContain("legacy"); + expect(mixedHtml.match(/>调整倍率<\/button>/g)).toHaveLength(2); + }); +}); diff --git a/tests/zip-download.test.ts b/tests/zip-download.test.ts new file mode 100644 index 0000000..bcf8a2e --- /dev/null +++ b/tests/zip-download.test.ts @@ -0,0 +1,81 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; + +import { downloadFilesAsZip } from "@/lib/client/zip-download"; + +const png = Uint8Array.from([0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, 1, 2]); +const jpeg = Uint8Array.from([0xff, 0xd8, 0xff, 0xe0, 3, 4]); +const webp = Uint8Array.from([0x52, 0x49, 0x46, 0x46, 4, 0, 0, 0, 0x57, 0x45, 0x42, 0x50, 0x56, 0x50, 0x38, 0x20, 5]); +const gif = Uint8Array.from([0x47, 0x49, 0x46, 0x38, 0x39, 0x61, 6]); +const bmp = Uint8Array.from([0x42, 0x4d, 7]); +const ico = Uint8Array.from([0, 0, 1, 0, 8]); + +function zipEntries(data: Uint8Array) { + const view = new DataView(data.buffer, data.byteOffset, data.byteLength); + const entries: { name: string; bytes: Uint8Array }[] = []; + let offset = 0; + while (view.getUint32(offset, true) === 0x04034b50) { + const size = view.getUint32(offset + 18, true); + const nameSize = view.getUint16(offset + 26, true); + const extraSize = view.getUint16(offset + 28, true); + const nameStart = offset + 30; + const bytesStart = nameStart + nameSize + extraSize; + entries.push({ + name: new TextDecoder().decode(data.subarray(nameStart, nameStart + nameSize)), + bytes: data.slice(bytesStart, bytesStart + size) + }); + offset = bytesStart + size; + } + return entries; +} + +describe("ZIP layer downloads", () => { + afterEach(() => { + vi.unstubAllGlobals(); + vi.restoreAllMocks(); + }); + + it("uses image signatures for entry extensions and keeps the original bytes", async () => { + const files = [ + { name: "portrait-face.image", url: "/png" }, + { name: "portrait-face.image", url: "/png-copy" }, + { name: "sketch.image", url: "/jpeg" }, + { name: "original.JPEG", url: "/jpeg-original" }, + { name: "another.JPE", url: "/jpeg-jpe" }, + { name: "mask.image", url: "/webp" }, + { name: "animated.image", url: "/gif" }, + { name: "bitmap.image", url: "/bmp" }, + { name: "icon.image", url: "/ico" }, + { name: "invalid-webp.image", url: "/invalid-webp" }, + { name: "unrecognized.image", url: "/unknown" } + ]; + const bytesByUrl = new Map([ + ["/png", png], ["/png-copy", png], ["/jpeg", jpeg], + ["/jpeg-original", jpeg], ["/jpeg-jpe", jpeg], ["/webp", webp], + ["/gif", gif], ["/bmp", bmp], ["/ico", ico], + ["/invalid-webp", Uint8Array.from([0x52, 0x49, 0x46, 0x46, 4, 0, 0, 0, 0x57, 0x45, 0x42, 0x50, 5])], + ["/unknown", Uint8Array.from([1, 2, 3])] + ]); + let archive: Blob | undefined; + const anchor = { href: "", download: "", click: vi.fn(), remove: vi.fn() }; + vi.stubGlobal("fetch", vi.fn(async (url: string) => new Response(bytesByUrl.get(url), { status: 200 }))); + vi.stubGlobal("document", { createElement: () => anchor, body: { appendChild: vi.fn() } }); + vi.stubGlobal("window", { setTimeout: vi.fn() }); + vi.stubGlobal("URL", { createObjectURL: (blob: Blob) => { archive = blob; return "blob:layers"; }, revokeObjectURL: vi.fn() }); + + await downloadFilesAsZip(files, "layers.zip"); + + expect(anchor.download).toBe("layers.zip"); + expect(anchor.click).toHaveBeenCalledOnce(); + expect(archive).toBeDefined(); + const entries = zipEntries(new Uint8Array(await archive!.arrayBuffer())); + expect(entries.map((entry) => entry.name)).toEqual([ + "portrait-face.png", "portrait-face-2.png", "sketch.jpg", "original.JPEG", "another.JPE", + "mask.webp", "animated.gif", "bitmap.bmp", "icon.ico", "invalid-webp.image", "unrecognized.image" + ]); + expect(entries.map((entry) => entry.bytes)).toEqual([ + png, png, jpeg, jpeg, jpeg, webp, gif, bmp, ico, + Uint8Array.from([0x52, 0x49, 0x46, 0x46, 4, 0, 0, 0, 0x57, 0x45, 0x42, 0x50, 5]), + Uint8Array.from([1, 2, 3]) + ]); + }); +}); diff --git a/vitest.config.ts b/vitest.config.ts index 3848a9a..b7543d6 100644 --- a/vitest.config.ts +++ b/vitest.config.ts @@ -1,6 +1,7 @@ import { defineConfig } from "vitest/config"; export default defineConfig({ + oxc: { jsx: { runtime: "automatic" } }, test: { environment: "node", include: ["tests/**/*.test.ts"]