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:
niksis02
2026-08-25 01:03:12 +04:00
parent 762e244b43
commit e012a6fd01
60 changed files with 9108 additions and 629 deletions
+56
View File
@@ -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))
}
+194
View File
@@ -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
}
+71
View File
@@ -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)
}
}
+44
View File
@@ -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
}
+316
View File
@@ -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
}
+231
View File
@@ -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()
}
+406
View File
@@ -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
}
+213
View File
@@ -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)
})
}