Files
NianAIGC/backend/internal/assets/oss_test.go
2026-08-19 10:52:55 +08:00

182 lines
6.6 KiB
Go

package assets
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"log"
"strings"
"testing"
"time"
)
type ossClientStub struct {
puts []OSSPutRequest
acls []OSSACLRequest
gets []OSSObjectRequest
deletes []OSSObjectRequest
body []byte
contentType string
getErr, errorToReturn error
signedRequest OSSObjectRequest
signedTTL time.Duration
signedURL string
}
func (s *ossClientStub) Put(_ context.Context, r OSSPutRequest) error {
s.puts = append(s.puts, r)
_, _ = io.ReadAll(r.Body)
return s.errorToReturn
}
func (s *ossClientStub) SetACL(_ context.Context, r OSSACLRequest) error {
s.acls = append(s.acls, r)
return s.errorToReturn
}
func (s *ossClientStub) Get(_ context.Context, r OSSObjectRequest) (Blob, error) {
s.gets = append(s.gets, r)
if s.getErr != nil {
return Blob{}, s.getErr
}
return Blob{Body: io.NopCloser(bytes.NewReader(s.body)), ContentType: s.contentType, Size: int64(len(s.body))}, nil
}
func (s *ossClientStub) Delete(_ context.Context, r OSSObjectRequest) error {
s.deletes = append(s.deletes, r)
return s.errorToReturn
}
func (s *ossClientStub) SignGetURL(request OSSObjectRequest, ttl time.Duration) (string, error) {
s.signedRequest, s.signedTTL = request, ttl
return s.signedURL, s.errorToReturn
}
func TestOSSPutUsesConfiguredAddressContentTypeAndPublicACL(t *testing.T) {
client := &ossClientStub{}
store, err := NewOSS(OSSConfig{Endpoint: "https://oss-cn.test", Bucket: "bucket-a", PublicBaseURL: "https://cdn.test/root", PublicRead: true}, client)
if err != nil {
t.Fatal(err)
}
got, err := store.Put(context.Background(), "uploads/a b.png", bytes.NewReader([]byte("png")), 3, "image/png")
if err != nil {
t.Fatal(err)
}
if got.Key != "uploads/a b.png" || got.URL != "https://cdn.test/root/uploads/a%20b.png" {
t.Fatalf("stored=%#v", got)
}
if len(client.puts) != 1 || client.puts[0].Endpoint != "https://oss-cn.test" || client.puts[0].Bucket != "bucket-a" || client.puts[0].ContentType != "image/png" || client.puts[0].Size != 3 {
t.Fatalf("puts=%#v", client.puts)
}
if len(client.acls) != 1 || client.acls[0].ACL != OSSACLPublicRead {
t.Fatalf("acls=%#v", client.acls)
}
}
func TestOSSReadDeleteAndMissingMapping(t *testing.T) {
client := &ossClientStub{body: []byte("x"), contentType: "image/png"}
store, _ := NewOSS(OSSConfig{Endpoint: "e", Bucket: "b", PublicBaseURL: "https://cdn.test"}, client)
blob, err := store.Read(context.Background(), "generated-results/x.png")
if err != nil {
t.Fatal(err)
}
_ = blob.Body.Close()
if len(client.gets) != 1 || client.gets[0].Key != "generated-results/x.png" {
t.Fatalf("gets=%#v", client.gets)
}
if err = store.Delete(context.Background(), "generated-results/x.png"); err != nil || len(client.deletes) != 1 {
t.Fatalf("delete=%v %#v", err, client.deletes)
}
client.getErr = &OSSError{Status: 404, Code: "NoSuchKey", Err: errors.New("secret response")}
if _, err = store.Read(context.Background(), "missing"); !errors.Is(err, ErrBlobNotFound) {
t.Fatalf("missing=%v", err)
}
client.errorToReturn = &OSSError{Status: 404, Code: "NoSuchKey"}
if err = store.Delete(context.Background(), "missing"); err != nil {
t.Fatalf("delete missing=%v", err)
}
}
func TestOSSSignsExternalReadsAgainstPublicBucketBaseURL(t *testing.T) {
client := &ossClientStub{signedURL: "https://bucket-a.oss-cn-test.aliyuncs.com/uploads/a.png?Signature=signed"}
store, err := NewOSS(OSSConfig{
Endpoint: "https://oss-cn-test-internal.aliyuncs.com", Bucket: "bucket-a", PublicBaseURL: "https://bucket-a.oss-cn-test.aliyuncs.com",
}, client)
if err != nil {
t.Fatal(err)
}
got, err := store.SignReadURL("uploads/a.png", time.Hour)
if err != nil || got != client.signedURL {
t.Fatalf("signed URL = %q, err=%v", got, err)
}
if client.signedRequest.Endpoint != "https://bucket-a.oss-cn-test.aliyuncs.com" || client.signedRequest.Bucket != "bucket-a" || client.signedRequest.Key != "uploads/a.png" || client.signedTTL != time.Hour {
t.Fatalf("signed request = %#v, ttl=%s", client.signedRequest, client.signedTTL)
}
}
func TestOSSRejectsUnsafeKeysAndIncompleteConfiguration(t *testing.T) {
if _, err := NewOSS(OSSConfig{Endpoint: "e", Bucket: "b"}, &ossClientStub{}); err == nil {
t.Fatal("expected config error")
}
store, _ := NewOSS(OSSConfig{Endpoint: "e", Bucket: "b", PublicBaseURL: "https://cdn.test"}, &ossClientStub{})
for _, key := range []string{"", "../secret", "/absolute", "a\\b"} {
if _, err := store.Put(context.Background(), key, bytes.NewReader(nil), 0, "x"); !errors.Is(err, ErrUnsafeBlobKey) {
t.Errorf("key %q: %v", key, err)
}
}
}
func TestOSSOperationDiagnosticIncludesStageAndSafeServiceError(t *testing.T) {
diagnostic := diagnoseOSSOperation("set_acl", &OSSError{
Status: 403,
Code: "AccessDenied",
Err: errors.New("private-key private-object secret response"),
})
if diagnostic.Operation != "set_acl" || diagnostic.Status != 403 || diagnostic.Code != "AccessDenied" || diagnostic.ErrorClass != "service" {
t.Fatalf("diagnostic = %#v", diagnostic)
}
serialized := fmt.Sprintf("%+v", diagnostic)
for _, secret := range []string{"private-key", "private-object", "secret response"} {
if strings.Contains(serialized, secret) {
t.Fatalf("diagnostic leaks %q: %s", secret, serialized)
}
}
}
func TestOSSPutFailureLogsStageWithoutSensitiveDetails(t *testing.T) {
var output bytes.Buffer
previousOutput, previousFlags := log.Writer(), log.Flags()
log.SetOutput(&output)
log.SetFlags(0)
t.Cleanup(func() {
log.SetOutput(previousOutput)
log.SetFlags(previousFlags)
})
client := &ossClientStub{errorToReturn: &OSSError{
Status: 403,
Code: "AccessDenied",
Err: errors.New("access-key private-object secret response"),
}}
store, err := NewOSS(OSSConfig{
Endpoint: "https://oss-cn-guangzhou.aliyuncs.com",
Bucket: "private-bucket",
PublicBaseURL: "https://private-bucket.oss-cn-guangzhou.aliyuncs.com",
PublicRead: true,
}, client)
if err != nil {
t.Fatal(err)
}
if _, err = store.Put(context.Background(), "uploads/private-file.png", bytes.NewReader([]byte("png")), 3, "image/png"); err == nil {
t.Fatal("Put error = nil")
}
got := output.String()
for _, expected := range []string{"operation=put", "status=403", `code="AccessDenied"`, "errorClass=service"} {
if !strings.Contains(got, expected) {
t.Fatalf("log %q does not contain %q", got, expected)
}
}
for _, secret := range []string{"access-key", "private-object", "secret response", "private-bucket", "private-file"} {
if strings.Contains(got, secret) {
t.Fatalf("log leaks %q: %s", secret, got)
}
}
}