Files
NianAIGC/backend/internal/assets/oss_http.go

238 lines
7.1 KiB
Go

package assets
import (
"context"
"crypto/hmac"
"crypto/sha1"
"encoding/base64"
"encoding/xml"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"path"
"sort"
"strings"
"time"
)
const maxOSSErrorBody = 64 << 10
// OSSHTTPClient implements OSSClient using Aliyun OSS's HTTP authorization
// protocol. The supplied HTTP client owns timeout and transport policy.
type OSSHTTPClient struct {
accessKeyID string
accessKeySecret string
client *http.Client
now func() time.Time
}
// NewOSSHTTPClient constructs an OSS adapter with injected credentials,
// transport, and clock. Callers should supply an http.Client with a bounded
// Timeout.
func NewOSSHTTPClient(accessKeyID, accessKeySecret string, client *http.Client, now func() time.Time) (*OSSHTTPClient, error) {
if strings.TrimSpace(accessKeyID) == "" || strings.TrimSpace(accessKeySecret) == "" {
return nil, errors.New("OSS access key ID and secret are required")
}
if client == nil {
return nil, errors.New("OSS HTTP client is required")
}
if now == nil {
now = time.Now
}
return &OSSHTTPClient{accessKeyID: accessKeyID, accessKeySecret: accessKeySecret, client: client, now: now}, nil
}
func (c *OSSHTTPClient) Put(ctx context.Context, request OSSPutRequest) error {
req, err := c.newRequest(ctx, http.MethodPut, request.Endpoint, request.Bucket, request.Key, "", request.Body)
if err != nil {
return err
}
if request.Size < 0 {
return errors.New("OSS object size must not be negative")
}
req.ContentLength = request.Size
req.Header.Set("Content-Length", fmt.Sprintf("%d", request.Size))
if request.ContentType != "" {
req.Header.Set("Content-Type", request.ContentType)
}
c.sign(req, request.Bucket, request.Key, "")
response, err := c.do(req)
if err != nil {
return err
}
return consumeOSSResponse(response)
}
func (c *OSSHTTPClient) SetACL(ctx context.Context, request OSSACLRequest) error {
req, err := c.newRequest(ctx, http.MethodPut, request.Endpoint, request.Bucket, request.Key, "acl", nil)
if err != nil {
return err
}
req.Header.Set("x-oss-object-acl", request.ACL)
c.sign(req, request.Bucket, request.Key, "acl")
response, err := c.do(req)
if err != nil {
return err
}
return consumeOSSResponse(response)
}
func (c *OSSHTTPClient) Get(ctx context.Context, request OSSObjectRequest) (Blob, error) {
req, err := c.newRequest(ctx, http.MethodGet, request.Endpoint, request.Bucket, request.Key, "", nil)
if err != nil {
return Blob{}, err
}
c.sign(req, request.Bucket, request.Key, "")
response, err := c.do(req)
if err != nil {
return Blob{}, err
}
if !successfulOSSStatus(response.StatusCode) {
return Blob{}, readOSSError(response)
}
return Blob{Body: response.Body, ContentType: response.Header.Get("Content-Type"), Size: response.ContentLength}, nil
}
func (c *OSSHTTPClient) Delete(ctx context.Context, request OSSObjectRequest) error {
req, err := c.newRequest(ctx, http.MethodDelete, request.Endpoint, request.Bucket, request.Key, "", nil)
if err != nil {
return err
}
c.sign(req, request.Bucket, request.Key, "")
response, err := c.do(req)
if err != nil {
return err
}
return consumeOSSResponse(response)
}
func (c *OSSHTTPClient) newRequest(ctx context.Context, method, endpoint, bucket, key, subresource string, body io.Reader) (*http.Request, error) {
requestURL, err := ossObjectURL(endpoint, bucket, key, subresource)
if err != nil {
return nil, err
}
req, err := http.NewRequestWithContext(ctx, method, requestURL.String(), body)
if err != nil {
return nil, errors.New("build OSS request")
}
return req, nil
}
func (c *OSSHTTPClient) sign(req *http.Request, bucket, key, subresource string) {
req.Header.Set("Date", c.now().UTC().Format(http.TimeFormat))
canonicalResource := "/" + bucket + "/" + key
if subresource != "" {
canonicalResource += "?" + subresource
}
stringToSign := strings.Join([]string{
req.Method,
req.Header.Get("Content-MD5"),
req.Header.Get("Content-Type"),
req.Header.Get("Date"),
canonicalOSSHeaders(req.Header) + canonicalResource,
}, "\n")
mac := hmac.New(sha1.New, []byte(c.accessKeySecret))
_, _ = mac.Write([]byte(stringToSign))
signature := base64.StdEncoding.EncodeToString(mac.Sum(nil))
req.Header.Set("Authorization", "OSS "+c.accessKeyID+":"+signature)
}
func (c *OSSHTTPClient) do(req *http.Request) (*http.Response, error) {
response, err := c.client.Do(req)
if err == nil {
return response, nil
}
if ctxErr := req.Context().Err(); ctxErr != nil {
return nil, ctxErr
}
return nil, &OSSError{Err: errors.New("OSS transport request failed")}
}
func ossObjectURL(endpoint, bucket, key, subresource string) (*url.URL, error) {
endpoint = strings.TrimSpace(endpoint)
if !strings.Contains(endpoint, "://") {
endpoint = "https://" + endpoint
}
u, err := url.Parse(endpoint)
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" {
return nil, errors.New("invalid OSS endpoint")
}
if bucket == "" || strings.ContainsAny(bucket, "/\\") || key == "" || strings.HasPrefix(key, "/") {
return nil, errors.New("invalid OSS bucket or object key")
}
basePath := strings.TrimRight(u.Path, "/")
pathStyle := basePath != "" && path.Base(basePath) == bucket
hostHasBucket := strings.HasPrefix(strings.ToLower(u.Hostname()), strings.ToLower(bucket)+".")
if !pathStyle && !hostHasBucket {
hostname := bucket + "." + u.Hostname()
if port := u.Port(); port != "" {
hostname += ":" + port
}
u.Host = hostname
}
u.Path = basePath + "/" + key
u.RawPath = strings.TrimRight(escapedURLPath(basePath), "/") + "/" + escapeOSSKey(key)
if subresource != "" {
u.RawQuery = subresource
}
return u, nil
}
func escapedURLPath(value string) string {
if value == "" {
return ""
}
parts := strings.Split(value, "/")
for i := range parts {
parts[i] = url.PathEscape(parts[i])
}
return strings.Join(parts, "/")
}
func canonicalOSSHeaders(header http.Header) string {
keys := make([]string, 0)
values := make(map[string]string)
for key, entries := range header {
lowerKey := strings.ToLower(key)
if !strings.HasPrefix(lowerKey, "x-oss-") {
continue
}
keys = append(keys, lowerKey)
trimmed := make([]string, len(entries))
for i, entry := range entries {
trimmed[i] = strings.TrimSpace(entry)
}
values[lowerKey] = strings.Join(trimmed, ",")
}
sort.Strings(keys)
var canonical strings.Builder
for _, key := range keys {
fmt.Fprintf(&canonical, "%s:%s\n", key, values[key])
}
return canonical.String()
}
func consumeOSSResponse(response *http.Response) error {
if successfulOSSStatus(response.StatusCode) {
_, _ = io.Copy(io.Discard, response.Body)
return response.Body.Close()
}
return readOSSError(response)
}
func successfulOSSStatus(status int) bool { return status >= 200 && status < 300 }
func readOSSError(response *http.Response) error {
defer response.Body.Close()
var wire struct {
Code string `xml:"Code"`
}
_ = xml.NewDecoder(io.LimitReader(response.Body, maxOSSErrorBody)).Decode(&wire)
return &OSSError{Status: response.StatusCode, Code: strings.TrimSpace(wire.Code)}
}
var _ OSSClient = (*OSSHTTPClient)(nil)