diff --git a/go.mod b/go.mod index 710f9b0..9f94c43 100644 --- a/go.mod +++ b/go.mod @@ -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 ) diff --git a/go.sum b/go.sum index 388b706..f67c1ef 100644 --- a/go.sum +++ b/go.sum @@ -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= diff --git a/internal/core/storage/storage.go b/internal/core/storage/storage.go index 1ad1bb3..e600116 100644 --- a/internal/core/storage/storage.go +++ b/internal/core/storage/storage.go @@ -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 } diff --git a/internal/platform/crypto/keyring.go b/internal/platform/crypto/keyring.go new file mode 100644 index 0000000..e9cc4de --- /dev/null +++ b/internal/platform/crypto/keyring.go @@ -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) diff --git a/internal/platform/crypto/keyring_test.go b/internal/platform/crypto/keyring_test.go new file mode 100644 index 0000000..fad38e1 --- /dev/null +++ b/internal/platform/crypto/keyring_test.go @@ -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) + } +} diff --git a/internal/platform/http/safehttp.go b/internal/platform/http/safehttp.go new file mode 100644 index 0000000..6901976 --- /dev/null +++ b/internal/platform/http/safehttp.go @@ -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 "" + } + return strings.ToLower(target.Scheme) + "://" + target.Host +} diff --git a/internal/platform/http/safehttp_test.go b/internal/platform/http/safehttp_test.go new file mode 100644 index 0000000..ef3bea3 --- /dev/null +++ b/internal/platform/http/safehttp_test.go @@ -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) + } +} diff --git a/internal/platform/storage/local.go b/internal/platform/storage/local.go new file mode 100644 index 0000000..4b4f810 --- /dev/null +++ b/internal/platform/storage/local.go @@ -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) diff --git a/internal/platform/storage/local_test.go b/internal/platform/storage/local_test.go new file mode 100644 index 0000000..84dcdcb --- /dev/null +++ b/internal/platform/storage/local_test.go @@ -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() +}