mirror of
https://github.com/seaweedfs/seaweedfs.git
synced 2026-10-06 22:55:51 +00:00
Compare commits
30
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9860ae582b | ||
|
|
c6f454fb9b | ||
|
|
abbd0207ba | ||
|
|
d49c2a7364 | ||
|
|
995dfc4d5d | ||
|
|
8fad85aed7 | ||
|
|
2e98902f29 | ||
|
|
d3cea714d0 | ||
|
|
91087c0737 | ||
|
|
d2d21cd26b | ||
|
|
0503311ded | ||
|
|
bb23939b36 | ||
|
|
059bee683f | ||
|
|
b8236a10d1 | ||
|
|
a4b896a224 | ||
|
|
7c59b639c9 | ||
|
|
772ad67f6b | ||
|
|
3a5016bcd7 | ||
|
|
0d8b024911 | ||
|
|
0f5e6b1f34 | ||
|
|
98714b9f70 | ||
|
|
9552e80b58 | ||
|
|
a974190cb1 | ||
|
|
e7fc243ee1 | ||
|
|
ab4e52ae2f | ||
|
|
888c32cbde | ||
|
|
efbed39e25 | ||
|
|
e93f4e3f39 | ||
|
|
647b46bd8a | ||
|
|
24805ff478 |
@@ -0,0 +1 @@
|
||||
{"sessionId":"d6574c47-eafc-4a94-9dce-f9ffea22b53c","pid":10111,"acquiredAt":1775248373916}
|
||||
@@ -39,7 +39,80 @@ concurrency:
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
|
||||
# ── Pre-build Rust volume server binaries natively ──────────────────
|
||||
# Cross-compiles for amd64 and arm64 without QEMU, turning a 5-hour
|
||||
# emulated cargo build into ~15 minutes of native compilation.
|
||||
build-rust-binaries:
|
||||
runs-on: ubuntu-22.04
|
||||
strategy:
|
||||
matrix:
|
||||
include:
|
||||
- target: x86_64-unknown-linux-musl
|
||||
arch: amd64
|
||||
- target: aarch64-unknown-linux-musl
|
||||
arch: arm64
|
||||
cross: true
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v6
|
||||
|
||||
- name: Install protobuf compiler
|
||||
run: sudo apt-get update && sudo apt-get install -y protobuf-compiler
|
||||
|
||||
- name: Install Rust toolchain
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
targets: ${{ matrix.target }}
|
||||
|
||||
- name: Install musl tools (amd64)
|
||||
if: ${{ !matrix.cross }}
|
||||
run: sudo apt-get install -y musl-tools
|
||||
|
||||
- name: Install cross-compilation tools (arm64)
|
||||
if: matrix.cross
|
||||
run: |
|
||||
sudo apt-get install -y gcc-aarch64-linux-gnu
|
||||
echo "CARGO_TARGET_AARCH64_UNKNOWN_LINUX_MUSL_LINKER=aarch64-linux-gnu-gcc" >> "$GITHUB_ENV"
|
||||
|
||||
- name: Cache cargo registry and target
|
||||
uses: actions/cache@v5
|
||||
with:
|
||||
path: |
|
||||
~/.cargo/registry
|
||||
~/.cargo/git
|
||||
seaweed-volume/target
|
||||
key: rust-docker-${{ matrix.target }}-${{ hashFiles('seaweed-volume/Cargo.lock') }}
|
||||
restore-keys: |
|
||||
rust-docker-${{ matrix.target }}-
|
||||
|
||||
- name: Build large-disk variant
|
||||
env:
|
||||
SEAWEEDFS_COMMIT: ${{ github.sha }}
|
||||
run: |
|
||||
cd seaweed-volume
|
||||
cargo build --release --target ${{ matrix.target }}
|
||||
cp target/${{ matrix.target }}/release/weed-volume ../weed-volume-large-disk-${{ matrix.arch }}
|
||||
|
||||
- name: Build normal variant
|
||||
env:
|
||||
SEAWEEDFS_COMMIT: ${{ github.sha }}
|
||||
run: |
|
||||
cd seaweed-volume
|
||||
cargo build --release --target ${{ matrix.target }} --no-default-features
|
||||
cp target/${{ matrix.target }}/release/weed-volume ../weed-volume-normal-${{ matrix.arch }}
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: rust-volume-${{ matrix.arch }}
|
||||
path: |
|
||||
weed-volume-large-disk-${{ matrix.arch }}
|
||||
weed-volume-normal-${{ matrix.arch }}
|
||||
|
||||
# ── Build Docker containers ─────────────────────────────────────────
|
||||
build:
|
||||
needs: [build-rust-binaries]
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
# Build sequentially to avoid rate limits
|
||||
@@ -52,20 +125,23 @@ jobs:
|
||||
dockerfile: ./docker/Dockerfile.go_build
|
||||
build_args: ""
|
||||
tag_suffix: ""
|
||||
|
||||
# Large disk - multi-arch
|
||||
rust_variant: normal
|
||||
|
||||
# Large disk - multi-arch
|
||||
- variant: large_disk
|
||||
platforms: linux/amd64,linux/arm64,linux/arm/v7,linux/386
|
||||
dockerfile: ./docker/Dockerfile.go_build
|
||||
build_args: TAGS=5BytesOffset
|
||||
tag_suffix: _large_disk
|
||||
|
||||
rust_variant: large-disk
|
||||
|
||||
# Full tags - multi-arch
|
||||
- variant: full
|
||||
platforms: linux/amd64,linux/arm64
|
||||
dockerfile: ./docker/Dockerfile.go_build
|
||||
build_args: TAGS=elastic,gocdk,rclone,sqlite,tarantool,tikv,ydb
|
||||
tag_suffix: _full
|
||||
rust_variant: normal
|
||||
|
||||
# Large disk + full tags - multi-arch
|
||||
- variant: large_disk_full
|
||||
@@ -73,19 +149,42 @@ jobs:
|
||||
dockerfile: ./docker/Dockerfile.go_build
|
||||
build_args: TAGS=5BytesOffset,elastic,gocdk,rclone,sqlite,tarantool,tikv,ydb
|
||||
tag_suffix: _large_disk_full
|
||||
|
||||
rust_variant: large-disk
|
||||
|
||||
# RocksDB large disk - amd64 only
|
||||
- variant: rocksdb
|
||||
platforms: linux/amd64
|
||||
dockerfile: ./docker/Dockerfile.rocksdb_large
|
||||
build_args: ""
|
||||
tag_suffix: _large_disk_rocksdb
|
||||
rust_variant: large-disk
|
||||
|
||||
steps:
|
||||
- name: Checkout
|
||||
if: github.event_name != 'workflow_dispatch' || github.event.inputs.variant == 'all' || github.event.inputs.variant == matrix.variant
|
||||
uses: actions/checkout@v6
|
||||
|
||||
|
||||
- name: Download pre-built Rust binaries
|
||||
if: github.event_name != 'workflow_dispatch' || github.event.inputs.variant == 'all' || github.event.inputs.variant == matrix.variant
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
pattern: rust-volume-*
|
||||
merge-multiple: true
|
||||
path: ./rust-bins
|
||||
|
||||
- name: Place Rust binaries in Docker context
|
||||
if: github.event_name != 'workflow_dispatch' || github.event.inputs.variant == 'all' || github.event.inputs.variant == matrix.variant
|
||||
run: |
|
||||
mkdir -p docker/weed-volume-prebuilt
|
||||
for arch in amd64 arm64; do
|
||||
src="./rust-bins/weed-volume-${{ matrix.rust_variant }}-${arch}"
|
||||
if [ -f "$src" ]; then
|
||||
cp "$src" "docker/weed-volume-prebuilt/weed-volume-${arch}"
|
||||
echo "Placed pre-built Rust binary for ${arch}"
|
||||
fi
|
||||
done
|
||||
ls -la docker/weed-volume-prebuilt/
|
||||
|
||||
- name: Free Disk Space
|
||||
if: github.event_name != 'workflow_dispatch' || github.event.inputs.variant == 'all' || github.event.inputs.variant == matrix.variant
|
||||
run: |
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
{
|
||||
"setup": [],
|
||||
"teardown": [],
|
||||
"run": []
|
||||
}
|
||||
@@ -16,15 +16,21 @@ RUN cd /go/src/github.com/seaweedfs/seaweedfs/weed \
|
||||
&& export LDFLAGS="-X github.com/seaweedfs/seaweedfs/weed/util/version.COMMIT=$(git rev-parse --short HEAD)" \
|
||||
&& CGO_ENABLED=0 go install -tags "$TAGS" -ldflags "-extldflags -static ${LDFLAGS}"
|
||||
|
||||
# Rust volume server builder. Alpine packages avoid depending on the
|
||||
# upstream rust:alpine manifest list, which no longer includes linux/386.
|
||||
# Rust volume server: use pre-built binary from CI when available (placed in
|
||||
# weed-volume-prebuilt/ by the build-rust-binaries job), otherwise compile
|
||||
# from source. Pre-building avoids a multi-hour QEMU-emulated cargo build
|
||||
# for non-native architectures.
|
||||
FROM alpine:3.23 as rust_builder
|
||||
ARG TARGETARCH
|
||||
ARG TAGS
|
||||
COPY weed-volume-prebuilt/ /prebuilt/
|
||||
COPY --from=builder /go/src/github.com/seaweedfs/seaweedfs/seaweed-volume /build/seaweed-volume
|
||||
COPY --from=builder /go/src/github.com/seaweedfs/seaweedfs/weed /build/weed
|
||||
WORKDIR /build/seaweed-volume
|
||||
ARG TAGS
|
||||
RUN if [ "$TARGETARCH" = "amd64" ] || [ "$TARGETARCH" = "arm64" ]; then \
|
||||
RUN if [ -f "/prebuilt/weed-volume-${TARGETARCH}" ]; then \
|
||||
echo "Using pre-built Rust binary for ${TARGETARCH}" && \
|
||||
cp "/prebuilt/weed-volume-${TARGETARCH}" /weed-volume; \
|
||||
elif [ "$TARGETARCH" = "amd64" ] || [ "$TARGETARCH" = "arm64" ]; then \
|
||||
apk add --no-cache musl-dev openssl-dev protobuf-dev git rust cargo; \
|
||||
if [ "$TAGS" = "5BytesOffset" ]; then \
|
||||
cargo build --release; \
|
||||
@@ -51,7 +57,7 @@ COPY --from=builder /go/src/github.com/seaweedfs/seaweedfs/docker/entrypoint.sh
|
||||
|
||||
# Install dependencies and create non-root user
|
||||
RUN apk upgrade --no-cache zlib && \
|
||||
apk add --no-cache fuse curl su-exec && \
|
||||
apk add --no-cache fuse curl su-exec libgcc && \
|
||||
addgroup -g 1000 seaweed && \
|
||||
adduser -D -u 1000 -G seaweed seaweed
|
||||
|
||||
|
||||
@@ -371,7 +371,7 @@ require (
|
||||
github.com/geoffgarside/ber v1.2.0 // indirect
|
||||
github.com/go-chi/chi/v5 v5.2.5 // indirect
|
||||
github.com/go-darwin/apfs v0.0.0-20211011131704-f84b94dbf348 // indirect
|
||||
github.com/go-jose/go-jose/v4 v4.1.3 // indirect
|
||||
github.com/go-jose/go-jose/v4 v4.1.4 // indirect
|
||||
github.com/go-logr/logr v1.4.3 // indirect
|
||||
github.com/go-logr/stdr v1.2.2 // indirect
|
||||
github.com/go-ole/go-ole v1.3.0 // indirect
|
||||
|
||||
@@ -1076,8 +1076,8 @@ github.com/go-git/go-billy/v5 v5.6.2/go.mod h1:rcFC2rAsp/erv7CMz9GczHcuD0D32fWzH
|
||||
github.com/go-gl/glfw v0.0.0-20190409004039-e6da0acd62b1/go.mod h1:vR7hzQXu2zJy9AVAgeJqvqgH9Q5CA+iKCZ2gyEVpxRU=
|
||||
github.com/go-gl/glfw/v3.3/glfw v0.0.0-20191125211704-12ad95a8df72/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8=
|
||||
github.com/go-gl/glfw/v3.3/glfw v0.0.0-20200222043503-6f7a984d4dc4/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8=
|
||||
github.com/go-jose/go-jose/v4 v4.1.3 h1:CVLmWDhDVRa6Mi/IgCgaopNosCaHz7zrMeF9MlZRkrs=
|
||||
github.com/go-jose/go-jose/v4 v4.1.3/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
|
||||
github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA=
|
||||
github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
|
||||
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
||||
|
||||
Generated
+11
@@ -2561,6 +2561,15 @@ version = "0.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe"
|
||||
|
||||
[[package]]
|
||||
name = "openssl-src"
|
||||
version = "300.5.5+3.5.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3f1787d533e03597a7934fd0a765f0d28e94ecc5fb7789f8053b1e699a56f709"
|
||||
dependencies = [
|
||||
"cc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "openssl-sys"
|
||||
version = "0.9.111"
|
||||
@@ -2569,6 +2578,7 @@ checksum = "82cab2d520aa75e3c58898289429321eb788c3106963d0dc886ec7a5f4adc321"
|
||||
dependencies = [
|
||||
"cc",
|
||||
"libc",
|
||||
"openssl-src",
|
||||
"pkg-config",
|
||||
"vcpkg",
|
||||
]
|
||||
@@ -4654,6 +4664,7 @@ dependencies = [
|
||||
"memmap2",
|
||||
"mime_guess",
|
||||
"multer",
|
||||
"openssl",
|
||||
"parking_lot 0.12.5",
|
||||
"pprof",
|
||||
"prometheus",
|
||||
|
||||
@@ -606,7 +606,11 @@ async fn run(
|
||||
let grpc_tls_acceptor = grpc_tls_acceptor.clone();
|
||||
let mut shutdown_rx = shutdown_tx.subscribe();
|
||||
tokio::spawn(async move {
|
||||
let addr = grpc_addr.parse().expect("Invalid gRPC address");
|
||||
let addr = tokio::net::lookup_host(&grpc_addr)
|
||||
.await
|
||||
.expect("Failed to resolve gRPC address")
|
||||
.next()
|
||||
.expect("No addresses found for gRPC bind address");
|
||||
let grpc_service = VolumeGrpcService {
|
||||
state: grpc_state.clone(),
|
||||
};
|
||||
|
||||
@@ -249,11 +249,18 @@ func testConcurrentDirectoryOperations(t *testing.T, framework *FuseTestFramewor
|
||||
return
|
||||
}
|
||||
|
||||
// Create file in subdirectory
|
||||
// Create file in subdirectory with retry for transient FUSE errors
|
||||
testFile := filepath.Join(subDir, "test.txt")
|
||||
content := []byte(fmt.Sprintf("Worker %d, Subdir %d", workerID, i))
|
||||
if err := os.WriteFile(testFile, content, 0644); err != nil {
|
||||
addError(fmt.Errorf("worker %d file %d: %v", workerID, i, err))
|
||||
var writeErr error
|
||||
for attempt := 0; attempt < 3; attempt++ {
|
||||
if writeErr = os.WriteFile(testFile, content, 0644); writeErr == nil {
|
||||
break
|
||||
}
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
if writeErr != nil {
|
||||
addError(fmt.Errorf("worker %d file %d: %v", workerID, i, writeErr))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
@@ -85,9 +85,19 @@ func testGitCloneAndPull(t *testing.T, mountPoint, localDir string) {
|
||||
branch := gitOutput(t, localClone, "rev-parse", "--abbrev-ref", "HEAD")
|
||||
gitRun(t, localClone, "push", "origin", branch)
|
||||
|
||||
// The bare repo lives on the FUSE mount and can briefly disappear after
|
||||
// the push completes, and pushed pack objects may not be immediately
|
||||
// consistent on the FUSE layer. Give the mount a chance to settle, then
|
||||
// recover from the local clone if the remote is still missing.
|
||||
if !waitForBareRepoEventually(t, bareRepo, 10*time.Second) {
|
||||
t.Logf("bare repo %s did not stabilise after push; forcing recovery before clone", bareRepo)
|
||||
}
|
||||
refreshDirEntry(t, bareRepo)
|
||||
time.Sleep(1 * time.Second)
|
||||
|
||||
// ---- Phase 3: Clone from mount bare repo into on-mount working dir ----
|
||||
t.Log("Phase 3: clone from mount bare repo to on-mount working dir")
|
||||
gitRun(t, "", "clone", bareRepo, mountClone)
|
||||
ensureMountCloneFromBareWithRecovery(t, bareRepo, localClone, mountClone)
|
||||
|
||||
assertFileContains(t, filepath.Join(mountClone, "README.md"), "# Updated")
|
||||
assertFileContains(t, filepath.Join(mountClone, "src/main.go"), "v2")
|
||||
@@ -290,7 +300,7 @@ func waitForBareRepoEventually(t *testing.T, bareRepo string, timeout time.Durat
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
if isBareRepo(bareRepo) {
|
||||
if isBareRepoAccessible(bareRepo) {
|
||||
return true
|
||||
}
|
||||
refreshDirEntry(t, bareRepo)
|
||||
@@ -312,6 +322,14 @@ func isBareRepo(bareRepo string) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func isBareRepoAccessible(bareRepo string) bool {
|
||||
if !isBareRepo(bareRepo) {
|
||||
return false
|
||||
}
|
||||
out, err := tryGitCommand("", "--git-dir="+bareRepo, "rev-parse", "--is-bare-repository")
|
||||
return err == nil && out == "true"
|
||||
}
|
||||
|
||||
func ensureMountClone(t *testing.T, bareRepo, mountClone string) {
|
||||
t.Helper()
|
||||
require.NoError(t, tryEnsureMountClone(bareRepo, mountClone))
|
||||
@@ -320,7 +338,7 @@ func ensureMountClone(t *testing.T, bareRepo, mountClone string) {
|
||||
// tryEnsureBareRepo verifies the bare repo on the FUSE mount exists.
|
||||
// If it has vanished, it re-creates it from the local clone.
|
||||
func tryEnsureBareRepo(bareRepo, localClone string) error {
|
||||
if _, err := os.Stat(filepath.Join(bareRepo, "HEAD")); err == nil {
|
||||
if isBareRepoAccessible(bareRepo) {
|
||||
return nil
|
||||
}
|
||||
branch, err := tryGitCommand(localClone, "rev-parse", "--abbrev-ref", "HEAD")
|
||||
@@ -342,6 +360,36 @@ func tryEnsureBareRepo(bareRepo, localClone string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensureMountCloneFromBareWithRecovery(t *testing.T, bareRepo, localClone, mountClone string) {
|
||||
t.Helper()
|
||||
const maxAttempts = 3
|
||||
var lastErr error
|
||||
for attempt := 1; attempt <= maxAttempts; attempt++ {
|
||||
if lastErr = tryEnsureMountCloneFromBare(bareRepo, localClone, mountClone); lastErr == nil {
|
||||
return
|
||||
}
|
||||
if attempt == maxAttempts {
|
||||
require.NoError(t, lastErr, "git clone %s %s failed after %d recovery attempts", bareRepo, mountClone, maxAttempts)
|
||||
}
|
||||
t.Logf("clone recovery attempt %d: %v — removing clone for re-create", attempt, lastErr)
|
||||
os.RemoveAll(mountClone)
|
||||
time.Sleep(2 * time.Second)
|
||||
}
|
||||
}
|
||||
|
||||
func tryEnsureMountCloneFromBare(bareRepo, localClone, mountClone string) error {
|
||||
if err := tryEnsureBareRepo(bareRepo, localClone); err != nil {
|
||||
return fmt.Errorf("ensure bare repo: %w", err)
|
||||
}
|
||||
if err := tryEnsureMountClone(bareRepo, mountClone); err != nil {
|
||||
return fmt.Errorf("ensure mount clone: %w", err)
|
||||
}
|
||||
if _, err := tryGitCommand(mountClone, "rev-parse", "HEAD"); err != nil {
|
||||
return fmt.Errorf("verify mount clone: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// tryEnsureMountClone is like ensureMountClone but returns an error instead
|
||||
// of failing the test, for use in recovery loops.
|
||||
func tryEnsureMountClone(bareRepo, mountClone string) error {
|
||||
@@ -550,3 +598,33 @@ func TestTryEnsureBareRepoPreservesCurrentBranch(t *testing.T) {
|
||||
assert.Equal(t, branch, restoredHead, "clone from recovered bare repo should check out the current branch")
|
||||
assertFileContains(t, filepath.Join(restoredClone, "README.md"), "hello recovery")
|
||||
}
|
||||
|
||||
func TestEnsureMountCloneFromBareWithRecoveryRecreatesMissingBareRepo(t *testing.T) {
|
||||
tempDir, err := os.MkdirTemp("", "git_mount_clone_recovery_")
|
||||
require.NoError(t, err)
|
||||
defer os.RemoveAll(tempDir)
|
||||
|
||||
bareRepo := filepath.Join(tempDir, "repo.git")
|
||||
localClone := filepath.Join(tempDir, "clone")
|
||||
mountClone := filepath.Join(tempDir, "mount-clone")
|
||||
|
||||
gitRun(t, "", "init", "--bare", bareRepo)
|
||||
gitRun(t, "", "clone", bareRepo, localClone)
|
||||
gitRun(t, localClone, "config", "user.email", "test@seaweedfs.test")
|
||||
gitRun(t, localClone, "config", "user.name", "Test")
|
||||
|
||||
writeFile(t, localClone, "README.md", "hello clone recovery\n")
|
||||
gitRun(t, localClone, "add", "README.md")
|
||||
gitRun(t, localClone, "commit", "-m", "initial commit")
|
||||
|
||||
branch := gitOutput(t, localClone, "rev-parse", "--abbrev-ref", "HEAD")
|
||||
gitRun(t, localClone, "push", "origin", branch)
|
||||
|
||||
require.NoError(t, os.RemoveAll(bareRepo))
|
||||
|
||||
ensureMountCloneFromBareWithRecovery(t, bareRepo, localClone, mountClone)
|
||||
|
||||
head := gitOutput(t, mountClone, "rev-parse", "--abbrev-ref", "HEAD")
|
||||
assert.Equal(t, branch, head, "recovered clone should stay on the pushed branch")
|
||||
assertFileContains(t, filepath.Join(mountClone, "README.md"), "hello clone recovery")
|
||||
}
|
||||
|
||||
+13
-13
@@ -43,25 +43,25 @@ require (
|
||||
github.com/andybalholm/cascadia v1.3.3 // indirect
|
||||
github.com/appscode/go-querystring v0.0.0-20170504095604-0126cfb3f1dc // indirect
|
||||
github.com/aws/aws-sdk-go v1.55.8 // indirect
|
||||
github.com/aws/aws-sdk-go-v2 v1.41.4 // indirect
|
||||
github.com/aws/aws-sdk-go-v2 v1.41.5 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.4 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/config v1.32.9 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/credentials v1.19.12 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.20 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/config v1.32.13 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/credentials v1.19.13 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.21 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/feature/s3/manager v1.20.12 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.20 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.20 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.4 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.21 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.21 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.6 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.17 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.7 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.8 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.20 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.21 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.17 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/s3 v1.96.0 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/signin v1.0.8 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/sso v1.30.13 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.17 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/sts v1.41.9 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/signin v1.0.9 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/sso v1.30.14 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.18 // indirect
|
||||
github.com/aws/aws-sdk-go-v2/service/sts v1.41.10 // indirect
|
||||
github.com/aws/smithy-go v1.24.2 // indirect
|
||||
github.com/bahlo/generic-list-go v0.2.0 // indirect
|
||||
github.com/beorn7/perks v1.0.1 // indirect
|
||||
@@ -104,7 +104,7 @@ require (
|
||||
github.com/go-chi/chi/v5 v5.2.5 // indirect
|
||||
github.com/go-darwin/apfs v0.0.0-20211011131704-f84b94dbf348 // indirect
|
||||
github.com/go-git/go-billy/v5 v5.6.2 // indirect
|
||||
github.com/go-jose/go-jose/v4 v4.1.3 // indirect
|
||||
github.com/go-jose/go-jose/v4 v4.1.4 // indirect
|
||||
github.com/go-logr/logr v1.4.3 // indirect
|
||||
github.com/go-logr/stdr v1.2.2 // indirect
|
||||
github.com/go-ole/go-ole v1.3.0 // indirect
|
||||
|
||||
+28
-28
@@ -112,44 +112,44 @@ github.com/appscode/go-querystring v0.0.0-20170504095604-0126cfb3f1dc h1:LoL75er
|
||||
github.com/appscode/go-querystring v0.0.0-20170504095604-0126cfb3f1dc/go.mod h1:w648aMHEgFYS6xb0KVMMtZ2uMeemhiKCuD2vj6gY52A=
|
||||
github.com/aws/aws-sdk-go v1.55.8 h1:JRmEUbU52aJQZ2AjX4q4Wu7t4uZjOu71uyNmaWlUkJQ=
|
||||
github.com/aws/aws-sdk-go v1.55.8/go.mod h1:ZkViS9AqA6otK+JBBNH2++sx1sgxrPKcSzPPvQkUtXk=
|
||||
github.com/aws/aws-sdk-go-v2 v1.41.4 h1:10f50G7WyU02T56ox1wWXq+zTX9I1zxG46HYuG1hH/k=
|
||||
github.com/aws/aws-sdk-go-v2 v1.41.4/go.mod h1:mwsPRE8ceUUpiTgF7QmQIJ7lgsKUPQOUl3o72QBrE1o=
|
||||
github.com/aws/aws-sdk-go-v2 v1.41.5 h1:dj5kopbwUsVUVFgO4Fi5BIT3t4WyqIDjGKCangnV/yY=
|
||||
github.com/aws/aws-sdk-go-v2 v1.41.5/go.mod h1:mwsPRE8ceUUpiTgF7QmQIJ7lgsKUPQOUl3o72QBrE1o=
|
||||
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.4 h1:489krEF9xIGkOaaX3CE/Be2uWjiXrkCH6gUX+bZA/BU=
|
||||
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.4/go.mod h1:IOAPF6oT9KCsceNTvvYMNHy0+kMF8akOjeDvPENWxp4=
|
||||
github.com/aws/aws-sdk-go-v2/config v1.32.9 h1:ktda/mtAydeObvJXlHzyGpK1xcsLaP16zfUPDGoW90A=
|
||||
github.com/aws/aws-sdk-go-v2/config v1.32.9/go.mod h1:U+fCQ+9QKsLW786BCfEjYRj34VVTbPdsLP3CHSYXMOI=
|
||||
github.com/aws/aws-sdk-go-v2/credentials v1.19.12 h1:oqtA6v+y5fZg//tcTWahyN9PEn5eDU/Wpvc2+kJ4aY8=
|
||||
github.com/aws/aws-sdk-go-v2/credentials v1.19.12/go.mod h1:U3R1RtSHx6NB0DvEQFGyf/0sbrpJrluENHdPy1j/3TE=
|
||||
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.20 h1:zOgq3uezl5nznfoK3ODuqbhVg1JzAGDUhXOsU0IDCAo=
|
||||
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.20/go.mod h1:z/MVwUARehy6GAg/yQ1GO2IMl0k++cu1ohP9zo887wE=
|
||||
github.com/aws/aws-sdk-go-v2/config v1.32.13 h1:5KgbxMaS2coSWRrx9TX/QtWbqzgQkOdEa3sZPhBhCSg=
|
||||
github.com/aws/aws-sdk-go-v2/config v1.32.13/go.mod h1:8zz7wedqtCbw5e9Mi2doEwDyEgHcEE9YOJp6a8jdSMY=
|
||||
github.com/aws/aws-sdk-go-v2/credentials v1.19.13 h1:mA59E3fokBvyEGHKFdnpNNrvaR351cqiHgRg+JzOSRI=
|
||||
github.com/aws/aws-sdk-go-v2/credentials v1.19.13/go.mod h1:yoTXOQKea18nrM69wGF9jBdG4WocSZA1h38A+t/MAsk=
|
||||
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.21 h1:NUS3K4BTDArQqNu2ih7yeDLaS3bmHD0YndtA6UP884g=
|
||||
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.21/go.mod h1:YWNWJQNjKigKY1RHVJCuupeWDrrHjRqHm0N9rdrWzYI=
|
||||
github.com/aws/aws-sdk-go-v2/feature/s3/manager v1.20.12 h1:Zy6Tme1AA13kX8x3CnkHx5cqdGWGaj/anwOiWGnA0Xo=
|
||||
github.com/aws/aws-sdk-go-v2/feature/s3/manager v1.20.12/go.mod h1:ql4uXYKoTM9WUAUSmthY4AtPVrlTBZOvnBJTiCUdPxI=
|
||||
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.20 h1:CNXO7mvgThFGqOFgbNAP2nol2qAWBOGfqR/7tQlvLmc=
|
||||
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.20/go.mod h1:oydPDJKcfMhgfcgBUZaG+toBbwy8yPWubJXBVERtI4o=
|
||||
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.20 h1:tN6W/hg+pkM+tf9XDkWUbDEjGLb+raoBMFsTodcoYKw=
|
||||
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.20/go.mod h1:YJ898MhD067hSHA6xYCx5ts/jEd8BSOLtQDL3iZsvbc=
|
||||
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.4 h1:WKuaxf++XKWlHWu9ECbMlha8WOEGm0OUEZqm4K/Gcfk=
|
||||
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.4/go.mod h1:ZWy7j6v1vWGmPReu0iSGvRiise4YI5SkR3OHKTZ6Wuc=
|
||||
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.21 h1:Rgg6wvjjtX8bNHcvi9OnXWwcE0a2vGpbwmtICOsvcf4=
|
||||
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.21/go.mod h1:A/kJFst/nm//cyqonihbdpQZwiUhhzpqTsdbhDdRF9c=
|
||||
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.21 h1:PEgGVtPoB6NTpPrBgqSE5hE/o47Ij9qk/SEZFbUOe9A=
|
||||
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.21/go.mod h1:p+hz+PRAYlY3zcpJhPwXlLC4C+kqn70WIHwnzAfs6ps=
|
||||
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.6 h1:qYQ4pzQ2Oz6WpQ8T3HvGHnZydA72MnLuFK9tJwmrbHw=
|
||||
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.6/go.mod h1:O3h0IK87yXci+kg6flUKzJnWeziQUKciKrLjcatSNcY=
|
||||
github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.17 h1:JqcdRG//czea7Ppjb+g/n4o8i/R50aTBHkA7vu0lK+k=
|
||||
github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.17/go.mod h1:CO+WeGmIdj/MlPel2KwID9Gt7CNq4M65HUfBW97liM0=
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.7 h1:5EniKhLZe4xzL7a+fU3C2tfUN4nWIqlLesfrjkuPFTY=
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.7/go.mod h1:x0nZssQ3qZSnIcePWLvcoFisRXJzcTVvYpAAdYX8+GI=
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.8 h1:Z5EiPIzXKewUQK0QTMkutjiaPVeVYXX7KIqhXu/0fXs=
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.8/go.mod h1:FsTpJtvC4U1fyDXk7c71XoDv3HlRm8V3NiYLeYLh5YE=
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.20 h1:2HvVAIq+YqgGotK6EkMf+KIEqTISmTYh5zLpYyeTo1Y=
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.20/go.mod h1:V4X406Y666khGa8ghKmphma/7C0DAtEQYhkq9z4vpbk=
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.21 h1:c31//R3xgIJMSC8S6hEVq+38DcvUlgFY0FM6mSI5oto=
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.21/go.mod h1:r6+pf23ouCB718FUxaqzZdbpYFyDtehyZcmP5KL9FkA=
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.17 h1:bGeHBsGZx0Dvu/eJC0Lh9adJa3M1xREcndxLNZlve2U=
|
||||
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.17/go.mod h1:dcW24lbU0CzHusTE8LLHhRLI42ejmINN8Lcr22bwh/g=
|
||||
github.com/aws/aws-sdk-go-v2/service/s3 v1.96.0 h1:oeu8VPlOre74lBA/PMhxa5vewaMIMmILM+RraSyB8KA=
|
||||
github.com/aws/aws-sdk-go-v2/service/s3 v1.96.0/go.mod h1:5jggDlZ2CLQhwJBiZJb4vfk4f0GxWdEDruWKEJ1xOdo=
|
||||
github.com/aws/aws-sdk-go-v2/service/signin v1.0.8 h1:0GFOLzEbOyZABS3PhYfBIx2rNBACYcKty+XGkTgw1ow=
|
||||
github.com/aws/aws-sdk-go-v2/service/signin v1.0.8/go.mod h1:LXypKvk85AROkKhOG6/YEcHFPoX+prKTowKnVdcaIxE=
|
||||
github.com/aws/aws-sdk-go-v2/service/sso v1.30.13 h1:kiIDLZ005EcKomYYITtfsjn7dtOwHDOFy7IbPXKek2o=
|
||||
github.com/aws/aws-sdk-go-v2/service/sso v1.30.13/go.mod h1:2h/xGEowcW/g38g06g3KpRWDlT+OTfxxI0o1KqayAB8=
|
||||
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.17 h1:jzKAXIlhZhJbnYwHbvUQZEB8KfgAEuG0dc08Bkda7NU=
|
||||
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.17/go.mod h1:Al9fFsXjv4KfbzQHGe6V4NZSZQXecFcvaIF4e70FoRA=
|
||||
github.com/aws/aws-sdk-go-v2/service/sts v1.41.9 h1:Cng+OOwCHmFljXIxpEVXAGMnBia8MSU6Ch5i9PgBkcU=
|
||||
github.com/aws/aws-sdk-go-v2/service/sts v1.41.9/go.mod h1:LrlIndBDdjA/EeXeyNBle+gyCwTlizzW5ycgWnvIxkk=
|
||||
github.com/aws/aws-sdk-go-v2/service/signin v1.0.9 h1:QKZH0S178gCmFEgst8hN0mCX1KxLgHBKKY/CLqwP8lg=
|
||||
github.com/aws/aws-sdk-go-v2/service/signin v1.0.9/go.mod h1:7yuQJoT+OoH8aqIxw9vwF+8KpvLZ8AWmvmUWHsGQZvI=
|
||||
github.com/aws/aws-sdk-go-v2/service/sso v1.30.14 h1:GcLE9ba5ehAQma6wlopUesYg/hbcOhFNWTjELkiWkh4=
|
||||
github.com/aws/aws-sdk-go-v2/service/sso v1.30.14/go.mod h1:WSvS1NLr7JaPunCXqpJnWk1Bjo7IxzZXrZi1QQCkuqM=
|
||||
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.18 h1:mP49nTpfKtpXLt5SLn8Uv8z6W+03jYVoOSAl/c02nog=
|
||||
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.18/go.mod h1:YO8TrYtFdl5w/4vmjL8zaBSsiNp3w0L1FfKVKenZT7w=
|
||||
github.com/aws/aws-sdk-go-v2/service/sts v1.41.10 h1:p8ogvvLugcR/zLBXTXrTkj0RYBUdErbMnAFFp12Lm/U=
|
||||
github.com/aws/aws-sdk-go-v2/service/sts v1.41.10/go.mod h1:60dv0eZJfeVXfbT1tFJinbHrDfSJ2GZl4Q//OSSNAVw=
|
||||
github.com/aws/smithy-go v1.24.2 h1:FzA3bu/nt/vDvmnkg+R8Xl46gmzEDam6mZ1hzmwXFng=
|
||||
github.com/aws/smithy-go v1.24.2/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc=
|
||||
github.com/bahlo/generic-list-go v0.2.0 h1:5sz/EEAK+ls5wF+NeqDpk5+iNdMDXrh3z3nPnH1Wvgk=
|
||||
@@ -284,8 +284,8 @@ github.com/go-git/go-billy/v5 v5.6.2/go.mod h1:rcFC2rAsp/erv7CMz9GczHcuD0D32fWzH
|
||||
github.com/go-gl/glfw v0.0.0-20190409004039-e6da0acd62b1/go.mod h1:vR7hzQXu2zJy9AVAgeJqvqgH9Q5CA+iKCZ2gyEVpxRU=
|
||||
github.com/go-gl/glfw/v3.3/glfw v0.0.0-20191125211704-12ad95a8df72/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8=
|
||||
github.com/go-gl/glfw/v3.3/glfw v0.0.0-20200222043503-6f7a984d4dc4/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8=
|
||||
github.com/go-jose/go-jose/v4 v4.1.3 h1:CVLmWDhDVRa6Mi/IgCgaopNosCaHz7zrMeF9MlZRkrs=
|
||||
github.com/go-jose/go-jose/v4 v4.1.3/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
|
||||
github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA=
|
||||
github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
|
||||
github.com/go-logr/logr v1.2.2/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A=
|
||||
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
|
||||
github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
|
||||
@@ -707,8 +707,8 @@ github.com/xanzy/ssh-agent v0.3.3 h1:+/15pJfg/RsTxqYcX6fHqOXZwwMP+2VyYWJeWM2qQFM
|
||||
github.com/xanzy/ssh-agent v0.3.3/go.mod h1:6dzNDKs0J9rVPHPhaGCukekBHKqfl+L3KghI1Bc68Uw=
|
||||
github.com/xdg-go/pbkdf2 v1.0.0 h1:Su7DPu48wXMwC3bs7MCNG+z4FhcyEuz5dlvchbq0B0c=
|
||||
github.com/xdg-go/pbkdf2 v1.0.0/go.mod h1:jrpuAogTd400dnrH08LKmI/xc1MbPOebTwRqcT5RDeI=
|
||||
github.com/xdg-go/scram v1.1.2 h1:FHX5I5B4i4hKRVRBCFRxq1iQRej7WO3hhBuJf+UUySY=
|
||||
github.com/xdg-go/scram v1.1.2/go.mod h1:RT/sEzTbU5y00aCK8UOx6R7YryM0iF1N2MOmC3kKLN4=
|
||||
github.com/xdg-go/scram v1.2.0 h1:bYKF2AEwG5rqd1BumT4gAnvwU/M9nBp2pTSxeZw7Wvs=
|
||||
github.com/xdg-go/scram v1.2.0/go.mod h1:3dlrS0iBaWKYVt2ZfA4cj48umJZ+cAEbR6/SjLA88I8=
|
||||
github.com/xdg-go/stringprep v1.0.4 h1:XLI/Ng3O1Atzq0oBs3TWm+5ZVgkq2aqdlvP9JtoZ6c8=
|
||||
github.com/xdg-go/stringprep v1.0.4/go.mod h1:mPGuuIYwz7CmR2bT9j4GbQqutWS1zV24gijq1dTyGkM=
|
||||
github.com/xeipuuv/gojsonpointer v0.0.0-20180127040702-4e3ac2762d5f/go.mod h1:N2zxlSyiKSe5eX1tZViRH5QA0qijqEDrYZiPEAiq3wU=
|
||||
|
||||
@@ -21,7 +21,9 @@ import (
|
||||
"github.com/aws/aws-sdk-go-v2/credentials"
|
||||
"github.com/aws/aws-sdk-go-v2/service/s3"
|
||||
"github.com/seaweedfs/seaweedfs/test/volume_server/framework"
|
||||
"github.com/seaweedfs/seaweedfs/weed/cluster/lock_manager"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/filer_pb"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/master_pb"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc"
|
||||
@@ -139,6 +141,7 @@ func startDistributedLockCluster(t *testing.T) *distributedLockCluster {
|
||||
require.NoError(t, cluster.waitForTCP(cluster.filerGRPCAddress(i), 30*time.Second), "wait for filer %d grpc\n%s", i, cluster.tailLog(fmt.Sprintf("filer%d.log", i)))
|
||||
}
|
||||
require.NoError(t, cluster.waitForFilerCount(2, 30*time.Second), "wait for filer group registration")
|
||||
require.NoError(t, cluster.waitForLockRingConverged(30*time.Second), "wait for lock ring convergence")
|
||||
|
||||
for i := 0; i < 2; i++ {
|
||||
require.NoError(t, cluster.startS3(i), "start s3 %d", i)
|
||||
@@ -411,6 +414,130 @@ func (c *distributedLockCluster) waitForFilerCount(expected int, timeout time.Du
|
||||
return fmt.Errorf("timed out waiting for %d filers in group %q", expected, c.filerGroup)
|
||||
}
|
||||
|
||||
// waitForLockRingConverged verifies that both filers have a consistent view of the
|
||||
// lock ring by acquiring the same lock through each filer and checking mutual exclusion.
|
||||
// This guards against the window between master seeing both filers and the filers
|
||||
// actually receiving the LockRingUpdate broadcast (delayed by the stabilization timer).
|
||||
//
|
||||
// A single arbitrary key could false-pass: if the key's real primary is filer0 and
|
||||
// filer0 has a stale ring (only sees itself), it still grants correctly because it IS
|
||||
// the primary. To catch the stale ring we must also test a key whose real primary is
|
||||
// filer1. So we generate one test key per primary filer using the same consistent-hash
|
||||
// ring as production, and require mutual exclusion for all of them.
|
||||
func (c *distributedLockCluster) waitForLockRingConverged(timeout time.Duration) error {
|
||||
deadline := time.Now().Add(timeout)
|
||||
|
||||
owners := make([]pb.ServerAddress, 0, len(c.filerPorts))
|
||||
for i := range c.filerPorts {
|
||||
owners = append(owners, c.filerServerAddress(i))
|
||||
}
|
||||
ring := lock_manager.NewHashRing(lock_manager.DefaultVnodeCount)
|
||||
ring.SetServers(owners)
|
||||
|
||||
attempt := 0
|
||||
for time.Now().Before(deadline) {
|
||||
// Generate one unique test key per primary filer
|
||||
testKeys := c.convergenceKeysPerPrimary(ring, owners, attempt)
|
||||
attempt++
|
||||
|
||||
allConverged := true
|
||||
for _, key := range testKeys {
|
||||
converged, _ := c.checkLockMutualExclusion(key)
|
||||
if !converged {
|
||||
allConverged = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if allConverged {
|
||||
return nil
|
||||
}
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
}
|
||||
return fmt.Errorf("lock ring did not converge: both filers independently grant the same lock")
|
||||
}
|
||||
|
||||
// convergenceKeysPerPrimary returns one lock key per distinct primary filer,
|
||||
// using the same consistent-hash ring as production routing.
|
||||
func (c *distributedLockCluster) convergenceKeysPerPrimary(ring *lock_manager.HashRing, owners []pb.ServerAddress, attempt int) []string {
|
||||
keysByOwner := make(map[pb.ServerAddress]string, len(owners))
|
||||
for i := 0; len(keysByOwner) < len(owners) && i < 1024; i++ {
|
||||
candidate := fmt.Sprintf("convergence-%d-%d", attempt, i)
|
||||
primary := ring.GetPrimary(candidate)
|
||||
if _, exists := keysByOwner[primary]; !exists {
|
||||
keysByOwner[primary] = candidate
|
||||
}
|
||||
}
|
||||
keys := make([]string, 0, len(keysByOwner))
|
||||
for _, k := range keysByOwner {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
// checkLockMutualExclusion acquires a lock via filer0, then tries the same lock via filer1.
|
||||
// Returns true if mutual exclusion holds (second attempt is denied).
|
||||
func (c *distributedLockCluster) checkLockMutualExclusion(testKey string) (bool, error) {
|
||||
conn0, err := grpc.NewClient(c.filerGRPCAddress(0), grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer conn0.Close()
|
||||
|
||||
conn1, err := grpc.NewClient(c.filerGRPCAddress(1), grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer conn1.Close()
|
||||
|
||||
client0 := filer_pb.NewSeaweedFilerClient(conn0)
|
||||
client1 := filer_pb.NewSeaweedFilerClient(conn1)
|
||||
|
||||
// Acquire lock via filer0
|
||||
ctx0, cancel0 := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel0()
|
||||
resp0, err := client0.DistributedLock(ctx0, &filer_pb.LockRequest{
|
||||
Name: testKey,
|
||||
SecondsToLock: 5,
|
||||
Owner: "convergence-filer0",
|
||||
})
|
||||
if err != nil || resp0.RenewToken == "" {
|
||||
return false, fmt.Errorf("filer0 lock failed: err=%v resp=%v", err, resp0)
|
||||
}
|
||||
defer func() {
|
||||
// Always release the lock we acquired
|
||||
client0.DistributedUnlock(context.Background(), &filer_pb.UnlockRequest{
|
||||
Name: testKey,
|
||||
RenewToken: resp0.RenewToken,
|
||||
})
|
||||
}()
|
||||
|
||||
// Try the same lock via filer1 - should be denied
|
||||
ctx1, cancel1 := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel1()
|
||||
resp1, err := client1.DistributedLock(ctx1, &filer_pb.LockRequest{
|
||||
Name: testKey,
|
||||
SecondsToLock: 5,
|
||||
Owner: "convergence-filer1",
|
||||
})
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if resp1.RenewToken != "" {
|
||||
// Both filers granted the lock - ring not converged. Release the second lock.
|
||||
client1.DistributedUnlock(context.Background(), &filer_pb.UnlockRequest{
|
||||
Name: testKey,
|
||||
RenewToken: resp1.RenewToken,
|
||||
})
|
||||
return false, nil
|
||||
}
|
||||
// Verify the denial is specifically because the lock is already held,
|
||||
// not due to a transient error that might give a false positive.
|
||||
if !strings.Contains(resp1.Error, "lock already owned") {
|
||||
return false, nil
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (c *distributedLockCluster) waitForHTTP(url string, timeout time.Duration) error {
|
||||
client := &net.Dialer{Timeout: time.Second}
|
||||
httpClient := &httpClientWithDialer{dialer: client}
|
||||
|
||||
@@ -6,7 +6,6 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
@@ -15,9 +14,9 @@ import (
|
||||
"github.com/aws/aws-sdk-go-v2/aws"
|
||||
"github.com/aws/aws-sdk-go-v2/service/s3"
|
||||
"github.com/aws/smithy-go"
|
||||
"github.com/seaweedfs/seaweedfs/weed/cluster/lock_manager"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb"
|
||||
"github.com/seaweedfs/seaweedfs/weed/s3api/s3_constants"
|
||||
"github.com/seaweedfs/seaweedfs/weed/util"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
@@ -149,14 +148,15 @@ func (c *distributedLockCluster) findLockOwnerKeys(bucket, prefix string) map[pb
|
||||
for i := range c.filerPorts {
|
||||
owners = append(owners, c.filerServerAddress(i))
|
||||
}
|
||||
sort.Slice(owners, func(i, j int) bool {
|
||||
return owners[i] < owners[j]
|
||||
})
|
||||
|
||||
ring := lock_manager.NewHashRing(lock_manager.DefaultVnodeCount)
|
||||
ring.SetServers(owners)
|
||||
|
||||
keysByOwner := make(map[pb.ServerAddress]string, len(owners))
|
||||
for i := 0; i < 1024 && len(keysByOwner) < len(owners); i++ {
|
||||
key := fmt.Sprintf("%s-%03d.txt", prefix, i)
|
||||
lockOwner := ownerForObjectLock(bucket, key, owners)
|
||||
lockKey := fmt.Sprintf("s3.object.write:/buckets/%s/%s", bucket, s3_constants.NormalizeObjectKey(key))
|
||||
lockOwner := ring.GetPrimary(lockKey)
|
||||
if _, exists := keysByOwner[lockOwner]; !exists {
|
||||
keysByOwner[lockOwner] = key
|
||||
}
|
||||
@@ -164,15 +164,6 @@ func (c *distributedLockCluster) findLockOwnerKeys(bucket, prefix string) map[pb
|
||||
return keysByOwner
|
||||
}
|
||||
|
||||
func ownerForObjectLock(bucket, object string, owners []pb.ServerAddress) pb.ServerAddress {
|
||||
lockKey := fmt.Sprintf("s3.object.write:/buckets/%s/%s", bucket, s3_constants.NormalizeObjectKey(object))
|
||||
hash := util.HashStringToLong(lockKey)
|
||||
if hash < 0 {
|
||||
hash = -hash
|
||||
}
|
||||
return owners[hash%int64(len(owners))]
|
||||
}
|
||||
|
||||
func lockOwnerLabel(owner pb.ServerAddress) string {
|
||||
replacer := strings.NewReplacer(":", "_", ".", "_")
|
||||
return "owner_" + replacer.Replace(string(owner))
|
||||
|
||||
@@ -0,0 +1,511 @@
|
||||
package iam
|
||||
|
||||
import (
|
||||
"encoding/xml"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/aws/aws-sdk-go/aws"
|
||||
"github.com/aws/aws-sdk-go/aws/credentials"
|
||||
"github.com/aws/aws-sdk-go/aws/session"
|
||||
"github.com/aws/aws-sdk-go/service/s3"
|
||||
v4 "github.com/aws/aws-sdk-go/aws/signer/v4"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// GetFederationTokenTestResponse represents the STS GetFederationToken response
|
||||
type GetFederationTokenTestResponse struct {
|
||||
XMLName xml.Name `xml:"GetFederationTokenResponse"`
|
||||
Result struct {
|
||||
Credentials struct {
|
||||
AccessKeyId string `xml:"AccessKeyId"`
|
||||
SecretAccessKey string `xml:"SecretAccessKey"`
|
||||
SessionToken string `xml:"SessionToken"`
|
||||
Expiration string `xml:"Expiration"`
|
||||
} `xml:"Credentials"`
|
||||
FederatedUser struct {
|
||||
FederatedUserId string `xml:"FederatedUserId"`
|
||||
Arn string `xml:"Arn"`
|
||||
} `xml:"FederatedUser"`
|
||||
} `xml:"GetFederationTokenResult"`
|
||||
}
|
||||
|
||||
func getTestCredentials() (string, string) {
|
||||
accessKey := os.Getenv("STS_TEST_ACCESS_KEY")
|
||||
if accessKey == "" {
|
||||
accessKey = "admin"
|
||||
}
|
||||
secretKey := os.Getenv("STS_TEST_SECRET_KEY")
|
||||
if secretKey == "" {
|
||||
secretKey = "admin"
|
||||
}
|
||||
return accessKey, secretKey
|
||||
}
|
||||
|
||||
// isGetFederationTokenImplemented checks if the running server supports GetFederationToken
|
||||
func isGetFederationTokenImplemented(t *testing.T) bool {
|
||||
accessKey, secretKey := getTestCredentials()
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"probe"},
|
||||
}, accessKey, secretKey)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var errResp STSErrorTestResponse
|
||||
if xml.Unmarshal(body, &errResp) == nil {
|
||||
if errResp.Error.Code == "InvalidAction" || errResp.Error.Code == "NotImplemented" {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// TestSTSGetFederationTokenValidation tests input validation for the GetFederationToken endpoint
|
||||
func TestSTSGetFederationTokenValidation(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("Skipping integration test in short mode")
|
||||
}
|
||||
|
||||
if !isSTSEndpointRunning(t) {
|
||||
t.Fatal("SeaweedFS STS endpoint is not running at", TestSTSEndpoint, "- please run 'make setup-all-tests' first")
|
||||
}
|
||||
|
||||
if !isGetFederationTokenImplemented(t) {
|
||||
t.Fatal("GetFederationToken action is not implemented in the running server")
|
||||
}
|
||||
|
||||
accessKey, secretKey := getTestCredentials()
|
||||
|
||||
t.Run("missing_name", func(t *testing.T) {
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
// Name is missing
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var errResp STSErrorTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &errResp), "Failed to parse: %s", string(body))
|
||||
assert.Equal(t, "MissingParameter", errResp.Error.Code)
|
||||
})
|
||||
|
||||
t.Run("name_too_short", func(t *testing.T) {
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"A"},
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var errResp STSErrorTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &errResp), "Failed to parse: %s", string(body))
|
||||
assert.Equal(t, "InvalidParameterValue", errResp.Error.Code)
|
||||
})
|
||||
|
||||
t.Run("name_too_long", func(t *testing.T) {
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {strings.Repeat("A", 33)},
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var errResp STSErrorTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &errResp), "Failed to parse: %s", string(body))
|
||||
assert.Equal(t, "InvalidParameterValue", errResp.Error.Code)
|
||||
})
|
||||
|
||||
t.Run("name_invalid_characters", func(t *testing.T) {
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"bad name"},
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var errResp STSErrorTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &errResp), "Failed to parse: %s", string(body))
|
||||
assert.Equal(t, "InvalidParameterValue", errResp.Error.Code)
|
||||
})
|
||||
|
||||
t.Run("duration_too_short", func(t *testing.T) {
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"TestApp"},
|
||||
"DurationSeconds": {"100"},
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var errResp STSErrorTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &errResp), "Failed to parse: %s", string(body))
|
||||
assert.Equal(t, "InvalidParameterValue", errResp.Error.Code)
|
||||
})
|
||||
|
||||
t.Run("duration_too_long", func(t *testing.T) {
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"TestApp"},
|
||||
"DurationSeconds": {"200000"},
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var errResp STSErrorTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &errResp), "Failed to parse: %s", string(body))
|
||||
assert.Equal(t, "InvalidParameterValue", errResp.Error.Code)
|
||||
})
|
||||
|
||||
t.Run("malformed_policy", func(t *testing.T) {
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"TestApp"},
|
||||
"Policy": {"not-valid-json"},
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var errResp STSErrorTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &errResp), "Failed to parse: %s", string(body))
|
||||
assert.Equal(t, "MalformedPolicyDocument", errResp.Error.Code)
|
||||
})
|
||||
|
||||
t.Run("anonymous_rejected", func(t *testing.T) {
|
||||
// GetFederationToken requires SigV4, anonymous should fail
|
||||
resp, err := callSTSAPI(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"TestApp"},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
assert.NotEqual(t, http.StatusOK, resp.StatusCode)
|
||||
})
|
||||
}
|
||||
|
||||
// TestSTSGetFederationTokenRejectTemporaryCredentials tests that temporary
|
||||
// credentials (session tokens) are rejected by GetFederationToken
|
||||
func TestSTSGetFederationTokenRejectTemporaryCredentials(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("Skipping integration test in short mode")
|
||||
}
|
||||
|
||||
if !isSTSEndpointRunning(t) {
|
||||
t.Skip("SeaweedFS STS endpoint is not running at", TestSTSEndpoint)
|
||||
}
|
||||
|
||||
if !isGetFederationTokenImplemented(t) {
|
||||
t.Skip("GetFederationToken not implemented")
|
||||
}
|
||||
|
||||
accessKey, secretKey := getTestCredentials()
|
||||
|
||||
// First, obtain temporary credentials via AssumeRole
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"AssumeRole"},
|
||||
"Version": {"2011-06-15"},
|
||||
"RoleArn": {"arn:aws:iam::role/admin"},
|
||||
"RoleSessionName": {"temp-session"},
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Skipf("AssumeRole failed (may not be configured): status=%d body=%s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var assumeResp AssumeRoleTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &assumeResp), "Parse AssumeRole response: %s", string(body))
|
||||
|
||||
tempAccessKey := assumeResp.Result.Credentials.AccessKeyId
|
||||
tempSecretKey := assumeResp.Result.Credentials.SecretAccessKey
|
||||
tempSessionToken := assumeResp.Result.Credentials.SessionToken
|
||||
require.NotEmpty(t, tempAccessKey)
|
||||
require.NotEmpty(t, tempSessionToken)
|
||||
|
||||
// Now try GetFederationToken with the temporary credentials
|
||||
// Include X-Amz-Security-Token header which marks this as a temp credential call
|
||||
params := url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"ShouldFail"},
|
||||
}
|
||||
|
||||
reqBody := params.Encode()
|
||||
req, err := http.NewRequest(http.MethodPost, TestSTSEndpoint+"/", strings.NewReader(reqBody))
|
||||
require.NoError(t, err)
|
||||
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
req.Header.Set("X-Amz-Security-Token", tempSessionToken)
|
||||
|
||||
creds := credentials.NewStaticCredentials(tempAccessKey, tempSecretKey, tempSessionToken)
|
||||
signer := v4.NewSigner(creds)
|
||||
_, err = signer.Sign(req, strings.NewReader(reqBody), "sts", "us-east-1", time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
client := &http.Client{Timeout: 30 * time.Second}
|
||||
resp2, err := client.Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp2.Body.Close()
|
||||
|
||||
body2, _ := io.ReadAll(resp2.Body)
|
||||
assert.Equal(t, http.StatusForbidden, resp2.StatusCode,
|
||||
"GetFederationToken should reject temporary credentials: %s", string(body2))
|
||||
assert.Contains(t, string(body2), "temporary credentials",
|
||||
"Error should mention temporary credentials")
|
||||
}
|
||||
|
||||
// TestSTSGetFederationTokenSuccess tests a successful GetFederationToken call
|
||||
// and verifies the returned credentials can be used to access S3
|
||||
func TestSTSGetFederationTokenSuccess(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("Skipping integration test in short mode")
|
||||
}
|
||||
|
||||
if !isSTSEndpointRunning(t) {
|
||||
t.Skip("SeaweedFS STS endpoint is not running at", TestSTSEndpoint)
|
||||
}
|
||||
|
||||
if !isGetFederationTokenImplemented(t) {
|
||||
t.Skip("GetFederationToken not implemented")
|
||||
}
|
||||
|
||||
accessKey, secretKey := getTestCredentials()
|
||||
|
||||
t.Run("basic_success", func(t *testing.T) {
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"AppClient"},
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
t.Logf("Response status: %d, body: %s", resp.StatusCode, string(body))
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
var errResp STSErrorTestResponse
|
||||
_ = xml.Unmarshal(body, &errResp)
|
||||
t.Fatalf("GetFederationToken failed: code=%s message=%s", errResp.Error.Code, errResp.Error.Message)
|
||||
}
|
||||
|
||||
var stsResp GetFederationTokenTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &stsResp), "Parse response: %s", string(body))
|
||||
|
||||
creds := stsResp.Result.Credentials
|
||||
assert.NotEmpty(t, creds.AccessKeyId)
|
||||
assert.NotEmpty(t, creds.SecretAccessKey)
|
||||
assert.NotEmpty(t, creds.SessionToken)
|
||||
assert.NotEmpty(t, creds.Expiration)
|
||||
|
||||
fedUser := stsResp.Result.FederatedUser
|
||||
assert.Contains(t, fedUser.Arn, "federated-user/AppClient")
|
||||
assert.Contains(t, fedUser.FederatedUserId, "AppClient")
|
||||
})
|
||||
|
||||
t.Run("with_custom_duration", func(t *testing.T) {
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"DurationTest"},
|
||||
"DurationSeconds": {"3600"},
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
t.Logf("Response status: %d, body: %s", resp.StatusCode, string(body))
|
||||
|
||||
if resp.StatusCode == http.StatusOK {
|
||||
var stsResp GetFederationTokenTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &stsResp))
|
||||
assert.NotEmpty(t, stsResp.Result.Credentials.AccessKeyId)
|
||||
|
||||
// Verify expiration is roughly 1 hour from now
|
||||
expTime, err := time.Parse(time.RFC3339, stsResp.Result.Credentials.Expiration)
|
||||
require.NoError(t, err)
|
||||
diff := time.Until(expTime)
|
||||
assert.InDelta(t, 3600, diff.Seconds(), 60,
|
||||
"Expiration should be ~1 hour from now")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("with_36_hour_duration", func(t *testing.T) {
|
||||
// GetFederationToken allows up to 36 hours (unlike AssumeRole's 12h max)
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"LongDuration"},
|
||||
"DurationSeconds": {"129600"}, // 36 hours
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode == http.StatusOK {
|
||||
var stsResp GetFederationTokenTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &stsResp))
|
||||
|
||||
expTime, err := time.Parse(time.RFC3339, stsResp.Result.Credentials.Expiration)
|
||||
require.NoError(t, err)
|
||||
diff := time.Until(expTime)
|
||||
assert.InDelta(t, 129600, diff.Seconds(), 60,
|
||||
"Expiration should be ~36 hours from now")
|
||||
} else {
|
||||
// Duration should not cause a rejection
|
||||
var errResp STSErrorTestResponse
|
||||
_ = xml.Unmarshal(body, &errResp)
|
||||
assert.NotContains(t, errResp.Error.Message, "DurationSeconds",
|
||||
"36-hour duration should be accepted by GetFederationToken")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSTSGetFederationTokenWithSessionPolicy tests that vended credentials
|
||||
// are scoped down by an inline session policy
|
||||
func TestSTSGetFederationTokenWithSessionPolicy(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("Skipping integration test in short mode")
|
||||
}
|
||||
|
||||
if !isSTSEndpointRunning(t) {
|
||||
t.Skip("SeaweedFS STS endpoint is not running at", TestSTSEndpoint)
|
||||
}
|
||||
|
||||
if !isGetFederationTokenImplemented(t) {
|
||||
t.Skip("GetFederationToken not implemented")
|
||||
}
|
||||
|
||||
accessKey, secretKey := getTestCredentials()
|
||||
|
||||
// Create a test bucket using admin credentials
|
||||
adminSess, err := session.NewSession(&aws.Config{
|
||||
Region: aws.String("us-east-1"),
|
||||
Endpoint: aws.String(TestSTSEndpoint),
|
||||
DisableSSL: aws.Bool(true),
|
||||
S3ForcePathStyle: aws.Bool(true),
|
||||
Credentials: credentials.NewStaticCredentials(accessKey, secretKey, ""),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
adminS3 := s3.New(adminSess)
|
||||
bucket := fmt.Sprintf("fed-token-test-%d", time.Now().UnixNano())
|
||||
|
||||
_, err = adminS3.CreateBucket(&s3.CreateBucketInput{Bucket: aws.String(bucket)})
|
||||
require.NoError(t, err)
|
||||
defer adminS3.DeleteBucket(&s3.DeleteBucketInput{Bucket: aws.String(bucket)})
|
||||
|
||||
_, err = adminS3.PutObject(&s3.PutObjectInput{
|
||||
Bucket: aws.String(bucket),
|
||||
Key: aws.String("test.txt"),
|
||||
Body: strings.NewReader("hello"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer adminS3.DeleteObject(&s3.DeleteObjectInput{Bucket: aws.String(bucket), Key: aws.String("test.txt")})
|
||||
|
||||
// Get federated credentials with a session policy that only allows GetObject
|
||||
sessionPolicy := fmt.Sprintf(`{
|
||||
"Version": "2012-10-17",
|
||||
"Statement": [{
|
||||
"Effect": "Allow",
|
||||
"Action": ["s3:GetObject"],
|
||||
"Resource": ["arn:aws:s3:::%s/*"]
|
||||
}]
|
||||
}`, bucket)
|
||||
|
||||
resp, err := callSTSAPIWithSigV4(t, url.Values{
|
||||
"Action": {"GetFederationToken"},
|
||||
"Version": {"2011-06-15"},
|
||||
"Name": {"ScopedClient"},
|
||||
"Policy": {sessionPolicy},
|
||||
}, accessKey, secretKey)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
t.Logf("GetFederationToken response: status=%d body=%s", resp.StatusCode, string(body))
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Skipf("GetFederationToken failed (may need IAM policy config): %s", string(body))
|
||||
}
|
||||
|
||||
var stsResp GetFederationTokenTestResponse
|
||||
require.NoError(t, xml.Unmarshal(body, &stsResp))
|
||||
|
||||
fedCreds := stsResp.Result.Credentials
|
||||
require.NotEmpty(t, fedCreds.AccessKeyId)
|
||||
require.NotEmpty(t, fedCreds.SessionToken)
|
||||
|
||||
// Create S3 client with the federated credentials
|
||||
fedSess, err := session.NewSession(&aws.Config{
|
||||
Region: aws.String("us-east-1"),
|
||||
Endpoint: aws.String(TestSTSEndpoint),
|
||||
DisableSSL: aws.Bool(true),
|
||||
S3ForcePathStyle: aws.Bool(true),
|
||||
Credentials: credentials.NewStaticCredentials(
|
||||
fedCreds.AccessKeyId, fedCreds.SecretAccessKey, fedCreds.SessionToken),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
fedS3 := s3.New(fedSess)
|
||||
|
||||
// GetObject should succeed (allowed by session policy)
|
||||
getResp, err := fedS3.GetObject(&s3.GetObjectInput{
|
||||
Bucket: aws.String(bucket),
|
||||
Key: aws.String("test.txt"),
|
||||
})
|
||||
if err == nil {
|
||||
defer getResp.Body.Close()
|
||||
t.Log("GetObject with federated credentials succeeded (as expected)")
|
||||
} else {
|
||||
t.Logf("GetObject with federated credentials: %v (session policy enforcement may vary)", err)
|
||||
}
|
||||
|
||||
// PutObject should be denied (not allowed by session policy)
|
||||
_, err = fedS3.PutObject(&s3.PutObjectInput{
|
||||
Bucket: aws.String(bucket),
|
||||
Key: aws.String("denied.txt"),
|
||||
Body: strings.NewReader("should fail"),
|
||||
})
|
||||
if err != nil {
|
||||
t.Log("PutObject correctly denied with federated credentials")
|
||||
assert.Contains(t, err.Error(), "AccessDenied",
|
||||
"PutObject should be denied by session policy")
|
||||
} else {
|
||||
// Clean up if unexpectedly succeeded
|
||||
adminS3.DeleteObject(&s3.DeleteObjectInput{Bucket: aws.String(bucket), Key: aws.String("denied.txt")})
|
||||
t.Log("PutObject unexpectedly succeeded — session policy enforcement may not be active")
|
||||
}
|
||||
}
|
||||
@@ -98,8 +98,8 @@ start-seaweedfs: check-binary
|
||||
# Create S3 configuration with SSE-KMS support
|
||||
@printf '{"identities":[{"name":"%s","credentials":[{"accessKey":"%s","secretKey":"%s"}],"actions":["Admin","Read","Write"]}],"kms":{"type":"%s","configs":{"keyId":"%s","encryptionContext":{},"bucketKey":false}}}' "$(ACCESS_KEY)" "$(ACCESS_KEY)" "$(SECRET_KEY)" "$(KMS_TYPE)" "$(KMS_KEY_ID)" > /tmp/seaweedfs-sse-s3.json
|
||||
|
||||
# Start weed mini
|
||||
@AWS_ACCESS_KEY_ID=$(ACCESS_KEY) AWS_SECRET_ACCESS_KEY=$(SECRET_KEY) $(SEAWEEDFS_BINARY) mini \
|
||||
# Start weed mini (WEED_S3_SSE_KEY enables SSE-S3 encryption)
|
||||
@AWS_ACCESS_KEY_ID=$(ACCESS_KEY) AWS_SECRET_ACCESS_KEY=$(SECRET_KEY) WEED_S3_SSE_KEY=test-sse-s3-key $(SEAWEEDFS_BINARY) mini \
|
||||
-dir=/tmp/seaweedfs-test-sse \
|
||||
-s3.port=$(S3_PORT) \
|
||||
-s3.config=/tmp/seaweedfs-sse-s3.json \
|
||||
@@ -337,8 +337,9 @@ start-seaweedfs-ci: check-binary
|
||||
s3-config-template.json > /tmp/seaweedfs-s3.json
|
||||
|
||||
# Start weed mini with embedded S3 using the JSON config (with verbose logging)
|
||||
# WEED_S3_SSE_KEY enables SSE-S3 encryption for testing (KEK derived via HKDF)
|
||||
@echo "Starting weed mini with embedded S3..."
|
||||
@AWS_ACCESS_KEY_ID=$(ACCESS_KEY) AWS_SECRET_ACCESS_KEY=$(SECRET_KEY) GLOG_v=4 $(SEAWEEDFS_BINARY) mini \
|
||||
@AWS_ACCESS_KEY_ID=$(ACCESS_KEY) AWS_SECRET_ACCESS_KEY=$(SECRET_KEY) WEED_S3_SSE_KEY=test-sse-s3-key GLOG_v=4 $(SEAWEEDFS_BINARY) mini \
|
||||
-dir=/tmp/seaweedfs-test-sse \
|
||||
-s3.port=$(S3_PORT) \
|
||||
-s3.config=/tmp/seaweedfs-s3.json \
|
||||
@@ -482,7 +483,7 @@ test-volume-encryption: build-weed
|
||||
-e 's/SECRET_KEY_PLACEHOLDER/$(SECRET_KEY)/g' \
|
||||
s3-config-template.json > /tmp/seaweedfs-s3.json
|
||||
@echo "Starting weed mini with S3 volume encryption..."
|
||||
@AWS_ACCESS_KEY_ID=$(ACCESS_KEY) AWS_SECRET_ACCESS_KEY=$(SECRET_KEY) GLOG_v=4 $(SEAWEEDFS_BINARY) mini \
|
||||
@AWS_ACCESS_KEY_ID=$(ACCESS_KEY) AWS_SECRET_ACCESS_KEY=$(SECRET_KEY) WEED_S3_SSE_KEY=test-sse-s3-key GLOG_v=4 $(SEAWEEDFS_BINARY) mini \
|
||||
-dir=/tmp/seaweedfs-test-sse \
|
||||
-s3.port=$(S3_PORT) \
|
||||
-s3.config=/tmp/seaweedfs-s3.json \
|
||||
|
||||
@@ -289,6 +289,157 @@ func TestVersioningPaginationMultipleObjectsManyVersions(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
// TestVersioningPaginationDeepDirectoryHierarchy tests that paginated ListObjectVersions
|
||||
// correctly skips directory subtrees before the key-marker. This reproduces the
|
||||
// real-world scenario where Veeam backup objects are spread across many subdirectories
|
||||
// (e.g., Mailboxes/<uuid>/ItemsData/<file>) and pagination becomes exponentially
|
||||
// slower as the marker advances through the tree.
|
||||
//
|
||||
// Run with: ENABLE_STRESS_TESTS=true go test -v -run TestVersioningPaginationDeepDirectoryHierarchy -timeout 10m
|
||||
func TestVersioningPaginationDeepDirectoryHierarchy(t *testing.T) {
|
||||
if os.Getenv("ENABLE_STRESS_TESTS") != "true" {
|
||||
t.Skip("Skipping stress test. Set ENABLE_STRESS_TESTS=true to run.")
|
||||
}
|
||||
|
||||
client := getS3Client(t)
|
||||
bucketName := getNewBucketName()
|
||||
|
||||
createBucket(t, client, bucketName)
|
||||
defer deleteBucket(t, client, bucketName)
|
||||
|
||||
enableVersioning(t, client, bucketName)
|
||||
checkVersioningStatus(t, client, bucketName, types.BucketVersioningStatusEnabled)
|
||||
|
||||
// Create a deep directory structure mimicking Veeam 365 backup layout:
|
||||
// Backup/Organizations/<org>/Mailboxes/<mailbox>/ItemsData/<file>
|
||||
numMailboxes := 20
|
||||
filesPerMailbox := 5
|
||||
totalObjects := numMailboxes * filesPerMailbox
|
||||
orgPrefix := "Backup/Organizations/org-001/Mailboxes"
|
||||
|
||||
t.Logf("Creating %d objects across %d subdirectories (depth=6)...", totalObjects, numMailboxes)
|
||||
startTime := time.Now()
|
||||
|
||||
allKeys := make([]string, 0, totalObjects)
|
||||
for i := 0; i < numMailboxes; i++ {
|
||||
mailboxId := fmt.Sprintf("mbx-%03d", i)
|
||||
for j := 0; j < filesPerMailbox; j++ {
|
||||
key := fmt.Sprintf("%s/%s/ItemsData/file-%03d.dat", orgPrefix, mailboxId, j)
|
||||
_, err := client.PutObject(context.TODO(), &s3.PutObjectInput{
|
||||
Bucket: aws.String(bucketName),
|
||||
Key: aws.String(key),
|
||||
Body: strings.NewReader(fmt.Sprintf("content-%d-%d", i, j)),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
allKeys = append(allKeys, key)
|
||||
}
|
||||
}
|
||||
t.Logf("Created %d objects in %v", totalObjects, time.Since(startTime))
|
||||
|
||||
// Test 1: Paginate through all versions with a broad prefix and small maxKeys.
|
||||
// This forces multiple pages that must skip earlier subdirectory trees.
|
||||
t.Run("PaginateAcrossSubdirectories", func(t *testing.T) {
|
||||
maxKeys := int32(10) // Force many pages to exercise marker skipping
|
||||
var allVersions []types.ObjectVersion
|
||||
var keyMarker, versionIdMarker *string
|
||||
pageCount := 0
|
||||
pageStartTimes := make([]time.Duration, 0)
|
||||
|
||||
for {
|
||||
pageStart := time.Now()
|
||||
resp, err := client.ListObjectVersions(context.TODO(), &s3.ListObjectVersionsInput{
|
||||
Bucket: aws.String(bucketName),
|
||||
Prefix: aws.String(orgPrefix + "/"),
|
||||
MaxKeys: aws.Int32(maxKeys),
|
||||
KeyMarker: keyMarker,
|
||||
VersionIdMarker: versionIdMarker,
|
||||
})
|
||||
pageDuration := time.Since(pageStart)
|
||||
require.NoError(t, err)
|
||||
pageCount++
|
||||
pageStartTimes = append(pageStartTimes, pageDuration)
|
||||
|
||||
allVersions = append(allVersions, resp.Versions...)
|
||||
t.Logf("Page %d: %d versions in %v (marker: %v)",
|
||||
pageCount, len(resp.Versions), pageDuration,
|
||||
keyMarker)
|
||||
|
||||
if resp.IsTruncated == nil || !*resp.IsTruncated {
|
||||
break
|
||||
}
|
||||
keyMarker = resp.NextKeyMarker
|
||||
versionIdMarker = resp.NextVersionIdMarker
|
||||
}
|
||||
|
||||
assert.Greater(t, pageCount, 1, "Should require multiple pages")
|
||||
|
||||
// Verify listed keys exactly match the keys we created (same elements, same order)
|
||||
listedKeys := make([]string, 0, len(allVersions))
|
||||
for _, v := range allVersions {
|
||||
listedKeys = append(listedKeys, *v.Key)
|
||||
}
|
||||
assert.Equal(t, allKeys, listedKeys,
|
||||
"Listed version keys should exactly match created keys")
|
||||
|
||||
// Check that later pages don't take dramatically longer than earlier ones.
|
||||
// Before the fix, later pages were exponentially slower because they
|
||||
// re-traversed the entire tree. With the fix, all pages should be similar.
|
||||
if len(pageStartTimes) >= 4 {
|
||||
firstQuarter := pageStartTimes[0]
|
||||
lastQuarter := pageStartTimes[len(pageStartTimes)-1]
|
||||
t.Logf("First page: %v, Last page: %v, Ratio: %.1fx",
|
||||
firstQuarter, lastQuarter,
|
||||
float64(lastQuarter)/float64(firstQuarter))
|
||||
// Allow generous 10x ratio to avoid flakiness; before the fix
|
||||
// the ratio was 100x+ on large datasets
|
||||
if lastQuarter > firstQuarter*10 && lastQuarter > 500*time.Millisecond {
|
||||
t.Errorf("Last page took %.1fx longer than first page (%v vs %v) — possible pagination regression",
|
||||
float64(lastQuarter)/float64(firstQuarter), lastQuarter, firstQuarter)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
// Test 2: Paginate with delimiter to verify CommonPrefixes interaction
|
||||
t.Run("PaginateWithDelimiterAcrossSubdirs", func(t *testing.T) {
|
||||
maxKeys := int32(5)
|
||||
var allPrefixes []string
|
||||
var keyMarker *string
|
||||
pageCount := 0
|
||||
|
||||
for {
|
||||
resp, err := client.ListObjectVersions(context.TODO(), &s3.ListObjectVersionsInput{
|
||||
Bucket: aws.String(bucketName),
|
||||
Prefix: aws.String(orgPrefix + "/"),
|
||||
Delimiter: aws.String("/"),
|
||||
MaxKeys: aws.Int32(maxKeys),
|
||||
KeyMarker: keyMarker,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
pageCount++
|
||||
|
||||
for _, cp := range resp.CommonPrefixes {
|
||||
allPrefixes = append(allPrefixes, *cp.Prefix)
|
||||
}
|
||||
|
||||
if resp.IsTruncated == nil || !*resp.IsTruncated {
|
||||
break
|
||||
}
|
||||
keyMarker = resp.NextKeyMarker
|
||||
}
|
||||
|
||||
assert.Greater(t, pageCount, 1, "Should require multiple pages with maxKeys=%d", maxKeys)
|
||||
|
||||
// Build the exact expected prefixes
|
||||
expectedPrefixes := make([]string, 0, numMailboxes)
|
||||
for i := 0; i < numMailboxes; i++ {
|
||||
expectedPrefixes = append(expectedPrefixes,
|
||||
fmt.Sprintf("%s/mbx-%03d/", orgPrefix, i))
|
||||
}
|
||||
assert.Equal(t, expectedPrefixes, allPrefixes,
|
||||
"CommonPrefixes should exactly match expected mailbox prefixes")
|
||||
})
|
||||
}
|
||||
|
||||
// listAllVersions is a helper to list all versions of a specific object using pagination
|
||||
func listAllVersions(t *testing.T, client *s3.Client, bucketName, objectKey string) []types.ObjectVersion {
|
||||
var allVersions []types.ObjectVersion
|
||||
|
||||
@@ -115,10 +115,13 @@ func NewTestEnvironment(t *testing.T) *TestEnvironment {
|
||||
|
||||
bindIP := testutil.FindBindIP()
|
||||
|
||||
masterPort, masterGrpcPort := testutil.MustFreePortPair(t, "Master")
|
||||
volumePort, volumeGrpcPort := testutil.MustFreePortPair(t, "Volume")
|
||||
filerPort, filerGrpcPort := testutil.MustFreePortPair(t, "Filer")
|
||||
s3Port, s3GrpcPort := testutil.MustFreePortPair(t, "S3")
|
||||
// Allocate all ports in a single batch to prevent the OS from recycling
|
||||
// a released port, which can cause two services to get the same port.
|
||||
ports := testutil.MustAllocatePorts(t, 8)
|
||||
masterPort, masterGrpcPort := ports[0], ports[1]
|
||||
volumePort, volumeGrpcPort := ports[2], ports[3]
|
||||
filerPort, filerGrpcPort := ports[4], ports[5]
|
||||
s3Port, s3GrpcPort := ports[6], ports[7]
|
||||
|
||||
return &TestEnvironment{
|
||||
seaweedDir: seaweedDir,
|
||||
|
||||
@@ -89,11 +89,14 @@ func NewTestEnvironment(t *testing.T) *TestEnvironment {
|
||||
|
||||
bindIP := testutil.FindBindIP()
|
||||
|
||||
masterPort, masterGrpcPort := testutil.MustFreePortPair(t, "Master")
|
||||
volumePort, volumeGrpcPort := testutil.MustFreePortPair(t, "Volume")
|
||||
filerPort, filerGrpcPort := testutil.MustFreePortPair(t, "Filer")
|
||||
s3Port, s3GrpcPort := testutil.MustFreePortPair(t, "S3")
|
||||
polarisPort, polarisAdminPort := testutil.MustFreePortPair(t, "Polaris")
|
||||
// Allocate all ports in a single batch to prevent the OS from recycling
|
||||
// a released port, which can cause two services to get the same port.
|
||||
ports := testutil.MustAllocatePorts(t, 10)
|
||||
masterPort, masterGrpcPort := ports[0], ports[1]
|
||||
volumePort, volumeGrpcPort := ports[2], ports[3]
|
||||
filerPort, filerGrpcPort := ports[4], ports[5]
|
||||
s3Port, s3GrpcPort := ports[6], ports[7]
|
||||
polarisPort, polarisAdminPort := ports[8], ports[9]
|
||||
|
||||
return &TestEnvironment{
|
||||
seaweedDir: seaweedDir,
|
||||
|
||||
@@ -93,10 +93,13 @@ func NewTestEnvironment(t *testing.T) *TestEnvironment {
|
||||
|
||||
bindIP := testutil.FindBindIP()
|
||||
|
||||
masterPort, masterGrpcPort := testutil.MustFreePortPair(t, "Master")
|
||||
volumePort, volumeGrpcPort := testutil.MustFreePortPair(t, "Volume")
|
||||
filerPort, filerGrpcPort := testutil.MustFreePortPair(t, "Filer")
|
||||
s3Port, s3GrpcPort := testutil.MustFreePortPair(t, "S3") // Changed to use testutil.MustFreePortPair
|
||||
// Allocate all ports in a single batch to prevent the OS from recycling
|
||||
// a released port, which can cause two services to get the same port.
|
||||
ports := testutil.MustAllocatePorts(t, 8)
|
||||
masterPort, masterGrpcPort := ports[0], ports[1]
|
||||
volumePort, volumeGrpcPort := ports[2], ports[3]
|
||||
filerPort, filerGrpcPort := ports[4], ports[5]
|
||||
s3Port, s3GrpcPort := ports[6], ports[7]
|
||||
|
||||
return &TestEnvironment{
|
||||
seaweedDir: seaweedDir,
|
||||
|
||||
@@ -15,33 +15,44 @@ func HasDocker() bool {
|
||||
return cmd.Run() == nil
|
||||
}
|
||||
|
||||
// MustFreePortPair is a convenience wrapper for tests that only need a single pair.
|
||||
// Prefer MustAllocatePorts when allocating multiple pairs to guarantee uniqueness.
|
||||
func MustFreePortPair(t *testing.T, name string) (int, int) {
|
||||
httpPort, grpcPort, err := findAvailablePortPair()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get free port pair for %s: %v", name, err)
|
||||
}
|
||||
return httpPort, grpcPort
|
||||
ports := MustAllocatePorts(t, 2)
|
||||
return ports[0], ports[1]
|
||||
}
|
||||
|
||||
func findAvailablePortPair() (int, int, error) {
|
||||
httpPort, err := GetFreePort()
|
||||
// MustAllocatePorts allocates count unique free ports atomically.
|
||||
// All listeners are held open until every port is obtained, preventing
|
||||
// the OS from recycling a port between successive allocations.
|
||||
func MustAllocatePorts(t *testing.T, count int) []int {
|
||||
t.Helper()
|
||||
ports, err := AllocatePorts(count)
|
||||
if err != nil {
|
||||
return 0, 0, err
|
||||
t.Fatalf("Failed to allocate %d free ports: %v", count, err)
|
||||
}
|
||||
grpcPort, err := GetFreePort()
|
||||
if err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
return httpPort, grpcPort, nil
|
||||
return ports
|
||||
}
|
||||
|
||||
func GetFreePort() (int, error) {
|
||||
listener, err := net.Listen("tcp", "0.0.0.0:0")
|
||||
if err != nil {
|
||||
return 0, err
|
||||
// AllocatePorts allocates count unique free ports atomically.
|
||||
func AllocatePorts(count int) ([]int, error) {
|
||||
listeners := make([]net.Listener, 0, count)
|
||||
ports := make([]int, 0, count)
|
||||
for i := 0; i < count; i++ {
|
||||
l, err := net.Listen("tcp", "0.0.0.0:0")
|
||||
if err != nil {
|
||||
for _, ll := range listeners {
|
||||
_ = ll.Close()
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
listeners = append(listeners, l)
|
||||
ports = append(ports, l.Addr().(*net.TCPAddr).Port)
|
||||
}
|
||||
defer listener.Close()
|
||||
return listener.Addr().(*net.TCPAddr).Port, nil
|
||||
for _, l := range listeners {
|
||||
_ = l.Close()
|
||||
}
|
||||
return ports, nil
|
||||
}
|
||||
|
||||
func WaitForService(url string, timeout time.Duration) bool {
|
||||
|
||||
@@ -447,11 +447,6 @@ type QueueStats = maintenance.QueueStats
|
||||
type WorkerDetailsData = maintenance.WorkerDetailsData
|
||||
type WorkerPerformance = maintenance.WorkerPerformance
|
||||
|
||||
// GetTaskIcon returns the icon CSS class for a task type from its UI provider
|
||||
func GetTaskIcon(taskType MaintenanceTaskType) string {
|
||||
return maintenance.GetTaskIcon(taskType)
|
||||
}
|
||||
|
||||
// Status constants (these are still static)
|
||||
const (
|
||||
TaskStatusPending = maintenance.TaskStatusPending
|
||||
|
||||
@@ -3,13 +3,13 @@ package dash
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/credential"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/iam_pb"
|
||||
)
|
||||
|
||||
@@ -435,10 +435,13 @@ func generateAccessKey() string {
|
||||
}
|
||||
|
||||
func generateSecretKey() string {
|
||||
// Generate 40-character secret key (AWS standard)
|
||||
b := make([]byte, 30) // 30 bytes = 40 characters in base64
|
||||
rand.Read(b)
|
||||
return base64.StdEncoding.EncodeToString(b)
|
||||
// Use the IAM helper to generate URL-safe secret keys (no +, / characters)
|
||||
// that won't break S3 signature authentication
|
||||
key, err := iam.GenerateSecretAccessKey()
|
||||
if err != nil {
|
||||
panic(fmt.Sprintf("failed to generate secret key: %v", err))
|
||||
}
|
||||
return key
|
||||
}
|
||||
|
||||
func generateAccountId() string {
|
||||
|
||||
@@ -74,9 +74,16 @@ func TestGenerateSecretKey(t *testing.T) {
|
||||
key1 := generateSecretKey()
|
||||
key2 := generateSecretKey()
|
||||
|
||||
// Check length (base64 encoding of 30 bytes = 40 characters)
|
||||
if len(key1) != 40 {
|
||||
t.Errorf("Expected secret key length 40, got %d", len(key1))
|
||||
// Check length (IAM standard secret key length)
|
||||
if len(key1) != 42 {
|
||||
t.Errorf("Expected secret key length 42, got %d", len(key1))
|
||||
}
|
||||
|
||||
// Check that key contains only URL-safe characters (no +, /)
|
||||
for _, c := range key1 {
|
||||
if c == '+' || c == '/' || c == '=' {
|
||||
t.Errorf("Secret key contains non-URL-safe character: %c", c)
|
||||
}
|
||||
}
|
||||
|
||||
// Check uniqueness
|
||||
|
||||
@@ -312,29 +312,6 @@ func (h *ClusterHandlers) ShowClusterFilers(w http.ResponseWriter, r *http.Reque
|
||||
}
|
||||
}
|
||||
|
||||
// ShowClusterBrokers renders the cluster message brokers page
|
||||
func (h *ClusterHandlers) ShowClusterBrokers(w http.ResponseWriter, r *http.Request) {
|
||||
// Get cluster brokers data
|
||||
brokersData, err := h.adminServer.GetClusterBrokers()
|
||||
if err != nil {
|
||||
writeJSONError(w, http.StatusInternalServerError, "Failed to get cluster brokers: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
username := usernameOrDefault(r)
|
||||
brokersData.Username = username
|
||||
|
||||
// Render HTML template
|
||||
w.Header().Set("Content-Type", "text/html")
|
||||
brokersComponent := app.ClusterBrokers(*brokersData)
|
||||
viewCtx := layout.NewViewContext(r, username, dash.CSRFTokenFromContext(r.Context()))
|
||||
layoutComponent := layout.Layout(viewCtx, brokersComponent)
|
||||
if err := layoutComponent.Render(r.Context(), w); err != nil {
|
||||
writeJSONError(w, http.StatusInternalServerError, "Failed to render template: "+err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// GetClusterTopology returns the cluster topology as JSON
|
||||
func (h *ClusterHandlers) GetClusterTopology(w http.ResponseWriter, r *http.Request) {
|
||||
topology, err := h.adminServer.GetClusterTopology()
|
||||
|
||||
@@ -78,34 +78,6 @@ func (h *MessageQueueHandlers) ShowTopics(w http.ResponseWriter, r *http.Request
|
||||
}
|
||||
}
|
||||
|
||||
// ShowSubscribers renders the message queue subscribers page
|
||||
func (h *MessageQueueHandlers) ShowSubscribers(w http.ResponseWriter, r *http.Request) {
|
||||
// Get subscribers data
|
||||
subscribersData, err := h.adminServer.GetSubscribers()
|
||||
if err != nil {
|
||||
writeJSONError(w, http.StatusInternalServerError, "Failed to get subscribers: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// Set username
|
||||
username := dash.UsernameFromContext(r.Context())
|
||||
if username == "" {
|
||||
username = "admin"
|
||||
}
|
||||
subscribersData.Username = username
|
||||
|
||||
// Render HTML template
|
||||
w.Header().Set("Content-Type", "text/html")
|
||||
subscribersComponent := app.Subscribers(*subscribersData)
|
||||
viewCtx := layout.NewViewContext(r, username, dash.CSRFTokenFromContext(r.Context()))
|
||||
layoutComponent := layout.Layout(viewCtx, subscribersComponent)
|
||||
err = layoutComponent.Render(r.Context(), w)
|
||||
if err != nil {
|
||||
writeJSONError(w, http.StatusInternalServerError, "Failed to render template: "+err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// ShowTopicDetails renders the topic details page
|
||||
func (h *MessageQueueHandlers) ShowTopicDetails(w http.ResponseWriter, r *http.Request) {
|
||||
// Get topic parameters from URL
|
||||
|
||||
@@ -1,124 +0,0 @@
|
||||
package maintenance
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/worker_pb"
|
||||
)
|
||||
|
||||
// VerifyProtobufConfig demonstrates that the protobuf configuration system is working
|
||||
func VerifyProtobufConfig() error {
|
||||
// Create configuration manager
|
||||
configManager := NewMaintenanceConfigManager()
|
||||
config := configManager.GetConfig()
|
||||
|
||||
// Verify basic configuration
|
||||
if !config.Enabled {
|
||||
return fmt.Errorf("expected config to be enabled by default")
|
||||
}
|
||||
|
||||
if config.ScanIntervalSeconds != 30*60 {
|
||||
return fmt.Errorf("expected scan interval to be 1800 seconds, got %d", config.ScanIntervalSeconds)
|
||||
}
|
||||
|
||||
// Verify policy configuration
|
||||
if config.Policy == nil {
|
||||
return fmt.Errorf("expected policy to be configured")
|
||||
}
|
||||
|
||||
if config.Policy.GlobalMaxConcurrent != 4 {
|
||||
return fmt.Errorf("expected global max concurrent to be 4, got %d", config.Policy.GlobalMaxConcurrent)
|
||||
}
|
||||
|
||||
// Verify task policies
|
||||
vacuumPolicy := config.Policy.TaskPolicies["vacuum"]
|
||||
if vacuumPolicy == nil {
|
||||
return fmt.Errorf("expected vacuum policy to be configured")
|
||||
}
|
||||
|
||||
if !vacuumPolicy.Enabled {
|
||||
return fmt.Errorf("expected vacuum policy to be enabled")
|
||||
}
|
||||
|
||||
// Verify typed configuration access
|
||||
vacuumConfig := vacuumPolicy.GetVacuumConfig()
|
||||
if vacuumConfig == nil {
|
||||
return fmt.Errorf("expected vacuum config to be accessible")
|
||||
}
|
||||
|
||||
if vacuumConfig.GarbageThreshold != 0.3 {
|
||||
return fmt.Errorf("expected garbage threshold to be 0.3, got %f", vacuumConfig.GarbageThreshold)
|
||||
}
|
||||
|
||||
// Verify helper functions work
|
||||
if !IsTaskEnabled(config.Policy, "vacuum") {
|
||||
return fmt.Errorf("expected vacuum task to be enabled via helper function")
|
||||
}
|
||||
|
||||
maxConcurrent := GetMaxConcurrent(config.Policy, "vacuum")
|
||||
if maxConcurrent != 2 {
|
||||
return fmt.Errorf("expected vacuum max concurrent to be 2, got %d", maxConcurrent)
|
||||
}
|
||||
|
||||
// Verify erasure coding configuration
|
||||
ecPolicy := config.Policy.TaskPolicies["erasure_coding"]
|
||||
if ecPolicy == nil {
|
||||
return fmt.Errorf("expected EC policy to be configured")
|
||||
}
|
||||
|
||||
ecConfig := ecPolicy.GetErasureCodingConfig()
|
||||
if ecConfig == nil {
|
||||
return fmt.Errorf("expected EC config to be accessible")
|
||||
}
|
||||
|
||||
// Verify configurable EC fields only
|
||||
if ecConfig.FullnessRatio <= 0 || ecConfig.FullnessRatio > 1 {
|
||||
return fmt.Errorf("expected EC config to have valid fullness ratio (0-1), got %f", ecConfig.FullnessRatio)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetProtobufConfigSummary returns a summary of the current protobuf configuration
|
||||
func GetProtobufConfigSummary() string {
|
||||
configManager := NewMaintenanceConfigManager()
|
||||
config := configManager.GetConfig()
|
||||
|
||||
summary := fmt.Sprintf("SeaweedFS Protobuf Maintenance Configuration:\n")
|
||||
summary += fmt.Sprintf(" Enabled: %v\n", config.Enabled)
|
||||
summary += fmt.Sprintf(" Scan Interval: %d seconds\n", config.ScanIntervalSeconds)
|
||||
summary += fmt.Sprintf(" Max Retries: %d\n", config.MaxRetries)
|
||||
summary += fmt.Sprintf(" Global Max Concurrent: %d\n", config.Policy.GlobalMaxConcurrent)
|
||||
summary += fmt.Sprintf(" Task Policies: %d configured\n", len(config.Policy.TaskPolicies))
|
||||
|
||||
for taskType, policy := range config.Policy.TaskPolicies {
|
||||
summary += fmt.Sprintf(" %s: enabled=%v, max_concurrent=%d\n",
|
||||
taskType, policy.Enabled, policy.MaxConcurrent)
|
||||
}
|
||||
|
||||
return summary
|
||||
}
|
||||
|
||||
// CreateCustomConfig demonstrates creating a custom protobuf configuration
|
||||
func CreateCustomConfig() *worker_pb.MaintenanceConfig {
|
||||
return &worker_pb.MaintenanceConfig{
|
||||
Enabled: true,
|
||||
ScanIntervalSeconds: 60 * 60, // 1 hour
|
||||
MaxRetries: 5,
|
||||
Policy: &worker_pb.MaintenancePolicy{
|
||||
GlobalMaxConcurrent: 8,
|
||||
TaskPolicies: map[string]*worker_pb.TaskPolicy{
|
||||
"custom_vacuum": {
|
||||
Enabled: true,
|
||||
MaxConcurrent: 4,
|
||||
TaskConfig: &worker_pb.TaskPolicy_VacuumConfig{
|
||||
VacuumConfig: &worker_pb.VacuumTaskConfig{
|
||||
GarbageThreshold: 0.5,
|
||||
MinVolumeAgeHours: 48,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -1,24 +1,9 @@
|
||||
package maintenance
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/worker_pb"
|
||||
)
|
||||
|
||||
// MaintenanceConfigManager handles protobuf-based configuration
|
||||
type MaintenanceConfigManager struct {
|
||||
config *worker_pb.MaintenanceConfig
|
||||
}
|
||||
|
||||
// NewMaintenanceConfigManager creates a new config manager with defaults
|
||||
func NewMaintenanceConfigManager() *MaintenanceConfigManager {
|
||||
return &MaintenanceConfigManager{
|
||||
config: DefaultMaintenanceConfigProto(),
|
||||
}
|
||||
}
|
||||
|
||||
// DefaultMaintenanceConfigProto returns default configuration as protobuf
|
||||
func DefaultMaintenanceConfigProto() *worker_pb.MaintenanceConfig {
|
||||
return &worker_pb.MaintenanceConfig{
|
||||
@@ -34,253 +19,3 @@ func DefaultMaintenanceConfigProto() *worker_pb.MaintenanceConfig {
|
||||
Policy: nil,
|
||||
}
|
||||
}
|
||||
|
||||
// GetConfig returns the current configuration
|
||||
func (mcm *MaintenanceConfigManager) GetConfig() *worker_pb.MaintenanceConfig {
|
||||
return mcm.config
|
||||
}
|
||||
|
||||
// Type-safe configuration accessors
|
||||
|
||||
// GetVacuumConfig returns vacuum-specific configuration for a task type
|
||||
func (mcm *MaintenanceConfigManager) GetVacuumConfig(taskType string) *worker_pb.VacuumTaskConfig {
|
||||
if policy := mcm.getTaskPolicy(taskType); policy != nil {
|
||||
if vacuumConfig := policy.GetVacuumConfig(); vacuumConfig != nil {
|
||||
return vacuumConfig
|
||||
}
|
||||
}
|
||||
// Return defaults if not configured
|
||||
return &worker_pb.VacuumTaskConfig{
|
||||
GarbageThreshold: 0.3,
|
||||
MinVolumeAgeHours: 24,
|
||||
}
|
||||
}
|
||||
|
||||
// GetErasureCodingConfig returns EC-specific configuration for a task type
|
||||
func (mcm *MaintenanceConfigManager) GetErasureCodingConfig(taskType string) *worker_pb.ErasureCodingTaskConfig {
|
||||
if policy := mcm.getTaskPolicy(taskType); policy != nil {
|
||||
if ecConfig := policy.GetErasureCodingConfig(); ecConfig != nil {
|
||||
return ecConfig
|
||||
}
|
||||
}
|
||||
// Return defaults if not configured
|
||||
return &worker_pb.ErasureCodingTaskConfig{
|
||||
FullnessRatio: 0.95,
|
||||
QuietForSeconds: 3600,
|
||||
MinVolumeSizeMb: 100,
|
||||
CollectionFilter: "",
|
||||
}
|
||||
}
|
||||
|
||||
// GetBalanceConfig returns balance-specific configuration for a task type
|
||||
func (mcm *MaintenanceConfigManager) GetBalanceConfig(taskType string) *worker_pb.BalanceTaskConfig {
|
||||
if policy := mcm.getTaskPolicy(taskType); policy != nil {
|
||||
if balanceConfig := policy.GetBalanceConfig(); balanceConfig != nil {
|
||||
return balanceConfig
|
||||
}
|
||||
}
|
||||
// Return defaults if not configured
|
||||
return &worker_pb.BalanceTaskConfig{
|
||||
ImbalanceThreshold: 0.2,
|
||||
MinServerCount: 2,
|
||||
}
|
||||
}
|
||||
|
||||
// GetReplicationConfig returns replication-specific configuration for a task type
|
||||
func (mcm *MaintenanceConfigManager) GetReplicationConfig(taskType string) *worker_pb.ReplicationTaskConfig {
|
||||
if policy := mcm.getTaskPolicy(taskType); policy != nil {
|
||||
if replicationConfig := policy.GetReplicationConfig(); replicationConfig != nil {
|
||||
return replicationConfig
|
||||
}
|
||||
}
|
||||
// Return defaults if not configured
|
||||
return &worker_pb.ReplicationTaskConfig{
|
||||
TargetReplicaCount: 2,
|
||||
}
|
||||
}
|
||||
|
||||
// Typed convenience methods for getting task configurations
|
||||
|
||||
// GetVacuumTaskConfigForType returns vacuum configuration for a specific task type
|
||||
func (mcm *MaintenanceConfigManager) GetVacuumTaskConfigForType(taskType string) *worker_pb.VacuumTaskConfig {
|
||||
return GetVacuumTaskConfig(mcm.config.Policy, MaintenanceTaskType(taskType))
|
||||
}
|
||||
|
||||
// GetErasureCodingTaskConfigForType returns erasure coding configuration for a specific task type
|
||||
func (mcm *MaintenanceConfigManager) GetErasureCodingTaskConfigForType(taskType string) *worker_pb.ErasureCodingTaskConfig {
|
||||
return GetErasureCodingTaskConfig(mcm.config.Policy, MaintenanceTaskType(taskType))
|
||||
}
|
||||
|
||||
// GetBalanceTaskConfigForType returns balance configuration for a specific task type
|
||||
func (mcm *MaintenanceConfigManager) GetBalanceTaskConfigForType(taskType string) *worker_pb.BalanceTaskConfig {
|
||||
return GetBalanceTaskConfig(mcm.config.Policy, MaintenanceTaskType(taskType))
|
||||
}
|
||||
|
||||
// GetReplicationTaskConfigForType returns replication configuration for a specific task type
|
||||
func (mcm *MaintenanceConfigManager) GetReplicationTaskConfigForType(taskType string) *worker_pb.ReplicationTaskConfig {
|
||||
return GetReplicationTaskConfig(mcm.config.Policy, MaintenanceTaskType(taskType))
|
||||
}
|
||||
|
||||
// Helper methods
|
||||
|
||||
func (mcm *MaintenanceConfigManager) getTaskPolicy(taskType string) *worker_pb.TaskPolicy {
|
||||
if mcm.config.Policy != nil && mcm.config.Policy.TaskPolicies != nil {
|
||||
return mcm.config.Policy.TaskPolicies[taskType]
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// IsTaskEnabled returns whether a task type is enabled
|
||||
func (mcm *MaintenanceConfigManager) IsTaskEnabled(taskType string) bool {
|
||||
if policy := mcm.getTaskPolicy(taskType); policy != nil {
|
||||
return policy.Enabled
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// GetMaxConcurrent returns the max concurrent limit for a task type
|
||||
func (mcm *MaintenanceConfigManager) GetMaxConcurrent(taskType string) int32 {
|
||||
if policy := mcm.getTaskPolicy(taskType); policy != nil {
|
||||
return policy.MaxConcurrent
|
||||
}
|
||||
return 1 // Default
|
||||
}
|
||||
|
||||
// GetRepeatInterval returns the repeat interval for a task type in seconds
|
||||
func (mcm *MaintenanceConfigManager) GetRepeatInterval(taskType string) int32 {
|
||||
if policy := mcm.getTaskPolicy(taskType); policy != nil {
|
||||
return policy.RepeatIntervalSeconds
|
||||
}
|
||||
return mcm.config.Policy.DefaultRepeatIntervalSeconds
|
||||
}
|
||||
|
||||
// GetCheckInterval returns the check interval for a task type in seconds
|
||||
func (mcm *MaintenanceConfigManager) GetCheckInterval(taskType string) int32 {
|
||||
if policy := mcm.getTaskPolicy(taskType); policy != nil {
|
||||
return policy.CheckIntervalSeconds
|
||||
}
|
||||
return mcm.config.Policy.DefaultCheckIntervalSeconds
|
||||
}
|
||||
|
||||
// Duration accessor methods
|
||||
|
||||
// GetScanInterval returns the scan interval as a time.Duration
|
||||
func (mcm *MaintenanceConfigManager) GetScanInterval() time.Duration {
|
||||
return time.Duration(mcm.config.ScanIntervalSeconds) * time.Second
|
||||
}
|
||||
|
||||
// GetWorkerTimeout returns the worker timeout as a time.Duration
|
||||
func (mcm *MaintenanceConfigManager) GetWorkerTimeout() time.Duration {
|
||||
return time.Duration(mcm.config.WorkerTimeoutSeconds) * time.Second
|
||||
}
|
||||
|
||||
// GetTaskTimeout returns the task timeout as a time.Duration
|
||||
func (mcm *MaintenanceConfigManager) GetTaskTimeout() time.Duration {
|
||||
return time.Duration(mcm.config.TaskTimeoutSeconds) * time.Second
|
||||
}
|
||||
|
||||
// GetRetryDelay returns the retry delay as a time.Duration
|
||||
func (mcm *MaintenanceConfigManager) GetRetryDelay() time.Duration {
|
||||
return time.Duration(mcm.config.RetryDelaySeconds) * time.Second
|
||||
}
|
||||
|
||||
// GetCleanupInterval returns the cleanup interval as a time.Duration
|
||||
func (mcm *MaintenanceConfigManager) GetCleanupInterval() time.Duration {
|
||||
return time.Duration(mcm.config.CleanupIntervalSeconds) * time.Second
|
||||
}
|
||||
|
||||
// GetTaskRetention returns the task retention period as a time.Duration
|
||||
func (mcm *MaintenanceConfigManager) GetTaskRetention() time.Duration {
|
||||
return time.Duration(mcm.config.TaskRetentionSeconds) * time.Second
|
||||
}
|
||||
|
||||
// ValidateMaintenanceConfigWithSchema validates protobuf maintenance configuration using ConfigField rules
|
||||
func ValidateMaintenanceConfigWithSchema(config *worker_pb.MaintenanceConfig) error {
|
||||
if config == nil {
|
||||
return fmt.Errorf("configuration cannot be nil")
|
||||
}
|
||||
|
||||
// Get the schema to access field validation rules
|
||||
schema := GetMaintenanceConfigSchema()
|
||||
|
||||
// Validate each field individually using the ConfigField rules
|
||||
if err := validateFieldWithSchema(schema, "enabled", config.Enabled); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := validateFieldWithSchema(schema, "scan_interval_seconds", int(config.ScanIntervalSeconds)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := validateFieldWithSchema(schema, "worker_timeout_seconds", int(config.WorkerTimeoutSeconds)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := validateFieldWithSchema(schema, "task_timeout_seconds", int(config.TaskTimeoutSeconds)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := validateFieldWithSchema(schema, "retry_delay_seconds", int(config.RetryDelaySeconds)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := validateFieldWithSchema(schema, "max_retries", int(config.MaxRetries)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := validateFieldWithSchema(schema, "cleanup_interval_seconds", int(config.CleanupIntervalSeconds)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := validateFieldWithSchema(schema, "task_retention_seconds", int(config.TaskRetentionSeconds)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Validate policy fields if present
|
||||
if config.Policy != nil {
|
||||
// Note: These field names might need to be adjusted based on the actual schema
|
||||
if err := validatePolicyField("global_max_concurrent", int(config.Policy.GlobalMaxConcurrent)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := validatePolicyField("default_repeat_interval_seconds", int(config.Policy.DefaultRepeatIntervalSeconds)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := validatePolicyField("default_check_interval_seconds", int(config.Policy.DefaultCheckIntervalSeconds)); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateFieldWithSchema validates a single field using its ConfigField definition
|
||||
func validateFieldWithSchema(schema *MaintenanceConfigSchema, fieldName string, value interface{}) error {
|
||||
field := schema.GetFieldByName(fieldName)
|
||||
if field == nil {
|
||||
// Field not in schema, skip validation
|
||||
return nil
|
||||
}
|
||||
|
||||
return field.ValidateValue(value)
|
||||
}
|
||||
|
||||
// validatePolicyField validates policy fields (simplified validation for now)
|
||||
func validatePolicyField(fieldName string, value int) error {
|
||||
switch fieldName {
|
||||
case "global_max_concurrent":
|
||||
if value < 1 || value > 20 {
|
||||
return fmt.Errorf("Global Max Concurrent must be between 1 and 20, got %d", value)
|
||||
}
|
||||
case "default_repeat_interval":
|
||||
if value < 1 || value > 168 {
|
||||
return fmt.Errorf("Default Repeat Interval must be between 1 and 168 hours, got %d", value)
|
||||
}
|
||||
case "default_check_interval":
|
||||
if value < 1 || value > 168 {
|
||||
return fmt.Errorf("Default Check Interval must be between 1 and 168 hours, got %d", value)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1055,28 +1055,6 @@ func (mq *MaintenanceQueue) getMaxConcurrentForTaskType(taskType MaintenanceTask
|
||||
return 1
|
||||
}
|
||||
|
||||
// getRunningTasks returns all currently running tasks
|
||||
func (mq *MaintenanceQueue) getRunningTasks() []*MaintenanceTask {
|
||||
var runningTasks []*MaintenanceTask
|
||||
for _, task := range mq.tasks {
|
||||
if task.Status == TaskStatusAssigned || task.Status == TaskStatusInProgress {
|
||||
runningTasks = append(runningTasks, task)
|
||||
}
|
||||
}
|
||||
return runningTasks
|
||||
}
|
||||
|
||||
// getAvailableWorkers returns all workers that can take more work
|
||||
func (mq *MaintenanceQueue) getAvailableWorkers() []*MaintenanceWorker {
|
||||
var availableWorkers []*MaintenanceWorker
|
||||
for _, worker := range mq.workers {
|
||||
if worker.Status == "active" && worker.CurrentLoad < worker.MaxConcurrent {
|
||||
availableWorkers = append(availableWorkers, worker)
|
||||
}
|
||||
}
|
||||
return availableWorkers
|
||||
}
|
||||
|
||||
// trackPendingOperation adds a task to the pending operations tracker
|
||||
func (mq *MaintenanceQueue) trackPendingOperation(task *MaintenanceTask) {
|
||||
if mq.integration == nil {
|
||||
|
||||
@@ -2,15 +2,11 @@ package maintenance
|
||||
|
||||
import (
|
||||
"html/template"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/glog"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/master_pb"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/worker_pb"
|
||||
"github.com/seaweedfs/seaweedfs/weed/worker/tasks"
|
||||
"github.com/seaweedfs/seaweedfs/weed/worker/types"
|
||||
)
|
||||
|
||||
// AdminClient interface defines what the maintenance system needs from the admin server
|
||||
@@ -21,51 +17,6 @@ type AdminClient interface {
|
||||
// MaintenanceTaskType represents different types of maintenance operations
|
||||
type MaintenanceTaskType string
|
||||
|
||||
// GetRegisteredMaintenanceTaskTypes returns all registered task types as MaintenanceTaskType values
|
||||
// sorted alphabetically for consistent menu ordering
|
||||
func GetRegisteredMaintenanceTaskTypes() []MaintenanceTaskType {
|
||||
typesRegistry := tasks.GetGlobalTypesRegistry()
|
||||
var taskTypes []MaintenanceTaskType
|
||||
|
||||
for workerTaskType := range typesRegistry.GetAllDetectors() {
|
||||
maintenanceTaskType := MaintenanceTaskType(string(workerTaskType))
|
||||
taskTypes = append(taskTypes, maintenanceTaskType)
|
||||
}
|
||||
|
||||
// Sort task types alphabetically to ensure consistent menu ordering
|
||||
sort.Slice(taskTypes, func(i, j int) bool {
|
||||
return string(taskTypes[i]) < string(taskTypes[j])
|
||||
})
|
||||
|
||||
return taskTypes
|
||||
}
|
||||
|
||||
// GetMaintenanceTaskType returns a specific task type if it's registered, or empty string if not found
|
||||
func GetMaintenanceTaskType(taskTypeName string) MaintenanceTaskType {
|
||||
typesRegistry := tasks.GetGlobalTypesRegistry()
|
||||
|
||||
for workerTaskType := range typesRegistry.GetAllDetectors() {
|
||||
if string(workerTaskType) == taskTypeName {
|
||||
return MaintenanceTaskType(taskTypeName)
|
||||
}
|
||||
}
|
||||
|
||||
return MaintenanceTaskType("")
|
||||
}
|
||||
|
||||
// IsMaintenanceTaskTypeRegistered checks if a task type is registered
|
||||
func IsMaintenanceTaskTypeRegistered(taskType MaintenanceTaskType) bool {
|
||||
typesRegistry := tasks.GetGlobalTypesRegistry()
|
||||
|
||||
for workerTaskType := range typesRegistry.GetAllDetectors() {
|
||||
if string(workerTaskType) == string(taskType) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// MaintenanceTaskPriority represents task execution priority
|
||||
type MaintenanceTaskPriority int
|
||||
|
||||
@@ -200,14 +151,6 @@ func GetTaskPolicy(mp *MaintenancePolicy, taskType MaintenanceTaskType) *TaskPol
|
||||
return mp.TaskPolicies[string(taskType)]
|
||||
}
|
||||
|
||||
// SetTaskPolicy sets the policy for a specific task type
|
||||
func SetTaskPolicy(mp *MaintenancePolicy, taskType MaintenanceTaskType, policy *TaskPolicy) {
|
||||
if mp.TaskPolicies == nil {
|
||||
mp.TaskPolicies = make(map[string]*TaskPolicy)
|
||||
}
|
||||
mp.TaskPolicies[string(taskType)] = policy
|
||||
}
|
||||
|
||||
// IsTaskEnabled returns whether a task type is enabled
|
||||
func IsTaskEnabled(mp *MaintenancePolicy, taskType MaintenanceTaskType) bool {
|
||||
policy := GetTaskPolicy(mp, taskType)
|
||||
@@ -235,84 +178,6 @@ func GetRepeatInterval(mp *MaintenancePolicy, taskType MaintenanceTaskType) int
|
||||
return int(policy.RepeatIntervalSeconds)
|
||||
}
|
||||
|
||||
// GetVacuumTaskConfig returns the vacuum task configuration
|
||||
func GetVacuumTaskConfig(mp *MaintenancePolicy, taskType MaintenanceTaskType) *worker_pb.VacuumTaskConfig {
|
||||
policy := GetTaskPolicy(mp, taskType)
|
||||
if policy == nil {
|
||||
return nil
|
||||
}
|
||||
return policy.GetVacuumConfig()
|
||||
}
|
||||
|
||||
// GetErasureCodingTaskConfig returns the erasure coding task configuration
|
||||
func GetErasureCodingTaskConfig(mp *MaintenancePolicy, taskType MaintenanceTaskType) *worker_pb.ErasureCodingTaskConfig {
|
||||
policy := GetTaskPolicy(mp, taskType)
|
||||
if policy == nil {
|
||||
return nil
|
||||
}
|
||||
return policy.GetErasureCodingConfig()
|
||||
}
|
||||
|
||||
// GetBalanceTaskConfig returns the balance task configuration
|
||||
func GetBalanceTaskConfig(mp *MaintenancePolicy, taskType MaintenanceTaskType) *worker_pb.BalanceTaskConfig {
|
||||
policy := GetTaskPolicy(mp, taskType)
|
||||
if policy == nil {
|
||||
return nil
|
||||
}
|
||||
return policy.GetBalanceConfig()
|
||||
}
|
||||
|
||||
// GetReplicationTaskConfig returns the replication task configuration
|
||||
func GetReplicationTaskConfig(mp *MaintenancePolicy, taskType MaintenanceTaskType) *worker_pb.ReplicationTaskConfig {
|
||||
policy := GetTaskPolicy(mp, taskType)
|
||||
if policy == nil {
|
||||
return nil
|
||||
}
|
||||
return policy.GetReplicationConfig()
|
||||
}
|
||||
|
||||
// Note: GetTaskConfig was removed - use typed getters: GetVacuumTaskConfig, GetErasureCodingTaskConfig, GetBalanceTaskConfig, or GetReplicationTaskConfig
|
||||
|
||||
// SetVacuumTaskConfig sets the vacuum task configuration
|
||||
func SetVacuumTaskConfig(mp *MaintenancePolicy, taskType MaintenanceTaskType, config *worker_pb.VacuumTaskConfig) {
|
||||
policy := GetTaskPolicy(mp, taskType)
|
||||
if policy != nil {
|
||||
policy.TaskConfig = &worker_pb.TaskPolicy_VacuumConfig{
|
||||
VacuumConfig: config,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// SetErasureCodingTaskConfig sets the erasure coding task configuration
|
||||
func SetErasureCodingTaskConfig(mp *MaintenancePolicy, taskType MaintenanceTaskType, config *worker_pb.ErasureCodingTaskConfig) {
|
||||
policy := GetTaskPolicy(mp, taskType)
|
||||
if policy != nil {
|
||||
policy.TaskConfig = &worker_pb.TaskPolicy_ErasureCodingConfig{
|
||||
ErasureCodingConfig: config,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// SetBalanceTaskConfig sets the balance task configuration
|
||||
func SetBalanceTaskConfig(mp *MaintenancePolicy, taskType MaintenanceTaskType, config *worker_pb.BalanceTaskConfig) {
|
||||
policy := GetTaskPolicy(mp, taskType)
|
||||
if policy != nil {
|
||||
policy.TaskConfig = &worker_pb.TaskPolicy_BalanceConfig{
|
||||
BalanceConfig: config,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// SetReplicationTaskConfig sets the replication task configuration
|
||||
func SetReplicationTaskConfig(mp *MaintenancePolicy, taskType MaintenanceTaskType, config *worker_pb.ReplicationTaskConfig) {
|
||||
policy := GetTaskPolicy(mp, taskType)
|
||||
if policy != nil {
|
||||
policy.TaskConfig = &worker_pb.TaskPolicy_ReplicationConfig{
|
||||
ReplicationConfig: config,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// SetTaskConfig sets a configuration value for a task type (legacy method - use typed setters above)
|
||||
// Note: SetTaskConfig was removed - use typed setters: SetVacuumTaskConfig, SetErasureCodingTaskConfig, SetBalanceTaskConfig, or SetReplicationTaskConfig
|
||||
|
||||
@@ -475,180 +340,6 @@ type ClusterReplicationTask struct {
|
||||
Metadata map[string]string `json:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
// BuildMaintenancePolicyFromTasks creates a maintenance policy with configurations
|
||||
// from all registered tasks using their UI providers
|
||||
func BuildMaintenancePolicyFromTasks() *MaintenancePolicy {
|
||||
policy := &MaintenancePolicy{
|
||||
TaskPolicies: make(map[string]*TaskPolicy),
|
||||
GlobalMaxConcurrent: 4,
|
||||
DefaultRepeatIntervalSeconds: 6 * 3600, // 6 hours in seconds
|
||||
DefaultCheckIntervalSeconds: 12 * 3600, // 12 hours in seconds
|
||||
}
|
||||
|
||||
// Get all registered task types from the UI registry
|
||||
uiRegistry := tasks.GetGlobalUIRegistry()
|
||||
typesRegistry := tasks.GetGlobalTypesRegistry()
|
||||
|
||||
for taskType, provider := range uiRegistry.GetAllProviders() {
|
||||
// Convert task type to maintenance task type
|
||||
maintenanceTaskType := MaintenanceTaskType(string(taskType))
|
||||
|
||||
// Get the default configuration from the UI provider
|
||||
defaultConfig := provider.GetCurrentConfig()
|
||||
|
||||
// Create task policy from UI configuration
|
||||
taskPolicy := &TaskPolicy{
|
||||
Enabled: true, // Default enabled
|
||||
MaxConcurrent: 2, // Default concurrency
|
||||
RepeatIntervalSeconds: policy.DefaultRepeatIntervalSeconds,
|
||||
CheckIntervalSeconds: policy.DefaultCheckIntervalSeconds,
|
||||
}
|
||||
|
||||
// Extract configuration using TaskConfig interface - no more map conversions!
|
||||
if taskConfig, ok := defaultConfig.(interface{ ToTaskPolicy() *worker_pb.TaskPolicy }); ok {
|
||||
// Use protobuf directly for clean, type-safe config extraction
|
||||
pbTaskPolicy := taskConfig.ToTaskPolicy()
|
||||
taskPolicy.Enabled = pbTaskPolicy.Enabled
|
||||
taskPolicy.MaxConcurrent = pbTaskPolicy.MaxConcurrent
|
||||
if pbTaskPolicy.RepeatIntervalSeconds > 0 {
|
||||
taskPolicy.RepeatIntervalSeconds = pbTaskPolicy.RepeatIntervalSeconds
|
||||
}
|
||||
if pbTaskPolicy.CheckIntervalSeconds > 0 {
|
||||
taskPolicy.CheckIntervalSeconds = pbTaskPolicy.CheckIntervalSeconds
|
||||
}
|
||||
}
|
||||
|
||||
// Also get defaults from scheduler if available (using types.TaskScheduler explicitly)
|
||||
var scheduler types.TaskScheduler = typesRegistry.GetScheduler(taskType)
|
||||
if scheduler != nil {
|
||||
if taskPolicy.MaxConcurrent <= 0 {
|
||||
taskPolicy.MaxConcurrent = int32(scheduler.GetMaxConcurrent())
|
||||
}
|
||||
// Convert default repeat interval to seconds
|
||||
if repeatInterval := scheduler.GetDefaultRepeatInterval(); repeatInterval > 0 {
|
||||
taskPolicy.RepeatIntervalSeconds = int32(repeatInterval.Seconds())
|
||||
}
|
||||
}
|
||||
|
||||
// Also get defaults from detector if available (using types.TaskDetector explicitly)
|
||||
var detector types.TaskDetector = typesRegistry.GetDetector(taskType)
|
||||
if detector != nil {
|
||||
// Convert scan interval to check interval (seconds)
|
||||
if scanInterval := detector.ScanInterval(); scanInterval > 0 {
|
||||
taskPolicy.CheckIntervalSeconds = int32(scanInterval.Seconds())
|
||||
}
|
||||
}
|
||||
|
||||
policy.TaskPolicies[string(maintenanceTaskType)] = taskPolicy
|
||||
glog.V(3).Infof("Built policy for task type %s: enabled=%v, max_concurrent=%d",
|
||||
maintenanceTaskType, taskPolicy.Enabled, taskPolicy.MaxConcurrent)
|
||||
}
|
||||
|
||||
glog.V(2).Infof("Built maintenance policy with %d task configurations", len(policy.TaskPolicies))
|
||||
return policy
|
||||
}
|
||||
|
||||
// SetPolicyFromTasks sets the maintenance policy from registered tasks
|
||||
func SetPolicyFromTasks(policy *MaintenancePolicy) {
|
||||
if policy == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// Build new policy from tasks
|
||||
newPolicy := BuildMaintenancePolicyFromTasks()
|
||||
|
||||
// Copy task policies
|
||||
policy.TaskPolicies = newPolicy.TaskPolicies
|
||||
|
||||
glog.V(1).Infof("Updated maintenance policy with %d task configurations from registered tasks", len(policy.TaskPolicies))
|
||||
}
|
||||
|
||||
// GetTaskIcon returns the icon CSS class for a task type from its UI provider
|
||||
func GetTaskIcon(taskType MaintenanceTaskType) string {
|
||||
typesRegistry := tasks.GetGlobalTypesRegistry()
|
||||
uiRegistry := tasks.GetGlobalUIRegistry()
|
||||
|
||||
// Convert MaintenanceTaskType to TaskType
|
||||
for workerTaskType := range typesRegistry.GetAllDetectors() {
|
||||
if string(workerTaskType) == string(taskType) {
|
||||
// Get the UI provider for this task type
|
||||
provider := uiRegistry.GetProvider(workerTaskType)
|
||||
if provider != nil {
|
||||
return provider.GetIcon()
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// Default icon if no UI provider found
|
||||
return "fas fa-cog text-muted"
|
||||
}
|
||||
|
||||
// GetTaskDisplayName returns the display name for a task type from its UI provider
|
||||
func GetTaskDisplayName(taskType MaintenanceTaskType) string {
|
||||
typesRegistry := tasks.GetGlobalTypesRegistry()
|
||||
uiRegistry := tasks.GetGlobalUIRegistry()
|
||||
|
||||
// Convert MaintenanceTaskType to TaskType
|
||||
for workerTaskType := range typesRegistry.GetAllDetectors() {
|
||||
if string(workerTaskType) == string(taskType) {
|
||||
// Get the UI provider for this task type
|
||||
provider := uiRegistry.GetProvider(workerTaskType)
|
||||
if provider != nil {
|
||||
return provider.GetDisplayName()
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback to the task type string
|
||||
return string(taskType)
|
||||
}
|
||||
|
||||
// GetTaskDescription returns the description for a task type from its UI provider
|
||||
func GetTaskDescription(taskType MaintenanceTaskType) string {
|
||||
typesRegistry := tasks.GetGlobalTypesRegistry()
|
||||
uiRegistry := tasks.GetGlobalUIRegistry()
|
||||
|
||||
// Convert MaintenanceTaskType to TaskType
|
||||
for workerTaskType := range typesRegistry.GetAllDetectors() {
|
||||
if string(workerTaskType) == string(taskType) {
|
||||
// Get the UI provider for this task type
|
||||
provider := uiRegistry.GetProvider(workerTaskType)
|
||||
if provider != nil {
|
||||
return provider.GetDescription()
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback to a generic description
|
||||
return "Configure detailed settings for " + string(taskType) + " tasks."
|
||||
}
|
||||
|
||||
// BuildMaintenanceMenuItems creates menu items for all registered task types
|
||||
func BuildMaintenanceMenuItems() []*MaintenanceMenuItem {
|
||||
var menuItems []*MaintenanceMenuItem
|
||||
|
||||
// Get all registered task types
|
||||
registeredTypes := GetRegisteredMaintenanceTaskTypes()
|
||||
|
||||
for _, taskType := range registeredTypes {
|
||||
menuItem := &MaintenanceMenuItem{
|
||||
TaskType: taskType,
|
||||
DisplayName: GetTaskDisplayName(taskType),
|
||||
Description: GetTaskDescription(taskType),
|
||||
Icon: GetTaskIcon(taskType),
|
||||
IsEnabled: IsMaintenanceTaskTypeRegistered(taskType),
|
||||
Path: "/maintenance/config/" + string(taskType),
|
||||
}
|
||||
|
||||
menuItems = append(menuItems, menuItem)
|
||||
}
|
||||
|
||||
return menuItems
|
||||
}
|
||||
|
||||
// Helper functions to extract configuration fields
|
||||
|
||||
// Note: Removed getVacuumConfigField, getErasureCodingConfigField, getBalanceConfigField, getReplicationConfigField
|
||||
|
||||
@@ -1,421 +0,0 @@
|
||||
package maintenance
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/glog"
|
||||
"github.com/seaweedfs/seaweedfs/weed/worker"
|
||||
"github.com/seaweedfs/seaweedfs/weed/worker/tasks"
|
||||
"github.com/seaweedfs/seaweedfs/weed/worker/types"
|
||||
|
||||
// Import task packages to trigger their auto-registration
|
||||
_ "github.com/seaweedfs/seaweedfs/weed/worker/tasks/balance"
|
||||
_ "github.com/seaweedfs/seaweedfs/weed/worker/tasks/erasure_coding"
|
||||
_ "github.com/seaweedfs/seaweedfs/weed/worker/tasks/vacuum"
|
||||
)
|
||||
|
||||
// MaintenanceWorkerService manages maintenance task execution
|
||||
// TaskExecutor defines the function signature for task execution
|
||||
type TaskExecutor func(*MaintenanceWorkerService, *MaintenanceTask) error
|
||||
|
||||
// TaskExecutorFactory creates a task executor for a given worker service
|
||||
type TaskExecutorFactory func() TaskExecutor
|
||||
|
||||
// Global registry for task executor factories
|
||||
var taskExecutorFactories = make(map[MaintenanceTaskType]TaskExecutorFactory)
|
||||
var executorRegistryMutex sync.RWMutex
|
||||
var executorRegistryInitOnce sync.Once
|
||||
|
||||
// initializeExecutorFactories dynamically registers executor factories for all auto-registered task types
|
||||
func initializeExecutorFactories() {
|
||||
executorRegistryInitOnce.Do(func() {
|
||||
// Get all registered task types from the global registry
|
||||
typesRegistry := tasks.GetGlobalTypesRegistry()
|
||||
|
||||
var taskTypes []MaintenanceTaskType
|
||||
for workerTaskType := range typesRegistry.GetAllDetectors() {
|
||||
// Convert types.TaskType to MaintenanceTaskType by string conversion
|
||||
maintenanceTaskType := MaintenanceTaskType(string(workerTaskType))
|
||||
taskTypes = append(taskTypes, maintenanceTaskType)
|
||||
}
|
||||
|
||||
// Register generic executor for all task types
|
||||
for _, taskType := range taskTypes {
|
||||
RegisterTaskExecutorFactory(taskType, createGenericTaskExecutor)
|
||||
}
|
||||
|
||||
glog.V(1).Infof("Dynamically registered generic task executor for %d task types: %v", len(taskTypes), taskTypes)
|
||||
})
|
||||
}
|
||||
|
||||
// RegisterTaskExecutorFactory registers a factory function for creating task executors
|
||||
func RegisterTaskExecutorFactory(taskType MaintenanceTaskType, factory TaskExecutorFactory) {
|
||||
executorRegistryMutex.Lock()
|
||||
defer executorRegistryMutex.Unlock()
|
||||
taskExecutorFactories[taskType] = factory
|
||||
glog.V(2).Infof("Registered executor factory for task type: %s", taskType)
|
||||
}
|
||||
|
||||
// GetTaskExecutorFactory returns the factory for a task type
|
||||
func GetTaskExecutorFactory(taskType MaintenanceTaskType) (TaskExecutorFactory, bool) {
|
||||
// Ensure executor factories are initialized
|
||||
initializeExecutorFactories()
|
||||
|
||||
executorRegistryMutex.RLock()
|
||||
defer executorRegistryMutex.RUnlock()
|
||||
factory, exists := taskExecutorFactories[taskType]
|
||||
return factory, exists
|
||||
}
|
||||
|
||||
// GetSupportedExecutorTaskTypes returns all task types with registered executor factories
|
||||
func GetSupportedExecutorTaskTypes() []MaintenanceTaskType {
|
||||
// Ensure executor factories are initialized
|
||||
initializeExecutorFactories()
|
||||
|
||||
executorRegistryMutex.RLock()
|
||||
defer executorRegistryMutex.RUnlock()
|
||||
|
||||
taskTypes := make([]MaintenanceTaskType, 0, len(taskExecutorFactories))
|
||||
for taskType := range taskExecutorFactories {
|
||||
taskTypes = append(taskTypes, taskType)
|
||||
}
|
||||
return taskTypes
|
||||
}
|
||||
|
||||
// createGenericTaskExecutor creates a generic task executor that uses the task registry
|
||||
func createGenericTaskExecutor() TaskExecutor {
|
||||
return func(mws *MaintenanceWorkerService, task *MaintenanceTask) error {
|
||||
return mws.executeGenericTask(task)
|
||||
}
|
||||
}
|
||||
|
||||
// init does minimal initialization - actual registration happens lazily
|
||||
func init() {
|
||||
// Executor factory registration will happen lazily when first accessed
|
||||
glog.V(1).Infof("Maintenance worker initialized - executor factories will be registered on first access")
|
||||
}
|
||||
|
||||
type MaintenanceWorkerService struct {
|
||||
workerID string
|
||||
address string
|
||||
adminServer string
|
||||
capabilities []MaintenanceTaskType
|
||||
maxConcurrent int
|
||||
currentTasks map[string]*MaintenanceTask
|
||||
queue *MaintenanceQueue
|
||||
adminClient AdminClient
|
||||
running bool
|
||||
stopChan chan struct{}
|
||||
|
||||
// Task execution registry
|
||||
taskExecutors map[MaintenanceTaskType]TaskExecutor
|
||||
|
||||
// Task registry for creating task instances
|
||||
taskRegistry *tasks.TaskRegistry
|
||||
}
|
||||
|
||||
// NewMaintenanceWorkerService creates a new maintenance worker service
|
||||
func NewMaintenanceWorkerService(workerID, address, adminServer string) *MaintenanceWorkerService {
|
||||
// Get all registered maintenance task types dynamically
|
||||
capabilities := GetRegisteredMaintenanceTaskTypes()
|
||||
|
||||
worker := &MaintenanceWorkerService{
|
||||
workerID: workerID,
|
||||
address: address,
|
||||
adminServer: adminServer,
|
||||
capabilities: capabilities,
|
||||
maxConcurrent: 2, // Default concurrent task limit
|
||||
currentTasks: make(map[string]*MaintenanceTask),
|
||||
stopChan: make(chan struct{}),
|
||||
taskExecutors: make(map[MaintenanceTaskType]TaskExecutor),
|
||||
taskRegistry: tasks.GetGlobalTaskRegistry(), // Use global registry with auto-registered tasks
|
||||
}
|
||||
|
||||
// Initialize task executor registry
|
||||
worker.initializeTaskExecutors()
|
||||
|
||||
glog.V(1).Infof("Created maintenance worker with %d registered task types", len(worker.taskRegistry.GetAll()))
|
||||
|
||||
return worker
|
||||
}
|
||||
|
||||
// executeGenericTask executes a task using the task registry instead of hardcoded methods
|
||||
func (mws *MaintenanceWorkerService) executeGenericTask(task *MaintenanceTask) error {
|
||||
glog.V(2).Infof("Executing generic task %s: %s for volume %d", task.ID, task.Type, task.VolumeID)
|
||||
|
||||
// Validate that task has proper typed parameters
|
||||
if task.TypedParams == nil {
|
||||
return fmt.Errorf("task %s has no typed parameters - task was not properly planned (insufficient destinations)", task.ID)
|
||||
}
|
||||
|
||||
// Convert MaintenanceTask to types.TaskType
|
||||
taskType := types.TaskType(string(task.Type))
|
||||
|
||||
// Create task instance using the registry
|
||||
taskInstance, err := mws.taskRegistry.Get(taskType).Create(task.TypedParams)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create task instance: %w", err)
|
||||
}
|
||||
|
||||
// Update progress to show task has started
|
||||
mws.updateTaskProgress(task.ID, 5)
|
||||
|
||||
// Execute the task
|
||||
err = taskInstance.Execute(context.Background(), task.TypedParams)
|
||||
if err != nil {
|
||||
return fmt.Errorf("task execution failed: %w", err)
|
||||
}
|
||||
|
||||
// Update progress to show completion
|
||||
mws.updateTaskProgress(task.ID, 100)
|
||||
|
||||
glog.V(2).Infof("Generic task %s completed successfully", task.ID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// initializeTaskExecutors sets up the task execution registry dynamically
|
||||
func (mws *MaintenanceWorkerService) initializeTaskExecutors() {
|
||||
mws.taskExecutors = make(map[MaintenanceTaskType]TaskExecutor)
|
||||
|
||||
// Get all registered executor factories and create executors
|
||||
executorRegistryMutex.RLock()
|
||||
defer executorRegistryMutex.RUnlock()
|
||||
|
||||
for taskType, factory := range taskExecutorFactories {
|
||||
executor := factory()
|
||||
mws.taskExecutors[taskType] = executor
|
||||
glog.V(3).Infof("Initialized executor for task type: %s", taskType)
|
||||
}
|
||||
|
||||
glog.V(2).Infof("Initialized %d task executors", len(mws.taskExecutors))
|
||||
}
|
||||
|
||||
// RegisterTaskExecutor allows dynamic registration of new task executors
|
||||
func (mws *MaintenanceWorkerService) RegisterTaskExecutor(taskType MaintenanceTaskType, executor TaskExecutor) {
|
||||
if mws.taskExecutors == nil {
|
||||
mws.taskExecutors = make(map[MaintenanceTaskType]TaskExecutor)
|
||||
}
|
||||
mws.taskExecutors[taskType] = executor
|
||||
glog.V(1).Infof("Registered executor for task type: %s", taskType)
|
||||
}
|
||||
|
||||
// GetSupportedTaskTypes returns all task types that this worker can execute
|
||||
func (mws *MaintenanceWorkerService) GetSupportedTaskTypes() []MaintenanceTaskType {
|
||||
return GetSupportedExecutorTaskTypes()
|
||||
}
|
||||
|
||||
// Start begins the worker service
|
||||
func (mws *MaintenanceWorkerService) Start() error {
|
||||
mws.running = true
|
||||
|
||||
// Register with admin server
|
||||
worker := &MaintenanceWorker{
|
||||
ID: mws.workerID,
|
||||
Address: mws.address,
|
||||
Capabilities: mws.capabilities,
|
||||
MaxConcurrent: mws.maxConcurrent,
|
||||
}
|
||||
|
||||
if mws.queue != nil {
|
||||
mws.queue.RegisterWorker(worker)
|
||||
}
|
||||
|
||||
// Start worker loop
|
||||
go mws.workerLoop()
|
||||
|
||||
glog.Infof("Maintenance worker %s started at %s", mws.workerID, mws.address)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stop terminates the worker service
|
||||
func (mws *MaintenanceWorkerService) Stop() {
|
||||
mws.running = false
|
||||
close(mws.stopChan)
|
||||
|
||||
// Wait for current tasks to complete or timeout
|
||||
timeout := time.NewTimer(30 * time.Second)
|
||||
defer timeout.Stop()
|
||||
|
||||
for len(mws.currentTasks) > 0 {
|
||||
select {
|
||||
case <-timeout.C:
|
||||
glog.Warningf("Worker %s stopping with %d tasks still running", mws.workerID, len(mws.currentTasks))
|
||||
return
|
||||
case <-time.After(time.Second):
|
||||
// Check again
|
||||
}
|
||||
}
|
||||
|
||||
glog.Infof("Maintenance worker %s stopped", mws.workerID)
|
||||
}
|
||||
|
||||
// workerLoop is the main worker event loop
|
||||
func (mws *MaintenanceWorkerService) workerLoop() {
|
||||
heartbeatTicker := time.NewTicker(30 * time.Second)
|
||||
defer heartbeatTicker.Stop()
|
||||
|
||||
taskRequestTicker := time.NewTicker(5 * time.Second)
|
||||
defer taskRequestTicker.Stop()
|
||||
|
||||
for mws.running {
|
||||
select {
|
||||
case <-mws.stopChan:
|
||||
return
|
||||
case <-heartbeatTicker.C:
|
||||
mws.sendHeartbeat()
|
||||
case <-taskRequestTicker.C:
|
||||
mws.requestTasks()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// sendHeartbeat sends heartbeat to admin server
|
||||
func (mws *MaintenanceWorkerService) sendHeartbeat() {
|
||||
if mws.queue != nil {
|
||||
mws.queue.UpdateWorkerHeartbeat(mws.workerID)
|
||||
}
|
||||
}
|
||||
|
||||
// requestTasks requests new tasks from the admin server
|
||||
func (mws *MaintenanceWorkerService) requestTasks() {
|
||||
if len(mws.currentTasks) >= mws.maxConcurrent {
|
||||
return // Already at capacity
|
||||
}
|
||||
|
||||
if mws.queue != nil {
|
||||
task := mws.queue.GetNextTask(mws.workerID, mws.capabilities)
|
||||
if task != nil {
|
||||
mws.executeTask(task)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// executeTask executes a maintenance task
|
||||
func (mws *MaintenanceWorkerService) executeTask(task *MaintenanceTask) {
|
||||
mws.currentTasks[task.ID] = task
|
||||
|
||||
go func() {
|
||||
defer func() {
|
||||
delete(mws.currentTasks, task.ID)
|
||||
}()
|
||||
|
||||
glog.Infof("Worker %s executing task %s: %s", mws.workerID, task.ID, task.Type)
|
||||
|
||||
// Execute task using dynamic executor registry
|
||||
var err error
|
||||
if executor, exists := mws.taskExecutors[task.Type]; exists {
|
||||
err = executor(mws, task)
|
||||
} else {
|
||||
err = fmt.Errorf("unsupported task type: %s", task.Type)
|
||||
glog.Errorf("No executor registered for task type: %s", task.Type)
|
||||
}
|
||||
|
||||
// Report task completion
|
||||
if mws.queue != nil {
|
||||
errorMsg := ""
|
||||
if err != nil {
|
||||
errorMsg = err.Error()
|
||||
}
|
||||
mws.queue.CompleteTask(task.ID, errorMsg)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
glog.Errorf("Worker %s failed to execute task %s: %v", mws.workerID, task.ID, err)
|
||||
} else {
|
||||
glog.Infof("Worker %s completed task %s successfully", mws.workerID, task.ID)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// updateTaskProgress updates the progress of a task
|
||||
func (mws *MaintenanceWorkerService) updateTaskProgress(taskID string, progress float64) {
|
||||
if mws.queue != nil {
|
||||
mws.queue.UpdateTaskProgress(taskID, progress)
|
||||
}
|
||||
}
|
||||
|
||||
// GetStatus returns the current status of the worker
|
||||
func (mws *MaintenanceWorkerService) GetStatus() map[string]interface{} {
|
||||
return map[string]interface{}{
|
||||
"worker_id": mws.workerID,
|
||||
"address": mws.address,
|
||||
"running": mws.running,
|
||||
"capabilities": mws.capabilities,
|
||||
"max_concurrent": mws.maxConcurrent,
|
||||
"current_tasks": len(mws.currentTasks),
|
||||
"task_details": mws.currentTasks,
|
||||
}
|
||||
}
|
||||
|
||||
// SetQueue sets the maintenance queue for the worker
|
||||
func (mws *MaintenanceWorkerService) SetQueue(queue *MaintenanceQueue) {
|
||||
mws.queue = queue
|
||||
}
|
||||
|
||||
// SetAdminClient sets the admin client for the worker
|
||||
func (mws *MaintenanceWorkerService) SetAdminClient(client AdminClient) {
|
||||
mws.adminClient = client
|
||||
}
|
||||
|
||||
// SetCapabilities sets the worker capabilities
|
||||
func (mws *MaintenanceWorkerService) SetCapabilities(capabilities []MaintenanceTaskType) {
|
||||
mws.capabilities = capabilities
|
||||
}
|
||||
|
||||
// SetMaxConcurrent sets the maximum concurrent tasks
|
||||
func (mws *MaintenanceWorkerService) SetMaxConcurrent(max int) {
|
||||
mws.maxConcurrent = max
|
||||
}
|
||||
|
||||
// SetHeartbeatInterval sets the heartbeat interval (placeholder for future use)
|
||||
func (mws *MaintenanceWorkerService) SetHeartbeatInterval(interval time.Duration) {
|
||||
// Future implementation for configurable heartbeat
|
||||
}
|
||||
|
||||
// SetTaskRequestInterval sets the task request interval (placeholder for future use)
|
||||
func (mws *MaintenanceWorkerService) SetTaskRequestInterval(interval time.Duration) {
|
||||
// Future implementation for configurable task requests
|
||||
}
|
||||
|
||||
// MaintenanceWorkerCommand represents a standalone maintenance worker command
|
||||
type MaintenanceWorkerCommand struct {
|
||||
workerService *MaintenanceWorkerService
|
||||
}
|
||||
|
||||
// NewMaintenanceWorkerCommand creates a new worker command
|
||||
func NewMaintenanceWorkerCommand(workerID, address, adminServer string) *MaintenanceWorkerCommand {
|
||||
return &MaintenanceWorkerCommand{
|
||||
workerService: NewMaintenanceWorkerService(workerID, address, adminServer),
|
||||
}
|
||||
}
|
||||
|
||||
// Run starts the maintenance worker as a standalone service
|
||||
func (mwc *MaintenanceWorkerCommand) Run() error {
|
||||
// Generate or load persistent worker ID if not provided
|
||||
if mwc.workerService.workerID == "" {
|
||||
// Get current working directory for worker ID persistence
|
||||
wd, err := os.Getwd()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get working directory: %w", err)
|
||||
}
|
||||
|
||||
workerID, err := worker.GenerateOrLoadWorkerID(wd)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to generate or load worker ID: %w", err)
|
||||
}
|
||||
mwc.workerService.workerID = workerID
|
||||
}
|
||||
|
||||
// Start the worker service
|
||||
err := mwc.workerService.Start()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to start maintenance worker: %w", err)
|
||||
}
|
||||
|
||||
// Wait for interrupt signal
|
||||
select {}
|
||||
}
|
||||
@@ -122,6 +122,7 @@ type Plugin struct {
|
||||
type streamSession struct {
|
||||
workerID string
|
||||
outgoing chan *plugin_pb.AdminToWorkerMessage
|
||||
done chan struct{}
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
@@ -274,6 +275,7 @@ func (r *Plugin) WorkerStream(stream plugin_pb.PluginControlService_WorkerStream
|
||||
session := &streamSession{
|
||||
workerID: workerID,
|
||||
outgoing: make(chan *plugin_pb.AdminToWorkerMessage, r.outgoingBuffer),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
r.putSession(session)
|
||||
defer r.cleanupSession(workerID)
|
||||
@@ -908,8 +910,10 @@ func (r *Plugin) sendLoop(
|
||||
return nil
|
||||
case <-r.shutdownCh:
|
||||
return nil
|
||||
case msg, ok := <-session.outgoing:
|
||||
if !ok {
|
||||
case <-session.done:
|
||||
return nil
|
||||
case msg := <-session.outgoing:
|
||||
if msg == nil {
|
||||
return nil
|
||||
}
|
||||
if err := stream.Send(msg); err != nil {
|
||||
@@ -930,6 +934,8 @@ func (r *Plugin) sendToWorker(workerID string, message *plugin_pb.AdminToWorkerM
|
||||
select {
|
||||
case <-r.shutdownCh:
|
||||
return fmt.Errorf("plugin is shutting down")
|
||||
case <-session.done:
|
||||
return fmt.Errorf("worker %s session is closed", workerID)
|
||||
case session.outgoing <- message:
|
||||
return nil
|
||||
case <-time.After(r.sendTimeout):
|
||||
@@ -1425,7 +1431,7 @@ func CloneConfigValueMap(in map[string]*plugin_pb.ConfigValue) map[string]*plugi
|
||||
|
||||
func (s *streamSession) close() {
|
||||
s.closeOnce.Do(func() {
|
||||
close(s.outgoing)
|
||||
close(s.done)
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -26,7 +26,7 @@ func TestRunDetectionSendsCancelOnContextDone(t *testing.T) {
|
||||
{JobType: jobType, CanDetect: true, MaxDetectionConcurrency: 1},
|
||||
},
|
||||
})
|
||||
session := &streamSession{workerID: workerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 4)}
|
||||
session := &streamSession{workerID: workerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 4), done: make(chan struct{})}
|
||||
pluginSvc.putSession(session)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
@@ -77,7 +77,7 @@ func TestExecuteJobSendsCancelOnContextDone(t *testing.T) {
|
||||
{JobType: jobType, CanExecute: true, MaxExecutionConcurrency: 1},
|
||||
},
|
||||
})
|
||||
session := &streamSession{workerID: workerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 4)}
|
||||
session := &streamSession{workerID: workerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 4), done: make(chan struct{})}
|
||||
pluginSvc.putSession(session)
|
||||
|
||||
job := &plugin_pb.JobSpec{JobId: "job-1", JobType: jobType}
|
||||
@@ -135,8 +135,8 @@ func TestAdminScriptExecutionBlocksOtherDetection(t *testing.T) {
|
||||
{JobType: "vacuum", CanDetect: true, MaxDetectionConcurrency: 1},
|
||||
},
|
||||
})
|
||||
adminSession := &streamSession{workerID: adminWorkerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 8)}
|
||||
otherSession := &streamSession{workerID: otherWorkerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 8)}
|
||||
adminSession := &streamSession{workerID: adminWorkerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 8), done: make(chan struct{})}
|
||||
otherSession := &streamSession{workerID: otherWorkerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 8), done: make(chan struct{})}
|
||||
pluginSvc.putSession(adminSession)
|
||||
pluginSvc.putSession(otherSession)
|
||||
|
||||
@@ -214,8 +214,8 @@ func TestAdminScriptExecutionBlocksOtherExecution(t *testing.T) {
|
||||
{JobType: "vacuum", CanExecute: true, MaxExecutionConcurrency: 1},
|
||||
},
|
||||
})
|
||||
adminSession := &streamSession{workerID: adminWorkerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 8)}
|
||||
otherSession := &streamSession{workerID: otherWorkerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 8)}
|
||||
adminSession := &streamSession{workerID: adminWorkerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 8), done: make(chan struct{})}
|
||||
otherSession := &streamSession{workerID: otherWorkerID, outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 8), done: make(chan struct{})}
|
||||
pluginSvc.putSession(adminSession)
|
||||
pluginSvc.putSession(otherSession)
|
||||
|
||||
|
||||
@@ -22,7 +22,7 @@ func TestRunDetectionIncludesLatestSuccessfulRun(t *testing.T) {
|
||||
{JobType: jobType, CanDetect: true, MaxDetectionConcurrency: 1},
|
||||
},
|
||||
})
|
||||
session := &streamSession{workerID: "worker-a", outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 1)}
|
||||
session := &streamSession{workerID: "worker-a", outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 1), done: make(chan struct{})}
|
||||
pluginSvc.putSession(session)
|
||||
|
||||
oldSuccess := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC)
|
||||
@@ -80,7 +80,7 @@ func TestRunDetectionOmitsLastSuccessfulRunWhenNoSuccessHistory(t *testing.T) {
|
||||
{JobType: jobType, CanDetect: true, MaxDetectionConcurrency: 1},
|
||||
},
|
||||
})
|
||||
session := &streamSession{workerID: "worker-a", outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 1)}
|
||||
session := &streamSession{workerID: "worker-a", outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 1), done: make(chan struct{})}
|
||||
pluginSvc.putSession(session)
|
||||
|
||||
if err := pluginSvc.store.AppendRunRecord(jobType, &JobRunRecord{
|
||||
@@ -130,7 +130,7 @@ func TestRunDetectionWithReportCapturesDetectionActivities(t *testing.T) {
|
||||
{JobType: jobType, CanDetect: true, MaxDetectionConcurrency: 1},
|
||||
},
|
||||
})
|
||||
session := &streamSession{workerID: "worker-a", outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 1)}
|
||||
session := &streamSession{workerID: "worker-a", outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 1), done: make(chan struct{})}
|
||||
pluginSvc.putSession(session)
|
||||
|
||||
reportCh := make(chan *DetectionReport, 1)
|
||||
@@ -210,7 +210,7 @@ func TestRunDetectionAdminScriptUsesLastCompletedRun(t *testing.T) {
|
||||
{JobType: jobType, CanDetect: true, MaxDetectionConcurrency: 1},
|
||||
},
|
||||
})
|
||||
session := &streamSession{workerID: "worker-admin-script", outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 1)}
|
||||
session := &streamSession{workerID: "worker-admin-script", outgoing: make(chan *plugin_pb.AdminToWorkerMessage, 1), done: make(chan struct{})}
|
||||
pluginSvc.putSession(session)
|
||||
|
||||
successCompleted := time.Date(2026, 2, 1, 10, 0, 0, 0, time.UTC)
|
||||
|
||||
@@ -95,16 +95,6 @@ func (r *Plugin) laneSchedulerLoop(ls *schedulerLaneState) {
|
||||
}
|
||||
}
|
||||
|
||||
// schedulerLoop is kept for backward compatibility; it delegates to
|
||||
// laneSchedulerLoop with the default lane. New code should not call this.
|
||||
func (r *Plugin) schedulerLoop() {
|
||||
ls := r.lanes[LaneDefault]
|
||||
if ls == nil {
|
||||
ls = newLaneState(LaneDefault)
|
||||
}
|
||||
r.laneSchedulerLoop(ls)
|
||||
}
|
||||
|
||||
// runLaneSchedulerIteration runs one scheduling pass for a single lane,
|
||||
// processing only the job types assigned to that lane.
|
||||
//
|
||||
@@ -229,82 +219,6 @@ func (r *Plugin) runLaneSchedulerIterationConcurrent(ls *schedulerLaneState, job
|
||||
return hadJobs.Load()
|
||||
}
|
||||
|
||||
// runSchedulerIteration is kept for backward compatibility. It runs a
|
||||
// single iteration across ALL job types (equivalent to the old single-loop
|
||||
// behavior). It is only used by the legacy schedulerLoop() fallback.
|
||||
func (r *Plugin) runSchedulerIteration() bool {
|
||||
ls := r.lanes[LaneDefault]
|
||||
if ls == nil {
|
||||
ls = newLaneState(LaneDefault)
|
||||
}
|
||||
// For backward compat, the old function processes all job types.
|
||||
r.expireStaleJobs(time.Now().UTC())
|
||||
|
||||
jobTypes := r.registry.DetectableJobTypes()
|
||||
if len(jobTypes) == 0 {
|
||||
r.setSchedulerLoopState("", "idle")
|
||||
return false
|
||||
}
|
||||
|
||||
r.setSchedulerLoopState("", "waiting_for_lock")
|
||||
releaseLock, err := r.acquireAdminLock("plugin scheduler iteration")
|
||||
if err != nil {
|
||||
glog.Warningf("Plugin scheduler failed to acquire lock: %v", err)
|
||||
r.setSchedulerLoopState("", "idle")
|
||||
return false
|
||||
}
|
||||
if releaseLock != nil {
|
||||
defer releaseLock()
|
||||
}
|
||||
|
||||
active := make(map[string]struct{}, len(jobTypes))
|
||||
hadJobs := false
|
||||
|
||||
for _, jobType := range jobTypes {
|
||||
active[jobType] = struct{}{}
|
||||
|
||||
policy, enabled, err := r.loadSchedulerPolicy(jobType)
|
||||
if err != nil {
|
||||
glog.Warningf("Plugin scheduler failed to load policy for %s: %v", jobType, err)
|
||||
continue
|
||||
}
|
||||
if !enabled {
|
||||
r.clearSchedulerJobType(jobType)
|
||||
continue
|
||||
}
|
||||
initialDelay := time.Duration(0)
|
||||
if runInfo := r.snapshotSchedulerRun(jobType); runInfo.lastRunStartedAt.IsZero() {
|
||||
initialDelay = 5 * time.Second
|
||||
}
|
||||
if !r.markDetectionDue(jobType, policy.DetectionInterval, initialDelay) {
|
||||
continue
|
||||
}
|
||||
|
||||
detected := r.runJobTypeIteration(jobType, policy)
|
||||
if detected {
|
||||
hadJobs = true
|
||||
}
|
||||
}
|
||||
|
||||
r.pruneSchedulerState(active)
|
||||
r.pruneDetectorLeases(active)
|
||||
r.setSchedulerLoopState("", "idle")
|
||||
return hadJobs
|
||||
}
|
||||
|
||||
// wakeLane wakes the scheduler goroutine for a specific lane.
|
||||
func (r *Plugin) wakeLane(lane SchedulerLane) {
|
||||
if r == nil {
|
||||
return
|
||||
}
|
||||
if ls, ok := r.lanes[lane]; ok {
|
||||
select {
|
||||
case ls.wakeCh <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// wakeAllLanes wakes all lane scheduler goroutines.
|
||||
func (r *Plugin) wakeAllLanes() {
|
||||
if r == nil {
|
||||
|
||||
@@ -210,16 +210,6 @@ func (r *Plugin) setSchedulerLoopStateForJobType(jobType, phase string) {
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Plugin) recordSchedulerIterationComplete(hadJobs bool) {
|
||||
if r == nil {
|
||||
return
|
||||
}
|
||||
r.schedulerLoopMu.Lock()
|
||||
r.schedulerLoopState.lastIterationHadJobs = hadJobs
|
||||
r.schedulerLoopState.lastIterationCompleted = time.Now().UTC()
|
||||
r.schedulerLoopMu.Unlock()
|
||||
}
|
||||
|
||||
func (r *Plugin) snapshotSchedulerLoopState() schedulerLoopState {
|
||||
if r == nil {
|
||||
return schedulerLoopState{}
|
||||
|
||||
@@ -163,7 +163,7 @@ templ S3Buckets(data dash.S3BucketsData) {
|
||||
for _, bucket := range data.Buckets {
|
||||
<tr>
|
||||
<td>
|
||||
<a href={templ.SafeURL(fmt.Sprintf("/files?path=/buckets/%s", bucket.Name))}
|
||||
<a href={dash.PUrl(ctx, fmt.Sprintf("/files?path=/buckets/%s", bucket.Name))}
|
||||
class="text-decoration-none">
|
||||
<i class="fas fa-cube me-2"></i>
|
||||
{bucket.Name}
|
||||
@@ -236,7 +236,7 @@ templ S3Buckets(data dash.S3BucketsData) {
|
||||
</td>
|
||||
<td>
|
||||
<div class="btn-group btn-group-sm" role="group">
|
||||
<a href={templ.SafeURL(fmt.Sprintf("/files?path=/buckets/%s", bucket.Name))}
|
||||
<a href={dash.PUrl(ctx, fmt.Sprintf("/files?path=/buckets/%s", bucket.Name))}
|
||||
class="btn btn-outline-success btn-sm"
|
||||
title="Browse Files">
|
||||
<i class="fas fa-folder-open"></i>
|
||||
|
||||
@@ -171,9 +171,9 @@ func S3Buckets(data dash.S3BucketsData) templ.Component {
|
||||
return templ_7745c5c3_Err
|
||||
}
|
||||
var templ_7745c5c3_Var5 templ.SafeURL
|
||||
templ_7745c5c3_Var5, templ_7745c5c3_Err = templ.JoinURLErrs(templ.SafeURL(fmt.Sprintf("/files?path=/buckets/%s", bucket.Name)))
|
||||
templ_7745c5c3_Var5, templ_7745c5c3_Err = templ.JoinURLErrs(dash.PUrl(ctx, fmt.Sprintf("/files?path=/buckets/%s", bucket.Name)))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3_buckets.templ`, Line: 166, Col: 123}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3_buckets.templ`, Line: 166, Col: 124}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var5))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -439,9 +439,9 @@ func S3Buckets(data dash.S3BucketsData) templ.Component {
|
||||
return templ_7745c5c3_Err
|
||||
}
|
||||
var templ_7745c5c3_Var19 templ.SafeURL
|
||||
templ_7745c5c3_Var19, templ_7745c5c3_Err = templ.JoinURLErrs(templ.SafeURL(fmt.Sprintf("/files?path=/buckets/%s", bucket.Name)))
|
||||
templ_7745c5c3_Var19, templ_7745c5c3_Err = templ.JoinURLErrs(dash.PUrl(ctx, fmt.Sprintf("/files?path=/buckets/%s", bucket.Name)))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3_buckets.templ`, Line: 239, Col: 127}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3_buckets.templ`, Line: 239, Col: 128}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var19))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
|
||||
@@ -2,6 +2,7 @@ package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/url"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/admin/dash"
|
||||
"github.com/seaweedfs/seaweedfs/weed/s3api/s3tables"
|
||||
@@ -152,7 +153,7 @@ templ S3TablesBuckets(data dash.S3TablesBucketsData) {
|
||||
<div class="btn-group btn-group-sm" role="group">
|
||||
{{ bucketName, parseErr := s3tables.ParseBucketNameFromARN(bucket.ARN) }}
|
||||
if parseErr == nil {
|
||||
<a class="btn btn-outline-primary btn-sm" href={ templ.SafeURL(fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces", bucketName)) }>
|
||||
<a class="btn btn-outline-primary btn-sm" href={ dash.PUrl(ctx, fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces", url.PathEscape(bucketName))) }>
|
||||
<i class="fas fa-folder-open"></i>
|
||||
</a>
|
||||
} else {
|
||||
|
||||
@@ -10,6 +10,7 @@ import templruntime "github.com/a-h/templ/runtime"
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/url"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/admin/dash"
|
||||
"github.com/seaweedfs/seaweedfs/weed/s3api/s3tables"
|
||||
@@ -48,7 +49,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var2 templ.SafeURL
|
||||
templ_7745c5c3_Var2, templ_7745c5c3_Err = templ.JoinURLErrs(templ.SafeURL(fmt.Sprintf("http://localhost:%d/v1/config", data.IcebergPort)))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 23, Col: 124}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 24, Col: 124}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var2))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -66,7 +67,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var3 string
|
||||
templ_7745c5c3_Var3, templ_7745c5c3_Err = templ.JoinStringErrs(fmt.Sprintf("%d", data.IcebergPort))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 30, Col: 91}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 31, Col: 91}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var3))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -84,7 +85,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var4 string
|
||||
templ_7745c5c3_Var4, templ_7745c5c3_Err = templ.JoinStringErrs(fmt.Sprintf("%d", data.IcebergPort))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 37, Col: 109}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 38, Col: 109}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var4))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -97,7 +98,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var5 string
|
||||
templ_7745c5c3_Var5, templ_7745c5c3_Err = templ.JoinStringErrs(fmt.Sprintf("%d", data.IcebergPort))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 39, Col: 107}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 40, Col: 107}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var5))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -120,7 +121,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var6 string
|
||||
templ_7745c5c3_Var6, templ_7745c5c3_Err = templ.JoinStringErrs(fmt.Sprintf("%d", data.TotalBuckets))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 69, Col: 47}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 70, Col: 47}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var6))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -133,7 +134,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var7 string
|
||||
templ_7745c5c3_Var7, templ_7745c5c3_Err = templ.JoinStringErrs(data.LastUpdated.Format("2006-01-02 15:04"))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 88, Col: 54}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 89, Col: 54}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var7))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -147,7 +148,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var8 string
|
||||
templ_7745c5c3_Var8, templ_7745c5c3_Err = templ.JoinStringErrs(fmt.Sprintf("%d", data.IcebergPort))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 108, Col: 47}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 109, Col: 47}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var8))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -171,7 +172,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var9 string
|
||||
templ_7745c5c3_Var9, templ_7745c5c3_Err = templ.JoinStringErrs(bucket.Name)
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 146, Col: 28}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 147, Col: 28}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var9))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -184,7 +185,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var10 string
|
||||
templ_7745c5c3_Var10, templ_7745c5c3_Err = templ.JoinStringErrs(bucket.OwnerAccountID)
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 147, Col: 38}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 148, Col: 38}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var10))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -197,7 +198,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var11 string
|
||||
templ_7745c5c3_Var11, templ_7745c5c3_Err = templ.JoinStringErrs(bucket.ARN)
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 148, Col: 52}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 149, Col: 52}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var11))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -210,7 +211,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var12 string
|
||||
templ_7745c5c3_Var12, templ_7745c5c3_Err = templ.JoinStringErrs(bucket.Name)
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 149, Col: 52}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 150, Col: 52}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var12))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -223,7 +224,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var13 string
|
||||
templ_7745c5c3_Var13, templ_7745c5c3_Err = templ.JoinStringErrs(bucket.CreatedAt.Format("2006-01-02 15:04"))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 150, Col: 60}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 151, Col: 60}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var13))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -240,9 +241,9 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
return templ_7745c5c3_Err
|
||||
}
|
||||
var templ_7745c5c3_Var14 templ.SafeURL
|
||||
templ_7745c5c3_Var14, templ_7745c5c3_Err = templ.JoinURLErrs(templ.SafeURL(fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces", bucketName)))
|
||||
templ_7745c5c3_Var14, templ_7745c5c3_Err = templ.JoinURLErrs(dash.PUrl(ctx, fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces", url.PathEscape(bucketName))))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 155, Col: 149}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 156, Col: 166}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var14))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -265,7 +266,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var15 string
|
||||
templ_7745c5c3_Var15, templ_7745c5c3_Err = templ.JoinStringErrs(bucket.ARN)
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 163, Col: 122}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 164, Col: 122}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var15))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -278,7 +279,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var16 string
|
||||
templ_7745c5c3_Var16, templ_7745c5c3_Err = templ.JoinStringErrs(bucket.ARN)
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 166, Col: 126}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 167, Col: 126}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var16))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -291,7 +292,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var17 string
|
||||
templ_7745c5c3_Var17, templ_7745c5c3_Err = templ.JoinStringErrs(bucket.ARN)
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 169, Col: 128}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 170, Col: 128}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var17))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -304,7 +305,7 @@ func S3TablesBuckets(data dash.S3TablesBucketsData) templ.Component {
|
||||
var templ_7745c5c3_Var18 string
|
||||
templ_7745c5c3_Var18, templ_7745c5c3_Err = templ.JoinStringErrs(bucket.Name)
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 169, Col: 161}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 170, Col: 161}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var18))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -342,7 +343,7 @@ CREATE SECRET (
|
||||
|
||||
SELECT * FROM iceberg_scan('s3://my-table-bucket/my-namespace/my-table');`)
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 219, Col: 74}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 220, Col: 74}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var19))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -365,7 +366,7 @@ catalog = load_catalog(
|
||||
|
||||
namespaces = catalog.list_namespaces()`)
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 235, Col: 39}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_buckets.templ`, Line: 236, Col: 39}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var20))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
|
||||
@@ -24,7 +24,7 @@ templ S3TablesTables(data dash.S3TablesTablesData) {
|
||||
<div class="mb-3">
|
||||
{{ bucketName, parseErr := s3tables.ParseBucketNameFromARN(data.BucketARN) }}
|
||||
if parseErr == nil {
|
||||
<a href={ templ.SafeURL(fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces", bucketName)) } class="btn btn-sm btn-outline-secondary">
|
||||
<a href={ dash.PUrl(ctx, fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces", url.PathEscape(bucketName))) } class="btn btn-sm btn-outline-secondary">
|
||||
<i class="fas fa-arrow-left me-1"></i>Back to Namespaces
|
||||
</a>
|
||||
} else {
|
||||
@@ -126,7 +126,7 @@ templ S3TablesTables(data dash.S3TablesTablesData) {
|
||||
<td>
|
||||
<div class="btn-group btn-group-sm" role="group">
|
||||
if parseErr == nil {
|
||||
<a class="btn btn-outline-primary btn-sm" href={ templ.SafeURL(fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces/%s/tables/%s", url.PathEscape(bucketName), url.PathEscape(data.Namespace), url.PathEscape(tableName))) } title="View Iceberg Details">
|
||||
<a class="btn btn-outline-primary btn-sm" href={ dash.PUrl(ctx, fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces/%s/tables/%s", url.PathEscape(bucketName), url.PathEscape(data.Namespace), url.PathEscape(tableName))) } title="View Iceberg Details">
|
||||
<i class="fas fa-eye"></i>
|
||||
</a>
|
||||
} else {
|
||||
|
||||
@@ -48,9 +48,9 @@ func S3TablesTables(data dash.S3TablesTablesData) templ.Component {
|
||||
return templ_7745c5c3_Err
|
||||
}
|
||||
var templ_7745c5c3_Var2 templ.SafeURL
|
||||
templ_7745c5c3_Var2, templ_7745c5c3_Err = templ.JoinURLErrs(templ.SafeURL(fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces", bucketName)))
|
||||
templ_7745c5c3_Var2, templ_7745c5c3_Err = templ.JoinURLErrs(dash.PUrl(ctx, fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces", url.PathEscape(bucketName))))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_tables.templ`, Line: 27, Col: 99}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_tables.templ`, Line: 27, Col: 116}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var2))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
@@ -308,9 +308,9 @@ func S3TablesTables(data dash.S3TablesTablesData) templ.Component {
|
||||
return templ_7745c5c3_Err
|
||||
}
|
||||
var templ_7745c5c3_Var17 templ.SafeURL
|
||||
templ_7745c5c3_Var17, templ_7745c5c3_Err = templ.JoinURLErrs(templ.SafeURL(fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces/%s/tables/%s", url.PathEscape(bucketName), url.PathEscape(data.Namespace), url.PathEscape(tableName))))
|
||||
templ_7745c5c3_Var17, templ_7745c5c3_Err = templ.JoinURLErrs(dash.PUrl(ctx, fmt.Sprintf("/object-store/s3tables/buckets/%s/namespaces/%s/tables/%s", url.PathEscape(bucketName), url.PathEscape(data.Namespace), url.PathEscape(tableName))))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_tables.templ`, Line: 129, Col: 237}
|
||||
return templ.Error{Err: templ_7745c5c3_Err, FileName: `view/app/s3tables_tables.templ`, Line: 129, Col: 238}
|
||||
}
|
||||
_, templ_7745c5c3_Err = templ_7745c5c3_Buffer.WriteString(templ.EscapeString(templ_7745c5c3_Var17))
|
||||
if templ_7745c5c3_Err != nil {
|
||||
|
||||
@@ -6,20 +6,6 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
// getStatusColor returns Bootstrap color class for status
|
||||
func getStatusColor(status string) string {
|
||||
switch status {
|
||||
case "active", "healthy":
|
||||
return "success"
|
||||
case "warning":
|
||||
return "warning"
|
||||
case "critical", "unreachable":
|
||||
return "danger"
|
||||
default:
|
||||
return "secondary"
|
||||
}
|
||||
}
|
||||
|
||||
// formatBytes converts bytes to human readable format
|
||||
func formatBytes(bytes int64) string {
|
||||
if bytes == 0 {
|
||||
|
||||
@@ -95,18 +95,6 @@ func NewCluster() *Cluster {
|
||||
}
|
||||
}
|
||||
|
||||
func (cluster *Cluster) getGroupMembers(filerGroup FilerGroupName, nodeType string, createIfNotFound bool) *GroupMembers {
|
||||
switch nodeType {
|
||||
case FilerType:
|
||||
return cluster.filerGroups.getGroupMembers(filerGroup, createIfNotFound)
|
||||
case BrokerType:
|
||||
return cluster.brokerGroups.getGroupMembers(filerGroup, createIfNotFound)
|
||||
case S3Type:
|
||||
return cluster.s3Groups.getGroupMembers(filerGroup, createIfNotFound)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cluster *Cluster) AddClusterNode(ns, nodeType string, dataCenter DataCenter, rack Rack, address pb.ServerAddress, version string) []*master_pb.KeepConnectedResponse {
|
||||
filerGroup := FilerGroupName(ns)
|
||||
switch nodeType {
|
||||
|
||||
@@ -511,11 +511,6 @@ func recoveryMiddleware(next http.Handler) http.Handler {
|
||||
})
|
||||
}
|
||||
|
||||
// GetAdminOptions returns the admin command options for testing
|
||||
func GetAdminOptions() *AdminOptions {
|
||||
return &AdminOptions{}
|
||||
}
|
||||
|
||||
// loadOrGenerateSessionKeys loads or creates authentication/encryption keys for session cookies.
|
||||
func loadOrGenerateSessionKeys(dataDir string) ([]byte, []byte, error) {
|
||||
const keyLen = 32
|
||||
|
||||
@@ -132,16 +132,3 @@ func fetchContent(masterFn operation.GetMasterFn, grpcDialOption grpc.DialOption
|
||||
content, e = io.ReadAll(rc.Body)
|
||||
return
|
||||
}
|
||||
|
||||
func WriteFile(filename string, data []byte, perm os.FileMode) error {
|
||||
f, err := os.OpenFile(filename, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, perm)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
n, err := f.Write(data)
|
||||
f.Close()
|
||||
if err == nil && n < len(data) {
|
||||
err = io.ErrShortWrite
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -146,6 +146,7 @@ func init() {
|
||||
filerS3Options.iamReadOnly = cmdFiler.Flag.Bool("s3.iam.readOnly", true, "disable IAM write operations on this server")
|
||||
filerS3Options.portIceberg = cmdFiler.Flag.Int("s3.port.iceberg", 8181, "Iceberg REST Catalog server listen port (0 to disable)")
|
||||
filerS3Options.externalUrl = cmdFiler.Flag.String("s3.externalUrl", "", "the external URL clients use to connect (e.g. https://api.example.com:9000). Used for S3 signature verification behind a reverse proxy. Falls back to S3_EXTERNAL_URL env var.")
|
||||
filerS3Options.defaultFileMode = cmdFiler.Flag.String("s3.defaultFileMode", "", "default file mode for S3 uploaded objects, e.g. 0660, 0644, 0666")
|
||||
|
||||
// start webdav on filer
|
||||
filerStartWebDav = cmdFiler.Flag.Bool("webdav", false, "whether to start webdav gateway")
|
||||
|
||||
@@ -250,6 +250,7 @@ func initMiniS3Flags() {
|
||||
miniS3Options.auditLogConfig = cmdMini.Flag.String("s3.auditLogConfig", "", "path to the audit log config file")
|
||||
miniS3Options.allowDeleteBucketNotEmpty = miniS3AllowDeleteBucketNotEmpty
|
||||
miniS3Options.externalUrl = cmdMini.Flag.String("s3.externalUrl", "", "the external URL clients use to connect (e.g. https://api.example.com:9000). Used for S3 signature verification behind a reverse proxy. Falls back to S3_EXTERNAL_URL env var.")
|
||||
miniS3Options.defaultFileMode = cmdMini.Flag.String("s3.defaultFileMode", "", "default file mode for S3 uploaded objects, e.g. 0660, 0644, 0666")
|
||||
// In mini mode, S3 uses the shared debug server started at line 681, not its own separate debug server
|
||||
miniS3Options.debug = new(bool) // explicitly false
|
||||
miniS3Options.debugPort = cmdMini.Flag.Int("s3.debug.port", 6060, "http port for debugging (unused in mini mode)")
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"net/http"
|
||||
"os"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -36,6 +37,9 @@ var (
|
||||
s3StandaloneOptions S3Options
|
||||
)
|
||||
|
||||
// S3Options holds CLI flags for the S3 gateway.
|
||||
// Flags are registered in multiple commands: s3.go (standalone), server.go, filer.go, and mini.go.
|
||||
// When adding a new field, update all four flag registration sites.
|
||||
type S3Options struct {
|
||||
filer *string
|
||||
bindIp *string
|
||||
@@ -68,6 +72,7 @@ type S3Options struct {
|
||||
debugPort *int
|
||||
cipher *bool
|
||||
externalUrl *string
|
||||
defaultFileMode *string
|
||||
}
|
||||
|
||||
func init() {
|
||||
@@ -103,6 +108,7 @@ func init() {
|
||||
s3StandaloneOptions.debugPort = cmdS3.Flag.Int("debug.port", 6060, "http port for debugging")
|
||||
s3StandaloneOptions.cipher = cmdS3.Flag.Bool("encryptVolumeData", false, "encrypt data on volume servers")
|
||||
s3StandaloneOptions.externalUrl = cmdS3.Flag.String("externalUrl", "", "the external URL clients use to connect (e.g. https://api.example.com:9000). Used for S3 signature verification behind a reverse proxy. Falls back to S3_EXTERNAL_URL env var.")
|
||||
s3StandaloneOptions.defaultFileMode = cmdS3.Flag.String("defaultFileMode", "", "default file mode for S3 uploaded objects, e.g. 0660, 0644, 0666")
|
||||
}
|
||||
|
||||
var cmdS3 = &Command{
|
||||
@@ -232,6 +238,17 @@ func (s3opt *S3Options) resolveExternalUrl() string {
|
||||
return os.Getenv("S3_EXTERNAL_URL")
|
||||
}
|
||||
|
||||
func (s3opt *S3Options) parseDefaultFileMode() (uint32, error) {
|
||||
if s3opt.defaultFileMode == nil || *s3opt.defaultFileMode == "" {
|
||||
return 0, nil
|
||||
}
|
||||
mode, err := strconv.ParseUint(*s3opt.defaultFileMode, 8, 32)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("invalid defaultFileMode %q: %v", *s3opt.defaultFileMode, err)
|
||||
}
|
||||
return uint32(mode), nil
|
||||
}
|
||||
|
||||
func (s3opt *S3Options) startS3Server() bool {
|
||||
|
||||
filerAddresses := pb.ServerAddresses(*s3opt.filer).ToAddresses()
|
||||
@@ -298,6 +315,11 @@ func (s3opt *S3Options) startS3Server() bool {
|
||||
*s3opt.bindIp = "0.0.0.0"
|
||||
}
|
||||
|
||||
defaultFileMode, fileModeErr := s3opt.parseDefaultFileMode()
|
||||
if fileModeErr != nil {
|
||||
glog.Fatalf("S3 API Server startup error: %v", fileModeErr)
|
||||
}
|
||||
|
||||
s3ApiServer, s3ApiServer_err = s3api.NewS3ApiServer(router, &s3api.S3ApiServerOption{
|
||||
Filers: filerAddresses,
|
||||
Masters: masterAddresses,
|
||||
@@ -320,6 +342,7 @@ func (s3opt *S3Options) startS3Server() bool {
|
||||
BindIp: *s3opt.bindIp,
|
||||
GrpcPort: *s3opt.portGrpc,
|
||||
ExternalUrl: s3opt.resolveExternalUrl(),
|
||||
DefaultFileMode: defaultFileMode,
|
||||
})
|
||||
if s3ApiServer_err != nil {
|
||||
glog.Fatalf("S3 API Server startup error: %v", s3ApiServer_err)
|
||||
|
||||
@@ -179,6 +179,19 @@ password = ""
|
||||
user = ""
|
||||
password = ""
|
||||
|
||||
# SSE-S3 server-side encryption key management
|
||||
# These settings configure the Key Encryption Key (KEK) for S3 SSE-S3 encryption.
|
||||
# Set exactly one of kek or key. If neither is set, SSE-S3 is disabled.
|
||||
# Can also be set via env vars: WEED_S3_SSE_KEK, WEED_S3_SSE_KEY
|
||||
[s3.sse]
|
||||
# hex-encoded 256-bit key, same format as the legacy /etc/s3/sse_kek filer file.
|
||||
# Use this to migrate from a filer-stored KEK: copy the value from /etc/s3/sse_kek.
|
||||
# Generate a new one with: openssl rand -hex 32
|
||||
kek = ""
|
||||
# any secret string; a 256-bit key is derived automatically via HKDF-SHA256.
|
||||
# Cannot be used while /etc/s3/sse_kek exists on the filer — delete it first.
|
||||
key = ""
|
||||
|
||||
# white list. It's checking request ip address.
|
||||
[guard]
|
||||
white_list = ""
|
||||
|
||||
@@ -180,6 +180,7 @@ func init() {
|
||||
s3Options.iamReadOnly = cmdServer.Flag.Bool("s3.iam.readOnly", true, "disable IAM write operations on this server")
|
||||
s3Options.cipher = cmdServer.Flag.Bool("s3.encryptVolumeData", false, "encrypt data on volume servers for S3 uploads")
|
||||
s3Options.externalUrl = cmdServer.Flag.String("s3.externalUrl", "", "the external URL clients use to connect (e.g. https://api.example.com:9000). Used for S3 signature verification behind a reverse proxy. Falls back to S3_EXTERNAL_URL env var.")
|
||||
s3Options.defaultFileMode = cmdServer.Flag.String("s3.defaultFileMode", "", "default file mode for S3 uploaded objects, e.g. 0660, 0644, 0666")
|
||||
|
||||
sftpOptions.port = cmdServer.Flag.Int("sftp.port", 2022, "SFTP server listen port")
|
||||
sftpOptions.sshPrivateKey = cmdServer.Flag.String("sftp.sshPrivateKey", "", "path to the SSH private key file for host authentication")
|
||||
|
||||
@@ -57,42 +57,6 @@ func LoadCredentialConfiguration() (*CredentialConfig, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetCredentialStoreConfig extracts credential store configuration from command line flags
|
||||
// This is used when credential store is configured via command line instead of credential.toml
|
||||
func GetCredentialStoreConfig(store string, config util.Configuration, prefix string) *CredentialConfig {
|
||||
if store == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
return &CredentialConfig{
|
||||
Store: store,
|
||||
Config: config,
|
||||
Prefix: prefix,
|
||||
}
|
||||
}
|
||||
|
||||
// MergeCredentialConfig merges command line credential config with credential.toml config
|
||||
// Command line flags take priority over credential.toml
|
||||
func MergeCredentialConfig(cmdLineStore string, cmdLineConfig util.Configuration, cmdLinePrefix string) (*CredentialConfig, error) {
|
||||
// If command line credential store is specified, use it
|
||||
if cmdLineStore != "" {
|
||||
glog.V(0).Infof("Using command line credential configuration: store=%s", cmdLineStore)
|
||||
return GetCredentialStoreConfig(cmdLineStore, cmdLineConfig, cmdLinePrefix), nil
|
||||
}
|
||||
|
||||
// Otherwise, try to load from credential.toml
|
||||
config, err := LoadCredentialConfiguration()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if config == nil {
|
||||
glog.V(1).Info("No credential store configured")
|
||||
}
|
||||
|
||||
return config, nil
|
||||
}
|
||||
|
||||
// NewCredentialManagerWithDefaults creates a credential manager with fallback to defaults
|
||||
// If explicitStore is provided, it will be used regardless of credential.toml
|
||||
// If explicitStore is empty, it tries credential.toml first, then defaults to "filer_etc"
|
||||
|
||||
@@ -207,32 +207,6 @@ func (store *FilerEtcStore) loadPoliciesFromMultiFile(ctx context.Context, polic
|
||||
})
|
||||
}
|
||||
|
||||
func (store *FilerEtcStore) migratePoliciesToMultiFile(ctx context.Context, policies map[string]policy_engine.PolicyDocument) error {
|
||||
glog.Infof("Migrating IAM policies to multi-file layout...")
|
||||
|
||||
// 1. Save all policies to individual files
|
||||
for name, policy := range policies {
|
||||
if err := store.savePolicy(ctx, name, policy); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Rename legacy file
|
||||
return store.withFilerClient(func(client filer_pb.SeaweedFilerClient) error {
|
||||
_, err := client.AtomicRenameEntry(ctx, &filer_pb.AtomicRenameEntryRequest{
|
||||
OldDirectory: filer.IamConfigDirectory,
|
||||
OldName: filer.IamPoliciesFile,
|
||||
NewDirectory: filer.IamConfigDirectory,
|
||||
NewName: IamLegacyPoliciesOldFile,
|
||||
})
|
||||
if err != nil {
|
||||
glog.Errorf("Failed to rename legacy IAM policies file %s/%s to %s: %v",
|
||||
filer.IamConfigDirectory, filer.IamPoliciesFile, IamLegacyPoliciesOldFile, err)
|
||||
}
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func (store *FilerEtcStore) savePolicy(ctx context.Context, name string, document policy_engine.PolicyDocument) error {
|
||||
if err := validatePolicyName(name); err != nil {
|
||||
return err
|
||||
|
||||
@@ -1,221 +0,0 @@
|
||||
package credential
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/glog"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/iam_pb"
|
||||
"github.com/seaweedfs/seaweedfs/weed/util"
|
||||
)
|
||||
|
||||
// MigrateCredentials migrates credentials from one store to another
|
||||
func MigrateCredentials(fromStoreName, toStoreName CredentialStoreTypeName, configuration util.Configuration, fromPrefix, toPrefix string) error {
|
||||
ctx := context.Background()
|
||||
|
||||
// Create source credential manager
|
||||
fromCM, err := NewCredentialManager(fromStoreName, configuration, fromPrefix)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create source credential manager (%s): %v", fromStoreName, err)
|
||||
}
|
||||
defer fromCM.Shutdown()
|
||||
|
||||
// Create destination credential manager
|
||||
toCM, err := NewCredentialManager(toStoreName, configuration, toPrefix)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create destination credential manager (%s): %v", toStoreName, err)
|
||||
}
|
||||
defer toCM.Shutdown()
|
||||
|
||||
// Load configuration from source
|
||||
glog.Infof("Loading configuration from %s store...", fromStoreName)
|
||||
config, err := fromCM.LoadConfiguration(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to load configuration from source store: %w", err)
|
||||
}
|
||||
|
||||
if config == nil || len(config.Identities) == 0 {
|
||||
glog.Info("No identities found in source store")
|
||||
return nil
|
||||
}
|
||||
|
||||
glog.Infof("Found %d identities in source store", len(config.Identities))
|
||||
|
||||
// Migrate each identity
|
||||
var migrated, failed int
|
||||
for _, identity := range config.Identities {
|
||||
glog.V(1).Infof("Migrating user: %s", identity.Name)
|
||||
|
||||
// Check if user already exists in destination
|
||||
existingUser, err := toCM.GetUser(ctx, identity.Name)
|
||||
if err != nil && err != ErrUserNotFound {
|
||||
glog.Errorf("Failed to check if user %s exists in destination: %v", identity.Name, err)
|
||||
failed++
|
||||
continue
|
||||
}
|
||||
|
||||
if existingUser != nil {
|
||||
glog.Warningf("User %s already exists in destination store, skipping", identity.Name)
|
||||
continue
|
||||
}
|
||||
|
||||
// Create user in destination
|
||||
err = toCM.CreateUser(ctx, identity)
|
||||
if err != nil {
|
||||
glog.Errorf("Failed to create user %s in destination store: %v", identity.Name, err)
|
||||
failed++
|
||||
continue
|
||||
}
|
||||
|
||||
migrated++
|
||||
glog.V(1).Infof("Successfully migrated user: %s", identity.Name)
|
||||
}
|
||||
|
||||
glog.Infof("Migration completed: %d migrated, %d failed", migrated, failed)
|
||||
|
||||
if failed > 0 {
|
||||
return fmt.Errorf("migration completed with %d failures", failed)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ExportCredentials exports credentials from a store to a configuration
|
||||
func ExportCredentials(storeName CredentialStoreTypeName, configuration util.Configuration, prefix string) (*iam_pb.S3ApiConfiguration, error) {
|
||||
ctx := context.Background()
|
||||
|
||||
// Create credential manager
|
||||
cm, err := NewCredentialManager(storeName, configuration, prefix)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create credential manager (%s): %v", storeName, err)
|
||||
}
|
||||
defer cm.Shutdown()
|
||||
|
||||
// Load configuration
|
||||
config, err := cm.LoadConfiguration(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to load configuration: %w", err)
|
||||
}
|
||||
|
||||
return config, nil
|
||||
}
|
||||
|
||||
// ImportCredentials imports credentials from a configuration to a store
|
||||
func ImportCredentials(storeName CredentialStoreTypeName, configuration util.Configuration, prefix string, config *iam_pb.S3ApiConfiguration) error {
|
||||
ctx := context.Background()
|
||||
|
||||
// Create credential manager
|
||||
cm, err := NewCredentialManager(storeName, configuration, prefix)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create credential manager (%s): %v", storeName, err)
|
||||
}
|
||||
defer cm.Shutdown()
|
||||
|
||||
// Import each identity
|
||||
var imported, failed int
|
||||
for _, identity := range config.Identities {
|
||||
glog.V(1).Infof("Importing user: %s", identity.Name)
|
||||
|
||||
// Check if user already exists
|
||||
existingUser, err := cm.GetUser(ctx, identity.Name)
|
||||
if err != nil && err != ErrUserNotFound {
|
||||
glog.Errorf("Failed to check if user %s exists: %v", identity.Name, err)
|
||||
failed++
|
||||
continue
|
||||
}
|
||||
|
||||
if existingUser != nil {
|
||||
glog.Warningf("User %s already exists, skipping", identity.Name)
|
||||
continue
|
||||
}
|
||||
|
||||
// Create user
|
||||
err = cm.CreateUser(ctx, identity)
|
||||
if err != nil {
|
||||
glog.Errorf("Failed to create user %s: %v", identity.Name, err)
|
||||
failed++
|
||||
continue
|
||||
}
|
||||
|
||||
imported++
|
||||
glog.V(1).Infof("Successfully imported user: %s", identity.Name)
|
||||
}
|
||||
|
||||
glog.Infof("Import completed: %d imported, %d failed", imported, failed)
|
||||
|
||||
if failed > 0 {
|
||||
return fmt.Errorf("import completed with %d failures", failed)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateCredentials validates that all credentials in a store are accessible
|
||||
func ValidateCredentials(storeName CredentialStoreTypeName, configuration util.Configuration, prefix string) error {
|
||||
ctx := context.Background()
|
||||
|
||||
// Create credential manager
|
||||
cm, err := NewCredentialManager(storeName, configuration, prefix)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create credential manager (%s): %v", storeName, err)
|
||||
}
|
||||
defer cm.Shutdown()
|
||||
|
||||
// Load configuration
|
||||
config, err := cm.LoadConfiguration(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to load configuration: %w", err)
|
||||
}
|
||||
|
||||
if config == nil || len(config.Identities) == 0 {
|
||||
glog.Info("No identities found in store")
|
||||
return nil
|
||||
}
|
||||
|
||||
glog.Infof("Validating %d identities...", len(config.Identities))
|
||||
|
||||
// Validate each identity
|
||||
var validated, failed int
|
||||
for _, identity := range config.Identities {
|
||||
// Check if user can be retrieved
|
||||
user, err := cm.GetUser(ctx, identity.Name)
|
||||
if err != nil {
|
||||
glog.Errorf("Failed to retrieve user %s: %v", identity.Name, err)
|
||||
failed++
|
||||
continue
|
||||
}
|
||||
|
||||
if user == nil {
|
||||
glog.Errorf("User %s not found", identity.Name)
|
||||
failed++
|
||||
continue
|
||||
}
|
||||
|
||||
// Validate access keys
|
||||
for _, credential := range identity.Credentials {
|
||||
accessKeyUser, err := cm.GetUserByAccessKey(ctx, credential.AccessKey)
|
||||
if err != nil {
|
||||
glog.Errorf("Failed to retrieve user by access key %s: %v", credential.AccessKey, err)
|
||||
failed++
|
||||
continue
|
||||
}
|
||||
|
||||
if accessKeyUser == nil || accessKeyUser.Name != identity.Name {
|
||||
glog.Errorf("Access key %s does not map to correct user %s", credential.AccessKey, identity.Name)
|
||||
failed++
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
validated++
|
||||
glog.V(1).Infof("Successfully validated user: %s", identity.Name)
|
||||
}
|
||||
|
||||
glog.Infof("Validation completed: %d validated, %d failed", validated, failed)
|
||||
|
||||
if failed > 0 {
|
||||
return fmt.Errorf("validation completed with %d failures", failed)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -246,10 +246,6 @@ func NewLogFileEntryCollector(f *Filer, startPosition log_buffer.MessagePosition
|
||||
}
|
||||
}
|
||||
|
||||
func (c *LogFileEntryCollector) hasMore() bool {
|
||||
return c.dayEntryQueue.Len() > 0
|
||||
}
|
||||
|
||||
func (c *LogFileEntryCollector) collectMore(v *OrderedLogVisitor) (err error) {
|
||||
dayEntry := c.dayEntryQueue.Dequeue()
|
||||
if dayEntry == nil {
|
||||
|
||||
@@ -128,6 +128,11 @@ func (fsw *FilerStoreWrapper) Initialize(configuration util.Configuration, prefi
|
||||
}
|
||||
|
||||
func (fsw *FilerStoreWrapper) InsertEntry(ctx context.Context, entry *Entry) error {
|
||||
// Fail fast if the context is already cancelled to prevent orphaned metadata
|
||||
// when the originating request has been abandoned.
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
actualStore := fsw.getActualStore(entry.FullPath)
|
||||
stats.FilerStoreCounter.WithLabelValues(actualStore.GetName(), "insert").Inc()
|
||||
@@ -155,6 +160,9 @@ func (fsw *FilerStoreWrapper) InsertEntry(ctx context.Context, entry *Entry) err
|
||||
// InsertEntryKnownAbsent skips the pre-insert FindEntry path when the caller has
|
||||
// already established that the target path does not exist.
|
||||
func (fsw *FilerStoreWrapper) InsertEntryKnownAbsent(ctx context.Context, entry *Entry) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
actualStore := fsw.getActualStore(entry.FullPath)
|
||||
stats.FilerStoreCounter.WithLabelValues(actualStore.GetName(), "insert").Inc()
|
||||
@@ -178,6 +186,9 @@ func (fsw *FilerStoreWrapper) InsertEntryKnownAbsent(ctx context.Context, entry
|
||||
}
|
||||
|
||||
func (fsw *FilerStoreWrapper) UpdateEntry(ctx context.Context, entry *Entry) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
actualStore := fsw.getActualStore(entry.FullPath)
|
||||
stats.FilerStoreCounter.WithLabelValues(actualStore.GetName(), "update").Inc()
|
||||
@@ -236,6 +247,9 @@ func (fsw *FilerStoreWrapper) FindEntry(ctx context.Context, fp util.FullPath) (
|
||||
}
|
||||
|
||||
func (fsw *FilerStoreWrapper) DeleteEntry(ctx context.Context, fp util.FullPath) (err error) {
|
||||
if err = ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
actualStore := fsw.getActualStore(fp)
|
||||
stats.FilerStoreCounter.WithLabelValues(actualStore.GetName(), "delete").Inc()
|
||||
@@ -264,6 +278,9 @@ func (fsw *FilerStoreWrapper) DeleteEntry(ctx context.Context, fp util.FullPath)
|
||||
}
|
||||
|
||||
func (fsw *FilerStoreWrapper) DeleteOneEntry(ctx context.Context, existingEntry *Entry) (err error) {
|
||||
if err = ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
actualStore := fsw.getActualStore(existingEntry.FullPath)
|
||||
stats.FilerStoreCounter.WithLabelValues(actualStore.GetName(), "delete").Inc()
|
||||
@@ -288,6 +305,9 @@ func (fsw *FilerStoreWrapper) DeleteOneEntry(ctx context.Context, existingEntry
|
||||
}
|
||||
|
||||
func (fsw *FilerStoreWrapper) DeleteFolderChildren(ctx context.Context, fp util.FullPath) (err error) {
|
||||
if err = ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
actualStore := fsw.getActualStore(fp + "/")
|
||||
stats.FilerStoreCounter.WithLabelValues(actualStore.GetName(), "deleteFolderChildren").Inc()
|
||||
@@ -394,11 +414,17 @@ func (fsw *FilerStoreWrapper) prefixFilterEntries(ctx context.Context, dirPath u
|
||||
}
|
||||
|
||||
func (fsw *FilerStoreWrapper) BeginTransaction(ctx context.Context) (context.Context, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
return fsw.getDefaultStore().BeginTransaction(ctx)
|
||||
}
|
||||
|
||||
func (fsw *FilerStoreWrapper) CommitTransaction(ctx context.Context) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
return fsw.getDefaultStore().CommitTransaction(ctx)
|
||||
}
|
||||
@@ -413,6 +439,9 @@ func (fsw *FilerStoreWrapper) Shutdown() {
|
||||
}
|
||||
|
||||
func (fsw *FilerStoreWrapper) KvPut(ctx context.Context, key []byte, value []byte) (err error) {
|
||||
if err = ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
return fsw.getDefaultStore().KvPut(ctx, key, value)
|
||||
}
|
||||
@@ -421,6 +450,9 @@ func (fsw *FilerStoreWrapper) KvGet(ctx context.Context, key []byte) (value []by
|
||||
return fsw.getDefaultStore().KvGet(ctx, key)
|
||||
}
|
||||
func (fsw *FilerStoreWrapper) KvDelete(ctx context.Context, key []byte) (err error) {
|
||||
if err = ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
return fsw.getDefaultStore().KvDelete(ctx, key)
|
||||
}
|
||||
|
||||
@@ -2,8 +2,10 @@ package filer
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/util"
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -69,3 +71,136 @@ func TestFilerStoreWrapperMimeNormalization(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// cancelledCtx returns a context that is already cancelled.
|
||||
func cancelledCtx() context.Context {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
return ctx
|
||||
}
|
||||
|
||||
// expiredCtx returns a context whose deadline has already passed.
|
||||
func expiredCtx() context.Context {
|
||||
ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-time.Second))
|
||||
_ = cancel // already expired, but keep the cancel func from leaking
|
||||
return ctx
|
||||
}
|
||||
|
||||
func TestFilerStoreWrapperWriteOpsRejectCancelledContext(t *testing.T) {
|
||||
newEntry := func(path string) *Entry {
|
||||
return &Entry{
|
||||
FullPath: util.FullPath(path),
|
||||
Attr: Attr{Mode: 0o660, Mime: "application/octet-stream"},
|
||||
}
|
||||
}
|
||||
|
||||
// Each write operation that should be guarded.
|
||||
writeOps := []struct {
|
||||
name string
|
||||
run func(*FilerStoreWrapper, context.Context) error
|
||||
}{
|
||||
{"InsertEntry", func(fsw *FilerStoreWrapper, ctx context.Context) error {
|
||||
return fsw.InsertEntry(ctx, newEntry("/test/a"))
|
||||
}},
|
||||
{"InsertEntryKnownAbsent", func(fsw *FilerStoreWrapper, ctx context.Context) error {
|
||||
return fsw.InsertEntryKnownAbsent(ctx, newEntry("/test/b"))
|
||||
}},
|
||||
{"UpdateEntry", func(fsw *FilerStoreWrapper, ctx context.Context) error {
|
||||
_ = fsw.InsertEntry(context.Background(), newEntry("/test/c"))
|
||||
return fsw.UpdateEntry(ctx, newEntry("/test/c"))
|
||||
}},
|
||||
{"DeleteEntry", func(fsw *FilerStoreWrapper, ctx context.Context) error {
|
||||
_ = fsw.InsertEntry(context.Background(), newEntry("/test/d"))
|
||||
return fsw.DeleteEntry(ctx, "/test/d")
|
||||
}},
|
||||
{"DeleteOneEntry", func(fsw *FilerStoreWrapper, ctx context.Context) error {
|
||||
e := newEntry("/test/e")
|
||||
_ = fsw.InsertEntry(context.Background(), e)
|
||||
return fsw.DeleteOneEntry(ctx, e)
|
||||
}},
|
||||
{"DeleteFolderChildren", func(fsw *FilerStoreWrapper, ctx context.Context) error {
|
||||
_ = fsw.InsertEntry(context.Background(), newEntry("/test/folder/child"))
|
||||
return fsw.DeleteFolderChildren(ctx, "/test/folder")
|
||||
}},
|
||||
{"BeginTransaction", func(fsw *FilerStoreWrapper, ctx context.Context) error {
|
||||
_, err := fsw.BeginTransaction(ctx)
|
||||
return err
|
||||
}},
|
||||
{"CommitTransaction", func(fsw *FilerStoreWrapper, ctx context.Context) error {
|
||||
return fsw.CommitTransaction(ctx)
|
||||
}},
|
||||
{"KvPut", func(fsw *FilerStoreWrapper, ctx context.Context) error {
|
||||
return fsw.KvPut(ctx, []byte("k"), []byte("v"))
|
||||
}},
|
||||
{"KvDelete", func(fsw *FilerStoreWrapper, ctx context.Context) error {
|
||||
_ = fsw.KvPut(context.Background(), []byte("k"), []byte("v"))
|
||||
return fsw.KvDelete(ctx, []byte("k"))
|
||||
}},
|
||||
}
|
||||
|
||||
badContexts := []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
wantError error
|
||||
}{
|
||||
{"cancelled", cancelledCtx(), context.Canceled},
|
||||
{"deadline exceeded", expiredCtx(), context.DeadlineExceeded},
|
||||
}
|
||||
|
||||
for _, op := range writeOps {
|
||||
for _, bc := range badContexts {
|
||||
t.Run(op.name+"/"+bc.name, func(t *testing.T) {
|
||||
wrapper := NewFilerStoreWrapper(newStubFilerStore())
|
||||
err := op.run(wrapper, bc.ctx)
|
||||
require.Error(t, err)
|
||||
assert.True(t, errors.Is(err, bc.wantError), "got %v, want %v", err, bc.wantError)
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFilerStoreWrapperWriteOpsSucceedWithActiveContext(t *testing.T) {
|
||||
wrapper := NewFilerStoreWrapper(newStubFilerStore())
|
||||
ctx := context.Background()
|
||||
entry := &Entry{
|
||||
FullPath: util.FullPath("/test/obj"),
|
||||
Attr: Attr{Mode: 0o660},
|
||||
}
|
||||
|
||||
require.NoError(t, wrapper.InsertEntry(ctx, entry))
|
||||
require.NoError(t, wrapper.UpdateEntry(ctx, entry))
|
||||
require.NoError(t, wrapper.DeleteOneEntry(ctx, entry))
|
||||
require.NoError(t, wrapper.InsertEntryKnownAbsent(ctx, entry))
|
||||
require.NoError(t, wrapper.DeleteEntry(ctx, entry.FullPath))
|
||||
require.NoError(t, wrapper.KvPut(ctx, []byte("k"), []byte("v")))
|
||||
require.NoError(t, wrapper.KvDelete(ctx, []byte("k")))
|
||||
|
||||
txCtx, err := wrapper.BeginTransaction(ctx)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, wrapper.CommitTransaction(txCtx))
|
||||
}
|
||||
|
||||
func TestFilerStoreWrapperReadOpsSucceedWithCancelledContext(t *testing.T) {
|
||||
wrapper := NewFilerStoreWrapper(newStubFilerStore())
|
||||
entry := &Entry{
|
||||
FullPath: util.FullPath("/test/readable"),
|
||||
Attr: Attr{Mode: 0o660},
|
||||
}
|
||||
require.NoError(t, wrapper.InsertEntry(context.Background(), entry))
|
||||
require.NoError(t, wrapper.KvPut(context.Background(), []byte("rk"), []byte("rv")))
|
||||
|
||||
ctx := cancelledCtx()
|
||||
|
||||
_, err := wrapper.FindEntry(ctx, entry.FullPath)
|
||||
assert.NoError(t, err)
|
||||
|
||||
_, err = wrapper.KvGet(ctx, []byte("rk"))
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
// RollbackTransaction must succeed even when the context is cancelled,
|
||||
// because it is a cleanup operation called after failures.
|
||||
func TestFilerStoreWrapperRollbackSucceedsWithCancelledContext(t *testing.T) {
|
||||
wrapper := NewFilerStoreWrapper(newStubFilerStore())
|
||||
assert.NoError(t, wrapper.RollbackTransaction(cancelledCtx()))
|
||||
}
|
||||
|
||||
@@ -2,7 +2,6 @@ package filer
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/glog"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/filer_pb"
|
||||
@@ -36,39 +35,3 @@ func Replay(filerStore FilerStore, resp *filer_pb.SubscribeMetadataResponse) err
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ParallelProcessDirectoryStructure processes each entry in parallel, and also ensure parent directories are processed first.
|
||||
// This also assumes the parent directories are in the entryChan already.
|
||||
func ParallelProcessDirectoryStructure(entryChan chan *Entry, concurrency int, eachEntryFn func(entry *Entry) error) (firstErr error) {
|
||||
|
||||
executors := util.NewLimitedConcurrentExecutor(concurrency)
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for entry := range entryChan {
|
||||
wg.Add(1)
|
||||
if entry.IsDirectory() {
|
||||
func() {
|
||||
defer wg.Done()
|
||||
if err := eachEntryFn(entry); err != nil {
|
||||
if firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
}()
|
||||
} else {
|
||||
executors.Execute(func() {
|
||||
defer wg.Done()
|
||||
if err := eachEntryFn(entry); err != nil {
|
||||
if firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
if firstErr != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
wg.Wait()
|
||||
return
|
||||
}
|
||||
|
||||
@@ -173,14 +173,16 @@ func (store *UniversalRedisStore) ListDirectoryEntries(ctx context.Context, dirP
|
||||
members = members[:limit]
|
||||
}
|
||||
|
||||
var entry *filer.Entry
|
||||
// fetch entry meta
|
||||
for _, fileName := range members {
|
||||
path := util.NewFullPath(string(dirPath), fileName)
|
||||
entry, err := store.FindEntry(ctx, path)
|
||||
entry, err = store.FindEntry(ctx, path)
|
||||
lastFileName = fileName
|
||||
if err != nil {
|
||||
glog.V(0).InfofCtx(ctx, "list %s : %v", path, err)
|
||||
if err == filer_pb.ErrNotFound {
|
||||
err = nil
|
||||
continue
|
||||
}
|
||||
} else {
|
||||
|
||||
@@ -16,15 +16,6 @@ type ItemList struct {
|
||||
prefix string
|
||||
}
|
||||
|
||||
func newItemList(client redis.UniversalClient, prefix string, store skiplist.ListStore, batchSize int) *ItemList {
|
||||
return &ItemList{
|
||||
skipList: skiplist.New(store),
|
||||
batchSize: batchSize,
|
||||
client: client,
|
||||
prefix: prefix,
|
||||
}
|
||||
}
|
||||
|
||||
/*
|
||||
Be reluctant to create new nodes. Try to fit into either previous node or next node.
|
||||
Prefer to add to previous node.
|
||||
|
||||
@@ -1,48 +0,0 @@
|
||||
package redis_lua
|
||||
|
||||
import (
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/seaweedfs/seaweedfs/weed/filer"
|
||||
"github.com/seaweedfs/seaweedfs/weed/util"
|
||||
)
|
||||
|
||||
func init() {
|
||||
filer.Stores = append(filer.Stores, &RedisLuaClusterStore{})
|
||||
}
|
||||
|
||||
type RedisLuaClusterStore struct {
|
||||
UniversalRedisLuaStore
|
||||
}
|
||||
|
||||
func (store *RedisLuaClusterStore) GetName() string {
|
||||
return "redis_lua_cluster"
|
||||
}
|
||||
|
||||
func (store *RedisLuaClusterStore) Initialize(configuration util.Configuration, prefix string) (err error) {
|
||||
|
||||
configuration.SetDefault(prefix+"useReadOnly", false)
|
||||
configuration.SetDefault(prefix+"routeByLatency", false)
|
||||
|
||||
return store.initialize(
|
||||
configuration.GetStringSlice(prefix+"addresses"),
|
||||
configuration.GetString(prefix+"username"),
|
||||
configuration.GetString(prefix+"password"),
|
||||
configuration.GetString(prefix+"keyPrefix"),
|
||||
configuration.GetBool(prefix+"useReadOnly"),
|
||||
configuration.GetBool(prefix+"routeByLatency"),
|
||||
configuration.GetStringSlice(prefix+"superLargeDirectories"),
|
||||
)
|
||||
}
|
||||
|
||||
func (store *RedisLuaClusterStore) initialize(addresses []string, username string, password string, keyPrefix string, readOnly, routeByLatency bool, superLargeDirectories []string) (err error) {
|
||||
store.Client = redis.NewClusterClient(&redis.ClusterOptions{
|
||||
Addrs: addresses,
|
||||
Username: username,
|
||||
Password: password,
|
||||
ReadOnly: readOnly,
|
||||
RouteByLatency: routeByLatency,
|
||||
})
|
||||
store.keyPrefix = keyPrefix
|
||||
store.loadSuperLargeDirectories(superLargeDirectories)
|
||||
return
|
||||
}
|
||||
@@ -1,48 +0,0 @@
|
||||
package redis_lua
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/seaweedfs/seaweedfs/weed/filer"
|
||||
"github.com/seaweedfs/seaweedfs/weed/util"
|
||||
)
|
||||
|
||||
func init() {
|
||||
filer.Stores = append(filer.Stores, &RedisLuaSentinelStore{})
|
||||
}
|
||||
|
||||
type RedisLuaSentinelStore struct {
|
||||
UniversalRedisLuaStore
|
||||
}
|
||||
|
||||
func (store *RedisLuaSentinelStore) GetName() string {
|
||||
return "redis_lua_sentinel"
|
||||
}
|
||||
|
||||
func (store *RedisLuaSentinelStore) Initialize(configuration util.Configuration, prefix string) (err error) {
|
||||
return store.initialize(
|
||||
configuration.GetStringSlice(prefix+"addresses"),
|
||||
configuration.GetString(prefix+"masterName"),
|
||||
configuration.GetString(prefix+"username"),
|
||||
configuration.GetString(prefix+"password"),
|
||||
configuration.GetInt(prefix+"database"),
|
||||
configuration.GetString(prefix+"keyPrefix"),
|
||||
)
|
||||
}
|
||||
|
||||
func (store *RedisLuaSentinelStore) initialize(addresses []string, masterName string, username string, password string, database int, keyPrefix string) (err error) {
|
||||
store.Client = redis.NewFailoverClient(&redis.FailoverOptions{
|
||||
MasterName: masterName,
|
||||
SentinelAddrs: addresses,
|
||||
Username: username,
|
||||
Password: password,
|
||||
DB: database,
|
||||
MinRetryBackoff: time.Millisecond * 100,
|
||||
MaxRetryBackoff: time.Minute * 1,
|
||||
ReadTimeout: time.Second * 30,
|
||||
WriteTimeout: time.Second * 5,
|
||||
})
|
||||
store.keyPrefix = keyPrefix
|
||||
return
|
||||
}
|
||||
@@ -1,42 +0,0 @@
|
||||
package redis_lua
|
||||
|
||||
import (
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/seaweedfs/seaweedfs/weed/filer"
|
||||
"github.com/seaweedfs/seaweedfs/weed/util"
|
||||
)
|
||||
|
||||
func init() {
|
||||
filer.Stores = append(filer.Stores, &RedisLuaStore{})
|
||||
}
|
||||
|
||||
type RedisLuaStore struct {
|
||||
UniversalRedisLuaStore
|
||||
}
|
||||
|
||||
func (store *RedisLuaStore) GetName() string {
|
||||
return "redis_lua"
|
||||
}
|
||||
|
||||
func (store *RedisLuaStore) Initialize(configuration util.Configuration, prefix string) (err error) {
|
||||
return store.initialize(
|
||||
configuration.GetString(prefix+"address"),
|
||||
configuration.GetString(prefix+"username"),
|
||||
configuration.GetString(prefix+"password"),
|
||||
configuration.GetInt(prefix+"database"),
|
||||
configuration.GetString(prefix+"keyPrefix"),
|
||||
configuration.GetStringSlice(prefix+"superLargeDirectories"),
|
||||
)
|
||||
}
|
||||
|
||||
func (store *RedisLuaStore) initialize(hostPort string, username string, password string, database int, keyPrefix string, superLargeDirectories []string) (err error) {
|
||||
store.Client = redis.NewClient(&redis.Options{
|
||||
Addr: hostPort,
|
||||
Username: username,
|
||||
Password: password,
|
||||
DB: database,
|
||||
})
|
||||
store.keyPrefix = keyPrefix
|
||||
store.loadSuperLargeDirectories(superLargeDirectories)
|
||||
return
|
||||
}
|
||||
@@ -1,19 +0,0 @@
|
||||
-- KEYS[1]: full path of entry
|
||||
local fullpath = KEYS[1]
|
||||
-- KEYS[2]: full path of entry
|
||||
local fullpath_list_key = KEYS[2]
|
||||
-- KEYS[3]: dir of the entry
|
||||
local dir_list_key = KEYS[3]
|
||||
|
||||
-- ARGV[1]: isSuperLargeDirectory
|
||||
local isSuperLargeDirectory = ARGV[1] == "1"
|
||||
-- ARGV[2]: name of the entry
|
||||
local name = ARGV[2]
|
||||
|
||||
redis.call("DEL", fullpath, fullpath_list_key)
|
||||
|
||||
if not isSuperLargeDirectory and name ~= "" then
|
||||
redis.call("ZREM", dir_list_key, name)
|
||||
end
|
||||
|
||||
return 0
|
||||
@@ -1,15 +0,0 @@
|
||||
-- KEYS[1]: full path of entry
|
||||
local fullpath = KEYS[1]
|
||||
|
||||
if fullpath ~= "" and string.sub(fullpath, -1) == "/" then
|
||||
fullpath = string.sub(fullpath, 0, -2)
|
||||
end
|
||||
|
||||
local files = redis.call("ZRANGE", fullpath .. "\0", "0", "-1")
|
||||
|
||||
for _, name in ipairs(files) do
|
||||
local file_path = fullpath .. "/" .. name
|
||||
redis.call("DEL", file_path, file_path .. "\0")
|
||||
end
|
||||
|
||||
return 0
|
||||
@@ -1,25 +0,0 @@
|
||||
package stored_procedure
|
||||
|
||||
import (
|
||||
_ "embed"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
func init() {
|
||||
InsertEntryScript = redis.NewScript(insertEntry)
|
||||
DeleteEntryScript = redis.NewScript(deleteEntry)
|
||||
DeleteFolderChildrenScript = redis.NewScript(deleteFolderChildren)
|
||||
}
|
||||
|
||||
//go:embed insert_entry.lua
|
||||
var insertEntry string
|
||||
var InsertEntryScript *redis.Script
|
||||
|
||||
//go:embed delete_entry.lua
|
||||
var deleteEntry string
|
||||
var DeleteEntryScript *redis.Script
|
||||
|
||||
//go:embed delete_folder_children.lua
|
||||
var deleteFolderChildren string
|
||||
var DeleteFolderChildrenScript *redis.Script
|
||||
@@ -1,27 +0,0 @@
|
||||
-- KEYS[1]: full path of entry
|
||||
local full_path = KEYS[1]
|
||||
-- KEYS[2]: dir of the entry
|
||||
local dir_list_key = KEYS[2]
|
||||
|
||||
-- ARGV[1]: content of the entry
|
||||
local entry = ARGV[1]
|
||||
-- ARGV[2]: TTL of the entry
|
||||
local ttlSec = tonumber(ARGV[2])
|
||||
-- ARGV[3]: isSuperLargeDirectory
|
||||
local isSuperLargeDirectory = ARGV[3] == "1"
|
||||
-- ARGV[4]: zscore of the entry in zset
|
||||
local zscore = tonumber(ARGV[4])
|
||||
-- ARGV[5]: name of the entry
|
||||
local name = ARGV[5]
|
||||
|
||||
if ttlSec > 0 then
|
||||
redis.call("SET", full_path, entry, "EX", ttlSec)
|
||||
else
|
||||
redis.call("SET", full_path, entry)
|
||||
end
|
||||
|
||||
if not isSuperLargeDirectory and name ~= "" then
|
||||
redis.call("ZADD", dir_list_key, "NX", zscore, name)
|
||||
end
|
||||
|
||||
return 0
|
||||
@@ -1,206 +0,0 @@
|
||||
package redis_lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/filer"
|
||||
"github.com/seaweedfs/seaweedfs/weed/filer/redis_lua/stored_procedure"
|
||||
"github.com/seaweedfs/seaweedfs/weed/glog"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/filer_pb"
|
||||
"github.com/seaweedfs/seaweedfs/weed/util"
|
||||
)
|
||||
|
||||
const (
|
||||
DIR_LIST_MARKER = "\x00"
|
||||
)
|
||||
|
||||
type UniversalRedisLuaStore struct {
|
||||
Client redis.UniversalClient
|
||||
keyPrefix string
|
||||
superLargeDirectoryHash map[string]bool
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) isSuperLargeDirectory(dir string) (isSuperLargeDirectory bool) {
|
||||
_, isSuperLargeDirectory = store.superLargeDirectoryHash[dir]
|
||||
return
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) loadSuperLargeDirectories(superLargeDirectories []string) {
|
||||
// set directory hash
|
||||
store.superLargeDirectoryHash = make(map[string]bool)
|
||||
for _, dir := range superLargeDirectories {
|
||||
store.superLargeDirectoryHash[dir] = true
|
||||
}
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) getKey(key string) string {
|
||||
if store.keyPrefix == "" {
|
||||
return key
|
||||
}
|
||||
return store.keyPrefix + key
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) BeginTransaction(ctx context.Context) (context.Context, error) {
|
||||
return ctx, nil
|
||||
}
|
||||
func (store *UniversalRedisLuaStore) CommitTransaction(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
func (store *UniversalRedisLuaStore) RollbackTransaction(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) InsertEntry(ctx context.Context, entry *filer.Entry) (err error) {
|
||||
|
||||
value, err := entry.EncodeAttributesAndChunks()
|
||||
if err != nil {
|
||||
return fmt.Errorf("encoding %s %+v: %v", entry.FullPath, entry.Attr, err)
|
||||
}
|
||||
|
||||
if len(entry.GetChunks()) > filer.CountEntryChunksForGzip {
|
||||
value = util.MaybeGzipData(value)
|
||||
}
|
||||
|
||||
dir, name := entry.FullPath.DirAndName()
|
||||
|
||||
err = stored_procedure.InsertEntryScript.Run(ctx, store.Client,
|
||||
[]string{store.getKey(string(entry.FullPath)), store.getKey(genDirectoryListKey(dir))},
|
||||
value, entry.TtlSec,
|
||||
store.isSuperLargeDirectory(dir), 0, name).Err()
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("persisting %s : %v", entry.FullPath, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) UpdateEntry(ctx context.Context, entry *filer.Entry) (err error) {
|
||||
|
||||
return store.InsertEntry(ctx, entry)
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) FindEntry(ctx context.Context, fullpath util.FullPath) (entry *filer.Entry, err error) {
|
||||
|
||||
data, err := store.Client.Get(ctx, store.getKey(string(fullpath))).Result()
|
||||
if err == redis.Nil {
|
||||
return nil, filer_pb.ErrNotFound
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get %s : %v", fullpath, err)
|
||||
}
|
||||
|
||||
entry = &filer.Entry{
|
||||
FullPath: fullpath,
|
||||
}
|
||||
err = entry.DecodeAttributesAndChunks(util.MaybeDecompressData([]byte(data)))
|
||||
if err != nil {
|
||||
return entry, fmt.Errorf("decode %s : %v", entry.FullPath, err)
|
||||
}
|
||||
|
||||
return entry, nil
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) DeleteEntry(ctx context.Context, fullpath util.FullPath) (err error) {
|
||||
|
||||
dir, name := fullpath.DirAndName()
|
||||
|
||||
err = stored_procedure.DeleteEntryScript.Run(ctx, store.Client,
|
||||
[]string{store.getKey(string(fullpath)), store.getKey(genDirectoryListKey(string(fullpath))), store.getKey(genDirectoryListKey(dir))},
|
||||
store.isSuperLargeDirectory(dir), name).Err()
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("DeleteEntry %s : %v", fullpath, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) DeleteFolderChildren(ctx context.Context, fullpath util.FullPath) (err error) {
|
||||
|
||||
if store.isSuperLargeDirectory(string(fullpath)) {
|
||||
return nil
|
||||
}
|
||||
|
||||
err = stored_procedure.DeleteFolderChildrenScript.Run(ctx, store.Client,
|
||||
[]string{store.getKey(string(fullpath))}).Err()
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("DeleteFolderChildren %s : %v", fullpath, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) ListDirectoryPrefixedEntries(ctx context.Context, dirPath util.FullPath, startFileName string, includeStartFile bool, limit int64, prefix string, eachEntryFunc filer.ListEachEntryFunc) (lastFileName string, err error) {
|
||||
return lastFileName, filer.ErrUnsupportedListDirectoryPrefixed
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) ListDirectoryEntries(ctx context.Context, dirPath util.FullPath, startFileName string, includeStartFile bool, limit int64, eachEntryFunc filer.ListEachEntryFunc) (lastFileName string, err error) {
|
||||
|
||||
dirListKey := store.getKey(genDirectoryListKey(string(dirPath)))
|
||||
|
||||
min := "-"
|
||||
if startFileName != "" {
|
||||
if includeStartFile {
|
||||
min = "[" + startFileName
|
||||
} else {
|
||||
min = "(" + startFileName
|
||||
}
|
||||
}
|
||||
|
||||
members, err := store.Client.ZRangeByLex(ctx, dirListKey, &redis.ZRangeBy{
|
||||
Min: min,
|
||||
Max: "+",
|
||||
Offset: 0,
|
||||
Count: limit,
|
||||
}).Result()
|
||||
if err != nil {
|
||||
return lastFileName, fmt.Errorf("list %s : %v", dirPath, err)
|
||||
}
|
||||
|
||||
// fetch entry meta
|
||||
for _, fileName := range members {
|
||||
path := util.NewFullPath(string(dirPath), fileName)
|
||||
entry, err := store.FindEntry(ctx, path)
|
||||
lastFileName = fileName
|
||||
if err != nil {
|
||||
glog.V(0).InfofCtx(ctx, "list %s : %v", path, err)
|
||||
if err == filer_pb.ErrNotFound {
|
||||
continue
|
||||
}
|
||||
} else {
|
||||
if entry.TtlSec > 0 {
|
||||
if entry.Attr.Crtime.Add(time.Duration(entry.TtlSec) * time.Second).Before(time.Now()) {
|
||||
store.DeleteEntry(ctx, path)
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
resEachEntryFunc, resEachEntryFuncErr := eachEntryFunc(entry)
|
||||
if resEachEntryFuncErr != nil {
|
||||
err = fmt.Errorf("failed to process eachEntryFunc: %w", resEachEntryFuncErr)
|
||||
break
|
||||
}
|
||||
|
||||
if !resEachEntryFunc {
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return lastFileName, err
|
||||
}
|
||||
|
||||
func genDirectoryListKey(dir string) (dirList string) {
|
||||
return dir + DIR_LIST_MARKER
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) Shutdown() {
|
||||
store.Client.Close()
|
||||
}
|
||||
@@ -1,42 +0,0 @@
|
||||
package redis_lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/seaweedfs/seaweedfs/weed/filer"
|
||||
)
|
||||
|
||||
func (store *UniversalRedisLuaStore) KvPut(ctx context.Context, key []byte, value []byte) (err error) {
|
||||
|
||||
_, err = store.Client.Set(ctx, string(key), value, 0).Result()
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("kv put: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) KvGet(ctx context.Context, key []byte) (value []byte, err error) {
|
||||
|
||||
data, err := store.Client.Get(ctx, string(key)).Result()
|
||||
|
||||
if err == redis.Nil {
|
||||
return nil, filer.ErrKvNotFound
|
||||
}
|
||||
|
||||
return []byte(data), err
|
||||
}
|
||||
|
||||
func (store *UniversalRedisLuaStore) KvDelete(ctx context.Context, key []byte) (err error) {
|
||||
|
||||
_, err = store.Client.Del(ctx, string(key)).Result()
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("kv delete: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -102,10 +102,6 @@ func PrepareStreamContent(masterClient wdclient.HasLookupFileIdFunction, jwtFunc
|
||||
|
||||
type VolumeServerJwtFunction func(fileId string) string
|
||||
|
||||
func noJwtFunc(string) string {
|
||||
return ""
|
||||
}
|
||||
|
||||
type CacheInvalidator interface {
|
||||
InvalidateCache(fileId string)
|
||||
}
|
||||
@@ -276,33 +272,6 @@ func writeZero(w io.Writer, size int64) (err error) {
|
||||
return
|
||||
}
|
||||
|
||||
func ReadAll(ctx context.Context, buffer []byte, masterClient *wdclient.MasterClient, chunks []*filer_pb.FileChunk) error {
|
||||
|
||||
lookupFileIdFn := func(ctx context.Context, fileId string) (targetUrls []string, err error) {
|
||||
return masterClient.LookupFileId(ctx, fileId)
|
||||
}
|
||||
|
||||
chunkViews := ViewFromChunks(ctx, lookupFileIdFn, chunks, 0, int64(len(buffer)))
|
||||
|
||||
idx := 0
|
||||
|
||||
for x := chunkViews.Front(); x != nil; x = x.Next {
|
||||
chunkView := x.Value
|
||||
urlStrings, err := lookupFileIdFn(ctx, chunkView.FileId)
|
||||
if err != nil {
|
||||
glog.V(1).InfofCtx(ctx, "operation LookupFileId %s failed, err: %v", chunkView.FileId, err)
|
||||
return err
|
||||
}
|
||||
|
||||
n, err := util_http.RetriedFetchChunkData(ctx, buffer[idx:idx+int(chunkView.ViewSize)], urlStrings, chunkView.CipherKey, chunkView.IsGzipped, chunkView.IsFullChunk(), chunkView.OffsetInChunk, chunkView.FileId)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
idx += n
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ---------------- ChunkStreamReader ----------------------------------
|
||||
type ChunkStreamReader struct {
|
||||
head *Interval[*ChunkView]
|
||||
|
||||
@@ -1,281 +0,0 @@
|
||||
package filer
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/filer_pb"
|
||||
"github.com/seaweedfs/seaweedfs/weed/wdclient"
|
||||
)
|
||||
|
||||
// mockMasterClient implements HasLookupFileIdFunction and CacheInvalidator
|
||||
type mockMasterClient struct {
|
||||
lookupFunc func(ctx context.Context, fileId string) ([]string, error)
|
||||
invalidatedFileIds []string
|
||||
}
|
||||
|
||||
func (m *mockMasterClient) GetLookupFileIdFunction() wdclient.LookupFileIdFunctionType {
|
||||
return m.lookupFunc
|
||||
}
|
||||
|
||||
func (m *mockMasterClient) InvalidateCache(fileId string) {
|
||||
m.invalidatedFileIds = append(m.invalidatedFileIds, fileId)
|
||||
}
|
||||
|
||||
// Test urlSlicesEqual helper function
|
||||
func TestUrlSlicesEqual(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
a []string
|
||||
b []string
|
||||
expected bool
|
||||
}{
|
||||
{
|
||||
name: "identical slices",
|
||||
a: []string{"http://server1", "http://server2"},
|
||||
b: []string{"http://server1", "http://server2"},
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "same URLs different order",
|
||||
a: []string{"http://server1", "http://server2"},
|
||||
b: []string{"http://server2", "http://server1"},
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "different URLs",
|
||||
a: []string{"http://server1", "http://server2"},
|
||||
b: []string{"http://server1", "http://server3"},
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "different lengths",
|
||||
a: []string{"http://server1"},
|
||||
b: []string{"http://server1", "http://server2"},
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "empty slices",
|
||||
a: []string{},
|
||||
b: []string{},
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "duplicates in both",
|
||||
a: []string{"http://server1", "http://server1"},
|
||||
b: []string{"http://server1", "http://server1"},
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "different duplicate counts",
|
||||
a: []string{"http://server1", "http://server1"},
|
||||
b: []string{"http://server1", "http://server2"},
|
||||
expected: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := urlSlicesEqual(tt.a, tt.b)
|
||||
if result != tt.expected {
|
||||
t.Errorf("urlSlicesEqual(%v, %v) = %v; want %v", tt.a, tt.b, result, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Test cache invalidation when read fails
|
||||
func TestStreamContentWithCacheInvalidation(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
fileId := "3,01234567890"
|
||||
|
||||
callCount := 0
|
||||
oldUrls := []string{"http://failed-server:8080"}
|
||||
newUrls := []string{"http://working-server:8080"}
|
||||
|
||||
mock := &mockMasterClient{
|
||||
lookupFunc: func(ctx context.Context, fid string) ([]string, error) {
|
||||
callCount++
|
||||
if callCount == 1 {
|
||||
// First call returns failing server
|
||||
return oldUrls, nil
|
||||
}
|
||||
// After invalidation, return working server
|
||||
return newUrls, nil
|
||||
},
|
||||
}
|
||||
|
||||
// Create a simple chunk
|
||||
chunks := []*filer_pb.FileChunk{
|
||||
{
|
||||
FileId: fileId,
|
||||
Offset: 0,
|
||||
Size: 10,
|
||||
},
|
||||
}
|
||||
|
||||
streamFn, err := PrepareStreamContentWithThrottler(ctx, mock, noJwtFunc, chunks, 0, 10, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("PrepareStreamContentWithThrottler failed: %v", err)
|
||||
}
|
||||
|
||||
// Note: This test can't fully execute streamFn because it would require actual HTTP servers
|
||||
// However, we can verify the setup was created correctly
|
||||
if streamFn == nil {
|
||||
t.Fatal("Expected non-nil stream function")
|
||||
}
|
||||
|
||||
// Verify the lookup was called
|
||||
if callCount != 1 {
|
||||
t.Errorf("Expected 1 lookup call, got %d", callCount)
|
||||
}
|
||||
}
|
||||
|
||||
// Test that InvalidateCache is called on read failure
|
||||
func TestCacheInvalidationInterface(t *testing.T) {
|
||||
mock := &mockMasterClient{
|
||||
lookupFunc: func(ctx context.Context, fileId string) ([]string, error) {
|
||||
return []string{"http://server:8080"}, nil
|
||||
},
|
||||
}
|
||||
|
||||
fileId := "3,test123"
|
||||
|
||||
// Simulate invalidation
|
||||
if invalidator, ok := interface{}(mock).(CacheInvalidator); ok {
|
||||
invalidator.InvalidateCache(fileId)
|
||||
} else {
|
||||
t.Fatal("mockMasterClient should implement CacheInvalidator")
|
||||
}
|
||||
|
||||
// Check that the file ID was recorded as invalidated
|
||||
if len(mock.invalidatedFileIds) != 1 {
|
||||
t.Fatalf("Expected 1 invalidated file ID, got %d", len(mock.invalidatedFileIds))
|
||||
}
|
||||
if mock.invalidatedFileIds[0] != fileId {
|
||||
t.Errorf("Expected invalidated file ID %s, got %s", fileId, mock.invalidatedFileIds[0])
|
||||
}
|
||||
}
|
||||
|
||||
// Test retry logic doesn't retry with same URLs
|
||||
func TestRetryLogicSkipsSameUrls(t *testing.T) {
|
||||
// This test verifies that the urlSlicesEqual check prevents infinite retries
|
||||
sameUrls := []string{"http://server1:8080", "http://server2:8080"}
|
||||
differentUrls := []string{"http://server3:8080", "http://server4:8080"}
|
||||
|
||||
// Same URLs should return true (and thus skip retry)
|
||||
if !urlSlicesEqual(sameUrls, sameUrls) {
|
||||
t.Error("Expected same URLs to be equal")
|
||||
}
|
||||
|
||||
// Different URLs should return false (and thus allow retry)
|
||||
if urlSlicesEqual(sameUrls, differentUrls) {
|
||||
t.Error("Expected different URLs to not be equal")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCanceledStreamSkipsCacheInvalidation(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
fileId := "3,canceled"
|
||||
|
||||
mock := &mockMasterClient{
|
||||
lookupFunc: func(ctx context.Context, fid string) ([]string, error) {
|
||||
return []string{"http://server:8080"}, nil
|
||||
},
|
||||
}
|
||||
|
||||
chunks := []*filer_pb.FileChunk{
|
||||
{
|
||||
FileId: fileId,
|
||||
Offset: 0,
|
||||
Size: 10,
|
||||
},
|
||||
}
|
||||
|
||||
streamFn, err := PrepareStreamContentWithThrottler(ctx, mock, noJwtFunc, chunks, 0, 10, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("PrepareStreamContentWithThrottler failed: %v", err)
|
||||
}
|
||||
|
||||
cancel()
|
||||
|
||||
err = streamFn(&bytes.Buffer{})
|
||||
if err != context.Canceled {
|
||||
t.Fatalf("expected context.Canceled, got %v", err)
|
||||
}
|
||||
if len(mock.invalidatedFileIds) != 0 {
|
||||
t.Fatalf("expected no cache invalidation on cancellation, got %v", mock.invalidatedFileIds)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrepareStreamContentSkipsLookupWhenContextAlreadyCanceled(t *testing.T) {
|
||||
oldSchedule := getLookupFileIdBackoffSchedule
|
||||
getLookupFileIdBackoffSchedule = []time.Duration{time.Millisecond}
|
||||
t.Cleanup(func() {
|
||||
getLookupFileIdBackoffSchedule = oldSchedule
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
lookupCalls := 0
|
||||
mock := &mockMasterClient{
|
||||
lookupFunc: func(ctx context.Context, fileId string) ([]string, error) {
|
||||
lookupCalls++
|
||||
return nil, errors.New("lookup should not run")
|
||||
},
|
||||
}
|
||||
|
||||
chunks := []*filer_pb.FileChunk{
|
||||
{
|
||||
FileId: "3,precanceled",
|
||||
Offset: 0,
|
||||
Size: 10,
|
||||
},
|
||||
}
|
||||
|
||||
_, err := PrepareStreamContentWithThrottler(ctx, mock, noJwtFunc, chunks, 0, 10, 0)
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context.Canceled, got %v", err)
|
||||
}
|
||||
if lookupCalls != 0 {
|
||||
t.Fatalf("expected no lookup calls after cancellation, got %d", lookupCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrepareStreamContentStopsLookupRetriesAfterContextCancellation(t *testing.T) {
|
||||
oldSchedule := getLookupFileIdBackoffSchedule
|
||||
getLookupFileIdBackoffSchedule = []time.Duration{time.Millisecond, time.Millisecond, time.Millisecond}
|
||||
t.Cleanup(func() {
|
||||
getLookupFileIdBackoffSchedule = oldSchedule
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
lookupCalls := 0
|
||||
mock := &mockMasterClient{
|
||||
lookupFunc: func(ctx context.Context, fileId string) ([]string, error) {
|
||||
lookupCalls++
|
||||
cancel()
|
||||
return nil, context.Canceled
|
||||
},
|
||||
}
|
||||
|
||||
chunks := []*filer_pb.FileChunk{
|
||||
{
|
||||
FileId: "3,cancel-during-lookup",
|
||||
Offset: 0,
|
||||
Size: 10,
|
||||
},
|
||||
}
|
||||
|
||||
_, err := PrepareStreamContentWithThrottler(ctx, mock, noJwtFunc, chunks, 0, 10, 0)
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context.Canceled, got %v", err)
|
||||
}
|
||||
if lookupCalls != 1 {
|
||||
t.Fatalf("expected lookup retries to stop after cancellation, got %d calls", lookupCalls)
|
||||
}
|
||||
}
|
||||
@@ -37,11 +37,6 @@ func GenerateRandomString(length int, charset string) (string, error) {
|
||||
return string(b), nil
|
||||
}
|
||||
|
||||
// GenerateAccessKeyId generates a new access key ID.
|
||||
func GenerateAccessKeyId() (string, error) {
|
||||
return GenerateRandomString(AccessKeyIdLength, CharsetUpper)
|
||||
}
|
||||
|
||||
// GenerateSecretAccessKey generates a new secret access key.
|
||||
func GenerateSecretAccessKey() (string, error) {
|
||||
return GenerateRandomString(SecretAccessKeyLength, Charset)
|
||||
@@ -179,11 +174,3 @@ func MapToIdentitiesAction(action string) string {
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
// MaskAccessKey masks an access key for logging, showing only the first 4 characters.
|
||||
func MaskAccessKey(accessKeyId string) string {
|
||||
if len(accessKeyId) > 4 {
|
||||
return accessKeyId[:4] + "***"
|
||||
}
|
||||
return accessKeyId
|
||||
}
|
||||
|
||||
@@ -1,164 +0,0 @@
|
||||
package iam
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/s3api/s3_constants"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestHash(t *testing.T) {
|
||||
input := "test"
|
||||
result := Hash(&input)
|
||||
assert.NotEmpty(t, result)
|
||||
assert.Len(t, result, 40) // SHA1 hex is 40 chars
|
||||
|
||||
// Same input should produce same hash
|
||||
result2 := Hash(&input)
|
||||
assert.Equal(t, result, result2)
|
||||
|
||||
// Different input should produce different hash
|
||||
different := "different"
|
||||
result3 := Hash(&different)
|
||||
assert.NotEqual(t, result, result3)
|
||||
}
|
||||
|
||||
func TestGenerateRandomString(t *testing.T) {
|
||||
// Valid generation
|
||||
result, err := GenerateRandomString(10, CharsetUpper)
|
||||
assert.NoError(t, err)
|
||||
assert.Len(t, result, 10)
|
||||
|
||||
// Different calls should produce different results (with high probability)
|
||||
result2, err := GenerateRandomString(10, CharsetUpper)
|
||||
assert.NoError(t, err)
|
||||
assert.NotEqual(t, result, result2)
|
||||
|
||||
// Invalid length
|
||||
_, err = GenerateRandomString(0, CharsetUpper)
|
||||
assert.Error(t, err)
|
||||
|
||||
_, err = GenerateRandomString(-1, CharsetUpper)
|
||||
assert.Error(t, err)
|
||||
|
||||
// Empty charset
|
||||
_, err = GenerateRandomString(10, "")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestGenerateAccessKeyId(t *testing.T) {
|
||||
keyId, err := GenerateAccessKeyId()
|
||||
assert.NoError(t, err)
|
||||
assert.Len(t, keyId, AccessKeyIdLength)
|
||||
}
|
||||
|
||||
func TestGenerateSecretAccessKey(t *testing.T) {
|
||||
secretKey, err := GenerateSecretAccessKey()
|
||||
assert.NoError(t, err)
|
||||
assert.Len(t, secretKey, SecretAccessKeyLength)
|
||||
}
|
||||
|
||||
func TestGenerateSecretAccessKey_URLSafe(t *testing.T) {
|
||||
// Generate multiple keys to increase probability of catching unsafe chars
|
||||
for i := 0; i < 100; i++ {
|
||||
secretKey, err := GenerateSecretAccessKey()
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify no URL-unsafe characters that would cause authentication issues
|
||||
assert.NotContains(t, secretKey, "/", "Secret key should not contain /")
|
||||
assert.NotContains(t, secretKey, "+", "Secret key should not contain +")
|
||||
|
||||
// Verify only expected characters are present
|
||||
for _, char := range secretKey {
|
||||
assert.Contains(t, Charset, string(char), "Secret key contains unexpected character: %c", char)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStringSlicesEqual(t *testing.T) {
|
||||
tests := []struct {
|
||||
a []string
|
||||
b []string
|
||||
expected bool
|
||||
}{
|
||||
{[]string{"a", "b", "c"}, []string{"a", "b", "c"}, true},
|
||||
{[]string{"c", "b", "a"}, []string{"a", "b", "c"}, true}, // Order independent
|
||||
{[]string{"a", "b"}, []string{"a", "b", "c"}, false},
|
||||
{[]string{}, []string{}, true},
|
||||
{nil, nil, true},
|
||||
{[]string{"a"}, []string{"b"}, false},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
result := StringSlicesEqual(test.a, test.b)
|
||||
assert.Equal(t, test.expected, result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapToStatementAction(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{StatementActionAdmin, s3_constants.ACTION_ADMIN},
|
||||
{StatementActionWrite, s3_constants.ACTION_WRITE},
|
||||
{StatementActionRead, s3_constants.ACTION_READ},
|
||||
{StatementActionList, s3_constants.ACTION_LIST},
|
||||
{StatementActionDelete, s3_constants.ACTION_DELETE_BUCKET},
|
||||
// Test fine-grained S3 action mappings (Issue #7864)
|
||||
{"DeleteObject", s3_constants.ACTION_WRITE},
|
||||
{"s3:DeleteObject", s3_constants.ACTION_WRITE},
|
||||
{"PutObject", s3_constants.ACTION_WRITE},
|
||||
{"s3:PutObject", s3_constants.ACTION_WRITE},
|
||||
{"GetObject", s3_constants.ACTION_READ},
|
||||
{"s3:GetObject", s3_constants.ACTION_READ},
|
||||
{"ListBucket", s3_constants.ACTION_LIST},
|
||||
{"s3:ListBucket", s3_constants.ACTION_LIST},
|
||||
{"PutObjectAcl", s3_constants.ACTION_WRITE_ACP},
|
||||
{"s3:PutObjectAcl", s3_constants.ACTION_WRITE_ACP},
|
||||
{"GetObjectAcl", s3_constants.ACTION_READ_ACP},
|
||||
{"s3:GetObjectAcl", s3_constants.ACTION_READ_ACP},
|
||||
{"unknown", ""},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
result := MapToStatementAction(test.input)
|
||||
assert.Equal(t, test.expected, result, "Failed for input: %s", test.input)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapToIdentitiesAction(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{s3_constants.ACTION_ADMIN, StatementActionAdmin},
|
||||
{s3_constants.ACTION_WRITE, StatementActionWrite},
|
||||
{s3_constants.ACTION_READ, StatementActionRead},
|
||||
{s3_constants.ACTION_LIST, StatementActionList},
|
||||
{s3_constants.ACTION_DELETE_BUCKET, StatementActionDelete},
|
||||
{"unknown", ""},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
result := MapToIdentitiesAction(test.input)
|
||||
assert.Equal(t, test.expected, result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMaskAccessKey(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{"AKIAIOSFODNN7EXAMPLE", "AKIA***"},
|
||||
{"AKIA", "AKIA"},
|
||||
{"AKI", "AKI"},
|
||||
{"", ""},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
result := MaskAccessKey(test.input)
|
||||
assert.Equal(t, test.expected, result)
|
||||
}
|
||||
}
|
||||
@@ -202,32 +202,6 @@ func (m *IAMManager) getFilerAddress() string {
|
||||
return "" // Fallback to empty string if no provider is set
|
||||
}
|
||||
|
||||
// createRoleStore creates a role store based on configuration
|
||||
func (m *IAMManager) createRoleStore(config *RoleStoreConfig) (RoleStore, error) {
|
||||
if config == nil {
|
||||
// Default to generic cached filer role store when no config provided
|
||||
return NewGenericCachedRoleStore(nil, nil)
|
||||
}
|
||||
|
||||
switch config.StoreType {
|
||||
case "", "filer":
|
||||
// Check if caching is explicitly disabled
|
||||
if config.StoreConfig != nil {
|
||||
if noCache, ok := config.StoreConfig["noCache"].(bool); ok && noCache {
|
||||
return NewFilerRoleStore(config.StoreConfig, nil)
|
||||
}
|
||||
}
|
||||
// Default to generic cached filer store for better performance
|
||||
return NewGenericCachedRoleStore(config.StoreConfig, nil)
|
||||
case "cached-filer", "generic-cached":
|
||||
return NewGenericCachedRoleStore(config.StoreConfig, nil)
|
||||
case "memory":
|
||||
return NewMemoryRoleStore(), nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported role store type: %s", config.StoreType)
|
||||
}
|
||||
}
|
||||
|
||||
// createRoleStoreWithProvider creates a role store with a filer address provider function
|
||||
func (m *IAMManager) createRoleStoreWithProvider(config *RoleStoreConfig, filerAddressProvider func() string) (RoleStore, error) {
|
||||
if config == nil {
|
||||
@@ -725,6 +699,23 @@ func (m *IAMManager) ExpireSessionForTesting(ctx context.Context, sessionToken s
|
||||
return m.stsService.ExpireSessionForTesting(ctx, sessionToken)
|
||||
}
|
||||
|
||||
// GetPoliciesForUser returns the policy names attached to an IAM user.
|
||||
// Returns an error if the user store is not configured or the lookup fails,
|
||||
// so callers can fail closed on policy-resolution failures.
|
||||
func (m *IAMManager) GetPoliciesForUser(ctx context.Context, username string) ([]string, error) {
|
||||
if m.userStore == nil {
|
||||
return nil, fmt.Errorf("user store not configured")
|
||||
}
|
||||
user, err := m.userStore.GetUser(ctx, username)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to look up user %q: %w", username, err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, nil
|
||||
}
|
||||
return user.PolicyNames, nil
|
||||
}
|
||||
|
||||
// GetSTSService returns the STS service instance
|
||||
func (m *IAMManager) GetSTSService() *sts.STSService {
|
||||
return m.stsService
|
||||
|
||||
@@ -388,157 +388,3 @@ type CachedFilerRoleStoreConfig struct {
|
||||
ListTTL string `json:"listTtl,omitempty"` // e.g., "1m", "30s"
|
||||
MaxCacheSize int `json:"maxCacheSize,omitempty"` // Maximum number of cached roles
|
||||
}
|
||||
|
||||
// NewCachedFilerRoleStore creates a new cached filer-based role store
|
||||
func NewCachedFilerRoleStore(config map[string]interface{}) (*CachedFilerRoleStore, error) {
|
||||
// Create underlying filer store
|
||||
filerStore, err := NewFilerRoleStore(config, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create filer role store: %w", err)
|
||||
}
|
||||
|
||||
// Parse cache configuration with defaults
|
||||
cacheTTL := 5 * time.Minute // Default 5 minutes for role cache
|
||||
listTTL := 1 * time.Minute // Default 1 minute for list cache
|
||||
maxCacheSize := 1000 // Default max 1000 cached roles
|
||||
|
||||
if config != nil {
|
||||
if ttlStr, ok := config["ttl"].(string); ok && ttlStr != "" {
|
||||
if parsed, err := time.ParseDuration(ttlStr); err == nil {
|
||||
cacheTTL = parsed
|
||||
}
|
||||
}
|
||||
if listTTLStr, ok := config["listTtl"].(string); ok && listTTLStr != "" {
|
||||
if parsed, err := time.ParseDuration(listTTLStr); err == nil {
|
||||
listTTL = parsed
|
||||
}
|
||||
}
|
||||
if maxSize, ok := config["maxCacheSize"].(int); ok && maxSize > 0 {
|
||||
maxCacheSize = maxSize
|
||||
}
|
||||
}
|
||||
|
||||
// Create ccache instances with appropriate configurations
|
||||
pruneCount := int64(maxCacheSize) >> 3
|
||||
if pruneCount <= 0 {
|
||||
pruneCount = 100
|
||||
}
|
||||
|
||||
store := &CachedFilerRoleStore{
|
||||
filerStore: filerStore,
|
||||
cache: ccache.New(ccache.Configure().MaxSize(int64(maxCacheSize)).ItemsToPrune(uint32(pruneCount))),
|
||||
listCache: ccache.New(ccache.Configure().MaxSize(100).ItemsToPrune(10)), // Smaller cache for lists
|
||||
ttl: cacheTTL,
|
||||
listTTL: listTTL,
|
||||
}
|
||||
|
||||
glog.V(2).Infof("Initialized CachedFilerRoleStore with TTL %v, List TTL %v, Max Cache Size %d",
|
||||
cacheTTL, listTTL, maxCacheSize)
|
||||
|
||||
return store, nil
|
||||
}
|
||||
|
||||
// StoreRole stores a role definition and invalidates the cache
|
||||
func (c *CachedFilerRoleStore) StoreRole(ctx context.Context, filerAddress string, roleName string, role *RoleDefinition) error {
|
||||
// Store in filer
|
||||
err := c.filerStore.StoreRole(ctx, filerAddress, roleName, role)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Invalidate cache entries
|
||||
c.cache.Delete(roleName)
|
||||
c.listCache.Clear() // Invalidate list cache
|
||||
|
||||
glog.V(3).Infof("Stored and invalidated cache for role %s", roleName)
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetRole retrieves a role definition with caching
|
||||
func (c *CachedFilerRoleStore) GetRole(ctx context.Context, filerAddress string, roleName string) (*RoleDefinition, error) {
|
||||
// Try to get from cache first
|
||||
item := c.cache.Get(roleName)
|
||||
if item != nil {
|
||||
// Cache hit - return cached role (DO NOT extend TTL)
|
||||
role := item.Value().(*RoleDefinition)
|
||||
glog.V(4).Infof("Cache hit for role %s", roleName)
|
||||
return copyRoleDefinition(role), nil
|
||||
}
|
||||
|
||||
// Cache miss - fetch from filer
|
||||
glog.V(4).Infof("Cache miss for role %s, fetching from filer", roleName)
|
||||
role, err := c.filerStore.GetRole(ctx, filerAddress, roleName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Cache the result with TTL
|
||||
c.cache.Set(roleName, copyRoleDefinition(role), c.ttl)
|
||||
glog.V(3).Infof("Cached role %s with TTL %v", roleName, c.ttl)
|
||||
return role, nil
|
||||
}
|
||||
|
||||
// ListRoles lists all role names with caching
|
||||
func (c *CachedFilerRoleStore) ListRoles(ctx context.Context, filerAddress string) ([]string, error) {
|
||||
// Use a constant key for the role list cache
|
||||
const listCacheKey = "role_list"
|
||||
|
||||
// Try to get from list cache first
|
||||
item := c.listCache.Get(listCacheKey)
|
||||
if item != nil {
|
||||
// Cache hit - return cached list (DO NOT extend TTL)
|
||||
roles := item.Value().([]string)
|
||||
glog.V(4).Infof("List cache hit, returning %d roles", len(roles))
|
||||
return append([]string(nil), roles...), nil // Return a copy
|
||||
}
|
||||
|
||||
// Cache miss - fetch from filer
|
||||
glog.V(4).Infof("List cache miss, fetching from filer")
|
||||
roles, err := c.filerStore.ListRoles(ctx, filerAddress)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Cache the result with TTL (store a copy)
|
||||
rolesCopy := append([]string(nil), roles...)
|
||||
c.listCache.Set(listCacheKey, rolesCopy, c.listTTL)
|
||||
glog.V(3).Infof("Cached role list with %d entries, TTL %v", len(roles), c.listTTL)
|
||||
return roles, nil
|
||||
}
|
||||
|
||||
// DeleteRole deletes a role definition and invalidates the cache
|
||||
func (c *CachedFilerRoleStore) DeleteRole(ctx context.Context, filerAddress string, roleName string) error {
|
||||
// Delete from filer
|
||||
err := c.filerStore.DeleteRole(ctx, filerAddress, roleName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Invalidate cache entries
|
||||
c.cache.Delete(roleName)
|
||||
c.listCache.Clear() // Invalidate list cache
|
||||
|
||||
glog.V(3).Infof("Deleted and invalidated cache for role %s", roleName)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClearCache clears all cached entries (for testing or manual cache invalidation)
|
||||
func (c *CachedFilerRoleStore) ClearCache() {
|
||||
c.cache.Clear()
|
||||
c.listCache.Clear()
|
||||
glog.V(2).Infof("Cleared all role cache entries")
|
||||
}
|
||||
|
||||
// GetCacheStats returns cache statistics
|
||||
func (c *CachedFilerRoleStore) GetCacheStats() map[string]interface{} {
|
||||
return map[string]interface{}{
|
||||
"roleCache": map[string]interface{}{
|
||||
"size": c.cache.ItemCount(),
|
||||
"ttl": c.ttl.String(),
|
||||
},
|
||||
"listCache": map[string]interface{}{
|
||||
"size": c.listCache.ItemCount(),
|
||||
"ttl": c.listTTL.String(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,687 +0,0 @@
|
||||
package policy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestConditionSetOperators(t *testing.T) {
|
||||
engine := setupTestPolicyEngine(t)
|
||||
|
||||
t.Run("ForAnyValue:StringEquals", func(t *testing.T) {
|
||||
trustPolicy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowOIDC",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRoleWithWebIdentity"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAnyValue:StringEquals": {
|
||||
"oidc:roles": []string{"Dev.SeaweedFS.TestBucket.ReadWrite", "Dev.SeaweedFS.Admin"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Match: Admin is in the requested roles
|
||||
evalCtxMatch := &EvaluationContext{
|
||||
Principal: "web-identity-user",
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"oidc:roles": []string{"Dev.SeaweedFS.Admin", "OtherRole"},
|
||||
},
|
||||
}
|
||||
resultMatch, err := engine.EvaluateTrustPolicy(context.Background(), trustPolicy, evalCtxMatch)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultMatch.Effect)
|
||||
|
||||
// No Match
|
||||
evalCtxNoMatch := &EvaluationContext{
|
||||
Principal: "web-identity-user",
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"oidc:roles": []string{"OtherRole1", "OtherRole2"},
|
||||
},
|
||||
}
|
||||
resultNoMatch, err := engine.EvaluateTrustPolicy(context.Background(), trustPolicy, evalCtxNoMatch)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultNoMatch.Effect)
|
||||
|
||||
// No Match: Empty context for ForAnyValue (should deny)
|
||||
evalCtxEmpty := &EvaluationContext{
|
||||
Principal: "web-identity-user",
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"oidc:roles": []string{},
|
||||
},
|
||||
}
|
||||
resultEmpty, err := engine.EvaluateTrustPolicy(context.Background(), trustPolicy, evalCtxEmpty)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultEmpty.Effect, "ForAnyValue should deny when context is empty")
|
||||
})
|
||||
|
||||
t.Run("ForAllValues:StringEquals", func(t *testing.T) {
|
||||
trustPolicyAll := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowOIDCAll",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRoleWithWebIdentity"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAllValues:StringEquals": {
|
||||
"oidc:roles": []string{"RoleA", "RoleB", "RoleC"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Match: All requested roles ARE in the allowed set
|
||||
evalCtxAllMatch := &EvaluationContext{
|
||||
Principal: "web-identity-user",
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"oidc:roles": []string{"RoleA", "RoleB"},
|
||||
},
|
||||
}
|
||||
resultAllMatch, err := engine.EvaluateTrustPolicy(context.Background(), trustPolicyAll, evalCtxAllMatch)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultAllMatch.Effect)
|
||||
|
||||
// Fail: RoleD is NOT in the allowed set
|
||||
evalCtxAllFail := &EvaluationContext{
|
||||
Principal: "web-identity-user",
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"oidc:roles": []string{"RoleA", "RoleD"},
|
||||
},
|
||||
}
|
||||
resultAllFail, err := engine.EvaluateTrustPolicy(context.Background(), trustPolicyAll, evalCtxAllFail)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultAllFail.Effect)
|
||||
|
||||
// Vacuously true: Request has NO roles
|
||||
evalCtxEmpty := &EvaluationContext{
|
||||
Principal: "web-identity-user",
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"oidc:roles": []string{},
|
||||
},
|
||||
}
|
||||
resultEmpty, err := engine.EvaluateTrustPolicy(context.Background(), trustPolicyAll, evalCtxEmpty)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultEmpty.Effect)
|
||||
})
|
||||
|
||||
t.Run("ForAllValues:NumericEqualsVacuouslyTrue", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowNumericAll",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRole"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAllValues:NumericEquals": {
|
||||
"aws:MultiFactorAuthAge": []string{"3600", "7200"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Vacuously true: Request has NO MFA age info
|
||||
evalCtxEmpty := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:MultiFactorAuthAge": []string{},
|
||||
},
|
||||
}
|
||||
resultEmpty, err := engine.EvaluateTrustPolicy(context.Background(), policy, evalCtxEmpty)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultEmpty.Effect, "Should allow when numeric context is empty for ForAllValues")
|
||||
})
|
||||
|
||||
t.Run("ForAllValues:BoolVacuouslyTrue", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowBoolAll",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRole"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAllValues:Bool": {
|
||||
"aws:SecureTransport": "true",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Vacuously true
|
||||
evalCtxEmpty := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SecureTransport": []interface{}{},
|
||||
},
|
||||
}
|
||||
resultEmpty, err := engine.EvaluateTrustPolicy(context.Background(), policy, evalCtxEmpty)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultEmpty.Effect, "Should allow when bool context is empty for ForAllValues")
|
||||
})
|
||||
|
||||
t.Run("ForAllValues:DateVacuouslyTrue", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowDateAll",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRole"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAllValues:DateGreaterThan": {
|
||||
"aws:CurrentTime": "2020-01-01T00:00:00Z",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Vacuously true
|
||||
evalCtxEmpty := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:CurrentTime": []interface{}{},
|
||||
},
|
||||
}
|
||||
resultEmpty, err := engine.EvaluateTrustPolicy(context.Background(), policy, evalCtxEmpty)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultEmpty.Effect, "Should allow when date context is empty for ForAllValues")
|
||||
})
|
||||
|
||||
t.Run("ForAllValues:DateWithLabelsAsStrings", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowDateStrings",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRole"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAllValues:DateGreaterThan": {
|
||||
"aws:CurrentTime": "2020-01-01T00:00:00Z",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
evalCtx := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:CurrentTime": []string{"2021-01-01T00:00:00Z", "2022-01-01T00:00:00Z"},
|
||||
},
|
||||
}
|
||||
result, err := engine.EvaluateTrustPolicy(context.Background(), policy, evalCtx)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, result.Effect, "Should allow when date context is a slice of strings")
|
||||
})
|
||||
|
||||
t.Run("ForAllValues:BoolWithLabelsAsStrings", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowBoolStrings",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRole"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAllValues:Bool": {
|
||||
"aws:SecureTransport": "true",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
evalCtx := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SecureTransport": []string{"true", "true"},
|
||||
},
|
||||
}
|
||||
result, err := engine.EvaluateTrustPolicy(context.Background(), policy, evalCtx)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, result.Effect, "Should allow when bool context is a slice of strings")
|
||||
})
|
||||
|
||||
t.Run("StringEqualsIgnoreCaseWithVariable", func(t *testing.T) {
|
||||
policyDoc := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowVar",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject"},
|
||||
Resource: []string{"arn:aws:s3:::bucket/*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"StringEqualsIgnoreCase": {
|
||||
"s3:prefix": "${aws:username}/",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "var-policy", policyDoc)
|
||||
require.NoError(t, err)
|
||||
|
||||
evalCtx := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/ALICE/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"s3:prefix": "ALICE/",
|
||||
"aws:username": "alice",
|
||||
},
|
||||
}
|
||||
|
||||
result, err := engine.Evaluate(context.Background(), "", evalCtx, []string{"var-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, result.Effect, "Should allow when variable expands and matches case-insensitively")
|
||||
})
|
||||
|
||||
t.Run("StringLike:CaseSensitivity", func(t *testing.T) {
|
||||
policyDoc := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowCaseSensitiveLike",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject"},
|
||||
Resource: []string{"arn:aws:s3:::bucket/*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"StringLike": {
|
||||
"s3:prefix": "Project/*",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "like-policy", policyDoc)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Match: Case sensitive match
|
||||
evalCtxMatch := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/Project/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"s3:prefix": "Project/data",
|
||||
},
|
||||
}
|
||||
resultMatch, err := engine.Evaluate(context.Background(), "", evalCtxMatch, []string{"like-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultMatch.Effect, "Should allow when case matches exactly")
|
||||
|
||||
// Fail: Case insensitive match (should fail for StringLike)
|
||||
evalCtxFail := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/project/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"s3:prefix": "project/data", // lowercase 'p'
|
||||
},
|
||||
}
|
||||
resultFail, err := engine.Evaluate(context.Background(), "", evalCtxFail, []string{"like-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultFail.Effect, "Should deny when case does not match for StringLike")
|
||||
})
|
||||
|
||||
t.Run("NumericNotEquals:Logic", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "DenySpecificAges",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRole"},
|
||||
Resource: []string{"*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAllValues:NumericNotEquals": {
|
||||
"aws:MultiFactorAuthAge": []string{"3600", "7200"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "numeric-not-equals-policy", policy)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Fail: One age matches an excluded value (3600)
|
||||
evalCtxFail := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:MultiFactorAuthAge": []string{"3600", "1800"},
|
||||
},
|
||||
}
|
||||
resultFail, err := engine.Evaluate(context.Background(), "", evalCtxFail, []string{"numeric-not-equals-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultFail.Effect, "Should deny when one age matches an excluded value")
|
||||
|
||||
// Pass: No age matches any excluded value
|
||||
evalCtxPass := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:MultiFactorAuthAge": []string{"1800", "900"},
|
||||
},
|
||||
}
|
||||
resultPass, err := engine.Evaluate(context.Background(), "", evalCtxPass, []string{"numeric-not-equals-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultPass.Effect, "Should allow when no age matches excluded values")
|
||||
})
|
||||
|
||||
t.Run("DateNotEquals:Logic", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "DenySpecificTimes",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRole"},
|
||||
Resource: []string{"*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAllValues:DateNotEquals": {
|
||||
"aws:CurrentTime": []string{"2024-01-01T00:00:00Z", "2024-01-02T00:00:00Z"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "date-not-equals-policy", policy)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Fail: One time matches an excluded value
|
||||
evalCtxFail := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:CurrentTime": []string{"2024-01-01T00:00:00Z", "2024-01-03T00:00:00Z"},
|
||||
},
|
||||
}
|
||||
resultFail, err := engine.Evaluate(context.Background(), "", evalCtxFail, []string{"date-not-equals-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultFail.Effect, "Should deny when one date matches an excluded value")
|
||||
})
|
||||
|
||||
t.Run("IpAddress:SetOperators", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowSpecificIPs",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject"},
|
||||
Resource: []string{"*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAllValues:IpAddress": {
|
||||
"aws:SourceIp": []string{"192.168.1.0/24", "10.0.0.1"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "ip-set-policy", policy)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Match: All source IPs are in allowed ranges
|
||||
evalCtxMatch := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SourceIp": []string{"192.168.1.10", "10.0.0.1"},
|
||||
},
|
||||
}
|
||||
resultMatch, err := engine.Evaluate(context.Background(), "", evalCtxMatch, []string{"ip-set-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultMatch.Effect)
|
||||
|
||||
// Fail: One source IP is NOT in allowed ranges
|
||||
evalCtxFail := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SourceIp": []string{"192.168.1.10", "172.16.0.1"},
|
||||
},
|
||||
}
|
||||
resultFail, err := engine.Evaluate(context.Background(), "", evalCtxFail, []string{"ip-set-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultFail.Effect)
|
||||
|
||||
// ForAnyValue: IPAddress
|
||||
policyAny := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowAnySpecificIPs",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject"},
|
||||
Resource: []string{"*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAnyValue:IpAddress": {
|
||||
"aws:SourceIp": []string{"192.168.1.0/24"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
err = engine.AddPolicy("", "ip-any-policy", policyAny)
|
||||
require.NoError(t, err)
|
||||
|
||||
evalCtxAnyMatch := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SourceIp": []string{"192.168.1.10", "172.16.0.1"},
|
||||
},
|
||||
}
|
||||
resultAnyMatch, err := engine.Evaluate(context.Background(), "", evalCtxAnyMatch, []string{"ip-any-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultAnyMatch.Effect)
|
||||
})
|
||||
|
||||
t.Run("IpAddress:SingleStringValue", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowSingleIP",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject"},
|
||||
Resource: []string{"*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"IpAddress": {
|
||||
"aws:SourceIp": "192.168.1.1",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "ip-single-policy", policy)
|
||||
require.NoError(t, err)
|
||||
|
||||
evalCtxMatch := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SourceIp": "192.168.1.1",
|
||||
},
|
||||
}
|
||||
resultMatch, err := engine.Evaluate(context.Background(), "", evalCtxMatch, []string{"ip-single-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultMatch.Effect)
|
||||
|
||||
evalCtxNoMatch := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SourceIp": "10.0.0.1",
|
||||
},
|
||||
}
|
||||
resultNoMatch, err := engine.Evaluate(context.Background(), "", evalCtxNoMatch, []string{"ip-single-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultNoMatch.Effect)
|
||||
})
|
||||
|
||||
t.Run("Bool:StringSlicePolicyValues", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowWithBoolStrings",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject"},
|
||||
Resource: []string{"*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"Bool": {
|
||||
"aws:SecureTransport": []string{"true", "false"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "bool-string-slice-policy", policy)
|
||||
require.NoError(t, err)
|
||||
|
||||
evalCtx := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SecureTransport": "true",
|
||||
},
|
||||
}
|
||||
result, err := engine.Evaluate(context.Background(), "", evalCtx, []string{"bool-string-slice-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, result.Effect)
|
||||
})
|
||||
|
||||
t.Run("StringEqualsIgnoreCase:StringSlicePolicyValues", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowWithIgnoreCaseStrings",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject"},
|
||||
Resource: []string{"*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"StringEqualsIgnoreCase": {
|
||||
"s3:x-amz-server-side-encryption": []string{"AES256", "aws:kms"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "string-ignorecase-slice-policy", policy)
|
||||
require.NoError(t, err)
|
||||
|
||||
evalCtx := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"s3:x-amz-server-side-encryption": "aes256",
|
||||
},
|
||||
}
|
||||
result, err := engine.Evaluate(context.Background(), "", evalCtx, []string{"string-ignorecase-slice-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, result.Effect)
|
||||
})
|
||||
|
||||
t.Run("IpAddress:CustomContextKey", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowCustomIPKey",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject"},
|
||||
Resource: []string{"*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"IpAddress": {
|
||||
"custom:VpcIp": "10.0.0.0/16",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "ip-custom-key-policy", policy)
|
||||
require.NoError(t, err)
|
||||
|
||||
evalCtxMatch := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"custom:VpcIp": "10.0.5.1",
|
||||
},
|
||||
}
|
||||
resultMatch, err := engine.Evaluate(context.Background(), "", evalCtxMatch, []string{"ip-custom-key-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultMatch.Effect)
|
||||
|
||||
evalCtxNoMatch := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::bucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"custom:VpcIp": "192.168.1.1",
|
||||
},
|
||||
}
|
||||
resultNoMatch, err := engine.Evaluate(context.Background(), "", evalCtxNoMatch, []string{"ip-custom-key-policy"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultNoMatch.Effect)
|
||||
})
|
||||
}
|
||||
@@ -1,101 +0,0 @@
|
||||
package policy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestNegationSetOperators(t *testing.T) {
|
||||
engine := setupTestPolicyEngine(t)
|
||||
|
||||
t.Run("ForAllValues:StringNotEquals", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "DenyAdmin",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRole"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAllValues:StringNotEquals": {
|
||||
"oidc:roles": []string{"Admin"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// All roles are NOT "Admin" -> Should Allow
|
||||
evalCtxAllow := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"oidc:roles": []string{"User", "Developer"},
|
||||
},
|
||||
}
|
||||
resultAllow, err := engine.EvaluateTrustPolicy(context.Background(), policy, evalCtxAllow)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultAllow.Effect, "Should allow when ALL roles satisfy StringNotEquals Admin")
|
||||
|
||||
// One role is "Admin" -> Should Deny
|
||||
evalCtxDeny := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"oidc:roles": []string{"Admin", "User"},
|
||||
},
|
||||
}
|
||||
resultDeny, err := engine.EvaluateTrustPolicy(context.Background(), policy, evalCtxDeny)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultDeny.Effect, "Should deny when one role is Admin and fails StringNotEquals")
|
||||
})
|
||||
|
||||
t.Run("ForAnyValue:StringNotEquals", func(t *testing.T) {
|
||||
policy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "Requirement",
|
||||
Effect: "Allow",
|
||||
Action: []string{"sts:AssumeRole"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"ForAnyValue:StringNotEquals": {
|
||||
"oidc:roles": []string{"Prohibited"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// At least one role is NOT prohibited -> Should Allow
|
||||
evalCtxAllow := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"oidc:roles": []string{"Prohibited", "Allowed"},
|
||||
},
|
||||
}
|
||||
resultAllow, err := engine.EvaluateTrustPolicy(context.Background(), policy, evalCtxAllow)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectAllow, resultAllow.Effect, "Should allow when at least one role is NOT Prohibited")
|
||||
|
||||
// All roles are Prohibited -> Should Deny
|
||||
evalCtxDeny := &EvaluationContext{
|
||||
Principal: "user",
|
||||
Action: "sts:AssumeRole",
|
||||
Resource: "arn:aws:iam::role/test-role",
|
||||
RequestContext: map[string]interface{}{
|
||||
"oidc:roles": []string{"Prohibited", "Prohibited"},
|
||||
},
|
||||
}
|
||||
resultDeny, err := engine.EvaluateTrustPolicy(context.Background(), policy, evalCtxDeny)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, EffectDeny, resultDeny.Effect, "Should deny when ALL roles are Prohibited")
|
||||
})
|
||||
}
|
||||
@@ -1155,11 +1155,6 @@ func ValidatePolicyDocumentWithType(policy *PolicyDocument, policyType string) e
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateStatement validates a single statement (for backward compatibility)
|
||||
func validateStatement(statement *Statement) error {
|
||||
return validateStatementWithType(statement, "resource")
|
||||
}
|
||||
|
||||
// validateStatementWithType validates a single statement based on policy type
|
||||
func validateStatementWithType(statement *Statement, policyType string) error {
|
||||
if statement.Effect != "Allow" && statement.Effect != "Deny" {
|
||||
@@ -1198,29 +1193,6 @@ func validateStatementWithType(statement *Statement, policyType string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// matchResource checks if a resource pattern matches a requested resource
|
||||
// Uses hybrid approach: simple suffix wildcards for compatibility, filepath.Match for complex patterns
|
||||
func matchResource(pattern, resource string) bool {
|
||||
if pattern == resource {
|
||||
return true
|
||||
}
|
||||
|
||||
// Handle simple suffix wildcard (backward compatibility)
|
||||
if strings.HasSuffix(pattern, "*") {
|
||||
prefix := pattern[:len(pattern)-1]
|
||||
return strings.HasPrefix(resource, prefix)
|
||||
}
|
||||
|
||||
// For complex patterns, use filepath.Match for advanced wildcard support (*, ?, [])
|
||||
matched, err := filepath.Match(pattern, resource)
|
||||
if err != nil {
|
||||
// Fallback to exact match if pattern is malformed
|
||||
return pattern == resource
|
||||
}
|
||||
|
||||
return matched
|
||||
}
|
||||
|
||||
// awsIAMMatch performs AWS IAM-compliant pattern matching with case-insensitivity and policy variable support
|
||||
func awsIAMMatch(pattern, value string, evalCtx *EvaluationContext) bool {
|
||||
// Step 1: Substitute policy variables (e.g., ${aws:username}, ${saml:username})
|
||||
@@ -1274,16 +1246,6 @@ func expandPolicyVariables(pattern string, evalCtx *EvaluationContext) string {
|
||||
return result
|
||||
}
|
||||
|
||||
// getContextValue safely gets a value from the evaluation context
|
||||
func getContextValue(evalCtx *EvaluationContext, key, defaultValue string) string {
|
||||
if value, exists := evalCtx.RequestContext[key]; exists {
|
||||
if str, ok := value.(string); ok {
|
||||
return str
|
||||
}
|
||||
}
|
||||
return defaultValue
|
||||
}
|
||||
|
||||
// AwsWildcardMatch performs case-insensitive wildcard matching like AWS IAM
|
||||
func AwsWildcardMatch(pattern, value string) bool {
|
||||
// Create regex pattern key for caching
|
||||
@@ -1322,29 +1284,6 @@ func AwsWildcardMatch(pattern, value string) bool {
|
||||
return regex.MatchString(value)
|
||||
}
|
||||
|
||||
// matchAction checks if an action pattern matches a requested action
|
||||
// Uses hybrid approach: simple suffix wildcards for compatibility, filepath.Match for complex patterns
|
||||
func matchAction(pattern, action string) bool {
|
||||
if pattern == action {
|
||||
return true
|
||||
}
|
||||
|
||||
// Handle simple suffix wildcard (backward compatibility)
|
||||
if strings.HasSuffix(pattern, "*") {
|
||||
prefix := pattern[:len(pattern)-1]
|
||||
return strings.HasPrefix(action, prefix)
|
||||
}
|
||||
|
||||
// For complex patterns, use filepath.Match for advanced wildcard support (*, ?, [])
|
||||
matched, err := filepath.Match(pattern, action)
|
||||
if err != nil {
|
||||
// Fallback to exact match if pattern is malformed
|
||||
return pattern == action
|
||||
}
|
||||
|
||||
return matched
|
||||
}
|
||||
|
||||
// evaluateStringConditionIgnoreCase evaluates string conditions with case insensitivity
|
||||
func (e *PolicyEngine) evaluateStringConditionIgnoreCase(block map[string]interface{}, evalCtx *EvaluationContext, shouldMatch bool, useWildcard bool, forAllValues bool) bool {
|
||||
for key, expectedValues := range block {
|
||||
|
||||
@@ -1,421 +0,0 @@
|
||||
package policy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestPrincipalMatching tests the matchesPrincipal method
|
||||
func TestPrincipalMatching(t *testing.T) {
|
||||
engine := setupTestPolicyEngine(t)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
principal interface{}
|
||||
evalCtx *EvaluationContext
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "plain wildcard principal",
|
||||
principal: "*",
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "structured wildcard federated principal",
|
||||
principal: map[string]interface{}{
|
||||
"Federated": "*",
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "wildcard in array",
|
||||
principal: map[string]interface{}{
|
||||
"Federated": []interface{}{"specific-provider", "*"},
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "specific federated provider match",
|
||||
principal: map[string]interface{}{
|
||||
"Federated": "https://example.com/oidc",
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://example.com/oidc",
|
||||
},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "specific federated provider no match",
|
||||
principal: map[string]interface{}{
|
||||
"Federated": "https://example.com/oidc",
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://other.com/oidc",
|
||||
},
|
||||
},
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "array with specific provider match",
|
||||
principal: map[string]interface{}{
|
||||
"Federated": []string{"https://provider1.com", "https://provider2.com"},
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://provider2.com",
|
||||
},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "AWS principal match",
|
||||
principal: map[string]interface{}{
|
||||
"AWS": "arn:aws:iam::123456789012:user/alice",
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:PrincipalArn": "arn:aws:iam::123456789012:user/alice",
|
||||
},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "Service principal match",
|
||||
principal: map[string]interface{}{
|
||||
"Service": "s3.amazonaws.com",
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:PrincipalServiceName": "s3.amazonaws.com",
|
||||
},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := engine.matchesPrincipal(tt.principal, tt.evalCtx)
|
||||
assert.Equal(t, tt.want, result, "Principal matching failed for: %s", tt.name)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestEvaluatePrincipalValue tests the evaluatePrincipalValue method
|
||||
func TestEvaluatePrincipalValue(t *testing.T) {
|
||||
engine := setupTestPolicyEngine(t)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
principalValue interface{}
|
||||
contextKey string
|
||||
evalCtx *EvaluationContext
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "wildcard string",
|
||||
principalValue: "*",
|
||||
contextKey: "aws:FederatedProvider",
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "specific string match",
|
||||
principalValue: "https://example.com",
|
||||
contextKey: "aws:FederatedProvider",
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://example.com",
|
||||
},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "specific string no match",
|
||||
principalValue: "https://example.com",
|
||||
contextKey: "aws:FederatedProvider",
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://other.com",
|
||||
},
|
||||
},
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "wildcard in array",
|
||||
principalValue: []interface{}{"provider1", "*"},
|
||||
contextKey: "aws:FederatedProvider",
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "array match",
|
||||
principalValue: []string{"provider1", "provider2", "provider3"},
|
||||
contextKey: "aws:FederatedProvider",
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "provider2",
|
||||
},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "array no match",
|
||||
principalValue: []string{"provider1", "provider2"},
|
||||
contextKey: "aws:FederatedProvider",
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "provider3",
|
||||
},
|
||||
},
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "missing context key",
|
||||
principalValue: "specific-value",
|
||||
contextKey: "aws:FederatedProvider",
|
||||
evalCtx: &EvaluationContext{
|
||||
RequestContext: map[string]interface{}{},
|
||||
},
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := engine.evaluatePrincipalValue(tt.principalValue, tt.evalCtx, tt.contextKey)
|
||||
assert.Equal(t, tt.want, result, "Principal value evaluation failed for: %s", tt.name)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestTrustPolicyEvaluation tests the EvaluateTrustPolicy method
|
||||
func TestTrustPolicyEvaluation(t *testing.T) {
|
||||
engine := setupTestPolicyEngine(t)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
trustPolicy *PolicyDocument
|
||||
evalCtx *EvaluationContext
|
||||
wantEffect Effect
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "wildcard federated principal allows any provider",
|
||||
trustPolicy: &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Effect: "Allow",
|
||||
Principal: map[string]interface{}{
|
||||
"Federated": "*",
|
||||
},
|
||||
Action: []string{"sts:AssumeRoleWithWebIdentity"},
|
||||
},
|
||||
},
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://any-provider.com",
|
||||
},
|
||||
},
|
||||
wantEffect: EffectAllow,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "specific federated principal matches",
|
||||
trustPolicy: &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Effect: "Allow",
|
||||
Principal: map[string]interface{}{
|
||||
"Federated": "https://example.com/oidc",
|
||||
},
|
||||
Action: []string{"sts:AssumeRoleWithWebIdentity"},
|
||||
},
|
||||
},
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://example.com/oidc",
|
||||
},
|
||||
},
|
||||
wantEffect: EffectAllow,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "specific federated principal does not match",
|
||||
trustPolicy: &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Effect: "Allow",
|
||||
Principal: map[string]interface{}{
|
||||
"Federated": "https://example.com/oidc",
|
||||
},
|
||||
Action: []string{"sts:AssumeRoleWithWebIdentity"},
|
||||
},
|
||||
},
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://other.com/oidc",
|
||||
},
|
||||
},
|
||||
wantEffect: EffectDeny,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "plain wildcard principal",
|
||||
trustPolicy: &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Effect: "Allow",
|
||||
Principal: "*",
|
||||
Action: []string{"sts:AssumeRoleWithWebIdentity"},
|
||||
},
|
||||
},
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://any-provider.com",
|
||||
},
|
||||
},
|
||||
wantEffect: EffectAllow,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "trust policy with conditions",
|
||||
trustPolicy: &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Effect: "Allow",
|
||||
Principal: map[string]interface{}{
|
||||
"Federated": "*",
|
||||
},
|
||||
Action: []string{"sts:AssumeRoleWithWebIdentity"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"StringEquals": {
|
||||
"oidc:aud": "my-app-id",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://provider.com",
|
||||
"oidc:aud": "my-app-id",
|
||||
},
|
||||
},
|
||||
wantEffect: EffectAllow,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "trust policy condition not met",
|
||||
trustPolicy: &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Effect: "Allow",
|
||||
Principal: map[string]interface{}{
|
||||
"Federated": "*",
|
||||
},
|
||||
Action: []string{"sts:AssumeRoleWithWebIdentity"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"StringEquals": {
|
||||
"oidc:aud": "my-app-id",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
evalCtx: &EvaluationContext{
|
||||
Action: "sts:AssumeRoleWithWebIdentity",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:FederatedProvider": "https://provider.com",
|
||||
"oidc:aud": "wrong-app-id",
|
||||
},
|
||||
},
|
||||
wantEffect: EffectDeny,
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result, err := engine.EvaluateTrustPolicy(context.Background(), tt.trustPolicy, tt.evalCtx)
|
||||
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tt.wantEffect, result.Effect, "Trust policy evaluation failed for: %s", tt.name)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestGetPrincipalContextKey tests the context key mapping
|
||||
func TestGetPrincipalContextKey(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
principalType string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "Federated principal",
|
||||
principalType: "Federated",
|
||||
want: "aws:FederatedProvider",
|
||||
},
|
||||
{
|
||||
name: "AWS principal",
|
||||
principalType: "AWS",
|
||||
want: "aws:PrincipalArn",
|
||||
},
|
||||
{
|
||||
name: "Service principal",
|
||||
principalType: "Service",
|
||||
want: "aws:PrincipalServiceName",
|
||||
},
|
||||
{
|
||||
name: "Custom principal type",
|
||||
principalType: "CustomType",
|
||||
want: "aws:PrincipalCustomType",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := getPrincipalContextKey(tt.principalType)
|
||||
assert.Equal(t, tt.want, result, "Context key mapping failed for: %s", tt.name)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,426 +0,0 @@
|
||||
package policy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestPolicyEngineInitialization tests policy engine initialization
|
||||
func TestPolicyEngineInitialization(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
config *PolicyEngineConfig
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "valid config",
|
||||
config: &PolicyEngineConfig{
|
||||
DefaultEffect: "Deny",
|
||||
StoreType: "memory",
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "invalid default effect",
|
||||
config: &PolicyEngineConfig{
|
||||
DefaultEffect: "Invalid",
|
||||
StoreType: "memory",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "nil config",
|
||||
config: nil,
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
engine := NewPolicyEngine()
|
||||
|
||||
err := engine.Initialize(tt.config)
|
||||
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, engine.IsInitialized())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestPolicyDocumentValidation tests policy document structure validation
|
||||
func TestPolicyDocumentValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
policy *PolicyDocument
|
||||
wantErr bool
|
||||
errorMsg string
|
||||
}{
|
||||
{
|
||||
name: "valid policy document",
|
||||
policy: &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowS3Read",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject", "s3:ListBucket"},
|
||||
Resource: []string{"arn:aws:s3:::mybucket/*"},
|
||||
},
|
||||
},
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "missing version",
|
||||
policy: &PolicyDocument{
|
||||
Statement: []Statement{
|
||||
{
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject"},
|
||||
Resource: []string{"arn:aws:s3:::mybucket/*"},
|
||||
},
|
||||
},
|
||||
},
|
||||
wantErr: true,
|
||||
errorMsg: "version is required",
|
||||
},
|
||||
{
|
||||
name: "empty statements",
|
||||
policy: &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{},
|
||||
},
|
||||
wantErr: true,
|
||||
errorMsg: "at least one statement is required",
|
||||
},
|
||||
{
|
||||
name: "invalid effect",
|
||||
policy: &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Effect: "Maybe",
|
||||
Action: []string{"s3:GetObject"},
|
||||
Resource: []string{"arn:aws:s3:::mybucket/*"},
|
||||
},
|
||||
},
|
||||
},
|
||||
wantErr: true,
|
||||
errorMsg: "invalid effect",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := ValidatePolicyDocument(tt.policy)
|
||||
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
if tt.errorMsg != "" {
|
||||
assert.Contains(t, err.Error(), tt.errorMsg)
|
||||
}
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestPolicyEvaluation tests policy evaluation logic
|
||||
func TestPolicyEvaluation(t *testing.T) {
|
||||
engine := setupTestPolicyEngine(t)
|
||||
|
||||
// Add test policies
|
||||
readPolicy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowS3Read",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:GetObject", "s3:ListBucket"},
|
||||
Resource: []string{
|
||||
"arn:aws:s3:::public-bucket/*", // For object operations
|
||||
"arn:aws:s3:::public-bucket", // For bucket operations
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "read-policy", readPolicy)
|
||||
require.NoError(t, err)
|
||||
|
||||
denyPolicy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "DenyS3Delete",
|
||||
Effect: "Deny",
|
||||
Action: []string{"s3:DeleteObject"},
|
||||
Resource: []string{"arn:aws:s3:::*"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err = engine.AddPolicy("", "deny-policy", denyPolicy)
|
||||
require.NoError(t, err)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
context *EvaluationContext
|
||||
policies []string
|
||||
want Effect
|
||||
}{
|
||||
{
|
||||
name: "allow read access",
|
||||
context: &EvaluationContext{
|
||||
Principal: "user:alice",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::public-bucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SourceIp": "192.168.1.100",
|
||||
},
|
||||
},
|
||||
policies: []string{"read-policy"},
|
||||
want: EffectAllow,
|
||||
},
|
||||
{
|
||||
name: "deny delete access (explicit deny)",
|
||||
context: &EvaluationContext{
|
||||
Principal: "user:alice",
|
||||
Action: "s3:DeleteObject",
|
||||
Resource: "arn:aws:s3:::public-bucket/file.txt",
|
||||
},
|
||||
policies: []string{"read-policy", "deny-policy"},
|
||||
want: EffectDeny,
|
||||
},
|
||||
{
|
||||
name: "deny by default (no matching policy)",
|
||||
context: &EvaluationContext{
|
||||
Principal: "user:alice",
|
||||
Action: "s3:PutObject",
|
||||
Resource: "arn:aws:s3:::public-bucket/file.txt",
|
||||
},
|
||||
policies: []string{"read-policy"},
|
||||
want: EffectDeny,
|
||||
},
|
||||
{
|
||||
name: "allow with wildcard action",
|
||||
context: &EvaluationContext{
|
||||
Principal: "user:admin",
|
||||
Action: "s3:ListBucket",
|
||||
Resource: "arn:aws:s3:::public-bucket",
|
||||
},
|
||||
policies: []string{"read-policy"},
|
||||
want: EffectAllow,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result, err := engine.Evaluate(context.Background(), "", tt.context, tt.policies)
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, tt.want, result.Effect)
|
||||
|
||||
// Verify evaluation details
|
||||
assert.NotNil(t, result.EvaluationDetails)
|
||||
assert.Equal(t, tt.context.Action, result.EvaluationDetails.Action)
|
||||
assert.Equal(t, tt.context.Resource, result.EvaluationDetails.Resource)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestConditionEvaluation tests policy conditions
|
||||
func TestConditionEvaluation(t *testing.T) {
|
||||
engine := setupTestPolicyEngine(t)
|
||||
|
||||
// Policy with IP address condition
|
||||
conditionalPolicy := &PolicyDocument{
|
||||
Version: "2012-10-17",
|
||||
Statement: []Statement{
|
||||
{
|
||||
Sid: "AllowFromOfficeIP",
|
||||
Effect: "Allow",
|
||||
Action: []string{"s3:*"},
|
||||
Resource: []string{"arn:aws:s3:::*"},
|
||||
Condition: map[string]map[string]interface{}{
|
||||
"IpAddress": {
|
||||
"aws:SourceIp": []string{"192.168.1.0/24", "10.0.0.0/8"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
err := engine.AddPolicy("", "ip-conditional", conditionalPolicy)
|
||||
require.NoError(t, err)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
context *EvaluationContext
|
||||
want Effect
|
||||
}{
|
||||
{
|
||||
name: "allow from office IP",
|
||||
context: &EvaluationContext{
|
||||
Principal: "user:alice",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::mybucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SourceIp": "192.168.1.100",
|
||||
},
|
||||
},
|
||||
want: EffectAllow,
|
||||
},
|
||||
{
|
||||
name: "deny from external IP",
|
||||
context: &EvaluationContext{
|
||||
Principal: "user:alice",
|
||||
Action: "s3:GetObject",
|
||||
Resource: "arn:aws:s3:::mybucket/file.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SourceIp": "8.8.8.8",
|
||||
},
|
||||
},
|
||||
want: EffectDeny,
|
||||
},
|
||||
{
|
||||
name: "allow from internal IP",
|
||||
context: &EvaluationContext{
|
||||
Principal: "user:alice",
|
||||
Action: "s3:PutObject",
|
||||
Resource: "arn:aws:s3:::mybucket/newfile.txt",
|
||||
RequestContext: map[string]interface{}{
|
||||
"aws:SourceIp": "10.1.2.3",
|
||||
},
|
||||
},
|
||||
want: EffectAllow,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result, err := engine.Evaluate(context.Background(), "", tt.context, []string{"ip-conditional"})
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, tt.want, result.Effect)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestResourceMatching tests resource ARN matching
|
||||
func TestResourceMatching(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
policyResource string
|
||||
requestResource string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "exact match",
|
||||
policyResource: "arn:aws:s3:::mybucket/file.txt",
|
||||
requestResource: "arn:aws:s3:::mybucket/file.txt",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "wildcard match",
|
||||
policyResource: "arn:aws:s3:::mybucket/*",
|
||||
requestResource: "arn:aws:s3:::mybucket/folder/file.txt",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "bucket wildcard",
|
||||
policyResource: "arn:aws:s3:::*",
|
||||
requestResource: "arn:aws:s3:::anybucket/file.txt",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "no match different bucket",
|
||||
policyResource: "arn:aws:s3:::mybucket/*",
|
||||
requestResource: "arn:aws:s3:::otherbucket/file.txt",
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "prefix match",
|
||||
policyResource: "arn:aws:s3:::mybucket/documents/*",
|
||||
requestResource: "arn:aws:s3:::mybucket/documents/secret.txt",
|
||||
want: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := matchResource(tt.policyResource, tt.requestResource)
|
||||
assert.Equal(t, tt.want, result)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestActionMatching tests action pattern matching
|
||||
func TestActionMatching(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
policyAction string
|
||||
requestAction string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "exact match",
|
||||
policyAction: "s3:GetObject",
|
||||
requestAction: "s3:GetObject",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "wildcard service",
|
||||
policyAction: "s3:*",
|
||||
requestAction: "s3:PutObject",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "wildcard all",
|
||||
policyAction: "*",
|
||||
requestAction: "filer:CreateEntry",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "prefix match",
|
||||
policyAction: "s3:Get*",
|
||||
requestAction: "s3:GetObject",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "no match different service",
|
||||
policyAction: "s3:GetObject",
|
||||
requestAction: "filer:GetEntry",
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := matchAction(tt.policyAction, tt.requestAction)
|
||||
assert.Equal(t, tt.want, result)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Helper function to set up test policy engine
|
||||
func setupTestPolicyEngine(t *testing.T) *PolicyEngine {
|
||||
engine := NewPolicyEngine()
|
||||
config := &PolicyEngineConfig{
|
||||
DefaultEffect: "Deny",
|
||||
StoreType: "memory",
|
||||
}
|
||||
|
||||
err := engine.Initialize(config)
|
||||
require.NoError(t, err)
|
||||
|
||||
return engine
|
||||
}
|
||||
@@ -1,246 +0,0 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestIdentityProviderInterface tests the core identity provider interface
|
||||
func TestIdentityProviderInterface(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
provider IdentityProvider
|
||||
wantErr bool
|
||||
}{
|
||||
// We'll add test cases as we implement providers
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
// Test provider name
|
||||
name := tt.provider.Name()
|
||||
assert.NotEmpty(t, name, "Provider name should not be empty")
|
||||
|
||||
// Test initialization
|
||||
err := tt.provider.Initialize(nil)
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
|
||||
// Test authentication with invalid token
|
||||
ctx := context.Background()
|
||||
_, err = tt.provider.Authenticate(ctx, "invalid-token")
|
||||
assert.Error(t, err, "Should fail with invalid token")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestExternalIdentityValidation tests external identity structure validation
|
||||
func TestExternalIdentityValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
identity *ExternalIdentity
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "valid identity",
|
||||
identity: &ExternalIdentity{
|
||||
UserID: "user123",
|
||||
Email: "user@example.com",
|
||||
DisplayName: "Test User",
|
||||
Groups: []string{"group1", "group2"},
|
||||
Attributes: map[string]string{"dept": "engineering"},
|
||||
Provider: "test-provider",
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "missing user id",
|
||||
identity: &ExternalIdentity{
|
||||
Email: "user@example.com",
|
||||
Provider: "test-provider",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "missing provider",
|
||||
identity: &ExternalIdentity{
|
||||
UserID: "user123",
|
||||
Email: "user@example.com",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid email",
|
||||
identity: &ExternalIdentity{
|
||||
UserID: "user123",
|
||||
Email: "invalid-email",
|
||||
Provider: "test-provider",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := tt.identity.Validate()
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenClaimsValidation tests token claims structure
|
||||
func TestTokenClaimsValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
claims *TokenClaims
|
||||
valid bool
|
||||
}{
|
||||
{
|
||||
name: "valid claims",
|
||||
claims: &TokenClaims{
|
||||
Subject: "user123",
|
||||
Issuer: "https://provider.example.com",
|
||||
Audience: "seaweedfs",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
IssuedAt: time.Now().Add(-time.Minute),
|
||||
Claims: map[string]interface{}{"email": "user@example.com"},
|
||||
},
|
||||
valid: true,
|
||||
},
|
||||
{
|
||||
name: "expired token",
|
||||
claims: &TokenClaims{
|
||||
Subject: "user123",
|
||||
Issuer: "https://provider.example.com",
|
||||
Audience: "seaweedfs",
|
||||
ExpiresAt: time.Now().Add(-time.Hour), // Expired
|
||||
IssuedAt: time.Now().Add(-time.Hour * 2),
|
||||
Claims: map[string]interface{}{"email": "user@example.com"},
|
||||
},
|
||||
valid: false,
|
||||
},
|
||||
{
|
||||
name: "future issued token",
|
||||
claims: &TokenClaims{
|
||||
Subject: "user123",
|
||||
Issuer: "https://provider.example.com",
|
||||
Audience: "seaweedfs",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
IssuedAt: time.Now().Add(time.Hour), // Future
|
||||
Claims: map[string]interface{}{"email": "user@example.com"},
|
||||
},
|
||||
valid: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
valid := tt.claims.IsValid()
|
||||
assert.Equal(t, tt.valid, valid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestProviderRegistry tests provider registration and discovery
|
||||
func TestProviderRegistry(t *testing.T) {
|
||||
// Clear registry for test
|
||||
registry := NewProviderRegistry()
|
||||
|
||||
t.Run("register provider", func(t *testing.T) {
|
||||
mockProvider := &MockProvider{name: "test-provider"}
|
||||
|
||||
err := registry.RegisterProvider(mockProvider)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Test duplicate registration
|
||||
err = registry.RegisterProvider(mockProvider)
|
||||
assert.Error(t, err, "Should not allow duplicate registration")
|
||||
})
|
||||
|
||||
t.Run("get provider", func(t *testing.T) {
|
||||
provider, exists := registry.GetProvider("test-provider")
|
||||
assert.True(t, exists)
|
||||
assert.Equal(t, "test-provider", provider.Name())
|
||||
|
||||
// Test non-existent provider
|
||||
_, exists = registry.GetProvider("non-existent")
|
||||
assert.False(t, exists)
|
||||
})
|
||||
|
||||
t.Run("list providers", func(t *testing.T) {
|
||||
providers := registry.ListProviders()
|
||||
assert.Len(t, providers, 1)
|
||||
assert.Equal(t, "test-provider", providers[0])
|
||||
})
|
||||
}
|
||||
|
||||
// MockProvider for testing
|
||||
type MockProvider struct {
|
||||
name string
|
||||
initialized bool
|
||||
shouldError bool
|
||||
}
|
||||
|
||||
func (m *MockProvider) Name() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
func (m *MockProvider) Initialize(config interface{}) error {
|
||||
if m.shouldError {
|
||||
return assert.AnError
|
||||
}
|
||||
m.initialized = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *MockProvider) Authenticate(ctx context.Context, token string) (*ExternalIdentity, error) {
|
||||
if !m.initialized {
|
||||
return nil, assert.AnError
|
||||
}
|
||||
if token == "invalid-token" {
|
||||
return nil, assert.AnError
|
||||
}
|
||||
return &ExternalIdentity{
|
||||
UserID: "test-user",
|
||||
Email: "test@example.com",
|
||||
DisplayName: "Test User",
|
||||
Provider: m.name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *MockProvider) GetUserInfo(ctx context.Context, userID string) (*ExternalIdentity, error) {
|
||||
if !m.initialized || userID == "" {
|
||||
return nil, assert.AnError
|
||||
}
|
||||
return &ExternalIdentity{
|
||||
UserID: userID,
|
||||
Email: userID + "@example.com",
|
||||
DisplayName: "User " + userID,
|
||||
Provider: m.name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *MockProvider) ValidateToken(ctx context.Context, token string) (*TokenClaims, error) {
|
||||
if !m.initialized || token == "invalid-token" {
|
||||
return nil, assert.AnError
|
||||
}
|
||||
return &TokenClaims{
|
||||
Subject: "test-user",
|
||||
Issuer: "test-issuer",
|
||||
Audience: "seaweedfs",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
IssuedAt: time.Now(),
|
||||
Claims: map[string]interface{}{"email": "test@example.com"},
|
||||
}, nil
|
||||
}
|
||||
@@ -1,109 +0,0 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// ProviderRegistry manages registered identity providers
|
||||
type ProviderRegistry struct {
|
||||
mu sync.RWMutex
|
||||
providers map[string]IdentityProvider
|
||||
}
|
||||
|
||||
// NewProviderRegistry creates a new provider registry
|
||||
func NewProviderRegistry() *ProviderRegistry {
|
||||
return &ProviderRegistry{
|
||||
providers: make(map[string]IdentityProvider),
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterProvider registers a new identity provider
|
||||
func (r *ProviderRegistry) RegisterProvider(provider IdentityProvider) error {
|
||||
if provider == nil {
|
||||
return fmt.Errorf("provider cannot be nil")
|
||||
}
|
||||
|
||||
name := provider.Name()
|
||||
if name == "" {
|
||||
return fmt.Errorf("provider name cannot be empty")
|
||||
}
|
||||
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
if _, exists := r.providers[name]; exists {
|
||||
return fmt.Errorf("provider %s is already registered", name)
|
||||
}
|
||||
|
||||
r.providers[name] = provider
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetProvider retrieves a provider by name
|
||||
func (r *ProviderRegistry) GetProvider(name string) (IdentityProvider, bool) {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
provider, exists := r.providers[name]
|
||||
return provider, exists
|
||||
}
|
||||
|
||||
// ListProviders returns all registered provider names
|
||||
func (r *ProviderRegistry) ListProviders() []string {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
var names []string
|
||||
for name := range r.providers {
|
||||
names = append(names, name)
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// UnregisterProvider removes a provider from the registry
|
||||
func (r *ProviderRegistry) UnregisterProvider(name string) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
if _, exists := r.providers[name]; !exists {
|
||||
return fmt.Errorf("provider %s is not registered", name)
|
||||
}
|
||||
|
||||
delete(r.providers, name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Clear removes all providers from the registry
|
||||
func (r *ProviderRegistry) Clear() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
r.providers = make(map[string]IdentityProvider)
|
||||
}
|
||||
|
||||
// GetProviderCount returns the number of registered providers
|
||||
func (r *ProviderRegistry) GetProviderCount() int {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
return len(r.providers)
|
||||
}
|
||||
|
||||
// Default global registry
|
||||
var defaultRegistry = NewProviderRegistry()
|
||||
|
||||
// RegisterProvider registers a provider in the default registry
|
||||
func RegisterProvider(provider IdentityProvider) error {
|
||||
return defaultRegistry.RegisterProvider(provider)
|
||||
}
|
||||
|
||||
// GetProvider retrieves a provider from the default registry
|
||||
func GetProvider(name string) (IdentityProvider, bool) {
|
||||
return defaultRegistry.GetProvider(name)
|
||||
}
|
||||
|
||||
// ListProviders returns all provider names from the default registry
|
||||
func ListProviders() []string {
|
||||
return defaultRegistry.ListProviders()
|
||||
}
|
||||
@@ -124,6 +124,7 @@ const (
|
||||
ActionAssumeRole = "sts:AssumeRole"
|
||||
ActionAssumeRoleWithWebIdentity = "sts:AssumeRoleWithWebIdentity"
|
||||
ActionAssumeRoleWithCredentials = "sts:AssumeRoleWithCredentials"
|
||||
ActionGetFederationToken = "sts:GetFederationToken"
|
||||
ActionValidateSession = "sts:ValidateSession"
|
||||
)
|
||||
|
||||
|
||||
@@ -1,503 +0,0 @@
|
||||
package sts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/oidc"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/providers"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// Test-only constants for mock providers
|
||||
const (
|
||||
ProviderTypeMock = "mock"
|
||||
)
|
||||
|
||||
// createMockOIDCProvider creates a mock OIDC provider for testing
|
||||
// This is only available in test builds
|
||||
func createMockOIDCProvider(name string, config map[string]interface{}) (providers.IdentityProvider, error) {
|
||||
// Convert config to OIDC format
|
||||
factory := NewProviderFactory()
|
||||
oidcConfig, err := factory.convertToOIDCConfig(config)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Set default values for mock provider if not provided
|
||||
if oidcConfig.Issuer == "" {
|
||||
oidcConfig.Issuer = "http://localhost:9999"
|
||||
}
|
||||
|
||||
provider := oidc.NewMockOIDCProvider(name)
|
||||
if err := provider.Initialize(oidcConfig); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Set up default test data for the mock provider
|
||||
provider.SetupDefaultTestData()
|
||||
|
||||
return provider, nil
|
||||
}
|
||||
|
||||
// createMockJWT creates a test JWT token with the specified issuer for mock provider testing
|
||||
func createMockJWT(t *testing.T, issuer, subject string) string {
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
|
||||
"iss": issuer,
|
||||
"sub": subject,
|
||||
"aud": "test-client",
|
||||
"exp": time.Now().Add(time.Hour).Unix(),
|
||||
"iat": time.Now().Unix(),
|
||||
})
|
||||
|
||||
tokenString, err := token.SignedString([]byte("test-signing-key"))
|
||||
require.NoError(t, err)
|
||||
return tokenString
|
||||
}
|
||||
|
||||
// TestCrossInstanceTokenUsage verifies that tokens generated by one STS instance
|
||||
// can be used and validated by other STS instances in a distributed environment
|
||||
func TestCrossInstanceTokenUsage(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
// Dummy filer address for testing
|
||||
|
||||
// Common configuration that would be shared across all instances in production
|
||||
sharedConfig := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "distributed-sts-cluster", // SAME across all instances
|
||||
SigningKey: []byte(TestSigningKey32Chars), // SAME across all instances
|
||||
Providers: []*ProviderConfig{
|
||||
{
|
||||
Name: "company-oidc",
|
||||
Type: ProviderTypeOIDC,
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
ConfigFieldIssuer: "https://sso.company.com/realms/production",
|
||||
ConfigFieldClientID: "seaweedfs-cluster",
|
||||
ConfigFieldJWKSUri: "https://sso.company.com/realms/production/protocol/openid-connect/certs",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Create multiple STS instances simulating different S3 gateway instances
|
||||
instanceA := NewSTSService() // e.g., s3-gateway-1
|
||||
instanceB := NewSTSService() // e.g., s3-gateway-2
|
||||
instanceC := NewSTSService() // e.g., s3-gateway-3
|
||||
|
||||
// Initialize all instances with IDENTICAL configuration
|
||||
err := instanceA.Initialize(sharedConfig)
|
||||
require.NoError(t, err, "Instance A should initialize")
|
||||
|
||||
err = instanceB.Initialize(sharedConfig)
|
||||
require.NoError(t, err, "Instance B should initialize")
|
||||
|
||||
err = instanceC.Initialize(sharedConfig)
|
||||
require.NoError(t, err, "Instance C should initialize")
|
||||
|
||||
// Set up mock trust policy validator for all instances (required for STS testing)
|
||||
mockValidator := &MockTrustPolicyValidator{}
|
||||
instanceA.SetTrustPolicyValidator(mockValidator)
|
||||
instanceB.SetTrustPolicyValidator(mockValidator)
|
||||
instanceC.SetTrustPolicyValidator(mockValidator)
|
||||
|
||||
// Manually register mock provider for testing (not available in production)
|
||||
mockProviderConfig := map[string]interface{}{
|
||||
ConfigFieldIssuer: "http://test-mock:9999",
|
||||
ConfigFieldClientID: TestClientID,
|
||||
}
|
||||
mockProviderA, err := createMockOIDCProvider("test-mock", mockProviderConfig)
|
||||
require.NoError(t, err)
|
||||
mockProviderB, err := createMockOIDCProvider("test-mock", mockProviderConfig)
|
||||
require.NoError(t, err)
|
||||
mockProviderC, err := createMockOIDCProvider("test-mock", mockProviderConfig)
|
||||
require.NoError(t, err)
|
||||
|
||||
instanceA.RegisterProvider(mockProviderA)
|
||||
instanceB.RegisterProvider(mockProviderB)
|
||||
instanceC.RegisterProvider(mockProviderC)
|
||||
|
||||
// Test 1: Token generated on Instance A can be validated on Instance B & C
|
||||
t.Run("cross_instance_token_validation", func(t *testing.T) {
|
||||
// Generate session token on Instance A
|
||||
sessionId := TestSessionID
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
|
||||
tokenFromA, err := instanceA.GetTokenGenerator().GenerateSessionToken(sessionId, expiresAt)
|
||||
require.NoError(t, err, "Instance A should generate token")
|
||||
|
||||
// Validate token on Instance B
|
||||
claimsFromB, err := instanceB.GetTokenGenerator().ValidateSessionToken(tokenFromA)
|
||||
require.NoError(t, err, "Instance B should validate token from Instance A")
|
||||
assert.Equal(t, sessionId, claimsFromB.SessionId, "Session ID should match")
|
||||
|
||||
// Validate same token on Instance C
|
||||
claimsFromC, err := instanceC.GetTokenGenerator().ValidateSessionToken(tokenFromA)
|
||||
require.NoError(t, err, "Instance C should validate token from Instance A")
|
||||
assert.Equal(t, sessionId, claimsFromC.SessionId, "Session ID should match")
|
||||
|
||||
// All instances should extract identical claims
|
||||
assert.Equal(t, claimsFromB.SessionId, claimsFromC.SessionId)
|
||||
assert.Equal(t, claimsFromB.ExpiresAt.Unix(), claimsFromC.ExpiresAt.Unix())
|
||||
assert.Equal(t, claimsFromB.IssuedAt.Unix(), claimsFromC.IssuedAt.Unix())
|
||||
})
|
||||
|
||||
// Test 2: Complete assume role flow across instances
|
||||
t.Run("cross_instance_assume_role_flow", func(t *testing.T) {
|
||||
// Step 1: User authenticates and assumes role on Instance A
|
||||
// Create a valid JWT token for the mock provider
|
||||
mockToken := createMockJWT(t, "http://test-mock:9999", "test-user")
|
||||
|
||||
assumeRequest := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/CrossInstanceTestRole",
|
||||
WebIdentityToken: mockToken, // JWT token for mock provider
|
||||
RoleSessionName: "cross-instance-test-session",
|
||||
DurationSeconds: int64ToPtr(3600),
|
||||
}
|
||||
|
||||
// Instance A processes assume role request
|
||||
responseFromA, err := instanceA.AssumeRoleWithWebIdentity(ctx, assumeRequest)
|
||||
require.NoError(t, err, "Instance A should process assume role")
|
||||
|
||||
sessionToken := responseFromA.Credentials.SessionToken
|
||||
accessKeyId := responseFromA.Credentials.AccessKeyId
|
||||
secretAccessKey := responseFromA.Credentials.SecretAccessKey
|
||||
|
||||
// Verify response structure
|
||||
assert.NotEmpty(t, sessionToken, "Should have session token")
|
||||
assert.NotEmpty(t, accessKeyId, "Should have access key ID")
|
||||
assert.NotEmpty(t, secretAccessKey, "Should have secret access key")
|
||||
assert.NotNil(t, responseFromA.AssumedRoleUser, "Should have assumed role user")
|
||||
|
||||
// Step 2: Use session token on Instance B (different instance)
|
||||
sessionInfoFromB, err := instanceB.ValidateSessionToken(ctx, sessionToken)
|
||||
require.NoError(t, err, "Instance B should validate session token from Instance A")
|
||||
|
||||
assert.Equal(t, assumeRequest.RoleSessionName, sessionInfoFromB.SessionName)
|
||||
assert.Equal(t, assumeRequest.RoleArn, sessionInfoFromB.RoleArn)
|
||||
|
||||
// Step 3: Use same session token on Instance C (yet another instance)
|
||||
sessionInfoFromC, err := instanceC.ValidateSessionToken(ctx, sessionToken)
|
||||
require.NoError(t, err, "Instance C should validate session token from Instance A")
|
||||
|
||||
// All instances should return identical session information
|
||||
assert.Equal(t, sessionInfoFromB.SessionId, sessionInfoFromC.SessionId)
|
||||
assert.Equal(t, sessionInfoFromB.SessionName, sessionInfoFromC.SessionName)
|
||||
assert.Equal(t, sessionInfoFromB.RoleArn, sessionInfoFromC.RoleArn)
|
||||
assert.Equal(t, sessionInfoFromB.Subject, sessionInfoFromC.Subject)
|
||||
assert.Equal(t, sessionInfoFromB.Provider, sessionInfoFromC.Provider)
|
||||
})
|
||||
|
||||
// Test 3: Session revocation across instances
|
||||
t.Run("cross_instance_session_revocation", func(t *testing.T) {
|
||||
// Create session on Instance A
|
||||
mockToken := createMockJWT(t, "http://test-mock:9999", "test-user")
|
||||
|
||||
assumeRequest := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/RevocationTestRole",
|
||||
WebIdentityToken: mockToken,
|
||||
RoleSessionName: "revocation-test-session",
|
||||
}
|
||||
|
||||
response, err := instanceA.AssumeRoleWithWebIdentity(ctx, assumeRequest)
|
||||
require.NoError(t, err)
|
||||
sessionToken := response.Credentials.SessionToken
|
||||
|
||||
// Verify token works on Instance B
|
||||
_, err = instanceB.ValidateSessionToken(ctx, sessionToken)
|
||||
require.NoError(t, err, "Token should be valid on Instance B initially")
|
||||
|
||||
// Validate session on Instance C to verify cross-instance token compatibility
|
||||
_, err = instanceC.ValidateSessionToken(ctx, sessionToken)
|
||||
require.NoError(t, err, "Instance C should be able to validate session token")
|
||||
|
||||
// In a stateless JWT system, tokens remain valid on all instances since they're self-contained
|
||||
// No revocation is possible without breaking the stateless architecture
|
||||
_, err = instanceA.ValidateSessionToken(ctx, sessionToken)
|
||||
assert.NoError(t, err, "Token should still be valid on Instance A (stateless system)")
|
||||
|
||||
// Verify token is still valid on Instance B
|
||||
_, err = instanceB.ValidateSessionToken(ctx, sessionToken)
|
||||
assert.NoError(t, err, "Token should still be valid on Instance B (stateless system)")
|
||||
})
|
||||
|
||||
// Test 4: Provider consistency across instances
|
||||
t.Run("provider_consistency_affects_token_generation", func(t *testing.T) {
|
||||
// All instances should have same providers and be able to process same OIDC tokens
|
||||
providerNamesA := instanceA.getProviderNames()
|
||||
providerNamesB := instanceB.getProviderNames()
|
||||
providerNamesC := instanceC.getProviderNames()
|
||||
|
||||
assert.ElementsMatch(t, providerNamesA, providerNamesB, "Instance A and B should have same providers")
|
||||
assert.ElementsMatch(t, providerNamesB, providerNamesC, "Instance B and C should have same providers")
|
||||
|
||||
// All instances should be able to process same web identity token
|
||||
testToken := createMockJWT(t, "http://test-mock:9999", "test-user")
|
||||
|
||||
// Try to assume role with same token on different instances
|
||||
assumeRequest := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/ProviderTestRole",
|
||||
WebIdentityToken: testToken,
|
||||
RoleSessionName: "provider-consistency-test",
|
||||
}
|
||||
|
||||
// Should work on any instance
|
||||
responseA, errA := instanceA.AssumeRoleWithWebIdentity(ctx, assumeRequest)
|
||||
responseB, errB := instanceB.AssumeRoleWithWebIdentity(ctx, assumeRequest)
|
||||
responseC, errC := instanceC.AssumeRoleWithWebIdentity(ctx, assumeRequest)
|
||||
|
||||
require.NoError(t, errA, "Instance A should process OIDC token")
|
||||
require.NoError(t, errB, "Instance B should process OIDC token")
|
||||
require.NoError(t, errC, "Instance C should process OIDC token")
|
||||
|
||||
// All should return valid responses (sessions will have different IDs but same structure)
|
||||
assert.NotEmpty(t, responseA.Credentials.SessionToken)
|
||||
assert.NotEmpty(t, responseB.Credentials.SessionToken)
|
||||
assert.NotEmpty(t, responseC.Credentials.SessionToken)
|
||||
})
|
||||
}
|
||||
|
||||
// TestSTSDistributedConfigurationRequirements tests the configuration requirements
|
||||
// for cross-instance token compatibility
|
||||
func TestSTSDistributedConfigurationRequirements(t *testing.T) {
|
||||
_ = "localhost:8888" // Dummy filer address for testing (not used in these tests)
|
||||
|
||||
t.Run("same_signing_key_required", func(t *testing.T) {
|
||||
// Instance A with signing key 1
|
||||
configA := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "test-sts",
|
||||
SigningKey: []byte("signing-key-1-32-characters-long"),
|
||||
}
|
||||
|
||||
// Instance B with different signing key
|
||||
configB := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "test-sts",
|
||||
SigningKey: []byte("signing-key-2-32-characters-long"), // DIFFERENT!
|
||||
}
|
||||
|
||||
instanceA := NewSTSService()
|
||||
instanceB := NewSTSService()
|
||||
|
||||
err := instanceA.Initialize(configA)
|
||||
require.NoError(t, err)
|
||||
|
||||
err = instanceB.Initialize(configB)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Generate token on Instance A
|
||||
sessionId := "test-session"
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
tokenFromA, err := instanceA.GetTokenGenerator().GenerateSessionToken(sessionId, expiresAt)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Instance A should validate its own token
|
||||
_, err = instanceA.GetTokenGenerator().ValidateSessionToken(tokenFromA)
|
||||
assert.NoError(t, err, "Instance A should validate own token")
|
||||
|
||||
// Instance B should REJECT token due to different signing key
|
||||
_, err = instanceB.GetTokenGenerator().ValidateSessionToken(tokenFromA)
|
||||
assert.Error(t, err, "Instance B should reject token with different signing key")
|
||||
assert.Contains(t, err.Error(), "invalid token", "Should be signature validation error")
|
||||
})
|
||||
|
||||
t.Run("same_issuer_required", func(t *testing.T) {
|
||||
sharedSigningKey := []byte("shared-signing-key-32-characters-lo")
|
||||
|
||||
// Instance A with issuer 1
|
||||
configA := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "sts-cluster-1",
|
||||
SigningKey: sharedSigningKey,
|
||||
}
|
||||
|
||||
// Instance B with different issuer
|
||||
configB := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "sts-cluster-2", // DIFFERENT!
|
||||
SigningKey: sharedSigningKey,
|
||||
}
|
||||
|
||||
instanceA := NewSTSService()
|
||||
instanceB := NewSTSService()
|
||||
|
||||
err := instanceA.Initialize(configA)
|
||||
require.NoError(t, err)
|
||||
|
||||
err = instanceB.Initialize(configB)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Generate token on Instance A
|
||||
sessionId := "test-session"
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
tokenFromA, err := instanceA.GetTokenGenerator().GenerateSessionToken(sessionId, expiresAt)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Instance B should REJECT token due to different issuer
|
||||
_, err = instanceB.GetTokenGenerator().ValidateSessionToken(tokenFromA)
|
||||
assert.Error(t, err, "Instance B should reject token with different issuer")
|
||||
assert.Contains(t, err.Error(), "invalid issuer", "Should be issuer validation error")
|
||||
})
|
||||
|
||||
t.Run("identical_configuration_required", func(t *testing.T) {
|
||||
// Identical configuration
|
||||
identicalConfig := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "production-sts-cluster",
|
||||
SigningKey: []byte("production-signing-key-32-chars-l"),
|
||||
}
|
||||
|
||||
// Create multiple instances with identical config
|
||||
instances := make([]*STSService, 5)
|
||||
for i := 0; i < 5; i++ {
|
||||
instances[i] = NewSTSService()
|
||||
err := instances[i].Initialize(identicalConfig)
|
||||
require.NoError(t, err, "Instance %d should initialize", i)
|
||||
}
|
||||
|
||||
// Generate token on Instance 0
|
||||
sessionId := "multi-instance-test"
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
token, err := instances[0].GetTokenGenerator().GenerateSessionToken(sessionId, expiresAt)
|
||||
require.NoError(t, err)
|
||||
|
||||
// All other instances should validate the token
|
||||
for i := 1; i < 5; i++ {
|
||||
claims, err := instances[i].GetTokenGenerator().ValidateSessionToken(token)
|
||||
require.NoError(t, err, "Instance %d should validate token", i)
|
||||
assert.Equal(t, sessionId, claims.SessionId, "Instance %d should extract correct session ID", i)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSTSRealWorldDistributedScenarios tests realistic distributed deployment scenarios
|
||||
func TestSTSRealWorldDistributedScenarios(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("load_balanced_s3_gateway_scenario", func(t *testing.T) {
|
||||
// Simulate real production scenario:
|
||||
// 1. User authenticates with OIDC provider
|
||||
// 2. User calls AssumeRoleWithWebIdentity on S3 Gateway 1
|
||||
// 3. User makes S3 requests that hit S3 Gateway 2 & 3 via load balancer
|
||||
// 4. All instances should handle the session token correctly
|
||||
|
||||
productionConfig := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{2 * time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{24 * time.Hour},
|
||||
Issuer: "seaweedfs-production-sts",
|
||||
SigningKey: []byte("prod-signing-key-32-characters-lon"),
|
||||
|
||||
Providers: []*ProviderConfig{
|
||||
{
|
||||
Name: "corporate-oidc",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://sso.company.com/realms/production",
|
||||
"clientId": "seaweedfs-prod-cluster",
|
||||
"clientSecret": "supersecret-prod-key",
|
||||
"scopes": []string{"openid", "profile", "email", "groups"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Create 3 S3 Gateway instances behind load balancer
|
||||
gateway1 := NewSTSService()
|
||||
gateway2 := NewSTSService()
|
||||
gateway3 := NewSTSService()
|
||||
|
||||
err := gateway1.Initialize(productionConfig)
|
||||
require.NoError(t, err)
|
||||
|
||||
err = gateway2.Initialize(productionConfig)
|
||||
require.NoError(t, err)
|
||||
|
||||
err = gateway3.Initialize(productionConfig)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set up mock trust policy validator for all gateway instances
|
||||
mockValidator := &MockTrustPolicyValidator{}
|
||||
gateway1.SetTrustPolicyValidator(mockValidator)
|
||||
gateway2.SetTrustPolicyValidator(mockValidator)
|
||||
gateway3.SetTrustPolicyValidator(mockValidator)
|
||||
|
||||
// Manually register mock provider for testing (not available in production)
|
||||
mockProviderConfig := map[string]interface{}{
|
||||
ConfigFieldIssuer: "http://test-mock:9999",
|
||||
ConfigFieldClientID: "test-client-id",
|
||||
}
|
||||
mockProvider1, err := createMockOIDCProvider("test-mock", mockProviderConfig)
|
||||
require.NoError(t, err)
|
||||
mockProvider2, err := createMockOIDCProvider("test-mock", mockProviderConfig)
|
||||
require.NoError(t, err)
|
||||
mockProvider3, err := createMockOIDCProvider("test-mock", mockProviderConfig)
|
||||
require.NoError(t, err)
|
||||
|
||||
gateway1.RegisterProvider(mockProvider1)
|
||||
gateway2.RegisterProvider(mockProvider2)
|
||||
gateway3.RegisterProvider(mockProvider3)
|
||||
|
||||
// Step 1: User authenticates and hits Gateway 1 for AssumeRole
|
||||
mockToken := createMockJWT(t, "http://test-mock:9999", "production-user")
|
||||
|
||||
assumeRequest := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/ProductionS3User",
|
||||
WebIdentityToken: mockToken, // JWT token from mock provider
|
||||
RoleSessionName: "user-production-session",
|
||||
DurationSeconds: int64ToPtr(7200), // 2 hours
|
||||
}
|
||||
|
||||
stsResponse, err := gateway1.AssumeRoleWithWebIdentity(ctx, assumeRequest)
|
||||
require.NoError(t, err, "Gateway 1 should handle AssumeRole")
|
||||
|
||||
sessionToken := stsResponse.Credentials.SessionToken
|
||||
accessKey := stsResponse.Credentials.AccessKeyId
|
||||
secretKey := stsResponse.Credentials.SecretAccessKey
|
||||
|
||||
// Step 2: User makes S3 requests that hit different gateways via load balancer
|
||||
// Simulate S3 request validation on Gateway 2
|
||||
sessionInfo2, err := gateway2.ValidateSessionToken(ctx, sessionToken)
|
||||
require.NoError(t, err, "Gateway 2 should validate session from Gateway 1")
|
||||
assert.Equal(t, "user-production-session", sessionInfo2.SessionName)
|
||||
assert.Equal(t, "arn:aws:iam::role/ProductionS3User", sessionInfo2.RoleArn)
|
||||
|
||||
// Simulate S3 request validation on Gateway 3
|
||||
sessionInfo3, err := gateway3.ValidateSessionToken(ctx, sessionToken)
|
||||
require.NoError(t, err, "Gateway 3 should validate session from Gateway 1")
|
||||
assert.Equal(t, sessionInfo2.SessionId, sessionInfo3.SessionId, "Should be same session")
|
||||
|
||||
// Step 3: Verify credentials are consistent
|
||||
assert.Equal(t, accessKey, stsResponse.Credentials.AccessKeyId, "Access key should be consistent")
|
||||
assert.Equal(t, secretKey, stsResponse.Credentials.SecretAccessKey, "Secret key should be consistent")
|
||||
|
||||
// Step 4: Session expiration should be honored across all instances
|
||||
assert.True(t, sessionInfo2.ExpiresAt.After(time.Now()), "Session should not be expired")
|
||||
assert.True(t, sessionInfo3.ExpiresAt.After(time.Now()), "Session should not be expired")
|
||||
|
||||
// Step 5: Token should be identical when parsed
|
||||
claims2, err := gateway2.GetTokenGenerator().ValidateSessionToken(sessionToken)
|
||||
require.NoError(t, err)
|
||||
|
||||
claims3, err := gateway3.GetTokenGenerator().ValidateSessionToken(sessionToken)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, claims2.SessionId, claims3.SessionId, "Session IDs should match")
|
||||
assert.Equal(t, claims2.ExpiresAt.Unix(), claims3.ExpiresAt.Unix(), "Expiration should match")
|
||||
})
|
||||
}
|
||||
|
||||
// Helper function to convert int64 to pointer
|
||||
func int64ToPtr(i int64) *int64 {
|
||||
return &i
|
||||
}
|
||||
@@ -1,340 +0,0 @@
|
||||
package sts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestDistributedSTSService verifies that multiple STS instances with identical configurations
|
||||
// behave consistently across distributed environments
|
||||
func TestDistributedSTSService(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// Common configuration for all instances
|
||||
commonConfig := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "distributed-sts-test",
|
||||
SigningKey: []byte("test-signing-key-32-characters-long"),
|
||||
|
||||
Providers: []*ProviderConfig{
|
||||
{
|
||||
Name: "keycloak-oidc",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "http://keycloak:8080/realms/seaweedfs-test",
|
||||
"clientId": "seaweedfs-s3",
|
||||
"jwksUri": "http://keycloak:8080/realms/seaweedfs-test/protocol/openid-connect/certs",
|
||||
},
|
||||
},
|
||||
|
||||
{
|
||||
Name: "disabled-ldap",
|
||||
Type: "oidc", // Use OIDC as placeholder since LDAP isn't implemented
|
||||
Enabled: false,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "ldap://company.com",
|
||||
"clientId": "ldap-client",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Create multiple STS instances simulating distributed deployment
|
||||
instance1 := NewSTSService()
|
||||
instance2 := NewSTSService()
|
||||
instance3 := NewSTSService()
|
||||
|
||||
// Initialize all instances with identical configuration
|
||||
err := instance1.Initialize(commonConfig)
|
||||
require.NoError(t, err, "Instance 1 should initialize successfully")
|
||||
|
||||
err = instance2.Initialize(commonConfig)
|
||||
require.NoError(t, err, "Instance 2 should initialize successfully")
|
||||
|
||||
err = instance3.Initialize(commonConfig)
|
||||
require.NoError(t, err, "Instance 3 should initialize successfully")
|
||||
|
||||
// Manually register mock providers for testing (not available in production)
|
||||
mockProviderConfig := map[string]interface{}{
|
||||
"issuer": "http://localhost:9999",
|
||||
"clientId": "test-client",
|
||||
}
|
||||
mockProvider1, err := createMockOIDCProvider("test-mock-provider", mockProviderConfig)
|
||||
require.NoError(t, err)
|
||||
mockProvider2, err := createMockOIDCProvider("test-mock-provider", mockProviderConfig)
|
||||
require.NoError(t, err)
|
||||
mockProvider3, err := createMockOIDCProvider("test-mock-provider", mockProviderConfig)
|
||||
require.NoError(t, err)
|
||||
|
||||
instance1.RegisterProvider(mockProvider1)
|
||||
instance2.RegisterProvider(mockProvider2)
|
||||
instance3.RegisterProvider(mockProvider3)
|
||||
|
||||
// Verify all instances have identical provider configurations
|
||||
t.Run("provider_consistency", func(t *testing.T) {
|
||||
// All instances should have same number of providers
|
||||
assert.Len(t, instance1.providers, 2, "Instance 1 should have 2 enabled providers")
|
||||
assert.Len(t, instance2.providers, 2, "Instance 2 should have 2 enabled providers")
|
||||
assert.Len(t, instance3.providers, 2, "Instance 3 should have 2 enabled providers")
|
||||
|
||||
// All instances should have same provider names
|
||||
instance1Names := instance1.getProviderNames()
|
||||
instance2Names := instance2.getProviderNames()
|
||||
instance3Names := instance3.getProviderNames()
|
||||
|
||||
assert.ElementsMatch(t, instance1Names, instance2Names, "Instance 1 and 2 should have same providers")
|
||||
assert.ElementsMatch(t, instance2Names, instance3Names, "Instance 2 and 3 should have same providers")
|
||||
|
||||
// Verify specific providers exist on all instances
|
||||
expectedProviders := []string{"keycloak-oidc", "test-mock-provider"}
|
||||
assert.ElementsMatch(t, instance1Names, expectedProviders, "Instance 1 should have expected providers")
|
||||
assert.ElementsMatch(t, instance2Names, expectedProviders, "Instance 2 should have expected providers")
|
||||
assert.ElementsMatch(t, instance3Names, expectedProviders, "Instance 3 should have expected providers")
|
||||
|
||||
// Verify disabled providers are not loaded
|
||||
assert.NotContains(t, instance1Names, "disabled-ldap", "Disabled providers should not be loaded")
|
||||
assert.NotContains(t, instance2Names, "disabled-ldap", "Disabled providers should not be loaded")
|
||||
assert.NotContains(t, instance3Names, "disabled-ldap", "Disabled providers should not be loaded")
|
||||
})
|
||||
|
||||
// Test token generation consistency across instances
|
||||
t.Run("token_generation_consistency", func(t *testing.T) {
|
||||
sessionId := "test-session-123"
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
|
||||
// Generate tokens from different instances
|
||||
token1, err1 := instance1.GetTokenGenerator().GenerateSessionToken(sessionId, expiresAt)
|
||||
token2, err2 := instance2.GetTokenGenerator().GenerateSessionToken(sessionId, expiresAt)
|
||||
token3, err3 := instance3.GetTokenGenerator().GenerateSessionToken(sessionId, expiresAt)
|
||||
|
||||
require.NoError(t, err1, "Instance 1 token generation should succeed")
|
||||
require.NoError(t, err2, "Instance 2 token generation should succeed")
|
||||
require.NoError(t, err3, "Instance 3 token generation should succeed")
|
||||
|
||||
// All tokens should be different (due to timestamp variations)
|
||||
// But they should all be valid JWTs with same signing key
|
||||
assert.NotEmpty(t, token1)
|
||||
assert.NotEmpty(t, token2)
|
||||
assert.NotEmpty(t, token3)
|
||||
})
|
||||
|
||||
// Test token validation consistency - any instance should validate tokens from any other instance
|
||||
t.Run("cross_instance_token_validation", func(t *testing.T) {
|
||||
sessionId := "cross-validation-session"
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
|
||||
// Generate token on instance 1
|
||||
token, err := instance1.GetTokenGenerator().GenerateSessionToken(sessionId, expiresAt)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Validate on all instances
|
||||
claims1, err1 := instance1.GetTokenGenerator().ValidateSessionToken(token)
|
||||
claims2, err2 := instance2.GetTokenGenerator().ValidateSessionToken(token)
|
||||
claims3, err3 := instance3.GetTokenGenerator().ValidateSessionToken(token)
|
||||
|
||||
require.NoError(t, err1, "Instance 1 should validate token from instance 1")
|
||||
require.NoError(t, err2, "Instance 2 should validate token from instance 1")
|
||||
require.NoError(t, err3, "Instance 3 should validate token from instance 1")
|
||||
|
||||
// All instances should extract same session ID
|
||||
assert.Equal(t, sessionId, claims1.SessionId)
|
||||
assert.Equal(t, sessionId, claims2.SessionId)
|
||||
assert.Equal(t, sessionId, claims3.SessionId)
|
||||
|
||||
assert.Equal(t, claims1.SessionId, claims2.SessionId)
|
||||
assert.Equal(t, claims2.SessionId, claims3.SessionId)
|
||||
})
|
||||
|
||||
// Test provider access consistency
|
||||
t.Run("provider_access_consistency", func(t *testing.T) {
|
||||
// All instances should be able to access the same providers
|
||||
provider1, exists1 := instance1.providers["test-mock-provider"]
|
||||
provider2, exists2 := instance2.providers["test-mock-provider"]
|
||||
provider3, exists3 := instance3.providers["test-mock-provider"]
|
||||
|
||||
assert.True(t, exists1, "Instance 1 should have test-mock-provider")
|
||||
assert.True(t, exists2, "Instance 2 should have test-mock-provider")
|
||||
assert.True(t, exists3, "Instance 3 should have test-mock-provider")
|
||||
|
||||
assert.Equal(t, provider1.Name(), provider2.Name())
|
||||
assert.Equal(t, provider2.Name(), provider3.Name())
|
||||
|
||||
// Test authentication with the mock provider on all instances
|
||||
testToken := "valid_test_token"
|
||||
|
||||
identity1, err1 := provider1.Authenticate(ctx, testToken)
|
||||
identity2, err2 := provider2.Authenticate(ctx, testToken)
|
||||
identity3, err3 := provider3.Authenticate(ctx, testToken)
|
||||
|
||||
require.NoError(t, err1, "Instance 1 provider should authenticate successfully")
|
||||
require.NoError(t, err2, "Instance 2 provider should authenticate successfully")
|
||||
require.NoError(t, err3, "Instance 3 provider should authenticate successfully")
|
||||
|
||||
// All instances should return identical identity information
|
||||
assert.Equal(t, identity1.UserID, identity2.UserID)
|
||||
assert.Equal(t, identity2.UserID, identity3.UserID)
|
||||
assert.Equal(t, identity1.Email, identity2.Email)
|
||||
assert.Equal(t, identity2.Email, identity3.Email)
|
||||
assert.Equal(t, identity1.Provider, identity2.Provider)
|
||||
assert.Equal(t, identity2.Provider, identity3.Provider)
|
||||
})
|
||||
}
|
||||
|
||||
// TestSTSConfigurationValidation tests configuration validation for distributed deployments
|
||||
func TestSTSConfigurationValidation(t *testing.T) {
|
||||
t.Run("consistent_signing_keys_required", func(t *testing.T) {
|
||||
// Different signing keys should result in incompatible token validation
|
||||
config1 := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "test-sts",
|
||||
SigningKey: []byte("signing-key-1-32-characters-long"),
|
||||
}
|
||||
|
||||
config2 := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "test-sts",
|
||||
SigningKey: []byte("signing-key-2-32-characters-long"), // Different key!
|
||||
}
|
||||
|
||||
instance1 := NewSTSService()
|
||||
instance2 := NewSTSService()
|
||||
|
||||
err1 := instance1.Initialize(config1)
|
||||
err2 := instance2.Initialize(config2)
|
||||
|
||||
require.NoError(t, err1)
|
||||
require.NoError(t, err2)
|
||||
|
||||
// Generate token on instance 1
|
||||
sessionId := "test-session"
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
token, err := instance1.GetTokenGenerator().GenerateSessionToken(sessionId, expiresAt)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Instance 1 should validate its own token
|
||||
_, err = instance1.GetTokenGenerator().ValidateSessionToken(token)
|
||||
assert.NoError(t, err, "Instance 1 should validate its own token")
|
||||
|
||||
// Instance 2 should reject token from instance 1 (different signing key)
|
||||
_, err = instance2.GetTokenGenerator().ValidateSessionToken(token)
|
||||
assert.Error(t, err, "Instance 2 should reject token with different signing key")
|
||||
})
|
||||
|
||||
t.Run("consistent_issuer_required", func(t *testing.T) {
|
||||
// Different issuers should result in incompatible tokens
|
||||
commonSigningKey := []byte("shared-signing-key-32-characters-lo")
|
||||
|
||||
config1 := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "sts-instance-1",
|
||||
SigningKey: commonSigningKey,
|
||||
}
|
||||
|
||||
config2 := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{12 * time.Hour},
|
||||
Issuer: "sts-instance-2", // Different issuer!
|
||||
SigningKey: commonSigningKey,
|
||||
}
|
||||
|
||||
instance1 := NewSTSService()
|
||||
instance2 := NewSTSService()
|
||||
|
||||
err1 := instance1.Initialize(config1)
|
||||
err2 := instance2.Initialize(config2)
|
||||
|
||||
require.NoError(t, err1)
|
||||
require.NoError(t, err2)
|
||||
|
||||
// Generate token on instance 1
|
||||
sessionId := "test-session"
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
token, err := instance1.GetTokenGenerator().GenerateSessionToken(sessionId, expiresAt)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Instance 2 should reject token due to issuer mismatch
|
||||
// (Even though signing key is the same, issuer validation will fail)
|
||||
_, err = instance2.GetTokenGenerator().ValidateSessionToken(token)
|
||||
assert.Error(t, err, "Instance 2 should reject token with different issuer")
|
||||
})
|
||||
}
|
||||
|
||||
// TestProviderFactoryDistributed tests the provider factory in distributed scenarios
|
||||
func TestProviderFactoryDistributed(t *testing.T) {
|
||||
factory := NewProviderFactory()
|
||||
|
||||
// Simulate configuration that would be identical across all instances
|
||||
configs := []*ProviderConfig{
|
||||
{
|
||||
Name: "production-keycloak",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://keycloak.company.com/realms/seaweedfs",
|
||||
"clientId": "seaweedfs-prod",
|
||||
"clientSecret": "super-secret-key",
|
||||
"jwksUri": "https://keycloak.company.com/realms/seaweedfs/protocol/openid-connect/certs",
|
||||
"scopes": []string{"openid", "profile", "email", "roles"},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "backup-oidc",
|
||||
Type: "oidc",
|
||||
Enabled: false, // Disabled by default
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://backup-oidc.company.com",
|
||||
"clientId": "seaweedfs-backup",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Create providers multiple times (simulating multiple instances)
|
||||
providers1, err1 := factory.LoadProvidersFromConfig(configs)
|
||||
providers2, err2 := factory.LoadProvidersFromConfig(configs)
|
||||
providers3, err3 := factory.LoadProvidersFromConfig(configs)
|
||||
|
||||
require.NoError(t, err1, "First load should succeed")
|
||||
require.NoError(t, err2, "Second load should succeed")
|
||||
require.NoError(t, err3, "Third load should succeed")
|
||||
|
||||
// All instances should have same provider counts
|
||||
assert.Len(t, providers1, 1, "First instance should have 1 enabled provider")
|
||||
assert.Len(t, providers2, 1, "Second instance should have 1 enabled provider")
|
||||
assert.Len(t, providers3, 1, "Third instance should have 1 enabled provider")
|
||||
|
||||
// All instances should have same provider names
|
||||
names1 := make([]string, 0, len(providers1))
|
||||
names2 := make([]string, 0, len(providers2))
|
||||
names3 := make([]string, 0, len(providers3))
|
||||
|
||||
for name := range providers1 {
|
||||
names1 = append(names1, name)
|
||||
}
|
||||
for name := range providers2 {
|
||||
names2 = append(names2, name)
|
||||
}
|
||||
for name := range providers3 {
|
||||
names3 = append(names3, name)
|
||||
}
|
||||
|
||||
assert.ElementsMatch(t, names1, names2, "Instance 1 and 2 should have same provider names")
|
||||
assert.ElementsMatch(t, names2, names3, "Instance 2 and 3 should have same provider names")
|
||||
|
||||
// Verify specific providers
|
||||
expectedProviders := []string{"production-keycloak"}
|
||||
assert.ElementsMatch(t, names1, expectedProviders, "Should have expected enabled providers")
|
||||
|
||||
// Verify disabled providers are not included
|
||||
assert.NotContains(t, names1, "backup-oidc", "Disabled providers should not be loaded")
|
||||
assert.NotContains(t, names2, "backup-oidc", "Disabled providers should not be loaded")
|
||||
assert.NotContains(t, names3, "backup-oidc", "Disabled providers should not be loaded")
|
||||
}
|
||||
@@ -274,69 +274,3 @@ func (f *ProviderFactory) convertToRoleMapping(value interface{}) (*providers.Ro
|
||||
|
||||
return roleMapping, nil
|
||||
}
|
||||
|
||||
// ValidateProviderConfig validates a provider configuration
|
||||
func (f *ProviderFactory) ValidateProviderConfig(config *ProviderConfig) error {
|
||||
if config == nil {
|
||||
return fmt.Errorf("provider config cannot be nil")
|
||||
}
|
||||
|
||||
if config.Name == "" {
|
||||
return fmt.Errorf("provider name cannot be empty")
|
||||
}
|
||||
|
||||
if config.Type == "" {
|
||||
return fmt.Errorf("provider type cannot be empty")
|
||||
}
|
||||
|
||||
if config.Config == nil {
|
||||
return fmt.Errorf("provider config cannot be nil")
|
||||
}
|
||||
|
||||
// Type-specific validation
|
||||
switch config.Type {
|
||||
case "oidc":
|
||||
return f.validateOIDCConfig(config.Config)
|
||||
case "ldap":
|
||||
return f.validateLDAPConfig(config.Config)
|
||||
case "saml":
|
||||
return f.validateSAMLConfig(config.Config)
|
||||
default:
|
||||
return fmt.Errorf("unsupported provider type: %s", config.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// validateOIDCConfig validates OIDC provider configuration
|
||||
func (f *ProviderFactory) validateOIDCConfig(config map[string]interface{}) error {
|
||||
if _, ok := config[ConfigFieldIssuer]; !ok {
|
||||
return fmt.Errorf("OIDC provider requires '%s' field", ConfigFieldIssuer)
|
||||
}
|
||||
|
||||
if _, ok := config[ConfigFieldClientID]; !ok {
|
||||
return fmt.Errorf("OIDC provider requires '%s' field", ConfigFieldClientID)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateLDAPConfig validates LDAP provider configuration
|
||||
func (f *ProviderFactory) validateLDAPConfig(config map[string]interface{}) error {
|
||||
if _, ok := config["server"]; !ok {
|
||||
return fmt.Errorf("LDAP provider requires 'server' field")
|
||||
}
|
||||
if _, ok := config["baseDN"]; !ok {
|
||||
return fmt.Errorf("LDAP provider requires 'baseDN' field")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateSAMLConfig validates SAML provider configuration
|
||||
func (f *ProviderFactory) validateSAMLConfig(config map[string]interface{}) error {
|
||||
// TODO: Implement when SAML provider is available
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetSupportedProviderTypes returns list of supported provider types
|
||||
func (f *ProviderFactory) GetSupportedProviderTypes() []string {
|
||||
return []string{ProviderTypeOIDC}
|
||||
}
|
||||
|
||||
@@ -1,312 +0,0 @@
|
||||
package sts
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestProviderFactory_CreateOIDCProvider(t *testing.T) {
|
||||
factory := NewProviderFactory()
|
||||
|
||||
config := &ProviderConfig{
|
||||
Name: "test-oidc",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://test-issuer.com",
|
||||
"clientId": "test-client",
|
||||
"clientSecret": "test-secret",
|
||||
"jwksUri": "https://test-issuer.com/.well-known/jwks.json",
|
||||
"scopes": []string{"openid", "profile", "email"},
|
||||
},
|
||||
}
|
||||
|
||||
provider, err := factory.CreateProvider(config)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, provider)
|
||||
assert.Equal(t, "test-oidc", provider.Name())
|
||||
}
|
||||
|
||||
// Note: Mock provider tests removed - mock providers are now test-only
|
||||
// and not available through the production ProviderFactory
|
||||
|
||||
func TestProviderFactory_DisabledProvider(t *testing.T) {
|
||||
factory := NewProviderFactory()
|
||||
|
||||
config := &ProviderConfig{
|
||||
Name: "disabled-provider",
|
||||
Type: "oidc",
|
||||
Enabled: false,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://test-issuer.com",
|
||||
"clientId": "test-client",
|
||||
},
|
||||
}
|
||||
|
||||
provider, err := factory.CreateProvider(config)
|
||||
require.NoError(t, err)
|
||||
assert.Nil(t, provider) // Should return nil for disabled providers
|
||||
}
|
||||
|
||||
func TestProviderFactory_InvalidProviderType(t *testing.T) {
|
||||
factory := NewProviderFactory()
|
||||
|
||||
config := &ProviderConfig{
|
||||
Name: "invalid-provider",
|
||||
Type: "unsupported-type",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{},
|
||||
}
|
||||
|
||||
provider, err := factory.CreateProvider(config)
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, provider)
|
||||
assert.Contains(t, err.Error(), "unsupported provider type")
|
||||
}
|
||||
|
||||
func TestProviderFactory_LoadMultipleProviders(t *testing.T) {
|
||||
factory := NewProviderFactory()
|
||||
|
||||
configs := []*ProviderConfig{
|
||||
{
|
||||
Name: "oidc-provider",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://oidc-issuer.com",
|
||||
"clientId": "oidc-client",
|
||||
},
|
||||
},
|
||||
|
||||
{
|
||||
Name: "disabled-provider",
|
||||
Type: "oidc",
|
||||
Enabled: false,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://disabled-issuer.com",
|
||||
"clientId": "disabled-client",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
providers, err := factory.LoadProvidersFromConfig(configs)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, providers, 1) // Only enabled providers should be loaded
|
||||
|
||||
assert.Contains(t, providers, "oidc-provider")
|
||||
assert.NotContains(t, providers, "disabled-provider")
|
||||
}
|
||||
|
||||
func TestProviderFactory_ValidateOIDCConfig(t *testing.T) {
|
||||
factory := NewProviderFactory()
|
||||
|
||||
t.Run("valid config", func(t *testing.T) {
|
||||
config := &ProviderConfig{
|
||||
Name: "valid-oidc",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://valid-issuer.com",
|
||||
"clientId": "valid-client",
|
||||
},
|
||||
}
|
||||
|
||||
err := factory.ValidateProviderConfig(config)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("missing issuer", func(t *testing.T) {
|
||||
config := &ProviderConfig{
|
||||
Name: "invalid-oidc",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"clientId": "valid-client",
|
||||
},
|
||||
}
|
||||
|
||||
err := factory.ValidateProviderConfig(config)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "issuer")
|
||||
})
|
||||
|
||||
t.Run("missing clientId", func(t *testing.T) {
|
||||
config := &ProviderConfig{
|
||||
Name: "invalid-oidc",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://valid-issuer.com",
|
||||
},
|
||||
}
|
||||
|
||||
err := factory.ValidateProviderConfig(config)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "clientId")
|
||||
})
|
||||
}
|
||||
|
||||
func TestProviderFactory_ConvertToStringSlice(t *testing.T) {
|
||||
factory := NewProviderFactory()
|
||||
|
||||
t.Run("string slice", func(t *testing.T) {
|
||||
input := []string{"a", "b", "c"}
|
||||
result, err := factory.convertToStringSlice(input)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"a", "b", "c"}, result)
|
||||
})
|
||||
|
||||
t.Run("interface slice", func(t *testing.T) {
|
||||
input := []interface{}{"a", "b", "c"}
|
||||
result, err := factory.convertToStringSlice(input)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"a", "b", "c"}, result)
|
||||
})
|
||||
|
||||
t.Run("invalid type", func(t *testing.T) {
|
||||
input := "not-a-slice"
|
||||
result, err := factory.convertToStringSlice(input)
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, result)
|
||||
})
|
||||
}
|
||||
|
||||
func TestProviderFactory_ConfigConversionErrors(t *testing.T) {
|
||||
factory := NewProviderFactory()
|
||||
|
||||
t.Run("invalid scopes type", func(t *testing.T) {
|
||||
config := &ProviderConfig{
|
||||
Name: "invalid-scopes",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://test-issuer.com",
|
||||
"clientId": "test-client",
|
||||
"scopes": "invalid-not-array", // Should be array
|
||||
},
|
||||
}
|
||||
|
||||
provider, err := factory.CreateProvider(config)
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, provider)
|
||||
assert.Contains(t, err.Error(), "failed to convert scopes")
|
||||
})
|
||||
|
||||
t.Run("invalid claimsMapping type", func(t *testing.T) {
|
||||
config := &ProviderConfig{
|
||||
Name: "invalid-claims",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://test-issuer.com",
|
||||
"clientId": "test-client",
|
||||
"claimsMapping": "invalid-not-map", // Should be map
|
||||
},
|
||||
}
|
||||
|
||||
provider, err := factory.CreateProvider(config)
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, provider)
|
||||
assert.Contains(t, err.Error(), "failed to convert claimsMapping")
|
||||
})
|
||||
|
||||
t.Run("invalid roleMapping type", func(t *testing.T) {
|
||||
config := &ProviderConfig{
|
||||
Name: "invalid-roles",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://test-issuer.com",
|
||||
"clientId": "test-client",
|
||||
"roleMapping": "invalid-not-map", // Should be map
|
||||
},
|
||||
}
|
||||
|
||||
provider, err := factory.CreateProvider(config)
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, provider)
|
||||
assert.Contains(t, err.Error(), "failed to convert roleMapping")
|
||||
})
|
||||
}
|
||||
|
||||
func TestProviderFactory_ConvertToStringMap(t *testing.T) {
|
||||
factory := NewProviderFactory()
|
||||
|
||||
t.Run("string map", func(t *testing.T) {
|
||||
input := map[string]string{"key1": "value1", "key2": "value2"}
|
||||
result, err := factory.convertToStringMap(input)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, map[string]string{"key1": "value1", "key2": "value2"}, result)
|
||||
})
|
||||
|
||||
t.Run("interface map", func(t *testing.T) {
|
||||
input := map[string]interface{}{"key1": "value1", "key2": "value2"}
|
||||
result, err := factory.convertToStringMap(input)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, map[string]string{"key1": "value1", "key2": "value2"}, result)
|
||||
})
|
||||
|
||||
t.Run("invalid type", func(t *testing.T) {
|
||||
input := "not-a-map"
|
||||
result, err := factory.convertToStringMap(input)
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, result)
|
||||
})
|
||||
}
|
||||
|
||||
func TestProviderFactory_GetSupportedProviderTypes(t *testing.T) {
|
||||
factory := NewProviderFactory()
|
||||
|
||||
supportedTypes := factory.GetSupportedProviderTypes()
|
||||
assert.Contains(t, supportedTypes, "oidc")
|
||||
assert.Len(t, supportedTypes, 1) // Currently only OIDC is supported in production
|
||||
}
|
||||
|
||||
func TestSTSService_LoadProvidersFromConfig(t *testing.T) {
|
||||
stsConfig := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{3600 * time.Second},
|
||||
MaxSessionLength: FlexibleDuration{43200 * time.Second},
|
||||
Issuer: "test-issuer",
|
||||
SigningKey: []byte("test-signing-key-32-characters-long"),
|
||||
Providers: []*ProviderConfig{
|
||||
{
|
||||
Name: "test-provider",
|
||||
Type: "oidc",
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://test-issuer.com",
|
||||
"clientId": "test-client",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
stsService := NewSTSService()
|
||||
err := stsService.Initialize(stsConfig)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Check that provider was loaded
|
||||
assert.Len(t, stsService.providers, 1)
|
||||
assert.Contains(t, stsService.providers, "test-provider")
|
||||
assert.Equal(t, "test-provider", stsService.providers["test-provider"].Name())
|
||||
}
|
||||
|
||||
func TestSTSService_NoProvidersConfig(t *testing.T) {
|
||||
stsConfig := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{3600 * time.Second},
|
||||
MaxSessionLength: FlexibleDuration{43200 * time.Second},
|
||||
Issuer: "test-issuer",
|
||||
SigningKey: []byte("test-signing-key-32-characters-long"),
|
||||
// No providers configured
|
||||
}
|
||||
|
||||
stsService := NewSTSService()
|
||||
err := stsService.Initialize(stsConfig)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Should initialize successfully with no providers
|
||||
assert.Len(t, stsService.providers, 0)
|
||||
}
|
||||
@@ -1,193 +0,0 @@
|
||||
package sts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/providers"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestSecurityIssuerToProviderMapping tests the security fix that ensures JWT tokens
|
||||
// with specific issuer claims can only be validated by the provider registered for that issuer
|
||||
func TestSecurityIssuerToProviderMapping(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// Create STS service with two mock providers
|
||||
service := NewSTSService()
|
||||
config := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{time.Hour * 12},
|
||||
Issuer: "test-sts",
|
||||
SigningKey: []byte("test-signing-key-32-characters-long"),
|
||||
}
|
||||
|
||||
err := service.Initialize(config)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set up mock trust policy validator
|
||||
mockValidator := &MockTrustPolicyValidator{}
|
||||
service.SetTrustPolicyValidator(mockValidator)
|
||||
|
||||
// Create two mock providers with different issuers
|
||||
providerA := &MockIdentityProviderWithIssuer{
|
||||
name: "provider-a",
|
||||
issuer: "https://provider-a.com",
|
||||
validTokens: map[string]bool{
|
||||
"token-for-provider-a": true,
|
||||
},
|
||||
}
|
||||
|
||||
providerB := &MockIdentityProviderWithIssuer{
|
||||
name: "provider-b",
|
||||
issuer: "https://provider-b.com",
|
||||
validTokens: map[string]bool{
|
||||
"token-for-provider-b": true,
|
||||
},
|
||||
}
|
||||
|
||||
// Register both providers
|
||||
err = service.RegisterProvider(providerA)
|
||||
require.NoError(t, err)
|
||||
err = service.RegisterProvider(providerB)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create JWT tokens with specific issuer claims
|
||||
tokenForProviderA := createTestJWT(t, "https://provider-a.com", "user-a")
|
||||
tokenForProviderB := createTestJWT(t, "https://provider-b.com", "user-b")
|
||||
|
||||
t.Run("jwt_token_with_issuer_a_only_validated_by_provider_a", func(t *testing.T) {
|
||||
// This should succeed - token has issuer A and provider A is registered
|
||||
identity, provider, err := service.validateWebIdentityToken(ctx, tokenForProviderA)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, identity)
|
||||
assert.Equal(t, "provider-a", provider.Name())
|
||||
})
|
||||
|
||||
t.Run("jwt_token_with_issuer_b_only_validated_by_provider_b", func(t *testing.T) {
|
||||
// This should succeed - token has issuer B and provider B is registered
|
||||
identity, provider, err := service.validateWebIdentityToken(ctx, tokenForProviderB)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, identity)
|
||||
assert.Equal(t, "provider-b", provider.Name())
|
||||
})
|
||||
|
||||
t.Run("jwt_token_with_unregistered_issuer_fails", func(t *testing.T) {
|
||||
// Create token with unregistered issuer
|
||||
tokenWithUnknownIssuer := createTestJWT(t, "https://unknown-issuer.com", "user-x")
|
||||
|
||||
// This should fail - no provider registered for this issuer
|
||||
identity, provider, err := service.validateWebIdentityToken(ctx, tokenWithUnknownIssuer)
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, identity)
|
||||
assert.Nil(t, provider)
|
||||
assert.Contains(t, err.Error(), "no identity provider registered for issuer: https://unknown-issuer.com")
|
||||
})
|
||||
|
||||
t.Run("non_jwt_tokens_are_rejected", func(t *testing.T) {
|
||||
// Non-JWT tokens should be rejected - no fallback mechanism exists for security
|
||||
identity, provider, err := service.validateWebIdentityToken(ctx, "token-for-provider-a")
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, identity)
|
||||
assert.Nil(t, provider)
|
||||
assert.Contains(t, err.Error(), "web identity token must be a valid JWT token")
|
||||
})
|
||||
}
|
||||
|
||||
// createTestJWT creates a test JWT token with the specified issuer and subject
|
||||
func createTestJWT(t *testing.T, issuer, subject string) string {
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
|
||||
"iss": issuer,
|
||||
"sub": subject,
|
||||
"aud": "test-client",
|
||||
"exp": time.Now().Add(time.Hour).Unix(),
|
||||
"iat": time.Now().Unix(),
|
||||
})
|
||||
|
||||
tokenString, err := token.SignedString([]byte("test-signing-key"))
|
||||
require.NoError(t, err)
|
||||
return tokenString
|
||||
}
|
||||
|
||||
// MockIdentityProviderWithIssuer is a mock provider that supports issuer mapping
|
||||
type MockIdentityProviderWithIssuer struct {
|
||||
name string
|
||||
issuer string
|
||||
validTokens map[string]bool
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithIssuer) Name() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithIssuer) GetIssuer() string {
|
||||
return m.issuer
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithIssuer) Initialize(config interface{}) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithIssuer) Authenticate(ctx context.Context, token string) (*providers.ExternalIdentity, error) {
|
||||
// For JWT tokens, parse and validate the token format
|
||||
if len(token) > 50 && strings.Contains(token, ".") {
|
||||
// This looks like a JWT - parse it to get the subject
|
||||
parsedToken, _, err := new(jwt.Parser).ParseUnverified(token, jwt.MapClaims{})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid JWT token")
|
||||
}
|
||||
|
||||
claims, ok := parsedToken.Claims.(jwt.MapClaims)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("invalid claims")
|
||||
}
|
||||
|
||||
issuer, _ := claims["iss"].(string)
|
||||
subject, _ := claims["sub"].(string)
|
||||
|
||||
// Verify the issuer matches what we expect
|
||||
if issuer != m.issuer {
|
||||
return nil, fmt.Errorf("token issuer %s does not match provider issuer %s", issuer, m.issuer)
|
||||
}
|
||||
|
||||
return &providers.ExternalIdentity{
|
||||
UserID: subject,
|
||||
Email: subject + "@" + m.name + ".com",
|
||||
Provider: m.name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// For non-JWT tokens, check our simple token list
|
||||
if m.validTokens[token] {
|
||||
return &providers.ExternalIdentity{
|
||||
UserID: "test-user",
|
||||
Email: "test@" + m.name + ".com",
|
||||
Provider: m.name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("invalid token")
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithIssuer) GetUserInfo(ctx context.Context, userID string) (*providers.ExternalIdentity, error) {
|
||||
return &providers.ExternalIdentity{
|
||||
UserID: userID,
|
||||
Email: userID + "@" + m.name + ".com",
|
||||
Provider: m.name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithIssuer) ValidateToken(ctx context.Context, token string) (*providers.TokenClaims, error) {
|
||||
if m.validTokens[token] {
|
||||
return &providers.TokenClaims{
|
||||
Subject: "test-user",
|
||||
Issuer: m.issuer,
|
||||
}, nil
|
||||
}
|
||||
return nil, fmt.Errorf("invalid token")
|
||||
}
|
||||
@@ -1,168 +0,0 @@
|
||||
package sts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// createSessionPolicyTestJWT creates a test JWT token for session policy tests
|
||||
func createSessionPolicyTestJWT(t *testing.T, issuer, subject string) string {
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
|
||||
"iss": issuer,
|
||||
"sub": subject,
|
||||
"aud": "test-client",
|
||||
"exp": time.Now().Add(time.Hour).Unix(),
|
||||
"iat": time.Now().Unix(),
|
||||
})
|
||||
|
||||
tokenString, err := token.SignedString([]byte("test-signing-key"))
|
||||
require.NoError(t, err)
|
||||
return tokenString
|
||||
}
|
||||
|
||||
// TestAssumeRoleWithWebIdentity_SessionPolicy verifies inline session policies are preserved in tokens.
|
||||
func TestAssumeRoleWithWebIdentity_SessionPolicy(t *testing.T) {
|
||||
service := setupTestSTSService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
sessionPolicy := `{"Version":"2012-10-17","Statement":[{"Effect":"Allow","Action":"s3:GetObject","Resource":"arn:aws:s3:::example-bucket/*"}]}`
|
||||
testToken := createSessionPolicyTestJWT(t, "test-issuer", "test-user")
|
||||
|
||||
request := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/TestRole",
|
||||
WebIdentityToken: testToken,
|
||||
RoleSessionName: "test-session",
|
||||
Policy: &sessionPolicy,
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithWebIdentity(ctx, request)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response)
|
||||
|
||||
sessionInfo, err := service.ValidateSessionToken(ctx, response.Credentials.SessionToken)
|
||||
require.NoError(t, err)
|
||||
|
||||
normalized, err := NormalizeSessionPolicy(sessionPolicy)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, normalized, sessionInfo.SessionPolicy)
|
||||
|
||||
t.Run("should_succeed_without_session_policy", func(t *testing.T) {
|
||||
request := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/TestRole",
|
||||
WebIdentityToken: createSessionPolicyTestJWT(t, "test-issuer", "test-user"),
|
||||
RoleSessionName: "test-session",
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithWebIdentity(ctx, request)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response)
|
||||
|
||||
sessionInfo, err := service.ValidateSessionToken(ctx, response.Credentials.SessionToken)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, sessionInfo.SessionPolicy)
|
||||
})
|
||||
}
|
||||
|
||||
// Test edge case scenarios for the Policy field handling
|
||||
func TestAssumeRoleWithWebIdentity_SessionPolicy_EdgeCases(t *testing.T) {
|
||||
service := setupTestSTSService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("malformed_json_policy_rejected", func(t *testing.T) {
|
||||
malformedPolicy := `{"Version": "2012-10-17", "Statement": [` // Incomplete JSON
|
||||
|
||||
request := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/TestRole",
|
||||
WebIdentityToken: createSessionPolicyTestJWT(t, "test-issuer", "test-user"),
|
||||
RoleSessionName: "test-session",
|
||||
Policy: &malformedPolicy,
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithWebIdentity(ctx, request)
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, response)
|
||||
assert.Contains(t, err.Error(), "invalid session policy JSON")
|
||||
})
|
||||
|
||||
t.Run("invalid_policy_document_rejected", func(t *testing.T) {
|
||||
invalidPolicy := `{"Version":"2012-10-17","Statement":[{"Effect":"Allow"}]}`
|
||||
|
||||
request := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/TestRole",
|
||||
WebIdentityToken: createSessionPolicyTestJWT(t, "test-issuer", "test-user"),
|
||||
RoleSessionName: "test-session",
|
||||
Policy: &invalidPolicy,
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithWebIdentity(ctx, request)
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, response)
|
||||
assert.Contains(t, err.Error(), "invalid session policy document")
|
||||
})
|
||||
|
||||
t.Run("whitespace_policy_ignored", func(t *testing.T) {
|
||||
whitespacePolicy := " \t\n "
|
||||
|
||||
request := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/TestRole",
|
||||
WebIdentityToken: createSessionPolicyTestJWT(t, "test-issuer", "test-user"),
|
||||
RoleSessionName: "test-session",
|
||||
Policy: &whitespacePolicy,
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithWebIdentity(ctx, request)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response)
|
||||
|
||||
sessionInfo, err := service.ValidateSessionToken(ctx, response.Credentials.SessionToken)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, sessionInfo.SessionPolicy)
|
||||
})
|
||||
}
|
||||
|
||||
// TestAssumeRoleWithWebIdentity_PolicyFieldDocumentation verifies that the struct field exists and is optional.
|
||||
func TestAssumeRoleWithWebIdentity_PolicyFieldDocumentation(t *testing.T) {
|
||||
request := &AssumeRoleWithWebIdentityRequest{}
|
||||
|
||||
assert.IsType(t, (*string)(nil), request.Policy,
|
||||
"Policy field should be *string type for optional JSON policy")
|
||||
assert.Nil(t, request.Policy,
|
||||
"Policy field should default to nil (no session policy)")
|
||||
|
||||
policyValue := `{"Version": "2012-10-17"}`
|
||||
request.Policy = &policyValue
|
||||
assert.NotNil(t, request.Policy, "Should be able to assign policy value")
|
||||
assert.Equal(t, policyValue, *request.Policy, "Policy value should be preserved")
|
||||
}
|
||||
|
||||
// TestAssumeRoleWithCredentials_SessionPolicy verifies session policy support for credentials-based flow.
|
||||
func TestAssumeRoleWithCredentials_SessionPolicy(t *testing.T) {
|
||||
service := setupTestSTSService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
sessionPolicy := `{"Version":"2012-10-17","Statement":[{"Effect":"Allow","Action":"filer:CreateEntry","Resource":"arn:aws:filer::path/user-docs/*"}]}`
|
||||
request := &AssumeRoleWithCredentialsRequest{
|
||||
RoleArn: "arn:aws:iam::role/TestRole",
|
||||
Username: "testuser",
|
||||
Password: "testpass",
|
||||
RoleSessionName: "test-session",
|
||||
ProviderName: "test-ldap",
|
||||
Policy: &sessionPolicy,
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithCredentials(ctx, request)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response)
|
||||
|
||||
sessionInfo, err := service.ValidateSessionToken(ctx, response.Credentials.SessionToken)
|
||||
require.NoError(t, err)
|
||||
|
||||
normalized, err := NormalizeSessionPolicy(sessionPolicy)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, normalized, sessionInfo.SessionPolicy)
|
||||
}
|
||||
@@ -879,21 +879,6 @@ func (s *STSService) calculateSessionDuration(durationSeconds *int64, tokenExpir
|
||||
return duration
|
||||
}
|
||||
|
||||
// extractSessionIdFromToken extracts session ID from JWT session token
|
||||
func (s *STSService) extractSessionIdFromToken(sessionToken string) string {
|
||||
// Validate JWT and extract session claims
|
||||
claims, err := s.tokenGenerator.ValidateJWTWithClaims(sessionToken)
|
||||
if err != nil {
|
||||
// For test compatibility, also handle direct session IDs
|
||||
if len(sessionToken) == 32 { // Typical session ID length
|
||||
return sessionToken
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
return claims.SessionId
|
||||
}
|
||||
|
||||
// validateAssumeRoleWithCredentialsRequest validates the credentials request parameters
|
||||
func (s *STSService) validateAssumeRoleWithCredentialsRequest(request *AssumeRoleWithCredentialsRequest) error {
|
||||
if request.RoleArn == "" {
|
||||
|
||||
@@ -1,778 +0,0 @@
|
||||
package sts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/providers"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// createSTSTestJWT creates a test JWT token for STS service tests
|
||||
func createSTSTestJWT(t *testing.T, issuer, subject string) string {
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
|
||||
"iss": issuer,
|
||||
"sub": subject,
|
||||
"aud": "test-client",
|
||||
"exp": time.Now().Add(time.Hour).Unix(),
|
||||
"iat": time.Now().Unix(),
|
||||
})
|
||||
|
||||
tokenString, err := token.SignedString([]byte("test-signing-key"))
|
||||
require.NoError(t, err)
|
||||
return tokenString
|
||||
}
|
||||
|
||||
// TestSTSServiceInitialization tests STS service initialization
|
||||
func TestSTSServiceInitialization(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
config *STSConfig
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "valid config",
|
||||
config: &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{time.Hour * 12},
|
||||
Issuer: "seaweedfs-sts",
|
||||
SigningKey: []byte("test-signing-key"),
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "missing signing key",
|
||||
config: &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
Issuer: "seaweedfs-sts",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid token duration",
|
||||
config: &STSConfig{
|
||||
TokenDuration: FlexibleDuration{-time.Hour},
|
||||
Issuer: "seaweedfs-sts",
|
||||
SigningKey: []byte("test-key"),
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
service := NewSTSService()
|
||||
|
||||
err := service.Initialize(tt.config)
|
||||
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, service.IsInitialized())
|
||||
|
||||
// Verify defaults if applicable
|
||||
if tt.config.Issuer == "" {
|
||||
assert.Equal(t, DefaultIssuer, service.Config.Issuer)
|
||||
}
|
||||
if tt.config.TokenDuration.Duration == 0 {
|
||||
assert.Equal(t, time.Duration(DefaultTokenDuration)*time.Second, service.Config.TokenDuration.Duration)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSTSServiceDefaults(t *testing.T) {
|
||||
service := NewSTSService()
|
||||
config := &STSConfig{
|
||||
SigningKey: []byte("test-signing-key"),
|
||||
// Missing duration and issuer
|
||||
}
|
||||
|
||||
err := service.Initialize(config)
|
||||
assert.NoError(t, err)
|
||||
|
||||
assert.Equal(t, DefaultIssuer, config.Issuer)
|
||||
assert.Equal(t, time.Duration(DefaultTokenDuration)*time.Second, config.TokenDuration.Duration)
|
||||
assert.Equal(t, time.Duration(DefaultMaxSessionLength)*time.Second, config.MaxSessionLength.Duration)
|
||||
}
|
||||
|
||||
// TestAssumeRoleWithWebIdentity tests role assumption with OIDC tokens
|
||||
func TestAssumeRoleWithWebIdentity(t *testing.T) {
|
||||
service := setupTestSTSService(t)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
roleArn string
|
||||
webIdentityToken string
|
||||
sessionName string
|
||||
durationSeconds *int64
|
||||
wantErr bool
|
||||
expectedSubject string
|
||||
}{
|
||||
{
|
||||
name: "successful role assumption",
|
||||
roleArn: "arn:aws:iam::role/TestRole",
|
||||
webIdentityToken: createSTSTestJWT(t, "test-issuer", "test-user-id"),
|
||||
sessionName: "test-session",
|
||||
durationSeconds: nil, // Use default
|
||||
wantErr: false,
|
||||
expectedSubject: "test-user-id",
|
||||
},
|
||||
{
|
||||
name: "invalid web identity token",
|
||||
roleArn: "arn:aws:iam::role/TestRole",
|
||||
webIdentityToken: "invalid-token",
|
||||
sessionName: "test-session",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "non-existent role",
|
||||
roleArn: "arn:aws:iam::role/NonExistentRole",
|
||||
webIdentityToken: createSTSTestJWT(t, "test-issuer", "test-user"),
|
||||
sessionName: "test-session",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "custom session duration",
|
||||
roleArn: "arn:aws:iam::role/TestRole",
|
||||
webIdentityToken: createSTSTestJWT(t, "test-issuer", "test-user"),
|
||||
sessionName: "test-session",
|
||||
durationSeconds: int64Ptr(7200), // 2 hours
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
request := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: tt.roleArn,
|
||||
WebIdentityToken: tt.webIdentityToken,
|
||||
RoleSessionName: tt.sessionName,
|
||||
DurationSeconds: tt.durationSeconds,
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithWebIdentity(ctx, request)
|
||||
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, response)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, response)
|
||||
assert.NotNil(t, response.Credentials)
|
||||
assert.NotNil(t, response.AssumedRoleUser)
|
||||
|
||||
// Verify credentials
|
||||
creds := response.Credentials
|
||||
assert.NotEmpty(t, creds.AccessKeyId)
|
||||
assert.NotEmpty(t, creds.SecretAccessKey)
|
||||
assert.NotEmpty(t, creds.SessionToken)
|
||||
assert.True(t, creds.Expiration.After(time.Now()))
|
||||
|
||||
// Verify assumed role user
|
||||
user := response.AssumedRoleUser
|
||||
assert.Equal(t, tt.roleArn, user.AssumedRoleId)
|
||||
assert.Contains(t, user.Arn, tt.sessionName)
|
||||
|
||||
if tt.expectedSubject != "" {
|
||||
assert.Equal(t, tt.expectedSubject, user.Subject)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAssumeRoleWithLDAP tests role assumption with LDAP credentials
|
||||
func TestAssumeRoleWithLDAP(t *testing.T) {
|
||||
service := setupTestSTSService(t)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
roleArn string
|
||||
username string
|
||||
password string
|
||||
sessionName string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "successful LDAP role assumption",
|
||||
roleArn: "arn:aws:iam::role/LDAPRole",
|
||||
username: "testuser",
|
||||
password: "testpass",
|
||||
sessionName: "ldap-session",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "invalid LDAP credentials",
|
||||
roleArn: "arn:aws:iam::role/LDAPRole",
|
||||
username: "testuser",
|
||||
password: "wrongpass",
|
||||
sessionName: "ldap-session",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
request := &AssumeRoleWithCredentialsRequest{
|
||||
RoleArn: tt.roleArn,
|
||||
Username: tt.username,
|
||||
Password: tt.password,
|
||||
RoleSessionName: tt.sessionName,
|
||||
ProviderName: "test-ldap",
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithCredentials(ctx, request)
|
||||
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, response)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, response)
|
||||
assert.NotNil(t, response.Credentials)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestSessionTokenValidation tests session token validation
|
||||
func TestSessionTokenValidation(t *testing.T) {
|
||||
service := setupTestSTSService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// First, create a session
|
||||
request := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/TestRole",
|
||||
WebIdentityToken: createSTSTestJWT(t, "test-issuer", "test-user"),
|
||||
RoleSessionName: "test-session",
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithWebIdentity(ctx, request)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response)
|
||||
|
||||
sessionToken := response.Credentials.SessionToken
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
token string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "valid session token",
|
||||
token: sessionToken,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "invalid session token",
|
||||
token: "invalid-session-token",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "empty session token",
|
||||
token: "",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
session, err := service.ValidateSessionToken(ctx, tt.token)
|
||||
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
assert.Nil(t, session)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, session)
|
||||
assert.Equal(t, "test-session", session.SessionName)
|
||||
assert.Equal(t, "arn:aws:iam::role/TestRole", session.RoleArn)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestSessionTokenPersistence tests that JWT tokens remain valid throughout their lifetime
|
||||
// Note: In the stateless JWT design, tokens cannot be revoked and remain valid until expiration
|
||||
func TestSessionTokenPersistence(t *testing.T) {
|
||||
service := setupTestSTSService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Create a session first
|
||||
request := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/TestRole",
|
||||
WebIdentityToken: createSTSTestJWT(t, "test-issuer", "test-user"),
|
||||
RoleSessionName: "test-session",
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithWebIdentity(ctx, request)
|
||||
require.NoError(t, err)
|
||||
|
||||
sessionToken := response.Credentials.SessionToken
|
||||
|
||||
// Verify token is valid initially
|
||||
session, err := service.ValidateSessionToken(ctx, sessionToken)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, session)
|
||||
assert.Equal(t, "test-session", session.SessionName)
|
||||
|
||||
// In a stateless JWT system, tokens remain valid throughout their lifetime
|
||||
// Multiple validations should all succeed as long as the token hasn't expired
|
||||
session2, err := service.ValidateSessionToken(ctx, sessionToken)
|
||||
assert.NoError(t, err, "Token should remain valid in stateless system")
|
||||
assert.NotNil(t, session2, "Session should be returned from JWT token")
|
||||
assert.Equal(t, session.SessionId, session2.SessionId, "Session ID should be consistent")
|
||||
}
|
||||
|
||||
// Helper functions
|
||||
|
||||
func setupTestSTSService(t *testing.T) *STSService {
|
||||
service := NewSTSService()
|
||||
|
||||
config := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{time.Hour * 12},
|
||||
Issuer: "test-sts",
|
||||
SigningKey: []byte("test-signing-key-32-characters-long"),
|
||||
}
|
||||
|
||||
err := service.Initialize(config)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set up mock trust policy validator (required for STS testing)
|
||||
mockValidator := &MockTrustPolicyValidator{}
|
||||
service.SetTrustPolicyValidator(mockValidator)
|
||||
|
||||
// Register test providers
|
||||
mockOIDCProvider := &MockIdentityProvider{
|
||||
name: "test-oidc",
|
||||
validTokens: map[string]*providers.TokenClaims{
|
||||
createSTSTestJWT(t, "test-issuer", "test-user"): {
|
||||
Subject: "test-user-id",
|
||||
Issuer: "test-issuer",
|
||||
Claims: map[string]interface{}{
|
||||
"email": "test@example.com",
|
||||
"name": "Test User",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
mockLDAPProvider := &MockIdentityProvider{
|
||||
name: "test-ldap",
|
||||
validCredentials: map[string]string{
|
||||
"testuser": "testpass",
|
||||
},
|
||||
}
|
||||
|
||||
service.RegisterProvider(mockOIDCProvider)
|
||||
service.RegisterProvider(mockLDAPProvider)
|
||||
|
||||
return service
|
||||
}
|
||||
|
||||
func int64Ptr(v int64) *int64 {
|
||||
return &v
|
||||
}
|
||||
|
||||
// Mock identity provider for testing
|
||||
type MockIdentityProvider struct {
|
||||
name string
|
||||
validTokens map[string]*providers.TokenClaims
|
||||
validCredentials map[string]string
|
||||
}
|
||||
|
||||
func (m *MockIdentityProvider) Name() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
func (m *MockIdentityProvider) GetIssuer() string {
|
||||
return "test-issuer" // This matches the issuer in the token claims
|
||||
}
|
||||
|
||||
func (m *MockIdentityProvider) Initialize(config interface{}) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *MockIdentityProvider) Authenticate(ctx context.Context, token string) (*providers.ExternalIdentity, error) {
|
||||
// First try to parse as JWT token
|
||||
if len(token) > 20 && strings.Count(token, ".") >= 2 {
|
||||
parsedToken, _, err := new(jwt.Parser).ParseUnverified(token, jwt.MapClaims{})
|
||||
if err == nil {
|
||||
if claims, ok := parsedToken.Claims.(jwt.MapClaims); ok {
|
||||
issuer, _ := claims["iss"].(string)
|
||||
subject, _ := claims["sub"].(string)
|
||||
|
||||
// Verify the issuer matches what we expect
|
||||
if issuer == "test-issuer" && subject != "" {
|
||||
return &providers.ExternalIdentity{
|
||||
UserID: subject,
|
||||
Email: subject + "@test-domain.com",
|
||||
DisplayName: "Test User " + subject,
|
||||
Provider: m.name,
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Handle legacy OIDC tokens (for backwards compatibility)
|
||||
if claims, exists := m.validTokens[token]; exists {
|
||||
email, _ := claims.GetClaimString("email")
|
||||
name, _ := claims.GetClaimString("name")
|
||||
|
||||
return &providers.ExternalIdentity{
|
||||
UserID: claims.Subject,
|
||||
Email: email,
|
||||
DisplayName: name,
|
||||
Provider: m.name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Handle LDAP credentials (username:password format)
|
||||
if m.validCredentials != nil {
|
||||
parts := strings.Split(token, ":")
|
||||
if len(parts) == 2 {
|
||||
username, password := parts[0], parts[1]
|
||||
if expectedPassword, exists := m.validCredentials[username]; exists && expectedPassword == password {
|
||||
return &providers.ExternalIdentity{
|
||||
UserID: username,
|
||||
Email: username + "@" + m.name + ".com",
|
||||
DisplayName: "Test User " + username,
|
||||
Provider: m.name,
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("unknown test token: %s", token)
|
||||
}
|
||||
|
||||
func (m *MockIdentityProvider) GetUserInfo(ctx context.Context, userID string) (*providers.ExternalIdentity, error) {
|
||||
return &providers.ExternalIdentity{
|
||||
UserID: userID,
|
||||
Email: userID + "@" + m.name + ".com",
|
||||
Provider: m.name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *MockIdentityProvider) ValidateToken(ctx context.Context, token string) (*providers.TokenClaims, error) {
|
||||
if claims, exists := m.validTokens[token]; exists {
|
||||
return claims, nil
|
||||
}
|
||||
return nil, fmt.Errorf("invalid token")
|
||||
}
|
||||
|
||||
// TestSessionDurationCappedByTokenExpiration tests that session duration is capped by the source token's exp claim
|
||||
func TestSessionDurationCappedByTokenExpiration(t *testing.T) {
|
||||
service := NewSTSService()
|
||||
|
||||
config := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour}, // Default: 1 hour
|
||||
MaxSessionLength: FlexibleDuration{time.Hour * 12},
|
||||
Issuer: "test-sts",
|
||||
SigningKey: []byte("test-signing-key-32-characters-long"),
|
||||
}
|
||||
|
||||
err := service.Initialize(config)
|
||||
require.NoError(t, err)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
durationSeconds *int64
|
||||
tokenExpiration *time.Time
|
||||
expectedMaxSeconds int64
|
||||
description string
|
||||
}{
|
||||
{
|
||||
name: "no token expiration - use default duration",
|
||||
durationSeconds: nil,
|
||||
tokenExpiration: nil,
|
||||
expectedMaxSeconds: 3600, // 1 hour default
|
||||
description: "When no token expiration is set, use the configured default duration",
|
||||
},
|
||||
{
|
||||
name: "token expires before default duration",
|
||||
durationSeconds: nil,
|
||||
tokenExpiration: timePtr(time.Now().Add(30 * time.Minute)),
|
||||
expectedMaxSeconds: 30 * 60, // 30 minutes
|
||||
description: "When token expires in 30 min, session should be capped at 30 min",
|
||||
},
|
||||
{
|
||||
name: "token expires after default duration - use default",
|
||||
durationSeconds: nil,
|
||||
tokenExpiration: timePtr(time.Now().Add(2 * time.Hour)),
|
||||
expectedMaxSeconds: 3600, // 1 hour default, since it's less than 2 hour token expiry
|
||||
description: "When token expires after default duration, use the default duration",
|
||||
},
|
||||
{
|
||||
name: "requested duration shorter than token expiry",
|
||||
durationSeconds: int64Ptr(1800), // 30 min requested
|
||||
tokenExpiration: timePtr(time.Now().Add(time.Hour)),
|
||||
expectedMaxSeconds: 1800, // 30 minutes as requested
|
||||
description: "When requested duration is shorter than token expiry, use requested duration",
|
||||
},
|
||||
{
|
||||
name: "requested duration longer than token expiry - cap at token expiry",
|
||||
durationSeconds: int64Ptr(3600), // 1 hour requested
|
||||
tokenExpiration: timePtr(time.Now().Add(15 * time.Minute)),
|
||||
expectedMaxSeconds: 15 * 60, // Capped at 15 minutes
|
||||
description: "When requested duration exceeds token expiry, cap at token expiry",
|
||||
},
|
||||
{
|
||||
name: "GitLab CI short-lived token scenario",
|
||||
durationSeconds: nil,
|
||||
tokenExpiration: timePtr(time.Now().Add(5 * time.Minute)),
|
||||
expectedMaxSeconds: 5 * 60, // 5 minutes
|
||||
description: "GitLab CI job with 5 minute timeout should result in 5 minute session",
|
||||
},
|
||||
{
|
||||
name: "already expired token - defense in depth",
|
||||
durationSeconds: nil,
|
||||
tokenExpiration: timePtr(time.Now().Add(-5 * time.Minute)), // Expired 5 minutes ago
|
||||
expectedMaxSeconds: 60, // 1 minute minimum
|
||||
description: "Already expired token should result in minimal 1 minute session",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
duration := service.calculateSessionDuration(tt.durationSeconds, tt.tokenExpiration)
|
||||
|
||||
// Allow 5 second tolerance for time calculations
|
||||
maxExpected := time.Duration(tt.expectedMaxSeconds+5) * time.Second
|
||||
minExpected := time.Duration(tt.expectedMaxSeconds-5) * time.Second
|
||||
|
||||
assert.GreaterOrEqual(t, duration, minExpected,
|
||||
"%s: duration %v should be >= %v", tt.description, duration, minExpected)
|
||||
assert.LessOrEqual(t, duration, maxExpected,
|
||||
"%s: duration %v should be <= %v", tt.description, duration, maxExpected)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAssumeRoleWithWebIdentityRespectsTokenExpiration tests end-to-end that session duration is capped
|
||||
func TestAssumeRoleWithWebIdentityRespectsTokenExpiration(t *testing.T) {
|
||||
service := NewSTSService()
|
||||
|
||||
config := &STSConfig{
|
||||
TokenDuration: FlexibleDuration{time.Hour},
|
||||
MaxSessionLength: FlexibleDuration{time.Hour * 12},
|
||||
Issuer: "test-sts",
|
||||
SigningKey: []byte("test-signing-key-32-characters-long"),
|
||||
}
|
||||
|
||||
err := service.Initialize(config)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set up mock trust policy validator
|
||||
mockValidator := &MockTrustPolicyValidator{}
|
||||
service.SetTrustPolicyValidator(mockValidator)
|
||||
|
||||
// Create a mock provider that returns tokens with short expiration
|
||||
shortLivedTokenExpiration := time.Now().Add(10 * time.Minute)
|
||||
mockProvider := &MockIdentityProviderWithExpiration{
|
||||
name: "short-lived-issuer",
|
||||
tokenExpiration: &shortLivedTokenExpiration,
|
||||
}
|
||||
service.RegisterProvider(mockProvider)
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Create a JWT token with short expiration
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
|
||||
"iss": "short-lived-issuer",
|
||||
"sub": "test-user",
|
||||
"aud": "test-client",
|
||||
"exp": shortLivedTokenExpiration.Unix(),
|
||||
"iat": time.Now().Unix(),
|
||||
})
|
||||
tokenString, err := token.SignedString([]byte("test-signing-key"))
|
||||
require.NoError(t, err)
|
||||
|
||||
request := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/TestRole",
|
||||
WebIdentityToken: tokenString,
|
||||
RoleSessionName: "test-session",
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithWebIdentity(ctx, request)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response)
|
||||
|
||||
// Verify the session expires at or before the token expiration
|
||||
// Allow 5 second tolerance
|
||||
assert.True(t, response.Credentials.Expiration.Before(shortLivedTokenExpiration.Add(5*time.Second)),
|
||||
"Session expiration (%v) should not exceed token expiration (%v)",
|
||||
response.Credentials.Expiration, shortLivedTokenExpiration)
|
||||
}
|
||||
|
||||
// MockIdentityProviderWithExpiration is a mock provider that returns tokens with configurable expiration
|
||||
type MockIdentityProviderWithExpiration struct {
|
||||
name string
|
||||
tokenExpiration *time.Time
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithExpiration) Name() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithExpiration) GetIssuer() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithExpiration) Initialize(config interface{}) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithExpiration) Authenticate(ctx context.Context, token string) (*providers.ExternalIdentity, error) {
|
||||
// Parse the token to get subject
|
||||
parsedToken, _, err := new(jwt.Parser).ParseUnverified(token, jwt.MapClaims{})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse token: %w", err)
|
||||
}
|
||||
|
||||
claims, ok := parsedToken.Claims.(jwt.MapClaims)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("invalid claims")
|
||||
}
|
||||
|
||||
subject, _ := claims["sub"].(string)
|
||||
|
||||
identity := &providers.ExternalIdentity{
|
||||
UserID: subject,
|
||||
Email: subject + "@example.com",
|
||||
DisplayName: "Test User",
|
||||
Provider: m.name,
|
||||
TokenExpiration: m.tokenExpiration,
|
||||
}
|
||||
|
||||
return identity, nil
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithExpiration) GetUserInfo(ctx context.Context, userID string) (*providers.ExternalIdentity, error) {
|
||||
return &providers.ExternalIdentity{
|
||||
UserID: userID,
|
||||
Provider: m.name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithExpiration) ValidateToken(ctx context.Context, token string) (*providers.TokenClaims, error) {
|
||||
claims := &providers.TokenClaims{
|
||||
Subject: "test-user",
|
||||
Issuer: m.name,
|
||||
}
|
||||
if m.tokenExpiration != nil {
|
||||
claims.ExpiresAt = *m.tokenExpiration
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
func timePtr(t time.Time) *time.Time {
|
||||
return &t
|
||||
}
|
||||
|
||||
// TestAssumeRoleWithWebIdentity_PreservesAttributes tests that attributes from the identity provider
|
||||
// are correctly propagated to the session token's request context
|
||||
func TestAssumeRoleWithWebIdentity_PreservesAttributes(t *testing.T) {
|
||||
service := setupTestSTSService(t)
|
||||
|
||||
// Create a mock provider that returns a user with attributes
|
||||
mockProvider := &MockIdentityProviderWithAttributes{
|
||||
name: "attr-provider",
|
||||
attributes: map[string]string{
|
||||
"preferred_username": "my-user",
|
||||
"department": "engineering",
|
||||
"project": "seaweedfs",
|
||||
},
|
||||
}
|
||||
service.RegisterProvider(mockProvider)
|
||||
|
||||
// Create a valid JWT token for the provider
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
|
||||
"iss": "attr-provider",
|
||||
"sub": "test-user-id",
|
||||
"aud": "test-client",
|
||||
"exp": time.Now().Add(time.Hour).Unix(),
|
||||
"iat": time.Now().Unix(),
|
||||
})
|
||||
tokenString, err := token.SignedString([]byte("test-signing-key"))
|
||||
require.NoError(t, err)
|
||||
|
||||
ctx := context.Background()
|
||||
request := &AssumeRoleWithWebIdentityRequest{
|
||||
RoleArn: "arn:aws:iam::role/TestRole",
|
||||
WebIdentityToken: tokenString,
|
||||
RoleSessionName: "test-session",
|
||||
}
|
||||
|
||||
response, err := service.AssumeRoleWithWebIdentity(ctx, request)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response)
|
||||
|
||||
// Validate the session token to check claims
|
||||
sessionInfo, err := service.ValidateSessionToken(ctx, response.Credentials.SessionToken)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Check that attributes are present in RequestContext
|
||||
require.NotNil(t, sessionInfo.RequestContext, "RequestContext should not be nil")
|
||||
assert.Equal(t, "my-user", sessionInfo.RequestContext["preferred_username"])
|
||||
assert.Equal(t, "engineering", sessionInfo.RequestContext["department"])
|
||||
assert.Equal(t, "seaweedfs", sessionInfo.RequestContext["project"])
|
||||
|
||||
// Check standard claims are also present
|
||||
assert.Equal(t, "test-user-id", sessionInfo.RequestContext["sub"])
|
||||
assert.Equal(t, "test@example.com", sessionInfo.RequestContext["email"])
|
||||
assert.Equal(t, "Test User", sessionInfo.RequestContext["name"])
|
||||
}
|
||||
|
||||
// MockIdentityProviderWithAttributes is a mock provider that returns configured attributes
|
||||
type MockIdentityProviderWithAttributes struct {
|
||||
name string
|
||||
attributes map[string]string
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithAttributes) Name() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithAttributes) GetIssuer() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithAttributes) Initialize(config interface{}) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithAttributes) Authenticate(ctx context.Context, token string) (*providers.ExternalIdentity, error) {
|
||||
return &providers.ExternalIdentity{
|
||||
UserID: "test-user-id",
|
||||
Email: "test@example.com",
|
||||
DisplayName: "Test User",
|
||||
Provider: m.name,
|
||||
Attributes: m.attributes,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithAttributes) GetUserInfo(ctx context.Context, userID string) (*providers.ExternalIdentity, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (m *MockIdentityProviderWithAttributes) ValidateToken(ctx context.Context, token string) (*providers.TokenClaims, error) {
|
||||
return &providers.TokenClaims{
|
||||
Subject: "test-user-id",
|
||||
Issuer: m.name,
|
||||
}, nil
|
||||
}
|
||||
@@ -1,53 +1,4 @@
|
||||
package sts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/providers"
|
||||
)
|
||||
|
||||
// MockTrustPolicyValidator is a simple mock for testing STS functionality
|
||||
type MockTrustPolicyValidator struct{}
|
||||
|
||||
// ValidateTrustPolicyForWebIdentity allows valid JWT test tokens for STS testing
|
||||
func (m *MockTrustPolicyValidator) ValidateTrustPolicyForWebIdentity(ctx context.Context, roleArn string, webIdentityToken string, durationSeconds *int64) error {
|
||||
// Reject non-existent roles for testing
|
||||
if strings.Contains(roleArn, "NonExistentRole") {
|
||||
return fmt.Errorf("trust policy validation failed: role does not exist")
|
||||
}
|
||||
|
||||
// For STS unit tests, allow JWT tokens that look valid (contain dots for JWT structure)
|
||||
// In real implementation, this would validate against actual trust policies
|
||||
if len(webIdentityToken) > 20 && strings.Count(webIdentityToken, ".") >= 2 {
|
||||
// This appears to be a JWT token - allow it for testing
|
||||
return nil
|
||||
}
|
||||
|
||||
// Legacy support for specific test tokens during migration
|
||||
if webIdentityToken == "valid_test_token" || webIdentityToken == "valid-oidc-token" {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Reject invalid tokens
|
||||
if webIdentityToken == "invalid_token" || webIdentityToken == "expired_token" || webIdentityToken == "invalid-token" {
|
||||
return fmt.Errorf("trust policy denies token")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateTrustPolicyForCredentials allows valid test identities for STS testing
|
||||
func (m *MockTrustPolicyValidator) ValidateTrustPolicyForCredentials(ctx context.Context, roleArn string, identity *providers.ExternalIdentity) error {
|
||||
// Reject non-existent roles for testing
|
||||
if strings.Contains(roleArn, "NonExistentRole") {
|
||||
return fmt.Errorf("trust policy validation failed: role does not exist")
|
||||
}
|
||||
|
||||
// For STS unit tests, allow test identities
|
||||
if identity != nil && identity.UserID != "" {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("invalid identity for role assumption")
|
||||
}
|
||||
|
||||
@@ -1,29 +0,0 @@
|
||||
package images
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
/*
|
||||
* Preprocess image files on client side.
|
||||
* 1. possibly adjust the orientation
|
||||
* 2. resize the image to a width or height limit
|
||||
* 3. remove the exif data
|
||||
* Call this function on any file uploaded to SeaweedFS
|
||||
*
|
||||
*/
|
||||
func MaybePreprocessImage(filename string, data []byte, width, height int) (resized io.ReadSeeker, w int, h int) {
|
||||
ext := filepath.Ext(filename)
|
||||
ext = strings.ToLower(ext)
|
||||
switch ext {
|
||||
case ".png", ".gif":
|
||||
return Resized(ext, bytes.NewReader(data), width, height, "")
|
||||
case ".jpg", ".jpeg":
|
||||
data = FixJpgOrientation(data)
|
||||
return Resized(ext, bytes.NewReader(data), width, height, "")
|
||||
}
|
||||
return bytes.NewReader(data), 0, 0
|
||||
}
|
||||
@@ -290,15 +290,6 @@ func (loader *ConfigLoader) ValidateConfiguration() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// LoadKMSFromFilerToml is a convenience function to load KMS configuration from filer.toml
|
||||
func LoadKMSFromFilerToml(v ViperConfig) error {
|
||||
loader := NewConfigLoader(v)
|
||||
if err := loader.LoadConfigurations(); err != nil {
|
||||
return err
|
||||
}
|
||||
return loader.ValidateConfiguration()
|
||||
}
|
||||
|
||||
// LoadKMSFromConfig loads KMS configuration directly from parsed JSON data
|
||||
func LoadKMSFromConfig(kmsConfig interface{}) error {
|
||||
kmsMap, ok := kmsConfig.(map[string]interface{})
|
||||
@@ -415,12 +406,3 @@ func getIntFromConfig(config map[string]interface{}, key string, defaultValue in
|
||||
}
|
||||
return defaultValue
|
||||
}
|
||||
|
||||
func getStringFromConfig(config map[string]interface{}, key string, defaultValue string) string {
|
||||
if value, exists := config[key]; exists {
|
||||
if stringValue, ok := value.(string); ok {
|
||||
return stringValue
|
||||
}
|
||||
}
|
||||
return defaultValue
|
||||
}
|
||||
|
||||
@@ -147,13 +147,6 @@ func (fh *FileHandle) ReleaseHandle() {
|
||||
}
|
||||
}
|
||||
|
||||
func lessThan(a, b *filer_pb.FileChunk) bool {
|
||||
if a.ModifiedTsNs == b.ModifiedTsNs {
|
||||
return a.Fid.FileKey < b.Fid.FileKey
|
||||
}
|
||||
return a.ModifiedTsNs < b.ModifiedTsNs
|
||||
}
|
||||
|
||||
// getCumulativeOffsets returns cached cumulative offsets for chunks, computing them if necessary
|
||||
func (fh *FileHandle) getCumulativeOffsets(chunks []*filer_pb.FileChunk) []int64 {
|
||||
fh.chunkCacheLock.RLock()
|
||||
|
||||
@@ -21,9 +21,3 @@ func min(x, y int64) int64 {
|
||||
}
|
||||
return y
|
||||
}
|
||||
func minInt(x, y int) int {
|
||||
if x < y {
|
||||
return x
|
||||
}
|
||||
return y
|
||||
}
|
||||
|
||||
@@ -119,13 +119,6 @@ func (c *RDMAMountClient) lookupVolumeLocationByFileID(ctx context.Context, file
|
||||
return bestAddress, nil
|
||||
}
|
||||
|
||||
// lookupVolumeLocation finds the best volume server for a given volume ID (legacy method)
|
||||
func (c *RDMAMountClient) lookupVolumeLocation(ctx context.Context, volumeID uint32, needleID uint64, cookie uint32) (string, error) {
|
||||
// Create a file ID for lookup (format: volumeId,needleId,cookie)
|
||||
fileID := fmt.Sprintf("%d,%x,%d", volumeID, needleID, cookie)
|
||||
return c.lookupVolumeLocationByFileID(ctx, fileID)
|
||||
}
|
||||
|
||||
// healthCheck verifies that the RDMA sidecar is available and functioning
|
||||
func (c *RDMAMountClient) healthCheck() error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), c.timeout)
|
||||
|
||||
@@ -117,11 +117,6 @@ func GetBrokerErrorInfo(code int32) BrokerErrorInfo {
|
||||
}
|
||||
}
|
||||
|
||||
// GetKafkaErrorCode returns the corresponding Kafka protocol error code for a broker error
|
||||
func GetKafkaErrorCode(brokerErrorCode int32) int16 {
|
||||
return GetBrokerErrorInfo(brokerErrorCode).KafkaCode
|
||||
}
|
||||
|
||||
// CreateBrokerError creates a structured broker error with both error code and message
|
||||
func CreateBrokerError(code int32, message string) (int32, string) {
|
||||
info := GetBrokerErrorInfo(code)
|
||||
|
||||
@@ -1,351 +0,0 @@
|
||||
package broker
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/mq/topic"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/schema_pb"
|
||||
)
|
||||
|
||||
func createTestTopic() topic.Topic {
|
||||
return topic.Topic{
|
||||
Namespace: "test",
|
||||
Name: "offset-test",
|
||||
}
|
||||
}
|
||||
|
||||
func createTestPartition() topic.Partition {
|
||||
return topic.Partition{
|
||||
RingSize: 1024,
|
||||
RangeStart: 0,
|
||||
RangeStop: 31,
|
||||
UnixTimeNs: time.Now().UnixNano(),
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokerOffsetManager_AssignOffset(t *testing.T) {
|
||||
storage := NewInMemoryOffsetStorageForTesting()
|
||||
manager := NewBrokerOffsetManagerWithStorage(storage)
|
||||
testTopic := createTestTopic()
|
||||
testPartition := createTestPartition()
|
||||
|
||||
// Test sequential offset assignment
|
||||
for i := int64(0); i < 10; i++ {
|
||||
assignedOffset, err := manager.AssignOffset(testTopic, testPartition)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to assign offset %d: %v", i, err)
|
||||
}
|
||||
|
||||
if assignedOffset != i {
|
||||
t.Errorf("Expected offset %d, got %d", i, assignedOffset)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokerOffsetManager_AssignBatchOffsets(t *testing.T) {
|
||||
storage := NewInMemoryOffsetStorageForTesting()
|
||||
manager := NewBrokerOffsetManagerWithStorage(storage)
|
||||
testTopic := createTestTopic()
|
||||
testPartition := createTestPartition()
|
||||
|
||||
// Assign batch of offsets
|
||||
baseOffset, lastOffset, err := manager.AssignBatchOffsets(testTopic, testPartition, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to assign batch offsets: %v", err)
|
||||
}
|
||||
|
||||
if baseOffset != 0 {
|
||||
t.Errorf("Expected base offset 0, got %d", baseOffset)
|
||||
}
|
||||
|
||||
if lastOffset != 4 {
|
||||
t.Errorf("Expected last offset 4, got %d", lastOffset)
|
||||
}
|
||||
|
||||
// Assign another batch
|
||||
baseOffset2, lastOffset2, err := manager.AssignBatchOffsets(testTopic, testPartition, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to assign second batch offsets: %v", err)
|
||||
}
|
||||
|
||||
if baseOffset2 != 5 {
|
||||
t.Errorf("Expected base offset 5, got %d", baseOffset2)
|
||||
}
|
||||
|
||||
if lastOffset2 != 7 {
|
||||
t.Errorf("Expected last offset 7, got %d", lastOffset2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokerOffsetManager_GetHighWaterMark(t *testing.T) {
|
||||
storage := NewInMemoryOffsetStorageForTesting()
|
||||
manager := NewBrokerOffsetManagerWithStorage(storage)
|
||||
testTopic := createTestTopic()
|
||||
testPartition := createTestPartition()
|
||||
|
||||
// Initially should be 0
|
||||
hwm, err := manager.GetHighWaterMark(testTopic, testPartition)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get initial high water mark: %v", err)
|
||||
}
|
||||
|
||||
if hwm != 0 {
|
||||
t.Errorf("Expected initial high water mark 0, got %d", hwm)
|
||||
}
|
||||
|
||||
// Assign some offsets
|
||||
manager.AssignBatchOffsets(testTopic, testPartition, 10)
|
||||
|
||||
// High water mark should be updated
|
||||
hwm, err = manager.GetHighWaterMark(testTopic, testPartition)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get high water mark after assignment: %v", err)
|
||||
}
|
||||
|
||||
if hwm != 10 {
|
||||
t.Errorf("Expected high water mark 10, got %d", hwm)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokerOffsetManager_CreateSubscription(t *testing.T) {
|
||||
storage := NewInMemoryOffsetStorageForTesting()
|
||||
manager := NewBrokerOffsetManagerWithStorage(storage)
|
||||
testTopic := createTestTopic()
|
||||
testPartition := createTestPartition()
|
||||
|
||||
// Assign some offsets first
|
||||
manager.AssignBatchOffsets(testTopic, testPartition, 5)
|
||||
|
||||
// Create subscription
|
||||
sub, err := manager.CreateSubscription(
|
||||
"test-sub",
|
||||
testTopic,
|
||||
testPartition,
|
||||
schema_pb.OffsetType_RESET_TO_EARLIEST,
|
||||
0,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create subscription: %v", err)
|
||||
}
|
||||
|
||||
if sub.ID != "test-sub" {
|
||||
t.Errorf("Expected subscription ID 'test-sub', got %s", sub.ID)
|
||||
}
|
||||
|
||||
if sub.StartOffset != 0 {
|
||||
t.Errorf("Expected start offset 0, got %d", sub.StartOffset)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokerOffsetManager_GetPartitionOffsetInfo(t *testing.T) {
|
||||
storage := NewInMemoryOffsetStorageForTesting()
|
||||
manager := NewBrokerOffsetManagerWithStorage(storage)
|
||||
testTopic := createTestTopic()
|
||||
testPartition := createTestPartition()
|
||||
|
||||
// Test empty partition
|
||||
info, err := manager.GetPartitionOffsetInfo(testTopic, testPartition)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get partition offset info: %v", err)
|
||||
}
|
||||
|
||||
if info.EarliestOffset != 0 {
|
||||
t.Errorf("Expected earliest offset 0, got %d", info.EarliestOffset)
|
||||
}
|
||||
|
||||
if info.LatestOffset != -1 {
|
||||
t.Errorf("Expected latest offset -1 for empty partition, got %d", info.LatestOffset)
|
||||
}
|
||||
|
||||
// Assign offsets and test again
|
||||
manager.AssignBatchOffsets(testTopic, testPartition, 5)
|
||||
|
||||
info, err = manager.GetPartitionOffsetInfo(testTopic, testPartition)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get partition offset info after assignment: %v", err)
|
||||
}
|
||||
|
||||
if info.LatestOffset != 4 {
|
||||
t.Errorf("Expected latest offset 4, got %d", info.LatestOffset)
|
||||
}
|
||||
|
||||
if info.HighWaterMark != 5 {
|
||||
t.Errorf("Expected high water mark 5, got %d", info.HighWaterMark)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokerOffsetManager_MultiplePartitions(t *testing.T) {
|
||||
storage := NewInMemoryOffsetStorageForTesting()
|
||||
manager := NewBrokerOffsetManagerWithStorage(storage)
|
||||
testTopic := createTestTopic()
|
||||
|
||||
// Create different partitions
|
||||
partition1 := topic.Partition{
|
||||
RingSize: 1024,
|
||||
RangeStart: 0,
|
||||
RangeStop: 31,
|
||||
UnixTimeNs: time.Now().UnixNano(),
|
||||
}
|
||||
|
||||
partition2 := topic.Partition{
|
||||
RingSize: 1024,
|
||||
RangeStart: 32,
|
||||
RangeStop: 63,
|
||||
UnixTimeNs: time.Now().UnixNano(),
|
||||
}
|
||||
|
||||
// Assign offsets to different partitions
|
||||
assignedOffset1, err := manager.AssignOffset(testTopic, partition1)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to assign offset to partition1: %v", err)
|
||||
}
|
||||
|
||||
assignedOffset2, err := manager.AssignOffset(testTopic, partition2)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to assign offset to partition2: %v", err)
|
||||
}
|
||||
|
||||
// Both should start at 0
|
||||
if assignedOffset1 != 0 {
|
||||
t.Errorf("Expected offset 0 for partition1, got %d", assignedOffset1)
|
||||
}
|
||||
|
||||
if assignedOffset2 != 0 {
|
||||
t.Errorf("Expected offset 0 for partition2, got %d", assignedOffset2)
|
||||
}
|
||||
|
||||
// Assign more offsets to partition1
|
||||
assignedOffset1_2, err := manager.AssignOffset(testTopic, partition1)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to assign second offset to partition1: %v", err)
|
||||
}
|
||||
|
||||
if assignedOffset1_2 != 1 {
|
||||
t.Errorf("Expected offset 1 for partition1, got %d", assignedOffset1_2)
|
||||
}
|
||||
|
||||
// Partition2 should still be at 0 for next assignment
|
||||
assignedOffset2_2, err := manager.AssignOffset(testTopic, partition2)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to assign second offset to partition2: %v", err)
|
||||
}
|
||||
|
||||
if assignedOffset2_2 != 1 {
|
||||
t.Errorf("Expected offset 1 for partition2, got %d", assignedOffset2_2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOffsetAwarePublisher(t *testing.T) {
|
||||
storage := NewInMemoryOffsetStorageForTesting()
|
||||
manager := NewBrokerOffsetManagerWithStorage(storage)
|
||||
testTopic := createTestTopic()
|
||||
testPartition := createTestPartition()
|
||||
|
||||
// Create a mock local partition (simplified for testing)
|
||||
localPartition := &topic.LocalPartition{}
|
||||
|
||||
// Create offset assignment function
|
||||
assignOffsetFn := func() (int64, error) {
|
||||
return manager.AssignOffset(testTopic, testPartition)
|
||||
}
|
||||
|
||||
// Create offset-aware publisher
|
||||
publisher := topic.NewOffsetAwarePublisher(localPartition, assignOffsetFn)
|
||||
|
||||
if publisher.GetPartition() != localPartition {
|
||||
t.Error("Publisher should return the correct partition")
|
||||
}
|
||||
|
||||
// Test would require more setup to actually publish messages
|
||||
// This tests the basic structure
|
||||
}
|
||||
|
||||
func TestBrokerOffsetManager_GetOffsetMetrics(t *testing.T) {
|
||||
storage := NewInMemoryOffsetStorageForTesting()
|
||||
manager := NewBrokerOffsetManagerWithStorage(storage)
|
||||
testTopic := createTestTopic()
|
||||
testPartition := createTestPartition()
|
||||
|
||||
// Initial metrics
|
||||
metrics := manager.GetOffsetMetrics()
|
||||
if metrics.TotalOffsets != 0 {
|
||||
t.Errorf("Expected 0 total offsets initially, got %d", metrics.TotalOffsets)
|
||||
}
|
||||
|
||||
// Assign some offsets
|
||||
manager.AssignBatchOffsets(testTopic, testPartition, 5)
|
||||
|
||||
// Create subscription
|
||||
manager.CreateSubscription("test-sub", testTopic, testPartition, schema_pb.OffsetType_RESET_TO_EARLIEST, 0)
|
||||
|
||||
// Check updated metrics
|
||||
metrics = manager.GetOffsetMetrics()
|
||||
if metrics.PartitionCount != 1 {
|
||||
t.Errorf("Expected 1 partition, got %d", metrics.PartitionCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokerOffsetManager_AssignOffsetsWithResult(t *testing.T) {
|
||||
storage := NewInMemoryOffsetStorageForTesting()
|
||||
manager := NewBrokerOffsetManagerWithStorage(storage)
|
||||
testTopic := createTestTopic()
|
||||
testPartition := createTestPartition()
|
||||
|
||||
// Assign offsets with result
|
||||
result := manager.AssignOffsetsWithResult(testTopic, testPartition, 3)
|
||||
|
||||
if result.Error != nil {
|
||||
t.Fatalf("Expected no error, got: %v", result.Error)
|
||||
}
|
||||
|
||||
if result.BaseOffset != 0 {
|
||||
t.Errorf("Expected base offset 0, got %d", result.BaseOffset)
|
||||
}
|
||||
|
||||
if result.LastOffset != 2 {
|
||||
t.Errorf("Expected last offset 2, got %d", result.LastOffset)
|
||||
}
|
||||
|
||||
if result.Count != 3 {
|
||||
t.Errorf("Expected count 3, got %d", result.Count)
|
||||
}
|
||||
|
||||
if result.Topic != testTopic {
|
||||
t.Error("Topic mismatch in result")
|
||||
}
|
||||
|
||||
if result.Partition != testPartition {
|
||||
t.Error("Partition mismatch in result")
|
||||
}
|
||||
|
||||
if result.Timestamp <= 0 {
|
||||
t.Error("Timestamp should be set")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokerOffsetManager_Shutdown(t *testing.T) {
|
||||
storage := NewInMemoryOffsetStorageForTesting()
|
||||
manager := NewBrokerOffsetManagerWithStorage(storage)
|
||||
testTopic := createTestTopic()
|
||||
testPartition := createTestPartition()
|
||||
|
||||
// Assign some offsets and create subscriptions
|
||||
manager.AssignBatchOffsets(testTopic, testPartition, 5)
|
||||
manager.CreateSubscription("test-sub", testTopic, testPartition, schema_pb.OffsetType_RESET_TO_EARLIEST, 0)
|
||||
|
||||
// Shutdown should not panic
|
||||
manager.Shutdown()
|
||||
|
||||
// After shutdown, operations should still work (using new managers)
|
||||
offset, err := manager.AssignOffset(testTopic, testPartition)
|
||||
if err != nil {
|
||||
t.Fatalf("Operations should still work after shutdown: %v", err)
|
||||
}
|
||||
|
||||
// Should start from 0 again (new manager)
|
||||
if offset != 0 {
|
||||
t.Errorf("Expected offset 0 after shutdown, got %d", offset)
|
||||
}
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user