feat: add secure HTTP crypto and local storage (#8)
This commit is contained in:
@@ -3,6 +3,7 @@ module git.ilapage.cn/OPC/chorus
|
||||
go 1.26.5
|
||||
|
||||
require (
|
||||
github.com/disintegration/imaging v1.6.2
|
||||
github.com/go-sql-driver/mysql v1.10.0
|
||||
golang.org/x/crypto v0.55.0
|
||||
gorm.io/gorm v1.31.2
|
||||
@@ -12,5 +13,6 @@ require (
|
||||
filippo.io/edwards25519 v1.2.0 // indirect
|
||||
github.com/jinzhu/inflection v1.0.0 // indirect
|
||||
github.com/jinzhu/now v1.1.5 // indirect
|
||||
golang.org/x/image v0.41.0 // indirect
|
||||
golang.org/x/text v0.41.0 // indirect
|
||||
)
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
filippo.io/edwards25519 v1.2.0 h1:crnVqOiS4jqYleHd9vaKZ+HKtHfllngJIiOpNpoJsjo=
|
||||
filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q77mc=
|
||||
github.com/disintegration/imaging v1.6.2 h1:w1LecBlG2Lnp8B3jk5zSuNqd7b4DXhcjwek1ei82L+c=
|
||||
github.com/disintegration/imaging v1.6.2/go.mod h1:44/5580QXChDfwIclfc/PCwrr44amcmDAg8hxG0Ewe4=
|
||||
github.com/go-sql-driver/mysql v1.10.0 h1:Q+1LV8DkHJvSYAdR83XzuhDaTykuDx0l6fkXxoWCWfw=
|
||||
github.com/go-sql-driver/mysql v1.10.0/go.mod h1:M+cqaI7+xxXGG9swrdeUIoPG3Y3KCkF0pZej+SK+nWk=
|
||||
github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E=
|
||||
@@ -8,6 +10,10 @@ github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
|
||||
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
|
||||
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
||||
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
||||
golang.org/x/image v0.0.0-20191009234506-e7c1f5e7dbb8/go.mod h1:FeLwcggjj3mMvU+oOTbSwawSJRM1uh48EjtB4UJZlP0=
|
||||
golang.org/x/image v0.41.0 h1:8wS72eGJMJaBxK6okTzd4WaXumUlTVlb753MlsSvTCo=
|
||||
golang.org/x/image v0.41.0/go.mod h1:uIc348UZMSvS5Z65CVZ7iDPaNobNFEPeJ4kbqTOszmA=
|
||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
|
||||
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
|
||||
gorm.io/gorm v1.31.2 h1:3o8FXNo9v9S858gil+3LlZA1LkCOzgb4g5BL64FgaCo=
|
||||
|
||||
@@ -6,13 +6,23 @@ import (
|
||||
)
|
||||
|
||||
type Object struct {
|
||||
Key string
|
||||
ContentType string
|
||||
Size int64
|
||||
Key string
|
||||
OwnerID uint64
|
||||
GenerationID uint64
|
||||
ContentType string
|
||||
Size int64
|
||||
}
|
||||
|
||||
type PutRequest struct {
|
||||
Key string
|
||||
OwnerID uint64
|
||||
GenerationID uint64
|
||||
ContentType string
|
||||
Source io.Reader
|
||||
}
|
||||
|
||||
type Store interface {
|
||||
Put(ctx context.Context, key string, source io.Reader, contentType string) (Object, error)
|
||||
Open(ctx context.Context, key string) (io.ReadCloser, error)
|
||||
Put(ctx context.Context, request PutRequest) (Object, error)
|
||||
Open(ctx context.Context, key string) (io.ReadCloser, Object, error)
|
||||
Delete(ctx context.Context, key string) error
|
||||
}
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
package crypto
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
corecrypto "git.ilapage.cn/OPC/chorus/internal/core/crypto"
|
||||
)
|
||||
|
||||
const envelopeVersion uint8 = 1
|
||||
|
||||
var (
|
||||
ErrInvalidKeyRing = errors.New("invalid key ring")
|
||||
ErrUnknownKey = errors.New("envelope key is unavailable")
|
||||
ErrInvalidEnvelope = errors.New("invalid encrypted envelope")
|
||||
)
|
||||
|
||||
type KeyRing struct {
|
||||
activeKeyID string
|
||||
keys map[string][]byte
|
||||
random io.Reader
|
||||
}
|
||||
|
||||
func NewKeyRing(activeKeyID string, keys map[string][]byte) (*KeyRing, error) {
|
||||
if activeKeyID == "" || len(keys) == 0 {
|
||||
return nil, ErrInvalidKeyRing
|
||||
}
|
||||
copied := make(map[string][]byte, len(keys))
|
||||
for keyID, key := range keys {
|
||||
if keyID == "" {
|
||||
return nil, ErrInvalidKeyRing
|
||||
}
|
||||
if _, err := aes.NewCipher(key); err != nil {
|
||||
return nil, fmt.Errorf("%w: key %q has an unsupported length", ErrInvalidKeyRing, keyID)
|
||||
}
|
||||
copied[keyID] = append([]byte(nil), key...)
|
||||
}
|
||||
if _, ok := copied[activeKeyID]; !ok {
|
||||
return nil, ErrInvalidKeyRing
|
||||
}
|
||||
return &KeyRing{activeKeyID: activeKeyID, keys: copied, random: rand.Reader}, nil
|
||||
}
|
||||
|
||||
func (r *KeyRing) Encrypt(ctx context.Context, plaintext []byte) (corecrypto.Envelope, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return corecrypto.Envelope{}, err
|
||||
}
|
||||
gcm, err := r.gcm(r.activeKeyID)
|
||||
if err != nil {
|
||||
return corecrypto.Envelope{}, err
|
||||
}
|
||||
nonce := make([]byte, gcm.NonceSize())
|
||||
if _, err := io.ReadFull(r.random, nonce); err != nil {
|
||||
return corecrypto.Envelope{}, fmt.Errorf("generate encryption nonce: %w", err)
|
||||
}
|
||||
envelope := corecrypto.Envelope{Version: envelopeVersion, KeyID: r.activeKeyID, Nonce: nonce}
|
||||
envelope.Ciphertext = gcm.Seal(nil, nonce, plaintext, associatedData(envelope.Version, envelope.KeyID))
|
||||
return envelope, nil
|
||||
}
|
||||
|
||||
func (r *KeyRing) Decrypt(ctx context.Context, envelope corecrypto.Envelope) ([]byte, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if envelope.Version != envelopeVersion || envelope.KeyID == "" {
|
||||
return nil, ErrInvalidEnvelope
|
||||
}
|
||||
gcm, err := r.gcm(envelope.KeyID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(envelope.Nonce) != gcm.NonceSize() || len(envelope.Ciphertext) < gcm.Overhead() {
|
||||
return nil, ErrInvalidEnvelope
|
||||
}
|
||||
plaintext, err := gcm.Open(nil, envelope.Nonce, envelope.Ciphertext, associatedData(envelope.Version, envelope.KeyID))
|
||||
if err != nil {
|
||||
return nil, ErrInvalidEnvelope
|
||||
}
|
||||
return plaintext, nil
|
||||
}
|
||||
|
||||
func (r *KeyRing) gcm(keyID string) (cipher.AEAD, error) {
|
||||
key, ok := r.keys[keyID]
|
||||
if !ok {
|
||||
return nil, ErrUnknownKey
|
||||
}
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return nil, ErrInvalidKeyRing
|
||||
}
|
||||
return cipher.NewGCM(block)
|
||||
}
|
||||
|
||||
func associatedData(version uint8, keyID string) []byte {
|
||||
data := make([]byte, 3+len(keyID))
|
||||
data[0] = version
|
||||
binary.BigEndian.PutUint16(data[1:3], uint16(len(keyID)))
|
||||
copy(data[3:], keyID)
|
||||
return data
|
||||
}
|
||||
|
||||
var _ corecrypto.KeyCipher = (*KeyRing)(nil)
|
||||
@@ -0,0 +1,63 @@
|
||||
package crypto
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
corecrypto "git.ilapage.cn/OPC/chorus/internal/core/crypto"
|
||||
)
|
||||
|
||||
func TestKeyRingRoundTripAndRotation(t *testing.T) {
|
||||
oldKey := []byte("0123456789abcdef0123456789abcdef")
|
||||
newKey := []byte("abcdef0123456789abcdef0123456789")
|
||||
oldRing, err := NewKeyRing("old", map[string][]byte{"old": oldKey})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
oldEnvelope, err := oldRing.Encrypt(context.Background(), []byte("synthetic-api-key"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
rotated, err := NewKeyRing("new", map[string][]byte{"old": oldKey, "new": newKey})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
plaintext, err := rotated.Decrypt(context.Background(), oldEnvelope)
|
||||
if err != nil || string(plaintext) != "synthetic-api-key" {
|
||||
t.Fatalf("decrypt old envelope = %q, %v", plaintext, err)
|
||||
}
|
||||
newEnvelope, err := rotated.Encrypt(context.Background(), []byte("new-synthetic-key"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if newEnvelope.KeyID != "new" || newEnvelope.Version != 1 {
|
||||
t.Fatalf("unexpected new envelope metadata: %#v", newEnvelope)
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeyRingRejectsTamperAndUnknownKeyWithoutPlaintext(t *testing.T) {
|
||||
ring, err := NewKeyRing("current", map[string][]byte{"current": []byte("0123456789abcdef0123456789abcdef")})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
secret := "do-not-leak-this-value"
|
||||
envelope, err := ring.Encrypt(context.Background(), []byte(secret))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
envelope.Ciphertext[0] ^= 0xff
|
||||
_, err = ring.Decrypt(context.Background(), envelope)
|
||||
if !errors.Is(err, ErrInvalidEnvelope) || strings.Contains(err.Error(), secret) {
|
||||
t.Fatalf("tamper error = %v", err)
|
||||
}
|
||||
|
||||
_, err = ring.Decrypt(context.Background(), corecrypto.Envelope{
|
||||
Version: 1, KeyID: "retired", Nonce: make([]byte, 12), Ciphertext: make([]byte, 16),
|
||||
})
|
||||
if !errors.Is(err, ErrUnknownKey) {
|
||||
t.Fatalf("unknown key error = %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,256 @@
|
||||
package safehttp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrInvalidTarget = errors.New("outbound target is invalid")
|
||||
ErrBlockedAddress = errors.New("outbound target address is blocked")
|
||||
ErrTooManyRedirects = errors.New("outbound redirect limit exceeded")
|
||||
ErrResponseTooLarge = errors.New("outbound response exceeds size limit")
|
||||
ErrRequestFailed = errors.New("safe outbound request failed")
|
||||
)
|
||||
|
||||
type Resolver interface {
|
||||
LookupIPAddr(ctx context.Context, host string) ([]net.IPAddr, error)
|
||||
}
|
||||
|
||||
type DialContextFunc func(ctx context.Context, network, address string) (net.Conn, error)
|
||||
|
||||
type Config struct {
|
||||
Resolver Resolver
|
||||
DialContext DialContextFunc
|
||||
Timeout time.Duration
|
||||
MaxRedirects int
|
||||
AllowedPorts []uint16
|
||||
}
|
||||
|
||||
type Client struct {
|
||||
httpClient *http.Client
|
||||
validator *validator
|
||||
}
|
||||
|
||||
type Response struct {
|
||||
StatusCode int
|
||||
ContentType string
|
||||
Body []byte
|
||||
}
|
||||
|
||||
func New(config Config) (*Client, error) {
|
||||
if config.Resolver == nil {
|
||||
config.Resolver = net.DefaultResolver
|
||||
}
|
||||
if config.DialContext == nil {
|
||||
dialer := &net.Dialer{Timeout: 10 * time.Second, KeepAlive: 30 * time.Second}
|
||||
config.DialContext = dialer.DialContext
|
||||
}
|
||||
if config.Timeout <= 0 {
|
||||
return nil, fmt.Errorf("safe HTTP timeout must be positive")
|
||||
}
|
||||
if config.MaxRedirects < 0 {
|
||||
return nil, fmt.Errorf("safe HTTP redirect limit cannot be negative")
|
||||
}
|
||||
if len(config.AllowedPorts) == 0 {
|
||||
config.AllowedPorts = []uint16{80, 443}
|
||||
}
|
||||
ports := make(map[string]struct{}, len(config.AllowedPorts))
|
||||
for _, port := range config.AllowedPorts {
|
||||
if port == 0 {
|
||||
return nil, fmt.Errorf("safe HTTP allowed port cannot be zero")
|
||||
}
|
||||
ports[strconv.Itoa(int(port))] = struct{}{}
|
||||
}
|
||||
v := &validator{resolver: config.Resolver, allowedPorts: ports}
|
||||
safeDialer := &dialer{validator: v, dial: config.DialContext}
|
||||
transport := &http.Transport{
|
||||
Proxy: nil,
|
||||
DialContext: safeDialer.DialContext,
|
||||
ForceAttemptHTTP2: true,
|
||||
MaxIdleConns: 20,
|
||||
IdleConnTimeout: 30 * time.Second,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ResponseHeaderTimeout: config.Timeout,
|
||||
}
|
||||
client := &http.Client{Transport: transport, Timeout: config.Timeout}
|
||||
client.CheckRedirect = func(request *http.Request, via []*http.Request) error {
|
||||
if len(via) > config.MaxRedirects {
|
||||
return ErrTooManyRedirects
|
||||
}
|
||||
return v.validateURL(request.Context(), request.URL, true)
|
||||
}
|
||||
return &Client{httpClient: client, validator: v}, nil
|
||||
}
|
||||
|
||||
func (c *Client) Fetch(ctx context.Context, rawURL string, maxBytes int64) (Response, error) {
|
||||
if maxBytes <= 0 {
|
||||
return Response{}, fmt.Errorf("maximum response bytes must be positive")
|
||||
}
|
||||
target, err := url.Parse(rawURL)
|
||||
if err != nil {
|
||||
return Response{}, ErrInvalidTarget
|
||||
}
|
||||
if err := c.validator.validateURL(ctx, target, true); err != nil {
|
||||
return Response{}, err
|
||||
}
|
||||
request, err := http.NewRequestWithContext(ctx, http.MethodGet, target.String(), nil)
|
||||
if err != nil {
|
||||
return Response{}, ErrInvalidTarget
|
||||
}
|
||||
response, err := c.httpClient.Do(request)
|
||||
if err != nil {
|
||||
var urlError *url.Error
|
||||
if errors.As(err, &urlError) && urlError.Err != nil {
|
||||
return Response{}, fmt.Errorf("%w: %w", ErrRequestFailed, urlError.Err)
|
||||
}
|
||||
return Response{}, ErrRequestFailed
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.ContentLength > maxBytes {
|
||||
return Response{}, ErrResponseTooLarge
|
||||
}
|
||||
body, err := io.ReadAll(io.LimitReader(response.Body, maxBytes+1))
|
||||
if err != nil {
|
||||
return Response{}, fmt.Errorf("read safe HTTP response: %w", err)
|
||||
}
|
||||
if int64(len(body)) > maxBytes {
|
||||
return Response{}, ErrResponseTooLarge
|
||||
}
|
||||
return Response{StatusCode: response.StatusCode, ContentType: response.Header.Get("Content-Type"), Body: body}, nil
|
||||
}
|
||||
|
||||
func (c *Client) CloseIdleConnections() {
|
||||
c.httpClient.CloseIdleConnections()
|
||||
}
|
||||
|
||||
type validator struct {
|
||||
resolver Resolver
|
||||
allowedPorts map[string]struct{}
|
||||
}
|
||||
|
||||
func (v *validator) validateURL(ctx context.Context, target *url.URL, resolve bool) error {
|
||||
if target == nil || (target.Scheme != "http" && target.Scheme != "https") || target.Hostname() == "" || target.User != nil {
|
||||
return ErrInvalidTarget
|
||||
}
|
||||
port := target.Port()
|
||||
if port == "" {
|
||||
if target.Scheme == "https" {
|
||||
port = "443"
|
||||
} else {
|
||||
port = "80"
|
||||
}
|
||||
}
|
||||
if _, ok := v.allowedPorts[port]; !ok {
|
||||
return ErrInvalidTarget
|
||||
}
|
||||
if resolve {
|
||||
_, err := v.resolvePublic(ctx, target.Hostname())
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (v *validator) resolvePublic(ctx context.Context, host string) ([]net.IPAddr, error) {
|
||||
if parsed := net.ParseIP(host); parsed != nil {
|
||||
addresses := []net.IPAddr{{IP: parsed}}
|
||||
if err := validateAddresses(addresses); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return addresses, nil
|
||||
}
|
||||
addresses, err := v.resolver.LookupIPAddr(ctx, host)
|
||||
if err != nil || len(addresses) == 0 {
|
||||
return nil, ErrInvalidTarget
|
||||
}
|
||||
if err := validateAddresses(addresses); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return addresses, nil
|
||||
}
|
||||
|
||||
type dialer struct {
|
||||
validator *validator
|
||||
dial DialContextFunc
|
||||
}
|
||||
|
||||
func (d *dialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
host, port, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return nil, ErrInvalidTarget
|
||||
}
|
||||
if _, ok := d.validator.allowedPorts[port]; !ok {
|
||||
return nil, ErrInvalidTarget
|
||||
}
|
||||
addresses, err := d.validator.resolvePublic(ctx, host)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var lastErr error
|
||||
for _, resolved := range addresses {
|
||||
ip := resolved.IP.String()
|
||||
connection, dialErr := d.dial(ctx, network, net.JoinHostPort(ip, port))
|
||||
if dialErr == nil {
|
||||
return connection, nil
|
||||
}
|
||||
lastErr = dialErr
|
||||
}
|
||||
if lastErr != nil {
|
||||
return nil, fmt.Errorf("safe outbound connection failed")
|
||||
}
|
||||
return nil, ErrInvalidTarget
|
||||
}
|
||||
|
||||
func validateAddresses(addresses []net.IPAddr) error {
|
||||
for _, address := range addresses {
|
||||
parsed, ok := netip.AddrFromSlice(address.IP)
|
||||
if !ok || blockedIP(parsed.Unmap()) {
|
||||
return ErrBlockedAddress
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
var blockedPrefixes = mustPrefixes(
|
||||
"0.0.0.0/8", "10.0.0.0/8", "100.64.0.0/10", "127.0.0.0/8", "169.254.0.0/16",
|
||||
"172.16.0.0/12", "192.0.0.0/24", "192.0.2.0/24", "192.88.99.0/24", "192.168.0.0/16",
|
||||
"198.18.0.0/15", "198.51.100.0/24", "203.0.113.0/24", "224.0.0.0/4", "240.0.0.0/4",
|
||||
"::/128", "::1/128", "64:ff9b::/96", "100::/64", "2001:db8::/32", "fc00::/7", "fe80::/10", "ff00::/8",
|
||||
)
|
||||
|
||||
func blockedIP(address netip.Addr) bool {
|
||||
if !address.IsValid() || !address.IsGlobalUnicast() || address.IsPrivate() || address.IsLoopback() || address.IsLinkLocalUnicast() || address.IsLinkLocalMulticast() || address.IsMulticast() || address.IsUnspecified() {
|
||||
return true
|
||||
}
|
||||
for _, prefix := range blockedPrefixes {
|
||||
if prefix.Contains(address) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func mustPrefixes(values ...string) []netip.Prefix {
|
||||
prefixes := make([]netip.Prefix, 0, len(values))
|
||||
for _, value := range values {
|
||||
prefixes = append(prefixes, netip.MustParsePrefix(value))
|
||||
}
|
||||
return prefixes
|
||||
}
|
||||
|
||||
func RedactedURL(rawURL string) string {
|
||||
target, err := url.Parse(rawURL)
|
||||
if err != nil || target.Hostname() == "" {
|
||||
return "<invalid-url>"
|
||||
}
|
||||
return strings.ToLower(target.Scheme) + "://" + target.Host
|
||||
}
|
||||
@@ -0,0 +1,178 @@
|
||||
package safehttp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
type fakeResolver struct {
|
||||
mu sync.Mutex
|
||||
addresses map[string][][]net.IPAddr
|
||||
}
|
||||
|
||||
func (r *fakeResolver) LookupIPAddr(_ context.Context, host string) ([]net.IPAddr, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
answers := r.addresses[host]
|
||||
if len(answers) == 0 {
|
||||
return nil, fmt.Errorf("no answer")
|
||||
}
|
||||
answer := answers[0]
|
||||
if len(answers) > 1 {
|
||||
r.addresses[host] = answers[1:]
|
||||
}
|
||||
return answer, nil
|
||||
}
|
||||
|
||||
func ips(values ...string) []net.IPAddr {
|
||||
result := make([]net.IPAddr, 0, len(values))
|
||||
for _, value := range values {
|
||||
result = append(result, net.IPAddr{IP: net.ParseIP(value)})
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func TestBlockedIPv4AndIPv6(t *testing.T) {
|
||||
blocked := []string{"127.0.0.1", "10.0.0.1", "169.254.169.254", "100.64.0.1", "0.0.0.0", "::1", "fc00::1", "fe80::1", "ff02::1", "::ffff:192.168.1.1"}
|
||||
for _, value := range blocked {
|
||||
t.Run(value, func(t *testing.T) {
|
||||
if err := validateAddresses(ips(value)); !errors.Is(err, ErrBlockedAddress) {
|
||||
t.Fatalf("validateAddresses(%s) error = %v", value, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
if err := validateAddresses(ips("93.184.216.34", "2606:4700:4700::1111")); err != nil {
|
||||
t.Fatalf("public addresses rejected: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDialUsesValidatedIPAndRejectsDNSRebinding(t *testing.T) {
|
||||
resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{
|
||||
"public.test": {ips("93.184.216.34")},
|
||||
"rebind.test": {ips("93.184.216.34"), ips("127.0.0.1")},
|
||||
}}
|
||||
var dialed string
|
||||
fakeDial := func(_ context.Context, _, address string) (net.Conn, error) {
|
||||
dialed = address
|
||||
return nil, errors.New("synthetic dial stop")
|
||||
}
|
||||
client, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, MaxRedirects: 2})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, _ = client.Fetch(context.Background(), "http://public.test/data", 10)
|
||||
if dialed != "93.184.216.34:80" {
|
||||
t.Fatalf("dialed address = %q", dialed)
|
||||
}
|
||||
dialed = ""
|
||||
_, err = client.Fetch(context.Background(), "http://rebind.test/data", 10)
|
||||
if !errors.Is(err, ErrBlockedAddress) || dialed != "" {
|
||||
t.Fatalf("rebind result: dialed=%q error=%v", dialed, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchPublicMockAndSizeLimit(t *testing.T) {
|
||||
resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{
|
||||
"public.test": {ips("93.184.216.34"), ips("93.184.216.34"), ips("93.184.216.34"), ips("93.184.216.34")},
|
||||
}}
|
||||
fakeDial := func(_ context.Context, _, _ string) (net.Conn, error) {
|
||||
client, server := net.Pipe()
|
||||
go func() {
|
||||
defer server.Close()
|
||||
buffer := make([]byte, 4096)
|
||||
_, _ = server.Read(buffer)
|
||||
_, _ = server.Write([]byte("HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 11\r\nConnection: close\r\n\r\nhello world"))
|
||||
}()
|
||||
return client, nil
|
||||
}
|
||||
client, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, MaxRedirects: 1})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
response, err := client.Fetch(context.Background(), "http://public.test/data", 11)
|
||||
if err != nil || string(response.Body) != "hello world" {
|
||||
t.Fatalf("Fetch() = %#v, %v", response, err)
|
||||
}
|
||||
_, err = client.Fetch(context.Background(), "http://public.test/data", 5)
|
||||
if !errors.Is(err, ErrResponseTooLarge) {
|
||||
t.Fatalf("size limit error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRedirectAndProxyBypassAreRejected(t *testing.T) {
|
||||
resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{
|
||||
"public.test": {ips("93.184.216.34"), ips("93.184.216.34")},
|
||||
"private.test": {ips("10.0.0.8")},
|
||||
}}
|
||||
fakeDial := func(_ context.Context, _, _ string) (net.Conn, error) {
|
||||
client, server := net.Pipe()
|
||||
go func() {
|
||||
defer server.Close()
|
||||
buffer := make([]byte, 4096)
|
||||
_, _ = server.Read(buffer)
|
||||
_, _ = server.Write([]byte("HTTP/1.1 302 Found\r\nLocation: http://private.test/result\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"))
|
||||
}()
|
||||
return client, nil
|
||||
}
|
||||
t.Setenv("HTTP_PROXY", "http://127.0.0.1:9999")
|
||||
client, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, MaxRedirects: 2})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = client.Fetch(context.Background(), "http://public.test/start", 100)
|
||||
if !errors.Is(err, ErrBlockedAddress) {
|
||||
t.Fatalf("redirect error = %v", err)
|
||||
}
|
||||
transport := client.httpClient.Transport.(*http.Transport)
|
||||
if transport.Proxy != nil {
|
||||
t.Fatal("safe transport unexpectedly configured a proxy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestURLRestrictionsAndRedaction(t *testing.T) {
|
||||
resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{"public.test": {ips("93.184.216.34")}}}
|
||||
client, err := New(Config{Resolver: resolver, Timeout: time.Second, MaxRedirects: 0})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, target := range []string{"file:///etc/passwd", "http://user:password@public.test/", "http://public.test:22/"} {
|
||||
if _, err := client.Fetch(context.Background(), target, 10); !errors.Is(err, ErrInvalidTarget) {
|
||||
t.Errorf("Fetch(%q) error = %v", target, err)
|
||||
}
|
||||
}
|
||||
redacted := RedactedURL("https://user:secret@Public.Test/path?q=secret")
|
||||
if strings.Contains(redacted, "secret") || redacted != "https://Public.Test" {
|
||||
t.Fatalf("RedactedURL() = %q", redacted)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMaliciousResultURLAndTimeout(t *testing.T) {
|
||||
resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{
|
||||
"slow.test": {ips("93.184.216.34"), ips("93.184.216.34")},
|
||||
}}
|
||||
client, err := New(Config{
|
||||
Resolver: resolver,
|
||||
DialContext: func(_ context.Context, _, _ string) (net.Conn, error) {
|
||||
client, _ := net.Pipe()
|
||||
return client, nil
|
||||
},
|
||||
Timeout: 20 * time.Millisecond, MaxRedirects: 0,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := client.Fetch(context.Background(), "http://169.254.169.254/latest/meta-data", 100); !errors.Is(err, ErrBlockedAddress) {
|
||||
t.Fatalf("malicious result URL error = %v", err)
|
||||
}
|
||||
_, err = client.Fetch(context.Background(), "http://slow.test/result?signature=do-not-leak", 100)
|
||||
if !errors.Is(err, ErrRequestFailed) || strings.Contains(err.Error(), "do-not-leak") {
|
||||
t.Fatalf("timeout error = %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,396 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/png"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage"
|
||||
"github.com/disintegration/imaging"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrInvalidKey = errors.New("storage key is invalid")
|
||||
ErrInvalidMetadata = errors.New("storage ownership metadata is invalid")
|
||||
ErrObjectExists = errors.New("storage object already exists")
|
||||
ErrObjectTooLarge = errors.New("storage object exceeds size limit")
|
||||
ErrInvalidImage = errors.New("image content is invalid")
|
||||
ErrImageTooLarge = errors.New("image dimensions exceed limit")
|
||||
ErrUnsupportedImage = errors.New("image format is not allowed")
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
Root string
|
||||
MaxObjectBytes int64
|
||||
MaxImagePixels uint64
|
||||
ThumbnailMaxSide int
|
||||
AllowedImageMIME map[string]bool
|
||||
}
|
||||
|
||||
type Local struct {
|
||||
root string
|
||||
temporaryRoot string
|
||||
maxObjectBytes int64
|
||||
maxImagePixels uint64
|
||||
thumbnailMaxSide int
|
||||
allowedImageMIME map[string]bool
|
||||
}
|
||||
|
||||
type metadata struct {
|
||||
Key string `json:"key"`
|
||||
OwnerID uint64 `json:"owner_id"`
|
||||
GenerationID uint64 `json:"generation_id"`
|
||||
ContentType string `json:"content_type"`
|
||||
Size int64 `json:"size"`
|
||||
}
|
||||
|
||||
type ImageRequest struct {
|
||||
Key string
|
||||
ThumbnailKey string
|
||||
OwnerID uint64
|
||||
GenerationID uint64
|
||||
ContentType string
|
||||
Source io.Reader
|
||||
}
|
||||
|
||||
type ImageObjects struct {
|
||||
Original corestorage.Object
|
||||
Thumbnail corestorage.Object
|
||||
}
|
||||
|
||||
func NewLocal(config Config) (*Local, error) {
|
||||
if strings.TrimSpace(config.Root) == "" || config.MaxObjectBytes <= 0 || config.MaxImagePixels == 0 || config.ThumbnailMaxSide <= 0 || len(config.AllowedImageMIME) == 0 {
|
||||
return nil, fmt.Errorf("local storage configuration is incomplete")
|
||||
}
|
||||
root, err := filepath.Abs(config.Root)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("resolve storage root: %w", err)
|
||||
}
|
||||
if err := os.MkdirAll(root, 0o700); err != nil {
|
||||
return nil, fmt.Errorf("create storage root: %w", err)
|
||||
}
|
||||
root, err = filepath.EvalSymlinks(root)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("resolve storage root links: %w", err)
|
||||
}
|
||||
temporaryRoot := filepath.Join(root, ".tmp")
|
||||
if err := secureMkdirAll(root, temporaryRoot); err != nil {
|
||||
return nil, fmt.Errorf("create storage temporary root: %w", err)
|
||||
}
|
||||
allowed := make(map[string]bool, len(config.AllowedImageMIME))
|
||||
for mimeType, enabled := range config.AllowedImageMIME {
|
||||
if enabled {
|
||||
allowed[strings.ToLower(strings.TrimSpace(mimeType))] = true
|
||||
}
|
||||
}
|
||||
return &Local{
|
||||
root: root, temporaryRoot: temporaryRoot, maxObjectBytes: config.MaxObjectBytes,
|
||||
maxImagePixels: config.MaxImagePixels, thumbnailMaxSide: config.ThumbnailMaxSide,
|
||||
allowedImageMIME: allowed,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Local) Put(ctx context.Context, request corestorage.PutRequest) (object corestorage.Object, err error) {
|
||||
if request.OwnerID == 0 || request.GenerationID == 0 || strings.TrimSpace(request.ContentType) == "" || request.Source == nil {
|
||||
return corestorage.Object{}, ErrInvalidMetadata
|
||||
}
|
||||
finalDirectory, err := s.objectDirectory(request.Key)
|
||||
if err != nil {
|
||||
return corestorage.Object{}, err
|
||||
}
|
||||
if _, err := os.Lstat(finalDirectory); err == nil {
|
||||
return corestorage.Object{}, ErrObjectExists
|
||||
} else if !errors.Is(err, os.ErrNotExist) {
|
||||
return corestorage.Object{}, fmt.Errorf("inspect storage target: %w", err)
|
||||
}
|
||||
if err := secureMkdirAll(s.root, filepath.Dir(finalDirectory)); err != nil {
|
||||
return corestorage.Object{}, fmt.Errorf("create storage parent: %w", err)
|
||||
}
|
||||
temporaryDirectory, err := os.MkdirTemp(s.temporaryRoot, "put-")
|
||||
if err != nil {
|
||||
return corestorage.Object{}, fmt.Errorf("create storage temporary directory: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
if temporaryDirectory != "" {
|
||||
_ = os.RemoveAll(temporaryDirectory)
|
||||
}
|
||||
}()
|
||||
|
||||
contentPath := filepath.Join(temporaryDirectory, "content")
|
||||
content, err := os.OpenFile(contentPath, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600)
|
||||
if err != nil {
|
||||
return corestorage.Object{}, fmt.Errorf("create storage content: %w", err)
|
||||
}
|
||||
size, copyErr := copyLimitedContext(ctx, content, request.Source, s.maxObjectBytes)
|
||||
syncErr := content.Sync()
|
||||
closeErr := content.Close()
|
||||
if copyErr != nil {
|
||||
return corestorage.Object{}, copyErr
|
||||
}
|
||||
if syncErr != nil || closeErr != nil {
|
||||
return corestorage.Object{}, fmt.Errorf("flush storage content")
|
||||
}
|
||||
|
||||
storedMetadata := metadata{Key: request.Key, OwnerID: request.OwnerID, GenerationID: request.GenerationID, ContentType: request.ContentType, Size: size}
|
||||
metadataBytes, err := json.Marshal(storedMetadata)
|
||||
if err != nil {
|
||||
return corestorage.Object{}, fmt.Errorf("encode storage metadata: %w", err)
|
||||
}
|
||||
if err := writeSyncedFile(filepath.Join(temporaryDirectory, "metadata.json"), metadataBytes); err != nil {
|
||||
return corestorage.Object{}, err
|
||||
}
|
||||
if err := os.Rename(temporaryDirectory, finalDirectory); err != nil {
|
||||
return corestorage.Object{}, fmt.Errorf("atomically place storage object: %w", err)
|
||||
}
|
||||
temporaryDirectory = ""
|
||||
return objectFromMetadata(storedMetadata), nil
|
||||
}
|
||||
|
||||
func (s *Local) Open(ctx context.Context, key string) (io.ReadCloser, corestorage.Object, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, corestorage.Object{}, err
|
||||
}
|
||||
directory, err := s.objectDirectory(key)
|
||||
if err != nil {
|
||||
return nil, corestorage.Object{}, err
|
||||
}
|
||||
if err := ensureNoSymlink(s.root, directory); err != nil {
|
||||
return nil, corestorage.Object{}, err
|
||||
}
|
||||
metadataBytes, err := os.ReadFile(filepath.Join(directory, "metadata.json"))
|
||||
if err != nil {
|
||||
return nil, corestorage.Object{}, fmt.Errorf("read storage metadata: %w", err)
|
||||
}
|
||||
var storedMetadata metadata
|
||||
if err := json.Unmarshal(metadataBytes, &storedMetadata); err != nil || storedMetadata.Key != key || storedMetadata.OwnerID == 0 || storedMetadata.GenerationID == 0 {
|
||||
return nil, corestorage.Object{}, ErrInvalidMetadata
|
||||
}
|
||||
content, err := os.Open(filepath.Join(directory, "content"))
|
||||
if err != nil {
|
||||
return nil, corestorage.Object{}, fmt.Errorf("open storage content: %w", err)
|
||||
}
|
||||
return content, objectFromMetadata(storedMetadata), nil
|
||||
}
|
||||
|
||||
func (s *Local) Delete(ctx context.Context, key string) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
directory, err := s.objectDirectory(key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := os.Lstat(directory); errors.Is(err, os.ErrNotExist) {
|
||||
return nil
|
||||
} else if err != nil {
|
||||
return fmt.Errorf("inspect storage object: %w", err)
|
||||
}
|
||||
if err := ensureNoSymlink(s.root, directory); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.RemoveAll(directory); err != nil {
|
||||
return fmt.Errorf("delete storage object: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Local) PutImage(ctx context.Context, request ImageRequest) (ImageObjects, error) {
|
||||
if request.Key == "" || request.ThumbnailKey == "" || request.Key == request.ThumbnailKey {
|
||||
return ImageObjects{}, ErrInvalidKey
|
||||
}
|
||||
data, err := readLimitedContext(ctx, request.Source, s.maxObjectBytes)
|
||||
if err != nil {
|
||||
return ImageObjects{}, err
|
||||
}
|
||||
config, format, err := image.DecodeConfig(bytes.NewReader(data))
|
||||
if err != nil || config.Width <= 0 || config.Height <= 0 {
|
||||
return ImageObjects{}, ErrInvalidImage
|
||||
}
|
||||
actualMIME := imageMIME(format)
|
||||
if actualMIME == "" || !s.allowedImageMIME[actualMIME] || !strings.EqualFold(strings.TrimSpace(request.ContentType), actualMIME) {
|
||||
return ImageObjects{}, ErrUnsupportedImage
|
||||
}
|
||||
pixels := uint64(config.Width) * uint64(config.Height)
|
||||
if pixels > s.maxImagePixels {
|
||||
return ImageObjects{}, ErrImageTooLarge
|
||||
}
|
||||
decoded, _, err := image.Decode(bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return ImageObjects{}, ErrInvalidImage
|
||||
}
|
||||
thumbnail := imaging.Fit(decoded, s.thumbnailMaxSide, s.thumbnailMaxSide, imaging.Lanczos)
|
||||
var thumbnailData bytes.Buffer
|
||||
if err := png.Encode(&thumbnailData, thumbnail); err != nil {
|
||||
return ImageObjects{}, fmt.Errorf("encode image thumbnail: %w", err)
|
||||
}
|
||||
|
||||
original, err := s.Put(ctx, corestorage.PutRequest{
|
||||
Key: request.Key, OwnerID: request.OwnerID, GenerationID: request.GenerationID,
|
||||
ContentType: actualMIME, Source: bytes.NewReader(data),
|
||||
})
|
||||
if err != nil {
|
||||
return ImageObjects{}, err
|
||||
}
|
||||
thumb, err := s.Put(ctx, corestorage.PutRequest{
|
||||
Key: request.ThumbnailKey, OwnerID: request.OwnerID, GenerationID: request.GenerationID,
|
||||
ContentType: "image/png", Source: bytes.NewReader(thumbnailData.Bytes()),
|
||||
})
|
||||
if err != nil {
|
||||
_ = s.Delete(context.Background(), request.Key)
|
||||
return ImageObjects{}, err
|
||||
}
|
||||
return ImageObjects{Original: original, Thumbnail: thumb}, nil
|
||||
}
|
||||
|
||||
func (s *Local) objectDirectory(key string) (string, error) {
|
||||
if key == "" || strings.Contains(key, "\\") || strings.Contains(key, ":") || strings.HasPrefix(key, "/") {
|
||||
return "", ErrInvalidKey
|
||||
}
|
||||
parts := strings.Split(key, "/")
|
||||
for _, part := range parts {
|
||||
if part == "" || part == "." || part == ".." {
|
||||
return "", ErrInvalidKey
|
||||
}
|
||||
}
|
||||
path := filepath.Join(append([]string{s.root, "objects"}, parts...)...)
|
||||
relative, err := filepath.Rel(s.root, path)
|
||||
if err != nil || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) {
|
||||
return "", ErrInvalidKey
|
||||
}
|
||||
return path, nil
|
||||
}
|
||||
|
||||
func ensureNoSymlink(root, target string) error {
|
||||
relative, err := filepath.Rel(root, target)
|
||||
if err != nil || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) {
|
||||
return ErrInvalidKey
|
||||
}
|
||||
current := root
|
||||
for _, part := range strings.Split(relative, string(filepath.Separator)) {
|
||||
if part == "." || part == "" {
|
||||
continue
|
||||
}
|
||||
current = filepath.Join(current, part)
|
||||
info, err := os.Lstat(current)
|
||||
if err != nil {
|
||||
return fmt.Errorf("inspect storage path: %w", err)
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 {
|
||||
return ErrInvalidKey
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func secureMkdirAll(root, target string) error {
|
||||
relative, err := filepath.Rel(root, target)
|
||||
if err != nil || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) {
|
||||
return ErrInvalidKey
|
||||
}
|
||||
current := root
|
||||
for _, part := range strings.Split(relative, string(filepath.Separator)) {
|
||||
if part == "." || part == "" {
|
||||
continue
|
||||
}
|
||||
current = filepath.Join(current, part)
|
||||
info, err := os.Lstat(current)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
if err := os.Mkdir(current, 0o700); err != nil && !errors.Is(err, os.ErrExist) {
|
||||
return err
|
||||
}
|
||||
info, err = os.Lstat(current)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
|
||||
return ErrInvalidKey
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func copyLimitedContext(ctx context.Context, destination io.Writer, source io.Reader, limit int64) (int64, error) {
|
||||
buffer := make([]byte, 32*1024)
|
||||
var total int64
|
||||
for {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return total, err
|
||||
}
|
||||
read, readErr := source.Read(buffer)
|
||||
if read > 0 {
|
||||
if total+int64(read) > limit {
|
||||
return total, ErrObjectTooLarge
|
||||
}
|
||||
written, writeErr := destination.Write(buffer[:read])
|
||||
total += int64(written)
|
||||
if writeErr != nil {
|
||||
return total, writeErr
|
||||
}
|
||||
if written != read {
|
||||
return total, io.ErrShortWrite
|
||||
}
|
||||
}
|
||||
if readErr != nil {
|
||||
if errors.Is(readErr, io.EOF) {
|
||||
return total, nil
|
||||
}
|
||||
return total, readErr
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func readLimitedContext(ctx context.Context, source io.Reader, limit int64) ([]byte, error) {
|
||||
var data bytes.Buffer
|
||||
_, err := copyLimitedContext(ctx, &data, source, limit)
|
||||
return data.Bytes(), err
|
||||
}
|
||||
|
||||
func writeSyncedFile(path string, data []byte) error {
|
||||
file, err := os.OpenFile(path, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600)
|
||||
if err != nil {
|
||||
return fmt.Errorf("create storage metadata: %w", err)
|
||||
}
|
||||
if _, err := file.Write(data); err != nil {
|
||||
file.Close()
|
||||
return fmt.Errorf("write storage metadata: %w", err)
|
||||
}
|
||||
if err := file.Sync(); err != nil {
|
||||
file.Close()
|
||||
return fmt.Errorf("flush storage metadata: %w", err)
|
||||
}
|
||||
if err := file.Close(); err != nil {
|
||||
return fmt.Errorf("close storage metadata: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func imageMIME(format string) string {
|
||||
switch strings.ToLower(format) {
|
||||
case "jpeg":
|
||||
return "image/jpeg"
|
||||
case "png":
|
||||
return "image/png"
|
||||
case "gif":
|
||||
return "image/gif"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func objectFromMetadata(stored metadata) corestorage.Object {
|
||||
return corestorage.Object{
|
||||
Key: stored.Key, OwnerID: stored.OwnerID, GenerationID: stored.GenerationID,
|
||||
ContentType: stored.ContentType, Size: stored.Size,
|
||||
}
|
||||
}
|
||||
|
||||
var _ corestorage.Store = (*Local)(nil)
|
||||
@@ -0,0 +1,151 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"image"
|
||||
"image/color"
|
||||
"image/png"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage"
|
||||
)
|
||||
|
||||
func newTestStore(t *testing.T, maxBytes int64, maxPixels uint64) *Local {
|
||||
t.Helper()
|
||||
store, err := NewLocal(Config{
|
||||
Root: t.TempDir(), MaxObjectBytes: maxBytes, MaxImagePixels: maxPixels, ThumbnailMaxSide: 256,
|
||||
AllowedImageMIME: map[string]bool{"image/png": true, "image/jpeg": true},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return store
|
||||
}
|
||||
|
||||
func TestLocalPutOpenAtomicMetadataAndTraversal(t *testing.T) {
|
||||
store := newTestStore(t, 1024, 1_000_000)
|
||||
object, err := store.Put(context.Background(), corestorage.PutRequest{
|
||||
Key: "users/7/generations/9/result", OwnerID: 7, GenerationID: 9,
|
||||
ContentType: "text/plain", Source: bytes.NewBufferString("synthetic result"),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if object.OwnerID != 7 || object.GenerationID != 9 || object.Size != 16 {
|
||||
t.Fatalf("unexpected object: %#v", object)
|
||||
}
|
||||
reader, metadata, err := store.Open(context.Background(), object.Key)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer reader.Close()
|
||||
data, _ := io.ReadAll(reader)
|
||||
if string(data) != "synthetic result" || metadata.OwnerID != 7 {
|
||||
t.Fatalf("open data=%q metadata=%#v", data, metadata)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(store.root, "objects", "users", "7", "generations", "9", "result", "metadata.json")); err != nil {
|
||||
t.Fatalf("metadata not atomically placed: %v", err)
|
||||
}
|
||||
for _, key := range []string{"../outside", "/absolute", "a\\b", "C:/escape", "a//b"} {
|
||||
_, err := store.Put(context.Background(), corestorage.PutRequest{Key: key, OwnerID: 1, GenerationID: 1, ContentType: "x", Source: bytes.NewReader(nil)})
|
||||
if !errors.Is(err, ErrInvalidKey) {
|
||||
t.Errorf("Put(%q) error = %v", key, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalFailureCleansTemporaryFiles(t *testing.T) {
|
||||
store := newTestStore(t, 4, 1_000_000)
|
||||
_, err := store.Put(context.Background(), corestorage.PutRequest{
|
||||
Key: "too-large", OwnerID: 1, GenerationID: 1, ContentType: "text/plain", Source: bytes.NewBufferString("12345"),
|
||||
})
|
||||
if !errors.Is(err, ErrObjectTooLarge) {
|
||||
t.Fatalf("Put() error = %v", err)
|
||||
}
|
||||
entries, err := os.ReadDir(store.temporaryRoot)
|
||||
if err != nil || len(entries) != 0 {
|
||||
t.Fatalf("temporary directory not cleaned: entries=%v error=%v", entries, err)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(store.root, "objects", "too-large")); !errors.Is(err, os.ErrNotExist) {
|
||||
t.Fatalf("failed object unexpectedly exists: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalRejectsSymlinkComponents(t *testing.T) {
|
||||
store := newTestStore(t, 1024, 1_000_000)
|
||||
outside := t.TempDir()
|
||||
objectsRoot := filepath.Join(store.root, "objects")
|
||||
if err := os.Mkdir(objectsRoot, 0o700); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.Symlink(outside, filepath.Join(objectsRoot, "linked")); err != nil {
|
||||
t.Skipf("symlinks unavailable on this platform: %v", err)
|
||||
}
|
||||
_, err := store.Put(context.Background(), corestorage.PutRequest{
|
||||
Key: "linked/escape", OwnerID: 1, GenerationID: 1, ContentType: "text/plain", Source: bytes.NewBufferString("x"),
|
||||
})
|
||||
if !errors.Is(err, ErrInvalidKey) {
|
||||
t.Fatalf("symlink component error = %v", err)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(outside, "escape")); !errors.Is(err, os.ErrNotExist) {
|
||||
t.Fatalf("storage wrote through symlink: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPutImageValidatesPixelsMIMEAndCreatesThumbnail(t *testing.T) {
|
||||
store := newTestStore(t, 2<<20, 512*512)
|
||||
imageData := encodePNG(t, 512, 256)
|
||||
objects, err := store.PutImage(context.Background(), ImageRequest{
|
||||
Key: "images/original", ThumbnailKey: "images/thumb", OwnerID: 2, GenerationID: 3,
|
||||
ContentType: "image/png", Source: bytes.NewReader(imageData),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
reader, _, err := store.Open(context.Background(), objects.Thumbnail.Key)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer reader.Close()
|
||||
config, format, err := image.DecodeConfig(reader)
|
||||
if err != nil || format != "png" || config.Width != 256 || config.Height != 128 {
|
||||
t.Fatalf("thumbnail config=%#v format=%s error=%v", config, format, err)
|
||||
}
|
||||
|
||||
_, err = store.PutImage(context.Background(), ImageRequest{
|
||||
Key: "images/mismatch", ThumbnailKey: "images/mismatch-thumb", OwnerID: 2, GenerationID: 3,
|
||||
ContentType: "image/jpeg", Source: bytes.NewReader(imageData),
|
||||
})
|
||||
if !errors.Is(err, ErrUnsupportedImage) {
|
||||
t.Fatalf("MIME mismatch error = %v", err)
|
||||
}
|
||||
|
||||
smallLimit := newTestStore(t, 2<<20, 10_000)
|
||||
_, err = smallLimit.PutImage(context.Background(), ImageRequest{
|
||||
Key: "images/bomb", ThumbnailKey: "images/bomb-thumb", OwnerID: 2, GenerationID: 3,
|
||||
ContentType: "image/png", Source: bytes.NewReader(imageData),
|
||||
})
|
||||
if !errors.Is(err, ErrImageTooLarge) {
|
||||
t.Fatalf("pixel limit error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func encodePNG(t *testing.T, width, height int) []byte {
|
||||
t.Helper()
|
||||
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
||||
for y := 0; y < height; y++ {
|
||||
for x := 0; x < width; x++ {
|
||||
img.Set(x, y, color.RGBA{R: uint8(x), G: uint8(y), B: 100, A: 255})
|
||||
}
|
||||
}
|
||||
var data bytes.Buffer
|
||||
if err := png.Encode(&data, img); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return data.Bytes()
|
||||
}
|
||||
Reference in New Issue
Block a user