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

70 lines
1.8 KiB
Go

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
}