feat: add scoped assets and storage core

This commit is contained in:
2026-08-13 16:30:15 +08:00
parent 5e60bb40e7
commit 7847c95538
10 changed files with 1345 additions and 0 deletions

View File

@@ -0,0 +1,348 @@
// Package assets owns asset visibility, metadata, and blob lifecycle policy.
package assets
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"path"
"regexp"
"strings"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/publicapi"
)
const PublicJobLookupLimit = 200
var (
ErrNotFound = errors.New("asset not found")
ErrBlobNotFound = errors.New("asset blob not found")
ErrUnsafeBlobKey = errors.New("unsafe blob key")
ErrRemoteProtocol = errors.New("remote asset protocol is not allowed")
ErrRemoteTooLarge = errors.New("remote asset exceeds size limit")
)
type Kind string
const (
KindImage Kind = "image"
KindVideo Kind = "video"
KindMask Kind = "mask"
KindReference Kind = "reference"
KindOther Kind = "other"
)
type Source string
const (
SourceUpload Source = "upload"
SourceGenerated Source = "generated"
SourceEdited Source = "edited"
SourceUpscaled Source = "upscaled"
SourceExternal Source = "external"
SourceSeed Source = "seed"
)
type Asset struct {
ID string `json:"id"`
OwnerID string `json:"ownerId"`
Kind Kind `json:"kind"`
Name string `json:"name"`
URL string `json:"url"`
StoragePath string `json:"storagePath,omitempty"`
Source Source `json:"source"`
Tags []string `json:"tags"`
Metadata map[string]any `json:"metadata"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
}
// MarshalJSON preserves JavaScript Date.toISOString() millisecond precision.
func (a Asset) MarshalJSON() ([]byte, error) {
type wireAsset struct {
ID string `json:"id"`
OwnerID string `json:"ownerId"`
Kind Kind `json:"kind"`
Name string `json:"name"`
URL string `json:"url"`
StoragePath string `json:"storagePath,omitempty"`
Source Source `json:"source"`
Tags []string `json:"tags"`
Metadata map[string]any `json:"metadata"`
CreatedAt string `json:"createdAt"`
UpdatedAt string `json:"updatedAt"`
}
return json.Marshal(wireAsset{
ID: a.ID, OwnerID: a.OwnerID, Kind: a.Kind, Name: a.Name, URL: a.URL,
StoragePath: a.StoragePath, Source: a.Source, Tags: a.Tags, Metadata: a.Metadata,
CreatedAt: a.CreatedAt.UTC().Format("2006-01-02T15:04:05.000Z"),
UpdatedAt: a.UpdatedAt.UTC().Format("2006-01-02T15:04:05.000Z"),
})
}
type scopeKind uint8
const (
platformScope scopeKind = iota
publicScope
)
type Scope struct {
kind scopeKind
ownerID, clientID string
}
func PlatformScope(ownerID string) Scope { return Scope{kind: platformScope, ownerID: ownerID} }
func PublicScope(clientID string) Scope {
return Scope{kind: publicScope, ownerID: publicapi.OwnerID(clientID), clientID: clientID}
}
func ClientTag(clientID string) string { return "api-client:" + clientID }
type Catalog interface {
ListOwner(context.Context, string) ([]Asset, error)
GetOwner(context.Context, string, string) (Asset, bool, error)
ListPublic(context.Context, string, string, int) ([]Asset, error)
GetPublic(context.Context, string, string, string, int) (Asset, bool, error)
Create(context.Context, Asset) (Asset, error)
DeleteOwner(context.Context, string, string) (Asset, bool, error)
}
type StoredObject struct{ Key, URL string }
type Blob struct {
Body io.ReadCloser
ContentType string
Size int64
}
type BlobStore interface {
Put(context.Context, string, io.Reader, int64, string) (StoredObject, error)
Read(context.Context, string) (Blob, error)
Delete(context.Context, string) error
}
type RemoteFetcher interface {
Fetch(context.Context, string) (Blob, error)
}
type Service struct {
catalog Catalog
blobs BlobStore
remote RemoteFetcher
now func() time.Time
id func(string) string
}
func NewService(c Catalog, b BlobStore, r RemoteFetcher, now func() time.Time, id func(string) string) *Service {
if now == nil {
now = time.Now
}
if id == nil {
id = func(prefix string) string { return fmt.Sprintf("%s-%d", prefix, now().UnixNano()) }
}
return &Service{catalog: c, blobs: b, remote: r, now: now, id: id}
}
func (s *Service) List(ctx context.Context, scope Scope) ([]Asset, error) {
if err := validScope(scope); err != nil {
return nil, err
}
if scope.kind == publicScope {
return s.catalog.ListPublic(ctx, scope.ownerID, scope.clientID, PublicJobLookupLimit)
}
return s.catalog.ListOwner(ctx, scope.ownerID)
}
func (s *Service) Get(ctx context.Context, scope Scope, id string) (Asset, error) {
if err := validScope(scope); err != nil {
return Asset{}, err
}
var a Asset
var found bool
var err error
if scope.kind == publicScope {
a, found, err = s.catalog.GetPublic(ctx, scope.ownerID, scope.clientID, id, PublicJobLookupLimit)
} else {
a, found, err = s.catalog.GetOwner(ctx, scope.ownerID, id)
}
if err != nil {
return Asset{}, err
}
if !found {
return Asset{}, ErrNotFound
}
return a, nil
}
type CreateExternalCommand struct {
URL, Name string
Kind Kind
Tags []string
}
func (s *Service) CreateExternal(ctx context.Context, scope Scope, cmd CreateExternalCommand) (Asset, error) {
if err := validScope(scope); err != nil {
return Asset{}, err
}
if strings.TrimSpace(cmd.URL) == "" {
return Asset{}, errors.New("url is required")
}
now := s.now().UTC()
kind := cmd.Kind
if kind == "" {
kind = KindImage
}
name := cmd.Name
tags := cloneStrings(cmd.Tags)
metadata := map[string]any{}
if scope.kind == publicScope {
name = defaultString(name, "外部素材")
tags = appendUnique(tags, "external", "public-api", ClientTag(scope.clientID))
metadata["registeredFrom"] = "public-api"
metadata["externalClientId"] = scope.clientID
} else {
name = defaultString(name, "外部图片")
if cmd.Tags == nil {
tags = []string{"external"}
}
metadata["registeredFrom"] = "api"
}
a := Asset{ID: s.id("asset"), OwnerID: scope.ownerID, Kind: kind, Name: name, URL: cmd.URL, Source: SourceExternal, Tags: tags, Metadata: metadata, CreatedAt: now, UpdatedAt: now}
return s.catalog.Create(ctx, a)
}
type UploadCommand struct {
Bytes []byte
FileName, ContentType, Origin string
Kind Kind
Tags []string
}
func (s *Service) Upload(ctx context.Context, scope Scope, cmd UploadCommand) (Asset, error) {
if err := validScope(scope); err != nil {
return Asset{}, err
}
if s.blobs == nil {
return Asset{}, errors.New("blob store is unavailable")
}
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)
if err != nil {
return Asset{}, err
}
kind := cmd.Kind
if kind == "" {
kind = inferKind(cmd.ContentType)
}
tags := cloneStrings(cmd.Tags)
if scope.kind == publicScope {
tags = appendUnique(tags, "upload", "public-api", ClientTag(scope.clientID))
} else if cmd.Tags == nil {
tags = []string{"upload"}
}
now := s.now().UTC()
a := Asset{ID: s.id("asset"), OwnerID: scope.ownerID, Kind: kind, Name: cmd.FileName, URL: stored.URL, StoragePath: stored.Key, Source: SourceUpload, Tags: tags, Metadata: map[string]any{"contentType": cmd.ContentType, "size": len(cmd.Bytes)}, CreatedAt: now, UpdatedAt: now}
created, err := s.catalog.Create(ctx, a)
if err != nil {
_ = s.blobs.Delete(context.WithoutCancel(ctx), stored.Key)
return Asset{}, err
}
return created, nil
}
func (s *Service) Delete(ctx context.Context, scope Scope, id string) (Asset, error) {
a, err := s.Get(ctx, scope, id)
if err != nil {
return Asset{}, err
}
if a.StoragePath != "" && s.blobs != nil {
if err := s.blobs.Delete(ctx, a.StoragePath); err != nil && !errors.Is(err, ErrBlobNotFound) {
return Asset{}, err
}
}
deleted, found, err := s.catalog.DeleteOwner(ctx, scope.ownerID, id)
if err != nil {
return Asset{}, err
}
if !found {
return Asset{}, ErrNotFound
}
return deleted, nil
}
func (s *Service) Download(ctx context.Context, scope Scope, id string) (Blob, error) {
a, err := s.Get(ctx, scope, id)
if err != nil {
return Blob{}, err
}
if a.StoragePath != "" && s.blobs != nil {
b, e := s.blobs.Read(ctx, a.StoragePath)
if e == nil {
return b, nil
}
if !errors.Is(e, ErrBlobNotFound) {
return Blob{}, e
}
}
if s.remote == nil {
return Blob{}, ErrBlobNotFound
}
return s.remote.Fetch(ctx, a.URL)
}
func validScope(s Scope) error {
if strings.TrimSpace(s.ownerID) == "" {
return errors.New("asset owner is required")
}
if s.kind == publicScope && strings.TrimSpace(s.clientID) == "" {
return errors.New("public client is required")
}
return nil
}
func defaultString(v, d string) string {
if strings.TrimSpace(v) == "" {
return d
}
return v
}
func appendUnique(values []string, add ...string) []string {
for _, v := range add {
if !contains(values, v) {
values = append(values, v)
}
}
return values
}
func cloneStrings(values []string) []string {
if values == nil {
return nil
}
return append([]string{}, values...)
}
func contains(values []string, want string) bool {
for _, v := range values {
if v == want {
return true
}
}
return false
}
var unsafeName = regexp.MustCompile(`[^A-Za-z0-9._-]+`)
func sanitizeFileName(v string) string {
v = strings.Trim(unsafeName.ReplaceAllString(v, "-"), "-")
if v == "" {
return "asset.bin"
}
return v
}
func inferKind(contentType string) Kind {
if strings.HasPrefix(contentType, "image/") {
return KindImage
}
if strings.HasPrefix(contentType, "video/") {
return KindVideo
}
return KindOther
}

View File

@@ -0,0 +1,178 @@
package assets
import (
"context"
"errors"
"fmt"
"io"
"mime"
"net/url"
"os"
"path/filepath"
"strings"
)
type LocalFS struct{ root, baseURL string }
func NewLocalFS(root, baseURL string) (*LocalFS, error) {
abs, err := filepath.Abs(root)
if err != nil {
return nil, err
}
if err = os.MkdirAll(abs, 0750); err != nil {
return nil, err
}
info, err := os.Lstat(abs)
if err != nil {
return nil, err
}
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
return nil, fmt.Errorf("%w: root", ErrUnsafeBlobKey)
}
return &LocalFS{root: abs, baseURL: strings.TrimRight(baseURL, "/")}, nil
}
func (s *LocalFS) Put(_ context.Context, key string, body io.Reader, _ int64, contentType string) (StoredObject, error) {
full, parts, err := s.resolve(key)
if err != nil {
return StoredObject{}, err
}
if err = s.ensureParents(parts[:len(parts)-1]); err != nil {
return StoredObject{}, err
}
if info, e := os.Lstat(full); e == nil && info.Mode()&os.ModeSymlink != 0 {
return StoredObject{}, ErrUnsafeBlobKey
} else if e != nil && !errors.Is(e, os.ErrNotExist) {
return StoredObject{}, e
}
tmp, err := os.CreateTemp(filepath.Dir(full), ".asset-*")
if err != nil {
return StoredObject{}, err
}
tmpName := tmp.Name()
defer os.Remove(tmpName)
if _, err = io.Copy(tmp, body); err == nil {
err = tmp.Sync()
}
if closeErr := tmp.Close(); err == nil {
err = closeErr
}
if err != nil {
return StoredObject{}, err
}
if err = os.Rename(tmpName, full); err != nil {
return StoredObject{}, err
}
_ = contentType
return StoredObject{Key: strings.Join(parts, "/"), URL: s.baseURL + "/" + escapeParts(parts)}, nil
}
func (s *LocalFS) Read(_ context.Context, key string) (Blob, error) {
full, parts, err := s.resolve(key)
if err != nil {
return Blob{}, err
}
if err = s.rejectSymlinks(parts); err != nil {
if errors.Is(err, os.ErrNotExist) {
return Blob{}, ErrBlobNotFound
}
return Blob{}, err
}
f, err := os.Open(full)
if errors.Is(err, os.ErrNotExist) {
return Blob{}, ErrBlobNotFound
}
if err != nil {
return Blob{}, err
}
info, err := f.Stat()
if err != nil {
f.Close()
return Blob{}, err
}
if !info.Mode().IsRegular() {
f.Close()
return Blob{}, ErrUnsafeBlobKey
}
contentType := mime.TypeByExtension(filepath.Ext(full))
if contentType == "" {
contentType = "application/octet-stream"
}
return Blob{Body: f, ContentType: contentType, Size: info.Size()}, nil
}
func (s *LocalFS) Delete(_ context.Context, key string) error {
full, parts, err := s.resolve(key)
if err != nil {
return err
}
if err = s.rejectSymlinks(parts); errors.Is(err, os.ErrNotExist) {
return nil
} else if err != nil {
return err
}
err = os.Remove(full)
if errors.Is(err, os.ErrNotExist) {
return nil
}
return err
}
func (s *LocalFS) resolve(key string) (string, []string, error) {
if key == "" || filepath.IsAbs(key) || strings.Contains(key, "\\") {
return "", nil, ErrUnsafeBlobKey
}
clean := filepath.ToSlash(filepath.Clean(key))
if clean == "." || clean == ".." || strings.HasPrefix(clean, "../") {
return "", nil, ErrUnsafeBlobKey
}
parts := strings.Split(clean, "/")
for _, p := range parts {
if p == "" || p == "." || p == ".." {
return "", nil, ErrUnsafeBlobKey
}
}
full := filepath.Join(append([]string{s.root}, parts...)...)
rel, err := filepath.Rel(s.root, full)
if err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
return "", nil, ErrUnsafeBlobKey
}
return full, parts, nil
}
func (s *LocalFS) ensureParents(parts []string) error {
current := s.root
for _, part := range parts {
current = filepath.Join(current, part)
info, err := os.Lstat(current)
if errors.Is(err, os.ErrNotExist) {
if err = os.Mkdir(current, 0750); err != nil {
return err
}
continue
}
if err != nil {
return err
}
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
return ErrUnsafeBlobKey
}
}
return nil
}
func (s *LocalFS) rejectSymlinks(parts []string) error {
current := s.root
for _, part := range parts {
current = filepath.Join(current, part)
info, err := os.Lstat(current)
if err != nil {
return err
}
if info.Mode()&os.ModeSymlink != 0 {
return ErrUnsafeBlobKey
}
}
return nil
}
func escapeParts(parts []string) string {
out := make([]string, len(parts))
for i, p := range parts {
out[i] = url.PathEscape(p)
}
return strings.Join(out, "/")
}

View File

@@ -0,0 +1,73 @@
package assets
import (
"bytes"
"context"
"errors"
"io"
"os"
"path/filepath"
"runtime"
"testing"
)
func TestLocalFSPutReadDelete(t *testing.T) {
store, err := NewLocalFS(t.TempDir(), "https://app.test")
if err != nil {
t.Fatal(err)
}
stored, err := store.Put(context.Background(), "uploads/day/a file.png", bytes.NewReader([]byte("asset")), 5, "image/png")
if err != nil {
t.Fatal(err)
}
if stored.URL != "https://app.test/uploads/day/a%20file.png" {
t.Fatalf("URL = %q", stored.URL)
}
blob, err := store.Read(context.Background(), stored.Key)
if err != nil {
t.Fatal(err)
}
b, _ := io.ReadAll(blob.Body)
blob.Body.Close()
if string(b) != "asset" || blob.ContentType != "image/png" {
t.Fatalf("blob = %q %#v", b, blob)
}
if err := store.Delete(context.Background(), stored.Key); err != nil {
t.Fatal(err)
}
if _, err := store.Read(context.Background(), stored.Key); !errors.Is(err, ErrBlobNotFound) {
t.Fatalf("Read deleted error = %v", err)
}
}
func TestLocalFSRejectsTraversalAbsoluteAndSymlinks(t *testing.T) {
root := t.TempDir()
outside := t.TempDir()
store, err := NewLocalFS(root, "http://local.test")
if err != nil {
t.Fatal(err)
}
for _, key := range []string{"../escape", "uploads/../../escape", filepath.Join(outside, "escape"), "uploads\\..\\escape"} {
if _, err := store.Put(context.Background(), key, bytes.NewReader(nil), 0, "application/octet-stream"); !errors.Is(err, ErrUnsafeBlobKey) {
t.Errorf("Put(%q) error = %v", key, err)
}
}
if runtime.GOOS == "windows" {
return
}
if err := os.Symlink(outside, filepath.Join(root, "linked")); err != nil {
t.Fatal(err)
}
if _, err := store.Put(context.Background(), "linked/escape", bytes.NewReader([]byte("x")), 1, "text/plain"); !errors.Is(err, ErrUnsafeBlobKey) {
t.Fatalf("symlink directory error = %v", err)
}
if err := os.WriteFile(filepath.Join(outside, "target"), []byte("secret"), 0600); err != nil {
t.Fatal(err)
}
if err := os.Symlink(filepath.Join(outside, "target"), filepath.Join(root, "final")); err != nil {
t.Fatal(err)
}
if _, err := store.Read(context.Background(), "final"); !errors.Is(err, ErrUnsafeBlobKey) {
t.Fatalf("symlink file error = %v", err)
}
}

View File

@@ -0,0 +1,69 @@
package assets
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"net/http"
"net/url"
)
type HTTPRemoteFetcher struct {
client *http.Client
maxBytes int64
}
func NewHTTPRemoteFetcher(client *http.Client, maxBytes int64) *HTTPRemoteFetcher {
if client == nil {
client = http.DefaultClient
}
base := *client
previousRedirect := base.CheckRedirect
base.CheckRedirect = func(req *http.Request, via []*http.Request) error {
if req.URL.Scheme != "http" && req.URL.Scheme != "https" {
return ErrRemoteProtocol
}
if previousRedirect != nil {
return previousRedirect(req, via)
}
if len(via) >= 10 {
return errors.New("stopped after 10 redirects")
}
return nil
}
return &HTTPRemoteFetcher{client: &base, maxBytes: maxBytes}
}
func (f *HTTPRemoteFetcher) Fetch(ctx context.Context, rawURL string) (Blob, error) {
u, err := url.Parse(rawURL)
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" {
return Blob{}, ErrRemoteProtocol
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil)
if err != nil {
return Blob{}, err
}
resp, err := f.client.Do(req)
if err != nil {
return Blob{}, err
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return Blob{}, fmt.Errorf("remote asset returned HTTP %d", resp.StatusCode)
}
if f.maxBytes <= 0 {
return Blob{}, errors.New("remote size limit must be positive")
}
if resp.ContentLength > f.maxBytes {
return Blob{}, ErrRemoteTooLarge
}
b, err := io.ReadAll(io.LimitReader(resp.Body, f.maxBytes+1))
if err != nil {
return Blob{}, err
}
if int64(len(b)) > f.maxBytes {
return Blob{}, ErrRemoteTooLarge
}
return Blob{Body: io.NopCloser(bytes.NewReader(b)), ContentType: resp.Header.Get("Content-Type"), Size: int64(len(b))}, nil
}

View File

@@ -0,0 +1,42 @@
package assets
import (
"context"
"errors"
"io"
"net/http"
"net/http/httptest"
"testing"
)
func TestHTTPRemoteFetcherRestrictsProtocolAndSize(t *testing.T) {
fetcher := NewHTTPRemoteFetcher(http.DefaultClient, 4)
if _, err := fetcher.Fetch(context.Background(), "file:///etc/passwd"); !errors.Is(err, ErrRemoteProtocol) {
t.Fatalf("file protocol error = %v", err)
}
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "image/png")
_, _ = io.WriteString(w, "12345")
}))
defer server.Close()
if _, err := fetcher.Fetch(context.Background(), server.URL); !errors.Is(err, ErrRemoteTooLarge) {
t.Fatalf("oversize error = %v", err)
}
}
func TestHTTPRemoteFetcherReturnsBoundedBody(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "image/png")
_, _ = io.WriteString(w, "1234")
}))
defer server.Close()
blob, err := NewHTTPRemoteFetcher(http.DefaultClient, 4).Fetch(context.Background(), server.URL)
if err != nil {
t.Fatal(err)
}
defer blob.Body.Close()
body, _ := io.ReadAll(blob.Body)
if string(body) != "1234" || blob.Size != 4 || blob.ContentType != "image/png" {
t.Fatalf("blob = %#v body=%q", blob, body)
}
}

View File

@@ -0,0 +1,194 @@
package assets
import (
"bytes"
"context"
"errors"
"io"
"reflect"
"testing"
"time"
)
func TestPlatformScopeOwnerIsolationAndDefaults(t *testing.T) {
cat := &memoryCatalog{}
svc := NewService(cat, nil, nil, func() time.Time { return time.Date(2026, 8, 13, 8, 0, 0, 0, time.UTC) }, func(string) string { return "asset-1" })
scope := PlatformScope("owner-a")
created, err := svc.CreateExternal(context.Background(), scope, CreateExternalCommand{URL: "https://example.test/a.png"})
if err != nil {
t.Fatal(err)
}
if created.OwnerID != "owner-a" || created.Name != "外部图片" || created.Kind != KindImage || created.Source != SourceExternal || !reflect.DeepEqual(created.Tags, []string{"external"}) || created.Metadata["registeredFrom"] != "api" {
t.Fatalf("created = %#v", created)
}
if _, err := svc.Get(context.Background(), PlatformScope("owner-b"), created.ID); !errors.Is(err, ErrNotFound) {
t.Fatalf("cross-owner Get error = %v", err)
}
if _, err := svc.Delete(context.Background(), PlatformScope("owner-b"), created.ID); !errors.Is(err, ErrNotFound) {
t.Fatalf("cross-owner Delete error = %v", err)
}
}
func TestPlatformExplicitEmptyTagsArePreserved(t *testing.T) {
cat := &memoryCatalog{}
svc := NewService(cat, nil, nil, time.Now, func(string) string { return "asset-empty-tags" })
created, err := svc.CreateExternal(context.Background(), PlatformScope("owner-a"), CreateExternalCommand{URL: "https://example.test/a.png", Tags: []string{}})
if err != nil {
t.Fatal(err)
}
if created.Tags == nil || len(created.Tags) != 0 {
t.Fatalf("tags = %#v, want explicit empty array", created.Tags)
}
}
func TestPublicScopeAddsClientIdentityAndUsesCatalogVisibility(t *testing.T) {
cat := &memoryCatalog{}
svc := NewService(cat, nil, nil, time.Now, func(string) string { return "asset-public" })
created, err := svc.CreateExternal(context.Background(), PublicScope("agent-a"), CreateExternalCommand{URL: "https://example.test/a.png", Tags: []string{"campaign"}})
if err != nil {
t.Fatal(err)
}
wantTags := []string{"campaign", "external", "public-api", "api-client:agent-a"}
if created.OwnerID != "api:agent-a" || created.Name != "外部素材" || !reflect.DeepEqual(created.Tags, wantTags) || created.Metadata["externalClientId"] != "agent-a" || cat.lastPublicLimit != 0 {
t.Fatalf("created = %#v", created)
}
cat.publicVisible = map[string]bool{created.ID: true}
if _, err := svc.Get(context.Background(), PublicScope("agent-a"), created.ID); err != nil {
t.Fatal(err)
}
if cat.lastPublicClient != "agent-a" || cat.lastPublicLimit != PublicJobLookupLimit {
t.Fatalf("public lookup = %q, %d", cat.lastPublicClient, cat.lastPublicLimit)
}
}
func TestUploadCompensatesBlobWhenCatalogCreateFails(t *testing.T) {
cat := &memoryCatalog{createErr: errors.New("database down")}
blobs := &memoryBlobs{}
svc := NewService(cat, blobs, nil, time.Now, func(prefix string) string { return prefix + "-1" })
_, err := svc.Upload(context.Background(), PlatformScope("owner-a"), UploadCommand{Bytes: []byte("png"), FileName: "poster.png", ContentType: "image/png", Origin: "https://app.test"})
if !errors.Is(err, cat.createErr) {
t.Fatalf("Upload error = %v", err)
}
if len(blobs.deleted) != 1 || blobs.deleted[0] != blobs.putKey {
t.Fatalf("compensating deletes = %#v, put key %q", blobs.deleted, blobs.putKey)
}
}
func TestDeleteLeavesCatalogWhenBlobDeletionFails(t *testing.T) {
cat := &memoryCatalog{assets: []Asset{{ID: "a", OwnerID: "o", StoragePath: "uploads/a"}}}
blobs := &memoryBlobs{deleteErr: errors.New("storage down")}
svc := NewService(cat, blobs, nil, time.Now, nil)
if _, err := svc.Delete(context.Background(), PlatformScope("o"), "a"); !errors.Is(err, blobs.deleteErr) {
t.Fatalf("Delete error = %v", err)
}
if _, found, _ := cat.GetOwner(context.Background(), "o", "a"); !found {
t.Fatal("catalog row deleted after blob failure")
}
}
func TestDownloadUsesBlobBeforeRemoteFetcher(t *testing.T) {
cat := &memoryCatalog{assets: []Asset{{ID: "a", OwnerID: "o", Name: "a.png", URL: "https://remote.test/a.png", StoragePath: "uploads/a.png", Metadata: map[string]any{"contentType": "image/png"}}}}
blobs := &memoryBlobs{read: Blob{Body: io.NopCloser(bytes.NewReader([]byte("local"))), ContentType: "image/png", Size: 5}}
remote := &memoryRemote{blob: Blob{Body: io.NopCloser(bytes.NewReader([]byte("remote")))}}
svc := NewService(cat, blobs, remote, time.Now, nil)
got, err := svc.Download(context.Background(), PlatformScope("o"), "a")
if err != nil {
t.Fatal(err)
}
defer got.Body.Close()
b, _ := io.ReadAll(got.Body)
if string(b) != "local" || remote.called {
t.Fatalf("download = %q remoteCalled=%v", b, remote.called)
}
}
type memoryCatalog struct {
assets []Asset
createErr error
publicVisible map[string]bool
lastPublicClient string
lastPublicLimit int
}
func (m *memoryCatalog) ListOwner(_ context.Context, owner string) ([]Asset, error) {
var out []Asset
for _, a := range m.assets {
if a.OwnerID == owner {
out = append(out, a)
}
}
return out, nil
}
func (m *memoryCatalog) GetOwner(_ context.Context, owner, id string) (Asset, bool, error) {
for _, a := range m.assets {
if a.OwnerID == owner && a.ID == id {
return a, true, nil
}
}
return Asset{}, false, nil
}
func (m *memoryCatalog) ListPublic(ctx context.Context, owner, client string, limit int) ([]Asset, error) {
m.lastPublicClient = client
m.lastPublicLimit = limit
all, _ := m.ListOwner(ctx, owner)
var out []Asset
for _, a := range all {
if m.publicVisible[a.ID] || contains(a.Tags, ClientTag(client)) {
out = append(out, a)
}
}
return out, nil
}
func (m *memoryCatalog) GetPublic(_ context.Context, owner, client, id string, limit int) (Asset, bool, error) {
m.lastPublicClient = client
m.lastPublicLimit = limit
for _, a := range m.assets {
if a.OwnerID == owner && a.ID == id && (m.publicVisible[id] || contains(a.Tags, ClientTag(client))) {
return a, true, nil
}
}
return Asset{}, false, nil
}
func (m *memoryCatalog) Create(_ context.Context, a Asset) (Asset, error) {
if m.createErr != nil {
return Asset{}, m.createErr
}
m.assets = append(m.assets, a)
return a, nil
}
func (m *memoryCatalog) DeleteOwner(_ context.Context, owner, id string) (Asset, bool, error) {
for i, a := range m.assets {
if a.OwnerID == owner && a.ID == id {
m.assets = append(m.assets[:i], m.assets[i+1:]...)
return a, true, nil
}
}
return Asset{}, false, nil
}
type memoryBlobs struct {
putKey string
deleted []string
read Blob
deleteErr error
}
func (m *memoryBlobs) Put(_ context.Context, key string, _ io.Reader, _ int64, _ string) (StoredObject, error) {
m.putKey = key
return StoredObject{Key: key, URL: "https://app.test/" + key}, nil
}
func (m *memoryBlobs) Read(context.Context, string) (Blob, error) { return m.read, nil }
func (m *memoryBlobs) Delete(_ context.Context, key string) error {
m.deleted = append(m.deleted, key)
return m.deleteErr
}
type memoryRemote struct {
called bool
blob Blob
}
func (m *memoryRemote) Fetch(context.Context, string) (Blob, error) {
m.called = true
return m.blob, nil
}

View File

@@ -0,0 +1,193 @@
package postgres
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"os"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/assets"
)
const assetFields = `a.id,
a.owner_id,
a.kind,
a.name,
a.url,
a.storage_path,
a.source,
a.tags,
a.metadata,
a.created_at,
a.updated_at`
const ListOwnerAssetsSQL = `SELECT ` + assetFields + `
FROM public.assets AS a
WHERE a.owner_id = $1::text
ORDER BY a.created_at DESC`
const GetOwnerAssetSQL = `SELECT ` + assetFields + `
FROM public.assets AS a
WHERE a.owner_id = $1::text AND a.id = $2::text
LIMIT 1`
const listPublicAccessibleAssetIDs = `SELECT unnest(j.input_asset_ids || j.output_asset_ids) AS asset_id
FROM (
SELECT input_asset_ids, output_asset_ids
FROM public.generation_jobs
WHERE owner_id = $1::text AND external_client_id = $3::text
ORDER BY created_at DESC
LIMIT $4::integer
) AS j`
const ListPublicAssetsSQL = `SELECT ` + assetFields + `
FROM public.assets AS a
WHERE a.owner_id = $1::text
AND ($2::text = ANY(a.tags) OR a.id IN (` + listPublicAccessibleAssetIDs + `))
ORDER BY a.created_at DESC`
const getPublicAccessibleAssetIDs = `SELECT unnest(j.input_asset_ids || j.output_asset_ids) AS asset_id
FROM (
SELECT input_asset_ids, output_asset_ids
FROM public.generation_jobs
WHERE owner_id = $1::text AND external_client_id = $4::text
ORDER BY created_at DESC
LIMIT $5::integer
) AS j`
const GetPublicAssetSQL = `SELECT ` + assetFields + `
FROM public.assets AS a
WHERE a.owner_id = $1::text
AND a.id = $2::text
AND ($3::text = ANY(a.tags) OR a.id IN (` + getPublicAccessibleAssetIDs + `))
LIMIT 1`
const CreateAssetSQL = `INSERT INTO public.assets (
id, owner_id, kind, name, url, storage_path, source, tags, metadata, created_at, updated_at
) VALUES ($1::text, $2::text, $3::text, $4::text, $5::text, $6::text, $7::text, $8::text[], $9::jsonb, $10::timestamptz, $11::timestamptz)
RETURNING id, owner_id, kind, name, url, storage_path, source, tags, metadata, created_at, updated_at`
const DeleteOwnerAssetSQL = `DELETE FROM public.assets
WHERE owner_id = $1::text AND id = $2::text
RETURNING id, owner_id, kind, name, url, storage_path, source, tags, metadata, created_at, updated_at`
func (db *Database) ListOwner(ctx context.Context, owner string) ([]assets.Asset, error) {
return db.listAssets(ctx, ListOwnerAssetsSQL, owner)
}
func (db *Database) GetOwner(ctx context.Context, owner, id string) (assets.Asset, bool, error) {
return db.oneAsset(ctx, GetOwnerAssetSQL, owner, id)
}
func (db *Database) ListPublic(ctx context.Context, owner, client string, limit int) ([]assets.Asset, error) {
return db.listAssets(ctx, ListPublicAssetsSQL, owner, assets.ClientTag(client), client, limit)
}
func (db *Database) GetPublic(ctx context.Context, owner, client, id string, limit int) (assets.Asset, bool, error) {
return db.oneAsset(ctx, GetPublicAssetSQL, owner, id, assets.ClientTag(client), client, limit)
}
func (db *Database) Create(ctx context.Context, a assets.Asset) (assets.Asset, error) {
if err := db.available(); err != nil {
return assets.Asset{}, err
}
if a.Tags == nil {
a.Tags = []string{}
}
if a.Metadata == nil {
a.Metadata = map[string]any{}
}
metadata, err := json.Marshal(a.Metadata)
if err != nil {
return assets.Asset{}, fmt.Errorf("encode asset metadata: %w", err)
}
return db.requiredAsset(ctx, CreateAssetSQL, a.ID, a.OwnerID, string(a.Kind), a.Name, a.URL, optionalDatabaseText(a.StoragePath), string(a.Source), a.Tags, metadata, a.CreatedAt, a.UpdatedAt)
}
func (db *Database) DeleteOwner(ctx context.Context, owner, id string) (assets.Asset, bool, error) {
return db.oneAsset(ctx, DeleteOwnerAssetSQL, owner, id)
}
func (db *Database) available() error {
if db.config.Backend != BackendPostgres || db.querier == nil {
return fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend)
}
return nil
}
func (db *Database) listAssets(ctx context.Context, query string, args ...any) ([]assets.Asset, error) {
if err := db.available(); err != nil {
return nil, err
}
rows, err := db.querier.Query(ctx, query, args...)
if err != nil {
return nil, fmt.Errorf("query assets: %w", err)
}
defer rows.Close()
out := []assets.Asset{}
for rows.Next() {
a, err := scanAsset(rows)
if err != nil {
return nil, err
}
out = append(out, a)
}
if err = rows.Err(); err != nil {
return nil, fmt.Errorf("read assets: %w", err)
}
return out, nil
}
func (db *Database) oneAsset(ctx context.Context, query string, args ...any) (assets.Asset, bool, error) {
if err := db.available(); err != nil {
return assets.Asset{}, false, err
}
rows, err := db.querier.Query(ctx, query, args...)
if err != nil {
return assets.Asset{}, false, fmt.Errorf("query asset: %w", err)
}
defer rows.Close()
if !rows.Next() {
if err = rows.Err(); err != nil {
return assets.Asset{}, false, fmt.Errorf("read asset: %w", err)
}
return assets.Asset{}, false, nil
}
a, err := scanAsset(rows)
if err != nil {
return assets.Asset{}, false, err
}
if err = rows.Err(); err != nil {
return assets.Asset{}, false, fmt.Errorf("read asset: %w", err)
}
return a, true, nil
}
func (db *Database) requiredAsset(ctx context.Context, query string, args ...any) (assets.Asset, error) {
a, found, err := db.oneAsset(ctx, query, args...)
if err != nil {
return assets.Asset{}, err
}
if !found {
return assets.Asset{}, fmt.Errorf("asset write returned no row")
}
return a, nil
}
func scanAsset(rows Rows) (assets.Asset, error) {
var a assets.Asset
var kind, source string
var storage sql.NullString
var metadata []byte
var created, updated time.Time
if err := rows.Scan(&a.ID, &a.OwnerID, &kind, &a.Name, &a.URL, &storage, &source, &a.Tags, &metadata, &created, &updated); err != nil {
return assets.Asset{}, fmt.Errorf("scan asset: %w", err)
}
a.Kind = assets.Kind(kind)
a.Source = assets.Source(source)
if storage.Valid {
a.StoragePath = storage.String
}
a.CreatedAt = created
a.UpdatedAt = updated
if len(metadata) == 0 {
metadata = []byte(`{}`)
}
if err := json.Unmarshal(metadata, &a.Metadata); err != nil {
return assets.Asset{}, fmt.Errorf("decode asset metadata: %w", err)
}
if a.Tags == nil {
a.Tags = []string{}
}
if a.Metadata == nil {
a.Metadata = map[string]any{}
}
return a, nil
}
func readContractFixture(path string) ([]byte, error) { return os.ReadFile(path) }

View File

@@ -0,0 +1,172 @@
package postgres
import (
"context"
"database/sql"
"encoding/json"
"errors"
"reflect"
"testing"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/assets"
)
func TestAssetCatalogUsesExplicitOwnerScopedQueries(t *testing.T) {
now := time.Date(2026, 8, 13, 8, 0, 0, 0, time.UTC)
row := []any{"asset-1", "owner-1", "image", "poster.png", "https://cdn.test/a.png", "uploads/a.png", "upload", []string{"upload"}, []byte(`{"size":4}`), now, now}
tests := []struct {
name string
call func(*Database) (assets.Asset, bool, error)
wantSQL string
wantArgs []any
}{
{name: "get owner", call: func(db *Database) (assets.Asset, bool, error) {
return db.GetOwner(context.Background(), "owner-1", "asset-1")
}, wantSQL: GetOwnerAssetSQL, wantArgs: []any{"owner-1", "asset-1"}},
{name: "get public", call: func(db *Database) (assets.Asset, bool, error) {
return db.GetPublic(context.Background(), "api:agent-a", "agent-a", "asset-1", 200)
}, wantSQL: GetPublicAssetSQL, wantArgs: []any{"api:agent-a", "asset-1", "api-client:agent-a", "agent-a", 200}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
q := &assetQuerier{rows: &assetRows{rows: [][]any{row}}}
db := NewDatabase(Config{Backend: BackendPostgres}, q)
got, found, err := tt.call(db)
if err != nil || !found {
t.Fatalf("result = %#v,%v,%v", got, found, err)
}
if q.sql != tt.wantSQL || !reflect.DeepEqual(q.args, tt.wantArgs) {
t.Fatalf("query=%q args=%#v", q.sql, q.args)
}
if got.Metadata["size"] != float64(4) || got.StoragePath != "uploads/a.png" {
t.Fatalf("asset=%#v", got)
}
})
}
}
func TestAssetCatalogCreateDefaultsCollectionsAndUsesAllFields(t *testing.T) {
now := time.Date(2026, 8, 13, 8, 0, 0, 0, time.UTC)
q := &assetQuerier{rows: &assetRows{rows: [][]any{{"a", "o", "image", "n", "u", nil, "external", []string{}, []byte(`{}`), now, now}}}}
db := NewDatabase(Config{Backend: BackendPostgres}, q)
got, err := db.Create(context.Background(), assets.Asset{ID: "a", OwnerID: "o", Kind: assets.KindImage, Name: "n", URL: "u", Source: assets.SourceExternal, CreatedAt: now, UpdatedAt: now})
if err != nil {
t.Fatal(err)
}
if q.sql != CreateAssetSQL || got.Tags == nil || got.Metadata == nil {
t.Fatalf("query=%q asset=%#v", q.sql, got)
}
if !reflect.DeepEqual(q.args[7], []string{}) || string(q.args[8].([]byte)) != "{}" {
t.Fatalf("args=%#v", q.args)
}
}
func TestAssetCatalogDeleteIsOwnerScopedAndNotFound(t *testing.T) {
q := &assetQuerier{rows: &assetRows{}}
db := NewDatabase(Config{Backend: BackendPostgres}, q)
_, found, err := db.DeleteOwner(context.Background(), "owner", "asset")
if err != nil || found {
t.Fatalf("found=%v err=%v", found, err)
}
if q.sql != DeleteOwnerAssetSQL || !reflect.DeepEqual(q.args, []any{"owner", "asset"}) {
t.Fatalf("query=%q args=%#v", q.sql, q.args)
}
}
func TestAssetCatalogFailsClosedWithoutPostgres(t *testing.T) {
db := NewDatabase(Config{Backend: BackendLocal}, nil)
if _, err := db.ListOwner(context.Background(), "owner"); err == nil {
t.Fatal("expected error")
}
}
var _ assets.Catalog = (*Database)(nil)
type assetQuerier struct {
rows *assetRows
err error
sql string
args []any
}
func (q *assetQuerier) Query(_ context.Context, sql string, args ...any) (Rows, error) {
q.sql = sql
q.args = args
return q.rows, q.err
}
type assetRows struct {
rows [][]any
idx int
err error
}
func (r *assetRows) Close() {}
func (r *assetRows) Err() error { return r.err }
func (r *assetRows) Next() bool { return r.idx < len(r.rows) }
func (r *assetRows) Scan(dest ...any) error {
if r.idx >= len(r.rows) {
return errors.New("past end")
}
row := r.rows[r.idx]
r.idx++
if len(row) != len(dest) {
return errors.New("arity")
}
for i, v := range row {
switch d := dest[i].(type) {
case *string:
if v == nil {
*d = ""
} else {
*d = v.(string)
}
case *sql.NullString:
if v == nil {
*d = sql.NullString{}
} else {
*d = sql.NullString{String: v.(string), Valid: true}
}
case *[]string:
*d = append([]string(nil), v.([]string)...)
case *[]byte:
*d = append([]byte(nil), v.([]byte)...)
case *time.Time:
*d = v.(time.Time)
default:
return errors.New("unsupported")
}
}
return nil
}
func TestAssetWireContractFixtureDecodes(t *testing.T) {
raw, err := readContractFixture("../../../contracts/assets/assets-v1.json")
if err != nil {
t.Fatal(err)
}
var fixture struct {
Asset assets.Asset `json:"asset"`
PublicVisibility struct {
JobLookupLimit int `json:"jobLookupLimit"`
} `json:"publicVisibility"`
}
if err = json.Unmarshal(raw, &fixture); err != nil {
t.Fatal(err)
}
if fixture.Asset.ID != "asset-contract-1" || fixture.PublicVisibility.JobLookupLimit != assets.PublicJobLookupLimit {
t.Fatalf("fixture=%#v", fixture)
}
wire, err := json.Marshal(fixture.Asset)
if err != nil {
t.Fatal(err)
}
var encoded map[string]any
if err := json.Unmarshal(wire, &encoded); err != nil {
t.Fatal(err)
}
if encoded["createdAt"] != "2026-08-13T08:00:00.000Z" || encoded["updatedAt"] != "2026-08-13T08:00:00.000Z" {
t.Fatalf("wire timestamps = %q, %q", encoded["createdAt"], encoded["updatedAt"])
}
}

View File

@@ -0,0 +1,28 @@
{
"version": 1,
"asset": {
"id": "asset-contract-1",
"ownerId": "owner-1",
"kind": "image",
"name": "poster.png",
"url": "https://cdn.example.test/uploads/poster.png",
"storagePath": "uploads/2026-08-13/poster.png",
"source": "upload",
"tags": ["upload"],
"metadata": { "contentType": "image/png", "size": 4 },
"createdAt": "2026-08-13T08:00:00.000Z",
"updatedAt": "2026-08-13T08:00:00.000Z"
},
"defaults": {
"platformExternalName": "外部图片",
"publicExternalName": "外部素材",
"platformExternalTags": ["external"],
"publicRequiredTags": ["external", "public-api", "api-client:agent-a"],
"platformRegisteredFrom": "api",
"publicRegisteredFrom": "public-api"
},
"publicVisibility": {
"jobLookupLimit": 200,
"clientTagPrefix": "api-client:"
}
}

View File

@@ -0,0 +1,48 @@
import { readFile } from "node:fs/promises";
import { describe, expect, it } from "vitest";
import type { Asset } from "@/lib/types";
type AssetsContract = {
version: number;
asset: Asset;
defaults: {
platformExternalName: string;
publicExternalName: string;
platformExternalTags: string[];
publicRequiredTags: string[];
platformRegisteredFrom: string;
publicRegisteredFrom: string;
};
publicVisibility: { jobLookupLimit: number; clientTagPrefix: string };
};
describe("assets v1 compatibility contract", () => {
it("freezes the TypeScript wire shape, defaults, and public visibility rule", async () => {
const contract = JSON.parse(await readFile(new URL("../contracts/assets/assets-v1.json", import.meta.url), "utf8")) as AssetsContract;
const asset: Asset = contract.asset;
expect(contract.version).toBe(1);
expect(asset).toEqual({
id: "asset-contract-1",
ownerId: "owner-1",
kind: "image",
name: "poster.png",
url: "https://cdn.example.test/uploads/poster.png",
storagePath: "uploads/2026-08-13/poster.png",
source: "upload",
tags: ["upload"],
metadata: { contentType: "image/png", size: 4 },
createdAt: "2026-08-13T08:00:00.000Z",
updatedAt: "2026-08-13T08:00:00.000Z"
});
expect(contract.defaults).toEqual({
platformExternalName: "外部图片",
publicExternalName: "外部素材",
platformExternalTags: ["external"],
publicRequiredTags: ["external", "public-api", "api-client:agent-a"],
platformRegisteredFrom: "api",
publicRegisteredFrom: "public-api"
});
expect(contract.publicVisibility).toEqual({ jobLookupLimit: 200, clientTagPrefix: "api-client:" });
});
});