feat: add secure HTTP crypto and local storage (#8)

This commit is contained in:
ila
2026-08-20 23:34:52 +08:00
parent 62017df9b0
commit ceac631638
9 changed files with 1175 additions and 5 deletions
+2
View File
@@ -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
)
+6
View File
@@ -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=
+15 -5
View File
@@ -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
}
+108
View File
@@ -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)
+63
View File
@@ -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)
}
}
+256
View File
@@ -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
}
+178
View File
@@ -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)
}
}
+396
View File
@@ -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)
+151
View File
@@ -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()
}