feat: add scoped assets and storage core
This commit is contained in:
348
backend/internal/assets/assets.go
Normal file
348
backend/internal/assets/assets.go
Normal 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
|
||||
}
|
||||
178
backend/internal/assets/localfs.go
Normal file
178
backend/internal/assets/localfs.go
Normal 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, "/")
|
||||
}
|
||||
73
backend/internal/assets/localfs_test.go
Normal file
73
backend/internal/assets/localfs_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
69
backend/internal/assets/remote.go
Normal file
69
backend/internal/assets/remote.go
Normal 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
|
||||
}
|
||||
42
backend/internal/assets/remote_test.go
Normal file
42
backend/internal/assets/remote_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
194
backend/internal/assets/service_test.go
Normal file
194
backend/internal/assets/service_test.go
Normal 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
|
||||
}
|
||||
193
backend/internal/postgres/assets.go
Normal file
193
backend/internal/postgres/assets.go
Normal 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) }
|
||||
172
backend/internal/postgres/assets_test.go
Normal file
172
backend/internal/postgres/assets_test.go
Normal 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"])
|
||||
}
|
||||
}
|
||||
28
contracts/assets/assets-v1.json
Normal file
28
contracts/assets/assets-v1.json
Normal 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:"
|
||||
}
|
||||
}
|
||||
48
tests/assets-contract.test.ts
Normal file
48
tests/assets-contract.test.ts
Normal 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:" });
|
||||
});
|
||||
});
|
||||
Reference in New Issue
Block a user