mirror of
https://github.com/versity/versitygw.git
synced 2026-09-24 17:04:16 +00:00
feat: add AWS-compatible standalone IAM service
Closes #1640 Add a standalone AWS IAM Query API implementation for managing IAM users through standard AWS SDKs and the AWS CLI. Server usage Start the IAM server with internal file-backed storage: mkdir -p /tmp/versitygw-iam ./versitygw --port 127.0.0.1:7070 --access user --secret pass iam --dir /tmp/versitygw-iam Start the IAM server with Vault KV v2 storage using AppRole: VGW_IAM_VAULT_ROLE_SECRET=<role-secret> ./versitygw --port 127.0.0.1:7070 --access user --secret pass iam --vault-endpoint-url http://127.0.0.1:8200 --vault-auth-method approle --vault-role-id <role-id> --vault-mount-path kv --vault-secret-storage-path iam Vault authentication also supports root tokens, separate authentication and secret-storage namespaces, custom mount paths, server certificate validation, and mutual TLS client certificates. Configure the AWS CLI credentials used by the IAM server: export AWS_ACCESS_KEY_ID=user export AWS_SECRET_ACCESS_KEY=pass export AWS_DEFAULT_REGION=us-east-1 Implemented IAM actions CreateUser creates an IAM user with an AWS-compatible ARN, generated AIDA user ID, creation timestamp, optional path, and tags. It validates usernames, paths, tag limits, reserved tag prefixes, duplicate tag keys, and existing users. aws --endpoint-url http://127.0.0.1:7070 iam create-user --user-name bob aws --endpoint-url http://127.0.0.1:7070 iam create-user --user-name bob --path /engineering/ --tags Key=team,Value=storage GetUser returns a stored user or the root identity when requested without a username through the IAM Query API. aws --endpoint-url http://127.0.0.1:7070 iam get-user --user-name bob ListUsers returns users in deterministic username order and supports path filtering, marker-based pagination, and MaxItems limits. aws --endpoint-url http://127.0.0.1:7070 iam list-users aws --endpoint-url http://127.0.0.1:7070 iam list-users --path-prefix /engineering/ --max-items 100 UpdateUser updates the username and/or path, recalculates the user ARN, and rejects conflicts with existing users. aws --endpoint-url http://127.0.0.1:7070 iam update-user --user-name bob --new-user-name robert --new-path /platform/ DeleteUser permanently removes an IAM user and returns AWS-compatible errors for missing users. aws --endpoint-url http://127.0.0.1:7070 iam delete-user --user-name robert IAM protocol and authentication - Support the AWS IAM Query protocol version 2010-05-08 over GET and POST form requests. - Return AWS-compatible XML responses, error documents, status codes, request IDs, user metadata, and pagination fields. - Authenticate root credentials with AWS Signature Version 4 for the IAM service in us-east-1. - Support both Authorization-header and query-string SigV4 authentication. - Validate credential scope, signed headers, timestamps, clock skew, content length, signatures, and unsupported signature or session-token modes. - Add IAM-specific validation and error mapping for malformed requests, invalid actions, duplicate entities, missing users, throttling, and internal failures. Storage implementations - Add an internal JSON-backed store using iam.json and iam.json.backup with atomic temporary-file replacement, concurrent access protection, stable ordering, pagination, and persistence across restarts. - Add a Vault KV v2 store with one secret per user, CAS-based duplicate protection, permanent deletion, AppRole reauthentication, namespace support, configurable authentication and KV mounts, root-token authentication, and TLS/mTLS configuration. - Introduce a common Storer interface and require exactly one storage backend to be configured. Server and embedding support - Register the new `versitygw iam` command with environment-variable and CLI configuration for both storage backends. - Add `embedgw.RunIAMAPI` and `IAMConfig` for embedding the IAM service in Go applications. Gateway-level internal packages - Add `internal/iamstore` as a reusable generic file-backed IAM persistence engine and migrate the existing gateway internal IAM service to it. - Add `internal/sigv4auth` for shared SigV4 header and presigned-query parsing, canonical request generation, signature verification, and structured authentication errors. - Refactor the S3 authentication paths to use the shared SigV4 implementation while preserving S3-specific error responses. - Add `internal/httpctx` for shared Fiber context keys and AWS-style request ID handling. - Add `internal/routekit` for shared query, form, and header route matchers. - Add `internal/netutil` for reusable certificate storage, hostname-aware listeners, multi-address serving, TLS listeners, and UNIX socket handling. - Update the custom SigV4 signer to honor an explicitly supplied signed-header list so unrelated headers do not alter IAM signatures. Testing and CI - Add AWS IAM SDK-based integration coverage for all supported user actions, header authentication, query authentication, validation, errors, filtering, and pagination. - Split standalone IAM tests into `versitygw test iam` and retain existing gateway IAM tests under `versitygw test gw-iam`. - Add unit coverage for controllers, authentication, routing, storage, embedding, listeners, request matching, persistence, and signing behavior. - Add `runiamtests.sh` to exercise internal storage over HTTP and HTTPS plus Vault storage through AppRole. - Add a dedicated IAM functional-test workflow with a Vault service and merged runtime coverage reporting. - Include the IAM test runner in shellcheck and add the AWS IAM SDK dependency.
This commit is contained in:
@@ -0,0 +1,56 @@
|
||||
// Copyright 2026 Versity Software
|
||||
// This file is licensed under the Apache License, Version 2.0
|
||||
// (the "License"); you may not use this file except in compliance
|
||||
// with the License. You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing,
|
||||
// software distributed under the License is distributed on an
|
||||
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
// KIND, either express or implied. See the License for the
|
||||
// specific language governing permissions and limitations
|
||||
// under the License.
|
||||
|
||||
package httpctx
|
||||
|
||||
import "github.com/gofiber/fiber/v3"
|
||||
|
||||
// ContextKey names a request-local value stored in fiber.Ctx locals.
|
||||
type ContextKey string
|
||||
|
||||
const (
|
||||
ContextKeyRegion ContextKey = "region"
|
||||
ContextKeyStartTime ContextKey = "start-time"
|
||||
ContextKeyIsRoot ContextKey = "is-root"
|
||||
ContextKeyRootAccessKey ContextKey = "root-access-key"
|
||||
ContextKeyAccount ContextKey = "account"
|
||||
ContextKeyAuthenticated ContextKey = "authenticated"
|
||||
ContextKeyPublicBucket ContextKey = "public-bucket"
|
||||
ContextKeyParsedAcl ContextKey = "parsed-acl"
|
||||
ContextKeySkipResBodyLog ContextKey = "skip-res-body-log"
|
||||
ContextKeyBodyReader ContextKey = "body-reader"
|
||||
ContextKeySkip ContextKey = "__skip"
|
||||
ContextKeyStack ContextKey = "stack"
|
||||
ContextKeyBucketOwner ContextKey = "bucket-owner"
|
||||
ContextKeyObjectPostResult ContextKey = "object-post-result"
|
||||
ContextKeyRequestID ContextKey = "request-id"
|
||||
ContextKeyHostID ContextKey = "host-id"
|
||||
ContextKeyWebsiteConfig ContextKey = "website-config"
|
||||
)
|
||||
|
||||
func (ck ContextKey) Set(ctx fiber.Ctx, val any) {
|
||||
ctx.Locals(string(ck), val)
|
||||
}
|
||||
|
||||
func (ck ContextKey) IsSet(ctx fiber.Ctx) bool {
|
||||
return ctx.Locals(string(ck)) != nil
|
||||
}
|
||||
|
||||
func (ck ContextKey) Delete(ctx fiber.Ctx) {
|
||||
ctx.Locals(string(ck), nil)
|
||||
}
|
||||
|
||||
func (ck ContextKey) Get(ctx fiber.Ctx) any {
|
||||
return ctx.Locals(string(ck))
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
// Copyright 2026 Versity Software
|
||||
// This file is licensed under the Apache License, Version 2.0
|
||||
// (the "License"); you may not use this file except in compliance
|
||||
// with the License. You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing,
|
||||
// software distributed under the License is distributed on an
|
||||
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
// KIND, either express or implied. See the License for the
|
||||
// specific language governing permissions and limitations
|
||||
// under the License.
|
||||
|
||||
package iamstore
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
iamMode = 0600
|
||||
backoff = 100 * time.Millisecond
|
||||
maxretry = 300
|
||||
)
|
||||
|
||||
// UpdateFunc accepts the current JSON data and returns the new JSON data to store.
|
||||
type UpdateFunc func([]byte) ([]byte, error)
|
||||
|
||||
type NormalizeFunc[T any] func(*T)
|
||||
|
||||
type Engine[T any] struct {
|
||||
dir string
|
||||
iamFile string
|
||||
iamBackupFile string
|
||||
defaultConfig T
|
||||
normalize NormalizeFunc[T]
|
||||
}
|
||||
|
||||
func New[T any](dir, iamFile, iamBackupFile string, defaultConfig T, normalize NormalizeFunc[T]) (*Engine[T], error) {
|
||||
engine := &Engine[T]{
|
||||
dir: dir,
|
||||
iamFile: iamFile,
|
||||
iamBackupFile: iamBackupFile,
|
||||
defaultConfig: defaultConfig,
|
||||
normalize: normalize,
|
||||
}
|
||||
|
||||
if err := engine.InitIAM(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return engine, nil
|
||||
}
|
||||
|
||||
func (e *Engine[T]) InitIAM() error {
|
||||
fname := filepath.Join(e.dir, e.iamFile)
|
||||
|
||||
_, err := os.ReadFile(fname)
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
b, err := json.Marshal(e.defaultConfig)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal default iam: %w", err)
|
||||
}
|
||||
err = os.WriteFile(fname, b, iamMode)
|
||||
if err != nil {
|
||||
return fmt.Errorf("write default iam: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *Engine[T]) GetIAM() (T, error) {
|
||||
b, err := e.ReadIAMData()
|
||||
if err != nil {
|
||||
var zero T
|
||||
return zero, err
|
||||
}
|
||||
|
||||
return e.ParseIAM(b)
|
||||
}
|
||||
|
||||
func (e *Engine[T]) ParseIAM(b []byte) (T, error) {
|
||||
return ParseIAM(b, e.normalize)
|
||||
}
|
||||
|
||||
func ParseIAM[T any](b []byte, normalize NormalizeFunc[T]) (T, error) {
|
||||
var conf T
|
||||
if err := json.Unmarshal(b, &conf); err != nil {
|
||||
return conf, fmt.Errorf("failed to parse the config file: %w", err)
|
||||
}
|
||||
|
||||
if normalize != nil {
|
||||
normalize(&conf)
|
||||
}
|
||||
|
||||
return conf, nil
|
||||
}
|
||||
|
||||
func (e *Engine[T]) ReadIAMData() ([]byte, error) {
|
||||
// We are going to be racing with other running gateways without any
|
||||
// coordination. So we might find the file does not exist at times.
|
||||
// For this case we need to retry for a while assuming the other gateway
|
||||
// will eventually write the file. If it doesn't after the max retries,
|
||||
// then we will return the error.
|
||||
|
||||
retries := 0
|
||||
|
||||
for {
|
||||
b, err := os.ReadFile(filepath.Join(e.dir, e.iamFile))
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
// racing with someone else updating
|
||||
// keep retrying after backoff
|
||||
retries++
|
||||
if retries < maxretry {
|
||||
time.Sleep(backoff)
|
||||
continue
|
||||
}
|
||||
return nil, fmt.Errorf("read iam file: %w", err)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return b, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (e *Engine[T]) StoreIAM(update UpdateFunc) error {
|
||||
// We are going to be racing with other running gateways without any
|
||||
// coordination. So the strategy here is to read the current file data,
|
||||
// update the data, write back out to a temp file, then rename the
|
||||
// temp file to the original file. This rename will replace the
|
||||
// original file with the new file. This is atomic and should always
|
||||
// allow for a consistent view of the data. There is a small
|
||||
// window where the file could be read and then updated by
|
||||
// another process. In this case any updates the other process did
|
||||
// will be lost. This is a limitation of the internal IAM service.
|
||||
// This should be rare, and even when it does happen should result
|
||||
// in a valid IAM file, just without the other process's updates.
|
||||
|
||||
iamFname := filepath.Join(e.dir, e.iamFile)
|
||||
backupFname := filepath.Join(e.dir, e.iamBackupFile)
|
||||
|
||||
b, err := os.ReadFile(iamFname)
|
||||
if err != nil && !errors.Is(err, fs.ErrNotExist) {
|
||||
return fmt.Errorf("read iam file: %w", err)
|
||||
}
|
||||
|
||||
err = e.writeUsingTempFile(b, backupFname)
|
||||
if err != nil {
|
||||
return fmt.Errorf("write backup iam file: %w", err)
|
||||
}
|
||||
|
||||
b, err = update(b)
|
||||
if err != nil {
|
||||
return fmt.Errorf("update iam data: %w", err)
|
||||
}
|
||||
|
||||
err = e.writeUsingTempFile(b, iamFname)
|
||||
if err != nil {
|
||||
return fmt.Errorf("write iam file: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *Engine[T]) writeUsingTempFile(b []byte, fname string) error {
|
||||
f, err := os.CreateTemp(e.dir, e.iamFile)
|
||||
if err != nil {
|
||||
return fmt.Errorf("create temp file: %w", err)
|
||||
}
|
||||
defer os.Remove(f.Name())
|
||||
|
||||
_, err = f.Write(b)
|
||||
f.Close()
|
||||
if err != nil {
|
||||
return fmt.Errorf("write temp file: %w", err)
|
||||
}
|
||||
|
||||
err = os.Rename(f.Name(), fname)
|
||||
if err != nil {
|
||||
return fmt.Errorf("rename temp file: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
// Copyright 2026 Versity Software
|
||||
// This file is licensed under the Apache License, Version 2.0
|
||||
// (the "License"); you may not use this file except in compliance
|
||||
// with the License. You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing,
|
||||
// software distributed under the License is distributed on an
|
||||
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
// KIND, either express or implied. See the License for the
|
||||
// specific language governing permissions and limitations
|
||||
// under the License.
|
||||
|
||||
package iamstore
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type testConfig struct {
|
||||
Users map[string]string `json:"users"`
|
||||
}
|
||||
|
||||
func TestEngineInitializesReadsParsesAndStoresJSON(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
|
||||
engine, err := New(dir, "users.json", "users.json.backup", testConfig{Users: map[string]string{}}, func(conf *testConfig) {
|
||||
if conf.Users == nil {
|
||||
conf.Users = map[string]string{}
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("New: %v", err)
|
||||
}
|
||||
|
||||
conf, err := engine.GetIAM()
|
||||
if err != nil {
|
||||
t.Fatalf("GetIAM: %v", err)
|
||||
}
|
||||
if conf.Users == nil {
|
||||
t.Fatal("GetIAM returned nil Users map")
|
||||
}
|
||||
|
||||
err = engine.StoreIAM(func(data []byte) ([]byte, error) {
|
||||
conf, err := engine.ParseIAM(data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
conf.Users["alice"] = "created"
|
||||
return json.Marshal(conf)
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("StoreIAM: %v", err)
|
||||
}
|
||||
|
||||
conf, err = engine.GetIAM()
|
||||
if err != nil {
|
||||
t.Fatalf("GetIAM after store: %v", err)
|
||||
}
|
||||
if conf.Users["alice"] != "created" {
|
||||
t.Fatalf("stored user = %q, want created", conf.Users["alice"])
|
||||
}
|
||||
|
||||
if _, err := os.Stat(filepath.Join(dir, "users.json.backup")); err != nil {
|
||||
t.Fatalf("stat backup file: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
// Copyright 2026 Versity Software
|
||||
// This file is licensed under the Apache License, Version 2.0
|
||||
// (the "License"); you may not use this file except in compliance
|
||||
// with the License. You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing,
|
||||
// software distributed under the License is distributed on an
|
||||
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
// KIND, either express or implied. See the License for the
|
||||
// specific language governing permissions and limitations
|
||||
// under the License.
|
||||
|
||||
package netutil
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"fmt"
|
||||
"sync/atomic"
|
||||
)
|
||||
|
||||
type CertStorage struct {
|
||||
cert atomic.Pointer[tls.Certificate]
|
||||
}
|
||||
|
||||
func NewCertStorage() *CertStorage {
|
||||
return &CertStorage{}
|
||||
}
|
||||
|
||||
func (cs *CertStorage) GetCertificate(_ *tls.ClientHelloInfo) (*tls.Certificate, error) {
|
||||
return cs.cert.Load(), nil
|
||||
}
|
||||
|
||||
func (cs *CertStorage) SetCertificate(certFile string, keyFile string) error {
|
||||
cert, err := tls.LoadX509KeyPair(certFile, keyFile)
|
||||
if err != nil {
|
||||
return fmt.Errorf("unable to set certificate: %w", err)
|
||||
}
|
||||
|
||||
cs.cert.Store(&cert)
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,316 @@
|
||||
// Copyright 2026 Versity Software
|
||||
// This file is licensed under the Apache License, Version 2.0
|
||||
// (the "License"); you may not use this file except in compliance
|
||||
// with the License. You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing,
|
||||
// software distributed under the License is distributed on an
|
||||
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
// KIND, either express or implied. See the License for the
|
||||
// specific language governing permissions and limitations
|
||||
// under the License.
|
||||
|
||||
package netutil
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// MultiListener implements net.Listener and accepts connections from multiple
|
||||
// underlying listeners.
|
||||
type MultiListener struct {
|
||||
listeners []net.Listener
|
||||
acceptCh chan acceptResult
|
||||
closeCh chan struct{}
|
||||
closeOnce sync.Once
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
type acceptResult struct {
|
||||
conn net.Conn
|
||||
err error
|
||||
}
|
||||
|
||||
func NewMultiListener(listeners ...net.Listener) *MultiListener {
|
||||
if len(listeners) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
ml := &MultiListener{
|
||||
listeners: listeners,
|
||||
acceptCh: make(chan acceptResult, 2*len(listeners)),
|
||||
closeCh: make(chan struct{}),
|
||||
}
|
||||
|
||||
for _, ln := range listeners {
|
||||
ml.wg.Add(1)
|
||||
go ml.acceptLoop(ln)
|
||||
}
|
||||
|
||||
return ml
|
||||
}
|
||||
|
||||
func (ml *MultiListener) acceptLoop(ln net.Listener) {
|
||||
defer ml.wg.Done()
|
||||
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
|
||||
select {
|
||||
case <-ml.closeCh:
|
||||
if conn != nil {
|
||||
conn.Close()
|
||||
}
|
||||
return
|
||||
case ml.acceptCh <- acceptResult{conn: conn, err: err}:
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (ml *MultiListener) Accept() (net.Conn, error) {
|
||||
select {
|
||||
case <-ml.closeCh:
|
||||
return nil, errors.New("listener closed")
|
||||
case result, ok := <-ml.acceptCh:
|
||||
if !ok {
|
||||
return nil, errors.New("listener closed")
|
||||
}
|
||||
return result.conn, result.err
|
||||
}
|
||||
}
|
||||
|
||||
func (ml *MultiListener) Close() error {
|
||||
var errs []error
|
||||
|
||||
ml.closeOnce.Do(func() {
|
||||
close(ml.closeCh)
|
||||
|
||||
for _, ln := range ml.listeners {
|
||||
if err := ln.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
ml.wg.Wait()
|
||||
|
||||
close(ml.acceptCh)
|
||||
for range ml.acceptCh {
|
||||
}
|
||||
})
|
||||
|
||||
if len(errs) > 0 {
|
||||
return fmt.Errorf("errors closing listeners: %v", errs)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ml *MultiListener) Addr() net.Addr {
|
||||
if len(ml.listeners) > 0 {
|
||||
return ml.listeners[0].Addr()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func IsUnixSocketPath(addr string) bool {
|
||||
_, _, err := net.SplitHostPort(addr)
|
||||
return err != nil
|
||||
}
|
||||
|
||||
func AbsSocketPaths(addrs []string) ([]string, error) {
|
||||
result := make([]string, len(addrs))
|
||||
for i, addr := range addrs {
|
||||
if strings.HasPrefix(addr, "./") {
|
||||
abs, err := filepath.Abs(addr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to resolve socket path %q: %w", addr, err)
|
||||
}
|
||||
result[i] = abs
|
||||
} else {
|
||||
result[i] = addr
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func isAbstractSocket(addr string) bool {
|
||||
return strings.HasPrefix(addr, "@")
|
||||
}
|
||||
|
||||
func removeStaleSocket(path string) error {
|
||||
fi, err := os.Stat(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("failed to stat socket path %q: %w", path, err)
|
||||
}
|
||||
if fi.Mode()&os.ModeSocket == 0 {
|
||||
return fmt.Errorf("path %q already exists and is not a socket (mode %s)", path, fi.Mode())
|
||||
}
|
||||
return os.Remove(path)
|
||||
}
|
||||
|
||||
func ResolveHostnameIPs(address string) ([]string, error) {
|
||||
if IsUnixSocketPath(address) {
|
||||
return []string{address}, nil
|
||||
}
|
||||
|
||||
host, _, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid address %q: %w", address, err)
|
||||
}
|
||||
|
||||
if host == "" {
|
||||
return []string{""}, nil
|
||||
}
|
||||
|
||||
if net.ParseIP(host) != nil {
|
||||
return []string{host}, nil
|
||||
}
|
||||
|
||||
ips, err := net.LookupIP(host)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to resolve hostname %q: %w", host, err)
|
||||
}
|
||||
if len(ips) == 0 {
|
||||
return nil, fmt.Errorf("no addresses found for hostname %q", host)
|
||||
}
|
||||
|
||||
result := make([]string, 0, len(ips))
|
||||
for _, ip := range ips {
|
||||
result = append(result, ip.String())
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func resolveHostnameAddrs(address string) ([]string, error) {
|
||||
if IsUnixSocketPath(address) {
|
||||
return []string{address}, nil
|
||||
}
|
||||
|
||||
host, port, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid address %q: %w", address, err)
|
||||
}
|
||||
|
||||
if host == "" || net.ParseIP(host) != nil {
|
||||
return []string{address}, nil
|
||||
}
|
||||
|
||||
ips, err := net.LookupIP(host)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to resolve hostname %q: %w", host, err)
|
||||
}
|
||||
if len(ips) == 0 {
|
||||
return nil, fmt.Errorf("no addresses found for hostname %q", host)
|
||||
}
|
||||
|
||||
addrs := make([]string, 0, len(ips))
|
||||
for _, ip := range ips {
|
||||
addrs = append(addrs, net.JoinHostPort(ip.String(), port))
|
||||
}
|
||||
|
||||
return addrs, nil
|
||||
}
|
||||
|
||||
type ListenerOptions struct {
|
||||
SocketPerm os.FileMode
|
||||
}
|
||||
|
||||
func NewMultiAddrListener(network, address string, opts ListenerOptions) (net.Listener, error) {
|
||||
if IsUnixSocketPath(address) {
|
||||
if !isAbstractSocket(address) {
|
||||
if err := removeStaleSocket(address); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
ln, err := net.Listen("unix", address)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to bind unix socket listener %s: %w", address, err)
|
||||
}
|
||||
if opts.SocketPerm != 0 && !isAbstractSocket(address) {
|
||||
if err := os.Chmod(address, opts.SocketPerm); err != nil {
|
||||
ln.Close()
|
||||
return nil, fmt.Errorf("failed to set permissions on socket %s: %w", address, err)
|
||||
}
|
||||
}
|
||||
return NewMultiListener(ln), nil
|
||||
}
|
||||
|
||||
addrs, err := resolveHostnameAddrs(address)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
listeners := make([]net.Listener, 0, len(addrs))
|
||||
for _, addr := range addrs {
|
||||
ln, err := net.Listen(network, addr)
|
||||
if err != nil {
|
||||
for _, l := range listeners {
|
||||
l.Close()
|
||||
}
|
||||
return nil, fmt.Errorf("failed to bind listener %s: %w", addr, err)
|
||||
}
|
||||
listeners = append(listeners, ln)
|
||||
}
|
||||
|
||||
return NewMultiListener(listeners...), nil
|
||||
}
|
||||
|
||||
func NewMultiAddrTLSListener(network, address string, getCertificateFunc func(*tls.ClientHelloInfo) (*tls.Certificate, error), opts ListenerOptions) (net.Listener, error) {
|
||||
config := &tls.Config{
|
||||
MinVersion: tls.VersionTLS12,
|
||||
GetCertificate: getCertificateFunc,
|
||||
}
|
||||
|
||||
if IsUnixSocketPath(address) {
|
||||
if !isAbstractSocket(address) {
|
||||
if err := removeStaleSocket(address); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
ln, err := net.Listen("unix", address)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to bind unix TLS socket listener %s: %w", address, err)
|
||||
}
|
||||
if opts.SocketPerm != 0 && !isAbstractSocket(address) {
|
||||
if err := os.Chmod(address, opts.SocketPerm); err != nil {
|
||||
ln.Close()
|
||||
return nil, fmt.Errorf("failed to set permissions on socket %s: %w", address, err)
|
||||
}
|
||||
}
|
||||
return NewMultiListener(tls.NewListener(ln, config)), nil
|
||||
}
|
||||
|
||||
addrs, err := resolveHostnameAddrs(address)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
listeners := make([]net.Listener, 0, len(addrs))
|
||||
for _, addr := range addrs {
|
||||
ln, err := net.Listen(network, addr)
|
||||
if err != nil {
|
||||
for _, l := range listeners {
|
||||
l.Close()
|
||||
}
|
||||
return nil, fmt.Errorf("failed to bind TLS listener %s: %w", addr, err)
|
||||
}
|
||||
listeners = append(listeners, tls.NewListener(ln, config))
|
||||
}
|
||||
|
||||
return NewMultiListener(listeners...), nil
|
||||
}
|
||||
@@ -0,0 +1,231 @@
|
||||
// Copyright 2026 Versity Software
|
||||
// This file is licensed under the Apache License, Version 2.0
|
||||
// (the "License"); you may not use this file except in compliance
|
||||
// with the License. You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package sigv4auth
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
const (
|
||||
AlgorithmHMACSHA256 = "AWS4-HMAC-SHA256"
|
||||
Terminal = "aws4_request"
|
||||
ServiceS3 = "s3"
|
||||
ServiceIAM = "iam"
|
||||
|
||||
ISO8601Format = "20060102T150405Z"
|
||||
YYYYMMDD = "20060102"
|
||||
)
|
||||
|
||||
type ParseErrorKind string
|
||||
|
||||
const (
|
||||
ErrInvalidAuthorizationHeader ParseErrorKind = "invalid_authorization_header"
|
||||
ErrUnsupportedAuthorizationVersion ParseErrorKind = "unsupported_authorization_version"
|
||||
ErrInvalidAuthorizationType ParseErrorKind = "invalid_authorization_type"
|
||||
ErrMissingComponents ParseErrorKind = "missing_components"
|
||||
ErrMissingCredential ParseErrorKind = "missing_credential"
|
||||
ErrMissingSignedHeaders ParseErrorKind = "missing_signed_headers"
|
||||
ErrMissingSignature ParseErrorKind = "missing_signature"
|
||||
ErrMalformedComponent ParseErrorKind = "malformed_component"
|
||||
ErrMalformedCredential ParseErrorKind = "malformed_credential"
|
||||
ErrIncorrectService ParseErrorKind = "incorrect_service"
|
||||
ErrIncorrectTerminal ParseErrorKind = "incorrect_terminal"
|
||||
ErrInvalidDateFormat ParseErrorKind = "invalid_date_format"
|
||||
)
|
||||
|
||||
type ParseError struct {
|
||||
Kind ParseErrorKind
|
||||
Input string
|
||||
Value string
|
||||
Expected string
|
||||
Actual string
|
||||
}
|
||||
|
||||
func (e *ParseError) Error() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
switch e.Kind {
|
||||
case ErrIncorrectService, ErrIncorrectTerminal:
|
||||
return fmt.Sprintf("sigv4 %s: expected %q, got %q", e.Kind, e.Expected, e.Actual)
|
||||
case ErrInvalidAuthorizationType, ErrMalformedComponent, ErrInvalidDateFormat:
|
||||
return fmt.Sprintf("sigv4 %s: %q", e.Kind, e.Value)
|
||||
default:
|
||||
return string(e.Kind)
|
||||
}
|
||||
}
|
||||
|
||||
// AuthData is the parsed authorization data from an AWS SigV4 Authorization header.
|
||||
type AuthData struct {
|
||||
Algorithm string
|
||||
Access string
|
||||
Region string
|
||||
Service string
|
||||
SignedHeaders string
|
||||
Signature string
|
||||
Date string
|
||||
}
|
||||
|
||||
type CredentialsScope struct {
|
||||
Access string
|
||||
Date string
|
||||
Region string
|
||||
Service string
|
||||
}
|
||||
|
||||
// HexBytes returns the hex byte representation used by AWS-style diagnostic
|
||||
// signature mismatch errors.
|
||||
func HexBytes(s string) string {
|
||||
b := []byte(s)
|
||||
|
||||
parts := make([]string, len(b))
|
||||
for i, v := range b {
|
||||
parts[i] = fmt.Sprintf("%02x", v)
|
||||
}
|
||||
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
func PayloadSHA256Hex(payload []byte) string {
|
||||
hashedPayload := sha256.Sum256(payload)
|
||||
return hex.EncodeToString(hashedPayload[:])
|
||||
}
|
||||
|
||||
// ParseAuthorization parses and validates an AWS SigV4 Authorization header.
|
||||
// The credential scope service must match expectedService.
|
||||
func ParseAuthorization(authorization, expectedService string) (AuthData, error) {
|
||||
a := AuthData{}
|
||||
|
||||
authParts := strings.SplitN(authorization, " ", 2)
|
||||
for i, el := range authParts {
|
||||
if strings.Contains(el, " ") {
|
||||
authParts[i] = removeSpace(el)
|
||||
}
|
||||
}
|
||||
|
||||
if len(authParts) < 2 {
|
||||
return a, &ParseError{Kind: ErrInvalidAuthorizationHeader, Input: authorization}
|
||||
}
|
||||
|
||||
algo := authParts[0]
|
||||
if algo == "AWS" {
|
||||
return a, &ParseError{Kind: ErrUnsupportedAuthorizationVersion, Value: algo}
|
||||
}
|
||||
if algo != AlgorithmHMACSHA256 {
|
||||
return a, &ParseError{Kind: ErrInvalidAuthorizationType, Value: algo}
|
||||
}
|
||||
|
||||
kvPairs := strings.Split(authParts[1], ",")
|
||||
if len(kvPairs) != 3 {
|
||||
return a, &ParseError{Kind: ErrMissingComponents, Input: authorization}
|
||||
}
|
||||
|
||||
var access, region, service, signedHeaders, signature, date string
|
||||
for i, kv := range kvPairs {
|
||||
keyValue := strings.Split(kv, "=")
|
||||
if len(keyValue) != 2 {
|
||||
return a, &ParseError{Kind: ErrMalformedComponent, Value: kv}
|
||||
}
|
||||
key, value := keyValue[0], keyValue[1]
|
||||
switch i {
|
||||
case 0:
|
||||
if key != "Credential" {
|
||||
return a, &ParseError{Kind: ErrMissingCredential}
|
||||
}
|
||||
case 1:
|
||||
if key != "SignedHeaders" {
|
||||
return a, &ParseError{Kind: ErrMissingSignedHeaders}
|
||||
}
|
||||
case 2:
|
||||
if key != "Signature" {
|
||||
return a, &ParseError{Kind: ErrMissingSignature}
|
||||
}
|
||||
}
|
||||
|
||||
switch key {
|
||||
case "Credential":
|
||||
creds, err := ParseCredentials(value, expectedService)
|
||||
if err != nil {
|
||||
return a, err
|
||||
}
|
||||
access = creds.Access
|
||||
date = creds.Date
|
||||
region = creds.Region
|
||||
service = creds.Service
|
||||
case "SignedHeaders":
|
||||
signedHeaders = value
|
||||
case "Signature":
|
||||
signature = value
|
||||
}
|
||||
}
|
||||
|
||||
return AuthData{
|
||||
Algorithm: algo,
|
||||
Access: access,
|
||||
Region: region,
|
||||
Service: service,
|
||||
SignedHeaders: signedHeaders,
|
||||
Signature: signature,
|
||||
Date: date,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func ParseCredentials(input, expectedService string) (*CredentialsScope, error) {
|
||||
creds := strings.Split(input, "/")
|
||||
if len(creds) != 5 {
|
||||
return nil, &ParseError{Kind: ErrMalformedCredential, Input: input}
|
||||
}
|
||||
if creds[3] != expectedService {
|
||||
return nil, &ParseError{
|
||||
Kind: ErrIncorrectService,
|
||||
Input: input,
|
||||
Expected: expectedService,
|
||||
Actual: creds[3],
|
||||
}
|
||||
}
|
||||
if creds[4] != Terminal {
|
||||
return nil, &ParseError{
|
||||
Kind: ErrIncorrectTerminal,
|
||||
Input: input,
|
||||
Expected: Terminal,
|
||||
Actual: creds[4],
|
||||
}
|
||||
}
|
||||
if _, err := time.Parse(YYYYMMDD, creds[1]); err != nil {
|
||||
return nil, &ParseError{Kind: ErrInvalidDateFormat, Input: input, Value: creds[1]}
|
||||
}
|
||||
|
||||
return &CredentialsScope{
|
||||
Access: creds[0],
|
||||
Date: creds[1],
|
||||
Region: creds[2],
|
||||
Service: creds[3],
|
||||
}, nil
|
||||
}
|
||||
|
||||
func removeSpace(str string) string {
|
||||
var b strings.Builder
|
||||
b.Grow(len(str))
|
||||
for _, ch := range str {
|
||||
if !unicode.IsSpace(ch) {
|
||||
b.WriteRune(ch)
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
@@ -0,0 +1,406 @@
|
||||
// Copyright 2026 Versity Software
|
||||
// This file is licensed under the Apache License, Version 2.0
|
||||
// (the "License"); you may not use this file except in compliance
|
||||
// with the License. You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package sigv4auth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/aws/aws-sdk-go-v2/aws"
|
||||
"github.com/aws/smithy-go/logging"
|
||||
"github.com/gofiber/fiber/v3"
|
||||
"github.com/versity/versitygw/aws/signer/v4"
|
||||
"github.com/versity/versitygw/debuglogger"
|
||||
)
|
||||
|
||||
const (
|
||||
AlgorithmECDSAP256SHA256 = "AWS4-ECDSA-P256-SHA256"
|
||||
|
||||
QueryAlgorithm = "X-Amz-Algorithm"
|
||||
QueryCredential = "X-Amz-Credential"
|
||||
QueryDate = "X-Amz-Date"
|
||||
QueryExpires = "X-Amz-Expires"
|
||||
QuerySignedHeaders = "X-Amz-SignedHeaders"
|
||||
QuerySignature = "X-Amz-Signature"
|
||||
QuerySecurityToken = "X-Amz-Security-Token"
|
||||
|
||||
maxQueryExpirationSeconds = 604800
|
||||
)
|
||||
|
||||
type QueryErrorKind string
|
||||
|
||||
const (
|
||||
ErrQueryMissingRequiredParams QueryErrorKind = "missing_required_query_parameters"
|
||||
ErrQueryUnsupportedAlgorithm QueryErrorKind = "unsupported_query_algorithm"
|
||||
ErrQueryUnsupportedECDSA QueryErrorKind = "unsupported_query_ecdsa"
|
||||
ErrQueryInvalidDateFormat QueryErrorKind = "invalid_query_date_format"
|
||||
ErrQueryDateMismatch QueryErrorKind = "query_date_mismatch"
|
||||
ErrQueryIncorrectRegion QueryErrorKind = "query_incorrect_region"
|
||||
ErrQueryExpiresNumber QueryErrorKind = "query_expires_number"
|
||||
ErrQueryExpiresNegative QueryErrorKind = "query_expires_negative"
|
||||
ErrQueryExpiresTooLarge QueryErrorKind = "query_expires_too_large"
|
||||
ErrQueryExpired QueryErrorKind = "query_expired"
|
||||
ErrQuerySecurityToken QueryErrorKind = "query_security_token"
|
||||
)
|
||||
|
||||
type QueryError struct {
|
||||
Kind QueryErrorKind
|
||||
Value string
|
||||
Expected string
|
||||
Actual string
|
||||
Expires int
|
||||
ExpiresAt time.Time
|
||||
ServerTime time.Time
|
||||
}
|
||||
|
||||
func (e *QueryError) Error() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
switch e.Kind {
|
||||
case ErrQueryIncorrectRegion:
|
||||
return fmt.Sprintf("sigv4 query %s: expected %q, got %q", e.Kind, e.Expected, e.Actual)
|
||||
case ErrQueryDateMismatch:
|
||||
return fmt.Sprintf("sigv4 query %s: expected %q, got %q", e.Kind, e.Expected, e.Actual)
|
||||
case ErrQueryExpired:
|
||||
return fmt.Sprintf("sigv4 query %s: expired at %s", e.Kind, e.ExpiresAt.Format(time.RFC3339))
|
||||
case ErrQueryUnsupportedAlgorithm, ErrQueryUnsupportedECDSA, ErrQueryExpiresNumber:
|
||||
return fmt.Sprintf("sigv4 query %s: %q", e.Kind, e.Value)
|
||||
default:
|
||||
return string(e.Kind)
|
||||
}
|
||||
}
|
||||
|
||||
type QueryAuthOptions struct {
|
||||
Service string
|
||||
Region string
|
||||
// RequireExpiration enables the X-Amz-Expires validation required by S3
|
||||
// presigned URLs. Other SigV4 query-auth services, including IAM, leave it
|
||||
// disabled.
|
||||
RequireExpiration bool
|
||||
Now func() time.Time
|
||||
}
|
||||
|
||||
type QueryAuthDetails struct {
|
||||
SigningTime time.Time
|
||||
Expires int
|
||||
ExpiresAt time.Time
|
||||
ServerTime time.Time
|
||||
}
|
||||
|
||||
// ParseQueryAuthorization parses and validates AWS SigV4 query-string
|
||||
// authentication parameters. The credential scope service must match
|
||||
// opts.Service. If opts.Region is set, the credential scope region must match
|
||||
// it as well.
|
||||
func ParseQueryAuthorization(ctx fiber.Ctx, opts QueryAuthOptions) (AuthData, QueryAuthDetails, error) {
|
||||
a := AuthData{}
|
||||
details := QueryAuthDetails{}
|
||||
|
||||
if err := ValidateQueryAlgorithm(ctx.Query(QueryAlgorithm)); err != nil {
|
||||
return a, details, err
|
||||
}
|
||||
|
||||
credsQuery := ctx.Query(QueryCredential)
|
||||
if credsQuery == "" {
|
||||
return a, details, missingQueryParameterError(QueryCredential)
|
||||
}
|
||||
|
||||
creds, err := ParseCredentials(credsQuery, opts.Service)
|
||||
if err != nil {
|
||||
return a, details, err
|
||||
}
|
||||
|
||||
if opts.Region != "" && creds.Region != opts.Region {
|
||||
return a, details, &QueryError{
|
||||
Kind: ErrQueryIncorrectRegion,
|
||||
Expected: opts.Region,
|
||||
Actual: creds.Region,
|
||||
}
|
||||
}
|
||||
|
||||
date := ctx.Query(QueryDate)
|
||||
if date == "" {
|
||||
return a, details, missingQueryParameterError(QueryDate)
|
||||
}
|
||||
|
||||
tdate, err := time.Parse(ISO8601Format, date)
|
||||
if err != nil {
|
||||
return a, details, &QueryError{Kind: ErrQueryInvalidDateFormat, Value: date}
|
||||
}
|
||||
|
||||
if date[:8] != creds.Date {
|
||||
return a, details, &QueryError{
|
||||
Kind: ErrQueryDateMismatch,
|
||||
Expected: creds.Date,
|
||||
Actual: date[:8],
|
||||
}
|
||||
}
|
||||
|
||||
signature := ctx.Query(QuerySignature)
|
||||
if signature == "" {
|
||||
return a, details, missingQueryParameterError(QuerySignature)
|
||||
}
|
||||
|
||||
signedHdrs := ctx.Query(QuerySignedHeaders)
|
||||
if signedHdrs == "" {
|
||||
return a, details, missingQueryParameterError(QuerySignedHeaders)
|
||||
}
|
||||
|
||||
expiration := QueryExpiration{}
|
||||
if opts.RequireExpiration {
|
||||
now := time.Now().UTC()
|
||||
if opts.Now != nil {
|
||||
now = opts.Now().UTC()
|
||||
}
|
||||
expiration, err = ValidateQueryExpiration(ctx.Query(QueryExpires), tdate, now)
|
||||
if err != nil {
|
||||
return a, details, err
|
||||
}
|
||||
}
|
||||
|
||||
a = AuthData{
|
||||
Algorithm: ctx.Query(QueryAlgorithm),
|
||||
Access: creds.Access,
|
||||
Region: creds.Region,
|
||||
Service: creds.Service,
|
||||
SignedHeaders: signedHdrs,
|
||||
Signature: signature,
|
||||
Date: date,
|
||||
}
|
||||
details = QueryAuthDetails{
|
||||
SigningTime: tdate,
|
||||
Expires: expiration.Expires,
|
||||
ExpiresAt: expiration.ExpiresAt,
|
||||
ServerTime: expiration.ServerTime,
|
||||
}
|
||||
|
||||
return a, details, nil
|
||||
}
|
||||
|
||||
func ValidateQueryAlgorithm(algo string) error {
|
||||
switch algo {
|
||||
case "":
|
||||
return missingQueryParameterError(QueryAlgorithm)
|
||||
case AlgorithmHMACSHA256:
|
||||
return nil
|
||||
case AlgorithmECDSAP256SHA256:
|
||||
return &QueryError{Kind: ErrQueryUnsupportedECDSA, Value: algo}
|
||||
default:
|
||||
return &QueryError{Kind: ErrQueryUnsupportedAlgorithm, Value: algo}
|
||||
}
|
||||
}
|
||||
|
||||
type QueryExpiration struct {
|
||||
Expires int
|
||||
ExpiresAt time.Time
|
||||
ServerTime time.Time
|
||||
}
|
||||
|
||||
func ValidateQueryExpiration(str string, date, now time.Time) (QueryExpiration, error) {
|
||||
if str == "" {
|
||||
return QueryExpiration{}, missingQueryParameterError(QueryExpires)
|
||||
}
|
||||
|
||||
exp, err := strconv.Atoi(str)
|
||||
if err != nil {
|
||||
return QueryExpiration{}, &QueryError{Kind: ErrQueryExpiresNumber, Value: str}
|
||||
}
|
||||
|
||||
if exp < 0 {
|
||||
return QueryExpiration{}, &QueryError{Kind: ErrQueryExpiresNegative, Value: str}
|
||||
}
|
||||
|
||||
if exp > maxQueryExpirationSeconds {
|
||||
return QueryExpiration{}, &QueryError{Kind: ErrQueryExpiresTooLarge, Value: str}
|
||||
}
|
||||
|
||||
now = now.UTC()
|
||||
expiresAt := date.Add(time.Duration(exp) * time.Second)
|
||||
expiration := QueryExpiration{
|
||||
Expires: exp,
|
||||
ExpiresAt: expiresAt,
|
||||
ServerTime: now,
|
||||
}
|
||||
|
||||
if expiresAt.Before(now) {
|
||||
return expiration, &QueryError{
|
||||
Kind: ErrQueryExpired,
|
||||
Expires: exp,
|
||||
ExpiresAt: expiresAt,
|
||||
ServerTime: now,
|
||||
}
|
||||
}
|
||||
|
||||
return expiration, nil
|
||||
}
|
||||
|
||||
func missingQueryParameterError(parameter string) *QueryError {
|
||||
return &QueryError{Kind: ErrQueryMissingRequiredParams, Value: parameter}
|
||||
}
|
||||
|
||||
// CheckQuerySignature rebuilds a SigV4 query-auth request and compares the
|
||||
// generated query signature to the signature presented by the client.
|
||||
func CheckQuerySignature(ctx fiber.Ctx, auth AuthData, secret, payloadHash string, tdate time.Time, contentLen int64, opts CheckOptions) (*CheckResult, error) {
|
||||
service := opts.Service
|
||||
if service == "" {
|
||||
service = auth.Service
|
||||
}
|
||||
signedHdrs := strings.Split(auth.SignedHeaders, ";")
|
||||
|
||||
req, err := createPresignedHTTPRequestFromCtx(ctx, signedHdrs, contentLen, opts.RequiredSignedHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
signer := v4.NewSigner()
|
||||
uri, _, signMeta, err := signer.PresignHTTP(ctx.RequestCtx(),
|
||||
aws.Credentials{
|
||||
AccessKeyID: auth.Access,
|
||||
SecretAccessKey: secret,
|
||||
},
|
||||
req, payloadHash, service, auth.Region, tdate, signedHdrs,
|
||||
func(options *v4.SignerOptions) {
|
||||
options.DisableURIPathEscaping = opts.DisableURIPathEscaping
|
||||
if debuglogger.IsDebugEnabled() {
|
||||
options.LogSigning = true
|
||||
options.Logger = logging.NewStandardLogger(os.Stderr)
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("presign generated http request: %w", err)
|
||||
}
|
||||
|
||||
urlParts, err := url.Parse(uri)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse presigned url: %w", err)
|
||||
}
|
||||
|
||||
signature := urlParts.Query().Get(QuerySignature)
|
||||
if signature != auth.Signature {
|
||||
return nil, &SignatureMismatchError{
|
||||
AccessKeyID: auth.Access,
|
||||
StringToSign: signMeta.StringToSign,
|
||||
SignatureProvided: auth.Signature,
|
||||
StringToSignBytes: HexBytes(signMeta.StringToSign),
|
||||
CanonicalRequest: signMeta.CanonicalString,
|
||||
CanonicalRequestBytes: HexBytes(signMeta.CanonicalString),
|
||||
}
|
||||
}
|
||||
|
||||
return &CheckResult{
|
||||
CanonicalString: signMeta.CanonicalString,
|
||||
StringToSign: signMeta.StringToSign,
|
||||
}, nil
|
||||
}
|
||||
|
||||
var generatedQueryAuthParams = map[string]struct{}{
|
||||
QueryAlgorithm: {},
|
||||
QueryCredential: {},
|
||||
QueryDate: {},
|
||||
QuerySignedHeaders: {},
|
||||
QuerySignature: {},
|
||||
}
|
||||
|
||||
func createPresignedHTTPRequestFromCtx(ctx fiber.Ctx, signedHdrs []string, contentLength int64, requiredSignedHdrs []string) (*http.Request, error) {
|
||||
req := ctx.Request()
|
||||
if err := validateRequiredSignedHeaders(signedHdrs, requiredSignedHdrs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
uri, _, _ := strings.Cut(ctx.OriginalURL(), "?")
|
||||
query := strings.Builder{}
|
||||
|
||||
for key, value := range ctx.Request().URI().QueryArgs().All() {
|
||||
keyStr := string(key)
|
||||
if _, ok := generatedQueryAuthParams[keyStr]; ok {
|
||||
continue
|
||||
}
|
||||
|
||||
if query.Len() > 0 {
|
||||
query.WriteByte('&')
|
||||
}
|
||||
query.WriteString(url.QueryEscape(keyStr))
|
||||
query.WriteByte('=')
|
||||
query.WriteString(url.QueryEscape(string(value)))
|
||||
}
|
||||
|
||||
if query.Len() > 0 {
|
||||
uri += "?" + query.String()
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequest(string(req.Header.Method()), uri, nil)
|
||||
if err != nil {
|
||||
return nil, errors.New("error in creating an http request")
|
||||
}
|
||||
if err := addRequestHeadersFromCtx(ctx, httpReq, signedHdrs, requiredSignedHdrs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !includeHeader("Content-Length", signedHdrs) {
|
||||
httpReq.ContentLength = 0
|
||||
} else {
|
||||
httpReq.ContentLength = contentLength
|
||||
}
|
||||
|
||||
httpReq.Host = string(req.Header.Host())
|
||||
|
||||
return httpReq, nil
|
||||
}
|
||||
|
||||
// IsQueryAuth determines if a request uses SigV4 query-string auth.
|
||||
func IsQueryAuth(ctx fiber.Ctx) bool {
|
||||
algo := ctx.Query(QueryAlgorithm)
|
||||
creds := ctx.Query(QueryCredential)
|
||||
date := ctx.Query(QueryDate)
|
||||
signature := ctx.Query(QuerySignature)
|
||||
signedHeaders := ctx.Query(QuerySignedHeaders)
|
||||
|
||||
return !allEmpty(algo, creds, date, signature, signedHeaders)
|
||||
}
|
||||
|
||||
// IsQueryAuthV2 determines if a request is query-string signed with the legacy
|
||||
// AWS Signature Version 2 signer.
|
||||
func IsQueryAuthV2(ctx fiber.Ctx) bool {
|
||||
expires := ctx.Query("Expires")
|
||||
access := ctx.Query("AWSAccessKeyId")
|
||||
signature := ctx.Query("Signature")
|
||||
|
||||
return anyNonEmpty(expires, access, signature)
|
||||
}
|
||||
|
||||
func allEmpty(args ...string) bool {
|
||||
for _, a := range args {
|
||||
if a != "" {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func anyNonEmpty(args ...string) bool {
|
||||
for _, a := range args {
|
||||
if a != "" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,213 @@
|
||||
// Copyright 2026 Versity Software
|
||||
// This file is licensed under the Apache License, Version 2.0
|
||||
// (the "License"); you may not use this file except in compliance
|
||||
// with the License. You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package sigv4auth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/aws/aws-sdk-go-v2/aws"
|
||||
"github.com/aws/smithy-go/logging"
|
||||
"github.com/gofiber/fiber/v3"
|
||||
"github.com/versity/versitygw/aws/signer/v4"
|
||||
"github.com/versity/versitygw/debuglogger"
|
||||
)
|
||||
|
||||
type CheckOptions struct {
|
||||
Service string
|
||||
DisableURIPathEscaping bool
|
||||
// RequiredSignedHeaders overrides the default AWS signed-header policy.
|
||||
// A nil slice requires every applicable X-Amz-* header to be signed.
|
||||
RequiredSignedHeaders []string
|
||||
}
|
||||
|
||||
type CheckResult struct {
|
||||
CanonicalString string
|
||||
StringToSign string
|
||||
}
|
||||
|
||||
type HeadersNotSignedError struct {
|
||||
Headers []string
|
||||
}
|
||||
|
||||
func (e *HeadersNotSignedError) Error() string {
|
||||
return fmt.Sprintf("headers not signed: %s", strings.Join(e.Headers, ", "))
|
||||
}
|
||||
|
||||
type SignatureMismatchError struct {
|
||||
AccessKeyID string
|
||||
StringToSign string
|
||||
SignatureProvided string
|
||||
StringToSignBytes string
|
||||
CanonicalRequest string
|
||||
CanonicalRequestBytes string
|
||||
}
|
||||
|
||||
func (e *SignatureMismatchError) Error() string {
|
||||
return "signature does not match"
|
||||
}
|
||||
|
||||
// CheckSignature rebuilds the canonical request with the supplied service,
|
||||
// region, payload hash, signing time, and signed headers, then compares the
|
||||
// generated signature to the signature presented by the client.
|
||||
func CheckSignature(ctx fiber.Ctx, auth AuthData, secret, payloadHash string, tdate time.Time, contentLen int64, opts CheckOptions) (*CheckResult, error) {
|
||||
service := opts.Service
|
||||
if service == "" {
|
||||
service = auth.Service
|
||||
}
|
||||
signedHdrs := strings.Split(auth.SignedHeaders, ";")
|
||||
|
||||
req, err := createHTTPRequestFromCtx(ctx, signedHdrs, contentLen, opts.RequiredSignedHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
signer := v4.NewSigner()
|
||||
|
||||
signMeta, err := signer.SignHTTP(req.Context(),
|
||||
aws.Credentials{
|
||||
AccessKeyID: auth.Access,
|
||||
SecretAccessKey: secret,
|
||||
},
|
||||
req, payloadHash, service, auth.Region, tdate, signedHdrs,
|
||||
func(options *v4.SignerOptions) {
|
||||
options.DisableURIPathEscaping = opts.DisableURIPathEscaping
|
||||
if debuglogger.IsDebugEnabled() {
|
||||
options.LogSigning = true
|
||||
options.Logger = logging.NewStandardLogger(os.Stderr)
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("sign generated http request: %w", err)
|
||||
}
|
||||
|
||||
genAuth, err := ParseAuthorization(req.Header.Get("Authorization"), service)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if auth.Signature != genAuth.Signature {
|
||||
return nil, &SignatureMismatchError{
|
||||
AccessKeyID: auth.Access,
|
||||
StringToSign: signMeta.StringToSign,
|
||||
SignatureProvided: auth.Signature,
|
||||
StringToSignBytes: HexBytes(signMeta.StringToSign),
|
||||
CanonicalRequest: signMeta.CanonicalString,
|
||||
CanonicalRequestBytes: HexBytes(signMeta.CanonicalString),
|
||||
}
|
||||
}
|
||||
|
||||
return &CheckResult{
|
||||
CanonicalString: signMeta.CanonicalString,
|
||||
StringToSign: signMeta.StringToSign,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func CreateHTTPRequestFromCtx(ctx fiber.Ctx, signedHdrs []string, contentLength int64) (*http.Request, error) {
|
||||
return createHTTPRequestFromCtx(ctx, signedHdrs, contentLength, nil)
|
||||
}
|
||||
|
||||
func createHTTPRequestFromCtx(ctx fiber.Ctx, signedHdrs []string, contentLength int64, requiredSignedHdrs []string) (*http.Request, error) {
|
||||
req := ctx.Request()
|
||||
if err := validateRequiredSignedHeaders(signedHdrs, requiredSignedHdrs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequest(string(req.Header.Method()), ctx.OriginalURL(), nil)
|
||||
if err != nil {
|
||||
return nil, errors.New("error in creating an http request")
|
||||
}
|
||||
|
||||
if err := addRequestHeadersFromCtx(ctx, httpReq, signedHdrs, requiredSignedHdrs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, header := range signedHdrs {
|
||||
if httpReq.Header.Get(header) == "" {
|
||||
httpReq.Header.Set(header, "")
|
||||
}
|
||||
}
|
||||
|
||||
if !includeHeader("Content-Length", signedHdrs) {
|
||||
httpReq.ContentLength = 0
|
||||
} else {
|
||||
httpReq.ContentLength = contentLength
|
||||
}
|
||||
|
||||
httpReq.Host = string(req.Header.Host())
|
||||
|
||||
return httpReq, nil
|
||||
}
|
||||
|
||||
func AddRequestHeadersFromCtx(ctx fiber.Ctx, httpReq *http.Request, signedHdrs []string) error {
|
||||
return addRequestHeadersFromCtx(ctx, httpReq, signedHdrs, nil)
|
||||
}
|
||||
|
||||
func addRequestHeadersFromCtx(ctx fiber.Ctx, httpReq *http.Request, signedHdrs, requiredSignedHdrs []string) error {
|
||||
headersNotSigned := []string{}
|
||||
for key, value := range ctx.Request().Header.All() {
|
||||
keyStr := string(key)
|
||||
if includeHeader(keyStr, signedHdrs) || v4.IsIgnoredHeader(keyStr) {
|
||||
httpReq.Header.Add(keyStr, string(value))
|
||||
continue
|
||||
}
|
||||
if isRequiredSignedHeader(keyStr, requiredSignedHdrs) {
|
||||
headersNotSigned = append(headersNotSigned, strings.ToLower(keyStr))
|
||||
}
|
||||
}
|
||||
|
||||
if len(headersNotSigned) != 0 {
|
||||
debuglogger.Logf("headers present in request but not included in SignedHeaders: %q", strings.Join(headersNotSigned, ", "))
|
||||
return &HeadersNotSignedError{Headers: headersNotSigned}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateRequiredSignedHeaders(signedHdrs, requiredSignedHdrs []string) error {
|
||||
if requiredSignedHdrs == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
headersNotSigned := []string{}
|
||||
for _, header := range requiredSignedHdrs {
|
||||
if !includeHeader(header, signedHdrs) {
|
||||
headersNotSigned = append(headersNotSigned, strings.ToLower(header))
|
||||
}
|
||||
}
|
||||
if len(headersNotSigned) != 0 {
|
||||
return &HeadersNotSignedError{Headers: headersNotSigned}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func isRequiredSignedHeader(header string, requiredSignedHdrs []string) bool {
|
||||
if requiredSignedHdrs == nil {
|
||||
return v4.IsRequiredSignedHeader(header)
|
||||
}
|
||||
|
||||
return includeHeader(header, requiredSignedHdrs)
|
||||
}
|
||||
|
||||
func includeHeader(hdr string, signedHdrs []string) bool {
|
||||
return slices.ContainsFunc(signedHdrs, func(shdr string) bool {
|
||||
return strings.EqualFold(hdr, shdr)
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user