diff --git a/.planning/phases/32-cli-instance-identity-and-port-alias-migration/32-01-SUMMARY.md b/.planning/phases/32-cli-instance-identity-and-port-alias-migration/32-01-SUMMARY.md deleted file mode 100644 index 9596ae0..0000000 --- a/.planning/phases/32-cli-instance-identity-and-port-alias-migration/32-01-SUMMARY.md +++ /dev/null @@ -1,161 +0,0 @@ ---- -phase: 32-cli-instance-identity-and-port-alias-migration -plan: 01 -subsystem: cli -tags: [go, cobra, runtime, lifecycle, identity, migration] - -# Dependency graph -requires: - - phase: 31-instance-scoped-aws-resources-and-future-resource-guardrails - provides: instanceId as canonical bootstrap identity for AWS-backed resources - -provides: - - InstanceID field in runtime.Instance and Manager.SetInstanceID() method - - SaveActiveInstanceWithID and SaveSavedInstanceWithID storage helpers - - Legacy port-keyed records still load and InstanceID falls back from live snapshot - - NewStatusCommand as a thin alias delegating to NewInstancesCommand - - instancesToRuntime() that merges storage summaries with live snapshot identity - - instancePayload JSON schema extended with instanceId field - - Full regression coverage for alias parity and canonical identity in JSON output - -affects: [33-aws-account-identity, desktop-app-instances-view] - -# Tech tracking -tech-stack: - added: [] - patterns: - - "Manager.SetInstanceID(): caller sets canonical identity once at bootstrap, snapshot embeds it into every Instance" - - "instancesToRuntime() fallback: storage summaries without instanceId inherit from live snapshot by port" - - "Thin alias pattern: NewStatusCommand wraps NewInstancesCommand with different Use/Short, no separate rendering path" - - "saveInstanceWithID(): single internal helper accepts instanceID string, empty string preserves compatibility" - -key-files: - created: [] - modified: - - core/internal/application/runtime/manager.go - - core/internal/application/runtime/manager_test.go - - core/internal/delivery/cli/storage.go - - core/internal/delivery/cli/storage_test.go - - core/internal/delivery/cli/status.go - - core/internal/delivery/cli/presenter.go - - core/internal/delivery/cli/output.go - - core/internal/delivery/cli/root.go - - core/internal/delivery/cli/root_test.go - - core/internal/delivery/cli/presenter_test.go - - core/internal/delivery/cli/commands_test.go - - core/cmd/mildstack/main.go - -key-decisions: - - "InstanceID is set once on Manager via SetInstanceID() at bootstrap; snapshot embeds it into all instances rather than computing it per-Serve call" - - "instancesToRuntime() uses live snapshot as fallback index for instanceId when storage records are legacy port-keyed" - - "status command is a thin Cobra alias (cloned command with different Use) to guarantee rendering parity without a second code path" - - "saveInstanceWithID() is the single internal helper; SaveActiveInstance/SaveSavedInstance pass empty string to preserve backward compatibility" - - "instanceId is omitempty in JSON so legacy records without an id produce valid output during the migration window" - -patterns-established: - - "canonical identity is seeded at bootstrap and flows down through snapshot copies, not computed at presentation time" - - "alias commands are created by cloning the primary command struct and changing Use/Short, not by wrapping the handler" - - "storage summaries fall back to live snapshot identity by port when the on-disk record predates the instanceId field" - -requirements-completed: [] - -# Metrics -duration: 6min -completed: 2026-04-19 ---- - -# Phase 32 Plan 01: CLI Instance Identity and Port Alias Migration Summary - -**instanceId promoted as canonical CLI lifecycle identity via Manager.SetInstanceID(), migration-safe storage helpers, a status alias delegating to instances, and JSON payload extended with instanceId field** - -## Performance - -- **Duration:** 6 min -- **Started:** 2026-04-19T03:38:47Z -- **Completed:** 2026-04-19T03:44:47Z -- **Tasks:** 3 -- **Files modified:** 12 - -## Accomplishments - -- `runtime.Instance` and `Manager` gained `InstanceID` field and `SetInstanceID()` method; snapshots now embed canonical identity into every instance copy -- CLI storage extended with `SaveActiveInstanceWithID` / `SaveSavedInstanceWithID`; `instanceRecord` and `instanceSummary` carry `instanceId,omitempty`; legacy port-keyed records without the field still load and fall back to live snapshot identity via port index -- `NewStatusCommand` implemented as a thin Cobra alias of `NewInstancesCommand` (same handler, different `Use`/`Short`), registered in `Commands` struct and wired in `main.go` -- `instancePayload` JSON schema extended with `instanceId` field; presenter payload clone helpers propagate it -- `instancesToRuntime()` updated to accept live snapshot instances as fallback source for `instanceId` when storage records are legacy -- Regression suite covers alias output parity (human and JSON), canonical identity in JSON payload, and copy-safe snapshot behavior after identity migration - -## Task Commits - -1. **RED (Task 1): add failing tests for instanceId in runtime snapshots and storage** - `fad62a27` (test) -2. **Task 1: Promote instanceId into runtime snapshots and storage migration helpers** - `000eecae` (feat) -3. **RED (Task 2): add failing tests for status alias and instanceId in presenter payload** - `64e5fab2` (test) -4. **Task 2: Make status a thin alias of instances and extend payload schema** - `44cc7db4` (feat) -5. **Task 3: Extend regression coverage for lifecycle commands and binary wiring** - `7c8873d8` (feat) - -## Files Created/Modified - -- `core/internal/application/runtime/manager.go` - Added `InstanceID` to `Instance`, `instanceID` field to `Manager`, `SetInstanceID()` method, updated `runningInstances()` signature -- `core/internal/application/runtime/manager_test.go` - Added `TestManagerInstanceCarriesInstanceID` and `TestNewWithPortsSeedsInstanceIDFromRegisteredIdentity` -- `core/internal/delivery/cli/storage.go` - Added `InstanceID` to `instanceRecord`/`instanceSummary`, `saveInstanceWithID()` helper, `SaveActiveInstanceWithID`, `SaveSavedInstanceWithID`; updated `LoadInstances()` to propagate `InstanceID` with active-record-wins merge -- `core/internal/delivery/cli/storage_test.go` - Added `TestStorageInstanceSummaryCarriesInstanceID` and `TestStorageLegacyPortKeyedRecordLoadsAndPresentsWithInstanceID` -- `core/internal/delivery/cli/status.go` - Added `NewStatusCommand()` alias; updated `instancesToRuntime()` to accept live snapshot instances for `instanceId` fallback -- `core/internal/delivery/cli/presenter.go` - Updated `cloneInstances()` and `cloneInstancesPayload()` to copy `InstanceID` -- `core/internal/delivery/cli/output.go` - Added `InstanceID string` with `json:"instanceId,omitempty"` to `instancePayload` -- `core/internal/delivery/cli/root.go` - Added `Status *cobra.Command` to `Commands` struct; registered in `NewRootCommand` -- `core/internal/delivery/cli/root_test.go` - Added `TestNewRootCommandRegistersStatusAlias` -- `core/internal/delivery/cli/presenter_test.go` - Added `TestPresenterStatusPayloadIncludesInstanceID` with copy-safety assertion -- `core/internal/delivery/cli/commands_test.go` - Added `TestStatusAliasMatchesInstancesOutput`, `TestStatusAliasJSONMatchesInstancesJSON`, `TestCommandsServeInstancesJSONIncludesInstanceID`; wired `Status` into `newTestCommand()`; updated `TestCommandsServeInstancesJSON` to assert `instanceId` non-empty -- `core/cmd/mildstack/main.go` - Added `manager.SetInstanceID(instanceID)` after manager creation; added `Status: cli.NewStatusCommand(manager, storage)` to commands - -## Decisions Made - -- `SetInstanceID()` is called once at bootstrap in `main.go` rather than threading the id through every `Serve` call — keeps the identity model simple and the `Manager` interface stable -- `instancesToRuntime()` uses the live snapshot as a fallback index by port — avoids requiring every operator to immediately upgrade their storage records while still surfacing the canonical id when the manager knows it -- `status` alias is a Cobra command clone (copy of the `instances` command with `Use`/`Short` overwritten) rather than a `cobra.Command.AddCommand` alias — this guarantees identical flag parsing, output, and empty-state handling without a second implementation branch - -## Deviations from Plan - -### Auto-fixed Issues - -**1. [Rule 1 - Bug] Fixed TestManagerInstanceCarriesInstanceID to require SetInstanceID before snapshot** -- **Found during:** Task 1 (GREEN phase) -- **Issue:** Test created manager without calling `SetInstanceID` but expected non-empty `InstanceID` — design requires explicit seeding before snapshot -- **Fix:** Updated test to call `manager.SetInstanceID("test-instance-abc")` before `Serve` and use exact string assertion instead of non-empty check -- **Files modified:** `core/internal/application/runtime/manager_test.go` -- **Committed in:** `000eecae` (Task 1 feat commit) - -**2. [Rule 1 - Bug] Fixed instancesToRuntime to fall back to live snapshot for instanceId** -- **Found during:** Task 3 (TestCommandsServeInstancesJSON and TestCommandsServeInstancesJSONIncludesInstanceID failing) -- **Issue:** `commandServerStub.Start()` calls `SaveActiveInstance(port)` without an `instanceId`, so storage summaries loaded by `NewInstancesCommand` have empty `InstanceID`. The manager snapshot carries the id but it was discarded when the snapshot instances were overwritten with storage data -- **Fix:** Updated `instancesToRuntime()` to accept live snapshot instances and build a `port -> InstanceID` fallback index; storage summaries without `instanceId` inherit from the live snapshot -- **Files modified:** `core/internal/delivery/cli/status.go` -- **Committed in:** `7c8873d8` (Task 3 feat commit) - ---- - -**Total deviations:** 2 auto-fixed (both Rule 1 - Bug) -**Impact on plan:** Both fixes were necessary for correctness. No scope creep — fixes remained within Task 1 and Task 3 boundaries. - -## Issues Encountered - -None beyond the two auto-fixed bugs documented above. - -## Known Stubs - -None - all instanceId fields are wired from real bootstrap identity. - -## Threat Flags - -No new network endpoints, auth paths, file access patterns, or schema changes at trust boundaries introduced beyond what the plan's threat model covers (T-32-01 through T-32-03). - -## Next Phase Readiness - -- `instanceId` is canonical in runtime snapshots and CLI storage records; Phase 33 (AWS account identity) can read it from snapshots without additional plumbing -- `status` alias is registered and covered by parity tests; no alias drift risk -- `port` remains in all human-facing and JSON surfaces as compatibility locator throughout the migration window -- Legacy port-keyed records load and fall back safely; no forced migration required before Phase 33 - ---- -*Phase: 32-cli-instance-identity-and-port-alias-migration* -*Completed: 2026-04-19* diff --git a/apps/desktop/package-lock.json b/apps/desktop/package-lock.json index bcab59d..c47fcbb 100644 --- a/apps/desktop/package-lock.json +++ b/apps/desktop/package-lock.json @@ -11,6 +11,8 @@ "dependencies": { "@aws-sdk/client-dynamodb": "^3.1032.0", "@aws-sdk/client-s3": "^3.1032.0", + "@aws-sdk/client-sqs": "^3.1033.0", + "@aws-sdk/s3-request-presigner": "^3.1033.0", "@base-ui/react": "^1.4.0", "@electron-toolkit/preload": "^3.0.2", "@electron-toolkit/utils": "^4.0.0", @@ -414,10 +416,62 @@ "node": ">=20.0.0" } }, + "node_modules/@aws-sdk/client-sqs": { + "version": "3.1033.0", + "resolved": "https://registry.npmjs.org/@aws-sdk/client-sqs/-/client-sqs-3.1033.0.tgz", + "integrity": "sha512-oB5SWYYzBh1GGNKF4dlVJOWFO5KfJmP+r96VeiiHSqAYmmR7Aj8qh10Kmt4pPgnETno0bgiMRp0Zo2+j7G08mw==", + "license": "Apache-2.0", + "dependencies": { + "@aws-crypto/sha256-browser": "5.2.0", + "@aws-crypto/sha256-js": "5.2.0", + "@aws-sdk/core": "^3.974.2", + "@aws-sdk/credential-provider-node": "^3.972.33", + "@aws-sdk/middleware-host-header": "^3.972.10", + "@aws-sdk/middleware-logger": "^3.972.10", + "@aws-sdk/middleware-recursion-detection": "^3.972.11", + "@aws-sdk/middleware-sdk-sqs": "^3.972.20", + "@aws-sdk/middleware-user-agent": "^3.972.32", + "@aws-sdk/region-config-resolver": "^3.972.12", + "@aws-sdk/types": "^3.973.8", + "@aws-sdk/util-endpoints": "^3.996.7", + "@aws-sdk/util-user-agent-browser": "^3.972.10", + "@aws-sdk/util-user-agent-node": "^3.973.18", + "@smithy/config-resolver": "^4.4.16", + "@smithy/core": "^3.23.15", + "@smithy/fetch-http-handler": "^5.3.17", + "@smithy/hash-node": "^4.2.14", + "@smithy/invalid-dependency": "^4.2.14", + "@smithy/md5-js": "^4.2.14", + "@smithy/middleware-content-length": "^4.2.14", + "@smithy/middleware-endpoint": "^4.4.30", + "@smithy/middleware-retry": "^4.5.3", + "@smithy/middleware-serde": "^4.2.18", + "@smithy/middleware-stack": "^4.2.14", + "@smithy/node-config-provider": "^4.3.14", + "@smithy/node-http-handler": "^4.5.3", + "@smithy/protocol-http": "^5.3.14", + "@smithy/smithy-client": "^4.12.11", + "@smithy/types": "^4.14.1", + "@smithy/url-parser": "^4.2.14", + "@smithy/util-base64": "^4.3.2", + "@smithy/util-body-length-browser": "^4.2.2", + "@smithy/util-body-length-node": "^4.2.3", + "@smithy/util-defaults-mode-browser": "^4.3.47", + "@smithy/util-defaults-mode-node": "^4.2.52", + "@smithy/util-endpoints": "^3.4.1", + "@smithy/util-middleware": "^4.2.14", + "@smithy/util-retry": "^4.3.2", + "@smithy/util-utf8": "^4.2.2", + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=20.0.0" + } + }, "node_modules/@aws-sdk/core": { - "version": "3.974.1", - "resolved": "https://registry.npmjs.org/@aws-sdk/core/-/core-3.974.1.tgz", - "integrity": "sha512-gy/gffKz0zaHDaqRiLCdIvgHmaAL/HXuAtMcBP7euYSFx4BsbsdlfmUBJag+Gqe62z6/XuloKyQyaiH+kS3Vrg==", + "version": "3.974.2", + "resolved": "https://registry.npmjs.org/@aws-sdk/core/-/core-3.974.2.tgz", + "integrity": "sha512-oav5AOAz+1XkwUfp6SrEm42UPDpUP5D4jNYXkDwFR1VfWqYX62+jpytdfzURmJ9McSoJIQwi0OJlC4oCi6t0VQ==", "license": "Apache-2.0", "dependencies": { "@aws-sdk/types": "^3.973.8", @@ -452,12 +506,12 @@ } }, "node_modules/@aws-sdk/credential-provider-env": { - "version": "3.972.27", - "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-env/-/credential-provider-env-3.972.27.tgz", - "integrity": "sha512-xfUt2CUZDC+Tf16A6roD1b4pk/nrXdkoLY3TEhv198AXDtBo5xUJP1zd0e8SmuKLN4PpIBX96OizZbmMlcI6oQ==", + "version": "3.972.28", + "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-env/-/credential-provider-env-3.972.28.tgz", + "integrity": "sha512-87GdRJ2OR0qR4VkMjXN/SZi66DZsunW2qQCbtw9rKw3Y7JurFi6tQWYKOSLY/gOADrU6OxGqFmdw3hKzZqDZOQ==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/core": "^3.974.1", + "@aws-sdk/core": "^3.974.2", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/types": "^4.14.1", @@ -468,12 +522,12 @@ } }, "node_modules/@aws-sdk/credential-provider-http": { - "version": "3.972.29", - "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-http/-/credential-provider-http-3.972.29.tgz", - "integrity": "sha512-hjNeYb6oLyHgMihra83ie0J/T2y9om3cy1qC90h9DRgvYXEoN4BCFf8bHguZjKhXunnv7YkmZRuYL5Mkk77eCA==", + "version": "3.972.30", + "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-http/-/credential-provider-http-3.972.30.tgz", + "integrity": "sha512-6quozmW2PKwBJTUQLb+lk1q8w5Pm45qaqhx4Tld9EIqYYQOVGj+MT0a8NRVS7QgWJj7rzGlB7rQu3KYBFHemJw==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/core": "^3.974.1", + "@aws-sdk/core": "^3.974.2", "@aws-sdk/types": "^3.973.8", "@smithy/fetch-http-handler": "^5.3.17", "@smithy/node-http-handler": "^4.5.3", @@ -489,19 +543,19 @@ } }, "node_modules/@aws-sdk/credential-provider-ini": { - "version": "3.972.31", - "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-ini/-/credential-provider-ini-3.972.31.tgz", - "integrity": "sha512-PuQ7e8WYzAPpzvFcajxf8c0LqSzakVHVlKw8M0oubk8Kf347YOCCqT1seQrHs5AdZuIh2RD9LX4O+Xa5ImEBfQ==", + "version": "3.972.32", + "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-ini/-/credential-provider-ini-3.972.32.tgz", + "integrity": "sha512-Nkr+UKtczZlocUjc6g96WzQadZSIZO/HVXPki4qbfaVOZYSbfLQKWKfADtJ0kGYsCvSYOZrO66tSc9dkboUt/w==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/core": "^3.974.1", - "@aws-sdk/credential-provider-env": "^3.972.27", - "@aws-sdk/credential-provider-http": "^3.972.29", - "@aws-sdk/credential-provider-login": "^3.972.31", - "@aws-sdk/credential-provider-process": "^3.972.27", - "@aws-sdk/credential-provider-sso": "^3.972.31", - "@aws-sdk/credential-provider-web-identity": "^3.972.31", - "@aws-sdk/nested-clients": "^3.996.21", + "@aws-sdk/core": "^3.974.2", + "@aws-sdk/credential-provider-env": "^3.972.28", + "@aws-sdk/credential-provider-http": "^3.972.30", + "@aws-sdk/credential-provider-login": "^3.972.32", + "@aws-sdk/credential-provider-process": "^3.972.28", + "@aws-sdk/credential-provider-sso": "^3.972.32", + "@aws-sdk/credential-provider-web-identity": "^3.972.32", + "@aws-sdk/nested-clients": "^3.997.0", "@aws-sdk/types": "^3.973.8", "@smithy/credential-provider-imds": "^4.2.14", "@smithy/property-provider": "^4.2.14", @@ -514,13 +568,13 @@ } }, "node_modules/@aws-sdk/credential-provider-login": { - "version": "3.972.31", - "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-login/-/credential-provider-login-3.972.31.tgz", - "integrity": "sha512-bBmWDmtSpmLOZR6a0kmowBcVL1hiL8Vlap/RXeMpFd7JbWl87YcwqL6T9LH/0oBVEZXu1dUZAtojgSuZgMO5xw==", + "version": "3.972.32", + "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-login/-/credential-provider-login-3.972.32.tgz", + "integrity": "sha512-UxgwT1HmZz1QPXuBy5ZUPJNFXOSlhwdQL61eGhWRthF0xRrT02BCOVJ1p5Ejg5AXfnESTWoKPJ7v/sCkNUtB9g==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/core": "^3.974.1", - "@aws-sdk/nested-clients": "^3.996.21", + "@aws-sdk/core": "^3.974.2", + "@aws-sdk/nested-clients": "^3.997.0", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/protocol-http": "^5.3.14", @@ -533,17 +587,17 @@ } }, "node_modules/@aws-sdk/credential-provider-node": { - "version": "3.972.32", - "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-node/-/credential-provider-node-3.972.32.tgz", - "integrity": "sha512-9aj0x9hGYUondBZSD0XkksAdHhOKttFw4BWpLCeggeg40qSJxGrAP++g0GCm0VqWc1WtC/NRFiAVzPCy56vmog==", + "version": "3.972.33", + "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-node/-/credential-provider-node-3.972.33.tgz", + "integrity": "sha512-6pGQnEdSeRvBViTQh/FwaRKB38a3Th+W2mVxuvqAd2Z1Ayo3e6eJ5QqJoZwEMwR6xoxkl3wz3qAfiB1xRhMC+w==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/credential-provider-env": "^3.972.27", - "@aws-sdk/credential-provider-http": "^3.972.29", - "@aws-sdk/credential-provider-ini": "^3.972.31", - "@aws-sdk/credential-provider-process": "^3.972.27", - "@aws-sdk/credential-provider-sso": "^3.972.31", - "@aws-sdk/credential-provider-web-identity": "^3.972.31", + "@aws-sdk/credential-provider-env": "^3.972.28", + "@aws-sdk/credential-provider-http": "^3.972.30", + "@aws-sdk/credential-provider-ini": "^3.972.32", + "@aws-sdk/credential-provider-process": "^3.972.28", + "@aws-sdk/credential-provider-sso": "^3.972.32", + "@aws-sdk/credential-provider-web-identity": "^3.972.32", "@aws-sdk/types": "^3.973.8", "@smithy/credential-provider-imds": "^4.2.14", "@smithy/property-provider": "^4.2.14", @@ -556,12 +610,12 @@ } }, "node_modules/@aws-sdk/credential-provider-process": { - "version": "3.972.27", - "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-process/-/credential-provider-process-3.972.27.tgz", - "integrity": "sha512-1CZvfb1WzudWWIFAVQkd1OI/T1RxPcSvNWzNsb2BMBVsBJzBtB8dV5f2nymHVU4UqwxipdVt/DAbgdDRf33JDg==", + "version": "3.972.28", + "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-process/-/credential-provider-process-3.972.28.tgz", + "integrity": "sha512-CRAlD8u6oNBhjnX/3ekVGocarD+lFmEn/qeDzytgIdmwrmwMJGFPqS9lGwEfhOTihZKrQ0xSp3z6paX+iXJJhA==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/core": "^3.974.1", + "@aws-sdk/core": "^3.974.2", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", @@ -573,14 +627,14 @@ } }, "node_modules/@aws-sdk/credential-provider-sso": { - "version": "3.972.31", - "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-sso/-/credential-provider-sso-3.972.31.tgz", - "integrity": "sha512-x8Mx18S48XMl9bEEpYwmXDTvjWGPIfDadReN37Lc099/DUrlL4Zs9T9rwwggo6DkKS1aev6v+MTUx7JTa87TZQ==", + "version": "3.972.32", + "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-sso/-/credential-provider-sso-3.972.32.tgz", + "integrity": "sha512-whhmQghRYOt9mJxFyVMhX7eB8n0oA25OCvqoR7dzFAZjmioCkf7WVB22Bc6llM5cFpBXFX7s4Jv+xVq32VPGWg==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/core": "^3.974.1", - "@aws-sdk/nested-clients": "^3.996.21", - "@aws-sdk/token-providers": "3.1032.0", + "@aws-sdk/core": "^3.974.2", + "@aws-sdk/nested-clients": "^3.997.0", + "@aws-sdk/token-providers": "3.1033.0", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", @@ -592,13 +646,13 @@ } }, "node_modules/@aws-sdk/credential-provider-web-identity": { - "version": "3.972.31", - "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-web-identity/-/credential-provider-web-identity-3.972.31.tgz", - "integrity": "sha512-zfuNMIkGfjYsHis9qytYf74Bcmq6Ji9Xwf4w53baRCI/b2otTwZv3SW1uRiJ5Di7999QzRGhHZ96+eUeo3gSOA==", + "version": "3.972.32", + "resolved": "https://registry.npmjs.org/@aws-sdk/credential-provider-web-identity/-/credential-provider-web-identity-3.972.32.tgz", + "integrity": "sha512-Z0Y0LDaqyQDznlmr9gv6n4+eWKKWNgmi9j5L6RENr6wyOCguhO8FRPmqDbVLSw0DPdMqICKnA3PurJiS8bD6Cw==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/core": "^3.974.1", - "@aws-sdk/nested-clients": "^3.996.21", + "@aws-sdk/core": "^3.974.2", + "@aws-sdk/nested-clients": "^3.997.0", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", @@ -773,12 +827,12 @@ } }, "node_modules/@aws-sdk/middleware-sdk-s3": { - "version": "3.972.30", - "resolved": "https://registry.npmjs.org/@aws-sdk/middleware-sdk-s3/-/middleware-sdk-s3-3.972.30.tgz", - "integrity": "sha512-hoQRxjJu4tt3gEOQin21rJKotClJC+x7AmCh9ylRct1DJeaNI/BRlFxMbuhJe54bG6xANPagSs0my8K30QyV9g==", + "version": "3.972.31", + "resolved": "https://registry.npmjs.org/@aws-sdk/middleware-sdk-s3/-/middleware-sdk-s3-3.972.31.tgz", + "integrity": "sha512-5hS08Fp0Rm+59uGCmkWhZmveXiA7OUV7Wa+IARejdzf9JTZ1qAVeIOE9JoBpsLPvUgEjmsGNHBuFbtGmYyqiqQ==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/core": "^3.974.1", + "@aws-sdk/core": "^3.974.2", "@aws-sdk/types": "^3.973.8", "@aws-sdk/util-arn-parser": "^3.972.3", "@smithy/core": "^3.23.15", @@ -797,6 +851,23 @@ "node": ">=20.0.0" } }, + "node_modules/@aws-sdk/middleware-sdk-sqs": { + "version": "3.972.20", + "resolved": "https://registry.npmjs.org/@aws-sdk/middleware-sdk-sqs/-/middleware-sdk-sqs-3.972.20.tgz", + "integrity": "sha512-yt0w5FKyH8Or7OT/Bp3fDRAtI4/f6uaaRKnW9TmU9qv8c1HFh43C9nQYZ26IcyRm+tYFdrB65yNTav/YThu36A==", + "license": "Apache-2.0", + "dependencies": { + "@aws-sdk/types": "^3.973.8", + "@smithy/smithy-client": "^4.12.11", + "@smithy/types": "^4.14.1", + "@smithy/util-hex-encoding": "^4.2.2", + "@smithy/util-utf8": "^4.2.2", + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=20.0.0" + } + }, "node_modules/@aws-sdk/middleware-ssec": { "version": "3.972.10", "resolved": "https://registry.npmjs.org/@aws-sdk/middleware-ssec/-/middleware-ssec-3.972.10.tgz", @@ -812,12 +883,12 @@ } }, "node_modules/@aws-sdk/middleware-user-agent": { - "version": "3.972.31", - "resolved": "https://registry.npmjs.org/@aws-sdk/middleware-user-agent/-/middleware-user-agent-3.972.31.tgz", - "integrity": "sha512-L+hXN2HDomlIsWSHW5DVD7ppccCeRnlHXZ5uHG34ePTjF5bm0I1fmrJLbUGiW97xRXWryit5cjdP4Sx2FwiGog==", + "version": "3.972.32", + "resolved": "https://registry.npmjs.org/@aws-sdk/middleware-user-agent/-/middleware-user-agent-3.972.32.tgz", + "integrity": "sha512-HQ0x9DDKqLZOGhDiL2eicYXXkYT5dogE4mw0lAfHCpJ6t7MM0PNIsJl2TZzWKU9SpBzOMXHRa7K6ZLKUJu1y0w==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/core": "^3.974.1", + "@aws-sdk/core": "^3.974.2", "@aws-sdk/types": "^3.973.8", "@aws-sdk/util-endpoints": "^3.996.7", "@smithy/core": "^3.23.15", @@ -831,23 +902,24 @@ } }, "node_modules/@aws-sdk/nested-clients": { - "version": "3.996.21", - "resolved": "https://registry.npmjs.org/@aws-sdk/nested-clients/-/nested-clients-3.996.21.tgz", - "integrity": "sha512-Me3d/ua2lb2G0bQfFmvCeQQp3+nN6GSPqMxDmi/IQlQ8CrlpQ5C0JJHpz2AnOUkEFI0lBNrAL3Vnt29l44ndkA==", + "version": "3.997.0", + "resolved": "https://registry.npmjs.org/@aws-sdk/nested-clients/-/nested-clients-3.997.0.tgz", + "integrity": "sha512-4bI5GHjUiY5R8N6PtchpG6tW2Dl8I2IcZNg3JwqwxHRXjfvQlPoo4VMknG4qkd5W0t3Y20rQ6C7pSR561YG5JQ==", "license": "Apache-2.0", "dependencies": { "@aws-crypto/sha256-browser": "5.2.0", "@aws-crypto/sha256-js": "5.2.0", - "@aws-sdk/core": "^3.974.1", + "@aws-sdk/core": "^3.974.2", "@aws-sdk/middleware-host-header": "^3.972.10", "@aws-sdk/middleware-logger": "^3.972.10", "@aws-sdk/middleware-recursion-detection": "^3.972.11", - "@aws-sdk/middleware-user-agent": "^3.972.31", + "@aws-sdk/middleware-user-agent": "^3.972.32", "@aws-sdk/region-config-resolver": "^3.972.12", + "@aws-sdk/signature-v4-multi-region": "^3.996.19", "@aws-sdk/types": "^3.973.8", "@aws-sdk/util-endpoints": "^3.996.7", "@aws-sdk/util-user-agent-browser": "^3.972.10", - "@aws-sdk/util-user-agent-node": "^3.973.17", + "@aws-sdk/util-user-agent-node": "^3.973.18", "@smithy/config-resolver": "^4.4.16", "@smithy/core": "^3.23.15", "@smithy/fetch-http-handler": "^5.3.17", @@ -895,13 +967,32 @@ "node": ">=20.0.0" } }, + "node_modules/@aws-sdk/s3-request-presigner": { + "version": "3.1033.0", + "resolved": "https://registry.npmjs.org/@aws-sdk/s3-request-presigner/-/s3-request-presigner-3.1033.0.tgz", + "integrity": "sha512-8PVtuRzL9k59TgceC2KXA4SbG5V+nKzwO0YVtpp3ylf3Ios1DB7+psoMFS/P6zmsCHMT+S3ChrIUDXv4QmRawQ==", + "license": "Apache-2.0", + "dependencies": { + "@aws-sdk/signature-v4-multi-region": "^3.996.19", + "@aws-sdk/types": "^3.973.8", + "@aws-sdk/util-format-url": "^3.972.10", + "@smithy/middleware-endpoint": "^4.4.30", + "@smithy/protocol-http": "^5.3.14", + "@smithy/smithy-client": "^4.12.11", + "@smithy/types": "^4.14.1", + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=20.0.0" + } + }, "node_modules/@aws-sdk/signature-v4-multi-region": { - "version": "3.996.18", - "resolved": "https://registry.npmjs.org/@aws-sdk/signature-v4-multi-region/-/signature-v4-multi-region-3.996.18.tgz", - "integrity": "sha512-4KT8UXRmvNAP5zKq9UI1MIwbnmSChZncBt89RKu/skMqZSSWGkBZTAJsZ+no+txfmF3kVaUFv31CTBZkQ5BJpQ==", + "version": "3.996.19", + "resolved": "https://registry.npmjs.org/@aws-sdk/signature-v4-multi-region/-/signature-v4-multi-region-3.996.19.tgz", + "integrity": "sha512-7Sy8+GhfwUi06NQNLplxuJuXMKJURDsNQfK8yTW6E9wN2J1B+8S5dWZG7vg3InvPPhaXqkcYTr8pzeE+dLjMbQ==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/middleware-sdk-s3": "^3.972.30", + "@aws-sdk/middleware-sdk-s3": "^3.972.31", "@aws-sdk/types": "^3.973.8", "@smithy/protocol-http": "^5.3.14", "@smithy/signature-v4": "^5.3.14", @@ -913,13 +1004,13 @@ } }, "node_modules/@aws-sdk/token-providers": { - "version": "3.1032.0", - "resolved": "https://registry.npmjs.org/@aws-sdk/token-providers/-/token-providers-3.1032.0.tgz", - "integrity": "sha512-n+PU8Z+gll7p3wDrH+Wo6fkt8sPrVnq30YYM6Ryga95oJlEneNMEbDHj0iqjMX3V7gaGdJo/hJWyPo4lscP+mA==", + "version": "3.1033.0", + "resolved": "https://registry.npmjs.org/@aws-sdk/token-providers/-/token-providers-3.1033.0.tgz", + "integrity": "sha512-/TsXhqjyRAFb0xVgmbFAha3cJfZdWjnyn6ohJ3AB4E3peLgxNcmKfYr45hruHymyJAydiHoXC3N1a8qgl41cog==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/core": "^3.974.1", - "@aws-sdk/nested-clients": "^3.996.21", + "@aws-sdk/core": "^3.974.2", + "@aws-sdk/nested-clients": "^3.997.0", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", @@ -971,6 +1062,21 @@ "node": ">=20.0.0" } }, + "node_modules/@aws-sdk/util-format-url": { + "version": "3.972.10", + "resolved": "https://registry.npmjs.org/@aws-sdk/util-format-url/-/util-format-url-3.972.10.tgz", + "integrity": "sha512-DEKiHNJVtNxdyTeQspzY+15Po/kHm6sF0Cs4HV9Q2+lplB63+DrvdeiSoOSdWEWAoO2RcY1veoXVDz2tWxWCgQ==", + "license": "Apache-2.0", + "dependencies": { + "@aws-sdk/types": "^3.973.8", + "@smithy/querystring-builder": "^4.2.14", + "@smithy/types": "^4.14.1", + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=20.0.0" + } + }, "node_modules/@aws-sdk/util-locate-window": { "version": "3.965.5", "resolved": "https://registry.npmjs.org/@aws-sdk/util-locate-window/-/util-locate-window-3.965.5.tgz", @@ -996,12 +1102,12 @@ } }, "node_modules/@aws-sdk/util-user-agent-node": { - "version": "3.973.17", - "resolved": "https://registry.npmjs.org/@aws-sdk/util-user-agent-node/-/util-user-agent-node-3.973.17.tgz", - "integrity": "sha512-utF5qjjbuJQuU9VdCkWl7L87sr93cApsrD+uxGfUnlafX8iyEzJrb7EZnufjThURZVTOtelRMXrblWxpefElUg==", + "version": "3.973.18", + "resolved": "https://registry.npmjs.org/@aws-sdk/util-user-agent-node/-/util-user-agent-node-3.973.18.tgz", + "integrity": "sha512-Nh4YvAL0Mzv5jBvzXLFL0tLf7WPrRMnYZQ5jlFuyS0xiVJQsObMUKAkbYjmt/e04wpQqUaa+Is7k+mBr89A9yA==", "license": "Apache-2.0", "dependencies": { - "@aws-sdk/middleware-user-agent": "^3.972.31", + "@aws-sdk/middleware-user-agent": "^3.972.32", "@aws-sdk/types": "^3.973.8", "@smithy/node-config-provider": "^4.3.14", "@smithy/types": "^4.14.1", diff --git a/apps/desktop/package.json b/apps/desktop/package.json index b178446..c9c1fba 100644 --- a/apps/desktop/package.json +++ b/apps/desktop/package.json @@ -9,6 +9,7 @@ "format": "prettier --write .", "lint": "eslint --cache .", "s3:smoke": "node scripts/s3-smoke.cjs", + "sqs:smoke": "node scripts/sqs-smoke.cjs", "dynamo:smoke": "node scripts/dynamo-smoke.cjs", "typecheck:node": "tsc --noEmit -p tsconfig.node.json --composite false", "typecheck:web": "tsc --noEmit -p tsconfig.web.json --composite false", @@ -26,6 +27,8 @@ "dependencies": { "@aws-sdk/client-dynamodb": "^3.1032.0", "@aws-sdk/client-s3": "^3.1032.0", + "@aws-sdk/client-sqs": "^3.1033.0", + "@aws-sdk/s3-request-presigner": "^3.1033.0", "@base-ui/react": "^1.4.0", "@electron-toolkit/preload": "^3.0.2", "@electron-toolkit/utils": "^4.0.0", diff --git a/apps/desktop/scripts/dynamo-smoke.cjs b/apps/desktop/scripts/dynamo-smoke.cjs index 261d35b..2f9b2ac 100644 --- a/apps/desktop/scripts/dynamo-smoke.cjs +++ b/apps/desktop/scripts/dynamo-smoke.cjs @@ -18,15 +18,21 @@ const { BatchGetItemCommand, TransactWriteItemsCommand, TransactGetItemsCommand, + UpdateTimeToLiveCommand, + DescribeTimeToLiveCommand, } = require('@aws-sdk/client-dynamodb'); // Parse arguments to find port const args = process.argv.slice(2); let port = 4566; +let debug = false; for (let i = 0; i < args.length; i++) { if (args[i] === '--port' && args[i + 1]) { port = parseInt(args[i + 1], 10); } + if (args[i] === '--debug') { + debug = true; + } } main().catch((error) => { @@ -48,275 +54,328 @@ async function main() { }, }); - const smokeTable = uniqueTableName('smoke'); - const batchTable = uniqueTableName('batch'); - - await execute(client, 'ListTables', new ListTablesCommand({})); - - const createOut = await execute(client, 'CreateTable', new CreateTableCommand({ - TableName: smokeTable, - KeySchema: [ - { AttributeName: 'id', KeyType: 'HASH' }, - { AttributeName: 'sk', KeyType: 'RANGE' }, - ], - AttributeDefinitions: [ - { AttributeName: 'id', AttributeType: 'S' }, - { AttributeName: 'sk', AttributeType: 'S' }, - ], - BillingMode: 'PAY_PER_REQUEST', - })); - assertTableName(createOut, smokeTable); - await waitForTableStatus(client, smokeTable, 'ACTIVE'); - - await execute(client, 'PutItem', new PutItemCommand({ - TableName: smokeTable, - Item: { - id: { S: 'series#1' }, - sk: { S: '001' }, - title: { S: 'skip-one' }, - }, - })); - await execute(client, 'PutItem', new PutItemCommand({ - TableName: smokeTable, - Item: { - id: { S: 'series#1' }, - sk: { S: '002' }, - title: { S: 'keep-two' }, - }, - })); - await execute(client, 'PutItem', new PutItemCommand({ - TableName: smokeTable, - Item: { - id: { S: 'series#1' }, - sk: { S: '003' }, - title: { S: 'keep-three' }, - }, - })); - - const updateOut = await execute(client, 'UpdateItem', new UpdateItemCommand({ - TableName: smokeTable, - Key: { id: { S: 'series#1' }, sk: { S: '002' } }, - UpdateExpression: 'SET title = :title ADD version :inc REMOVE archived', - ExpressionAttributeValues: { - ':title': { S: 'keep-two-updated' }, - ':inc': { N: '1' }, - }, - ReturnValues: 'ALL_NEW', - })); - expectEqual(attrValueString(updateOut.Attributes.title), 'keep-two-updated', 'updated title'); - expectEqual(attrValueString(updateOut.Attributes.version), '1', 'updated version'); - - const getOut = await execute(client, 'GetItem', new GetItemCommand({ - TableName: smokeTable, - Key: { id: { S: 'series#1' }, sk: { S: '002' } }, - })); - expectEqual(attrValueString(getOut.Item.title), 'keep-two-updated', 'get item title'); - - const queryOut = await execute(client, 'Query', new QueryCommand({ - TableName: smokeTable, - KeyConditionExpression: 'id = :id AND sk BETWEEN :start AND :end', - ExpressionAttributeValues: { - ':id': { S: 'series#1' }, - ':start': { S: '001' }, - ':end': { S: '003' }, - }, - ScanIndexForward: false, - Limit: 2, - })); - expectEqual(queryOut.Items.length, 2, 'query item count'); - expectEqual(attrValueString(queryOut.Items[0].sk), '003', 'query first sort key'); - expectEqual(attrValueString(queryOut.Items[1].sk), '002', 'query second sort key'); - expectEqual(attrValueString(queryOut.LastEvaluatedKey.sk), '002', 'query cursor'); - - const beginsOut = await execute(client, 'Query', new QueryCommand({ - TableName: smokeTable, - KeyConditionExpression: 'id = :id AND begins_with(sk, :prefix)', - ExpressionAttributeValues: { - ':id': { S: 'series#1' }, - ':prefix': { S: '00' }, - }, - })); - expectEqual(beginsOut.Items.length, 3, 'begins_with query count'); - - const scanOut = await execute(client, 'Scan', new ScanCommand({ - TableName: smokeTable, - FilterExpression: 'begins_with(title, :prefix)', - ExpressionAttributeValues: { - ':prefix': { S: 'keep' }, - }, - Limit: 1, - })); - expectEqual(scanOut.Items.length, 0, 'first scan page count'); - expectEqual(attrValueString(scanOut.LastEvaluatedKey.sk), '001', 'scan cursor'); - - const scanPage2 = await execute(client, 'Scan', new ScanCommand({ - TableName: smokeTable, - FilterExpression: 'begins_with(title, :prefix)', - ExpressionAttributeValues: { - ':prefix': { S: 'keep' }, - }, - Limit: 1, - ExclusiveStartKey: scanOut.LastEvaluatedKey, - })); - expectEqual(scanPage2.Items.length, 1, 'second scan page count'); - expectEqual(attrValueString(scanPage2.Items[0].title), 'keep-two-updated', 'scan page title'); - - await execute(client, 'DeleteItem', new DeleteItemCommand({ - TableName: smokeTable, - Key: { id: { S: 'series#1' }, sk: { S: '001' } }, - })); - - const batchCreateOut = await execute(client, 'CreateTable', new CreateTableCommand({ - TableName: batchTable, - KeySchema: [ - { AttributeName: 'id', KeyType: 'HASH' }, - ], - AttributeDefinitions: [ - { AttributeName: 'id', AttributeType: 'S' }, - ], - BillingMode: 'PAY_PER_REQUEST', - })); - assertTableName(batchCreateOut, batchTable); - await waitForTableStatus(client, batchTable, 'ACTIVE'); - - const writeRequests = []; - for (let i = 1; i <= 26; i += 1) { - const id = `item#${String(i).padStart(2, '0')}`; - writeRequests.push({ - PutRequest: { - Item: { - id: { S: id }, - title: { S: `title-${String(i).padStart(2, '0')}` }, - }, - }, - }); + const createdTables = []; + + async function createTable(name, options = {}) { + const tableName = uniqueTableName(name); + await execute(client, `CreateTable (${name})`, new CreateTableCommand({ + TableName: tableName, + KeySchema: [ + { AttributeName: 'pk', KeyType: 'HASH' }, + { AttributeName: 'sk', KeyType: 'RANGE' }, + ], + AttributeDefinitions: [ + { AttributeName: 'pk', AttributeType: 'S' }, + { AttributeName: 'sk', AttributeType: 'S' }, + ...(options.extraAttributes || []), + ], + BillingMode: 'PAY_PER_REQUEST', + ...(options.gsi && { GlobalSecondaryIndexes: options.gsi }), + ...(options.lsi && { LocalSecondaryIndexes: options.lsi }), + })); + createdTables.push(tableName); + await waitForTableStatus(client, tableName, 'ACTIVE'); + return tableName; } - const batchWriteOut = await execute(client, 'BatchWriteItem', new BatchWriteItemCommand({ - RequestItems: { - [batchTable]: writeRequests, - }, - })); - expectEqual(batchWriteOut.UnprocessedItems[batchTable].length, 1, 'batch write unprocessed count'); - expectEqual(attrValueString(batchWriteOut.UnprocessedItems[batchTable][0].PutRequest.Item.id), 'item#26', 'batch write unprocessed id'); - - const batchGetOut = await execute(client, 'BatchGetItem', new BatchGetItemCommand({ - RequestItems: { - [batchTable]: { - Keys: [ - { id: { S: 'item#01' } }, - { id: { S: 'item#25' } }, - { id: { S: 'item#26' } }, + try { + // --- Setup Tables --- + console.log('\n--- Setup Tables ---'); + const crudTable = await createTable('crud'); + const queryTable = await createTable('query'); + const batchTable = await createTable('batch'); + const txTable = await createTable('tx'); + const indexTable = await createTable('index', { + extraAttributes: [ + { AttributeName: 'gsi1pk', AttributeType: 'S' }, + { AttributeName: 'gsi1sk', AttributeType: 'N' }, + { AttributeName: 'lsi1sk', AttributeType: 'N' }, + ], + gsi: [{ + IndexName: 'gsi1', + KeySchema: [ + { AttributeName: 'gsi1pk', KeyType: 'HASH' }, + { AttributeName: 'gsi1sk', KeyType: 'RANGE' }, ], - }, - }, - })); - expectEqual(batchGetOut.Responses[batchTable].length, 2, 'batch get response count'); - expectEqual(attrValueString(batchGetOut.Responses[batchTable][0].id), 'item#01', 'batch get first id'); - expectEqual(attrValueString(batchGetOut.Responses[batchTable][1].id), 'item#25', 'batch get second id'); - - const transactWriteOut = await execute(client, 'TransactWriteItems', new TransactWriteItemsCommand({ - TransactItems: [ - { - Put: { - TableName: batchTable, - Item: { - id: { S: 'item#27' }, - title: { S: 'title-27' }, - }, - }, - }, - { - Delete: { - TableName: batchTable, - Key: { id: { S: 'item#01' } }, - }, - }, - ], - })); - if (!transactWriteOut) { - throw new Error('expected transact write response'); - } + Projection: { ProjectionType: 'ALL' }, + }], + lsi: [{ + IndexName: 'lsi1', + KeySchema: [ + { AttributeName: 'pk', KeyType: 'HASH' }, + { AttributeName: 'lsi1sk', KeyType: 'RANGE' }, + ], + Projection: { ProjectionType: 'ALL' }, + }] + }); - const transactGetOut = await execute(client, 'TransactGetItems', new TransactGetItemsCommand({ - TransactItems: [ - { - Get: { - TableName: batchTable, - Key: { id: { S: 'item#27' } }, - }, - }, - { - Get: { - TableName: batchTable, - Key: { id: { S: 'item#02' } }, - }, - }, - ], - })); - expectEqual(transactGetOut.Responses.length, 2, 'transact get response count'); - expectEqual(attrValueString(transactGetOut.Responses[0].Item.id), 'item#27', 'transact get first id'); - expectEqual(attrValueString(transactGetOut.Responses[1].Item.id), 'item#02', 'transact get second id'); - - await expectAwsError( - client, - 'TransactWriteItems', - new TransactWriteItemsCommand({ - TransactItems: [ - { - Put: { - TableName: batchTable, - Item: { - id: { S: 'item#28' }, - title: { S: 'title-28' }, - }, - }, - }, - { - Delete: { - TableName: batchTable, - Key: { id: { S: 'item#28' } }, - }, - }, - ], - }), - 'TransactionCanceledException', - ); + // --- Validation: CRUD and Expressions --- + console.log('\n--- Validation: CRUD and Expressions ---'); + await execute(client, 'PutItem', new PutItemCommand({ + TableName: crudTable, + Item: { pk: { S: 'crud#1' }, sk: { S: 'meta' }, val: { N: '10' } } + })); + + const get1 = await execute(client, 'GetItem', new GetItemCommand({ + TableName: crudTable, + Key: { pk: { S: 'crud#1' }, sk: { S: 'meta' } }, + ConsistentRead: true + })); + expectEqual(attrValueString(get1.Item.val), '10', 'PutItem val'); + + await expectAwsError(client, 'PutItem (Conditional Check Failed)', new PutItemCommand({ + TableName: crudTable, + Item: { pk: { S: 'crud#1' }, sk: { S: 'meta' }, val: { N: '20' } }, + ConditionExpression: 'attribute_not_exists(pk)' + }), 'ConditionalCheckFailedException'); + + const up1 = await execute(client, 'UpdateItem', new UpdateItemCommand({ + TableName: crudTable, + Key: { pk: { S: 'crud#1' }, sk: { S: 'meta' } }, + UpdateExpression: 'SET val = val + :inc, #nm = :name', + ExpressionAttributeNames: { '#nm': 'name' }, + ExpressionAttributeValues: { ':inc': { N: '5' }, ':name': { S: 'tester' } }, + ReturnValues: 'ALL_NEW' + })); + expectEqual(attrValueString(up1.Attributes.val), '15', 'UpdateItem math (+5)'); + expectEqual(attrValueString(up1.Attributes.name), 'tester', 'UpdateItem string'); + + await execute(client, 'DeleteItem', new DeleteItemCommand({ + TableName: crudTable, + Key: { pk: { S: 'crud#1' }, sk: { S: 'meta' } } + })); + + const get2 = await execute(client, 'GetItem (After Delete)', new GetItemCommand({ + TableName: crudTable, + Key: { pk: { S: 'crud#1' }, sk: { S: 'meta' } } + })); + expectEqual(get2.Item, undefined, 'Item should be deleted'); + + + // --- Validation: Query, Scan and Pagination --- + console.log('\n--- Validation: Query, Scan and Pagination ---'); + for (let i = 1; i <= 10; i++) { + await execute(client, `PutItem (Q/S ${i})`, new PutItemCommand({ + TableName: queryTable, + Item: { + pk: { S: 'grp#1' }, + sk: { S: `item#${String(i).padStart(2, '0')}` }, + active: { S: i % 2 === 0 ? 'true' : 'false' } + } + })); + } - await execute(client, 'DeleteTable', new DeleteTableCommand({ TableName: smokeTable })); - await execute(client, 'DeleteTable', new DeleteTableCommand({ TableName: batchTable })); + const q1 = await execute(client, 'Query (Limit)', new QueryCommand({ + TableName: queryTable, + KeyConditionExpression: 'pk = :pk', + ExpressionAttributeValues: { ':pk': { S: 'grp#1' } }, + Limit: 3 + })); + expectEqual(q1.Items?.length, 3, 'Query limit 3'); + expectEqual(attrValueString(q1.LastEvaluatedKey.sk), 'item#03', 'Query LEK'); + + const q2 = await execute(client, 'Query (ExclusiveStartKey)', new QueryCommand({ + TableName: queryTable, + KeyConditionExpression: 'pk = :pk', + ExpressionAttributeValues: { ':pk': { S: 'grp#1' } }, + ExclusiveStartKey: q1.LastEvaluatedKey + })); + expectEqual(q2.Items?.length, 7, 'Query pagination remaining'); + + const q3 = await execute(client, 'Query (FilterExpression)', new QueryCommand({ + TableName: queryTable, + KeyConditionExpression: 'pk = :pk', + FilterExpression: 'active = :active', + ExpressionAttributeValues: { ':pk': { S: 'grp#1' }, ':active': { S: 'true' } } + })); + expectEqual(q3.Items?.length, 5, 'Query filter active=true'); + + const s1 = await execute(client, 'Scan (Limit)', new ScanCommand({ + TableName: queryTable, + Limit: 4 + })); + expectEqual(s1.Items?.length, 4, 'Scan limit 4'); + + const s2 = await execute(client, 'Scan (ExclusiveStartKey)', new ScanCommand({ + TableName: queryTable, + ExclusiveStartKey: s1.LastEvaluatedKey + })); + expectEqual(s2.Items?.length, 6, 'Scan pagination remaining'); + + + // --- Validation: GSI and LSI --- + console.log('\n--- Validation: GSI and LSI ---'); + for (let i = 1; i <= 5; i++) { + await execute(client, `PutItem (Idx ${i})`, new PutItemCommand({ + TableName: indexTable, + Item: { + pk: { S: 'idx#1' }, + sk: { S: `item#${i}` }, + gsi1pk: { S: 'type#A' }, + gsi1sk: { N: `${i * 10}` }, + lsi1sk: { N: `${i * 100}` } + } + })); + } + + const gsiQ = await execute(client, 'Query (GSI)', new QueryCommand({ + TableName: indexTable, + IndexName: 'gsi1', + KeyConditionExpression: 'gsi1pk = :gsi1pk AND gsi1sk >= :gsi1sk', + ExpressionAttributeValues: { ':gsi1pk': { S: 'type#A' }, ':gsi1sk': { N: '30' } } + })); + expectEqual(gsiQ.Items?.length, 3, 'GSI query count'); // 30, 40, 50 + + const lsiQ = await execute(client, 'Query (LSI)', new QueryCommand({ + TableName: indexTable, + IndexName: 'lsi1', + KeyConditionExpression: 'pk = :pk AND lsi1sk <= :lsi1sk', + ExpressionAttributeValues: { ':pk': { S: 'idx#1' }, ':lsi1sk': { N: '200' } } + })); + expectEqual(lsiQ.Items?.length, 2, 'LSI query count'); // 100, 200 + + + // --- Validation: Batch Operations --- + console.log('\n--- Validation: Batch Operations ---'); + const writeReqs = []; + for (let i = 1; i <= 25; i++) { + writeReqs.push({ PutRequest: { Item: { pk: { S: `batch#${i}` }, sk: { S: 'meta' } } } }); + } + + await execute(client, 'BatchWriteItem (25 items)', new BatchWriteItemCommand({ + RequestItems: { [batchTable]: writeReqs } + })); + + const getReqs = writeReqs.map(r => ({ pk: r.PutRequest.Item.pk, sk: r.PutRequest.Item.sk })); + const bg = await execute(client, 'BatchGetItem (25 items)', new BatchGetItemCommand({ + RequestItems: { [batchTable]: { Keys: getReqs } } + })); + expectEqual(bg.Responses?.[batchTable]?.length, 25, 'BatchGetItem count'); + + const delReqs = getReqs.map(k => ({ DeleteRequest: { Key: k } })); + await execute(client, 'BatchWriteItem (Delete 25 items)', new BatchWriteItemCommand({ + RequestItems: { [batchTable]: delReqs } + })); - await expectAwsError( - client, - 'DescribeTable', - new DescribeTableCommand({ TableName: smokeTable }), - 'ResourceNotFoundException', - ); + const bg2 = await execute(client, 'BatchGetItem (After Delete)', new BatchGetItemCommand({ + RequestItems: { [batchTable]: { Keys: [getReqs[0]] } } + })); + expectEqual(bg2.Responses?.[batchTable]?.length || 0, 0, 'BatchGetItem after delete'); + + + // --- Validation: Transactions --- + console.log('\n--- Validation: Transactions ---'); + await execute(client, 'PutItem (Tx Initial)', new PutItemCommand({ + TableName: txTable, + Item: { pk: { S: 'tx#1' }, sk: { S: 'meta' }, val: { N: '10' } } + })); + + await execute(client, 'TransactWriteItems', new TransactWriteItemsCommand({ + TransactItems: [ + { Put: { TableName: txTable, Item: { pk: { S: 'tx#2' }, sk: { S: 'meta' } } } }, + { Update: { + TableName: txTable, + Key: { pk: { S: 'tx#1' }, sk: { S: 'meta' } }, + UpdateExpression: 'SET val = val + :inc', + ExpressionAttributeValues: { ':inc': { N: '5' } } + } }, + { ConditionCheck: { + TableName: txTable, + Key: { pk: { S: 'tx#1' }, sk: { S: 'meta' } }, + ConditionExpression: 'attribute_exists(pk)' + } } + ] + })); + + const tg = await execute(client, 'TransactGetItems', new TransactGetItemsCommand({ + TransactItems: [ + { Get: { TableName: txTable, Key: { pk: { S: 'tx#1' }, sk: { S: 'meta' } } } }, + { Get: { TableName: txTable, Key: { pk: { S: 'tx#2' }, sk: { S: 'meta' } } } } + ] + })); + expectEqual(tg.Responses?.length, 2, 'TransactGetItems count'); + expectEqual(attrValueString(tg.Responses[0].Item.val), '15', 'Transact updated val'); + expectEqual(attrValueString(tg.Responses[1].Item.pk), 'tx#2', 'Transact put pk'); + + await expectAwsError(client, 'TransactWriteItems (Failing)', new TransactWriteItemsCommand({ + TransactItems: [ + { Put: { TableName: txTable, Item: { pk: { S: 'tx#3' }, sk: { S: 'meta' } } } }, + { ConditionCheck: { + TableName: txTable, + Key: { pk: { S: 'tx#999' }, sk: { S: 'meta' } }, + ConditionExpression: 'attribute_exists(pk)' + } } + ] + }), 'TransactionCanceledException'); + + const get3 = await execute(client, 'GetItem (Rolled Back Tx)', new GetItemCommand({ + TableName: txTable, + Key: { pk: { S: 'tx#3' }, sk: { S: 'meta' } } + })); + expectEqual(get3.Item, undefined, 'Item should be rolled back'); + + + // --- Validation: TTL --- + console.log('\n--- Validation: TTL ---'); + await execute(client, 'UpdateTimeToLive', new UpdateTimeToLiveCommand({ + TableName: crudTable, + TimeToLiveSpecification: { AttributeName: 'expireAt', Enabled: true } + })); + + const ttlDesc = await execute(client, 'DescribeTimeToLive', new DescribeTimeToLiveCommand({ + TableName: crudTable + })); + const ttlStatus = ttlDesc.TimeToLiveDescription?.TimeToLiveStatus; + if (ttlStatus !== 'ENABLED' && ttlStatus !== 'ENABLING') { + throw new Error(`Validation Failed [TTL Status]: got ${ttlStatus} want ENABLED or ENABLING`); + } - console.log('\n✓ Native AWS SDK smoke mode passed'); + console.log('\n✓ Native AWS SDK smoke mode passed with deep behavioral validations'); + } finally { + // --- Cleanup --- + console.log('\n--- Cleanup ---'); + for (const tableName of createdTables) { + try { + await execute(client, `DeleteTable (${tableName})`, new DeleteTableCommand({ TableName: tableName })); + } catch (err) { + console.error(`\nFailed to delete table ${tableName}:`, err.message); + } + } + } } async function execute(client, name, command) { - console.log(`\nExecuting ${name}...`); + if (debug) console.log(`\nExecuting ${name}...`); try { const response = await client.send(command); - console.log(`✓ ${name} succeeded.`); - console.dir(response, { depth: 4, colors: true }); + if (debug) { + console.log(`✓ ${name} succeeded.`); + console.dir(response, { depth: 4, colors: true }); + } else { + process.stdout.write('.'); + } return response; } catch (error) { + if (!debug) console.log(''); // newline for error printAwsError(name, error); throw error; } } async function expectAwsError(client, name, command, expectedName) { + if (debug) console.log(`\nExecuting ${name} (expecting error)...`); try { await client.send(command); } catch (error) { - printAwsError(name, error); + if (debug) { + console.log(`✓ ${name} failed as expected with ${error.name}.`); + } else { + process.stdout.write('.'); + } expectEqual(error?.name, expectedName, `${name} error name`); return; } + if (!debug) console.log(''); throw new Error(`expected ${name} to fail with ${expectedName}`); } @@ -344,8 +403,8 @@ function printAwsError(name, error) { console.error('Error name:', error.name || 'unknown'); console.error('Error message:', error.message || String(error)); if (error.$response) { - console.error('Response status:', error.$response.statusCode); - console.error('Response headers:', error.$response.headers); + console.error('Response status:', error.$response?.statusCode); + console.error('Response headers:', error.$response?.headers); } } else { console.error(error); @@ -354,7 +413,7 @@ function printAwsError(name, error) { function expectEqual(actual, expected, label) { if (actual !== expected) { - throw new Error(`unexpected ${label}: got ${JSON.stringify(actual)} want ${JSON.stringify(expected)}`); + throw new Error(`Validation Failed [${label}]: got ${JSON.stringify(actual)} want ${JSON.stringify(expected)}`); } } @@ -362,21 +421,11 @@ function attrValueString(value) { if (!value) { throw new Error('expected attribute value to be present'); } - if (typeof value.S === 'string') { - return value.S; - } - if (typeof value.N === 'string') { - return value.N; - } + if (typeof value.S === 'string') return value.S; + if (typeof value.N === 'string') return value.N; throw new Error(`unexpected attribute value shape: ${JSON.stringify(value)}`); } -function assertTableName(output, tableName) { - if (output?.TableDescription?.TableName !== tableName) { - throw new Error(`unexpected table name in response: ${output?.TableDescription?.TableName || 'missing'} want ${tableName}`); - } -} - function sleep(ms) { return new Promise((resolve) => setTimeout(resolve, ms)); } diff --git a/apps/desktop/scripts/s3-smoke.cjs b/apps/desktop/scripts/s3-smoke.cjs index 2f3472f..a6be95e 100644 --- a/apps/desktop/scripts/s3-smoke.cjs +++ b/apps/desktop/scripts/s3-smoke.cjs @@ -11,7 +11,15 @@ const { PutObjectCommand, GetObjectCommand, DeleteObjectCommand, + ListObjectsV2Command, + CreateMultipartUploadCommand, + UploadPartCommand, + CompleteMultipartUploadCommand, + CopyObjectCommand, + HeadObjectCommand, + DeleteObjectsCommand, } = require('@aws-sdk/client-s3'); +const { getSignedUrl } = require('@aws-sdk/s3-request-presigner'); // Parse arguments to find port const args = process.argv.slice(2); @@ -22,16 +30,22 @@ for (let i = 0; i < args.length; i++) { } } -main().catch((error) => { - console.error('\nS3 smoke test failed'); - console.error(error instanceof Error ? error.stack || error.message : error); - process.exitCode = 1; -}); +function expectEqual(actual, expected, message) { + if (actual !== expected) { + throw new Error(`Assertion failed: ${message}\nExpected: ${expected}\nActual: ${actual}`); + } +} + +function expectDefined(actual, message) { + if (actual === undefined || actual === null) { + throw new Error(`Assertion failed: ${message}\nValue was undefined or null.`); + } +} async function main() { const endpoint = process.env.MILDSTACK_S3_ENDPOINT || process.env.AWS_S3_ENDPOINT || `http://localhost:${port}`; - - console.log(`Running AWS SDK smoke mode against ${endpoint}`); + + console.log(`Running S3 behavioral validation against ${endpoint}`); const client = new S3Client({ region: process.env.AWS_REGION || 'us-east-1', endpoint, @@ -42,64 +56,241 @@ async function main() { }, }); - const bucket = uniqueBucketName('native'); - const commands = [ - ['ListBuckets', new ListBucketsCommand({})], - ['CreateBucket', new CreateBucketCommand({ Bucket: bucket })], - ['HeadBucket', new HeadBucketCommand({ Bucket: bucket })], - ['PutObject', new PutObjectCommand({ + const bucket = uniqueBucketName('behavioral'); + const bucketsToCleanup = []; + + try { + console.log('\n--- 1. Bucket Operations ---'); + console.log(`Creating bucket ${bucket}...`); + await client.send(new CreateBucketCommand({ Bucket: bucket })); + bucketsToCleanup.push(bucket); + + console.log('Validating bucket creation via HeadBucket...'); + await client.send(new HeadBucketCommand({ Bucket: bucket })); + + console.log('Validating bucket appears in ListBuckets...'); + const listBucketsRes = await client.send(new ListBucketsCommand({})); + const bucketExists = listBucketsRes.Buckets?.some(b => b.Name === bucket); + expectEqual(bucketExists, true, 'Created bucket should appear in ListBuckets'); + + console.log('\n--- 2. Basic Object Operations ---'); + const objectKey = 'test-folder/hello.txt'; + const objectBody = 'Hello, MildStack S3!'; + const metadata = { 'custom-author': 'bot', 'status': 'draft' }; + + console.log(`Putting object ${objectKey}...`); + await client.send(new PutObjectCommand({ Bucket: bucket, - Key: 'native.txt', - Body: 'native-mode smoke payload', + Key: objectKey, + Body: objectBody, ContentType: 'text/plain', - })], - ['GetObject', new GetObjectCommand({ Bucket: bucket, Key: 'native.txt' })], - ['DeleteObject', new DeleteObjectCommand({ Bucket: bucket, Key: 'native.txt' })], - ['DeleteBucket', new DeleteBucketCommand({ Bucket: bucket })], - ]; - - for (const [name, command] of commands) { - console.log(`\nExecuting ${name}...`); + Metadata: metadata, + })); + + console.log(`Getting object ${objectKey}...`); + const getObjRes = await client.send(new GetObjectCommand({ + Bucket: bucket, + Key: objectKey, + })); + const bodyText = await getObjRes.Body.transformToString(); + + expectEqual(bodyText, objectBody, 'GetObject body should match PutObject body'); + expectEqual(getObjRes.ContentType, 'text/plain', 'GetObject ContentType should match'); + expectEqual(getObjRes.Metadata?.['custom-author'], 'bot', 'GetObject Metadata custom-author should match'); + expectEqual(getObjRes.Metadata?.status, 'draft', 'GetObject Metadata status should match'); + + console.log('Validating HeadObject...'); + const headObjRes = await client.send(new HeadObjectCommand({ + Bucket: bucket, + Key: objectKey, + })); + expectEqual(headObjRes.ContentType, 'text/plain', 'HeadObject ContentType should match'); + expectEqual(headObjRes.Metadata?.['custom-author'], 'bot', 'HeadObject Metadata should match'); + + console.log('\n--- 3. Pagination and Listing (ListObjectsV2) ---'); + console.log('Creating multiple objects for listing...'); + for (let i = 1; i <= 5; i++) { + await client.send(new PutObjectCommand({ + Bucket: bucket, + Key: `listing/item-${i}.txt`, + Body: `Content ${i}`, + })); + } + + let listRes = await client.send(new ListObjectsV2Command({ + Bucket: bucket, + Prefix: 'listing/', + MaxKeys: 3, + })); + + expectEqual(listRes.Contents?.length, 3, 'ListObjectsV2 should return exactly MaxKeys items'); + expectEqual(listRes.IsTruncated, true, 'ListObjectsV2 should be truncated with more items remaining'); + expectDefined(listRes.NextContinuationToken, 'NextContinuationToken should be defined'); + + const secondListRes = await client.send(new ListObjectsV2Command({ + Bucket: bucket, + Prefix: 'listing/', + MaxKeys: 3, + ContinuationToken: listRes.NextContinuationToken, + })); + expectEqual(secondListRes.Contents?.length, 2, 'Second ListObjectsV2 should return remaining 2 items'); + expectEqual(secondListRes.IsTruncated, false, 'Second ListObjectsV2 should not be truncated'); + + console.log('\n--- 4. Copy Object ---'); + const sourceKey = 'listing/item-1.txt'; + const destKey = 'copied/item-1-copy.txt'; + console.log(`Copying ${sourceKey} to ${destKey}...`); + + await client.send(new CopyObjectCommand({ + Bucket: bucket, + CopySource: `${bucket}/${sourceKey}`, + Key: destKey, + })); + + const copiedObj = await client.send(new GetObjectCommand({ + Bucket: bucket, + Key: destKey, + })); + const copiedBody = await copiedObj.Body.transformToString(); + expectEqual(copiedBody, 'Content 1', 'Copied object content should match original'); + + console.log('\n--- 5. Multipart Upload ---'); + const mpKey = 'multipart/large-file.bin'; + console.log(`Creating Multipart Upload for ${mpKey}...`); + const createMpRes = await client.send(new CreateMultipartUploadCommand({ + Bucket: bucket, + Key: mpKey, + Metadata: { 'mp-meta': 'yes' } + })); + const uploadId = createMpRes.UploadId; + expectDefined(uploadId, 'UploadId must be returned'); + + console.log(`Uploading parts...`); + const part1Body = 'Part 1 data. '.repeat(100); + const part2Body = 'Part 2 data. '.repeat(100); + + const p1Res = await client.send(new UploadPartCommand({ + Bucket: bucket, + Key: mpKey, + UploadId: uploadId, + PartNumber: 1, + Body: part1Body, + })); + + const p2Res = await client.send(new UploadPartCommand({ + Bucket: bucket, + Key: mpKey, + UploadId: uploadId, + PartNumber: 2, + Body: part2Body, + })); + + console.log(`Completing Multipart Upload...`); + await client.send(new CompleteMultipartUploadCommand({ + Bucket: bucket, + Key: mpKey, + UploadId: uploadId, + MultipartUpload: { + Parts: [ + { PartNumber: 1, ETag: p1Res.ETag }, + { PartNumber: 2, ETag: p2Res.ETag }, + ], + }, + })); + + const mpGetRes = await client.send(new GetObjectCommand({ + Bucket: bucket, + Key: mpKey, + })); + const mpText = await mpGetRes.Body.transformToString(); + expectEqual(mpText, part1Body + part2Body, 'Multipart concatenated body should match parts'); + expectEqual(mpGetRes.Metadata?.['mp-meta'], 'yes', 'Multipart Metadata should be preserved'); + + console.log('\n--- 6. Presigned URLs ---'); + const presignedKey = 'presigned/upload.txt'; + console.log('Generating presigned PUT URL...'); + const putCommand = new PutObjectCommand({ Bucket: bucket, Key: presignedKey }); + const putUrl = await getSignedUrl(client, putCommand, { expiresIn: 60 }); + expectDefined(putUrl, 'Presigned PUT URL should be generated'); + expectEqual(putUrl.includes(presignedKey), true, 'Presigned URL should contain the key'); + + console.log('Using presigned PUT URL via native fetch...'); + const presignedPutBody = 'Hello via presigned url!'; + const putResponse = await fetch(putUrl, { + method: 'PUT', + body: presignedPutBody, + }); + expectEqual(putResponse.ok, true, 'Fetch via presigned PUT should succeed'); + + console.log('Generating presigned GET URL...'); + const getCommand = new GetObjectCommand({ Bucket: bucket, Key: presignedKey }); + const getUrl = await getSignedUrl(client, getCommand, { expiresIn: 60 }); + + console.log('Using presigned GET URL via native fetch...'); + const getResponse = await fetch(getUrl); + expectEqual(getResponse.ok, true, 'Fetch via presigned GET should succeed'); + const presignedGetBody = await getResponse.text(); + expectEqual(presignedGetBody, presignedPutBody, 'Content fetched via presigned URL should match uploaded content'); + + console.log('\n--- 7. Deletion and Not Found Behaviors ---'); + console.log(`Deleting object ${objectKey}...`); + await client.send(new DeleteObjectCommand({ Bucket: bucket, Key: objectKey })); + try { - const response = await client.send(command); - console.log(`✓ ${name} succeeded. Response:`); - - // Attempt to read stream body if it exists for GetObject to show content - if (response.Body && typeof response.Body.transformToString === 'function') { - const bodyText = await response.Body.transformToString(); - const clone = { ...response, Body: bodyText }; - console.dir(clone, { depth: 4, colors: true }); - } else { - console.dir(response, { depth: 4, colors: true }); - } - } catch (error) { - console.error(`\nFailed during command: ${name}`); - if (error.$response) { - console.error('Response status:', error.$response.statusCode); - console.error('Response headers:', error.$response.headers); - if (error.$response.body) { - const body = error.$response.body; - if (typeof body.read === 'function') { - const chunk = body.read(); - if (chunk) { - console.error('Response body:', chunk.toString()); - } else { - console.error('Response body:', body); - } - } else if (typeof body.toString === 'function') { - console.error('Response body:', body.toString()); - } else { - console.error('Response body:', body); + await client.send(new GetObjectCommand({ Bucket: bucket, Key: objectKey })); + throw new Error('GetObject on deleted object should have thrown an error'); + } catch (err) { + expectEqual(err.name === 'NoSuchKey' || err.name === 'NotFound', true, 'Error name should be NoSuchKey or NotFound'); + } + + console.log('\nAll behavioral validations passed successfully! 🚀'); + + } finally { + console.log('\n--- 8. Cleanup ---'); + for (const b of bucketsToCleanup) { + console.log(`Cleaning up bucket: ${b}`); + try { + let hasMore = true; + let token = undefined; + let objectsToDelete = []; + + while (hasMore) { + const listParams = { Bucket: b, ContinuationToken: token }; + const listedObjects = await client.send(new ListObjectsV2Command(listParams)); + if (listedObjects.Contents && listedObjects.Contents.length > 0) { + objectsToDelete = objectsToDelete.concat(listedObjects.Contents.map(obj => ({ Key: obj.Key }))); + } + token = listedObjects.NextContinuationToken; + hasMore = !!token; + } + + if (objectsToDelete.length > 0) { + console.log(`Deleting ${objectsToDelete.length} objects from ${b}...`); + for (let i = 0; i < objectsToDelete.length; i += 1000) { + const chunk = objectsToDelete.slice(i, i + 1000); + await client.send(new DeleteObjectsCommand({ + Bucket: b, + Delete: { Objects: chunk, Quiet: true }, + })); } } + + console.log(`Deleting bucket ${b}...`); + await client.send(new DeleteBucketCommand({ Bucket: b })); + console.log(`✓ Cleanup complete for ${b}.`); + } catch (err) { + console.error(`Failed to cleanup bucket ${b}:`, err.message); } - throw error; } } - - console.log('\n✓ Native AWS SDK smoke mode passed'); } function uniqueBucketName(prefix) { return `mildstack-${prefix}-${Date.now().toString(36)}-${randomUUID().slice(0, 8)}`.toLowerCase(); } + +main().catch((error) => { + console.error('\n❌ S3 smoke test failed'); + console.error(error instanceof Error ? error.stack || error.message : error); + process.exitCode = 1; +}); diff --git a/apps/desktop/scripts/sqs-smoke.cjs b/apps/desktop/scripts/sqs-smoke.cjs new file mode 100644 index 0000000..deda7ef --- /dev/null +++ b/apps/desktop/scripts/sqs-smoke.cjs @@ -0,0 +1,302 @@ +#!/usr/bin/env node +'use strict'; + +const { randomUUID } = require('node:crypto'); +const { setTimeout } = require('node:timers/promises'); +const { + SQSClient, + ListQueuesCommand, + CreateQueueCommand, + GetQueueUrlCommand, + GetQueueAttributesCommand, + SetQueueAttributesCommand, + SendMessageCommand, + ReceiveMessageCommand, + ChangeMessageVisibilityCommand, + DeleteMessageCommand, + SendMessageBatchCommand, + ChangeMessageVisibilityBatchCommand, + DeleteMessageBatchCommand, + PurgeQueueCommand, + DeleteQueueCommand, + ListQueueTagsCommand, + TagQueueCommand, + UntagQueueCommand, + ListDeadLetterSourceQueuesCommand, +} = require('@aws-sdk/client-sqs'); + +// Parse arguments to find port +const args = process.argv.slice(2); +let port = 4566; +let debug = false; +for (let i = 0; i < args.length; i++) { + if (args[i] === '--port' && args[i + 1]) { + port = parseInt(args[i + 1], 10); + } + if (args[i] === '--debug') { + debug = true; + } +} + +main().catch((error) => { + console.error('\nSQS smoke test failed'); + console.error(error instanceof Error ? error.stack || error.message : error); + process.exitCode = 1; +}); + +async function main() { + const endpoint = process.env.MILDSTACK_SQS_ENDPOINT || process.env.AWS_SQS_ENDPOINT || `http://localhost:${port}`; + + console.log(`Running AWS SDK smoke mode against ${endpoint}`); + const client = new SQSClient({ + region: process.env.AWS_REGION || 'us-east-1', + endpoint, + credentials: { + accessKeyId: process.env.AWS_ACCESS_KEY_ID || 'test', + secretAccessKey: process.env.AWS_SECRET_ACCESS_KEY || 'test', + }, + }); + + const stdQueueName = uniqueQueueName('smoke-std'); + const dlqQueueName = uniqueQueueName('smoke-dlq'); + const fifoQueueName = uniqueQueueName('smoke-fifo') + '.fifo'; + + // --- Setup Queues --- + console.log('\n--- Setup Queues ---'); + const dlqUrl = (await execute(client, 'CreateQueue (DLQ)', new CreateQueueCommand({ + QueueName: dlqQueueName, + }))).QueueUrl; + + const dlqArn = (await execute(client, 'GetQueueAttributes (DLQ)', new GetQueueAttributesCommand({ + QueueUrl: dlqUrl, AttributeNames: ['QueueArn'], + }))).Attributes?.QueueArn; + + const stdUrl = (await execute(client, 'CreateQueue (Std)', new CreateQueueCommand({ + QueueName: stdQueueName, + Attributes: { VisibilityTimeout: '2', MessageRetentionPeriod: '345600' }, // short visibility timeout for tests + }))).QueueUrl; + + const fifoUrl = (await execute(client, 'CreateQueue (FIFO)', new CreateQueueCommand({ + QueueName: fifoQueueName, + Attributes: { FifoQueue: 'true', ContentBasedDeduplication: 'true', VisibilityTimeout: '10' }, + }))).QueueUrl; + + // Verify GetQueueUrl + const getUrlOut = await execute(client, 'GetQueueUrl (Std)', new GetQueueUrlCommand({ QueueName: stdQueueName })); + expectEqual(getUrlOut.QueueUrl, stdUrl, 'GetQueueUrl QueueUrl'); + + // Verify ListQueues + const listQueuesOut = await execute(client, 'ListQueues', new ListQueuesCommand({ QueueNamePrefix: 'mildstack-smoke' })); + if (!listQueuesOut.QueueUrls || listQueuesOut.QueueUrls.length < 3) { + throw new Error('ListQueues did not return created queues'); + } + + // --- Validation: Visibility Timeout and DLQ --- + console.log('\n--- Validation: Visibility Timeout and DLQ ---'); + await execute(client, 'SetQueueAttributes (Std -> DLQ)', new SetQueueAttributesCommand({ + QueueUrl: stdUrl, + Attributes: { + RedrivePolicy: JSON.stringify({ deadLetterTargetArn: dlqArn, maxReceiveCount: 2 }), + }, + })); + + const dlqSources = await execute(client, 'ListDeadLetterSourceQueues', new ListDeadLetterSourceQueuesCommand({ QueueUrl: dlqUrl })); + if (!dlqSources.queueUrls || !dlqSources.queueUrls.includes(stdUrl)) { + console.warn('\nWarning: ListDeadLetterSourceQueues did not include the source queue. This might be unsupported by the emulator.'); + } + + await execute(client, 'SendMessage (for DLQ test)', new SendMessageCommand({ + QueueUrl: stdUrl, MessageBody: 'dlq test message', + })); + + const r1 = await execute(client, 'Receive 1 (DLQ test)', new ReceiveMessageCommand({ QueueUrl: stdUrl, MaxNumberOfMessages: 1 })); + expectEqual(r1.Messages?.length, 1, 'Should receive message first time'); + + const r2 = await execute(client, 'Receive 2 immediately (hidden)', new ReceiveMessageCommand({ QueueUrl: stdUrl, WaitTimeSeconds: 0 })); + expectEqual(r2.Messages?.length || 0, 0, 'Message should be invisible'); + + console.log('\nWaiting for visibility timeout (3s)...'); + await setTimeout(3000); + + const r3 = await execute(client, 'Receive 3 (DLQ test)', new ReceiveMessageCommand({ QueueUrl: stdUrl, MaxNumberOfMessages: 1 })); + expectEqual(r3.Messages?.length, 1, 'Should receive message second time'); + + console.log('\nWaiting for visibility timeout again (3s)...'); + await setTimeout(3000); + + const r4 = await execute(client, 'Receive 4 from Std (Should be empty)', new ReceiveMessageCommand({ QueueUrl: stdUrl, MaxNumberOfMessages: 1, WaitTimeSeconds: 0 })); + expectEqual(r4.Messages?.length || 0, 0, 'Message should have been moved to DLQ'); + + const r5 = await execute(client, 'Receive from DLQ', new ReceiveMessageCommand({ QueueUrl: dlqUrl, MaxNumberOfMessages: 1 })); + expectEqual(r5.Messages?.length, 1, 'Message should be in DLQ'); + await execute(client, 'Delete from DLQ', new DeleteMessageCommand({ QueueUrl: dlqUrl, ReceiptHandle: r5.Messages[0].ReceiptHandle })); + + + // --- Validation: DelaySeconds --- + console.log('\n--- Validation: DelaySeconds ---'); + await execute(client, 'SetQueueAttributes (DelaySeconds)', new SetQueueAttributesCommand({ + QueueUrl: stdUrl, Attributes: { DelaySeconds: '2' }, + })); + + await execute(client, 'SendMessage (Queue Delay)', new SendMessageCommand({ + QueueUrl: stdUrl, MessageBody: 'queue delay message', + })); + + await execute(client, 'SendMessage (Message Delay Override)', new SendMessageCommand({ + QueueUrl: stdUrl, MessageBody: 'message delay message', DelaySeconds: 4, + })); + + const rd1 = await execute(client, 'Receive immediately (both hidden)', new ReceiveMessageCommand({ QueueUrl: stdUrl, WaitTimeSeconds: 0 })); + expectEqual(rd1.Messages?.length || 0, 0, 'Messages should be delayed'); + + console.log('\nWaiting for queue delay (3s)...'); + await setTimeout(3000); + + const rd2 = await execute(client, 'Receive after 3s', new ReceiveMessageCommand({ QueueUrl: stdUrl, WaitTimeSeconds: 0 })); + expectEqual(rd2.Messages?.length, 1, 'Should receive only the queue-delayed message'); + expectEqual(rd2.Messages[0].Body, 'queue delay message', 'Check correct message body'); + await execute(client, 'Delete message', new DeleteMessageCommand({ QueueUrl: stdUrl, ReceiptHandle: rd2.Messages[0].ReceiptHandle })); + + console.log('\nWaiting for message delay (2s more)...'); + await setTimeout(2000); + const rd3 = await execute(client, 'Receive after 5s total', new ReceiveMessageCommand({ QueueUrl: stdUrl, WaitTimeSeconds: 0 })); + expectEqual(rd3.Messages?.length, 1, 'Should receive the message-delayed message'); + expectEqual(rd3.Messages[0].Body, 'message delay message', 'Check correct message body'); + await execute(client, 'Delete message', new DeleteMessageCommand({ QueueUrl: stdUrl, ReceiptHandle: rd3.Messages[0].ReceiptHandle })); + + await execute(client, 'SetQueueAttributes (Reset Delay)', new SetQueueAttributesCommand({ + QueueUrl: stdUrl, Attributes: { DelaySeconds: '0' }, + })); + + + // --- Validation: Batch Send, Receive, and Visibility --- + console.log('\n--- Validation: Batch Send and Visibility ---'); + const batchSendOut = await execute(client, 'SendMessageBatch', new SendMessageBatchCommand({ + QueueUrl: stdUrl, + Entries: [ + { Id: 'm1', MessageBody: 'batch 1' }, + { Id: 'm2', MessageBody: 'batch 2' }, + { Id: 'm3', MessageBody: 'batch 3', DelaySeconds: 3 }, // delayed message + ], + })); + expectEqual(batchSendOut.Successful?.length, 3, 'All 3 messages should be sent successfully'); + + const rb1 = await execute(client, 'Receive Batch', new ReceiveMessageCommand({ QueueUrl: stdUrl, MaxNumberOfMessages: 10, WaitTimeSeconds: 0 })); + expectEqual(rb1.Messages?.length, 2, 'Should receive 2 batch messages immediately'); + + const m1 = rb1.Messages.find(m => m.Body === 'batch 1'); + const m2 = rb1.Messages.find(m => m.Body === 'batch 2'); + + const visOut = await execute(client, 'ChangeMessageVisibilityBatch', new ChangeMessageVisibilityBatchCommand({ + QueueUrl: stdUrl, + Entries: [ + { Id: 'v1', ReceiptHandle: m1.ReceiptHandle, VisibilityTimeout: 0 }, + ], + })); + expectEqual(visOut.Successful?.length, 1, 'Visibility change should succeed'); + + const rb2 = await execute(client, 'Receive after visibility change', new ReceiveMessageCommand({ QueueUrl: stdUrl, WaitTimeSeconds: 0 })); + expectEqual(rb2.Messages?.length, 1, 'Should receive only m1 since m2 is still invisible'); + expectEqual(rb2.Messages[0].Body, 'batch 1', 'Check body'); + + const delOut = await execute(client, 'DeleteMessageBatch', new DeleteMessageBatchCommand({ + QueueUrl: stdUrl, + Entries: [ + { Id: 'd1', ReceiptHandle: rb2.Messages[0].ReceiptHandle }, + { Id: 'd2', ReceiptHandle: m2.ReceiptHandle }, + ], + })); + expectEqual(delOut.Successful?.length, 2, 'Deletion should succeed'); + + console.log('\nWaiting for delayed batch message (4s)...'); + await setTimeout(4000); + const rb3 = await execute(client, 'Receive delayed batch message', new ReceiveMessageCommand({ QueueUrl: stdUrl, WaitTimeSeconds: 0 })); + expectEqual(rb3.Messages?.length, 1, 'Should receive delayed batch message'); + await execute(client, 'Delete delayed batch message', new DeleteMessageCommand({ QueueUrl: stdUrl, ReceiptHandle: rb3.Messages[0].ReceiptHandle })); + + + // --- Validation: FIFO Deduplication and Ordering --- + console.log('\n--- Validation: FIFO Deduplication ---'); + const f1 = await execute(client, 'SendMessage (FIFO 1)', new SendMessageCommand({ + QueueUrl: fifoUrl, MessageBody: 'dup test', MessageGroupId: 'g1', MessageDeduplicationId: 'dup1', + })); + const f2 = await execute(client, 'SendMessage (FIFO 2 - Duplicate)', new SendMessageCommand({ + QueueUrl: fifoUrl, MessageBody: 'dup test', MessageGroupId: 'g1', MessageDeduplicationId: 'dup1', + })); + + const rf1 = await execute(client, 'Receive FIFO', new ReceiveMessageCommand({ QueueUrl: fifoUrl, MaxNumberOfMessages: 10, WaitTimeSeconds: 0 })); + expectEqual(rf1.Messages?.length, 1, 'Should only receive 1 message due to deduplication'); + expectEqual(rf1.Messages[0].MessageId, f1.MessageId, 'MessageId should match the first one'); + + await execute(client, 'Delete FIFO message', new DeleteMessageCommand({ + QueueUrl: fifoUrl, ReceiptHandle: rf1.Messages[0].ReceiptHandle, + })); + + + // --- Validation: Tags --- + console.log('\n--- Validation: Tags ---'); + await execute(client, 'TagQueue', new TagQueueCommand({ + QueueUrl: stdUrl, Tags: { Env: 'Test', App: 'MildStack' }, + })); + const tagsOut = await execute(client, 'ListQueueTags', new ListQueueTagsCommand({ QueueUrl: stdUrl })); + expectEqual(tagsOut.Tags?.Env, 'Test', 'Tag Env should match'); + + await execute(client, 'UntagQueue', new UntagQueueCommand({ + QueueUrl: stdUrl, TagKeys: ['Env'], + })); + const tagsOut2 = await execute(client, 'ListQueueTags (After untag)', new ListQueueTagsCommand({ QueueUrl: stdUrl })); + expectEqual(tagsOut2.Tags?.Env, undefined, 'Tag Env should be removed'); + expectEqual(tagsOut2.Tags?.App, 'MildStack', 'Tag App should remain'); + + + // --- Cleanup --- + console.log('\n--- Cleanup ---'); + await execute(client, 'PurgeQueue', new PurgeQueueCommand({ QueueUrl: stdUrl })); + await execute(client, 'DeleteQueue (Std)', new DeleteQueueCommand({ QueueUrl: stdUrl })); + await execute(client, 'DeleteQueue (DLQ)', new DeleteQueueCommand({ QueueUrl: dlqUrl })); + await execute(client, 'DeleteQueue (FIFO)', new DeleteQueueCommand({ QueueUrl: fifoUrl })); + + console.log('\n✓ Native AWS SDK smoke mode passed with deep behavioral validations'); +} + +async function execute(client, name, command) { + if (debug) console.log(`\nExecuting ${name}...`); + try { + const response = await client.send(command); + if (debug) { + console.log(`✓ ${name} succeeded.`); + console.dir(response, { depth: 4, colors: true }); + } else { + process.stdout.write('.'); + } + return response; + } catch (error) { + if (!debug) console.log(''); // newline for error + printAwsError(name, error); + throw error; + } +} + +function printAwsError(name, error) { + console.error(`\nFailed during command: ${name}`); + if (error && typeof error === 'object') { + console.error('Error name:', error.name || 'unknown'); + console.error('Error message:', error.message || String(error)); + if (error.$response) { + console.error('Response status:', error.$response.statusCode); + console.error('Response headers:', error.$response.headers); + } + } else { + console.error(error); + } +} + +function expectEqual(actual, expected, label) { + if (actual !== expected) { + throw new Error(`Validation Failed [${label}]: got ${JSON.stringify(actual)} want ${JSON.stringify(expected)}`); + } +} + +function uniqueQueueName(prefix) { + return `mildstack-${prefix}-${Date.now().toString(36)}-${randomUUID().slice(0, 8)}`; +} diff --git a/core/cmd/mildstack/main_test.go b/core/cmd/mildstack/main_test.go index 5d39a09..3773f78 100644 --- a/core/cmd/mildstack/main_test.go +++ b/core/cmd/mildstack/main_test.go @@ -1,11 +1,9 @@ package main import ( - "bytes" "context" "errors" "fmt" - "io" "net/http" "net/http/httptest" "os" @@ -712,14 +710,11 @@ func TestRegisterNativeSQSRoutesExposesAwsCompatibleSmokeSurface(t *testing.T) { rootRecorder := httptest.NewRecorder() rootRequest := httptest.NewRequest(http.MethodGet, "/?Action=ListQueues&Version=2012-11-05", nil) router.Engine().ServeHTTP(rootRecorder, rootRequest) - if got, want := rootRecorder.Code, http.StatusBadRequest; got != want { + if got, want := rootRecorder.Code, http.StatusOK; got != want { t.Fatalf("unexpected sqs root status: got %d want %d", got, want) } - if !strings.Contains(rootRecorder.Body.String(), "") { - t.Fatalf("expected sqs error response xml, got %q", rootRecorder.Body.String()) - } - if !strings.Contains(rootRecorder.Body.String(), "UnsupportedOperation") { - t.Fatalf("expected unsupported operation xml, got %q", rootRecorder.Body.String()) + if !strings.Contains(rootRecorder.Body.String(), "ListQueuesResponse") { + t.Fatalf("expected sqs list queues xml response, got %q", rootRecorder.Body.String()) } server := httptest.NewServer(router.Engine()) @@ -728,11 +723,9 @@ func TestRegisterNativeSQSRoutesExposesAwsCompatibleSmokeSurface(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) t.Cleanup(cancel) - transport := &captureTransport{base: http.DefaultTransport} cfg, err := config.LoadDefaultConfig(ctx, config.WithRegion("us-east-1"), config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider("test", "test", "test")), - config.WithHTTPClient(&http.Client{Transport: transport}), ) if err != nil { t.Fatalf("load aws config: %v", err) @@ -742,42 +735,16 @@ func TestRegisterNativeSQSRoutesExposesAwsCompatibleSmokeSurface(t *testing.T) { o.BaseEndpoint = aws.String(server.URL) }) - _, err = client.ListQueues(ctx, &sqssdk.ListQueuesInput{}) - if err == nil { - t.Fatal("expected list queues to return an error") - } - if !strings.Contains(string(transport.body), "") { - t.Fatalf("expected captured sqs xml body, got %q", string(transport.body)) - } - if !strings.Contains(string(transport.body), "UnsupportedOperation") && !strings.Contains(string(transport.body), "InvalidQueryParameter") { - t.Fatalf("expected captured sqs xml body to contain an SQS error code, got %q", string(transport.body)) - } -} - -type captureTransport struct { - base http.RoundTripper - body []byte -} - -func (t *captureTransport) RoundTrip(req *http.Request) (*http.Response, error) { - resp, err := t.base.RoundTrip(req) + result, err := client.ListQueues(ctx, &sqssdk.ListQueuesInput{}) if err != nil { - return nil, err + t.Fatalf("expected list queues to succeed, got: %v", err) } - if resp.Body == nil { - return resp, nil + if result == nil { + t.Fatal("expected non-nil list queues result") } - - data, readErr := io.ReadAll(resp.Body) - _ = resp.Body.Close() - if readErr != nil { - return nil, readErr - } - t.body = append([]byte(nil), data...) - resp.Body = io.NopCloser(bytes.NewReader(data)) - return resp, nil } + func newDynamoDBSmokeClient(t *testing.T, endpoint string) *dynamodb.Client { t.Helper() diff --git a/core/internal/delivery/http/dynamodb_native.go b/core/internal/delivery/http/dynamodb_native.go index 32fb399..cc3114d 100644 --- a/core/internal/delivery/http/dynamodb_native.go +++ b/core/internal/delivery/http/dynamodb_native.go @@ -20,13 +20,13 @@ import ( type DynamoDBNativeService interface { ListTables() []dynamodbdomain.Table - CreateTable(name, partitionKey, sortKey, billingMode string) (dynamodbdomain.Table, error) + CreateTable(name, partitionKey, sortKey, billingMode string, specs ...dynamodbdomain.CreateTableSpec) (dynamodbdomain.Table, error) DescribeTable(name string) (dynamodbdomain.Table, error) DeleteTable(name string) (dynamodbdomain.Table, error) GetItem(table, key string) (dynamodbdomain.Item, error) PutItem(table, key string, attributes map[string]dynamodbdomain.AttributeValue) (dynamodbdomain.Item, error) UpdateItem(table, key, updateExpression, conditionExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]dynamodbdomain.AttributeValue) (dynamodbdomain.Item, error) - Query(table, keyConditionExpression, filterExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]dynamodbdomain.AttributeValue, limit *int, exclusiveStartKey map[string]dynamodbdomain.AttributeValue, scanIndexForward *bool) (dynamodbdomain.ReadPage, error) + Query(table, keyConditionExpression, filterExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]dynamodbdomain.AttributeValue, limit *int, exclusiveStartKey map[string]dynamodbdomain.AttributeValue, scanIndexForward *bool, options ...dynamodbdomain.QueryOptions) (dynamodbdomain.ReadPage, error) Scan(table, filterExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]dynamodbdomain.AttributeValue, limit *int, exclusiveStartKey map[string]dynamodbdomain.AttributeValue) (dynamodbdomain.ReadPage, error) DeleteItem(table, key string) error BatchWriteItem(request ddbcontracts.BatchWriteItemRequest) (ddbcontracts.BatchWriteItemResult, error) @@ -57,8 +57,25 @@ func RegisterDynamoDBNativeRoutes(engine *gin.Engine, service DynamoDBNativeServ } type dynamoDBNativeHandler struct { - service DynamoDBNativeService + service serviceWithIndexes registry map[string]dynamoTargetSpec + indexes map[string]map[string]nativeDynamoIndex + ttl map[string]nativeTTLDescription +} + +type serviceWithIndexes interface { + DynamoDBNativeService +} + +type nativeDynamoIndex struct { + Name string + PartitionKey string + SortKey string +} + +type nativeTTLDescription struct { + AttributeName string + Status string } type dynamoTargetSpec struct { @@ -70,6 +87,8 @@ func newDynamoDBNativeHandler(service DynamoDBNativeService) dynamoDBNativeHandl return dynamoDBNativeHandler{ service: service, registry: newDynamoDBTargetRegistry(), + indexes: make(map[string]map[string]nativeDynamoIndex), + ttl: make(map[string]nativeTTLDescription), } } @@ -131,6 +150,14 @@ func newDynamoDBTargetRegistry() map[string]dynamoTargetSpec { supported: true, execute: (*dynamoDBNativeHandler).handleTransactWriteItems, }, + "UpdateTimeToLive": { + supported: true, + execute: (*dynamoDBNativeHandler).handleUpdateTimeToLive, + }, + "DescribeTimeToLive": { + supported: true, + execute: (*dynamoDBNativeHandler).handleDescribeTimeToLive, + }, "UpdateTable": {supported: false}, } } @@ -156,7 +183,12 @@ func (h dynamoDBNativeHandler) dispatch(c *gin.Context) bool { return false } - targetName, err := parseDynamoDBTarget(c.Request.Header.Get("X-Amz-Target")) + rawTarget := c.Request.Header.Get("X-Amz-Target") + if !strings.HasPrefix(rawTarget, dynamoDBTargetPrefix) { + return false + } + + targetName, err := parseDynamoDBTarget(rawTarget) if err != nil { writeDynamoDBError(c, http.StatusBadRequest, "ValidationException", err.Error()) return true @@ -218,22 +250,19 @@ func (h *dynamoDBNativeHandler) handleCreateTable(c *gin.Context, body []byte) e partitionKey, sortKey := partitionAndSortKeys(request.KeySchema) billingMode := strings.TrimSpace(request.BillingMode) + spec := dynamodbdomain.CreateTableSpec{ + AttributeDefinitions: attributeDefinitionsToDomain(request.AttributeDefinitions), + GlobalSecondaryIndexes: secondaryIndexesToDomain(request.GlobalSecondaryIndexes), + LocalSecondaryIndexes: secondaryIndexesToDomain(request.LocalSecondaryIndexes), + } - table, err := h.service.CreateTable(tableName, partitionKey, sortKey, billingMode) + table, err := h.service.CreateTable(tableName, partitionKey, sortKey, billingMode, spec) if err != nil { return err } writeDynamoDBJSON(c, http.StatusOK, createTableResponse{ - TableDescription: tableDescription{ - TableName: table.Name, - TableStatus: table.Status, - TableArn: awscontext.Default().DynamoDBTableARN(table.Name), - CreationDateTime: awsTimestamp(table.CreatedAt), - KeySchema: cloneKeySchema(request.KeySchema, partitionKey, sortKey), - AttributeDefinitions: cloneAttributeDefinitions(request.AttributeDefinitions), - BillingModeSummary: billingModeSummaryFor(table.BillingMode), - }, + TableDescription: tableDescriptionFromDomain(table), }) return nil } @@ -293,13 +322,22 @@ func (h *dynamoDBNativeHandler) handleGetItem(c *gin.Context, body []byte) error return fmt.Errorf("dynamodb: table name is required") } - key, err := keyFromAttributeValueMap(request.Key) + table, err := h.service.DescribeTable(tableName) + if err != nil { + return err + } + + key, err := keyFromAttributeValueMap(table, request.Key) if err != nil { return err } item, err := h.service.GetItem(tableName, key) if err != nil { + if strings.Contains(err.Error(), "item ") && strings.Contains(err.Error(), " not found") { + writeDynamoDBJSON(c, http.StatusOK, getItemResponse{}) + return nil + } return err } @@ -320,11 +358,33 @@ func (h *dynamoDBNativeHandler) handlePutItem(c *gin.Context, body []byte) error return fmt.Errorf("dynamodb: table name is required") } - key, attributes, err := itemFromAttributeValueMap(request.Item) + table, err := h.service.DescribeTable(tableName) if err != nil { return err } + key, attributes, err := itemFromAttributeValueMap(table, request.Item) + if err != nil { + return err + } + + expressionAttributeValues := make(map[string]dynamodbdomain.AttributeValue, len(request.ExpressionAttributeValues)) + for name, value := range request.ExpressionAttributeValues { + converted, err := attributeValueToDomain(value) + if err != nil { + return err + } + expressionAttributeValues[name] = converted + } + + current, err := h.service.GetItem(tableName, key) + if err != nil && !strings.Contains(err.Error(), "not found") { + return err + } + if err := evaluateNativeCondition(current.Attributes, request.ConditionExpression, request.ExpressionAttributeNames, expressionAttributeValues); err != nil { + return err + } + item, err := h.service.PutItem(tableName, key, attributes) if err != nil { return err @@ -347,7 +407,12 @@ func (h *dynamoDBNativeHandler) handleDeleteItem(c *gin.Context, body []byte) er return fmt.Errorf("dynamodb: table name is required") } - key, err := keyFromAttributeValueMap(request.Key) + table, err := h.service.DescribeTable(tableName) + if err != nil { + return err + } + + key, err := keyFromAttributeValueMap(table, request.Key) if err != nil { return err } @@ -371,7 +436,12 @@ func (h *dynamoDBNativeHandler) handleUpdateItem(c *gin.Context, body []byte) er return fmt.Errorf("dynamodb: table name is required") } - key, err := keyFromAttributeValueMap(request.Key) + table, err := h.service.DescribeTable(tableName) + if err != nil { + return err + } + + key, err := keyFromAttributeValueMap(table, request.Key) if err != nil { return err } @@ -422,12 +492,6 @@ func (h *dynamoDBNativeHandler) handleQuery(c *gin.Context, body []byte) error { if strings.TrimSpace(request.KeyConditionExpression) == "" { return fmt.Errorf("dynamodb: key condition expression is required") } - if strings.TrimSpace(request.IndexName) != "" { - return fmt.Errorf("dynamodb: index queries are not supported") - } - if strings.TrimSpace(request.ProjectionExpression) != "" { - return fmt.Errorf("dynamodb: projection expressions are not supported") - } if selectValue := strings.ToUpper(strings.TrimSpace(request.Select)); selectValue != "" && selectValue != "ALL_ATTRIBUTES" { return fmt.Errorf("dynamodb: unsupported select value %q", request.Select) } @@ -446,16 +510,34 @@ func (h *dynamoDBNativeHandler) handleQuery(c *gin.Context, body []byte) error { return err } - result, err := h.service.Query( - tableName, - request.KeyConditionExpression, - request.FilterExpression, - request.ExpressionAttributeNames, - expressionAttributeValues, - request.Limit, - exclusiveStartKey, - request.ScanIndexForward, - ) + var result dynamodbdomain.ReadPage + if strings.TrimSpace(request.IndexName) != "" { + result, err = h.service.Query( + tableName, + request.KeyConditionExpression, + request.FilterExpression, + request.ExpressionAttributeNames, + expressionAttributeValues, + request.Limit, + exclusiveStartKey, + request.ScanIndexForward, + dynamodbdomain.QueryOptions{ + IndexName: strings.TrimSpace(request.IndexName), + ProjectionExpression: request.ProjectionExpression, + }, + ) + } else { + result, err = h.service.Query( + tableName, + request.KeyConditionExpression, + request.FilterExpression, + request.ExpressionAttributeNames, + expressionAttributeValues, + request.Limit, + exclusiveStartKey, + request.ScanIndexForward, + ) + } if err != nil { return err } @@ -703,8 +785,59 @@ func (h *dynamoDBNativeHandler) handleTransactWriteItems(c *gin.Context, body [] Table: item.Delete.TableName, DeleteKey: deleteKey, }) - case item.Update != nil || item.ConditionCheck != nil: - return fmt.Errorf("dynamodb: update and condition check transaction items are not supported") + case item.Update != nil: + if item.Put != nil || item.Delete != nil || item.ConditionCheck != nil { + return fmt.Errorf("dynamodb: each transaction item must contain exactly one operation") + } + if strings.TrimSpace(item.Update.TableName) == "" { + return fmt.Errorf("dynamodb: table name is required") + } + updateKey, err := attributeValueMapToDomain(item.Update.Key) + if err != nil { + return err + } + updateExpressionValues := make(map[string]dynamodbdomain.AttributeValue, len(item.Update.ExpressionAttributeValues)) + for name, value := range item.Update.ExpressionAttributeValues { + converted, err := attributeValueToDomain(value) + if err != nil { + return err + } + updateExpressionValues[name] = converted + } + appRequest.Items = append(appRequest.Items, ddbcontracts.TransactWriteItem{ + Table: item.Update.TableName, + UpdateKey: updateKey, + UpdateExpression: item.Update.UpdateExpression, + ConditionExpression: item.Update.ConditionExpression, + ExpressionAttributeNames: item.Update.ExpressionAttributeNames, + ExpressionAttributeValues: updateExpressionValues, + }) + case item.ConditionCheck != nil: + if item.Put != nil || item.Delete != nil || item.Update != nil { + return fmt.Errorf("dynamodb: each transaction item must contain exactly one operation") + } + if strings.TrimSpace(item.ConditionCheck.TableName) == "" { + return fmt.Errorf("dynamodb: table name is required") + } + checkKey, err := attributeValueMapToDomain(item.ConditionCheck.Key) + if err != nil { + return err + } + checkExpressionValues := make(map[string]dynamodbdomain.AttributeValue, len(item.ConditionCheck.ExpressionAttributeValues)) + for name, value := range item.ConditionCheck.ExpressionAttributeValues { + converted, err := attributeValueToDomain(value) + if err != nil { + return err + } + checkExpressionValues[name] = converted + } + appRequest.Items = append(appRequest.Items, ddbcontracts.TransactWriteItem{ + Table: item.ConditionCheck.TableName, + ConditionCheckKey: checkKey, + ConditionExpression: item.ConditionCheck.ConditionExpression, + ExpressionAttributeNames: item.ConditionCheck.ExpressionAttributeNames, + ExpressionAttributeValues: checkExpressionValues, + }) default: return fmt.Errorf("dynamodb: each transaction item must contain exactly one operation") } @@ -770,6 +903,67 @@ func (h *dynamoDBNativeHandler) handleTransactGetItems(c *gin.Context, body []by return nil } +func (h *dynamoDBNativeHandler) handleUpdateTimeToLive(c *gin.Context, body []byte) error { + request := updateTimeToLiveRequest{} + if err := json.Unmarshal(body, &request); err != nil { + return fmt.Errorf("dynamodb: invalid UpdateTimeToLive request: %w", err) + } + + tableName := strings.TrimSpace(request.TableName) + if tableName == "" { + return fmt.Errorf("dynamodb: table name is required") + } + if _, err := h.service.DescribeTable(tableName); err != nil { + return err + } + if strings.TrimSpace(request.TimeToLiveSpecification.AttributeName) == "" { + return fmt.Errorf("dynamodb: time to live attribute name is required") + } + + status := "DISABLED" + if request.TimeToLiveSpecification.Enabled { + status = "ENABLED" + } + h.ttl[tableName] = nativeTTLDescription{ + AttributeName: request.TimeToLiveSpecification.AttributeName, + Status: status, + } + + writeDynamoDBJSON(c, http.StatusOK, updateTimeToLiveResponse{ + TimeToLiveSpecification: timeToLiveSpecification{ + AttributeName: request.TimeToLiveSpecification.AttributeName, + Enabled: request.TimeToLiveSpecification.Enabled, + }, + }) + return nil +} + +func (h *dynamoDBNativeHandler) handleDescribeTimeToLive(c *gin.Context, body []byte) error { + request := describeTimeToLiveRequest{} + if err := json.Unmarshal(body, &request); err != nil { + return fmt.Errorf("dynamodb: invalid DescribeTimeToLive request: %w", err) + } + + tableName := strings.TrimSpace(request.TableName) + if tableName == "" { + return fmt.Errorf("dynamodb: table name is required") + } + if _, err := h.service.DescribeTable(tableName); err != nil { + return err + } + + description := timeToLiveDescription{TimeToLiveStatus: "DISABLED"} + if ttl, ok := h.ttl[tableName]; ok { + description.AttributeName = ttl.AttributeName + description.TimeToLiveStatus = ttl.Status + } + + writeDynamoDBJSON(c, http.StatusOK, describeTimeToLiveResponse{ + TimeToLiveDescription: description, + }) + return nil +} + func isDynamoDBJSONRequest(contentType string) bool { mediaType, _, err := mime.ParseMediaType(strings.TrimSpace(contentType)) if err != nil { @@ -911,16 +1105,24 @@ type listTablesResponse struct { } type createTableRequest struct { - TableName string `json:"TableName"` - BillingMode string `json:"BillingMode,omitempty"` - KeySchema []dynamoKeySchemaElement `json:"KeySchema,omitempty"` - AttributeDefinitions []dynamoAttributeDefinition `json:"AttributeDefinitions,omitempty"` + TableName string `json:"TableName"` + BillingMode string `json:"BillingMode,omitempty"` + KeySchema []dynamoKeySchemaElement `json:"KeySchema,omitempty"` + AttributeDefinitions []dynamoAttributeDefinition `json:"AttributeDefinitions,omitempty"` + GlobalSecondaryIndexes []dynamoSecondaryIndexDefinition `json:"GlobalSecondaryIndexes,omitempty"` + LocalSecondaryIndexes []dynamoSecondaryIndexDefinition `json:"LocalSecondaryIndexes,omitempty"` } type createTableResponse struct { TableDescription tableDescription `json:"TableDescription"` } +type dynamoSecondaryIndexDefinition struct { + IndexName string `json:"IndexName"` + KeySchema []dynamoKeySchemaElement `json:"KeySchema,omitempty"` + Projection dynamoProjection `json:"Projection,omitempty"` +} + type describeTableRequest struct { TableName string `json:"TableName"` } @@ -947,8 +1149,11 @@ type getItemResponse struct { } type putItemRequest struct { - TableName string `json:"TableName"` - Item map[string]dynamoAttributeValue `json:"Item"` + TableName string `json:"TableName"` + Item map[string]dynamoAttributeValue `json:"Item"` + ConditionExpression string `json:"ConditionExpression,omitempty"` + ExpressionAttributeNames map[string]string `json:"ExpressionAttributeNames,omitempty"` + ExpressionAttributeValues map[string]dynamoAttributeValue `json:"ExpressionAttributeValues,omitempty"` } type putItemResponse struct { @@ -1091,11 +1296,22 @@ type transactWriteDeleteRequest struct { } type transactWriteUpdateRequest struct { - TableName string `json:"TableName"` + TableName string `json:"TableName"` + Key map[string]dynamoAttributeValue `json:"Key"` + UpdateExpression string `json:"UpdateExpression"` + ConditionExpression string `json:"ConditionExpression,omitempty"` + ExpressionAttributeNames map[string]string `json:"ExpressionAttributeNames,omitempty"` + ExpressionAttributeValues map[string]dynamoAttributeValue `json:"ExpressionAttributeValues,omitempty"` + ReturnValuesOnConditionCheckFailure string `json:"ReturnValuesOnConditionCheckFailure,omitempty"` } type transactWriteConditionCheckRequest struct { - TableName string `json:"TableName"` + TableName string `json:"TableName"` + Key map[string]dynamoAttributeValue `json:"Key"` + ConditionExpression string `json:"ConditionExpression"` + ExpressionAttributeNames map[string]string `json:"ExpressionAttributeNames,omitempty"` + ExpressionAttributeValues map[string]dynamoAttributeValue `json:"ExpressionAttributeValues,omitempty"` + ReturnValuesOnConditionCheckFailure string `json:"ReturnValuesOnConditionCheckFailure,omitempty"` } type transactWriteItemsResponse struct{} @@ -1124,6 +1340,33 @@ type transactGetItemResponse struct { Item map[string]dynamoAttributeValue `json:"Item,omitempty"` } +type updateTimeToLiveRequest struct { + TableName string `json:"TableName"` + TimeToLiveSpecification timeToLiveSpecification `json:"TimeToLiveSpecification"` +} + +type describeTimeToLiveRequest struct { + TableName string `json:"TableName"` +} + +type updateTimeToLiveResponse struct { + TimeToLiveSpecification timeToLiveSpecification `json:"TimeToLiveSpecification"` +} + +type describeTimeToLiveResponse struct { + TimeToLiveDescription timeToLiveDescription `json:"TimeToLiveDescription"` +} + +type timeToLiveSpecification struct { + AttributeName string `json:"AttributeName,omitempty"` + Enabled bool `json:"Enabled"` +} + +type timeToLiveDescription struct { + AttributeName string `json:"AttributeName,omitempty"` + TimeToLiveStatus string `json:"TimeToLiveStatus,omitempty"` +} + type dynamoAttributeValue struct { S string `json:"S,omitempty"` N string `json:"N,omitempty"` @@ -1143,14 +1386,21 @@ type dynamoAttributeDefinition struct { AttributeType string `json:"AttributeType"` } +type dynamoProjection struct { + ProjectionType string `json:"ProjectionType,omitempty"` + NonKeyAttributes []string `json:"NonKeyAttributes,omitempty"` +} + type tableDescription struct { - TableName string `json:"TableName"` - TableStatus string `json:"TableStatus"` - TableArn string `json:"TableArn,omitempty"` - CreationDateTime int64 `json:"CreationDateTime,omitempty"` - KeySchema []dynamoKeySchemaElement `json:"KeySchema,omitempty"` - AttributeDefinitions []dynamoAttributeDefinition `json:"AttributeDefinitions,omitempty"` - BillingModeSummary *billingModeSummary `json:"BillingModeSummary,omitempty"` + TableName string `json:"TableName"` + TableStatus string `json:"TableStatus"` + TableArn string `json:"TableArn,omitempty"` + CreationDateTime int64 `json:"CreationDateTime,omitempty"` + KeySchema []dynamoKeySchemaElement `json:"KeySchema,omitempty"` + AttributeDefinitions []dynamoAttributeDefinition `json:"AttributeDefinitions,omitempty"` + GlobalSecondaryIndexes []dynamoSecondaryIndexDefinition `json:"GlobalSecondaryIndexes,omitempty"` + LocalSecondaryIndexes []dynamoSecondaryIndexDefinition `json:"LocalSecondaryIndexes,omitempty"` + BillingModeSummary *billingModeSummary `json:"BillingModeSummary,omitempty"` } type billingModeSummary struct { @@ -1221,13 +1471,108 @@ func cloneAttributeDefinitions(source []dynamoAttributeDefinition) []dynamoAttri func tableDescriptionFromDomain(table dynamodbdomain.Table) tableDescription { aws := awscontext.Default() return tableDescription{ - TableName: table.Name, - TableStatus: table.Status, - TableArn: aws.DynamoDBTableARN(table.Name), - CreationDateTime: awsTimestamp(table.CreatedAt), - KeySchema: cloneKeySchema(nil, table.PartitionKey, table.SortKey), - BillingModeSummary: billingModeSummaryFor(table.BillingMode), + TableName: table.Name, + TableStatus: table.Status, + TableArn: aws.DynamoDBTableARN(table.Name), + CreationDateTime: awsTimestamp(table.CreatedAt), + KeySchema: cloneKeySchema(nil, table.PartitionKey, table.SortKey), + AttributeDefinitions: attributeDefinitionsFromDomain(table.AttributeDefinitions), + GlobalSecondaryIndexes: secondaryIndexesFromDomain(table.GlobalSecondaryIndexes), + LocalSecondaryIndexes: secondaryIndexesFromDomain(table.LocalSecondaryIndexes), + BillingModeSummary: billingModeSummaryFor(table.BillingMode), + } +} + +func attributeDefinitionsFromDomain(source []dynamodbdomain.AttributeDefinition) []dynamoAttributeDefinition { + if len(source) == 0 { + return nil + } + cloned := make([]dynamoAttributeDefinition, len(source)) + for i, definition := range source { + cloned[i] = dynamoAttributeDefinition{ + AttributeName: definition.Name, + AttributeType: definition.Type, + } + } + return cloned +} + +func attributeDefinitionsToDomain(source []dynamoAttributeDefinition) []dynamodbdomain.AttributeDefinition { + if len(source) == 0 { + return nil + } + cloned := make([]dynamodbdomain.AttributeDefinition, len(source)) + for i, definition := range source { + cloned[i] = dynamodbdomain.AttributeDefinition{ + Name: strings.TrimSpace(definition.AttributeName), + Type: strings.ToUpper(strings.TrimSpace(definition.AttributeType)), + } } + return cloned +} + +func secondaryIndexesFromDomain(source []dynamodbdomain.SecondaryIndex) []dynamoSecondaryIndexDefinition { + if len(source) == 0 { + return nil + } + cloned := make([]dynamoSecondaryIndexDefinition, len(source)) + for i, index := range source { + cloned[i] = dynamoSecondaryIndexDefinition{ + IndexName: index.Name, + KeySchema: keySchemaFromDomain(index.KeySchema), + Projection: dynamoProjection{ + ProjectionType: index.Projection.Type, + NonKeyAttributes: append([]string(nil), index.Projection.NonKeyAttributes...), + }, + } + } + return cloned +} + +func secondaryIndexesToDomain(source []dynamoSecondaryIndexDefinition) []dynamodbdomain.SecondaryIndex { + if len(source) == 0 { + return nil + } + cloned := make([]dynamodbdomain.SecondaryIndex, len(source)) + for i, index := range source { + cloned[i] = dynamodbdomain.SecondaryIndex{ + Name: strings.TrimSpace(index.IndexName), + KeySchema: keySchemaToDomain(index.KeySchema), + Projection: dynamodbdomain.Projection{ + Type: strings.ToUpper(strings.TrimSpace(index.Projection.ProjectionType)), + NonKeyAttributes: append([]string(nil), index.Projection.NonKeyAttributes...), + }, + } + } + return cloned +} + +func keySchemaFromDomain(source []dynamodbdomain.KeySchemaElement) []dynamoKeySchemaElement { + if len(source) == 0 { + return nil + } + cloned := make([]dynamoKeySchemaElement, len(source)) + for i, element := range source { + cloned[i] = dynamoKeySchemaElement{ + AttributeName: element.AttributeName, + KeyType: element.KeyType, + } + } + return cloned +} + +func keySchemaToDomain(source []dynamoKeySchemaElement) []dynamodbdomain.KeySchemaElement { + if len(source) == 0 { + return nil + } + cloned := make([]dynamodbdomain.KeySchemaElement, len(source)) + for i, element := range source { + cloned[i] = dynamodbdomain.KeySchemaElement{ + AttributeName: strings.TrimSpace(element.AttributeName), + KeyType: strings.ToUpper(strings.TrimSpace(element.KeyType)), + } + } + return cloned } func awsTimestamp(value time.Time) int64 { @@ -1237,29 +1582,51 @@ func awsTimestamp(value time.Time) int64 { return value.Unix() } -func keyFromAttributeValueMap(values map[string]dynamoAttributeValue) (string, error) { +func keyFromAttributeValueMap(table dynamodbdomain.Table, values map[string]dynamoAttributeValue) (string, error) { if len(values) == 0 { return "", fmt.Errorf("dynamodb: key is required") } - if key, ok, err := syntheticItemKey(values); ok || err != nil { - return key, err + if strings.TrimSpace(table.PartitionKey) == "" { + return "", fmt.Errorf("dynamodb: table %q has no partition key", table.Name) } - if len(values) == 1 { - for _, value := range values { - return attributeValueToString(value) - } + expectedCount := 1 + if strings.TrimSpace(table.SortKey) != "" { + expectedCount++ + } + if len(values) != expectedCount { + return "", fmt.Errorf("dynamodb: unsupported key attributes %q", strings.Join(sortedDynamoAttributeKeys(values), ", ")) } - keys := make([]string, 0, len(values)) - for key := range values { - keys = append(keys, key) + partitionValue, ok := values[table.PartitionKey] + if !ok { + return "", fmt.Errorf("dynamodb: missing key attribute %q", table.PartitionKey) + } + + partitionKey, err := attributeValueToString(partitionValue) + if err != nil { + return "", err + } + + if strings.TrimSpace(table.SortKey) == "" { + return partitionKey, nil + } + + sortValue, ok := values[table.SortKey] + if !ok { + return "", fmt.Errorf("dynamodb: missing key attribute %q", table.SortKey) + } + + sortKey, err := attributeValueToString(sortValue) + if err != nil { + return "", err } - return "", fmt.Errorf("dynamodb: unsupported key attributes %q", strings.Join(keys, ", ")) + + return partitionKey + "|" + sortKey, nil } -func itemFromAttributeValueMap(values map[string]dynamoAttributeValue) (string, map[string]dynamodbdomain.AttributeValue, error) { +func itemFromAttributeValueMap(table dynamodbdomain.Table, values map[string]dynamoAttributeValue) (string, map[string]dynamodbdomain.AttributeValue, error) { if len(values) == 0 { return "", nil, fmt.Errorf("dynamodb: item is required") } @@ -1275,7 +1642,7 @@ func itemFromAttributeValueMap(values map[string]dynamoAttributeValue) (string, attributes[name] = copied } - key, _, err = syntheticItemKey(values) + key, err = itemKeyFromAttributeValueMap(table, values) if err != nil { return "", nil, err } @@ -1286,40 +1653,423 @@ func itemFromAttributeValueMap(values map[string]dynamoAttributeValue) (string, return key, attributes, nil } -func syntheticItemKey(values map[string]dynamoAttributeValue) (string, bool, error) { +func itemKeyFromAttributeValueMap(table dynamodbdomain.Table, values map[string]dynamoAttributeValue) (string, error) { if len(values) == 0 { - return "", false, nil + return "", fmt.Errorf("dynamodb: item is required") + } + if strings.TrimSpace(table.PartitionKey) == "" { + return "", fmt.Errorf("dynamodb: table %q has no partition key", table.Name) + } + + partitionValue, ok := values[table.PartitionKey] + if !ok { + return "", fmt.Errorf("dynamodb: missing key attribute %q", table.PartitionKey) + } + + partitionKey, err := attributeValueToString(partitionValue) + if err != nil { + return "", err + } + + if strings.TrimSpace(table.SortKey) == "" { + return partitionKey, nil + } + + sortValue, ok := values[table.SortKey] + if !ok { + return "", fmt.Errorf("dynamodb: missing key attribute %q", table.SortKey) } - if idValue, ok := values["id"]; ok { - id, err := attributeValueToString(idValue) + sortKey, err := attributeValueToString(sortValue) + if err != nil { + return "", err + } + + return partitionKey + "|" + sortKey, nil +} + +func sortedDynamoAttributeKeys(values map[string]dynamoAttributeValue) []string { + keys := make([]string, 0, len(values)) + for key := range values { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + +func evaluateNativeCondition(attributes map[string]dynamodbdomain.AttributeValue, conditionExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]dynamodbdomain.AttributeValue) error { + expression := strings.TrimSpace(conditionExpression) + if expression == "" { + return nil + } + if strings.ContainsAny(expression, "[]") { + return fmt.Errorf("dynamodb: unsupported condition expression %q", conditionExpression) + } + + if attributes == nil { + attributes = map[string]dynamodbdomain.AttributeValue{} + } + + lower := strings.ToLower(expression) + switch { + case strings.HasPrefix(lower, "attribute_exists(") && strings.HasSuffix(expression, ")"): + path, err := resolveNativeConditionPath(expression[len("attribute_exists("):len(expression)-1], expressionAttributeNames) if err != nil { - return "", false, err + return err + } + if _, ok := attributes[path]; !ok { + return fmt.Errorf("dynamodb: conditional check failed") + } + return nil + case strings.HasPrefix(lower, "attribute_not_exists(") && strings.HasSuffix(expression, ")"): + path, err := resolveNativeConditionPath(expression[len("attribute_not_exists("):len(expression)-1], expressionAttributeNames) + if err != nil { + return err + } + if _, ok := attributes[path]; ok { + return fmt.Errorf("dynamodb: conditional check failed") + } + return nil + case strings.Contains(expression, "="): + parts := strings.SplitN(expression, "=", 2) + path, err := resolveNativeConditionPath(parts[0], expressionAttributeNames) + if err != nil { + return err + } + value, err := resolveNativeConditionValue(parts[1], expressionAttributeValues) + if err != nil { + return err + } + existing, ok := attributes[path] + if !ok || !attributeValueEqual(existing, value) { + return fmt.Errorf("dynamodb: conditional check failed") + } + return nil + default: + return fmt.Errorf("dynamodb: unsupported condition expression %q", conditionExpression) + } +} + +func resolveNativeConditionPath(raw string, expressionAttributeNames map[string]string) (string, error) { + path := strings.TrimSpace(raw) + if path == "" { + return "", fmt.Errorf("dynamodb: condition path is required") + } + if strings.ContainsAny(path, ".[]") { + return "", fmt.Errorf("dynamodb: unsupported nested condition path %q", path) + } + if strings.HasPrefix(path, "#") { + resolved, ok := expressionAttributeNames[path] + if !ok { + return "", fmt.Errorf("dynamodb: unresolved expression attribute name %q", path) } - if skValue, ok := values["sk"]; ok { - sk, err := attributeValueToString(skValue) + path = strings.TrimSpace(resolved) + } + if path == "" { + return "", fmt.Errorf("dynamodb: condition path is required") + } + return path, nil +} + +func resolveNativeConditionValue(raw string, expressionAttributeValues map[string]dynamodbdomain.AttributeValue) (dynamodbdomain.AttributeValue, error) { + token := strings.TrimSpace(raw) + if token == "" { + return dynamodbdomain.AttributeValue{}, fmt.Errorf("dynamodb: condition value is required") + } + if !strings.HasPrefix(token, ":") { + return dynamodbdomain.AttributeValue{}, fmt.Errorf("dynamodb: unsupported literal condition value %q", token) + } + value, ok := expressionAttributeValues[token] + if !ok { + return dynamodbdomain.AttributeValue{}, fmt.Errorf("dynamodb: unresolved expression attribute value %q", token) + } + return value.Clone(), nil +} + +func (h *dynamoDBNativeHandler) storeIndexDefinitions(tableName string, gsi []dynamoSecondaryIndexDefinition, lsi []dynamoSecondaryIndexDefinition) { + if len(gsi) == 0 && len(lsi) == 0 { + return + } + + if h.indexes == nil { + h.indexes = make(map[string]map[string]nativeDynamoIndex) + } + tableIndexes := make(map[string]nativeDynamoIndex, len(gsi)+len(lsi)) + for _, index := range append(cloneSecondaryIndexDefinitions(gsi), cloneSecondaryIndexDefinitions(lsi)...) { + partitionKey, sortKey := partitionAndSortKeys(index.KeySchema) + tableIndexes[index.IndexName] = nativeDynamoIndex{ + Name: index.IndexName, + PartitionKey: partitionKey, + SortKey: sortKey, + } + } + h.indexes[tableName] = tableIndexes +} + +func (h *dynamoDBNativeHandler) queryByIndex(tableName, indexName, keyConditionExpression, filterExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]dynamodbdomain.AttributeValue, limit *int, exclusiveStartKey map[string]dynamodbdomain.AttributeValue, scanIndexForward *bool) (dynamodbdomain.ReadPage, error) { + if limit != nil || len(exclusiveStartKey) > 0 { + return dynamodbdomain.ReadPage{}, fmt.Errorf("dynamodb: indexed query pagination is not supported yet") + } + if strings.TrimSpace(filterExpression) != "" { + return dynamodbdomain.ReadPage{}, fmt.Errorf("dynamodb: indexed query filters are not supported yet") + } + + tableIndexes := h.indexes[tableName] + index, ok := tableIndexes[indexName] + if !ok { + return dynamodbdomain.ReadPage{}, fmt.Errorf("dynamodb: invalid index %q for table %q", indexName, tableName) + } + + all, err := h.service.Scan(tableName, "", nil, nil, nil, nil) + if err != nil { + return dynamodbdomain.ReadPage{}, err + } + + predicate, err := buildNativeIndexPredicate(index, keyConditionExpression, expressionAttributeNames, expressionAttributeValues) + if err != nil { + return dynamodbdomain.ReadPage{}, err + } + + items := make([]dynamodbdomain.Item, 0, len(all.Items)) + for _, item := range all.Items { + match, err := predicate(item) + if err != nil { + return dynamodbdomain.ReadPage{}, err + } + if match { + items = append(items, item) + } + } + + sort.SliceStable(items, func(i, j int) bool { + cmp := compareNativeIndexItems(items[i], items[j], index) + forward := scanIndexForward == nil || *scanIndexForward + if forward { + return cmp < 0 + } + return cmp > 0 + }) + + return dynamodbdomain.ReadPage{ + Items: items, + Count: len(items), + ScannedCount: len(items), + }, nil +} + +func buildNativeIndexPredicate(index nativeDynamoIndex, expression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]dynamodbdomain.AttributeValue) (func(dynamodbdomain.Item) (bool, error), error) { + expression = strings.TrimSpace(expression) + if expression == "" { + return nil, fmt.Errorf("dynamodb: key condition expression is required") + } + + partitionExpr, sortExpr := splitNativeKeyCondition(expression) + if partitionExpr == "" { + return nil, fmt.Errorf("dynamodb: unsupported key condition expression %q", expression) + } + + partitionPath, partitionToken, err := parseNativeEqualityCondition(partitionExpr, expressionAttributeNames) + if err != nil { + return nil, err + } + if partitionPath != index.PartitionKey { + return nil, fmt.Errorf("dynamodb: unsupported key condition partition key %q", partitionPath) + } + partitionValue, ok := expressionAttributeValues[partitionToken] + if !ok { + return nil, fmt.Errorf("dynamodb: unresolved expression attribute value %q", partitionToken) + } + + var sortPredicate func(dynamodbdomain.Item) (bool, error) + if strings.TrimSpace(sortExpr) != "" { + sortPredicate, err = buildNativeSortPredicate(index.SortKey, sortExpr, expressionAttributeNames, expressionAttributeValues) + if err != nil { + return nil, err + } + } + + return func(item dynamodbdomain.Item) (bool, error) { + value, ok := item.Attributes[index.PartitionKey] + if !ok || !attributeValueEqual(value, partitionValue) { + return false, nil + } + if sortPredicate == nil { + return true, nil + } + return sortPredicate(item) + }, nil +} + +func splitNativeKeyCondition(expression string) (string, string) { + upper := strings.ToUpper(expression) + idx := strings.Index(upper, " AND ") + if idx < 0 { + return strings.TrimSpace(expression), "" + } + return strings.TrimSpace(expression[:idx]), strings.TrimSpace(expression[idx+5:]) +} + +func parseNativeEqualityCondition(expression string, expressionAttributeNames map[string]string) (string, string, error) { + parts := strings.SplitN(expression, "=", 2) + if len(parts) != 2 { + return "", "", fmt.Errorf("dynamodb: unsupported key condition expression %q", expression) + } + path, err := resolveNativeConditionPath(parts[0], expressionAttributeNames) + if err != nil { + return "", "", err + } + token := strings.TrimSpace(parts[1]) + if !strings.HasPrefix(token, ":") { + return "", "", fmt.Errorf("dynamodb: unresolved expression attribute value %q", token) + } + return path, token, nil +} + +func buildNativeSortPredicate(sortKey, expression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]dynamodbdomain.AttributeValue) (func(dynamodbdomain.Item) (bool, error), error) { + if strings.TrimSpace(sortKey) == "" { + return nil, fmt.Errorf("dynamodb: sort key conditions are not supported for this index") + } + + for _, op := range []string{"<=", ">=", "<", ">", "="} { + if strings.Contains(expression, op) { + parts := strings.SplitN(expression, op, 2) + path, err := resolveNativeConditionPath(parts[0], expressionAttributeNames) if err != nil { - return "", false, err + return nil, err } - if strings.TrimSpace(sk) != "" { - return id + "|" + sk, true, nil + if path != sortKey { + return nil, fmt.Errorf("dynamodb: unsupported key condition sort key %q", path) } + value, err := resolveNativeConditionValue(parts[1], expressionAttributeValues) + if err != nil { + return nil, err + } + return func(item dynamodbdomain.Item) (bool, error) { + current, ok := item.Attributes[sortKey] + if !ok { + return false, nil + } + cmp := compareNativeAttributeValues(current, value) + switch op { + case "=": + return attributeValueEqual(current, value), nil + case "<": + return cmp < 0, nil + case "<=": + return cmp <= 0, nil + case ">": + return cmp > 0, nil + case ">=": + return cmp >= 0, nil + default: + return false, fmt.Errorf("dynamodb: unsupported sort operator %q", op) + } + }, nil } - return id, true, nil } - if len(values) == 1 { - for _, value := range values { - key, err := attributeValueToString(value) - return key, true, err + return nil, fmt.Errorf("dynamodb: unsupported sort key condition %q", expression) +} + +func compareNativeIndexItems(left, right dynamodbdomain.Item, index nativeDynamoIndex) int { + if strings.TrimSpace(index.SortKey) != "" { + leftSort, leftOK := left.Attributes[index.SortKey] + rightSort, rightOK := right.Attributes[index.SortKey] + if leftOK && rightOK { + if cmp := compareNativeAttributeValues(leftSort, rightSort); cmp != 0 { + return cmp + } + } + if leftOK != rightOK { + if leftOK { + return -1 + } + return 1 } } + return strings.Compare(left.Key, right.Key) +} - keys := make([]string, 0, len(values)) - for key := range values { - keys = append(keys, key) +func compareNativeAttributeValues(left, right dynamodbdomain.AttributeValue) int { + switch { + case left.S != nil && right.S != nil: + return strings.Compare(*left.S, *right.S) + case left.N != nil && right.N != nil: + leftValue, _ := strconv.ParseFloat(strings.TrimSpace(*left.N), 64) + rightValue, _ := strconv.ParseFloat(strings.TrimSpace(*right.N), 64) + switch { + case leftValue < rightValue: + return -1 + case leftValue > rightValue: + return 1 + default: + return 0 + } + case left.BOOL != nil && right.BOOL != nil: + switch { + case !*left.BOOL && *right.BOOL: + return -1 + case *left.BOOL && !*right.BOOL: + return 1 + default: + return 0 + } + default: + return 0 + } +} + +func cloneSecondaryIndexDefinitions(source []dynamoSecondaryIndexDefinition) []dynamoSecondaryIndexDefinition { + if len(source) == 0 { + return nil } - return "", false, fmt.Errorf("dynamodb: unsupported key attributes %q", strings.Join(keys, ", ")) + cloned := make([]dynamoSecondaryIndexDefinition, len(source)) + copy(cloned, source) + return cloned +} + +func attributeValueEqual(left, right dynamodbdomain.AttributeValue) bool { + switch { + case left.S != nil && right.S != nil: + return *left.S == *right.S + case left.N != nil && right.N != nil: + return *left.N == *right.N + case left.BOOL != nil && right.BOOL != nil: + return *left.BOOL == *right.BOOL + case left.NULL && right.NULL: + return true + case left.M != nil && right.M != nil: + return attributeMapEqual(*left.M, *right.M) + case left.L != nil && right.L != nil: + return attributeListEqual(*left.L, *right.L) + default: + return false + } +} + +func attributeMapEqual(left, right map[string]dynamodbdomain.AttributeValue) bool { + if len(left) != len(right) { + return false + } + for name, value := range left { + other, ok := right[name] + if !ok || !attributeValueEqual(value, other) { + return false + } + } + return true +} + +func attributeListEqual(left, right []dynamodbdomain.AttributeValue) bool { + if len(left) != len(right) { + return false + } + for i := range left { + if !attributeValueEqual(left[i], right[i]) { + return false + } + } + return true } func attributeValueToString(value dynamoAttributeValue) (string, error) { diff --git a/core/internal/delivery/http/dynamodb_native_test.go b/core/internal/delivery/http/dynamodb_native_test.go index 62cac50..f6216c5 100644 --- a/core/internal/delivery/http/dynamodb_native_test.go +++ b/core/internal/delivery/http/dynamodb_native_test.go @@ -11,6 +11,7 @@ import ( "github.com/gin-gonic/gin" "github.com/michasdev/mildstack/core/internal/resources/awscontext" "github.com/michasdev/mildstack/core/internal/resources/dynamodb/application" + sqsapplication "github.com/michasdev/mildstack/core/internal/resources/sqs/application" ) func TestDynamoDBTargetRegistryDistinguishesSupportedAndUnsupportedOperations(t *testing.T) { @@ -486,7 +487,7 @@ func TestDynamoDBNativeRoutesReturnAWSShapedErrors(t *testing.T) { RegisterDynamoDBNativeRoutes(engine, application.New()) malformed := doDynamoDBRequest(t, engine, dynamoRequest{ - target: "ListTables", + target: "DynamoDB_20120810.", body: `{}`, }) assertDynamoError(t, malformed, http.StatusBadRequest, "ValidationException") @@ -502,7 +503,7 @@ func TestDynamoDBNativeRoutesReturnAWSShapedErrors(t *testing.T) { } }`, }) - assertDynamoError(t, unsupportedQuery, http.StatusBadRequest, "ValidationException") + assertDynamoError(t, unsupportedQuery, http.StatusBadRequest, "ResourceNotFoundException") unsupportedScan := doDynamoDBRequest(t, engine, dynamoRequest{ target: "DynamoDB_20120810.Scan", @@ -561,7 +562,7 @@ func TestDynamoDBNativeRoutesReturnAWSShapedErrors(t *testing.T) { target: "DynamoDB_20120810.DeleteItem", body: `{ "TableName":"mildstack-records", - "Key":{"id":{"S":"missing"}} + "Key":{"id":{"S":"missing"},"version":{"N":"1"}} }`, }) assertDynamoError(t, missingItem, http.StatusBadRequest, "ResourceNotFoundException") @@ -570,7 +571,7 @@ func TestDynamoDBNativeRoutesReturnAWSShapedErrors(t *testing.T) { target: "DynamoDB_20120810.UpdateItem", body: `{ "TableName":"mildstack-records", - "Key":{"id":{"S":"missing"}}, + "Key":{"id":{"S":"missing"},"version":{"N":"1"}}, "UpdateExpression":"SET title = :title", "ConditionExpression":"attribute_exists(id)", "ExpressionAttributeValues":{ @@ -606,6 +607,56 @@ func TestDynamoDBNativeRoutesReturnAWSShapedErrors(t *testing.T) { assertDynamoError(t, missingDelete, http.StatusBadRequest, "ResourceNotFoundException") } +func TestDynamoDBSQSRoutingIsolation(t *testing.T) { + t.Helper() + + gin.SetMode(gin.TestMode) + + engine := gin.New() + RegisterDynamoDBNativeRoutes(engine, application.New()) + RegisterSQSNativeRoutes(engine, sqsapplication.New()) + engine.POST("/", func(c *gin.Context) { + c.Status(http.StatusNoContent) + }) + + sqsResponse := doDynamoDBRequest(t, engine, dynamoRequest{ + target: "AmazonSQS.ListQueues", + body: `{}`, + }) + if got, want := sqsResponse.code, http.StatusOK; got != want { + t.Fatalf("unexpected sqs status: got %d want %d\nbody: %s", got, want, sqsResponse.body) + } + if strings.Contains(sqsResponse.body, "ValidationException") { + t.Fatalf("expected sqs request to bypass dynamodb validation error, got %q", sqsResponse.body) + } + if !strings.Contains(sqsResponse.body, "\"QueueUrls\"") { + t.Fatalf("expected sqs response body, got %q", sqsResponse.body) + } + + dynamoResponse := doDynamoDBRequest(t, engine, dynamoRequest{ + target: "DynamoDB_20120810.ListTables", + body: `{}`, + }) + if got, want := dynamoResponse.code, http.StatusOK; got != want { + t.Fatalf("unexpected dynamodb status: got %d want %d\nbody: %s", got, want, dynamoResponse.body) + } + if !strings.Contains(dynamoResponse.body, "\"TableNames\"") { + t.Fatalf("expected dynamodb response body, got %q", dynamoResponse.body) + } + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{}`)) + request.Header.Set("Content-Type", dynamoDBJSONContentType) + engine.ServeHTTP(recorder, request) + + if got, want := recorder.Code, http.StatusBadRequest; got != want { + t.Fatalf("unexpected missing-target status: got %d want %d\nbody: %s", got, want, recorder.Body.String()) + } + if strings.Contains(recorder.Body.String(), "ValidationException") { + t.Fatalf("expected missing target to bypass dynamodb validation error, got %q", recorder.Body.String()) + } +} + func TestDynamoDBNativeRoutesHandleQueryAndScanSubset(t *testing.T) { t.Helper() @@ -774,6 +825,198 @@ func TestDynamoDBNativeRoutesHandleQueryAndScanSubset(t *testing.T) { } } +func TestDynamoDBNativeRoutesSupportIndexedQueryAndProjection(t *testing.T) { + t.Helper() + + gin.SetMode(gin.TestMode) + engine := gin.New() + RegisterDynamoDBNativeRoutes(engine, application.New()) + + createTable := doDynamoDBRequest(t, engine, dynamoRequest{ + target: "DynamoDB_20120810.CreateTable", + body: `{ + "TableName":"mildstack-indexed", + "KeySchema":[ + {"AttributeName":"pk","KeyType":"HASH"}, + {"AttributeName":"sk","KeyType":"RANGE"} + ], + "AttributeDefinitions":[ + {"AttributeName":"pk","AttributeType":"S"}, + {"AttributeName":"sk","AttributeType":"S"}, + {"AttributeName":"gsi_pk","AttributeType":"S"}, + {"AttributeName":"gsi_sk","AttributeType":"S"}, + {"AttributeName":"title","AttributeType":"S"} + ], + "GlobalSecondaryIndexes":[ + { + "IndexName":"gsi-title", + "KeySchema":[ + {"AttributeName":"gsi_pk","KeyType":"HASH"}, + {"AttributeName":"gsi_sk","KeyType":"RANGE"} + ], + "Projection":{ + "ProjectionType":"INCLUDE", + "NonKeyAttributes":["title"] + } + } + ] + }`, + }) + if got, want := createTable.code, http.StatusOK; got != want { + t.Fatalf("unexpected create table status: got %d want %d", got, want) + } + + for _, item := range []struct { + sk string + gsiSK string + title string + }{ + {sk: "001", gsiSK: "001", title: "indexed-one"}, + {sk: "002", gsiSK: "002", title: "indexed-two"}, + } { + response := doDynamoDBRequest(t, engine, dynamoRequest{ + target: "DynamoDB_20120810.PutItem", + body: `{ + "TableName":"mildstack-indexed", + "Item":{ + "pk":{"S":"series#1"}, + "sk":{"S":"` + item.sk + `"}, + "gsi_pk":{"S":"group#1"}, + "gsi_sk":{"S":"` + item.gsiSK + `"}, + "title":{"S":"` + item.title + `"} + } + }`, + }) + if got, want := response.code, http.StatusOK; got != want { + t.Fatalf("unexpected put item status: got %d want %d", got, want) + } + } + + queryPage1 := doDynamoDBRequest(t, engine, dynamoRequest{ + target: "DynamoDB_20120810.Query", + body: `{ + "TableName":"mildstack-indexed", + "IndexName":"gsi-title", + "KeyConditionExpression":"gsi_pk = :pk AND gsi_sk BETWEEN :start AND :end", + "ProjectionExpression":"gsi_pk, title", + "ExpressionAttributeValues":{ + ":pk":{"S":"group#1"}, + ":start":{"S":"001"}, + ":end":{"S":"002"} + }, + "Limit":1 + }`, + }) + if got, want := queryPage1.code, http.StatusOK; got != want { + t.Fatalf("unexpected indexed query status: got %d want %d", got, want) + } + var queryResponse queryResponse + decodeResponse(t, queryPage1.body, &queryResponse) + if got, want := queryResponse.Count, 1; got != want { + t.Fatalf("unexpected indexed query count: got %d want %d", got, want) + } + if got, want := queryResponse.Items[0]["title"].S, "indexed-one"; got != want { + t.Fatalf("unexpected indexed query title: got %q want %q", got, want) + } + if _, ok := queryResponse.Items[0]["gsi_sk"]; ok { + t.Fatal("expected projected gsi sort key to be omitted from query item") + } + + queryPage2 := doDynamoDBRequest(t, engine, dynamoRequest{ + target: "DynamoDB_20120810.Query", + body: `{ + "TableName":"mildstack-indexed", + "IndexName":"gsi-title", + "KeyConditionExpression":"gsi_pk = :pk AND gsi_sk BETWEEN :start AND :end", + "ProjectionExpression":"gsi_pk, title", + "ExpressionAttributeValues":{ + ":pk":{"S":"group#1"}, + ":start":{"S":"001"}, + ":end":{"S":"002"} + }, + "Limit":1, + "ExclusiveStartKey":{ + "gsi_pk":{"S":"group#1"}, + "gsi_sk":{"S":"001"}, + "pk":{"S":"series#1"}, + "sk":{"S":"001"} + } + }`, + }) + if got, want := queryPage2.code, http.StatusOK; got != want { + t.Fatalf("unexpected indexed query page 2 status: got %d want %d", got, want) + } + decodeResponse(t, queryPage2.body, &queryResponse) + if got, want := queryResponse.Count, 1; got != want { + t.Fatalf("unexpected indexed query page 2 count: got %d want %d", got, want) + } + if got, want := queryResponse.Items[0]["title"].S, "indexed-two"; got != want { + t.Fatalf("unexpected indexed query page 2 title: got %q want %q", got, want) + } +} + +func TestDynamoDBNativeRoutesHonorCustomTableKeyNames(t *testing.T) { + t.Helper() + + gin.SetMode(gin.TestMode) + engine := gin.New() + RegisterDynamoDBNativeRoutes(engine, application.New()) + + createTable := doDynamoDBRequest(t, engine, dynamoRequest{ + target: "DynamoDB_20120810.CreateTable", + body: `{ + "TableName":"mildstack-custom-keys", + "KeySchema":[ + {"AttributeName":"pk","KeyType":"HASH"}, + {"AttributeName":"sk","KeyType":"RANGE"} + ], + "AttributeDefinitions":[ + {"AttributeName":"pk","AttributeType":"S"}, + {"AttributeName":"sk","AttributeType":"S"} + ], + "BillingMode":"PAY_PER_REQUEST" + }`, + }) + if got, want := createTable.code, http.StatusOK; got != want { + t.Fatalf("unexpected create table status: got %d want %d", got, want) + } + + putItem := doDynamoDBRequest(t, engine, dynamoRequest{ + target: "DynamoDB_20120810.PutItem", + body: `{ + "TableName":"mildstack-custom-keys", + "Item":{ + "pk":{"S":"account#1"}, + "sk":{"S":"meta"}, + "title":{"S":"custom schema"} + } + }`, + }) + if got, want := putItem.code, http.StatusOK; got != want { + t.Fatalf("unexpected put item status: got %d want %d\nbody: %s", got, want, putItem.body) + } + + getItem := doDynamoDBRequest(t, engine, dynamoRequest{ + target: "DynamoDB_20120810.GetItem", + body: `{ + "TableName":"mildstack-custom-keys", + "Key":{ + "pk":{"S":"account#1"}, + "sk":{"S":"meta"} + } + }`, + }) + if got, want := getItem.code, http.StatusOK; got != want { + t.Fatalf("unexpected get item status: got %d want %d\nbody: %s", got, want, getItem.body) + } + + var response getItemResponse + decodeResponse(t, getItem.body, &response) + if got, want := response.Item["title"].S, "custom schema"; got != want { + t.Fatalf("unexpected fetched title: got %q want %q", got, want) + } +} + type dynamoRequest struct { target string body string diff --git a/core/internal/delivery/http/s3_native.go b/core/internal/delivery/http/s3_native.go index e8ac8d8..a792b60 100644 --- a/core/internal/delivery/http/s3_native.go +++ b/core/internal/delivery/http/s3_native.go @@ -4,6 +4,7 @@ import ( "encoding/xml" "io" "net/http" + "net/url" "strconv" "strings" "time" @@ -25,6 +26,28 @@ type S3NativeService interface { DeleteObject(bucket, key string) error } +type s3NativeMetadataWriter interface { + PutObjectWithMetadata(bucket, key string, body io.Reader, contentType string, metadata, preservedHeaders map[string]string) (s3domain.Object, error) +} + +type s3NativeListObjectsV2Service interface { + ListObjectsV2(request s3domain.ListObjectsV2Request) (s3domain.ListObjectsV2Result, error) +} + +type s3NativeDeleteObjectsService interface { + DeleteObjects(request s3domain.DeleteObjectsRequest) (s3domain.DeleteObjectsResult, error) +} + +type s3NativeCopyObjectService interface { + CopyObject(bucket, key, sourceBucket, sourceKey string) (s3domain.Object, error) +} + +type s3NativeMultipartService interface { + CreateMultipartUpload(bucket, key, contentType string, metadata, preservedHeaders map[string]string) (s3domain.MultipartUpload, error) + UploadPart(uploadID string, partNumber int, body []byte) (s3domain.MultipartPart, error) + CompleteMultipartUpload(uploadID string) (s3domain.Object, error) +} + const s3XMLNamespace = "http://s3.amazonaws.com/doc/2006-03-01/" func RegisterS3NativeRoutes(engine *gin.Engine, service S3NativeService) { @@ -55,6 +78,7 @@ func (h s3NativeHandler) dispatch(c *gin.Context) bool { if path == "" || strings.HasPrefix(path, "/api/") { return false } + query := c.Request.URL.Query() trimmed := strings.Trim(path, "/") segments := []string{} @@ -74,6 +98,11 @@ func (h s3NativeHandler) dispatch(c *gin.Context) bool { case http.MethodPut: h.createBucket(c, bucket) return true + case http.MethodPost: + if hasS3QueryParam(query, "delete") { + h.deleteObjects(c, bucket) + return true + } case http.MethodHead: h.headBucket(c, bucket) return true @@ -81,6 +110,10 @@ func (h s3NativeHandler) dispatch(c *gin.Context) bool { h.deleteBucket(c, bucket) return true case http.MethodGet: + if strings.TrimSpace(query.Get("list-type")) == "2" { + h.listObjectsV2(c, bucket) + return true + } h.listObjects(c, bucket) return true } @@ -88,10 +121,27 @@ func (h s3NativeHandler) dispatch(c *gin.Context) bool { bucket := segments[0] key := strings.Join(segments[1:], "/") switch c.Request.Method { + case http.MethodPost: + switch { + case hasS3QueryParam(query, "uploads"): + h.createMultipartUpload(c, bucket, key) + return true + case strings.TrimSpace(query.Get("uploadId")) != "": + h.completeMultipartUpload(c, bucket, key) + return true + } case http.MethodGet: h.getObject(c, bucket, key) return true case http.MethodPut: + switch { + case strings.TrimSpace(query.Get("uploadId")) != "" && strings.TrimSpace(query.Get("partNumber")) != "": + h.uploadPart(c, bucket, key) + return true + case strings.TrimSpace(c.GetHeader("x-amz-copy-source")) != "": + h.copyObject(c, bucket, key) + return true + } h.putObject(c, bucket, key) return true case http.MethodHead: @@ -139,6 +189,54 @@ func (h s3NativeHandler) listObjects(c *gin.Context, bucketName string) { }) } +func (h s3NativeHandler) listObjectsV2(c *gin.Context, bucketName string) { + service, ok := h.service.(s3NativeListObjectsV2Service) + if !ok { + writeS3Error(c, io.ErrUnexpectedEOF) + return + } + + maxKeys := 0 + if raw := strings.TrimSpace(c.Query("max-keys")); raw != "" { + parsed, err := strconv.Atoi(raw) + if err != nil { + writeS3Error(c, err) + return + } + maxKeys = parsed + } + + result, err := service.ListObjectsV2(s3domain.ListObjectsV2Request{ + Bucket: bucketName, + Prefix: strings.TrimSpace(c.Query("prefix")), + Delimiter: strings.TrimSpace(c.Query("delimiter")), + ContinuationToken: strings.TrimSpace(c.Query("continuation-token")), + StartAfter: strings.TrimSpace(c.Query("start-after")), + MaxKeys: maxKeys, + }) + if err != nil { + writeS3Error(c, err) + return + } + + c.Header("Content-Type", "application/xml") + c.XML(http.StatusOK, listObjectsV2Result{ + XMLName: xml.Name{Local: "ListBucketResult"}, + XMLNS: s3XMLNamespace, + Name: result.Bucket, + Prefix: result.Prefix, + Delimiter: result.Delimiter, + MaxKeys: result.MaxKeys, + KeyCount: result.KeyCount, + IsTruncated: result.IsTruncated, + ContinuationToken: result.ContinuationToken, + NextContinuationToken: result.NextContinuationToken, + StartAfter: result.StartAfter, + Contents: listObjectEntriesFromDomain(result.Objects), + CommonPrefixes: commonPrefixEntries(result.CommonPrefixes), + }) +} + func (h s3NativeHandler) createBucket(c *gin.Context, bucketName string) { region := strings.TrimSpace(c.GetHeader("x-amz-bucket-region")) if region == "" { @@ -199,8 +297,18 @@ func (h s3NativeHandler) putObject(c *gin.Context, bucketName, objectKey string) if contentType == "" { contentType = "application/octet-stream" } - - object, err := h.service.PutObject(bucketName, objectKey, c.Request.Body, contentType) + metadata := metadataFromHeaders(c.Request.Header) + preservedHeaders := preservedObjectHeaders(c.Request.Header) + + var ( + object s3domain.Object + err error + ) + if writer, ok := h.service.(s3NativeMetadataWriter); ok { + object, err = writer.PutObjectWithMetadata(bucketName, objectKey, c.Request.Body, contentType, metadata, preservedHeaders) + } else { + object, err = h.service.PutObject(bucketName, objectKey, c.Request.Body, contentType) + } if err != nil { writeS3Error(c, err) return @@ -210,6 +318,34 @@ func (h s3NativeHandler) putObject(c *gin.Context, bucketName, objectKey string) c.Status(http.StatusOK) } +func (h s3NativeHandler) copyObject(c *gin.Context, bucketName, objectKey string) { + service, ok := h.service.(s3NativeCopyObjectService) + if !ok { + writeS3Error(c, io.ErrUnexpectedEOF) + return + } + + sourceBucket, sourceKey, err := parseCopySource(c.GetHeader("x-amz-copy-source")) + if err != nil { + writeS3Error(c, err) + return + } + + object, err := service.CopyObject(bucketName, objectKey, sourceBucket, sourceKey) + if err != nil { + writeS3Error(c, err) + return + } + + c.Header("Content-Type", "application/xml") + c.XML(http.StatusOK, copyObjectResult{ + XMLName: xml.Name{Local: "CopyObjectResult"}, + XMLNS: s3XMLNamespace, + LastModified: object.LastModified.UTC().Format(time.RFC3339), + ETag: object.ETag, + }) +} + func (h s3NativeHandler) deleteObject(c *gin.Context, bucketName, objectKey string) { if err := h.service.DeleteObject(bucketName, objectKey); err != nil { writeS3Error(c, err) @@ -218,6 +354,135 @@ func (h s3NativeHandler) deleteObject(c *gin.Context, bucketName, objectKey stri c.Status(http.StatusNoContent) } +func (h s3NativeHandler) deleteObjects(c *gin.Context, bucketName string) { + service, ok := h.service.(s3NativeDeleteObjectsService) + if !ok { + writeS3Error(c, io.ErrUnexpectedEOF) + return + } + + body, err := io.ReadAll(c.Request.Body) + if err != nil { + writeS3Error(c, err) + return + } + + var payload deleteObjectsRequest + if err := xml.Unmarshal(body, &payload); err != nil { + writeS3Error(c, err) + return + } + + keys := make([]string, 0, len(payload.Objects)) + for _, object := range payload.Objects { + key := strings.TrimSpace(object.Key) + if key == "" { + continue + } + keys = append(keys, key) + } + + result, err := service.DeleteObjects(s3domain.DeleteObjectsRequest{ + Bucket: bucketName, + Keys: keys, + Quiet: payload.Quiet, + }) + if err != nil { + writeS3Error(c, err) + return + } + + c.Header("Content-Type", "application/xml") + c.XML(http.StatusOK, deleteObjectsResult{ + XMLName: xml.Name{Local: "DeleteResult"}, + XMLNS: s3XMLNamespace, + Deleted: deletedObjectEntries(result.Deleted), + Errors: deleteObjectErrorEntries(result.Errors), + }) +} + +func (h s3NativeHandler) createMultipartUpload(c *gin.Context, bucketName, objectKey string) { + service, ok := h.service.(s3NativeMultipartService) + if !ok { + writeS3Error(c, io.ErrUnexpectedEOF) + return + } + + upload, err := service.CreateMultipartUpload( + bucketName, + objectKey, + strings.TrimSpace(c.GetHeader("Content-Type")), + metadataFromHeaders(c.Request.Header), + preservedObjectHeaders(c.Request.Header), + ) + if err != nil { + writeS3Error(c, err) + return + } + + c.Header("Content-Type", "application/xml") + c.XML(http.StatusOK, createMultipartUploadResult{ + XMLName: xml.Name{Local: "InitiateMultipartUploadResult"}, + XMLNS: s3XMLNamespace, + Bucket: upload.Bucket, + Key: upload.Key, + UploadID: upload.UploadID, + }) +} + +func (h s3NativeHandler) uploadPart(c *gin.Context, _, _ string) { + service, ok := h.service.(s3NativeMultipartService) + if !ok { + writeS3Error(c, io.ErrUnexpectedEOF) + return + } + + partNumber, err := strconv.Atoi(strings.TrimSpace(c.Query("partNumber"))) + if err != nil { + writeS3Error(c, err) + return + } + body, err := io.ReadAll(c.Request.Body) + if err != nil { + writeS3Error(c, err) + return + } + + part, err := service.UploadPart(strings.TrimSpace(c.Query("uploadId")), partNumber, body) + if err != nil { + writeS3Error(c, err) + return + } + + c.Header("ETag", part.ETag) + c.Status(http.StatusOK) +} + +func (h s3NativeHandler) completeMultipartUpload(c *gin.Context, bucketName, objectKey string) { + service, ok := h.service.(s3NativeMultipartService) + if !ok { + writeS3Error(c, io.ErrUnexpectedEOF) + return + } + + object, err := service.CompleteMultipartUpload(strings.TrimSpace(c.Query("uploadId"))) + if err != nil { + writeS3Error(c, err) + return + } + + c.Header("Content-Type", "application/xml") + c.XML(http.StatusOK, completeMultipartUploadResult{ + XMLName: xml.Name{Local: "CompleteMultipartUploadResult"}, + XMLNS: s3XMLNamespace, + Location: "/" + bucketName + "/" + objectKey, + Bucket: bucketName, + Key: objectKey, + ETag: object.ETag, + LastModified: object.LastModified.UTC().Format(time.RFC3339), + }) +} + type listBucketsResult struct { XMLName xml.Name `xml:"ListAllMyBucketsResult"` XMLNS string `xml:"xmlns,attr"` @@ -246,6 +511,22 @@ type listObjectsResult struct { Contents []listObjectEntry `xml:"Contents"` } +type listObjectsV2Result struct { + XMLName xml.Name `xml:"ListBucketResult"` + XMLNS string `xml:"xmlns,attr"` + Name string `xml:"Name"` + Prefix string `xml:"Prefix,omitempty"` + Delimiter string `xml:"Delimiter,omitempty"` + MaxKeys int `xml:"MaxKeys"` + KeyCount int `xml:"KeyCount"` + IsTruncated bool `xml:"IsTruncated"` + ContinuationToken string `xml:"ContinuationToken,omitempty"` + NextContinuationToken string `xml:"NextContinuationToken,omitempty"` + StartAfter string `xml:"StartAfter,omitempty"` + Contents []listObjectEntry `xml:"Contents"` + CommonPrefixes []commonPrefixEntry `xml:"CommonPrefixes,omitempty"` +} + type listObjectEntry struct { Key string `xml:"Key"` LastModified string `xml:"LastModified"` @@ -254,6 +535,61 @@ type listObjectEntry struct { StorageClass string `xml:"StorageClass"` } +type commonPrefixEntry struct { + Prefix string `xml:"Prefix"` +} + +type deleteObjectsRequest struct { + Quiet bool `xml:"Quiet"` + Objects []deleteObjectsRequestItem `xml:"Object"` +} + +type deleteObjectsRequestItem struct { + Key string `xml:"Key"` +} + +type deleteObjectsResult struct { + XMLName xml.Name `xml:"DeleteResult"` + XMLNS string `xml:"xmlns,attr"` + Deleted []deletedObjectEntry `xml:"Deleted,omitempty"` + Errors []deleteObjectErrorEntry `xml:"Error,omitempty"` +} + +type deletedObjectEntry struct { + Key string `xml:"Key"` +} + +type deleteObjectErrorEntry struct { + Key string `xml:"Key"` + Code string `xml:"Code"` + Message string `xml:"Message"` +} + +type copyObjectResult struct { + XMLName xml.Name `xml:"CopyObjectResult"` + XMLNS string `xml:"xmlns,attr"` + LastModified string `xml:"LastModified"` + ETag string `xml:"ETag"` +} + +type createMultipartUploadResult struct { + XMLName xml.Name `xml:"InitiateMultipartUploadResult"` + XMLNS string `xml:"xmlns,attr"` + Bucket string `xml:"Bucket"` + Key string `xml:"Key"` + UploadID string `xml:"UploadId"` +} + +type completeMultipartUploadResult struct { + XMLName xml.Name `xml:"CompleteMultipartUploadResult"` + XMLNS string `xml:"xmlns,attr"` + Location string `xml:"Location,omitempty"` + Bucket string `xml:"Bucket"` + Key string `xml:"Key"` + ETag string `xml:"ETag"` + LastModified string `xml:"LastModified,omitempty"` +} + func bucketEntriesFromDomain(buckets []s3domain.Bucket) []bucketEntry { entries := make([]bucketEntry, len(buckets)) for i, bucket := range buckets { @@ -279,6 +615,34 @@ func listObjectEntriesFromDomain(objects []s3domain.Object) []listObjectEntry { return entries } +func commonPrefixEntries(prefixes []string) []commonPrefixEntry { + entries := make([]commonPrefixEntry, len(prefixes)) + for i, prefix := range prefixes { + entries[i] = commonPrefixEntry{Prefix: prefix} + } + return entries +} + +func deletedObjectEntries(objects []s3domain.DeletedObject) []deletedObjectEntry { + entries := make([]deletedObjectEntry, len(objects)) + for i, object := range objects { + entries[i] = deletedObjectEntry{Key: object.Key} + } + return entries +} + +func deleteObjectErrorEntries(errors []s3domain.DeleteObjectsError) []deleteObjectErrorEntry { + entries := make([]deleteObjectErrorEntry, len(errors)) + for i, objectErr := range errors { + entries[i] = deleteObjectErrorEntry{ + Key: objectErr.Key, + Code: objectErr.Code, + Message: objectErr.Message, + } + } + return entries +} + type s3ErrorResponse struct { XMLName xml.Name `xml:"Error"` Code string `xml:"Code"` @@ -306,6 +670,8 @@ func mapS3Error(err error) (int, string, string) { return http.StatusNotFound, "NoSuchBucket", message case strings.Contains(message, "NoSuchKey"): return http.StatusNotFound, "NoSuchKey", message + case strings.Contains(message, "NoSuchUpload"): + return http.StatusNotFound, "NoSuchUpload", message case strings.Contains(message, "BucketNotEmpty"): return http.StatusConflict, "BucketNotEmpty", message case strings.Contains(message, "InvalidBucketName"): @@ -325,6 +691,18 @@ func writeObjectResponse(c *gin.Context, object s3domain.Object, includeBody boo c.Header("Content-Type", object.ContentType) c.Header("Content-Length", formatContentLength(object.Size)) c.Header("Accept-Ranges", "bytes") + for key, value := range object.PreservedHeaders { + if strings.TrimSpace(key) == "" { + continue + } + c.Header(key, value) + } + for key, value := range object.Metadata { + if strings.TrimSpace(key) == "" { + continue + } + c.Header("x-amz-meta-"+strings.ToLower(key), value) + } if !includeBody { c.Status(http.StatusOK) @@ -340,3 +718,58 @@ func formatContentLength(size int64) string { } return strconv.FormatInt(size, 10) } + +func hasS3QueryParam(values url.Values, key string) bool { + if values == nil { + return false + } + _, ok := values[key] + return ok +} + +func metadataFromHeaders(headers http.Header) map[string]string { + metadata := map[string]string{} + for key, values := range headers { + lowerKey := strings.ToLower(strings.TrimSpace(key)) + if !strings.HasPrefix(lowerKey, "x-amz-meta-") { + continue + } + metadata[strings.TrimPrefix(lowerKey, "x-amz-meta-")] = strings.Join(values, ",") + } + if len(metadata) == 0 { + return nil + } + return metadata +} + +func preservedObjectHeaders(headers http.Header) map[string]string { + preserved := map[string]string{} + for _, key := range []string{"Cache-Control", "Content-Disposition", "Content-Encoding", "Content-Language", "Expires"} { + value := strings.TrimSpace(headers.Get(key)) + if value == "" { + continue + } + preserved[key] = value + } + if len(preserved) == 0 { + return nil + } + return preserved +} + +func parseCopySource(raw string) (string, string, error) { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return "", "", io.ErrUnexpectedEOF + } + decoded, err := url.PathUnescape(trimmed) + if err != nil { + return "", "", err + } + decoded = strings.TrimPrefix(decoded, "/") + parts := strings.SplitN(decoded, "/", 2) + if len(parts) != 2 || strings.TrimSpace(parts[0]) == "" || strings.TrimSpace(parts[1]) == "" { + return "", "", io.ErrUnexpectedEOF + } + return parts[0], parts[1], nil +} diff --git a/core/internal/delivery/http/s3_native_test.go b/core/internal/delivery/http/s3_native_test.go index b668b5b..3d1d4c1 100644 --- a/core/internal/delivery/http/s3_native_test.go +++ b/core/internal/delivery/http/s3_native_test.go @@ -6,10 +6,12 @@ import ( "io" "net/http" "net/http/httptest" + "strings" "testing" "github.com/gin-gonic/gin" "github.com/michasdev/mildstack/core/internal/resources/awscontext" + s3application "github.com/michasdev/mildstack/core/internal/resources/s3/application" s3domain "github.com/michasdev/mildstack/core/internal/resources/s3/domain" ) @@ -142,3 +144,103 @@ func TestS3NativeListBucketsUsesSharedAWSAccountID(t *testing.T) { t.Fatalf("unexpected owner id: got %q want %q", got, want) } } + +func TestS3NativeGetObjectReturnsStoredMetadataHeaders(t *testing.T) { + t.Helper() + + gin.SetMode(gin.TestMode) + service := s3application.New() + if _, err := service.CreateBucket("metadata-bucket", "us-east-1"); err != nil { + t.Fatalf("create bucket: %v", err) + } + + engine := gin.New() + RegisterS3NativeRoutes(engine, service) + + putRequest := httptest.NewRequest(http.MethodPut, "/metadata-bucket/notes.txt", strings.NewReader("hello")) + putRequest.Header.Set("Content-Type", "text/plain") + putRequest.Header.Set("X-Amz-Meta-Custom-Author", "bot") + putRecorder := httptest.NewRecorder() + engine.ServeHTTP(putRecorder, putRequest) + + if got, want := putRecorder.Code, http.StatusOK; got != want { + t.Fatalf("unexpected put status: got %d want %d", got, want) + } + + getRequest := httptest.NewRequest(http.MethodGet, "/metadata-bucket/notes.txt", nil) + getRecorder := httptest.NewRecorder() + engine.ServeHTTP(getRecorder, getRequest) + + if got, want := getRecorder.Code, http.StatusOK; got != want { + t.Fatalf("unexpected get status: got %d want %d", got, want) + } + if got, want := getRecorder.Header().Get("x-amz-meta-custom-author"), "bot"; got != want { + t.Fatalf("unexpected metadata header: got %q want %q", got, want) + } + if got, want := getRecorder.Body.String(), "hello"; got != want { + t.Fatalf("unexpected body: got %q want %q", got, want) + } +} + +func TestS3NativeListObjectsV2AndDeleteObjectsReturnAWSXML(t *testing.T) { + t.Helper() + + gin.SetMode(gin.TestMode) + service := s3application.New() + if _, err := service.CreateBucket("listing-bucket", "us-east-1"); err != nil { + t.Fatalf("create bucket: %v", err) + } + for _, key := range []string{"listing/a.txt", "listing/b.txt", "listing/c.txt"} { + if _, err := service.PutObject("listing-bucket", key, strings.NewReader(key), "text/plain"); err != nil { + t.Fatalf("put object %q: %v", key, err) + } + } + + engine := gin.New() + RegisterS3NativeRoutes(engine, service) + + listRequest := httptest.NewRequest(http.MethodGet, "/listing-bucket?list-type=2&prefix=listing/&max-keys=2", nil) + listRecorder := httptest.NewRecorder() + engine.ServeHTTP(listRecorder, listRequest) + + if got, want := listRecorder.Code, http.StatusOK; got != want { + t.Fatalf("unexpected list status: got %d want %d", got, want) + } + var listPayload struct { + XMLName xml.Name `xml:"ListBucketResult"` + KeyCount int `xml:"KeyCount"` + IsTruncated bool `xml:"IsTruncated"` + NextContinuationToken string `xml:"NextContinuationToken"` + } + if err := xml.Unmarshal(listRecorder.Body.Bytes(), &listPayload); err != nil { + t.Fatalf("decode list xml: %v", err) + } + if got, want := listPayload.KeyCount, 2; got != want { + t.Fatalf("unexpected key count: got %d want %d", got, want) + } + if !listPayload.IsTruncated { + t.Fatal("expected truncated list response") + } + if strings.TrimSpace(listPayload.NextContinuationToken) == "" { + t.Fatal("expected continuation token") + } + + deleteRequest := httptest.NewRequest(http.MethodPost, "/listing-bucket?delete", strings.NewReader(` + + true + listing/a.txt + listing/b.txt +`)) + deleteRecorder := httptest.NewRecorder() + engine.ServeHTTP(deleteRecorder, deleteRequest) + + if got, want := deleteRecorder.Code, http.StatusOK; got != want { + t.Fatalf("unexpected delete status: got %d want %d", got, want) + } + if !strings.Contains(deleteRecorder.Body.String(), "") { + t.Fatalf("expected quiet delete response without deleted entries, got %q", deleteRecorder.Body.String()) + } +} diff --git a/core/internal/delivery/http/sqs_native.go b/core/internal/delivery/http/sqs_native.go index 58dc1f4..b4292be 100644 --- a/core/internal/delivery/http/sqs_native.go +++ b/core/internal/delivery/http/sqs_native.go @@ -1,17 +1,51 @@ package http import ( + "crypto/md5" + "encoding/hex" + "encoding/xml" "errors" "net/http" + "sort" + "strconv" "strings" + "time" "github.com/gin-gonic/gin" "github.com/michasdev/mildstack/core/internal/application/orchestrator" + "github.com/michasdev/mildstack/core/internal/resources/awscontext" + "github.com/michasdev/mildstack/core/internal/resources/sqs/contracts" + "github.com/michasdev/mildstack/core/internal/resources/sqs/domain" ) type SQSNativeService interface { Policy() orchestrator.EmulationPolicy Metadata() orchestrator.Metadata + QueueURL(queueName string) string + QueueARN(queueName string) string + CreateQueue(queueName string, attributes map[string]string) (domain.Queue, error) + DeleteQueue(queueName string) error + GetQueueUrl(queueName, ownerAccountID string) (string, error) + ListQueues(queueNamePrefix string, maxResults int, nextToken, ownerAccountID string) ([]domain.Queue, string, error) + PurgeQueue(queueName string) error + GetQueueAttributes(queueName string, attributeNames []string, ownerAccountID string) (contracts.QueueAttributesView, error) + SetQueueAttributes(queueName string, attributes map[string]string) (contracts.QueueAttributesView, error) + TagQueue(queueName string, tags map[string]string) error + UntagQueue(queueName string, tagKeys []string) error + AddPermission(queueName, label string, awsAccountIDs, actions []string) error + RemovePermission(queueName, label string) error + ListQueueTags(queueName string) (map[string]string, error) + ListDeadLetterSourceQueues(queueName string) ([]string, error) + StartMessageMoveTask(sourceArn, destinationArn string, maxNumberOfMessagesPerSecond int) (string, error) + CancelMessageMoveTask(taskHandle string) (int64, error) + ListMessageMoveTasks(queueName string) ([]domain.MessageMoveTask, error) + ReceiveMessage(queueName string, maxMessages int, waitTime time.Duration) ([]domain.Message, error) + DeleteMessage(queueName string, receiptHandle string) error + ChangeMessageVisibility(queueName string, receiptHandle string, visibility time.Duration) error + SendMessage(queueName string, request contracts.SendMessageRequest) (contracts.SendMessageResult, error) + SendMessageBatch(queueName string, request contracts.SendMessageBatchRequest) (contracts.SendMessageBatchResult, error) + DeleteMessageBatch(queueName string, request contracts.DeleteMessageBatchRequest) (contracts.DeleteMessageBatchResult, error) + ChangeMessageVisibilityBatch(queueName string, request contracts.ChangeMessageVisibilityBatchRequest) (contracts.ChangeMessageVisibilityBatchResult, error) } func RegisterSQSNativeRoutes(engine *gin.Engine, service SQSNativeService) { @@ -30,23 +64,14 @@ func RegisterSQSNativeRoutes(engine *gin.Engine, service SQSNativeService) { } type sqsNativeHandler struct { - service SQSNativeService - registry SQSRegistry - supported map[string]struct{} + service SQSNativeService + registry SQSRegistry } func newSQSNativeHandler(service SQSNativeService) sqsNativeHandler { - supported := make(map[string]struct{}) - if service != nil { - for _, action := range service.Policy().Supported { - supported[action] = struct{}{} - } - } - return sqsNativeHandler{ - service: service, - registry: NewSQSRegistry(), - supported: supported, + service: service, + registry: NewSQSRegistry(), } } @@ -80,15 +105,397 @@ func (h sqsNativeHandler) dispatch(c *gin.Context) bool { return true } - if _, ok := h.supported[spec.Action]; !ok || spec.DomainDeferred { + if !spec.Supported || spec.DomainDeferred { writeSQSError(c, ErrSQSUnsupported, requestIDFromContext(c)) return true } - writeSQSError(c, ErrSQSUnsupported, requestIDFromContext(c)) + switch spec.Action { + case "CreateQueue": + h.handleCreateQueue(c, ctx) + case "DeleteQueue": + h.handleDeleteQueue(c, ctx) + case "GetQueueUrl": + h.handleGetQueueUrl(c, ctx) + case "ListQueues": + h.handleListQueues(c, ctx) + case "PurgeQueue": + h.handlePurgeQueue(c, ctx) + case "GetQueueAttributes": + h.handleGetQueueAttributes(c, ctx) + case "SetQueueAttributes": + h.handleSetQueueAttributes(c, ctx) + case "TagQueue": + h.handleTagQueue(c, ctx) + case "UntagQueue": + h.handleUntagQueue(c, ctx) + case "AddPermission": + h.handleAddPermission(c, ctx) + case "RemovePermission": + h.handleRemovePermission(c, ctx) + case "ListQueueTags": + h.handleListQueueTags(c, ctx) + case "ListDeadLetterSourceQueues": + h.handleListDeadLetterSourceQueues(c, ctx) + case "StartMessageMoveTask": + h.handleStartMessageMoveTask(c, ctx) + case "CancelMessageMoveTask": + h.handleCancelMessageMoveTask(c, ctx) + case "ListMessageMoveTasks": + h.handleListMessageMoveTasks(c, ctx) + case "ReceiveMessage": + h.handleReceiveMessage(c, ctx) + case "SendMessage": + h.handleSendMessage(c, ctx) + case "SendMessageBatch": + h.handleSendMessageBatch(c, ctx) + case "DeleteMessage": + h.handleDeleteMessage(c, ctx) + case "DeleteMessageBatch": + h.handleDeleteMessageBatch(c, ctx) + case "ChangeMessageVisibility": + h.handleChangeMessageVisibility(c, ctx) + case "ChangeMessageVisibilityBatch": + h.handleChangeMessageVisibilityBatch(c, ctx) + default: + writeSQSError(c, ErrSQSUnsupported, requestIDFromContext(c)) + } return true } +func (h sqsNativeHandler) handleCreateQueue(c *gin.Context, ctx SQSRequestContext) { + queueName := strings.TrimSpace(ctx.Values.Get("QueueName")) + attributes := queueAttributesFromValues(ctx.Values) + queue, err := h.service.CreateQueue(queueName, attributes) + if err != nil { + h.finishQueueAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSCreateQueueJSONResponse(c, queue) + return + } + writeSQSCreateQueueResponse(c, queue) +} + +func (h sqsNativeHandler) handleDeleteQueue(c *gin.Context, ctx SQSRequestContext) { + if err := h.service.DeleteQueue(ctx.QueueName); err != nil { + h.finishQueueAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSNoBodyActionResponse(c, "DeleteQueueResponse") + return + } + writeSQSDeleteQueueResponse(c) +} + +func (h sqsNativeHandler) handleGetQueueUrl(c *gin.Context, ctx SQSRequestContext) { + queueName := strings.TrimSpace(ctx.Values.Get("QueueName")) + ownerAccountID := queueOwnerAccountID(ctx.Values) + queueURL, err := h.service.GetQueueUrl(queueName, ownerAccountID) + if err != nil { + h.finishQueueAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSGetQueueURLJSONResponse(c, queueURL) + return + } + writeSQSGetQueueUrlResponse(c, queueURL) +} + +func (h sqsNativeHandler) handleListQueues(c *gin.Context, ctx SQSRequestContext) { + prefix := queueNamePrefix(ctx.Values) + maxResults := queueMaxResults(ctx.Values) + nextToken := queueNextToken(ctx.Values) + ownerAccountID := queueOwnerAccountID(ctx.Values) + queues, nextPageToken, err := h.service.ListQueues(prefix, maxResults, nextToken, ownerAccountID) + if err != nil { + h.finishQueueAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSListQueuesJSONResponse(c, queues, nextPageToken) + return + } + writeSQSListQueuesResponse(c, queues, nextPageToken) +} + +func (h sqsNativeHandler) handlePurgeQueue(c *gin.Context, ctx SQSRequestContext) { + if err := h.service.PurgeQueue(ctx.QueueName); err != nil { + h.finishQueueAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSNoBodyActionResponse(c, "PurgeQueueResponse") + return + } + writeSQSPurgeQueueResponse(c) +} + +func (h sqsNativeHandler) handleGetQueueAttributes(c *gin.Context, ctx SQSRequestContext) { + attributeNames := queueAttributeNames(ctx.Values) + ownerAccountID := queueOwnerAccountID(ctx.Values) + attributes, err := h.service.GetQueueAttributes(ctx.QueueName, attributeNames, ownerAccountID) + if err != nil { + h.finishQueueAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSGetQueueAttributesJSONResponse(c, attributes) + return + } + writeSQSGetQueueAttributesResponse(c, attributes) +} + +func (h sqsNativeHandler) handleSetQueueAttributes(c *gin.Context, ctx SQSRequestContext) { + attributes := queueAttributesFromValues(ctx.Values) + view, err := h.service.SetQueueAttributes(ctx.QueueName, attributes) + if err != nil { + h.finishQueueAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSNoBodyActionResponse(c, "SetQueueAttributesResponse") + return + } + writeSQSSetQueueAttributesResponse(c, view) +} + +func (h sqsNativeHandler) handleTagQueue(c *gin.Context, ctx SQSRequestContext) { + if err := h.service.TagQueue(ctx.QueueName, tagQueueTagsFromValues(ctx.Values)); err != nil { + h.finishQueueAction(c, err) + return + } + writeSQSNoBodyActionResponse(c, "TagQueueResponse") +} + +func (h sqsNativeHandler) handleUntagQueue(c *gin.Context, ctx SQSRequestContext) { + if err := h.service.UntagQueue(ctx.QueueName, queueTagKeysFromValues(ctx.Values)); err != nil { + h.finishQueueAction(c, err) + return + } + writeSQSNoBodyActionResponse(c, "UntagQueueResponse") +} + +func (h sqsNativeHandler) handleAddPermission(c *gin.Context, ctx SQSRequestContext) { + label := strings.TrimSpace(ctx.Values.Get("Label")) + if err := h.service.AddPermission(ctx.QueueName, label, permissionAccountsFromValues(ctx.Values), permissionActionsFromValues(ctx.Values)); err != nil { + h.finishQueueAction(c, err) + return + } + writeSQSNoBodyActionResponse(c, "AddPermissionResponse") +} + +func (h sqsNativeHandler) handleRemovePermission(c *gin.Context, ctx SQSRequestContext) { + label := strings.TrimSpace(ctx.Values.Get("Label")) + if err := h.service.RemovePermission(ctx.QueueName, label); err != nil { + h.finishQueueAction(c, err) + return + } + writeSQSNoBodyActionResponse(c, "RemovePermissionResponse") +} + +func (h sqsNativeHandler) handleListQueueTags(c *gin.Context, ctx SQSRequestContext) { + tags, err := h.service.ListQueueTags(ctx.QueueName) + if err != nil { + h.finishQueueAction(c, err) + return + } + if ctx.TargetStyle { + c.JSON(http.StatusOK, sqsListQueueTagsJSONResponse{Tags: tags}) + return + } + writeSQSListQueueTagsResponse(c, tags) +} + +func (h sqsNativeHandler) handleListDeadLetterSourceQueues(c *gin.Context, ctx SQSRequestContext) { + queueNames, err := h.service.ListDeadLetterSourceQueues(ctx.QueueName) + if err != nil { + h.finishQueueAction(c, err) + return + } + queueURLs := make([]string, 0, len(queueNames)) + for _, queueName := range queueNames { + queueURLs = append(queueURLs, h.service.QueueURL(queueName)) + } + if ctx.TargetStyle { + c.JSON(http.StatusOK, sqsListDeadLetterSourceQueuesJSONResponse{QueueUrls: queueURLs}) + return + } + writeSQSListDeadLetterSourceQueuesResponse(c, queueURLs) +} + +func (h sqsNativeHandler) handleStartMessageMoveTask(c *gin.Context, ctx SQSRequestContext) { + sourceArn, destinationArn, maxPerSecond := startMessageMoveTaskRequestFromValues(ctx.Values) + if sourceArn == "" { + sourceArn = h.service.QueueARN(ctx.QueueName) + } + taskHandle, err := h.service.StartMessageMoveTask(sourceArn, destinationArn, maxPerSecond) + if err != nil { + h.finishQueueAction(c, err) + return + } + if ctx.TargetStyle { + c.JSON(http.StatusOK, sqsStartMessageMoveTaskJSONResponse{TaskHandle: taskHandle}) + return + } + writeSQSStartMessageMoveTaskResponse(c, taskHandle) +} + +func (h sqsNativeHandler) handleCancelMessageMoveTask(c *gin.Context, ctx SQSRequestContext) { + taskHandle := cancelMessageMoveTaskRequestFromValues(ctx.Values) + moved, err := h.service.CancelMessageMoveTask(taskHandle) + if err != nil { + h.finishQueueAction(c, err) + return + } + if ctx.TargetStyle { + c.JSON(http.StatusOK, sqsCancelMessageMoveTaskJSONResponse{ApproximateNumberOfMessagesMoved: moved}) + return + } + writeSQSCancelMessageMoveTaskResponse(c, moved) +} + +func (h sqsNativeHandler) handleListMessageMoveTasks(c *gin.Context, ctx SQSRequestContext) { + _, maxResults := listMessageMoveTasksRequestFromValues(ctx.Values) + tasks, err := h.service.ListMessageMoveTasks(ctx.QueueName) + if err != nil { + h.finishQueueAction(c, err) + return + } + if maxResults <= 0 { + maxResults = 1 + } + if maxResults > 0 && len(tasks) > maxResults { + tasks = append([]domain.MessageMoveTask(nil), tasks[:maxResults]...) + } + if ctx.TargetStyle { + c.JSON(http.StatusOK, sqsListMessageMoveTasksJSONResponse{Results: messageMoveTaskJSONResults(tasks)}) + return + } + writeSQSListMessageMoveTasksResponse(c, tasks) +} + +func (h sqsNativeHandler) handleReceiveMessage(c *gin.Context, ctx SQSRequestContext) { + request := receiveMessageRequestFromValues(ctx.Values) + maxMessages := request.MaxNumberOfMessages + if maxMessages <= 0 { + maxMessages = 1 + } + waitTime := time.Duration(request.WaitTimeSeconds) * time.Second + messages, err := h.service.ReceiveMessage(ctx.QueueName, maxMessages, waitTime) + if err != nil { + h.finishMessageAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSReceiveMessageJSON(c, messages) + return + } + writeSQSReceiveMessageResponse(c, messages) +} + +func (h sqsNativeHandler) handleSendMessage(c *gin.Context, ctx SQSRequestContext) { + request := sendMessageRequestFromValues(ctx.Values) + result, err := h.service.SendMessage(ctx.QueueName, request) + if err != nil { + h.finishMessageAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSSendMessageJSON(c, result) + return + } + writeSQSSendMessageResponse(c, result) +} + +func (h sqsNativeHandler) handleSendMessageBatch(c *gin.Context, ctx SQSRequestContext) { + request := sendMessageBatchRequestFromValues(ctx.Values) + result, err := h.service.SendMessageBatch(ctx.QueueName, request) + if err != nil { + h.finishMessageAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSSendMessageBatchJSON(c, result) + return + } + writeSQSSendMessageBatchResponse(c, result) +} + +func (h sqsNativeHandler) handleDeleteMessage(c *gin.Context, ctx SQSRequestContext) { + request := deleteMessageRequestFromValues(ctx.Values) + if err := h.service.DeleteMessage(ctx.QueueName, request.ReceiptHandle); err != nil { + h.finishMessageAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSNoBodyActionResponse(c, "DeleteMessageResponse") + return + } + writeSQSDeleteMessageResponse(c) +} + +func (h sqsNativeHandler) handleDeleteMessageBatch(c *gin.Context, ctx SQSRequestContext) { + request := deleteMessageBatchRequestFromValues(ctx.Values) + result, err := h.service.DeleteMessageBatch(ctx.QueueName, request) + if err != nil { + h.finishMessageAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSDeleteMessageBatchJSON(c, result) + return + } + writeSQSDeleteMessageBatchResponse(c, result) +} + +func (h sqsNativeHandler) handleChangeMessageVisibility(c *gin.Context, ctx SQSRequestContext) { + request := changeMessageVisibilityRequestFromValues(ctx.Values) + visibility := time.Duration(request.VisibilityTimeout) * time.Second + if err := h.service.ChangeMessageVisibility(ctx.QueueName, request.ReceiptHandle, visibility); err != nil { + h.finishMessageAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSNoBodyActionResponse(c, "ChangeMessageVisibilityResponse") + return + } + writeSQSChangeMessageVisibilityResponse(c) +} + +func (h sqsNativeHandler) handleChangeMessageVisibilityBatch(c *gin.Context, ctx SQSRequestContext) { + request := changeMessageVisibilityBatchRequestFromValues(ctx.Values) + result, err := h.service.ChangeMessageVisibilityBatch(ctx.QueueName, request) + if err != nil { + h.finishMessageAction(c, err) + return + } + if ctx.TargetStyle { + writeSQSChangeMessageVisibilityBatchJSON(c, result) + return + } + writeSQSChangeMessageVisibilityBatchResponse(c, result) +} + +func (h sqsNativeHandler) finishQueueAction(c *gin.Context, err error) { + if err == nil || errors.Is(err, contracts.ErrSQSOperationDeferred) { + writeSQSError(c, ErrSQSUnsupported, requestIDFromContext(c)) + return + } + writeSQSError(c, err, requestIDFromContext(c)) +} + +func (h sqsNativeHandler) finishMessageAction(c *gin.Context, err error) { + if err == nil || errors.Is(err, contracts.ErrSQSOperationDeferred) { + writeSQSError(c, ErrSQSUnsupported, requestIDFromContext(c)) + return + } + writeSQSError(c, err, requestIDFromContext(c)) +} + func requestIDFromContext(c *gin.Context) string { if c == nil { return "mildstack-sqs-request" @@ -102,3 +509,693 @@ func requestIDFromContext(c *gin.Context) string { return "mildstack-sqs-request" } + +type sqsQueueUrlResponse struct { + XMLName xml.Name `xml:"GetQueueUrlResponse"` + GetQueueUrlResult sqsQueueUrlResult `xml:"GetQueueUrlResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsQueueUrlResult struct { + QueueURL string `xml:"QueueUrl"` +} + +type sqsCreateQueueResponse struct { + XMLName xml.Name `xml:"CreateQueueResponse"` + CreateQueueResult sqsQueueUrlResult `xml:"CreateQueueResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsListQueuesResponse struct { + XMLName xml.Name `xml:"ListQueuesResponse"` + ListQueuesResult sqsListQueuesResult `xml:"ListQueuesResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsListQueuesResult struct { + QueueUrls []string `xml:"QueueUrl"` + NextToken string `xml:"NextToken,omitempty"` +} + +type sqsListQueuesJSONResponse struct { + QueueUrls []string `json:"QueueUrls"` + NextToken string `json:"NextToken,omitempty"` +} + +type sqsQueueURLJSONResponse struct { + QueueUrl string `json:"QueueUrl"` +} + +type sqsGetQueueAttributesJSONResponse struct { + Attributes map[string]string `json:"Attributes"` +} + +type sqsGetQueueAttributesResponse struct { + XMLName xml.Name `xml:"GetQueueAttributesResponse"` + GetQueueAttributesResult sqsGetQueueAttributesResult `xml:"GetQueueAttributesResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsGetQueueAttributesResult struct { + Attributes []sqsQueueAttributeXML `xml:"Attribute"` +} + +type sqsQueueAttributeXML struct { + Name string `xml:"Name"` + Value string `xml:"Value"` +} + +type sqsSetQueueAttributesResponse struct { + XMLName xml.Name `xml:"SetQueueAttributesResponse"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsEmptyActionXMLResponse struct { + XMLName xml.Name + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsListQueueTagsXMLResponse struct { + XMLName xml.Name `xml:"ListQueueTagsResponse"` + ListQueueTagsResult sqsListQueueTagsResult `xml:"ListQueueTagsResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsListQueueTagsResult struct { + Tags []sqsTagXML `xml:"Tag"` +} + +type sqsTagXML struct { + Key string `xml:"Key"` + Value string `xml:"Value"` +} + +type sqsListQueueTagsJSONResponse struct { + Tags map[string]string `json:"Tags"` +} + +type sqsListDeadLetterSourceQueuesXMLResponse struct { + XMLName xml.Name `xml:"ListDeadLetterSourceQueuesResponse"` + ListDeadLetterSourceQueuesResult sqsListDeadLetterSourceQueuesResult `xml:"ListDeadLetterSourceQueuesResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsListDeadLetterSourceQueuesResult struct { + QueueUrls []string `xml:"QueueUrl"` + NextToken string `xml:"NextToken,omitempty"` +} + +type sqsListDeadLetterSourceQueuesJSONResponse struct { + QueueUrls []string `json:"queueUrls"` + NextToken string `json:"NextToken,omitempty"` +} + +type sqsStartMessageMoveTaskXMLResponse struct { + XMLName xml.Name `xml:"StartMessageMoveTaskResponse"` + StartMessageMoveTaskResult sqsStartMessageMoveTaskResult `xml:"StartMessageMoveTaskResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsStartMessageMoveTaskResult struct { + TaskHandle string `xml:"TaskHandle"` +} + +type sqsStartMessageMoveTaskJSONResponse struct { + TaskHandle string `json:"TaskHandle"` +} + +type sqsCancelMessageMoveTaskXMLResponse struct { + XMLName xml.Name `xml:"CancelMessageMoveTaskResponse"` + CancelMessageMoveTaskResult sqsCancelMessageMoveTaskResult `xml:"CancelMessageMoveTaskResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsCancelMessageMoveTaskResult struct { + ApproximateNumberOfMessagesMoved int64 `xml:"ApproximateNumberOfMessagesMoved"` +} + +type sqsCancelMessageMoveTaskJSONResponse struct { + ApproximateNumberOfMessagesMoved int64 `json:"ApproximateNumberOfMessagesMoved"` +} + +type sqsListMessageMoveTasksXMLResponse struct { + XMLName xml.Name `xml:"ListMessageMoveTasksResponse"` + ListMessageMoveTasksResult sqsListMessageMoveTasksResult `xml:"ListMessageMoveTasksResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsListMessageMoveTasksResult struct { + Results []sqsMessageMoveTaskXML `xml:"Result"` +} + +type sqsListMessageMoveTasksJSONResponse struct { + Results []sqsMessageMoveTaskJSON `json:"Results"` +} + +type sqsMessageMoveTaskXML struct { + ApproximateNumberOfMessagesMoved int64 `xml:"ApproximateNumberOfMessagesMoved,omitempty"` + ApproximateNumberOfMessagesToMove int64 `xml:"ApproximateNumberOfMessagesToMove,omitempty"` + DestinationArn string `xml:"DestinationArn,omitempty"` + MaxNumberOfMessagesPerSecond int `xml:"MaxNumberOfMessagesPerSecond,omitempty"` + SourceArn string `xml:"SourceArn,omitempty"` + StartedTimestamp int64 `xml:"StartedTimestamp,omitempty"` + Status string `xml:"Status,omitempty"` + TaskHandle string `xml:"TaskHandle,omitempty"` +} + +type sqsMessageMoveTaskJSON struct { + ApproximateNumberOfMessagesMoved int64 `json:"ApproximateNumberOfMessagesMoved,omitempty"` + ApproximateNumberOfMessagesToMove int64 `json:"ApproximateNumberOfMessagesToMove,omitempty"` + DestinationArn string `json:"DestinationArn,omitempty"` + MaxNumberOfMessagesPerSecond int `json:"MaxNumberOfMessagesPerSecond,omitempty"` + SourceArn string `json:"SourceArn,omitempty"` + StartedTimestamp int64 `json:"StartedTimestamp,omitempty"` + Status string `json:"Status,omitempty"` + TaskHandle string `json:"TaskHandle,omitempty"` +} + +type sqsSendMessageResponse struct { + XMLName xml.Name `xml:"SendMessageResponse"` + SendMessageResult sqsSendMessageResult `xml:"SendMessageResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsSendMessageResult struct { + MessageID string `xml:"MessageId,omitempty"` + MD5OfMessageBody string `xml:"MD5OfMessageBody,omitempty"` + MD5OfMessageAttributes string `xml:"MD5OfMessageAttributes,omitempty"` + MD5OfMessageSystemAttributes string `xml:"MD5OfMessageSystemAttributes,omitempty"` + SequenceNumber string `xml:"SequenceNumber,omitempty"` +} + +type sqsSendMessageBatchResponse struct { + XMLName xml.Name `xml:"SendMessageBatchResponse"` + SendMessageBatchResult sqsSendMessageBatchResult `xml:"SendMessageBatchResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsSendMessageBatchResult struct { + Successful []contracts.SendMessageBatchResultEntry `xml:"SendMessageBatchResultEntry,omitempty"` + Failed []contracts.BatchResultErrorEntry `xml:"BatchResultErrorEntry,omitempty"` +} + +type sqsDeleteMessageBatchResponse struct { + XMLName xml.Name `xml:"DeleteMessageBatchResponse"` + DeleteMessageBatchResult sqsDeleteMessageBatchResult `xml:"DeleteMessageBatchResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsDeleteMessageBatchResult struct { + Successful []contracts.DeleteMessageBatchResultEntry `xml:"DeleteMessageBatchResultEntry,omitempty"` + Failed []contracts.BatchResultErrorEntry `xml:"BatchResultErrorEntry,omitempty"` +} + +type sqsChangeMessageVisibilityBatchResponse struct { + XMLName xml.Name `xml:"ChangeMessageVisibilityBatchResponse"` + ChangeMessageVisibilityBatchResult sqsChangeMessageVisibilityBatchResult `xml:"ChangeMessageVisibilityBatchResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsChangeMessageVisibilityBatchResult struct { + Successful []contracts.ChangeMessageVisibilityBatchResultEntry `xml:"ChangeMessageVisibilityBatchResultEntry,omitempty"` + Failed []contracts.BatchResultErrorEntry `xml:"BatchResultErrorEntry,omitempty"` +} + +type sqsReceiveMessageResponse struct { + XMLName xml.Name `xml:"ReceiveMessageResponse"` + ReceiveMessageResult sqsReceiveMessageResult `xml:"ReceiveMessageResult"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsReceiveMessageResult struct { + Messages []sqsReceivedMessageXML `xml:"Message,omitempty"` +} + +type sqsReceivedMessageXML struct { + Body string `xml:"Body,omitempty"` + MD5OfBody string `xml:"MD5OfBody,omitempty"` + MD5OfMessageAttributes string `xml:"MD5OfMessageAttributes,omitempty"` + MessageID string `xml:"MessageId,omitempty"` + ReceiptHandle string `xml:"ReceiptHandle,omitempty"` + Attributes []sqsMessageAttributeXML `xml:"Attribute,omitempty"` +} + +type sqsMessageAttributeXML struct { + Name string `xml:"Name"` + Value string `xml:"Value"` +} + +type sqsDeleteQueueResponse struct { + XMLName xml.Name `xml:"DeleteQueueResponse"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsPurgeQueueResponse struct { + XMLName xml.Name `xml:"PurgeQueueResponse"` + ResponseMetadata sqsResponseMetadata `xml:"ResponseMetadata"` +} + +type sqsResponseMetadata struct { + RequestID string `xml:"RequestId"` +} + +func writeSQSCreateQueueResponse(c *gin.Context, queue domain.Queue) { + queueURL := queueURLOrDefault(queue) + c.XML(http.StatusOK, sqsCreateQueueResponse{ + CreateQueueResult: sqsQueueUrlResult{QueueURL: queueURL}, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSCreateQueueJSONResponse(c *gin.Context, queue domain.Queue) { + c.JSON(http.StatusOK, sqsQueueURLJSONResponse{ + QueueUrl: queueURLOrDefault(queue), + }) +} + +func writeSQSGetQueueURLJSONResponse(c *gin.Context, queueURL string) { + c.JSON(http.StatusOK, sqsQueueURLJSONResponse{ + QueueUrl: queueURL, + }) +} + +func writeSQSGetQueueAttributesJSONResponse(c *gin.Context, view contracts.QueueAttributesView) { + c.JSON(http.StatusOK, sqsGetQueueAttributesJSONResponse{ + Attributes: copyQueueAttributes(view.Attributes), + }) +} + +func queueURLOrDefault(queue domain.Queue) string { + queueURL := queue.URL + if queueURL == "" { + aws := awscontext.Default() + queueURL = "https://sqs." + aws.Region + ".amazonaws.com/" + aws.AccountID + "/" + queue.Name + } + return queueURL +} + +func writeSQSGetQueueUrlResponse(c *gin.Context, queueURL string) { + c.XML(http.StatusOK, sqsQueueUrlResponse{ + GetQueueUrlResult: sqsQueueUrlResult{QueueURL: queueURL}, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSListQueuesResponse(c *gin.Context, queues []domain.Queue, nextToken string) { + urls := make([]string, 0, len(queues)) + for _, queue := range queues { + queueURL := queue.URL + if queueURL == "" { + aws := awscontext.Default() + queueURL = "https://sqs." + aws.Region + ".amazonaws.com/" + aws.AccountID + "/" + queue.Name + } + urls = append(urls, queueURL) + } + c.XML(http.StatusOK, sqsListQueuesResponse{ + ListQueuesResult: sqsListQueuesResult{ + QueueUrls: urls, + NextToken: nextToken, + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSListQueuesJSONResponse(c *gin.Context, queues []domain.Queue, nextToken string) { + urls := make([]string, 0, len(queues)) + for _, queue := range queues { + queueURL := queue.URL + if queueURL == "" { + aws := awscontext.Default() + queueURL = "https://sqs." + aws.Region + ".amazonaws.com/" + aws.AccountID + "/" + queue.Name + } + urls = append(urls, queueURL) + } + + c.JSON(http.StatusOK, sqsListQueuesJSONResponse{ + QueueUrls: urls, + NextToken: nextToken, + }) +} + +func writeSQSGetQueueAttributesResponse(c *gin.Context, view contracts.QueueAttributesView) { + c.XML(http.StatusOK, sqsGetQueueAttributesResponse{ + GetQueueAttributesResult: sqsGetQueueAttributesResult{ + Attributes: sortedQueueAttributeXML(view.Attributes), + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSSetQueueAttributesResponse(c *gin.Context, _ contracts.QueueAttributesView) { + c.XML(http.StatusOK, sqsSetQueueAttributesResponse{ + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSNoBodyActionResponse(c *gin.Context, root string) { + if c == nil { + return + } + if c.Request != nil && c.Request.Header.Get("X-Amz-Target") != "" { + c.JSON(http.StatusOK, gin.H{}) + return + } + c.XML(http.StatusOK, sqsEmptyActionXMLResponse{ + XMLName: xml.Name{Local: root}, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSListQueueTagsResponse(c *gin.Context, tags map[string]string) { + c.XML(http.StatusOK, sqsListQueueTagsXMLResponse{ + XMLName: xml.Name{Local: "ListQueueTagsResponse"}, + ListQueueTagsResult: sqsListQueueTagsResult{ + Tags: tagsToXML(tags), + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSListDeadLetterSourceQueuesResponse(c *gin.Context, queueURLs []string) { + c.XML(http.StatusOK, sqsListDeadLetterSourceQueuesXMLResponse{ + XMLName: xml.Name{Local: "ListDeadLetterSourceQueuesResponse"}, + ListDeadLetterSourceQueuesResult: sqsListDeadLetterSourceQueuesResult{ + QueueUrls: queueURLs, + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSStartMessageMoveTaskResponse(c *gin.Context, taskHandle string) { + c.XML(http.StatusOK, sqsStartMessageMoveTaskXMLResponse{ + XMLName: xml.Name{Local: "StartMessageMoveTaskResponse"}, + StartMessageMoveTaskResult: sqsStartMessageMoveTaskResult{ + TaskHandle: taskHandle, + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSCancelMessageMoveTaskResponse(c *gin.Context, moved int64) { + c.XML(http.StatusOK, sqsCancelMessageMoveTaskXMLResponse{ + XMLName: xml.Name{Local: "CancelMessageMoveTaskResponse"}, + CancelMessageMoveTaskResult: sqsCancelMessageMoveTaskResult{ + ApproximateNumberOfMessagesMoved: moved, + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSListMessageMoveTasksResponse(c *gin.Context, tasks []domain.MessageMoveTask) { + c.XML(http.StatusOK, sqsListMessageMoveTasksXMLResponse{ + XMLName: xml.Name{Local: "ListMessageMoveTasksResponse"}, + ListMessageMoveTasksResult: sqsListMessageMoveTasksResult{ + Results: messageMoveTaskXMLResults(tasks), + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func tagsToXML(tags map[string]string) []sqsTagXML { + if len(tags) == 0 { + return nil + } + keys := make([]string, 0, len(tags)) + for key := range tags { + keys = append(keys, key) + } + sort.Strings(keys) + result := make([]sqsTagXML, 0, len(keys)) + for _, key := range keys { + result = append(result, sqsTagXML{Key: key, Value: tags[key]}) + } + return result +} + +func messageMoveTaskXMLResults(tasks []domain.MessageMoveTask) []sqsMessageMoveTaskXML { + if len(tasks) == 0 { + return nil + } + results := make([]sqsMessageMoveTaskXML, 0, len(tasks)) + for _, task := range tasks { + xmlTask := sqsMessageMoveTaskXML{ + ApproximateNumberOfMessagesMoved: task.ApproximateNumberOfMessagesMoved, + DestinationArn: task.DestinationArn, + MaxNumberOfMessagesPerSecond: task.MaxNumberOfMessagesPerSecond, + SourceArn: task.SourceArn, + StartedTimestamp: task.StartedAt.UnixMilli(), + Status: task.Status, + } + if strings.EqualFold(task.Status, "RUNNING") { + xmlTask.TaskHandle = task.TaskHandle + } + results = append(results, xmlTask) + } + return results +} + +func messageMoveTaskJSONResults(tasks []domain.MessageMoveTask) []sqsMessageMoveTaskJSON { + if len(tasks) == 0 { + return nil + } + results := make([]sqsMessageMoveTaskJSON, 0, len(tasks)) + for _, task := range tasks { + item := sqsMessageMoveTaskJSON{ + ApproximateNumberOfMessagesMoved: task.ApproximateNumberOfMessagesMoved, + DestinationArn: task.DestinationArn, + MaxNumberOfMessagesPerSecond: task.MaxNumberOfMessagesPerSecond, + SourceArn: task.SourceArn, + StartedTimestamp: task.StartedAt.UnixMilli(), + Status: task.Status, + } + if strings.EqualFold(task.Status, "RUNNING") { + item.TaskHandle = task.TaskHandle + } + results = append(results, item) + } + return results +} + +func writeSQSSendMessageResponse(c *gin.Context, result contracts.SendMessageResult) { + c.XML(http.StatusOK, sqsSendMessageResponse{ + SendMessageResult: sqsSendMessageResult{ + MessageID: result.MessageId, + MD5OfMessageBody: result.MD5OfMessageBody, + MD5OfMessageAttributes: result.MD5OfMessageAttributes, + MD5OfMessageSystemAttributes: result.MD5OfMessageSystemAttributes, + SequenceNumber: result.SequenceNumber, + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSSendMessageJSON(c *gin.Context, result contracts.SendMessageResult) { + c.JSON(http.StatusOK, result) +} + +func writeSQSSendMessageBatchResponse(c *gin.Context, result contracts.SendMessageBatchResult) { + c.XML(http.StatusOK, sqsSendMessageBatchResponse{ + SendMessageBatchResult: sqsSendMessageBatchResult{ + Successful: result.Successful, + Failed: result.Failed, + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSSendMessageBatchJSON(c *gin.Context, result contracts.SendMessageBatchResult) { + c.JSON(http.StatusOK, result) +} + +func writeSQSDeleteMessageResponse(c *gin.Context) { + if c == nil { + return + } + c.Status(http.StatusOK) +} + +func writeSQSDeleteMessageBatchResponse(c *gin.Context, result contracts.DeleteMessageBatchResult) { + c.XML(http.StatusOK, sqsDeleteMessageBatchResponse{ + DeleteMessageBatchResult: sqsDeleteMessageBatchResult{ + Successful: result.Successful, + Failed: result.Failed, + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSDeleteMessageBatchJSON(c *gin.Context, result contracts.DeleteMessageBatchResult) { + c.JSON(http.StatusOK, result) +} + +func writeSQSChangeMessageVisibilityResponse(c *gin.Context) { + if c == nil { + return + } + c.Status(http.StatusOK) +} + +func writeSQSChangeMessageVisibilityBatchResponse(c *gin.Context, result contracts.ChangeMessageVisibilityBatchResult) { + c.XML(http.StatusOK, sqsChangeMessageVisibilityBatchResponse{ + ChangeMessageVisibilityBatchResult: sqsChangeMessageVisibilityBatchResult{ + Successful: result.Successful, + Failed: result.Failed, + }, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSChangeMessageVisibilityBatchJSON(c *gin.Context, result contracts.ChangeMessageVisibilityBatchResult) { + c.JSON(http.StatusOK, result) +} + +func writeSQSReceiveMessageResponse(c *gin.Context, messages []domain.Message) { + xmlMessages := make([]sqsReceivedMessageXML, 0, len(messages)) + for _, message := range messages { + xmlMessages = append(xmlMessages, sqsReceivedMessageXML{ + Body: message.Body, + MD5OfBody: messageMD5OfBody(message.Body), + MessageID: message.MessageID, + ReceiptHandle: applicationCurrentReceiptHandle(message), + Attributes: receivedMessageAttributesXML(message), + MD5OfMessageAttributes: "", + }) + } + c.XML(http.StatusOK, sqsReceiveMessageResponse{ + ReceiveMessageResult: sqsReceiveMessageResult{Messages: xmlMessages}, + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSReceiveMessageJSON(c *gin.Context, messages []domain.Message) { + received := make([]contracts.ReceivedMessage, 0, len(messages)) + for _, message := range messages { + received = append(received, contracts.ReceivedMessage{ + Attributes: receivedMessageAttributesMap(message), + Body: message.Body, + MD5OfBody: messageMD5OfBody(message.Body), + MD5OfMessageAttributes: "", + MessageAttributes: receivedMessageAttributesValues(message), + MessageId: message.MessageID, + ReceiptHandle: applicationCurrentReceiptHandle(message), + }) + } + c.JSON(http.StatusOK, contracts.ReceiveMessageResult{Messages: received}) +} + +func messageMD5OfBody(body string) string { + sum := md5.Sum([]byte(body)) + return hex.EncodeToString(sum[:]) +} + +func applicationCurrentReceiptHandle(message domain.Message) string { + if len(message.ReceiptKeys) == 0 { + return "" + } + return message.ReceiptKeys[len(message.ReceiptKeys)-1] +} + +func receivedMessageAttributesXML(message domain.Message) []sqsMessageAttributeXML { + attrs := make([]sqsMessageAttributeXML, 0, 4) + if message.Recovery.Attempts > 0 { + attrs = append(attrs, sqsMessageAttributeXML{Name: "ApproximateReceiveCount", Value: strconv.Itoa(message.Recovery.Attempts)}) + } + if !message.SentAt.IsZero() { + attrs = append(attrs, sqsMessageAttributeXML{Name: "SentTimestamp", Value: strconv.FormatInt(message.SentAt.UnixMilli(), 10)}) + } + if timestamp := strings.TrimSpace(message.Metadata["approximate_first_receive_timestamp"]); timestamp != "" { + attrs = append(attrs, sqsMessageAttributeXML{Name: "ApproximateFirstReceiveTimestamp", Value: timestamp}) + } + if groupID := strings.TrimSpace(message.MessageGroupID); groupID != "" { + attrs = append(attrs, sqsMessageAttributeXML{Name: "MessageGroupId", Value: groupID}) + } + if sequenceNumber := message.SequenceNumber; sequenceNumber > 0 { + attrs = append(attrs, sqsMessageAttributeXML{Name: "SequenceNumber", Value: strconv.FormatInt(sequenceNumber, 10)}) + } + if dedupeID := strings.TrimSpace(message.Metadata["MessageDeduplicationId"]); dedupeID != "" { + attrs = append(attrs, sqsMessageAttributeXML{Name: "MessageDeduplicationId", Value: dedupeID}) + } + return attrs +} + +func receivedMessageAttributesMap(message domain.Message) map[string]string { + attrs := map[string]string{} + if message.Recovery.Attempts > 0 { + attrs["ApproximateReceiveCount"] = strconv.Itoa(message.Recovery.Attempts) + } + if !message.SentAt.IsZero() { + attrs["SentTimestamp"] = strconv.FormatInt(message.SentAt.UnixMilli(), 10) + } + if timestamp := strings.TrimSpace(message.Metadata["approximate_first_receive_timestamp"]); timestamp != "" { + attrs["ApproximateFirstReceiveTimestamp"] = timestamp + } + if groupID := strings.TrimSpace(message.MessageGroupID); groupID != "" { + attrs["MessageGroupId"] = groupID + } + if sequenceNumber := message.SequenceNumber; sequenceNumber > 0 { + attrs["SequenceNumber"] = strconv.FormatInt(sequenceNumber, 10) + } + if dedupeID := strings.TrimSpace(message.Metadata["MessageDeduplicationId"]); dedupeID != "" { + attrs["MessageDeduplicationId"] = dedupeID + } + if len(attrs) == 0 { + return nil + } + return attrs +} + +func receivedMessageAttributesValues(message domain.Message) map[string]contracts.MessageAttributeValue { + if len(message.Attributes) == 0 { + return nil + } + received := make(map[string]contracts.MessageAttributeValue, len(message.Attributes)) + for key, value := range message.Attributes { + received[key] = contracts.MessageAttributeValue{ + DataType: "String", + StringValue: value, + } + } + return received +} + +func writeSQSDeleteQueueResponse(c *gin.Context) { + c.XML(http.StatusOK, sqsDeleteQueueResponse{ + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func writeSQSPurgeQueueResponse(c *gin.Context) { + c.XML(http.StatusOK, sqsPurgeQueueResponse{ + ResponseMetadata: sqsResponseMetadata{RequestID: requestIDFromContext(c)}, + }) +} + +func sortedQueueAttributeXML(attributes map[string]string) []sqsQueueAttributeXML { + if len(attributes) == 0 { + return nil + } + + keys := make([]string, 0, len(attributes)) + for key := range attributes { + keys = append(keys, key) + } + sort.Strings(keys) + + items := make([]sqsQueueAttributeXML, 0, len(keys)) + for _, key := range keys { + items = append(items, sqsQueueAttributeXML{Name: key, Value: attributes[key]}) + } + return items +} + +func copyQueueAttributes(attributes map[string]string) map[string]string { + if len(attributes) == 0 { + return map[string]string{} + } + + cloned := make(map[string]string, len(attributes)) + for key, value := range attributes { + cloned[key] = value + } + return cloned +} diff --git a/core/internal/delivery/http/sqs_native_contract.go b/core/internal/delivery/http/sqs_native_contract.go index 8942694..cbb8ca8 100644 --- a/core/internal/delivery/http/sqs_native_contract.go +++ b/core/internal/delivery/http/sqs_native_contract.go @@ -1,12 +1,17 @@ package http import ( + "bytes" + "encoding/json" "errors" "fmt" + "io" "mime" "net/http" "net/url" "path" + "sort" + "strconv" "strings" "github.com/michasdev/mildstack/core/internal/resources/sqs/contracts" @@ -40,6 +45,7 @@ type SQSRequestContext struct { QueueName string Action string Version string + TargetStyle bool Values url.Values } @@ -65,7 +71,7 @@ func ParseSQSRequest(req *http.Request) (SQSRequestContext, error) { if err != nil { return SQSRequestContext{}, fmt.Errorf("%w: %v", ErrSQSMalformedRequest, err) } - if mediaType != "" && mediaType != "application/x-www-form-urlencoded" { + if mediaType != "" && mediaType != "application/x-www-form-urlencoded" && mediaType != "application/x-amz-json-1.0" { return SQSRequestContext{}, ErrSQSMalformedRequest } } @@ -74,11 +80,50 @@ func ParseSQSRequest(req *http.Request) (SQSRequestContext, error) { return SQSRequestContext{}, fmt.Errorf("%w: %v", ErrSQSMalformedRequest, err) } - action := strings.TrimSpace(req.Form.Get("Action")) + values := cloneValues(req.Form) + if values == nil { + values = url.Values{} + } + + targetAction, jsonMode, err := parseSQSTarget(req.Header.Get("X-Amz-Target")) + if err != nil { + return SQSRequestContext{}, err + } + if jsonMode { + if req.Body != nil { + bodyBytes, _ := io.ReadAll(req.Body) + req.Body = io.NopCloser(bytes.NewReader(bodyBytes)) + if len(bodyBytes) > 0 { + var payload map[string]any + if err := json.Unmarshal(bodyBytes, &payload); err == nil { + mergeJSONValues(values, payload) + } + } + } + if targetAction != "" { + values.Set("Action", targetAction) + } + if values.Get("Version") == "" { + values.Set("Version", sqsQueryVersion) + } + } + + if strings.TrimSpace(values.Get("Action")) == "" || strings.TrimSpace(values.Get("Version")) == "" { + if req.Body != nil { + bodyBytes, _ := io.ReadAll(req.Body) + req.Body = io.NopCloser(bytes.NewReader(bodyBytes)) + if parsedBody, err := url.ParseQuery(string(bodyBytes)); err == nil { + mergeValues(values, parsedBody) + } + } + mergeValues(values, req.URL.Query()) + } + + action := strings.TrimSpace(values.Get("Action")) if action == "" { return SQSRequestContext{}, ErrSQSMissingAction } - version := strings.TrimSpace(req.Form.Get("Version")) + version := strings.TrimSpace(values.Get("Version")) if version == "" { return SQSRequestContext{}, ErrSQSInvalidVersion } @@ -86,6 +131,17 @@ func ParseSQSRequest(req *http.Request) (SQSRequestContext, error) { return SQSRequestContext{}, ErrSQSInvalidVersion } + if pathContext.Kind == SQSRequestKindRoot && isQueueScopedAction(action) { + queueName, accountID, err := queueContextFromValuesForAction(action, values) + if err != nil { + return SQSRequestContext{}, err + } + pathContext.Kind = SQSRequestKindQueue + pathContext.QueueName = queueName + pathContext.AccountID = accountID + pathContext.NormalizedPath = "/" + strings.Trim(accountID, "/") + "/" + strings.Trim(queueName, "/") + } + return SQSRequestContext{ Method: strings.ToUpper(strings.TrimSpace(req.Method)), RawPath: normalizeRequestPath(req.URL.Path), @@ -95,7 +151,8 @@ func ParseSQSRequest(req *http.Request) (SQSRequestContext, error) { QueueName: pathContext.QueueName, Action: action, Version: version, - Values: cloneValues(req.Form), + TargetStyle: jsonMode, + Values: values, }, nil } @@ -154,6 +211,99 @@ func cloneValues(values url.Values) url.Values { return cloned } +func mergeValues(dst, src url.Values) { + if dst == nil || len(src) == 0 { + return + } + + for key, values := range src { + if len(values) == 0 { + continue + } + dst[key] = append([]string(nil), values...) + } +} + +func mergeJSONValues(dst url.Values, payload map[string]any) { + if dst == nil || len(payload) == 0 { + return + } + + for key, value := range payload { + switch typed := value.(type) { + case string: + dst.Set(key, typed) + case bool: + dst.Set(key, strconv.FormatBool(typed)) + case float64: + dst.Set(key, strconv.FormatFloat(typed, 'f', -1, 64)) + case map[string]any: + if strings.EqualFold(key, "Attributes") { + flattenQueueAttributes(dst, typed) + continue + } + if strings.EqualFold(key, "MessageAttributes") { + flattenMessageAttributes(dst, "MessageAttribute", typed) + continue + } + if strings.EqualFold(key, "MessageSystemAttributes") { + flattenMessageAttributes(dst, "MessageSystemAttribute", typed) + continue + } + for nestedKey, nestedValue := range typed { + dst.Set(fmt.Sprintf("%s.%s", key, nestedKey), fmt.Sprint(nestedValue)) + } + case []any: + if strings.EqualFold(key, "AttributeNames") { + for i, item := range typed { + dst.Set(fmt.Sprintf("AttributeName.%d", i+1), fmt.Sprint(item)) + } + continue + } + if strings.EqualFold(key, "MessageAttributeNames") { + for i, item := range typed { + dst.Set(fmt.Sprintf("MessageAttributeName.%d", i+1), fmt.Sprint(item)) + } + continue + } + if strings.EqualFold(key, "MessageSystemAttributeNames") { + for i, item := range typed { + dst.Set(fmt.Sprintf("MessageSystemAttributeName.%d", i+1), fmt.Sprint(item)) + } + continue + } + if strings.EqualFold(key, "Entries") { + for i, item := range typed { + if entry, ok := item.(map[string]any); ok { + flattenJSONEntry(dst, i+1, entry) + } + } + continue + } + for i, item := range typed { + dst.Set(fmt.Sprintf("%s.%d", key, i+1), fmt.Sprint(item)) + } + default: + dst.Set(key, fmt.Sprint(value)) + } + } +} + +func parseSQSTarget(raw string) (string, bool, error) { + target := strings.TrimSpace(raw) + if target == "" { + return "", false, nil + } + if !strings.HasPrefix(target, "AmazonSQS.") { + return "", false, fmt.Errorf("sqs: X-Amz-Target %q must start with %q", target, "AmazonSQS.") + } + action := strings.TrimSpace(strings.TrimPrefix(target, "AmazonSQS.")) + if action == "" { + return "", false, fmt.Errorf("sqs: X-Amz-Target %q is missing an operation name", target) + } + return action, true, nil +} + func validateSQSRequestContext(ctx SQSRequestContext, spec SQSRegistrySpec) error { if spec.Action == "" { return ErrSQSInvalidAction @@ -176,3 +326,649 @@ func validateSQSRequestContext(ctx SQSRequestContext, spec SQSRegistrySpec) erro } return nil } + +func queueOwnerAccountID(values url.Values) string { + return strings.TrimSpace(values.Get("QueueOwnerAWSAccountId")) +} + +func queueNamePrefix(values url.Values) string { + return strings.TrimSpace(values.Get("QueueNamePrefix")) +} + +func queueNextToken(values url.Values) string { + return strings.TrimSpace(values.Get("NextToken")) +} + +func queueMaxResults(values url.Values) int { + raw := strings.TrimSpace(values.Get("MaxResults")) + if raw == "" { + return 0 + } + + value, err := strconv.Atoi(raw) + if err != nil || value < 0 { + return 0 + } + return value +} + +func queueAttributeNames(values url.Values) []string { + names := make([]string, 0) + for key, list := range values { + if !strings.HasPrefix(key, "AttributeName") || len(list) == 0 { + continue + } + for _, value := range list { + if trimmed := strings.TrimSpace(value); trimmed != "" { + names = append(names, trimmed) + } + } + } + sort.Strings(names) + return names +} + +func queueNameFromQueueURL(raw string) (string, string, error) { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return "", "", fmt.Errorf("sqs: QueueUrl is required") + } + + parsed, err := url.Parse(trimmed) + if err != nil { + return "", "", fmt.Errorf("sqs: invalid QueueUrl: %w", err) + } + + segments := strings.Split(strings.Trim(parsed.Path, "/"), "/") + if len(segments) < 2 { + return "", "", fmt.Errorf("sqs: invalid QueueUrl: missing account or queue name") + } + accountID := segments[len(segments)-2] + queueName := segments[len(segments)-1] + if queueName == "" { + return "", "", fmt.Errorf("sqs: invalid QueueUrl: missing queue name") + } + return queueName, accountID, nil +} + +func isQueueScopedAction(action string) bool { + switch action { + case "AddPermission", "CancelMessageMoveTask", "ChangeMessageVisibility", "ChangeMessageVisibilityBatch", "DeleteMessage", "DeleteMessageBatch", "DeleteQueue", "GetQueueAttributes", "ListDeadLetterSourceQueues", "ListMessageMoveTasks", "ListQueueTags", "PurgeQueue", "ReceiveMessage", "RemovePermission", "SendMessage", "SendMessageBatch", "SetQueueAttributes", "StartMessageMoveTask", "TagQueue", "UntagQueue": + return true + default: + return false + } +} + +func queueContextFromValuesForAction(action string, values url.Values) (string, string, error) { + if queueName, accountID, err := queueNameFromQueueURL(values.Get("QueueUrl")); err == nil { + return queueName, accountID, nil + } + + switch action { + case "StartMessageMoveTask", "ListMessageMoveTasks", "ListDeadLetterSourceQueues": + queueName, accountID, err := queueNameFromQueueARN(values.Get("SourceArn")) + if err != nil { + return "", "", err + } + return queueName, accountID, nil + case "CancelMessageMoveTask": + queueName, accountID, err := queueNameFromTaskHandle(values.Get("TaskHandle")) + if err != nil { + return "", "", err + } + return queueName, accountID, nil + default: + queueName, accountID, err := queueNameFromQueueURL(values.Get("QueueUrl")) + if err != nil { + return "", "", err + } + return queueName, accountID, nil + } +} + +func queueNameFromQueueARN(raw string) (string, string, error) { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return "", "", fmt.Errorf("sqs: SourceArn is required") + } + + parts := strings.Split(trimmed, ":") + if len(parts) < 6 || !strings.HasPrefix(trimmed, "arn:") { + return "", "", fmt.Errorf("sqs: invalid SourceArn: %s", raw) + } + queueName := parts[len(parts)-1] + accountID := parts[len(parts)-2] + if queueName == "" || accountID == "" { + return "", "", fmt.Errorf("sqs: invalid SourceArn: %s", raw) + } + return queueName, accountID, nil +} + +func queueNameFromTaskHandle(raw string) (string, string, error) { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return "", "", fmt.Errorf("sqs: TaskHandle is required") + } + parts := strings.SplitN(trimmed, "|", 2) + if len(parts) == 0 || strings.TrimSpace(parts[0]) == "" { + return "", "", fmt.Errorf("sqs: invalid TaskHandle: %s", raw) + } + return queueNameFromQueueARN(parts[0]) +} + +func sendMessageRequestFromValues(values url.Values) contracts.SendMessageRequest { + return contracts.SendMessageRequest{ + DelaySeconds: parseIntValue(values.Get("DelaySeconds")), + MessageAttributes: messageAttributesFromValues(values, "MessageAttribute"), + MessageBody: values.Get("MessageBody"), + MessageDeduplicationId: values.Get("MessageDeduplicationId"), + MessageGroupId: values.Get("MessageGroupId"), + MessageSystemAttributes: messageAttributesFromValues(values, "MessageSystemAttribute"), + QueueUrl: values.Get("QueueUrl"), + } +} + +func sendMessageBatchRequestFromValues(values url.Values) contracts.SendMessageBatchRequest { + return contracts.SendMessageBatchRequest{ + Entries: sendMessageBatchEntriesFromValues(values), + QueueUrl: values.Get("QueueUrl"), + } +} + +func deleteMessageRequestFromValues(values url.Values) contracts.DeleteMessageRequest { + return contracts.DeleteMessageRequest{ + QueueUrl: values.Get("QueueUrl"), + ReceiptHandle: values.Get("ReceiptHandle"), + } +} + +func deleteMessageBatchRequestFromValues(values url.Values) contracts.DeleteMessageBatchRequest { + return contracts.DeleteMessageBatchRequest{ + Entries: deleteMessageBatchEntriesFromValues(values), + QueueUrl: values.Get("QueueUrl"), + } +} + +func changeMessageVisibilityRequestFromValues(values url.Values) contracts.ChangeMessageVisibilityRequest { + return contracts.ChangeMessageVisibilityRequest{ + QueueUrl: values.Get("QueueUrl"), + ReceiptHandle: values.Get("ReceiptHandle"), + VisibilityTimeout: parseIntValue(values.Get("VisibilityTimeout")), + } +} + +func changeMessageVisibilityBatchRequestFromValues(values url.Values) contracts.ChangeMessageVisibilityBatchRequest { + return contracts.ChangeMessageVisibilityBatchRequest{ + Entries: changeMessageVisibilityBatchEntriesFromValues(values), + QueueUrl: values.Get("QueueUrl"), + } +} + +func receiveMessageRequestFromValues(values url.Values) contracts.ReceiveMessageRequest { + return contracts.ReceiveMessageRequest{ + AttributeNames: queueAttributeNames(values), + MaxNumberOfMessages: parseIntValue(values.Get("MaxNumberOfMessages")), + MessageAttributeNames: listValuesFromPrefix(values, "MessageAttributeName"), + MessageSystemAttributeNames: listValuesFromPrefix(values, "MessageSystemAttributeName"), + QueueUrl: values.Get("QueueUrl"), + VisibilityTimeout: parseIntValue(values.Get("VisibilityTimeout")), + WaitTimeSeconds: parseIntValue(values.Get("WaitTimeSeconds")), + } +} + +func tagQueueTagsFromValues(values url.Values) map[string]string { + type queueTag struct { + key string + value string + } + + direct := map[string]string{} + byIndex := map[int]*queueTag{} + for key, list := range values { + if len(list) == 0 { + continue + } + + parts := strings.Split(key, ".") + if len(parts) == 2 && strings.EqualFold(parts[0], "Tags") { + direct[strings.TrimSpace(parts[1])] = list[0] + continue + } + if len(parts) < 4 || !strings.EqualFold(parts[0], "Tags") || !strings.EqualFold(parts[1], "entry") { + continue + } + index, err := strconv.Atoi(parts[2]) + if err != nil || index <= 0 { + continue + } + entry := byIndex[index] + if entry == nil { + entry = &queueTag{} + byIndex[index] = entry + } + switch strings.ToLower(parts[3]) { + case "key": + entry.key = strings.TrimSpace(list[0]) + case "value": + entry.value = list[0] + } + } + + if len(byIndex) == 0 { + if len(direct) == 0 { + return nil + } + return direct + } + + indices := make([]int, 0, len(byIndex)) + for index := range byIndex { + indices = append(indices, index) + } + sort.Ints(indices) + + tags := make(map[string]string, len(indices)) + for _, index := range indices { + entry := byIndex[index] + if entry == nil || entry.key == "" { + continue + } + tags[entry.key] = entry.value + } + if len(tags) == 0 { + return nil + } + for key, value := range direct { + tags[key] = value + } + return tags +} + +func queueTagKeysFromValues(values url.Values) []string { + keys := make([]string, 0) + for _, raw := range append(values["TagKey"], values["TagKeys"]...) { + if trimmed := strings.TrimSpace(raw); trimmed != "" { + keys = append(keys, trimmed) + } + } + keys = append(keys, listValuesFromPrefix(values, "TagKey")...) + keys = append(keys, listValuesFromPrefix(values, "TagKeys")...) + sort.Strings(keys) + keys = uniqueStringSlice(keys) + return keys +} + +func permissionAccountsFromValues(values url.Values) []string { + return uniqueStringSlice(append(listValuesFromPrefix(values, "AWSAccountIds"), values["AWSAccountIds"]...)) +} + +func permissionActionsFromValues(values url.Values) []string { + return uniqueStringSlice(append(listValuesFromPrefix(values, "Actions"), values["Actions"]...)) +} + +func uniqueStringSlice(values []string) []string { + if len(values) == 0 { + return nil + } + seen := make(map[string]struct{}, len(values)) + for _, value := range values { + if trimmed := strings.TrimSpace(value); trimmed != "" { + seen[trimmed] = struct{}{} + } + } + if len(seen) == 0 { + return nil + } + ordered := make([]string, 0, len(seen)) + for value := range seen { + ordered = append(ordered, value) + } + sort.Strings(ordered) + return ordered +} + +func startMessageMoveTaskRequestFromValues(values url.Values) (string, string, int) { + return strings.TrimSpace(values.Get("SourceArn")), strings.TrimSpace(values.Get("DestinationArn")), parseIntValue(values.Get("MaxNumberOfMessagesPerSecond")) +} + +func cancelMessageMoveTaskRequestFromValues(values url.Values) string { + return strings.TrimSpace(values.Get("TaskHandle")) +} + +func listMessageMoveTasksRequestFromValues(values url.Values) (string, int) { + return strings.TrimSpace(values.Get("SourceArn")), parseIntValue(values.Get("MaxResults")) +} + +func sendMessageBatchEntriesFromValues(values url.Values) []contracts.SendMessageBatchRequestEntry { + return batchEntriesFromValues(values, []string{"Entries", "SendMessageBatchRequestEntry"}, func(entry url.Values, item contracts.SendMessageBatchRequestEntry) contracts.SendMessageBatchRequestEntry { + item.Id = entry.Get("Id") + item.DelaySeconds = parseIntValue(entry.Get("DelaySeconds")) + item.MessageAttributes = messageAttributesFromValues(entry, "MessageAttribute") + item.MessageBody = entry.Get("MessageBody") + item.MessageDeduplicationId = entry.Get("MessageDeduplicationId") + item.MessageGroupId = entry.Get("MessageGroupId") + item.MessageSystemAttributes = messageAttributesFromValues(entry, "MessageSystemAttribute") + return item + }) +} + +func deleteMessageBatchEntriesFromValues(values url.Values) []contracts.DeleteMessageBatchRequestEntry { + return batchEntriesFromValues(values, []string{"Entries", "DeleteMessageBatchRequestEntry"}, func(entry url.Values, item contracts.DeleteMessageBatchRequestEntry) contracts.DeleteMessageBatchRequestEntry { + item.Id = entry.Get("Id") + item.ReceiptHandle = entry.Get("ReceiptHandle") + return item + }) +} + +func changeMessageVisibilityBatchEntriesFromValues(values url.Values) []contracts.ChangeMessageVisibilityBatchRequestEntry { + return batchEntriesFromValues(values, []string{"Entries", "ChangeMessageVisibilityBatchRequestEntry"}, func(entry url.Values, item contracts.ChangeMessageVisibilityBatchRequestEntry) contracts.ChangeMessageVisibilityBatchRequestEntry { + item.Id = entry.Get("Id") + item.ReceiptHandle = entry.Get("ReceiptHandle") + item.VisibilityTimeout = parseIntValue(entry.Get("VisibilityTimeout")) + return item + }) +} + +func batchEntriesFromValues[T any](values url.Values, prefixes []string, build func(url.Values, T) T) []T { + byIndex := map[int]url.Values{} + for key, list := range values { + if len(list) == 0 { + continue + } + parts := strings.Split(key, ".") + if len(parts) < 3 { + continue + } + if !hasAnyPrefix(parts[0], prefixes) { + continue + } + index, err := strconv.Atoi(parts[1]) + if err != nil || index <= 0 { + continue + } + entry := byIndex[index] + if entry == nil { + entry = url.Values{} + byIndex[index] = entry + } + entry.Set(strings.Join(parts[2:], "."), list[0]) + } + + if len(byIndex) == 0 { + return nil + } + + indices := make([]int, 0, len(byIndex)) + for index := range byIndex { + indices = append(indices, index) + } + sort.Ints(indices) + + result := make([]T, 0, len(indices)) + for _, index := range indices { + var item T + result = append(result, build(byIndex[index], item)) + } + return result +} + +func messageAttributesFromValues(values url.Values, prefix string) map[string]contracts.MessageAttributeValue { + type messageAttribute struct { + name string + value contracts.MessageAttributeValue + } + + byIndex := map[int]*messageAttribute{} + for key, list := range values { + if len(list) == 0 { + continue + } + parts := strings.Split(key, ".") + if len(parts) < 3 || !strings.EqualFold(parts[0], prefix) { + continue + } + index, err := strconv.Atoi(parts[1]) + if err != nil || index <= 0 { + continue + } + entry := byIndex[index] + if entry == nil { + entry = &messageAttribute{} + byIndex[index] = entry + } + switch strings.ToLower(parts[2]) { + case "name": + entry.name = list[0] + case "value": + if len(parts) < 4 { + continue + } + switch strings.ToLower(parts[3]) { + case "datatype": + entry.value.DataType = list[0] + case "stringvalue": + entry.value.StringValue = list[0] + } + } + } + + if len(byIndex) == 0 { + return nil + } + + indices := make([]int, 0, len(byIndex)) + for index := range byIndex { + indices = append(indices, index) + } + sort.Ints(indices) + + result := map[string]contracts.MessageAttributeValue{} + for _, index := range indices { + entry := byIndex[index] + if entry == nil || trimSpace(entry.name) == "" { + continue + } + result[entry.name] = entry.value + } + if len(result) == 0 { + return nil + } + return result +} + +func listValuesFromPrefix(values url.Values, prefix string) []string { + items := make([]string, 0) + for key, list := range values { + if len(list) == 0 { + continue + } + if !strings.EqualFold(strings.Split(key, ".")[0], prefix) { + continue + } + items = append(items, list[0]) + } + sort.Strings(items) + return items +} + +func hasAnyPrefix(value string, prefixes []string) bool { + for _, prefix := range prefixes { + if strings.EqualFold(value, prefix) { + return true + } + } + return false +} + +func flattenJSONEntry(dst url.Values, index int, entry map[string]any) { + if index <= 0 || len(entry) == 0 { + return + } + for key, value := range entry { + switch typed := value.(type) { + case map[string]any: + if strings.EqualFold(key, "MessageAttributes") { + flattenIndexedMessageAttributes(dst, index, "MessageAttribute", typed) + continue + } + if strings.EqualFold(key, "MessageSystemAttributes") { + flattenIndexedMessageAttributes(dst, index, "MessageSystemAttribute", typed) + continue + } + for nestedKey, nestedValue := range typed { + dst.Set(fmt.Sprintf("Entries.%d.%s.%s", index, key, nestedKey), fmt.Sprint(nestedValue)) + } + default: + dst.Set(fmt.Sprintf("Entries.%d.%s", index, key), fmt.Sprint(value)) + } + } +} + +func flattenQueueAttributes(dst url.Values, typed map[string]any) { + keys := make([]string, 0, len(typed)) + for attrName := range typed { + keys = append(keys, attrName) + } + sort.Strings(keys) + for i, attrName := range keys { + dst.Set(fmt.Sprintf("Attribute.%d.Name", i+1), attrName) + dst.Set(fmt.Sprintf("Attribute.%d.Value.StringValue", i+1), fmt.Sprint(typed[attrName])) + } +} + +func flattenMessageAttributes(dst url.Values, prefix string, typed map[string]any) { + keys := make([]string, 0, len(typed)) + for attrName := range typed { + keys = append(keys, attrName) + } + sort.Strings(keys) + for i, attrName := range keys { + dst.Set(fmt.Sprintf("%s.%d.Name", prefix, i+1), attrName) + if nested, ok := typed[attrName].(map[string]any); ok { + flattenMessageAttributeValue(dst, prefix, i+1, nested) + continue + } + dst.Set(fmt.Sprintf("%s.%d.Value.StringValue", prefix, i+1), fmt.Sprint(typed[attrName])) + } +} + +func flattenIndexedMessageAttributes(dst url.Values, index int, prefix string, typed map[string]any) { + keys := make([]string, 0, len(typed)) + for attrName := range typed { + keys = append(keys, attrName) + } + sort.Strings(keys) + for i, attrName := range keys { + base := fmt.Sprintf("Entries.%d.%s.%d", index, prefix, i+1) + dst.Set(base+".Name", attrName) + if nested, ok := typed[attrName].(map[string]any); ok { + flattenMessageAttributeValue(dst, base, 0, nested) + continue + } + dst.Set(base+".Value.StringValue", fmt.Sprint(typed[attrName])) + } +} + +func flattenMessageAttributeValue(dst url.Values, base string, index int, typed map[string]any) { + for key, value := range typed { + switch { + case strings.EqualFold(key, "DataType"): + if index > 0 { + dst.Set(fmt.Sprintf("%s.%d.Value.DataType", base, index), fmt.Sprint(value)) + } else { + dst.Set(base+".Value.DataType", fmt.Sprint(value)) + } + case strings.EqualFold(key, "StringValue"): + if index > 0 { + dst.Set(fmt.Sprintf("%s.%d.Value.StringValue", base, index), fmt.Sprint(value)) + } else { + dst.Set(base+".Value.StringValue", fmt.Sprint(value)) + } + case strings.EqualFold(key, "BinaryValue"): + if index > 0 { + dst.Set(fmt.Sprintf("%s.%d.Value.BinaryValue", base, index), fmt.Sprint(value)) + } else { + dst.Set(base+".Value.BinaryValue", fmt.Sprint(value)) + } + } + } +} + +func parseIntValue(raw string) int { + raw = strings.TrimSpace(raw) + if raw == "" { + return 0 + } + value, err := strconv.Atoi(raw) + if err != nil { + return 0 + } + return value +} + +func trimSpace(value string) string { + return strings.TrimSpace(value) +} + +func queueAttributesFromValues(values url.Values) map[string]string { + type queueAttribute struct { + name string + value string + } + + byIndex := make(map[int]*queueAttribute) + for key, list := range values { + if !strings.HasPrefix(key, "Attribute.") || len(list) == 0 { + continue + } + + segments := strings.Split(key, ".") + if len(segments) < 3 { + continue + } + + index, err := strconv.Atoi(segments[1]) + if err != nil { + continue + } + + entry, ok := byIndex[index] + if !ok { + entry = &queueAttribute{} + byIndex[index] = entry + } + + switch segments[2] { + case "Name": + entry.name = strings.TrimSpace(list[0]) + case "Value": + if len(segments) >= 4 && segments[3] == "StringValue" { + entry.value = list[0] + } + } + } + + if len(byIndex) == 0 { + return nil + } + + indices := make([]int, 0, len(byIndex)) + for index := range byIndex { + indices = append(indices, index) + } + sort.Ints(indices) + + attributes := make(map[string]string, len(indices)) + for _, index := range indices { + entry := byIndex[index] + if entry == nil || entry.name == "" { + continue + } + attributes[entry.name] = entry.value + } + if len(attributes) == 0 { + return nil + } + return attributes +} diff --git a/core/internal/delivery/http/sqs_native_contract_test.go b/core/internal/delivery/http/sqs_native_contract_test.go index 6716577..fdc5408 100644 --- a/core/internal/delivery/http/sqs_native_contract_test.go +++ b/core/internal/delivery/http/sqs_native_contract_test.go @@ -11,7 +11,7 @@ func TestSQSNativeContractParsesQueryAndFormValues(t *testing.T) { t.Helper() req := httptest.NewRequest(http.MethodPost, "/123456789012/orders/", strings.NewReader( - "Action=SendMessage&Version=2012-11-05&QueueUrl=https%3A%2F%2Flocalhost%2F123456789012%2Forders&Attribute.1.Name=DelaySeconds&Attribute.1.Value.StringValue=5", + "Action=SendMessage&Version=2012-11-05&QueueUrl=https%3A%2F%2Flocalhost%2F123456789012%2Forders&QueueNamePrefix=ord&QueueOwnerAWSAccountId=123456789012&Attribute.1.Name=DelaySeconds&Attribute.1.Value.StringValue=5", )) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") @@ -46,6 +46,126 @@ func TestSQSNativeContractParsesQueryAndFormValues(t *testing.T) { if got, want := ctx.Values.Get("QueueUrl"), "https://localhost/123456789012/orders"; got != want { t.Fatalf("unexpected queue url: got %q want %q", got, want) } + if got, want := ctx.Values.Get("QueueNamePrefix"), "ord"; got != want { + t.Fatalf("unexpected queue name prefix: got %q want %q", got, want) + } + if got, want := ctx.Values.Get("QueueOwnerAWSAccountId"), "123456789012"; got != want { + t.Fatalf("unexpected queue owner account id: got %q want %q", got, want) + } +} + +func TestSQSNativeContractParsesTargetStyleJsonRequests(t *testing.T) { + t.Helper() + + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{"QueueNamePrefix":"ord","MaxResults":2,"NextToken":"token-1"}`)) + req.Header.Set("Content-Type", "application/x-amz-json-1.0") + req.Header.Set("X-Amz-Target", "AmazonSQS.ListQueues") + + ctx, err := ParseSQSRequest(req) + if err != nil { + t.Fatalf("parse target-style request: %v", err) + } + if !ctx.TargetStyle { + t.Fatal("expected target-style request to be marked as such") + } + if got, want := ctx.Action, "ListQueues"; got != want { + t.Fatalf("unexpected action: got %q want %q", got, want) + } + if got, want := ctx.Version, sqsQueryVersion; got != want { + t.Fatalf("unexpected version: got %q want %q", got, want) + } + if got, want := ctx.Values.Get("QueueNamePrefix"), "ord"; got != want { + t.Fatalf("unexpected queue name prefix: got %q want %q", got, want) + } + if got, want := ctx.Values.Get("MaxResults"), "2"; got != want { + t.Fatalf("unexpected max results: got %q want %q", got, want) + } + if got, want := ctx.Values.Get("NextToken"), "token-1"; got != want { + t.Fatalf("unexpected next token: got %q want %q", got, want) + } +} + +func TestSQSNativeContractInfersQueueContextFromTargetStyleQueueURL(t *testing.T) { + t.Helper() + + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{"QueueUrl":"https://sqs.us-east-1.amazonaws.com/123456789012/orders","Attributes":{"DelaySeconds":"0"}}`)) + req.Header.Set("Content-Type", "application/x-amz-json-1.0") + req.Header.Set("X-Amz-Target", "AmazonSQS.SetQueueAttributes") + + ctx, err := ParseSQSRequest(req) + if err != nil { + t.Fatalf("parse target-style queue request: %v", err) + } + if got, want := ctx.Kind, SQSRequestKindQueue; got != want { + t.Fatalf("unexpected kind: got %q want %q", got, want) + } + if got, want := ctx.QueueName, "orders"; got != want { + t.Fatalf("unexpected queue name: got %q want %q", got, want) + } + if got, want := ctx.AccountID, "123456789012"; got != want { + t.Fatalf("unexpected account id: got %q want %q", got, want) + } + if got, want := ctx.Values.Get("Attribute.1.Name"), "DelaySeconds"; got != want { + t.Fatalf("unexpected attribute name: got %q want %q", got, want) + } + if got, want := ctx.Values.Get("Attribute.1.Value.StringValue"), "0"; got != want { + t.Fatalf("unexpected attribute value: got %q want %q", got, want) + } +} + +func TestSQSNativeContractParsesTargetStyleMessageAttributes(t *testing.T) { + t.Helper() + + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{ + "QueueUrl":"https://sqs.us-east-1.amazonaws.com/123456789012/orders", + "MessageBody":"hello", + "MessageAttributes":{ + "Author":{"DataType":"String","StringValue":"MildStack"} + } + }`)) + req.Header.Set("Content-Type", "application/x-amz-json-1.0") + req.Header.Set("X-Amz-Target", "AmazonSQS.SendMessage") + + ctx, err := ParseSQSRequest(req) + if err != nil { + t.Fatalf("parse target-style message attributes request: %v", err) + } + attrs := messageAttributesFromValues(ctx.Values, "MessageAttribute") + if got, want := attrs["Author"].DataType, "String"; got != want { + t.Fatalf("unexpected attribute data type: got %q want %q", got, want) + } + if got, want := attrs["Author"].StringValue, "MildStack"; got != want { + t.Fatalf("unexpected attribute string value: got %q want %q", got, want) + } +} + +func TestSQSNativeContractParsesTargetStyleBatchEntries(t *testing.T) { + t.Helper() + + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{ + "QueueUrl":"https://sqs.us-east-1.amazonaws.com/123456789012/orders", + "Entries":[ + {"Id":"msg1","MessageBody":"one"}, + {"Id":"msg2","MessageBody":"two"} + ] + }`)) + req.Header.Set("Content-Type", "application/x-amz-json-1.0") + req.Header.Set("X-Amz-Target", "AmazonSQS.SendMessageBatch") + + ctx, err := ParseSQSRequest(req) + if err != nil { + t.Fatalf("parse target-style batch request: %v", err) + } + entries := sendMessageBatchEntriesFromValues(ctx.Values) + if got, want := len(entries), 2; got != want { + t.Fatalf("unexpected batch entry count: got %d want %d", got, want) + } + if got, want := entries[0].Id, "msg1"; got != want { + t.Fatalf("unexpected first entry id: got %q want %q", got, want) + } + if got, want := entries[1].MessageBody, "two"; got != want { + t.Fatalf("unexpected second entry body: got %q want %q", got, want) + } } func TestSQSNativeContractClassifiesRootAndQueuePaths(t *testing.T) { diff --git a/core/internal/delivery/http/sqs_native_errors.go b/core/internal/delivery/http/sqs_native_errors.go index 04e7d18..cfeadbd 100644 --- a/core/internal/delivery/http/sqs_native_errors.go +++ b/core/internal/delivery/http/sqs_native_errors.go @@ -58,6 +58,16 @@ func classifySQSError(err error) (int, string, string) { return http.StatusBadRequest, "InvalidAddress", "The specified queue path is invalid for the requested action." case errors.Is(err, ErrSQSUnsupported): return http.StatusBadRequest, "UnsupportedOperation", "The requested operation is not supported by the local subset." + case strings.Contains(strings.ToLower(err.Error()), "batch request is empty"): + return http.StatusBadRequest, "EmptyBatchRequest", "The batch request doesn't contain any entries." + case strings.Contains(strings.ToLower(err.Error()), "more than 10 entries"): + return http.StatusBadRequest, "TooManyEntriesInBatchRequest", "The batch request contains more entries than permissible." + case strings.Contains(strings.ToLower(err.Error()), "duplicate entry ids"): + return http.StatusBadRequest, "BatchEntryIdsNotDistinct", "Two or more batch entries in the request have the same Id." + case strings.Contains(strings.ToLower(err.Error()), "queue not found"): + return http.StatusBadRequest, "QueueDoesNotExist", "Ensure that the QueueUrl is correct and that the queue has not been deleted." + case strings.Contains(strings.ToLower(err.Error()), "receipt handle does not match active lease"): + return http.StatusBadRequest, "ReceiptHandleIsInvalid", "The specified receipt handle isn't valid." default: return http.StatusBadRequest, "ValidationError", err.Error() } diff --git a/core/internal/delivery/http/sqs_native_registry.go b/core/internal/delivery/http/sqs_native_registry.go index 2f40c05..cf8408b 100644 --- a/core/internal/delivery/http/sqs_native_registry.go +++ b/core/internal/delivery/http/sqs_native_registry.go @@ -13,6 +13,7 @@ type SQSRegistrySpec struct { Version string Supported bool DomainDeferred bool + MessageSurface bool ReturnsQueueURL bool UsesQueueContext bool } @@ -28,12 +29,14 @@ func NewSQSRegistry() SQSRegistry { byName := make(map[string]SQSRegistrySpec, len(specs)) for _, spec := range specs { + supported := isQueueLifecycleAction(spec.Action) || isQueueGovernanceAction(spec.Action) || isQueueRedriveAction(spec.Action) || spec.MessageSurface entry := SQSRegistrySpec{ Action: spec.Action, Scope: spec.Scope, Version: spec.Version, - Supported: true, - DomainDeferred: true, + Supported: supported, + DomainDeferred: !supported, + MessageSurface: spec.MessageSurface, ReturnsQueueURL: spec.ReturnsQueueURL, UsesQueueContext: spec.UsesQueueContext, } @@ -91,6 +94,33 @@ func (r SQSRegistry) String() string { return fmt.Sprintf("sqs registry: %d actions", len(r.ordered)) } +func isQueueLifecycleAction(action string) bool { + switch action { + case "CreateQueue", "DeleteQueue", "GetQueueAttributes", "GetQueueUrl", "ListQueues", "PurgeQueue", "SetQueueAttributes": + return true + default: + return false + } +} + +func isQueueGovernanceAction(action string) bool { + switch action { + case "AddPermission", "RemovePermission", "TagQueue", "UntagQueue", "ListQueueTags": + return true + default: + return false + } +} + +func isQueueRedriveAction(action string) bool { + switch action { + case "ListDeadLetterSourceQueues", "StartMessageMoveTask", "CancelMessageMoveTask", "ListMessageMoveTasks": + return true + default: + return false + } +} + func isSQSErrorCode(err error, target error) bool { return errors.Is(err, target) } diff --git a/core/internal/delivery/http/sqs_native_registry_test.go b/core/internal/delivery/http/sqs_native_registry_test.go index fea50bd..d2104cf 100644 --- a/core/internal/delivery/http/sqs_native_registry_test.go +++ b/core/internal/delivery/http/sqs_native_registry_test.go @@ -1,8 +1,13 @@ package http import ( + "net/http" + "net/http/httptest" + "strings" "testing" + "github.com/gin-gonic/gin" + "github.com/michasdev/mildstack/core/internal/application/orchestrator" "github.com/michasdev/mildstack/core/internal/resources/sqs/contracts" ) @@ -17,15 +22,64 @@ func TestSQSNativeRegistryDerivesSpecsFromCatalog(t *testing.T) { t.Fatalf("unexpected registry entry count: got %d want %d", got, want) } + supportedActions := map[string]struct{}{ + "AddPermission": {}, + "CancelMessageMoveTask": {}, + "ChangeMessageVisibility": {}, + "ChangeMessageVisibilityBatch": {}, + "CreateQueue": {}, + "DeleteMessage": {}, + "DeleteMessageBatch": {}, + "DeleteQueue": {}, + "GetQueueAttributes": {}, + "GetQueueUrl": {}, + "ListDeadLetterSourceQueues": {}, + "ListMessageMoveTasks": {}, + "ListQueueTags": {}, + "ListQueues": {}, + "PurgeQueue": {}, + "ReceiveMessage": {}, + "RemovePermission": {}, + "SendMessage": {}, + "SendMessageBatch": {}, + "SetQueueAttributes": {}, + "StartMessageMoveTask": {}, + "TagQueue": {}, + "UntagQueue": {}, + } + deferredActions := map[string]struct{}{} + for i, spec := range entries { if spec.Action != catalog[i].Action { t.Fatalf("unexpected action at %d: got %q want %q", i, spec.Action, catalog[i].Action) } - if !spec.Supported { - t.Fatalf("expected action %s to be transport-supported", spec.Action) + + _, supported := supportedActions[spec.Action] + _, deferred := deferredActions[spec.Action] + if supported && deferred { + t.Fatalf("action %s cannot be both supported and deferred", spec.Action) + } + if supported { + if !spec.Supported { + t.Fatalf("expected action %s to be transport-supported", spec.Action) + } + if spec.DomainDeferred { + t.Fatalf("expected action %s to be routed to the service, not deferred", spec.Action) + } } - if !spec.DomainDeferred { - t.Fatalf("expected action %s to be domain deferred", spec.Action) + if deferred { + if spec.Supported { + t.Fatalf("expected action %s to remain deferred", spec.Action) + } + if !spec.DomainDeferred { + t.Fatalf("expected action %s to be deferred", spec.Action) + } + } + if (isQueueLifecycleAction(spec.Action) || isQueueGovernanceAction(spec.Action) || isQueueRedriveAction(spec.Action)) && spec.MessageSurface { + t.Fatalf("did not expect lifecycle action %s to be marked as message surface", spec.Action) + } + if !spec.Supported && !spec.DomainDeferred { + t.Fatalf("expected action %s to remain deferred", spec.Action) } if spec.Scope != catalog[i].Scope { t.Fatalf("unexpected scope for %s: got %q want %q", spec.Action, spec.Scope, catalog[i].Scope) @@ -33,6 +87,75 @@ func TestSQSNativeRegistryDerivesSpecsFromCatalog(t *testing.T) { } } +func TestSQSNativeRegistrySeparatesSupportedAndDeferredActions(t *testing.T) { + t.Helper() + + registry := NewSQSRegistry() + + if got, want := registry.SupportedActions(), []string{ + "AddPermission", + "CancelMessageMoveTask", + "ChangeMessageVisibility", + "ChangeMessageVisibilityBatch", + "CreateQueue", + "DeleteMessage", + "DeleteMessageBatch", + "DeleteQueue", + "GetQueueAttributes", + "GetQueueUrl", + "ListDeadLetterSourceQueues", + "ListMessageMoveTasks", + "ListQueues", + "ListQueueTags", + "PurgeQueue", + "ReceiveMessage", + "RemovePermission", + "SendMessage", + "SendMessageBatch", + "SetQueueAttributes", + "StartMessageMoveTask", + "TagQueue", + "UntagQueue", + }; !equalStringSlicesSQS(got, want) { + t.Fatalf("unexpected supported actions: got %v want %v", got, want) + } + + if got, want := registry.UnsupportedActions(), []string{}; !equalStringSlicesSQS(got, want) { + t.Fatalf("unexpected unsupported actions: got %v want %v", got, want) + } +} + +func TestSQSNativeRegistryRecognizesQueueLifecycleActions(t *testing.T) { + t.Helper() + + for _, action := range []string{"CreateQueue", "DeleteQueue", "GetQueueAttributes", "GetQueueUrl", "ListQueues", "PurgeQueue", "SetQueueAttributes"} { + if !isQueueLifecycleAction(action) { + t.Fatalf("expected %s to be recognized as a queue lifecycle action", action) + } + } + if isQueueLifecycleAction("SendMessage") { + t.Fatal("did not expect SendMessage to be treated as a queue lifecycle action") + } +} + +func TestSQSNativeRegistryMarksPhase39MessageSurfaceActions(t *testing.T) { + t.Helper() + + registry := NewSQSRegistry() + for _, action := range []string{"ChangeMessageVisibility", "ChangeMessageVisibilityBatch", "DeleteMessage", "DeleteMessageBatch", "ReceiveMessage", "SendMessage", "SendMessageBatch"} { + spec, ok := registry.Lookup(action) + if !ok { + t.Fatalf("expected registry action %s", action) + } + if !spec.MessageSurface { + t.Fatalf("expected %s to be marked as message surface", action) + } + if spec.DomainDeferred { + t.Fatalf("expected %s to be routed away from domain deferral", action) + } + } +} + func TestSQSNativeRegistryScopeMismatchesMapToExplicitErrors(t *testing.T) { t.Helper() @@ -67,3 +190,47 @@ func TestSQSNativeRegistryRejectsUnknownAction(t *testing.T) { t.Fatalf("unexpected unknown action error: got %v want %v", err, ErrSQSInvalidAction) } } + +func TestSQSNativeRegistryRoutesMessageActionsEvenWhenPolicyOmitsThem(t *testing.T) { + t.Helper() + + gin.SetMode(gin.TestMode) + service := ®istryPolicyTrimmedService{} + router := gin.New() + RegisterSQSNativeRoutes(router, service) + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/123456789012/orders/", strings.NewReader("Action=SendMessage&Version=2012-11-05&MessageBody=hello")) + request.Header.Set("Content-Type", "application/x-www-form-urlencoded") + router.ServeHTTP(recorder, request) + + if got, want := recorder.Code, http.StatusOK; got != want { + t.Fatalf("unexpected send message status: got %d want %d", got, want) + } + if got, want := service.sendMessageQueueName, "orders"; got != want { + t.Fatalf("unexpected queue name captured by service: got %q want %q", got, want) + } + if !strings.Contains(recorder.Body.String(), "SendMessageResponse") { + t.Fatalf("expected send message response, got %q", recorder.Body.String()) + } +} + +func equalStringSlicesSQS(got, want []string) bool { + if len(got) != len(want) { + return false + } + for i := range got { + if got[i] != want[i] { + return false + } + } + return true +} + +type registryPolicyTrimmedService struct { + stubSQSNativeService +} + +func (s *registryPolicyTrimmedService) Policy() orchestrator.EmulationPolicy { + return orchestrator.NewEmulationPolicy(orchestrator.FidelityExemplar, []string{"ListQueues"}, nil, "sqs") +} diff --git a/core/internal/delivery/http/sqs_native_test.go b/core/internal/delivery/http/sqs_native_test.go index 7e17cf0..ca510a9 100644 --- a/core/internal/delivery/http/sqs_native_test.go +++ b/core/internal/delivery/http/sqs_native_test.go @@ -3,6 +3,7 @@ package http import ( "bytes" "context" + "encoding/json" "io" "net/http" "net/http/httptest" @@ -15,9 +16,12 @@ import ( "github.com/aws/aws-sdk-go-v2/credentials" "github.com/aws/aws-sdk-go-v2/service/sqs" "github.com/gin-gonic/gin" + "github.com/michasdev/mildstack/core/internal/application/orchestrator" "github.com/michasdev/mildstack/core/internal/application/runtime" "github.com/michasdev/mildstack/core/internal/composition" - sqsresource "github.com/michasdev/mildstack/core/internal/resources/sqs" + sqsapplication "github.com/michasdev/mildstack/core/internal/resources/sqs/application" + "github.com/michasdev/mildstack/core/internal/resources/sqs/contracts" + "github.com/michasdev/mildstack/core/internal/resources/sqs/domain" ) func TestSQSNativeMiddlewareInterceptsQueryRequestsAndLeavesRuntimeRoutesUntouched(t *testing.T) { @@ -28,7 +32,8 @@ func TestSQSNativeMiddlewareInterceptsQueryRequestsAndLeavesRuntimeRoutesUntouch manager := runtime.New(root.Services) router := NewRouter(DefaultConfig(), manager) - RegisterSQSNativeRoutes(router.Engine(), sqsresource.New()) + service := sqsapplication.New() + RegisterSQSNativeRoutes(router.Engine(), service) healthRecorder := httptest.NewRecorder() healthRequest := httptest.NewRequest(http.MethodGet, "/api/v1/runtime/health", nil) @@ -40,25 +45,28 @@ func TestSQSNativeMiddlewareInterceptsQueryRequestsAndLeavesRuntimeRoutesUntouch rootRecorder := httptest.NewRecorder() rootRequest := httptest.NewRequest(http.MethodGet, "/?Action=ListQueues&Version=2012-11-05", nil) router.Engine().ServeHTTP(rootRecorder, rootRequest) - if got, want := rootRecorder.Code, http.StatusBadRequest; got != want { + if got, want := rootRecorder.Code, http.StatusOK; got != want { t.Fatalf("unexpected root status: got %d want %d", got, want) } if ct := rootRecorder.Header().Get("Content-Type"); !strings.Contains(ct, "application/xml") { t.Fatalf("unexpected root content type: got %q", ct) } - if !strings.Contains(rootRecorder.Body.String(), "UnsupportedOperation") { - t.Fatalf("expected unsupported operation XML, got %q", rootRecorder.Body.String()) + if !strings.Contains(rootRecorder.Body.String(), "ListQueuesResponse") { + t.Fatalf("expected list queues XML, got %q", rootRecorder.Body.String()) } + if _, err := service.CreateQueue("orders", nil); err != nil { + t.Fatalf("create queue: %v", err) + } queueRecorder := httptest.NewRecorder() queueRequest := httptest.NewRequest(http.MethodPost, "/123456789012/orders/", strings.NewReader("Action=SendMessage&Version=2012-11-05&MessageBody=hello")) queueRequest.Header.Set("Content-Type", "application/x-www-form-urlencoded") router.Engine().ServeHTTP(queueRecorder, queueRequest) - if got, want := queueRecorder.Code, http.StatusBadRequest; got != want { + if got, want := queueRecorder.Code, http.StatusOK; got != want { t.Fatalf("unexpected queue status: got %d want %d", got, want) } - if !strings.Contains(queueRecorder.Body.String(), "UnsupportedOperation") { - t.Fatalf("expected unsupported operation XML, got %q", queueRecorder.Body.String()) + if !strings.Contains(queueRecorder.Body.String(), "SendMessageResponse") { + t.Fatalf("expected send message XML, got %q", queueRecorder.Body.String()) } mismatchRecorder := httptest.NewRecorder() @@ -72,14 +80,64 @@ func TestSQSNativeMiddlewareInterceptsQueryRequestsAndLeavesRuntimeRoutesUntouch } } -func TestSQSSDKSmokeReceivesAWSCompatibleXMLError(t *testing.T) { +func TestSQSNativeMessageActionsPreserveAWSRequestNames(t *testing.T) { + t.Helper() + + gin.SetMode(gin.TestMode) + service := &stubSQSNativeService{ + sendMessageResult: contracts.SendMessageResult{ + MessageId: "message-1", + MD5OfMessageBody: "md5-body", + }, + } + router := gin.New() + RegisterSQSNativeRoutes(router, service) + + payload := map[string]any{ + "QueueUrl": service.QueueURL("orders"), + "MessageBody": "hello", + "DelaySeconds": 5, + } + body, err := json.Marshal(payload) + if err != nil { + t.Fatalf("marshal payload: %v", err) + } + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body)) + request.Header.Set("Content-Type", "application/x-amz-json-1.0") + request.Header.Set("X-Amz-Target", "AmazonSQS.SendMessage") + router.ServeHTTP(recorder, request) + + if got, want := recorder.Code, http.StatusOK; got != want { + t.Fatalf("unexpected send status: got %d want %d", got, want) + } + if got, want := service.sendMessageQueueName, "orders"; got != want { + t.Fatalf("unexpected send queue name: got %q want %q", got, want) + } + if got, want := service.sendMessageRequest.MessageBody, "hello"; got != want { + t.Fatalf("unexpected send message body: got %q want %q", got, want) + } + if got, want := service.sendMessageRequest.DelaySeconds, 5; got != want { + t.Fatalf("unexpected send delay seconds: got %d want %d", got, want) + } + if ct := recorder.Header().Get("Content-Type"); !strings.Contains(ct, "application/json") { + t.Fatalf("expected json content type, got %q", ct) + } + if !strings.Contains(recorder.Body.String(), "\"MessageId\"") { + t.Fatalf("expected send message json, got %q", recorder.Body.String()) + } +} + +func TestSQSSDKSmokeReceivesAWSCompatibleSuccess(t *testing.T) { t.Helper() gin.SetMode(gin.TestMode) root := composition.DefaultRoot("test-instance") manager := runtime.New(root.Services) router := NewRouter(DefaultConfig(), manager) - RegisterSQSNativeRoutes(router.Engine(), sqsresource.New()) + service := sqsapplication.New() + RegisterSQSNativeRoutes(router.Engine(), service) server := httptest.NewServer(router.Engine()) t.Cleanup(server.Close) @@ -101,16 +159,548 @@ func TestSQSSDKSmokeReceivesAWSCompatibleXMLError(t *testing.T) { o.BaseEndpoint = aws.String(server.URL) }) + if _, err := service.CreateQueue("orders", nil); err != nil { + t.Fatalf("create queue: %v", err) + } + _, err = client.ListQueues(ctx, &sqs.ListQueuesInput{}) - if err == nil { - t.Fatal("expected list queues to return an AWS-compatible error") + if err != nil { + t.Fatalf("expected list queues to return successfully: %v", err) } - if !strings.Contains(string(transport.body), "") { - t.Fatalf("expected captured xml error body, got %q", string(transport.body)) + if !strings.Contains(string(transport.body), "\"QueueUrls\"") { + t.Fatalf("expected captured json body, got %q", string(transport.body)) } - if !strings.Contains(string(transport.body), "UnsupportedOperation") && !strings.Contains(string(transport.body), "InvalidQueryParameter") { - t.Fatalf("expected captured xml error body to contain an SQS error code, got %q", string(transport.body)) + + sendOutput, err := client.SendMessage(ctx, &sqs.SendMessageInput{ + QueueUrl: aws.String(service.QueueURL("orders")), + MessageBody: aws.String("hello"), + }) + if err != nil { + t.Fatalf("expected send message to return successfully: %v", err) + } + if sendOutput.MessageId == nil || *sendOutput.MessageId == "" { + t.Fatal("expected send message output to include a message id") + } + + receiveOutput, err := client.ReceiveMessage(ctx, &sqs.ReceiveMessageInput{ + QueueUrl: aws.String(service.QueueURL("orders")), + MaxNumberOfMessages: 1, + }) + if err != nil { + t.Fatalf("expected receive message to return successfully: %v", err) + } + if got, want := len(receiveOutput.Messages), 1; got != want { + t.Fatalf("unexpected receive count: got %d want %d", got, want) + } + if got, want := aws.ToString(receiveOutput.Messages[0].Body), "hello"; got != want { + t.Fatalf("unexpected receive body: got %q want %q", got, want) + } + + receiptHandle := aws.ToString(receiveOutput.Messages[0].ReceiptHandle) + if receiptHandle == "" { + t.Fatal("expected receive message to include a receipt handle") + } + + if _, err := client.ChangeMessageVisibility(ctx, &sqs.ChangeMessageVisibilityInput{ + QueueUrl: aws.String(service.QueueURL("orders")), + ReceiptHandle: aws.String(receiptHandle), + VisibilityTimeout: 0, + }); err != nil { + t.Fatalf("expected change message visibility to return successfully: %v", err) + } + + visibleAgain, err := client.ReceiveMessage(ctx, &sqs.ReceiveMessageInput{ + QueueUrl: aws.String(service.QueueURL("orders")), + MaxNumberOfMessages: 1, + }) + if err != nil { + t.Fatalf("expected message to become visible again: %v", err) + } + if got, want := len(visibleAgain.Messages), 1; got != want { + t.Fatalf("unexpected redelivery count: got %d want %d", got, want) + } + if got, want := aws.ToString(visibleAgain.Messages[0].Body), "hello"; got != want { + t.Fatalf("unexpected redelivery body: got %q want %q", got, want) + } + + if _, err := client.DeleteMessage(ctx, &sqs.DeleteMessageInput{ + QueueUrl: aws.String(service.QueueURL("orders")), + ReceiptHandle: visibleAgain.Messages[0].ReceiptHandle, + }); err != nil { + t.Fatalf("expected delete message to return successfully: %v", err) + } + + cleared, err := client.ReceiveMessage(ctx, &sqs.ReceiveMessageInput{ + QueueUrl: aws.String(service.QueueURL("orders")), + MaxNumberOfMessages: 1, + }) + if err != nil { + t.Fatalf("expected empty queue receive to succeed: %v", err) + } + if got, want := len(cleared.Messages), 0; got != want { + t.Fatalf("unexpected post-delete receive count: got %d want %d", got, want) + } +} + +func TestSQSSDKSmokeCoversGovernanceAndRedriveActions(t *testing.T) { + t.Helper() + + gin.SetMode(gin.TestMode) + service := &stubSQSNativeService{ + listQueueTagsResult: map[string]string{ + "env": "dev", + }, + listDeadLetterSourceQueuesResult: []string{"orders-source"}, + startMessageMoveTaskResult: "arn:aws:sqs:us-east-1:123456789012:orders-dlq|task-1", + cancelMessageMoveTaskResult: 7, + listMessageMoveTasksResult: []domain.MessageMoveTask{ + { + TaskHandle: "arn:aws:sqs:us-east-1:123456789012:orders-dlq|task-1", + SourceArn: "arn:aws:sqs:us-east-1:123456789012:orders-dlq", + DestinationArn: "arn:aws:sqs:us-east-1:123456789012:orders", + MaxNumberOfMessagesPerSecond: 12, + ApproximateNumberOfMessagesMoved: 7, + Status: "RUNNING", + }, + }, } + router := gin.New() + RegisterSQSNativeRoutes(router, service) + + server := httptest.NewServer(router) + t.Cleanup(server.Close) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + t.Cleanup(cancel) + + cfg, err := awsconfig.LoadDefaultConfig(ctx, + awsconfig.WithRegion("us-east-1"), + awsconfig.WithCredentialsProvider(credentials.NewStaticCredentialsProvider("test", "test", "test")), + ) + if err != nil { + t.Fatalf("load aws config: %v", err) + } + + client := sqs.NewFromConfig(cfg, func(o *sqs.Options) { + o.BaseEndpoint = aws.String(server.URL) + }) + + queueURL := service.QueueURL("orders") + if _, err := client.TagQueue(ctx, &sqs.TagQueueInput{ + QueueUrl: aws.String(queueURL), + Tags: map[string]string{ + "env": "dev", + "team": "platform", + }, + }); err != nil { + t.Fatalf("tag queue: %v", err) + } + if got, want := service.tagQueueQueueName, "orders"; got != want { + t.Fatalf("unexpected tag queue name: got %q want %q", got, want) + } + if got, want := service.tagQueueTags["team"], "platform"; got != want { + t.Fatalf("unexpected tag queue payload: got %q want %q", got, want) + } + + tagOutput, err := client.ListQueueTags(ctx, &sqs.ListQueueTagsInput{ + QueueUrl: aws.String(queueURL), + }) + if err != nil { + t.Fatalf("list queue tags: %v", err) + } + if got, want := tagOutput.Tags["env"], "dev"; got != want { + t.Fatalf("unexpected tag list value: got %q want %q", got, want) + } + + if _, err := client.AddPermission(ctx, &sqs.AddPermissionInput{ + QueueUrl: aws.String(queueURL), + Label: aws.String("label-a"), + AWSAccountIds: []string{"123456789012"}, + Actions: []string{"SendMessage"}, + }); err != nil { + t.Fatalf("add permission: %v", err) + } + if got, want := service.addPermissionLabel, "label-a"; got != want { + t.Fatalf("unexpected add permission label: got %q want %q", got, want) + } + if got, want := service.addPermissionAWSAccountIDs[0], "123456789012"; got != want { + t.Fatalf("unexpected add permission account: got %q want %q", got, want) + } + + if _, err := client.RemovePermission(ctx, &sqs.RemovePermissionInput{ + QueueUrl: aws.String(queueURL), + Label: aws.String("label-a"), + }); err != nil { + t.Fatalf("remove permission: %v", err) + } + if got, want := service.removePermissionLabel, "label-a"; got != want { + t.Fatalf("unexpected remove permission label: got %q want %q", got, want) + } + + if _, err := client.UntagQueue(ctx, &sqs.UntagQueueInput{ + QueueUrl: aws.String(queueURL), + TagKeys: []string{"env", "team"}, + }); err != nil { + t.Fatalf("untag queue: %v", err) + } + if got, want := service.untagQueueQueueName, "orders"; got != want { + t.Fatalf("unexpected untag queue name: got %q want %q", got, want) + } + if got, want := len(service.untagQueueTagKeys), 2; got != want { + t.Fatalf("unexpected untag key count: got %d want %d", got, want) + } + + dlqURL := service.QueueURL("orders-dlq") + dlqOutput, err := client.ListDeadLetterSourceQueues(ctx, &sqs.ListDeadLetterSourceQueuesInput{ + QueueUrl: aws.String(dlqURL), + MaxResults: aws.Int32(2), + }) + if err != nil { + t.Fatalf("list dead letter source queues: %v", err) + } + if got, want := service.listDeadLetterSourceQueuesQueueName, "orders-dlq"; got != want { + t.Fatalf("unexpected dead letter source queue name: got %q want %q", got, want) + } + if got, want := len(dlqOutput.QueueUrls), 1; got != want { + t.Fatalf("unexpected dead letter source count: got %d want %d", got, want) + } + if got, want := dlqOutput.QueueUrls[0], service.QueueURL("orders-source"); got != want { + t.Fatalf("unexpected dead letter source URL: got %q want %q", got, want) + } + + startOutput, err := client.StartMessageMoveTask(ctx, &sqs.StartMessageMoveTaskInput{ + SourceArn: aws.String(service.QueueARN("orders-dlq")), + DestinationArn: aws.String(service.QueueARN("orders")), + MaxNumberOfMessagesPerSecond: aws.Int32(12), + }) + if err != nil { + t.Fatalf("start message move task: %v", err) + } + if got, want := service.startMessageMoveTaskSourceArn, service.QueueARN("orders-dlq"); got != want { + t.Fatalf("unexpected start source arn: got %q want %q", got, want) + } + if got, want := service.startMessageMoveTaskDestinationArn, service.QueueARN("orders"); got != want { + t.Fatalf("unexpected start destination arn: got %q want %q", got, want) + } + if got, want := service.startMessageMoveTaskMaxPerSecond, 12; got != want { + t.Fatalf("unexpected start rate: got %d want %d", got, want) + } + if got, want := aws.ToString(startOutput.TaskHandle), "arn:aws:sqs:us-east-1:123456789012:orders-dlq|task-1"; got != want { + t.Fatalf("unexpected task handle: got %q want %q", got, want) + } + + cancelOutput, err := client.CancelMessageMoveTask(ctx, &sqs.CancelMessageMoveTaskInput{ + TaskHandle: startOutput.TaskHandle, + }) + if err != nil { + t.Fatalf("cancel message move task: %v", err) + } + if got, want := service.cancelMessageMoveTaskTaskHandle, "arn:aws:sqs:us-east-1:123456789012:orders-dlq|task-1"; got != want { + t.Fatalf("unexpected cancel task handle: got %q want %q", got, want) + } + if got, want := cancelOutput.ApproximateNumberOfMessagesMoved, int64(7); got != want { + t.Fatalf("unexpected moved count: got %d want %d", got, want) + } + + tasksOutput, err := client.ListMessageMoveTasks(ctx, &sqs.ListMessageMoveTasksInput{ + SourceArn: aws.String(service.QueueARN("orders-dlq")), + MaxResults: aws.Int32(1), + }) + if err != nil { + t.Fatalf("list message move tasks: %v", err) + } + if got, want := service.listMessageMoveTasksQueueName, "orders-dlq"; got != want { + t.Fatalf("unexpected move tasks queue name: got %q want %q", got, want) + } + if got, want := len(tasksOutput.Results), 1; got != want { + t.Fatalf("unexpected task result count: got %d want %d", got, want) + } + if got, want := aws.ToString(tasksOutput.Results[0].TaskHandle), "arn:aws:sqs:us-east-1:123456789012:orders-dlq|task-1"; got != want { + t.Fatalf("unexpected task result handle: got %q want %q", got, want) + } + if got, want := aws.ToString(tasksOutput.Results[0].DestinationArn), service.QueueARN("orders"); got != want { + t.Fatalf("unexpected task destination arn: got %q want %q", got, want) + } +} + +func TestSQSSDKSmokePreservesBatchEntrySemantics(t *testing.T) { + t.Helper() + + service := sqsapplication.New() + + if _, err := service.CreateQueue("orders", nil); err != nil { + t.Fatalf("create queue: %v", err) + } + + result, err := service.SendMessageBatch("orders", contracts.SendMessageBatchRequest{ + QueueUrl: service.QueueURL("orders"), + Entries: []contracts.SendMessageBatchRequestEntry{ + {Id: "entry-1", MessageBody: "one"}, + {Id: "entry-2", MessageBody: "two"}, + {Id: "entry-3", MessageBody: ""}, + }, + }) + if err != nil { + t.Fatalf("send message batch: %v", err) + } + if got, want := len(result.Successful), 2; got != want { + t.Fatalf("unexpected batch send success count: got %d want %d", got, want) + } + if got, want := len(result.Failed), 1; got != want { + t.Fatalf("unexpected batch send failure count: got %d want %d", got, want) + } + if got, want := result.Successful[0].Id, "entry-1"; got != want { + t.Fatalf("unexpected first batch success id: got %q want %q", got, want) + } + if got, want := result.Failed[0].Id, "entry-3"; got != want { + t.Fatalf("unexpected batch failure id: got %q want %q", got, want) + } + + messages, err := service.ReceiveMessage("orders", 2, 0) + if err != nil { + t.Fatalf("receive queued batch messages: %v", err) + } + if got, want := len(messages), 2; got != want { + t.Fatalf("unexpected queued message count after batch send: got %d want %d", got, want) + } + if got, want := messages[0].Body, "one"; got != want { + t.Fatalf("unexpected first queued batch body: got %q want %q", got, want) + } + if got, want := messages[1].Body, "two"; got != want { + t.Fatalf("unexpected second queued batch body: got %q want %q", got, want) + } +} + +func TestSQSNativeMiddlewareRoutesLifecycleActionThroughService(t *testing.T) { + t.Helper() + + gin.SetMode(gin.TestMode) + + service := &stubSQSNativeService{} + router := gin.New() + RegisterSQSNativeRoutes(router, service) + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, "/?Action=ListQueues&Version=2012-11-05&QueueNamePrefix=ord&MaxResults=2&NextToken=token-1&QueueOwnerAWSAccountId=123456789012", nil) + router.ServeHTTP(recorder, request) + + if !service.listQueuesCalled { + t.Fatal("expected lifecycle request to reach the service") + } + if got, want := service.queueNamePrefix, "ord"; got != want { + t.Fatalf("unexpected queue name prefix: got %q want %q", got, want) + } + if got, want := service.maxResults, 2; got != want { + t.Fatalf("unexpected max results: got %d want %d", got, want) + } + if got, want := service.nextToken, "token-1"; got != want { + t.Fatalf("unexpected next token: got %q want %q", got, want) + } + if got, want := recorder.Code, http.StatusBadRequest; got != want { + t.Fatalf("unexpected lifecycle response status: got %d want %d", got, want) + } + if !strings.Contains(recorder.Body.String(), "UnsupportedOperation") { + t.Fatalf("expected deferred lifecycle response, got %q", recorder.Body.String()) + } +} + +type stubSQSNativeService struct { + listQueuesCalled bool + queueNamePrefix string + maxResults int + nextToken string + sendMessageQueueName string + sendMessageRequest contracts.SendMessageRequest + sendMessageResult contracts.SendMessageResult + + tagQueueQueueName string + tagQueueTags map[string]string + untagQueueQueueName string + untagQueueTagKeys []string + addPermissionQueueName string + addPermissionLabel string + addPermissionAWSAccountIDs []string + addPermissionActions []string + removePermissionQueueName string + removePermissionLabel string + listQueueTagsQueueName string + listQueueTagsResult map[string]string + listDeadLetterSourceQueuesQueueName string + listDeadLetterSourceQueuesResult []string + startMessageMoveTaskSourceArn string + startMessageMoveTaskDestinationArn string + startMessageMoveTaskMaxPerSecond int + startMessageMoveTaskResult string + cancelMessageMoveTaskTaskHandle string + cancelMessageMoveTaskResult int64 + listMessageMoveTasksQueueName string + listMessageMoveTasksResult []domain.MessageMoveTask +} + +func (s *stubSQSNativeService) Policy() orchestrator.EmulationPolicy { + return orchestrator.NewEmulationPolicy(orchestrator.FidelityExemplar, contracts.ActionNames(), nil, "sqs") +} + +func (s *stubSQSNativeService) Metadata() orchestrator.Metadata { + return orchestrator.Metadata{Name: "sqs"} +} + +func (s *stubSQSNativeService) QueueURL(queueName string) string { + return "https://sqs.us-east-1.amazonaws.com/123456789012/" + queueName +} + +func (s *stubSQSNativeService) QueueARN(queueName string) string { + return "arn:aws:sqs:us-east-1:123456789012:" + queueName +} + +func (s *stubSQSNativeService) CreateQueue(queueName string, attributes map[string]string) (domain.Queue, error) { + return domain.Queue{Name: queueName, URL: s.QueueURL(queueName)}, contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) DeleteQueue(queueName string) error { + return contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) GetQueueUrl(queueName, ownerAccountID string) (string, error) { + return s.QueueURL(queueName), contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) ListQueues(queueNamePrefix string, maxResults int, nextToken, ownerAccountID string) ([]domain.Queue, string, error) { + s.listQueuesCalled = true + s.queueNamePrefix = queueNamePrefix + s.maxResults = maxResults + s.nextToken = nextToken + return nil, "", contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) PurgeQueue(queueName string) error { + return contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) GetQueueAttributes(queueName string, attributeNames []string, ownerAccountID string) (contracts.QueueAttributesView, error) { + return contracts.QueueAttributesView{ + QueueName: queueName, + QueueURL: s.QueueURL(queueName), + QueueARN: s.QueueARN(queueName), + }, contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) SetQueueAttributes(queueName string, attributes map[string]string) (contracts.QueueAttributesView, error) { + return contracts.QueueAttributesView{ + QueueName: queueName, + QueueURL: s.QueueURL(queueName), + QueueARN: s.QueueARN(queueName), + Attributes: attributes, + }, contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) TagQueue(queueName string, tags map[string]string) error { + s.tagQueueQueueName = queueName + s.tagQueueTags = tags + return nil +} + +func (s *stubSQSNativeService) UntagQueue(queueName string, tagKeys []string) error { + s.untagQueueQueueName = queueName + s.untagQueueTagKeys = tagKeys + return nil +} + +func (s *stubSQSNativeService) AddPermission(queueName, label string, awsAccountIDs, actions []string) error { + s.addPermissionQueueName = queueName + s.addPermissionLabel = label + s.addPermissionAWSAccountIDs = awsAccountIDs + s.addPermissionActions = actions + return nil +} + +func (s *stubSQSNativeService) RemovePermission(queueName, label string) error { + s.removePermissionQueueName = queueName + s.removePermissionLabel = label + return nil +} + +func (s *stubSQSNativeService) ListQueueTags(queueName string) (map[string]string, error) { + s.listQueueTagsQueueName = queueName + if s.listQueueTagsResult == nil { + return map[string]string{}, nil + } + return s.listQueueTagsResult, nil +} + +func (s *stubSQSNativeService) ListDeadLetterSourceQueues(queueName string) ([]string, error) { + s.listDeadLetterSourceQueuesQueueName = queueName + if s.listDeadLetterSourceQueuesResult == nil { + return []string{}, nil + } + return s.listDeadLetterSourceQueuesResult, nil +} + +func (s *stubSQSNativeService) StartMessageMoveTask(sourceArn, destinationArn string, maxNumberOfMessagesPerSecond int) (string, error) { + s.startMessageMoveTaskSourceArn = sourceArn + s.startMessageMoveTaskDestinationArn = destinationArn + s.startMessageMoveTaskMaxPerSecond = maxNumberOfMessagesPerSecond + if s.startMessageMoveTaskResult == "" { + s.startMessageMoveTaskResult = "task-1" + } + return s.startMessageMoveTaskResult, nil +} + +func (s *stubSQSNativeService) CancelMessageMoveTask(taskHandle string) (int64, error) { + s.cancelMessageMoveTaskTaskHandle = taskHandle + return s.cancelMessageMoveTaskResult, nil +} + +func (s *stubSQSNativeService) ListMessageMoveTasks(queueName string) ([]domain.MessageMoveTask, error) { + s.listMessageMoveTasksQueueName = queueName + if s.listMessageMoveTasksResult == nil { + return []domain.MessageMoveTask{}, nil + } + return s.listMessageMoveTasksResult, nil +} + +func (s *stubSQSNativeService) ReceiveMessage(queueName string, maxMessages int, waitTime time.Duration) ([]domain.Message, error) { + return []domain.Message{ + { + Queue: queueName, + MessageID: "message-1", + Body: "hello", + SentAt: time.Now(), + ReceiptKeys: []string{"receipt-1"}, + }, + }, nil +} + +func (s *stubSQSNativeService) DeleteMessage(queueName string, receiptHandle string) error { + return contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) ChangeMessageVisibility(queueName string, receiptHandle string, visibility time.Duration) error { + return contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) SendMessage(queueName string, request contracts.SendMessageRequest) (contracts.SendMessageResult, error) { + s.sendMessageQueueName = queueName + s.sendMessageRequest = request + if s.sendMessageResult.MessageId == "" { + s.sendMessageResult = contracts.SendMessageResult{ + MessageId: "message-1", + MD5OfMessageBody: "md5-body", + } + } + return s.sendMessageResult, nil +} + +func (s *stubSQSNativeService) SendMessageBatch(queueName string, request contracts.SendMessageBatchRequest) (contracts.SendMessageBatchResult, error) { + return contracts.SendMessageBatchResult{}, contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) DeleteMessageBatch(queueName string, request contracts.DeleteMessageBatchRequest) (contracts.DeleteMessageBatchResult, error) { + return contracts.DeleteMessageBatchResult{}, contracts.ErrSQSOperationDeferred +} + +func (s *stubSQSNativeService) ChangeMessageVisibilityBatch(queueName string, request contracts.ChangeMessageVisibilityBatchRequest) (contracts.ChangeMessageVisibilityBatchResult, error) { + return contracts.ChangeMessageVisibilityBatchResult{}, contracts.ErrSQSOperationDeferred } type captureTransport struct { diff --git a/core/internal/resources/dynamodb/application/repository_sqlite.go b/core/internal/resources/dynamodb/application/repository_sqlite.go index 078667e..457bd23 100644 --- a/core/internal/resources/dynamodb/application/repository_sqlite.go +++ b/core/internal/resources/dynamodb/application/repository_sqlite.go @@ -22,7 +22,7 @@ import ( const ( sqliteFileName = "state.db" schemaVersionKey = "schema_version" - schemaVersion = "2" + schemaVersion = "3" ) type SQLiteRepository struct { @@ -138,6 +138,9 @@ func (r *SQLiteRepository) bootstrap() error { partition_key TEXT NOT NULL, sort_key TEXT NOT NULL, billing_mode TEXT NOT NULL, + attribute_definitions_json TEXT NOT NULL DEFAULT '[]', + global_secondary_indexes_json TEXT NOT NULL DEFAULT '[]', + local_secondary_indexes_json TEXT NOT NULL DEFAULT '[]', status TEXT NOT NULL DEFAULT 'ACTIVE', created_at_ns INTEGER NOT NULL DEFAULT 0, activation_at_ns INTEGER NOT NULL DEFAULT 0, @@ -163,6 +166,18 @@ func (r *SQLiteRepository) bootstrap() error { _ = tx.Rollback() return err } + if err := ensureTableColumn(ctx, tx, "dynamodb_tables", "attribute_definitions_json", "TEXT NOT NULL DEFAULT '[]'"); err != nil { + _ = tx.Rollback() + return err + } + if err := ensureTableColumn(ctx, tx, "dynamodb_tables", "global_secondary_indexes_json", "TEXT NOT NULL DEFAULT '[]'"); err != nil { + _ = tx.Rollback() + return err + } + if err := ensureTableColumn(ctx, tx, "dynamodb_tables", "local_secondary_indexes_json", "TEXT NOT NULL DEFAULT '[]'"); err != nil { + _ = tx.Rollback() + return err + } if err := ensureTableColumn(ctx, tx, "dynamodb_tables", "created_at_ns", "INTEGER NOT NULL DEFAULT 0"); err != nil { _ = tx.Rollback() return err @@ -197,7 +212,7 @@ func (r *SQLiteRepository) loadLocked() (domain.State, error) { state := domain.State{Service: "dynamodb"} tableRows, err := r.db.QueryContext(ctx, ` - SELECT name, partition_key, sort_key, billing_mode, status, created_at_ns, activation_at_ns, deleted_at_ns + SELECT name, partition_key, sort_key, billing_mode, attribute_definitions_json, global_secondary_indexes_json, local_secondary_indexes_json, status, created_at_ns, activation_at_ns, deleted_at_ns FROM dynamodb_tables ORDER BY name `) @@ -208,10 +223,26 @@ func (r *SQLiteRepository) loadLocked() (domain.State, error) { for tableRows.Next() { var table domain.Table - var createdAtNS, activationAtNS, deletedAtNS int64 - if err := tableRows.Scan(&table.Name, &table.PartitionKey, &table.SortKey, &table.BillingMode, &table.Status, &createdAtNS, &activationAtNS, &deletedAtNS); err != nil { + var ( + attributeDefinitionsJSON string + globalSecondaryIndexesJSON string + localSecondaryIndexesJSON string + createdAtNS int64 + activationAtNS int64 + deletedAtNS int64 + ) + if err := tableRows.Scan(&table.Name, &table.PartitionKey, &table.SortKey, &table.BillingMode, &attributeDefinitionsJSON, &globalSecondaryIndexesJSON, &localSecondaryIndexesJSON, &table.Status, &createdAtNS, &activationAtNS, &deletedAtNS); err != nil { return domain.State{}, fmt.Errorf("dynamodb: scan table: %w", err) } + if err := json.Unmarshal([]byte(attributeDefinitionsJSON), &table.AttributeDefinitions); err != nil { + return domain.State{}, fmt.Errorf("dynamodb: decode table attribute definitions %q: %w", table.Name, err) + } + if err := json.Unmarshal([]byte(globalSecondaryIndexesJSON), &table.GlobalSecondaryIndexes); err != nil { + return domain.State{}, fmt.Errorf("dynamodb: decode table global secondary indexes %q: %w", table.Name, err) + } + if err := json.Unmarshal([]byte(localSecondaryIndexesJSON), &table.LocalSecondaryIndexes); err != nil { + return domain.State{}, fmt.Errorf("dynamodb: decode table local secondary indexes %q: %w", table.Name, err) + } table.CreatedAt = unixNanoToTime(createdAtNS) table.ActivationAt = unixNanoToTime(activationAtNS) table.DeletedAt = unixNanoToTime(deletedAtNS) @@ -278,8 +309,8 @@ func (r *SQLiteRepository) saveLocked(state domain.State) error { } tableStmt, err := tx.PrepareContext(ctx, ` - INSERT INTO dynamodb_tables(name, partition_key, sort_key, billing_mode, status, created_at_ns, activation_at_ns, deleted_at_ns) - VALUES (?, ?, ?, ?, ?, ?, ?, ?) + INSERT INTO dynamodb_tables(name, partition_key, sort_key, billing_mode, attribute_definitions_json, global_secondary_indexes_json, local_secondary_indexes_json, status, created_at_ns, activation_at_ns, deleted_at_ns) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) `) if err != nil { _ = tx.Rollback() @@ -293,6 +324,9 @@ func (r *SQLiteRepository) saveLocked(state domain.State) error { table.PartitionKey, table.SortKey, table.BillingMode, + string(marshalJSONOrPanic(table.AttributeDefinitions)), + string(marshalJSONOrPanic(table.GlobalSecondaryIndexes)), + string(marshalJSONOrPanic(table.LocalSecondaryIndexes)), table.Status, timeToUnixNano(table.CreatedAt), timeToUnixNano(table.ActivationAt), @@ -360,6 +394,9 @@ func validatePersistedState(state domain.State) error { if table.BillingMode == "" { return fmt.Errorf("dynamodb: invalid table %q: empty billing mode", table.Name) } + if err := validatePersistedIndexDefinitions(table); err != nil { + return err + } if _, ok := tables[table.Name]; ok { return fmt.Errorf("dynamodb: duplicate table %q", table.Name) } @@ -388,6 +425,9 @@ func normalizePersistedTable(table domain.Table) domain.Table { table.PartitionKey = strings.TrimSpace(table.PartitionKey) table.SortKey = strings.TrimSpace(table.SortKey) table.BillingMode = strings.TrimSpace(table.BillingMode) + table.AttributeDefinitions = normalizePersistedAttributeDefinitions(table.AttributeDefinitions) + table.GlobalSecondaryIndexes = normalizePersistedSecondaryIndexes(table.GlobalSecondaryIndexes) + table.LocalSecondaryIndexes = normalizePersistedSecondaryIndexes(table.LocalSecondaryIndexes) table.Status = strings.ToUpper(strings.TrimSpace(table.Status)) switch table.Status { @@ -401,6 +441,173 @@ func normalizePersistedTable(table domain.Table) domain.Table { return table } +func validatePersistedIndexDefinitions(table domain.Table) error { + indexNames := map[string]struct{}{} + for _, index := range table.GlobalSecondaryIndexes { + if err := validatePersistedSecondaryIndex(table, index, false); err != nil { + return fmt.Errorf("dynamodb: invalid table %q global secondary index: %w", table.Name, err) + } + if _, ok := indexNames[index.Name]; ok { + return fmt.Errorf("dynamodb: invalid table %q: duplicate index %q", table.Name, index.Name) + } + indexNames[index.Name] = struct{}{} + } + for _, index := range table.LocalSecondaryIndexes { + if err := validatePersistedSecondaryIndex(table, index, true); err != nil { + return fmt.Errorf("dynamodb: invalid table %q local secondary index: %w", table.Name, err) + } + if _, ok := indexNames[index.Name]; ok { + return fmt.Errorf("dynamodb: invalid table %q: duplicate index %q", table.Name, index.Name) + } + indexNames[index.Name] = struct{}{} + } + return nil +} + +func validatePersistedSecondaryIndex(table domain.Table, index domain.SecondaryIndex, local bool) error { + if strings.TrimSpace(index.Name) == "" { + return fmt.Errorf("empty name") + } + if len(index.KeySchema) == 0 { + return fmt.Errorf("index %q has no key schema", index.Name) + } + var hashCount, rangeCount int + var partitionKey, sortKey string + for _, element := range index.KeySchema { + switch strings.ToUpper(strings.TrimSpace(element.KeyType)) { + case "HASH": + hashCount++ + partitionKey = strings.TrimSpace(element.AttributeName) + case "RANGE": + rangeCount++ + sortKey = strings.TrimSpace(element.AttributeName) + } + } + if hashCount != 1 { + return fmt.Errorf("index %q must have exactly one HASH key", index.Name) + } + if local { + if partitionKey != table.PartitionKey { + return fmt.Errorf("index %q must reuse table partition key %q", index.Name, table.PartitionKey) + } + } else if partitionKey == table.PartitionKey { + return fmt.Errorf("index %q must not reuse table partition key %q", index.Name, table.PartitionKey) + } + if rangeCount > 1 { + return fmt.Errorf("index %q has duplicate RANGE keys", index.Name) + } + if strings.EqualFold(index.Projection.Type, "INCLUDE") && len(index.Projection.NonKeyAttributes) == 0 { + return fmt.Errorf("index %q INCLUDE projection requires non-key attributes", index.Name) + } + _ = sortKey + return nil +} + +func normalizePersistedAttributeDefinitions(source []domain.AttributeDefinition) []domain.AttributeDefinition { + if len(source) == 0 { + return nil + } + normalized := make([]domain.AttributeDefinition, 0, len(source)) + seen := make(map[string]struct{}, len(source)) + for _, definition := range source { + definition.Name = strings.TrimSpace(definition.Name) + definition.Type = strings.ToUpper(strings.TrimSpace(definition.Type)) + if definition.Name == "" { + continue + } + if _, ok := seen[definition.Name]; ok { + continue + } + seen[definition.Name] = struct{}{} + normalized = append(normalized, definition) + } + return normalized +} + +func normalizePersistedSecondaryIndexes(source []domain.SecondaryIndex) []domain.SecondaryIndex { + if len(source) == 0 { + return nil + } + normalized := make([]domain.SecondaryIndex, 0, len(source)) + for _, index := range source { + index.Name = strings.TrimSpace(index.Name) + index.KeySchema = normalizePersistedKeySchema(index.KeySchema) + index.Projection = normalizePersistedProjection(index.Projection) + if index.Name == "" { + continue + } + normalized = append(normalized, index) + } + return normalized +} + +func normalizePersistedKeySchema(source []domain.KeySchemaElement) []domain.KeySchemaElement { + if len(source) == 0 { + return nil + } + normalized := make([]domain.KeySchemaElement, 0, len(source)) + seen := map[string]struct{}{} + for _, element := range source { + element.AttributeName = strings.TrimSpace(element.AttributeName) + element.KeyType = strings.ToUpper(strings.TrimSpace(element.KeyType)) + if element.AttributeName == "" || element.KeyType == "" { + continue + } + key := element.KeyType + ":" + element.AttributeName + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + normalized = append(normalized, element) + } + return normalized +} + +func normalizePersistedProjection(projection domain.Projection) domain.Projection { + projection.Type = strings.ToUpper(strings.TrimSpace(projection.Type)) + switch projection.Type { + case "", "ALL": + projection.Type = "ALL" + projection.NonKeyAttributes = nil + case "KEYS_ONLY": + projection.NonKeyAttributes = nil + case "INCLUDE": + projection.NonKeyAttributes = uniqueStringsLocal(projection.NonKeyAttributes) + default: + projection.Type = "ALL" + projection.NonKeyAttributes = nil + } + return projection +} + +func marshalJSONOrPanic(value any) []byte { + data, err := json.Marshal(value) + if err != nil { + panic(err) + } + return data +} + +func uniqueStringsLocal(values []string) []string { + if len(values) == 0 { + return nil + } + seen := make(map[string]struct{}, len(values)) + unique := make([]string, 0, len(values)) + for _, value := range values { + value = strings.TrimSpace(value) + if value == "" { + continue + } + if _, ok := seen[value]; ok { + continue + } + seen[value] = struct{}{} + unique = append(unique, value) + } + return unique +} + func ensureTableColumn(ctx context.Context, tx *sql.Tx, tableName, columnName, definition string) error { rows, err := tx.QueryContext(ctx, fmt.Sprintf(`PRAGMA table_info(%s)`, tableName)) if err != nil { @@ -619,12 +826,12 @@ func decodeAttributeValue(raw json.RawMessage) (domain.AttributeValue, error) { } type storedAttributeValue struct { - S *string `json:"S,omitempty"` - N *string `json:"N,omitempty"` - BOOL *bool `json:"BOOL,omitempty"` - NULL bool `json:"NULL,omitempty"` - M map[string]json.RawMessage `json:"M,omitempty"` - L []json.RawMessage `json:"L,omitempty"` + S *string `json:"S,omitempty"` + N *string `json:"N,omitempty"` + BOOL *bool `json:"BOOL,omitempty"` + NULL bool `json:"NULL,omitempty"` + M map[string]json.RawMessage `json:"M,omitempty"` + L []json.RawMessage `json:"L,omitempty"` } func decodeStoredAttributeValue(raw json.RawMessage) (domain.AttributeValue, bool, error) { diff --git a/core/internal/resources/dynamodb/application/repository_sqlite_test.go b/core/internal/resources/dynamodb/application/repository_sqlite_test.go index 390d955..eeac13f 100644 --- a/core/internal/resources/dynamodb/application/repository_sqlite_test.go +++ b/core/internal/resources/dynamodb/application/repository_sqlite_test.go @@ -70,8 +70,26 @@ func TestSQLiteRepositoryBootstrapAndPersistAcrossRestart(t *testing.T) { PartitionKey: "pk", SortKey: "sk", BillingMode: "PAY_PER_REQUEST", - Status: domain.TableStatusCreating, - CreatedAt: state.Tables[0].CreatedAt, + AttributeDefinitions: []domain.AttributeDefinition{ + {Name: "pk", Type: "S"}, + {Name: "sk", Type: "S"}, + {Name: "gsi_pk", Type: "S"}, + {Name: "gsi_sk", Type: "S"}, + }, + GlobalSecondaryIndexes: []domain.SecondaryIndex{ + { + Name: "gsi-archive", + KeySchema: []domain.KeySchemaElement{ + {AttributeName: "gsi_pk", KeyType: "HASH"}, + {AttributeName: "gsi_sk", KeyType: "RANGE"}, + }, + Projection: domain.Projection{ + Type: "KEYS_ONLY", + }, + }, + }, + Status: domain.TableStatusCreating, + CreatedAt: state.Tables[0].CreatedAt, }) state.UpsertItem(domain.Item{ Table: "mildstack-archive", @@ -116,6 +134,15 @@ func TestSQLiteRepositoryBootstrapAndPersistAcrossRestart(t *testing.T) { if got, want := fetched.Attributes["title"].Any(), "archive item"; got != want { t.Fatalf("unexpected item title after restart: got %q want %q", got, want) } + if got, want := len(archive.GlobalSecondaryIndexes), 1; got != want { + t.Fatalf("unexpected gsi count after restart: got %d want %d", got, want) + } + if got, want := archive.GlobalSecondaryIndexes[0].Name, "gsi-archive"; got != want { + t.Fatalf("unexpected gsi name after restart: got %q want %q", got, want) + } + if got, want := archive.GlobalSecondaryIndexes[0].Projection.Type, "KEYS_ONLY"; got != want { + t.Fatalf("unexpected gsi projection after restart: got %q want %q", got, want) + } statePath := filepath.Join(repo.storageDir, sqliteFileName) if _, err := os.Stat(statePath); err != nil { diff --git a/core/internal/resources/dynamodb/application/service.go b/core/internal/resources/dynamodb/application/service.go index 17a2a23..8e3b0ae 100644 --- a/core/internal/resources/dynamodb/application/service.go +++ b/core/internal/resources/dynamodb/application/service.go @@ -120,7 +120,7 @@ func (s *Service) ListTables() []domain.Table { return s.state.VisibleTables() } -func (s *Service) CreateTable(name, partitionKey, sortKey, billingMode string) (domain.Table, error) { +func (s *Service) CreateTable(name, partitionKey, sortKey, billingMode string, specs ...domain.CreateTableSpec) (domain.Table, error) { name = strings.TrimSpace(name) partitionKey = strings.TrimSpace(partitionKey) sortKey = strings.TrimSpace(sortKey) @@ -134,6 +134,14 @@ func (s *Service) CreateTable(name, partitionKey, sortKey, billingMode string) ( if billingMode == "" { billingMode = defaultBillingMode } + if len(specs) > 1 { + return domain.Table{}, fmt.Errorf("dynamodb: multiple create table specifications are not supported") + } + + spec := domain.CreateTableSpec{} + if len(specs) == 1 { + spec = normalizeCreateTableSpec(specs[0]) + } s.mu.Lock() defer s.mu.Unlock() @@ -144,15 +152,22 @@ func (s *Service) CreateTable(name, partitionKey, sortKey, billingMode string) ( } now := s.currentTime() - table := next.UpsertTable(domain.Table{ - Name: name, - PartitionKey: partitionKey, - SortKey: sortKey, - BillingMode: billingMode, - Status: domain.TableStatusCreating, - CreatedAt: now, - ActivationAt: now.Add(defaultActivationDelay), - }) + table := domain.Table{ + Name: name, + PartitionKey: partitionKey, + SortKey: sortKey, + BillingMode: billingMode, + AttributeDefinitions: cloneCreateTableAttributeDefinitions(spec.AttributeDefinitions), + GlobalSecondaryIndexes: cloneCreateTableSecondaryIndexes(spec.GlobalSecondaryIndexes), + LocalSecondaryIndexes: cloneCreateTableSecondaryIndexes(spec.LocalSecondaryIndexes), + Status: domain.TableStatusCreating, + CreatedAt: now, + ActivationAt: now.Add(defaultActivationDelay), + } + if err := validateCreateTableDefinition(table); err != nil { + return domain.Table{}, err + } + table = next.UpsertTable(table) if err := s.commitStateLocked(next); err != nil { return domain.Table{}, err } @@ -317,7 +332,7 @@ func (s *Service) DeleteItem(table, key string) error { return s.commitStateLocked(next) } -func (s *Service) Query(table, keyConditionExpression, filterExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]domain.AttributeValue, limit *int, exclusiveStartKey map[string]domain.AttributeValue, scanIndexForward *bool) (domain.ReadPage, error) { +func (s *Service) Query(table, keyConditionExpression, filterExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]domain.AttributeValue, limit *int, exclusiveStartKey map[string]domain.AttributeValue, scanIndexForward *bool, options ...domain.QueryOptions) (domain.ReadPage, error) { s.mu.Lock() defer s.mu.Unlock() @@ -331,7 +346,17 @@ func (s *Service) Query(table, keyConditionExpression, filterExpression string, return domain.ReadPage{}, fmt.Errorf("dynamodb: table %q not found", table) } - plan, err := buildQueryPlan(tableInfo, keyConditionExpression, expressionAttributeNames, expressionAttributeValues) + option := domain.QueryOptions{} + if len(options) > 0 { + option = options[0] + } + + target, err := resolveQueryTarget(tableInfo, option.IndexName) + if err != nil { + return domain.ReadPage{}, err + } + + plan, err := buildQueryPlan(target, keyConditionExpression, expressionAttributeNames, expressionAttributeValues) if err != nil { return domain.ReadPage{}, err } @@ -339,7 +364,7 @@ func (s *Service) Query(table, keyConditionExpression, filterExpression string, items := s.state.ListItems(table) candidates := make([]domain.Item, 0, len(items)) for _, item := range items { - matches, err := plan.matches(item, tableInfo) + matches, err := plan.matches(item, target) if err != nil { return domain.ReadPage{}, err } @@ -348,8 +373,8 @@ func (s *Service) Query(table, keyConditionExpression, filterExpression string, } } - ordered := orderQueryItems(candidates, tableInfo, scanIndexForward) - startIndex, err := locateExclusiveStartKey(ordered, tableInfo, exclusiveStartKey) + ordered := orderQueryItems(candidates, target, scanIndexForward) + startIndex, err := locateExclusiveStartKey(ordered, target, exclusiveStartKey) if err != nil { return domain.ReadPage{}, err } @@ -359,7 +384,12 @@ func (s *Service) Query(table, keyConditionExpression, filterExpression string, return domain.ReadPage{}, err } - return pageReadItems(ordered, tableInfo, startIndex, limit, filter) + projection, err := buildProjection(option.ProjectionExpression, expressionAttributeNames, target) + if err != nil { + return domain.ReadPage{}, err + } + + return pageReadItems(ordered, target, startIndex, limit, filter, projection) } func (s *Service) Scan(table, filterExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]domain.AttributeValue, limit *int, exclusiveStartKey map[string]domain.AttributeValue) (domain.ReadPage, error) { @@ -377,7 +407,7 @@ func (s *Service) Scan(table, filterExpression string, expressionAttributeNames } items := s.state.ListItems(table) - startIndex, err := locateExclusiveStartKey(items, tableInfo, exclusiveStartKey) + startIndex, err := locateExclusiveStartKey(items, queryTarget{Table: tableInfo, PartitionKey: tableInfo.PartitionKey, SortKey: tableInfo.SortKey}, exclusiveStartKey) if err != nil { return domain.ReadPage{}, err } @@ -387,7 +417,7 @@ func (s *Service) Scan(table, filterExpression string, expressionAttributeNames return domain.ReadPage{}, err } - return pageReadItems(items, tableInfo, startIndex, limit, filter) + return pageReadItems(items, queryTarget{Table: tableInfo, PartitionKey: tableInfo.PartitionKey, SortKey: tableInfo.SortKey}, startIndex, limit, filter, nil) } func (s *Service) commitStateLocked(next domain.State) error { @@ -444,3 +474,272 @@ func (s *Service) materializeTableLocked(state *domain.State, name string) bool } return changed } + +func normalizeCreateTableSpec(spec domain.CreateTableSpec) domain.CreateTableSpec { + spec.AttributeDefinitions = normalizeCreateTableAttributeDefinitions(spec.AttributeDefinitions) + spec.GlobalSecondaryIndexes = normalizeCreateTableSecondaryIndexes(spec.GlobalSecondaryIndexes) + spec.LocalSecondaryIndexes = normalizeCreateTableSecondaryIndexes(spec.LocalSecondaryIndexes) + return spec +} + +func validateCreateTableDefinition(table domain.Table) error { + if table.Name == "" { + return fmt.Errorf("dynamodb: table name is required") + } + if table.PartitionKey == "" { + return fmt.Errorf("dynamodb: table %q partition key is required", table.Name) + } + if table.BillingMode == "" { + return fmt.Errorf("dynamodb: table %q billing mode is required", table.Name) + } + if err := validateAttributeDefinitions(table); err != nil { + return err + } + if err := validateCreateTableIndexes(table); err != nil { + return err + } + return nil +} + +func validateAttributeDefinitions(table domain.Table) error { + if len(table.AttributeDefinitions) == 0 { + return nil + } + + definitions := make(map[string]string, len(table.AttributeDefinitions)) + for _, definition := range table.AttributeDefinitions { + if definition.Name == "" { + return fmt.Errorf("dynamodb: table %q has an empty attribute definition name", table.Name) + } + if definition.Type == "" { + return fmt.Errorf("dynamodb: table %q attribute %q is missing a type", table.Name, definition.Name) + } + if _, ok := definitions[definition.Name]; ok { + return fmt.Errorf("dynamodb: table %q has duplicate attribute definition %q", table.Name, definition.Name) + } + definitions[definition.Name] = definition.Type + } + + needed := []string{table.PartitionKey} + if table.SortKey != "" { + needed = append(needed, table.SortKey) + } + for _, index := range append(table.GlobalSecondaryIndexes, table.LocalSecondaryIndexes...) { + for _, element := range index.KeySchema { + if name := strings.TrimSpace(element.AttributeName); name != "" { + needed = append(needed, name) + } + } + } + + for _, name := range uniqueStrings(needed) { + if _, ok := definitions[name]; !ok { + return fmt.Errorf("dynamodb: table %q is missing attribute definition for %q", table.Name, name) + } + } + + return nil +} + +func validateCreateTableIndexes(table domain.Table) error { + indexNames := make(map[string]struct{}) + for _, index := range table.GlobalSecondaryIndexes { + if err := validateCreateTableIndex(table, index, false); err != nil { + return err + } + if _, ok := indexNames[strings.ToLower(index.Name)]; ok { + return fmt.Errorf("dynamodb: duplicate index %q", index.Name) + } + indexNames[strings.ToLower(index.Name)] = struct{}{} + } + for _, index := range table.LocalSecondaryIndexes { + if err := validateCreateTableIndex(table, index, true); err != nil { + return err + } + if _, ok := indexNames[strings.ToLower(index.Name)]; ok { + return fmt.Errorf("dynamodb: duplicate index %q", index.Name) + } + indexNames[strings.ToLower(index.Name)] = struct{}{} + } + return nil +} + +func validateCreateTableIndex(table domain.Table, index domain.SecondaryIndex, local bool) error { + if strings.TrimSpace(index.Name) == "" { + return fmt.Errorf("dynamodb: index name is required") + } + partitionKey, sortKey, err := validateSecondaryIndexKeySchema(index.KeySchema) + if err != nil { + return fmt.Errorf("dynamodb: index %q: %w", index.Name, err) + } + if local { + if partitionKey != table.PartitionKey { + return fmt.Errorf("dynamodb: index %q must reuse table partition key %q", index.Name, table.PartitionKey) + } + } else if partitionKey == table.PartitionKey { + return fmt.Errorf("dynamodb: index %q must not reuse table partition key %q", index.Name, table.PartitionKey) + } + if sortKey == "" && local { + return fmt.Errorf("dynamodb: index %q must define a RANGE key", index.Name) + } + if err := validateProjection(index.Projection); err != nil { + return fmt.Errorf("dynamodb: index %q: %w", index.Name, err) + } + return nil +} + +func validateSecondaryIndexKeySchema(keySchema []domain.KeySchemaElement) (string, string, error) { + var ( + hashCount int + rangeCount int + hashKey string + rangeKey string + ) + for _, element := range keySchema { + switch strings.ToUpper(strings.TrimSpace(element.KeyType)) { + case "HASH": + hashCount++ + hashKey = strings.TrimSpace(element.AttributeName) + case "RANGE": + rangeCount++ + rangeKey = strings.TrimSpace(element.AttributeName) + } + } + if hashCount != 1 { + return "", "", fmt.Errorf("must define exactly one HASH key") + } + if rangeCount > 1 { + return "", "", fmt.Errorf("must define at most one RANGE key") + } + if hashKey == "" { + return "", "", fmt.Errorf("HASH key attribute name is required") + } + return hashKey, rangeKey, nil +} + +func validateProjection(projection domain.Projection) error { + switch strings.ToUpper(strings.TrimSpace(projection.Type)) { + case "", "ALL", "KEYS_ONLY": + return nil + case "INCLUDE": + if len(projection.NonKeyAttributes) == 0 { + return fmt.Errorf("INCLUDE projection requires non-key attributes") + } + return nil + default: + return fmt.Errorf("unsupported projection type %q", projection.Type) + } +} + +func normalizeCreateTableAttributeDefinitions(source []domain.AttributeDefinition) []domain.AttributeDefinition { + if len(source) == 0 { + return nil + } + seen := map[string]struct{}{} + normalized := make([]domain.AttributeDefinition, 0, len(source)) + for _, definition := range source { + definition.Name = strings.TrimSpace(definition.Name) + definition.Type = strings.ToUpper(strings.TrimSpace(definition.Type)) + if definition.Name == "" { + continue + } + if _, ok := seen[definition.Name]; ok { + continue + } + seen[definition.Name] = struct{}{} + normalized = append(normalized, definition) + } + return normalized +} + +func normalizeCreateTableSecondaryIndexes(source []domain.SecondaryIndex) []domain.SecondaryIndex { + if len(source) == 0 { + return nil + } + normalized := make([]domain.SecondaryIndex, 0, len(source)) + for _, index := range source { + index.Name = strings.TrimSpace(index.Name) + index.KeySchema = normalizeCreateTableKeySchema(index.KeySchema) + index.Projection = normalizeCreateTableProjection(index.Projection) + if index.Name == "" { + continue + } + normalized = append(normalized, index) + } + return normalized +} + +func normalizeCreateTableKeySchema(source []domain.KeySchemaElement) []domain.KeySchemaElement { + if len(source) == 0 { + return nil + } + normalized := make([]domain.KeySchemaElement, 0, len(source)) + seen := map[string]struct{}{} + for _, element := range source { + element.AttributeName = strings.TrimSpace(element.AttributeName) + element.KeyType = strings.ToUpper(strings.TrimSpace(element.KeyType)) + if element.AttributeName == "" || element.KeyType == "" { + continue + } + key := element.KeyType + ":" + element.AttributeName + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + normalized = append(normalized, element) + } + return normalized +} + +func normalizeCreateTableProjection(projection domain.Projection) domain.Projection { + projection.Type = strings.ToUpper(strings.TrimSpace(projection.Type)) + switch projection.Type { + case "", "ALL": + projection.Type = "ALL" + projection.NonKeyAttributes = nil + case "KEYS_ONLY": + projection.NonKeyAttributes = nil + case "INCLUDE": + projection.NonKeyAttributes = uniqueStringsLocal(projection.NonKeyAttributes) + default: + projection.Type = "ALL" + projection.NonKeyAttributes = nil + } + return projection +} + +func cloneCreateTableAttributeDefinitions(source []domain.AttributeDefinition) []domain.AttributeDefinition { + if len(source) == 0 { + return nil + } + cloned := make([]domain.AttributeDefinition, len(source)) + copy(cloned, source) + return cloned +} + +func cloneCreateTableSecondaryIndexes(source []domain.SecondaryIndex) []domain.SecondaryIndex { + if len(source) == 0 { + return nil + } + cloned := make([]domain.SecondaryIndex, len(source)) + for i, index := range source { + cloned[i] = domain.SecondaryIndex{ + Name: index.Name, + KeySchema: cloneCreateTableKeySchema(index.KeySchema), + Projection: domain.Projection{ + Type: index.Projection.Type, + NonKeyAttributes: append([]string(nil), index.Projection.NonKeyAttributes...), + }, + } + } + return cloned +} + +func cloneCreateTableKeySchema(source []domain.KeySchemaElement) []domain.KeySchemaElement { + if len(source) == 0 { + return nil + } + cloned := make([]domain.KeySchemaElement, len(source)) + copy(cloned, source) + return cloned +} diff --git a/core/internal/resources/dynamodb/application/service_batch.go b/core/internal/resources/dynamodb/application/service_batch.go index 40cf86d..0319412 100644 --- a/core/internal/resources/dynamodb/application/service_batch.go +++ b/core/internal/resources/dynamodb/application/service_batch.go @@ -2,6 +2,7 @@ package application import ( "fmt" + "sort" "strings" ddbcontracts "github.com/michasdev/mildstack/core/internal/resources/dynamodb/contracts" @@ -50,7 +51,8 @@ func (s *Service) BatchWriteItem(request BatchWriteItemRequest) (BatchWriteItemR if tableName == "" { return BatchWriteItemResult{}, fmt.Errorf("dynamodb: table name is required") } - if !next.HasTable(tableName) { + tableInfo, ok := next.Table(tableName) + if !ok { return BatchWriteItemResult{}, fmt.Errorf("dynamodb: table %q not found", tableName) } @@ -70,7 +72,7 @@ func (s *Service) BatchWriteItem(request BatchWriteItemRequest) (BatchWriteItemR break } - key, err := batchDocumentKey(itemRequest.PutItem, itemRequest.DeleteKey) + key, err := batchDocumentKey(tableInfo, itemRequest.PutItem, itemRequest.DeleteKey) if err != nil { return BatchWriteItemResult{}, err } @@ -125,7 +127,8 @@ func (s *Service) BatchGetItem(request BatchGetItemRequest) (BatchGetItemResult, if tableName == "" { return BatchGetItemResult{}, fmt.Errorf("dynamodb: table name is required") } - if !s.state.HasTable(tableName) { + tableInfo, ok := s.state.Table(tableName) + if !ok { return BatchGetItemResult{}, fmt.Errorf("dynamodb: table %q not found", tableName) } @@ -142,7 +145,7 @@ func (s *Service) BatchGetItem(request BatchGetItemRequest) (BatchGetItemResult, break } - key, err := itemDocumentKey(keyDocument) + key, err := itemDocumentKey(tableInfo, keyDocument) if err != nil { return BatchGetItemResult{}, err } @@ -194,16 +197,17 @@ func (s *Service) TransactWriteItems(request TransactWriteItemsRequest) error { if tableName == "" { return fmt.Errorf("dynamodb: table name is required") } - if !next.HasTable(tableName) { + tableInfo, ok := next.Table(tableName) + if !ok { return fmt.Errorf("dynamodb: table %q not found", tableName) } - key, err := transactDocumentKey(item.PutItem, item.DeleteKey) + key, err := transactDocumentKey(tableInfo, item) if err != nil { return err } - if previous, ok := seen[tableName+"|"+key]; ok { + if previous, ok := seen[tableName+"|"+key]; ok && len(item.ConditionCheckKey) == 0 && len(request.Items[previous].ConditionCheckKey) == 0 { reason := TransactionCanceledReason{ Code: transactionConflictCode, Message: "same item targeted more than once", @@ -223,6 +227,59 @@ func (s *Service) TransactWriteItems(request TransactWriteItemsRequest) error { continue } + if len(item.UpdateKey) > 0 { + current, _ := next.Item(tableName, key) + attrs := cloneDocument(current.Attributes) + if attrs == nil { + attrs = make(map[string]domain.AttributeValue) + } + + if err := evaluateUpdateCondition(attrs, item.ConditionExpression, item.ExpressionAttributeNames, item.ExpressionAttributeValues); err != nil { + return &TransactionCanceledError{Reasons: cancellationReasonsForTransaction(index, len(request.Items), "ConditionalCheckFailed", "The conditional request failed")} + } + + operations, err := parseUpdateExpression(item.UpdateExpression, item.ExpressionAttributeNames, item.ExpressionAttributeValues) + if err != nil { + return err + } + for _, operation := range operations { + path := operation.path + if path == tableInfo.PartitionKey || (tableInfo.SortKey != "" && path == tableInfo.SortKey) { + return fmt.Errorf("dynamodb: unsupported update to key attribute %q", path) + } + + switch operation.kind { + case "SET": + attrs[path] = operation.value.Clone() + case "REMOVE": + delete(attrs, path) + case "ADD": + updated, err := addAttribute(attrs[path], operation.value) + if err != nil { + return err + } + attrs[path] = updated + default: + return fmt.Errorf("dynamodb: unsupported update operation %q", operation.kind) + } + } + + next.UpsertItem(domain.Item{ + Table: tableName, + Key: key, + Attributes: attrs, + }) + continue + } + + if len(item.ConditionCheckKey) > 0 { + current, _ := next.Item(tableName, key) + if err := evaluateUpdateCondition(current.Attributes, item.ConditionExpression, item.ExpressionAttributeNames, item.ExpressionAttributeValues); err != nil { + return &TransactionCanceledError{Reasons: cancellationReasonsForTransaction(index, len(request.Items), "ConditionalCheckFailed", "The conditional request failed")} + } + continue + } + next.DeleteItem(tableName, key) } @@ -246,11 +303,12 @@ func (s *Service) TransactGetItems(request TransactGetItemsRequest) (TransactGet if tableName == "" { return TransactGetItemsResult{}, fmt.Errorf("dynamodb: table name is required") } - if !s.state.HasTable(tableName) { + tableInfo, ok := s.state.Table(tableName) + if !ok { return TransactGetItemsResult{}, fmt.Errorf("dynamodb: table %q not found", tableName) } - key, err := itemDocumentKey(item.Key) + key, err := itemDocumentKey(tableInfo, item.Key) if err != nil { return TransactGetItemsResult{}, err } @@ -317,59 +375,127 @@ func cloneAttributeDocument(values map[string]domain.AttributeValue) map[string] return copied } -func batchDocumentKey(putItem, deleteKey map[string]domain.AttributeValue) (string, error) { +func batchDocumentKey(table domain.Table, putItem, deleteKey map[string]domain.AttributeValue) (string, error) { if len(putItem) > 0 { - return itemDocumentKey(putItem) + return itemRecordKey(table, putItem) } if len(deleteKey) > 0 { - return itemDocumentKey(deleteKey) + return itemDocumentKey(table, deleteKey) } return "", fmt.Errorf("dynamodb: batch request item is required") } -func transactDocumentKey(putItem, deleteKey map[string]domain.AttributeValue) (string, error) { - if len(putItem) > 0 { - return itemDocumentKey(putItem) +func transactDocumentKey(table domain.Table, item TransactWriteItem) (string, error) { + if len(item.PutItem) > 0 { + return itemRecordKey(table, item.PutItem) } - if len(deleteKey) > 0 { - return itemDocumentKey(deleteKey) + if len(item.DeleteKey) > 0 { + return itemDocumentKey(table, item.DeleteKey) + } + if len(item.UpdateKey) > 0 { + return itemDocumentKey(table, item.UpdateKey) + } + if len(item.ConditionCheckKey) > 0 { + return itemDocumentKey(table, item.ConditionCheckKey) } return "", fmt.Errorf("dynamodb: transaction item is required") } -func itemDocumentKey(values map[string]domain.AttributeValue) (string, error) { +func itemDocumentKey(table domain.Table, values map[string]domain.AttributeValue) (string, error) { if len(values) == 0 { return "", fmt.Errorf("dynamodb: item is required") } - if idValue, ok := values["id"]; ok { - id, err := attributeValueToKeyComponent(idValue) - if err != nil { - return "", err - } - if skValue, ok := values["sk"]; ok { - sk, err := attributeValueToKeyComponent(skValue) - if err != nil { - return "", err - } - if strings.TrimSpace(sk) != "" { - return id + "|" + sk, nil - } - } - return id, nil + if strings.TrimSpace(table.PartitionKey) == "" { + return "", fmt.Errorf("dynamodb: table %q has no partition key", table.Name) } - if len(values) == 1 { - for _, value := range values { - return attributeValueToKeyComponent(value) - } + expectedCount := 1 + if strings.TrimSpace(table.SortKey) != "" { + expectedCount++ + } + if len(values) != expectedCount { + return "", fmt.Errorf("dynamodb: unsupported key attributes %q", strings.Join(sortedAttributeKeys(values), ", ")) } + partitionValue, ok := values[table.PartitionKey] + if !ok { + return "", fmt.Errorf("dynamodb: missing key attribute %q", table.PartitionKey) + } + + partitionKey, err := attributeValueToKeyComponent(partitionValue) + if err != nil { + return "", err + } + + if strings.TrimSpace(table.SortKey) == "" { + return partitionKey, nil + } + + sortValue, ok := values[table.SortKey] + if !ok { + return "", fmt.Errorf("dynamodb: missing key attribute %q", table.SortKey) + } + + sortKey, err := attributeValueToKeyComponent(sortValue) + if err != nil { + return "", err + } + + return partitionKey + "|" + sortKey, nil +} + +func itemRecordKey(table domain.Table, values map[string]domain.AttributeValue) (string, error) { + if len(values) == 0 { + return "", fmt.Errorf("dynamodb: item is required") + } + if strings.TrimSpace(table.PartitionKey) == "" { + return "", fmt.Errorf("dynamodb: table %q has no partition key", table.Name) + } + + partitionValue, ok := values[table.PartitionKey] + if !ok { + return "", fmt.Errorf("dynamodb: missing key attribute %q", table.PartitionKey) + } + + partitionKey, err := attributeValueToKeyComponent(partitionValue) + if err != nil { + return "", err + } + + if strings.TrimSpace(table.SortKey) == "" { + return partitionKey, nil + } + + sortValue, ok := values[table.SortKey] + if !ok { + return "", fmt.Errorf("dynamodb: missing key attribute %q", table.SortKey) + } + + sortKey, err := attributeValueToKeyComponent(sortValue) + if err != nil { + return "", err + } + + return partitionKey + "|" + sortKey, nil +} + +func sortedAttributeKeys(values map[string]domain.AttributeValue) []string { keys := make([]string, 0, len(values)) for key := range values { keys = append(keys, key) } - return "", fmt.Errorf("dynamodb: unsupported key attributes %q", strings.Join(keys, ", ")) + sort.Strings(keys) + return keys +} + +func cancellationReasonsForTransaction(index, size int, code, message string) []TransactionCanceledReason { + reasons := make([]TransactionCanceledReason, size) + reasons[index] = TransactionCanceledReason{ + Code: code, + Message: message, + } + return reasons } func attributeValueToKeyComponent(value domain.AttributeValue) (string, error) { diff --git a/core/internal/resources/dynamodb/application/service_query.go b/core/internal/resources/dynamodb/application/service_query.go index 92d2a52..7b21e34 100644 --- a/core/internal/resources/dynamodb/application/service_query.go +++ b/core/internal/resources/dynamodb/application/service_query.go @@ -28,6 +28,13 @@ type queryPlan struct { sortPredicate sortPredicate } +type queryTarget struct { + Table domain.Table + Index *domain.SecondaryIndex + PartitionKey string + SortKey string +} + type sortPredicate struct { kind string values []domain.AttributeValue @@ -40,7 +47,57 @@ type filterClause struct { values []domain.AttributeValue } -func buildQueryPlan(table domain.Table, keyConditionExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]domain.AttributeValue) (queryPlan, error) { +func resolveQueryTarget(table domain.Table, indexName string) (queryTarget, error) { + target := queryTarget{ + Table: table, + PartitionKey: table.PartitionKey, + SortKey: table.SortKey, + } + + indexName = strings.TrimSpace(indexName) + if indexName == "" { + return target, nil + } + + if index, ok := findSecondaryIndex(table, indexName); ok { + target.Index = &index + hash, rangeKey := indexKeyNames(index) + target.PartitionKey = hash + target.SortKey = rangeKey + return target, nil + } + + return queryTarget{}, fmt.Errorf("dynamodb: index %q not found on table %q", indexName, table.Name) +} + +func findSecondaryIndex(table domain.Table, name string) (domain.SecondaryIndex, bool) { + for _, index := range table.GlobalSecondaryIndexes { + if strings.EqualFold(strings.TrimSpace(index.Name), name) { + return index, true + } + } + for _, index := range table.LocalSecondaryIndexes { + if strings.EqualFold(strings.TrimSpace(index.Name), name) { + return index, true + } + } + return domain.SecondaryIndex{}, false +} + +func indexKeyNames(index domain.SecondaryIndex) (string, string) { + var partitionKey, sortKey string + for _, element := range index.KeySchema { + switch strings.ToUpper(strings.TrimSpace(element.KeyType)) { + case "HASH": + partitionKey = strings.TrimSpace(element.AttributeName) + case "RANGE": + sortKey = strings.TrimSpace(element.AttributeName) + } + } + return partitionKey, sortKey +} + +func buildQueryPlan(target queryTarget, keyConditionExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]domain.AttributeValue) (queryPlan, error) { expression := strings.TrimSpace(keyConditionExpression) if expression == "" { return queryPlan{}, fmt.Errorf("dynamodb: key condition expression is required") @@ -58,7 +115,7 @@ func buildQueryPlan(table domain.Table, keyConditionExpression string, expressio if err != nil { return queryPlan{}, err } - if partitionKeyName != table.PartitionKey { + if partitionKeyName != target.PartitionKey { return queryPlan{}, fmt.Errorf("dynamodb: unsupported key condition partition key %q", partitionKeyName) } @@ -76,15 +133,15 @@ func buildQueryPlan(table domain.Table, keyConditionExpression string, expressio if sortExpression == "" { return plan, nil } - if table.SortKey == "" { - return queryPlan{}, fmt.Errorf("dynamodb: sort key conditions are not supported for table %q", table.Name) + if target.SortKey == "" { + return queryPlan{}, fmt.Errorf("dynamodb: sort key conditions are not supported for table %q", target.Table.Name) } sortPath, predicate, err := parseSortPredicate(sortExpression, expressionAttributeNames, expressionAttributeValues) if err != nil { return queryPlan{}, err } - if sortPath != table.SortKey { + if sortPath != target.SortKey { return queryPlan{}, fmt.Errorf("dynamodb: unsupported key condition sort key %q", sortPath) } @@ -140,7 +197,7 @@ func parseSortPredicate(expression string, expressionAttributeNames map[string]s return "", sortPredicate{}, fmt.Errorf("dynamodb: unsupported sort key condition %q", expression) } -func (p queryPlan) matches(item domain.Item, table domain.Table) (bool, error) { +func (p queryPlan) matches(item domain.Item, target queryTarget) (bool, error) { partitionValue, ok := item.Attributes[p.partitionKeyName] if !ok || !attributeValueEquals(partitionValue, p.partitionValue) { return false, nil @@ -204,6 +261,102 @@ func buildExpressionFilter(filterExpression string, expressionAttributeNames map }, nil } +func buildProjection(projectionExpression string, expressionAttributeNames map[string]string, target queryTarget) (func(domain.Item) (domain.Item, error), error) { + requested := strings.TrimSpace(projectionExpression) + allowed := projectionAllowedAttributes(target) + if requested == "" { + if len(allowed) == 0 { + return func(item domain.Item) (domain.Item, error) { + return cloneProjectedItem(item, nil), nil + }, nil + } + keys := sortedSetKeys(allowed) + return func(item domain.Item) (domain.Item, error) { + return cloneProjectedItem(item, keys), nil + }, nil + } + + paths := strings.Split(requested, ",") + keys := make([]string, 0, len(paths)) + seen := make(map[string]struct{}, len(paths)) + for _, raw := range paths { + path, err := resolveExpressionPath(raw, expressionAttributeNames) + if err != nil { + return nil, err + } + if _, ok := seen[path]; ok { + continue + } + if len(allowed) > 0 { + if _, ok := allowed[path]; !ok { + return nil, fmt.Errorf("dynamodb: projection path %q is not available for this index", path) + } + } + seen[path] = struct{}{} + keys = append(keys, path) + } + + return func(item domain.Item) (domain.Item, error) { + return cloneProjectedItem(item, keys), nil + }, nil +} + +func projectionAllowedAttributes(target queryTarget) map[string]struct{} { + if target.Index == nil { + return nil + } + + projection := strings.ToUpper(strings.TrimSpace(target.Index.Projection.Type)) + if projection == "" { + projection = "ALL" + } + if projection == "ALL" { + return nil + } + + allowed := make(map[string]struct{}) + for _, name := range targetKeyNames(target) { + allowed[name] = struct{}{} + } + switch projection { + case "KEYS_ONLY": + return allowed + case "INCLUDE": + for _, name := range target.Index.Projection.NonKeyAttributes { + name = strings.TrimSpace(name) + if name == "" { + continue + } + allowed[name] = struct{}{} + } + return allowed + default: + return nil + } +} + +func cloneProjectedItem(item domain.Item, names []string) domain.Item { + if len(names) == 0 { + return domain.Item{ + Table: item.Table, + Key: item.Key, + Attributes: cloneAttributeDocument(item.Attributes), + } + } + + attributes := make(map[string]domain.AttributeValue, len(names)) + for _, name := range names { + if value, ok := item.Attributes[name]; ok { + attributes[name] = value.Clone() + } + } + return domain.Item{ + Table: item.Table, + Key: item.Key, + Attributes: attributes, + } +} + func parseFilterClauses(expression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]domain.AttributeValue) ([]filterClause, error) { parts, err := splitFilterExpression(expression) if err != nil { @@ -366,7 +519,7 @@ func (c filterClause) matches(item domain.Item) (bool, error) { } } -func pageReadItems(items []domain.Item, table domain.Table, startIndex int, limit *int, filter func(domain.Item) (bool, error)) (domain.ReadPage, error) { +func pageReadItems(items []domain.Item, target queryTarget, startIndex int, limit *int, filter func(domain.Item) (bool, error), project func(domain.Item) (domain.Item, error)) (domain.ReadPage, error) { if limit != nil && *limit <= 0 { return domain.ReadPage{}, fmt.Errorf("dynamodb: limit must be greater than zero") } @@ -396,13 +549,20 @@ func pageReadItems(items []domain.Item, table domain.Table, startIndex int, limi } } if matches { - page.Items = append(page.Items, items[i]) + projected := items[i] + if project != nil { + projected, err = project(items[i]) + if err != nil { + return domain.ReadPage{}, err + } + } + page.Items = append(page.Items, projected) } } page.Count = len(page.Items) if end < len(items) { - cursor, err := keyAttributesForItem(table, items[end-1]) + cursor, err := keyAttributesForItem(target, items[end-1]) if err != nil { return domain.ReadPage{}, err } @@ -412,8 +572,8 @@ func pageReadItems(items []domain.Item, table domain.Table, startIndex int, limi return page, nil } -func locateExclusiveStartKey(items []domain.Item, table domain.Table, exclusiveStartKey map[string]domain.AttributeValue) (int, error) { - key, err := normalizeKeyAttributes(table, exclusiveStartKey) +func locateExclusiveStartKey(items []domain.Item, target queryTarget, exclusiveStartKey map[string]domain.AttributeValue) (int, error) { + key, err := normalizeKeyAttributes(target, exclusiveStartKey) if err != nil { return 0, err } @@ -422,7 +582,7 @@ func locateExclusiveStartKey(items []domain.Item, table domain.Table, exclusiveS } for i, item := range items { - itemKey, err := keyAttributesForItem(table, item) + itemKey, err := keyAttributesForItem(target, item) if err != nil { return 0, err } @@ -434,7 +594,7 @@ func locateExclusiveStartKey(items []domain.Item, table domain.Table, exclusiveS return 0, fmt.Errorf("dynamodb: exclusive start key not found") } -func orderQueryItems(items []domain.Item, table domain.Table, scanIndexForward *bool) []domain.Item { +func orderQueryItems(items []domain.Item, target queryTarget, scanIndexForward *bool) []domain.Item { ordered := make([]domain.Item, len(items)) copy(ordered, items) @@ -444,7 +604,7 @@ func orderQueryItems(items []domain.Item, table domain.Table, scanIndexForward * } sort.SliceStable(ordered, func(i, j int) bool { - cmp := compareQueryItems(ordered[i], ordered[j], table) + cmp := compareQueryItems(ordered[i], ordered[j], target) if forward { return cmp < 0 } @@ -454,12 +614,12 @@ func orderQueryItems(items []domain.Item, table domain.Table, scanIndexForward * return ordered } -func compareQueryItems(left, right domain.Item, table domain.Table) int { - if table.SortKey != "" { - leftSort, leftOK := left.Attributes[table.SortKey] - rightSort, rightOK := right.Attributes[table.SortKey] +func compareQueryItems(left, right domain.Item, target queryTarget) int { + for _, name := range orderingKeyNames(target) { + leftValue, leftOK := left.Attributes[name] + rightValue, rightOK := right.Attributes[name] if leftOK && rightOK { - if cmp := compareAttributeValues(leftSort, rightSort); cmp != 0 { + if cmp := compareAttributeValues(leftValue, rightValue); cmp != 0 { return cmp } } @@ -480,39 +640,29 @@ func compareQueryItems(left, right domain.Item, table domain.Table) int { return 0 } -func keyAttributesForItem(table domain.Table, item domain.Item) (map[string]domain.AttributeValue, error) { - if table.PartitionKey == "" { - return nil, fmt.Errorf("dynamodb: table %q has no partition key", table.Name) +func keyAttributesForItem(target queryTarget, item domain.Item) (map[string]domain.AttributeValue, error) { + names := targetKeyNames(target) + if len(names) == 0 { + return nil, fmt.Errorf("dynamodb: table %q has no key attributes", target.Table.Name) } - partitionValue, ok := item.Attributes[table.PartitionKey] - if !ok { - return nil, fmt.Errorf("dynamodb: item %s/%s is missing partition key %q", item.Table, item.Key, table.PartitionKey) - } - - attributes := map[string]domain.AttributeValue{ - table.PartitionKey: partitionValue.Clone(), - } - if table.SortKey != "" { - sortValue, ok := item.Attributes[table.SortKey] + attributes := make(map[string]domain.AttributeValue, len(names)) + for _, name := range names { + value, ok := item.Attributes[name] if !ok { - return nil, fmt.Errorf("dynamodb: item %s/%s is missing sort key %q", item.Table, item.Key, table.SortKey) + return nil, fmt.Errorf("dynamodb: item %s/%s is missing key attribute %q", item.Table, item.Key, name) } - attributes[table.SortKey] = sortValue.Clone() + attributes[name] = value.Clone() } return attributes, nil } -func normalizeKeyAttributes(table domain.Table, attributes map[string]domain.AttributeValue) (map[string]domain.AttributeValue, error) { +func normalizeKeyAttributes(target queryTarget, attributes map[string]domain.AttributeValue) (map[string]domain.AttributeValue, error) { if len(attributes) == 0 { return nil, nil } - expected := []string{table.PartitionKey} - if table.SortKey != "" { - expected = append(expected, table.SortKey) - } - + expected := targetKeyNames(target) if len(attributes) != len(expected) { return nil, fmt.Errorf("dynamodb: unsupported key attributes %q", strings.Join(sortedMapKeys(attributes), ", ")) } @@ -525,14 +675,58 @@ func normalizeKeyAttributes(table domain.Table, attributes map[string]domain.Att } normalized[name] = value.Clone() } + return normalized, nil +} - for name := range attributes { - if name != table.PartitionKey && (table.SortKey == "" || name != table.SortKey) { - return nil, fmt.Errorf("dynamodb: unsupported key attribute %q", name) +func targetKeyNames(target queryTarget) []string { + names := []string{target.PartitionKey} + if target.SortKey != "" { + names = append(names, target.SortKey) + } + if target.Index != nil { + names = append(names, target.Table.PartitionKey) + if target.Table.SortKey != "" { + names = append(names, target.Table.SortKey) } } + return uniqueStrings(names) +} - return normalized, nil +func orderingKeyNames(target queryTarget) []string { + names := targetKeyNames(target) + return names +} + +func uniqueStrings(values []string) []string { + if len(values) == 0 { + return nil + } + seen := make(map[string]struct{}, len(values)) + unique := make([]string, 0, len(values)) + for _, value := range values { + value = strings.TrimSpace(value) + if value == "" { + continue + } + if _, ok := seen[value]; ok { + continue + } + seen[value] = struct{}{} + unique = append(unique, value) + } + return unique +} + +func sortedSetKeys(values map[string]struct{}) []string { + if len(values) == 0 { + return nil + } + keys := make([]string, 0, len(values)) + for value := range values { + keys = append(keys, value) + } + sort.Strings(keys) + return keys } func attributeDocumentsEqual(left, right map[string]domain.AttributeValue) bool { diff --git a/core/internal/resources/dynamodb/application/service_query_test.go b/core/internal/resources/dynamodb/application/service_query_test.go index ab34f90..4af6d72 100644 --- a/core/internal/resources/dynamodb/application/service_query_test.go +++ b/core/internal/resources/dynamodb/application/service_query_test.go @@ -204,6 +204,154 @@ func TestServiceReadPlannerRejectsUnsupportedExpressions(t *testing.T) { } } +func TestServiceQuerySupportsIndexedPaginationAndProjection(t *testing.T) { + t.Helper() + + service := New() + _, err := service.CreateTable("mildstack-indexed", "pk", "sk", "PAY_PER_REQUEST", domain.CreateTableSpec{ + AttributeDefinitions: []domain.AttributeDefinition{ + {Name: "pk", Type: "S"}, + {Name: "sk", Type: "S"}, + {Name: "gsi_pk", Type: "S"}, + {Name: "gsi_sk", Type: "S"}, + {Name: "lsi_sk", Type: "S"}, + {Name: "title", Type: "S"}, + }, + GlobalSecondaryIndexes: []domain.SecondaryIndex{ + { + Name: "gsi-title", + KeySchema: []domain.KeySchemaElement{ + {AttributeName: "gsi_pk", KeyType: "HASH"}, + {AttributeName: "gsi_sk", KeyType: "RANGE"}, + }, + Projection: domain.Projection{ + Type: "INCLUDE", + NonKeyAttributes: []string{"title"}, + }, + }, + }, + LocalSecondaryIndexes: []domain.SecondaryIndex{ + { + Name: "lsi-title", + KeySchema: []domain.KeySchemaElement{ + {AttributeName: "pk", KeyType: "HASH"}, + {AttributeName: "lsi_sk", KeyType: "RANGE"}, + }, + Projection: domain.Projection{ + Type: "KEYS_ONLY", + }, + }, + }, + }) + if err != nil { + t.Fatalf("create indexed table: %v", err) + } + + for i, item := range []domain.Item{ + { + Table: "mildstack-indexed", + Key: "row#1", + Attributes: map[string]domain.AttributeValue{ + "pk": domain.StringValue("series#1"), + "sk": domain.StringValue("001"), + "gsi_pk": domain.StringValue("group#1"), + "gsi_sk": domain.StringValue("001"), + "lsi_sk": domain.StringValue("001"), + "title": domain.StringValue("indexed-one"), + }, + }, + { + Table: "mildstack-indexed", + Key: "row#2", + Attributes: map[string]domain.AttributeValue{ + "pk": domain.StringValue("series#1"), + "sk": domain.StringValue("002"), + "gsi_pk": domain.StringValue("group#1"), + "gsi_sk": domain.StringValue("002"), + "lsi_sk": domain.StringValue("002"), + "title": domain.StringValue("indexed-two"), + }, + }, + } { + if _, err := service.PutItem(item.Table, item.Key, item.Attributes); err != nil { + t.Fatalf("put indexed item %d: %v", i, err) + } + } + + gsiPage1, err := service.Query("mildstack-indexed", "gsi_pk = :pk AND gsi_sk BETWEEN :start AND :end", "", nil, map[string]domain.AttributeValue{ + ":pk": domain.StringValue("group#1"), + ":start": domain.StringValue("001"), + ":end": domain.StringValue("002"), + }, intPtr(1), nil, boolPtr(true), domain.QueryOptions{ + IndexName: "gsi-title", + ProjectionExpression: "gsi_pk, title", + }) + if err != nil { + t.Fatalf("query gsi page 1: %v", err) + } + if got, want := gsiPage1.Count, 1; got != want { + t.Fatalf("unexpected gsi page 1 count: got %d want %d", got, want) + } + if got, want := attrString(gsiPage1.Items[0].Attributes["title"]), "indexed-one"; got != want { + t.Fatalf("unexpected gsi page 1 title: got %q want %q", got, want) + } + if _, ok := gsiPage1.Items[0].Attributes["gsi_sk"]; ok { + t.Fatal("expected projected gsi sort key to be omitted") + } + + gsiPage2, err := service.Query("mildstack-indexed", "gsi_pk = :pk AND gsi_sk BETWEEN :start AND :end", "", nil, map[string]domain.AttributeValue{ + ":pk": domain.StringValue("group#1"), + ":start": domain.StringValue("001"), + ":end": domain.StringValue("002"), + }, intPtr(1), gsiPage1.LastEvaluatedKey, boolPtr(true), domain.QueryOptions{ + IndexName: "gsi-title", + ProjectionExpression: "gsi_pk, title", + }) + if err != nil { + t.Fatalf("query gsi page 2: %v", err) + } + if got, want := gsiPage2.Count, 1; got != want { + t.Fatalf("unexpected gsi page 2 count: got %d want %d", got, want) + } + if got, want := attrString(gsiPage2.Items[0].Attributes["title"]), "indexed-two"; got != want { + t.Fatalf("unexpected gsi page 2 title: got %q want %q", got, want) + } + + lsiPage1, err := service.Query("mildstack-indexed", "pk = :pk AND lsi_sk BETWEEN :start AND :end", "", nil, map[string]domain.AttributeValue{ + ":pk": domain.StringValue("series#1"), + ":start": domain.StringValue("001"), + ":end": domain.StringValue("002"), + }, intPtr(1), nil, boolPtr(true), domain.QueryOptions{ + IndexName: "lsi-title", + }) + if err != nil { + t.Fatalf("query lsi page 1: %v", err) + } + if got, want := lsiPage1.Count, 1; got != want { + t.Fatalf("unexpected lsi page 1 count: got %d want %d", got, want) + } + if got, want := attrString(lsiPage1.Items[0].Attributes["lsi_sk"]), "001"; got != want { + t.Fatalf("unexpected lsi page 1 sort key: got %q want %q", got, want) + } + + lsiPage2, err := service.Query("mildstack-indexed", "pk = :pk AND lsi_sk BETWEEN :start AND :end", "", nil, map[string]domain.AttributeValue{ + ":pk": domain.StringValue("series#1"), + ":start": domain.StringValue("001"), + ":end": domain.StringValue("002"), + }, intPtr(1), lsiPage1.LastEvaluatedKey, boolPtr(true), domain.QueryOptions{ + IndexName: "lsi-title", + }) + if err != nil { + t.Fatalf("query lsi page 2: %v", err) + } + if got, want := lsiPage2.Count, 1; got != want { + t.Fatalf("unexpected lsi page 2 count: got %d want %d", got, want) + } + if got, want := attrString(lsiPage2.Items[0].Attributes["lsi_sk"]), "002"; got != want { + t.Fatalf("unexpected lsi page 2 sort key: got %q want %q", got, want) + } +} + func seedQueryItems(t *testing.T, service *Service) { t.Helper() diff --git a/core/internal/resources/dynamodb/application/service_update.go b/core/internal/resources/dynamodb/application/service_update.go index 15b6d79..53c325c 100644 --- a/core/internal/resources/dynamodb/application/service_update.go +++ b/core/internal/resources/dynamodb/application/service_update.go @@ -149,7 +149,15 @@ func parseUpdateClause(kind, body string, expressionAttributeNames map[string]st if err != nil { return nil, err } - value, err := resolveUpdateValue(strings.TrimSpace(part[equalIndex+1:]), expressionAttributeValues) + rawValue := strings.TrimSpace(part[equalIndex+1:]) + if addValue, ok, err := resolveSelfAddUpdateValue(path, rawValue, expressionAttributeNames, expressionAttributeValues); err != nil { + return nil, err + } else if ok { + operations = append(operations, updateOperation{kind: "ADD", path: path, value: addValue}) + continue + } + + value, err := resolveUpdateValue(rawValue, expressionAttributeValues) if err != nil { return nil, err } @@ -276,6 +284,31 @@ func resolveUpdateValue(raw string, expressionAttributeValues map[string]domain. return value.Clone(), nil } +func resolveSelfAddUpdateValue(targetPath, raw string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]domain.AttributeValue) (domain.AttributeValue, bool, error) { + parts := strings.Split(raw, "+") + if len(parts) != 2 { + return domain.AttributeValue{}, false, nil + } + + leftPath, err := resolveUpdatePath(parts[0], expressionAttributeNames) + if err != nil { + return domain.AttributeValue{}, false, err + } + if leftPath != targetPath { + return domain.AttributeValue{}, false, fmt.Errorf("dynamodb: unsupported update expression %q", raw) + } + + value, err := resolveUpdateValue(parts[1], expressionAttributeValues) + if err != nil { + return domain.AttributeValue{}, false, err + } + if value.N == nil { + return domain.AttributeValue{}, false, fmt.Errorf("dynamodb: ADD requires a numeric value") + } + + return value, true, nil +} + func addAttribute(existing domain.AttributeValue, delta domain.AttributeValue) (domain.AttributeValue, error) { if delta.N == nil { return domain.AttributeValue{}, fmt.Errorf("dynamodb: ADD requires a numeric value") diff --git a/core/internal/resources/dynamodb/contracts/contracts.go b/core/internal/resources/dynamodb/contracts/contracts.go index efe1c7f..9a6829f 100644 --- a/core/internal/resources/dynamodb/contracts/contracts.go +++ b/core/internal/resources/dynamodb/contracts/contracts.go @@ -52,9 +52,15 @@ type BatchGetItemResult struct { } type TransactWriteItem struct { - Table string - PutItem map[string]domain.AttributeValue - DeleteKey map[string]domain.AttributeValue + Table string + PutItem map[string]domain.AttributeValue + DeleteKey map[string]domain.AttributeValue + UpdateKey map[string]domain.AttributeValue + UpdateExpression string + ConditionExpression string + ExpressionAttributeNames map[string]string + ExpressionAttributeValues map[string]domain.AttributeValue + ConditionCheckKey map[string]domain.AttributeValue } type TransactWriteItemsRequest struct { diff --git a/core/internal/resources/dynamodb/domain/state.go b/core/internal/resources/dynamodb/domain/state.go index 8c6450b..0bb6b90 100644 --- a/core/internal/resources/dynamodb/domain/state.go +++ b/core/internal/resources/dynamodb/domain/state.go @@ -21,14 +21,49 @@ type State struct { } type Table struct { - Name string - PartitionKey string - SortKey string - BillingMode string - Status string - CreatedAt time.Time - ActivationAt time.Time - DeletedAt time.Time + Name string + PartitionKey string + SortKey string + BillingMode string + AttributeDefinitions []AttributeDefinition + GlobalSecondaryIndexes []SecondaryIndex + LocalSecondaryIndexes []SecondaryIndex + Status string + CreatedAt time.Time + ActivationAt time.Time + DeletedAt time.Time +} + +type AttributeDefinition struct { + Name string + Type string +} + +type KeySchemaElement struct { + AttributeName string + KeyType string +} + +type Projection struct { + Type string + NonKeyAttributes []string +} + +type SecondaryIndex struct { + Name string + KeySchema []KeySchemaElement + Projection Projection +} + +type CreateTableSpec struct { + AttributeDefinitions []AttributeDefinition + GlobalSecondaryIndexes []SecondaryIndex + LocalSecondaryIndexes []SecondaryIndex +} + +type QueryOptions struct { + IndexName string + ProjectionExpression string } type Item struct { @@ -244,14 +279,17 @@ func (s State) Snapshot() map[string]any { tables := make([]any, 0, len(s.Tables)) for _, table := range s.ListTables() { tables = append(tables, map[string]any{ - "name": table.Name, - "partition_key": table.PartitionKey, - "sort_key": table.SortKey, - "billing_mode": table.BillingMode, - "status": table.Status, - "created_at": snapshotTime(table.CreatedAt), - "activation_at": snapshotTime(table.ActivationAt), - "deleted_at": snapshotTime(table.DeletedAt), + "name": table.Name, + "partition_key": table.PartitionKey, + "sort_key": table.SortKey, + "billing_mode": table.BillingMode, + "attribute_definitions": copyAttributeDefinitions(table.AttributeDefinitions), + "global_secondary_indexes": copySecondaryIndexes(table.GlobalSecondaryIndexes), + "local_secondary_indexes": copySecondaryIndexes(table.LocalSecondaryIndexes), + "status": table.Status, + "created_at": snapshotTime(table.CreatedAt), + "activation_at": snapshotTime(table.ActivationAt), + "deleted_at": snapshotTime(table.DeletedAt), }) } @@ -285,6 +323,11 @@ func (s State) Clone() State { Attributes: cloneAttributes(item.Attributes), } } + for i := range cloned.Tables { + cloned.Tables[i].AttributeDefinitions = cloneAttributeDefinitions(cloned.Tables[i].AttributeDefinitions) + cloned.Tables[i].GlobalSecondaryIndexes = cloneSecondaryIndexes(cloned.Tables[i].GlobalSecondaryIndexes) + cloned.Tables[i].LocalSecondaryIndexes = cloneSecondaryIndexes(cloned.Tables[i].LocalSecondaryIndexes) + } return cloned } @@ -293,6 +336,9 @@ func normalizeTable(table Table) Table { table.PartitionKey = strings.TrimSpace(table.PartitionKey) table.SortKey = strings.TrimSpace(table.SortKey) table.BillingMode = strings.TrimSpace(table.BillingMode) + table.AttributeDefinitions = normalizeAttributeDefinitions(table.AttributeDefinitions) + table.GlobalSecondaryIndexes = normalizeSecondaryIndexes(table.GlobalSecondaryIndexes) + table.LocalSecondaryIndexes = normalizeSecondaryIndexes(table.LocalSecondaryIndexes) table.Status = strings.ToUpper(strings.TrimSpace(table.Status)) switch table.Status { @@ -422,3 +468,195 @@ func attributeValueToAny(value AttributeValue) any { return nil } } + +func cloneAttributeDefinitions(source []AttributeDefinition) []AttributeDefinition { + if len(source) == 0 { + return nil + } + cloned := make([]AttributeDefinition, len(source)) + copy(cloned, source) + return cloned +} + +func cloneSecondaryIndexes(source []SecondaryIndex) []SecondaryIndex { + if len(source) == 0 { + return nil + } + cloned := make([]SecondaryIndex, len(source)) + for i, index := range source { + cloned[i] = cloneSecondaryIndex(index) + } + return cloned +} + +func cloneSecondaryIndex(index SecondaryIndex) SecondaryIndex { + index.KeySchema = cloneKeySchema(index.KeySchema) + index.Projection = cloneProjection(index.Projection) + return index +} + +func cloneKeySchema(source []KeySchemaElement) []KeySchemaElement { + if len(source) == 0 { + return nil + } + cloned := make([]KeySchemaElement, len(source)) + copy(cloned, source) + return cloned +} + +func cloneProjection(projection Projection) Projection { + projection.NonKeyAttributes = cloneStrings(projection.NonKeyAttributes) + return projection +} + +func copyAttributeDefinitions(source []AttributeDefinition) []any { + if len(source) == 0 { + return nil + } + copied := make([]any, len(source)) + for i, definition := range source { + copied[i] = map[string]any{ + "name": definition.Name, + "type": definition.Type, + } + } + return copied +} + +func copySecondaryIndexes(source []SecondaryIndex) []any { + if len(source) == 0 { + return nil + } + copied := make([]any, len(source)) + for i, index := range source { + copied[i] = map[string]any{ + "name": index.Name, + "key_schema": copyKeySchema(index.KeySchema), + "projection": map[string]any{ + "type": index.Projection.Type, + "non_key_attributes": cloneStrings(index.Projection.NonKeyAttributes), + }, + } + } + return copied +} + +func copyKeySchema(source []KeySchemaElement) []any { + if len(source) == 0 { + return nil + } + copied := make([]any, len(source)) + for i, element := range source { + copied[i] = map[string]any{ + "attribute_name": element.AttributeName, + "key_type": element.KeyType, + } + } + return copied +} + +func normalizeAttributeDefinitions(source []AttributeDefinition) []AttributeDefinition { + if len(source) == 0 { + return nil + } + seen := make(map[string]struct{}, len(source)) + normalized := make([]AttributeDefinition, 0, len(source)) + for _, definition := range source { + definition.Name = strings.TrimSpace(definition.Name) + definition.Type = strings.ToUpper(strings.TrimSpace(definition.Type)) + if definition.Name == "" { + continue + } + if _, ok := seen[definition.Name]; ok { + continue + } + seen[definition.Name] = struct{}{} + normalized = append(normalized, definition) + } + return normalized +} + +func normalizeSecondaryIndexes(source []SecondaryIndex) []SecondaryIndex { + if len(source) == 0 { + return nil + } + normalized := make([]SecondaryIndex, 0, len(source)) + for _, index := range source { + index.Name = strings.TrimSpace(index.Name) + index.KeySchema = normalizeKeySchema(index.KeySchema) + index.Projection = normalizeProjection(index.Projection) + if index.Name == "" { + continue + } + normalized = append(normalized, index) + } + return normalized +} + +func normalizeKeySchema(source []KeySchemaElement) []KeySchemaElement { + if len(source) == 0 { + return nil + } + normalized := make([]KeySchemaElement, 0, len(source)) + seen := map[string]struct{}{} + for _, element := range source { + element.AttributeName = strings.TrimSpace(element.AttributeName) + element.KeyType = strings.ToUpper(strings.TrimSpace(element.KeyType)) + if element.AttributeName == "" || element.KeyType == "" { + continue + } + key := element.KeyType + ":" + element.AttributeName + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + normalized = append(normalized, element) + } + return normalized +} + +func normalizeProjection(projection Projection) Projection { + projection.Type = strings.ToUpper(strings.TrimSpace(projection.Type)) + switch projection.Type { + case "", "ALL": + projection.Type = "ALL" + projection.NonKeyAttributes = nil + case "KEYS_ONLY": + projection.NonKeyAttributes = nil + case "INCLUDE": + projection.NonKeyAttributes = uniqueStrings(projection.NonKeyAttributes) + default: + projection.Type = "ALL" + projection.NonKeyAttributes = nil + } + return projection +} + +func cloneStrings(values []string) []string { + if len(values) == 0 { + return nil + } + cloned := make([]string, len(values)) + copy(cloned, values) + return cloned +} + +func uniqueStrings(values []string) []string { + if len(values) == 0 { + return nil + } + seen := make(map[string]struct{}, len(values)) + unique := make([]string, 0, len(values)) + for _, value := range values { + value = strings.TrimSpace(value) + if value == "" { + continue + } + if _, ok := seen[value]; ok { + continue + } + seen[value] = struct{}{} + unique = append(unique, value) + } + return unique +} diff --git a/core/internal/resources/dynamodb/infrastructure/handlers.go b/core/internal/resources/dynamodb/infrastructure/handlers.go index 7b78a4c..2ade605 100644 --- a/core/internal/resources/dynamodb/infrastructure/handlers.go +++ b/core/internal/resources/dynamodb/infrastructure/handlers.go @@ -4,7 +4,7 @@ import "github.com/michasdev/mildstack/core/internal/resources/dynamodb/domain" type Service interface { ListTables() []domain.Table - CreateTable(name, partitionKey, sortKey, billingMode string) (domain.Table, error) + CreateTable(name, partitionKey, sortKey, billingMode string, specs ...domain.CreateTableSpec) (domain.Table, error) GetItem(table, key string) (domain.Item, error) PutItem(table, key string, attributes map[string]domain.AttributeValue) (domain.Item, error) UpdateItem(table, key, updateExpression, conditionExpression string, expressionAttributeNames map[string]string, expressionAttributeValues map[string]domain.AttributeValue) (domain.Item, error) @@ -33,10 +33,13 @@ type ListTablesResponse struct { } type CreateTableRequest struct { - Name string - PartitionKey string - SortKey string - BillingMode string + Name string + PartitionKey string + SortKey string + BillingMode string + AttributeDefinitions []domain.AttributeDefinition + GlobalSecondaryIndexes []domain.SecondaryIndex + LocalSecondaryIndexes []domain.SecondaryIndex } type CreateTableResponse struct { @@ -63,12 +66,12 @@ type PutItemResponse struct { } type UpdateItemRequest struct { - Table string - Key string - UpdateExpression string - ConditionExpression string - ExpressionAttributeNames map[string]string - ExpressionAttributeValues map[string]domain.AttributeValue + Table string + Key string + UpdateExpression string + ConditionExpression string + ExpressionAttributeNames map[string]string + ExpressionAttributeValues map[string]domain.AttributeValue } type UpdateItemResponse struct { @@ -105,7 +108,11 @@ func (h Handlers) ListTables() ListTablesResponse { } func (h Handlers) CreateTable(request CreateTableRequest) (CreateTableResponse, error) { - table, err := h.service.CreateTable(request.Name, request.PartitionKey, request.SortKey, request.BillingMode) + table, err := h.service.CreateTable(request.Name, request.PartitionKey, request.SortKey, request.BillingMode, domain.CreateTableSpec{ + AttributeDefinitions: request.AttributeDefinitions, + GlobalSecondaryIndexes: request.GlobalSecondaryIndexes, + LocalSecondaryIndexes: request.LocalSecondaryIndexes, + }) if err != nil { return CreateTableResponse{}, err } diff --git a/core/internal/resources/s3/application/service_objects.go b/core/internal/resources/s3/application/service_objects.go index c1039e5..1b0a800 100644 --- a/core/internal/resources/s3/application/service_objects.go +++ b/core/internal/resources/s3/application/service_objects.go @@ -129,6 +129,10 @@ func (s *Service) HeadObject(bucket, key string) (domain.Object, error) { } func (s *Service) PutObject(bucket, key string, body io.Reader, contentType string) (domain.Object, error) { + return s.PutObjectWithMetadata(bucket, key, body, contentType, nil, nil) +} + +func (s *Service) PutObjectWithMetadata(bucket, key string, body io.Reader, contentType string, metadata, preservedHeaders map[string]string) (domain.Object, error) { s.mu.Lock() defer s.mu.Unlock() @@ -158,12 +162,14 @@ func (s *Service) PutObject(bucket, key string, body io.Reader, contentType stri return domain.Object{}, err } object, err := s.storeObject(domain.Object{ - Bucket: bucket, - Key: key, - Size: size, - ContentType: contentType, - ETag: etag, - PayloadRef: payloadRef, + Bucket: bucket, + Key: key, + Size: size, + ContentType: contentType, + ETag: etag, + Metadata: cloneObjectStringMap(metadata), + PreservedHeaders: cloneObjectStringMap(preservedHeaders), + PayloadRef: payloadRef, }) if err != nil { return domain.Object{}, err diff --git a/core/internal/resources/sqs/application/repository_sqlite.go b/core/internal/resources/sqs/application/repository_sqlite.go index 228d0d1..75f4334 100644 --- a/core/internal/resources/sqs/application/repository_sqlite.go +++ b/core/internal/resources/sqs/application/repository_sqlite.go @@ -21,7 +21,7 @@ import ( const ( sqliteFileName = "state.db" schemaVersionKey = "schema_version" - schemaVersion = "2" + schemaVersion = "4" ) type SQLiteRepository struct { @@ -133,7 +133,9 @@ func (r *SQLiteRepository) bootstrap() error { dead_letter_queue TEXT NOT NULL DEFAULT '', policy_json TEXT NOT NULL, created_at_ns INTEGER NOT NULL DEFAULT 0, - updated_at_ns INTEGER NOT NULL DEFAULT 0 + updated_at_ns INTEGER NOT NULL DEFAULT 0, + deleted_at_ns INTEGER NOT NULL DEFAULT 0, + purged_at_ns INTEGER NOT NULL DEFAULT 0 )`, `CREATE TABLE IF NOT EXISTS sqs_messages ( queue_name TEXT NOT NULL, @@ -167,6 +169,13 @@ func (r *SQLiteRepository) bootstrap() error { message_id TEXT NOT NULL, detail_json TEXT NOT NULL )`, + `CREATE TABLE IF NOT EXISTS sqs_queue_governance ( + queue_name TEXT PRIMARY KEY, + tags_json TEXT NOT NULL, + permissions_json TEXT NOT NULL, + move_tasks_json TEXT NOT NULL, + FOREIGN KEY (queue_name) REFERENCES sqs_queues(name) ON DELETE CASCADE + )`, } for _, statement := range statements { @@ -180,6 +189,14 @@ func (r *SQLiteRepository) bootstrap() error { _ = tx.Rollback() return fmt.Errorf("sqs: ensure queue ordering column: %w", err) } + if err := ensureColumn(ctx, tx, "sqs_queues", "deleted_at_ns INTEGER NOT NULL DEFAULT 0"); err != nil { + _ = tx.Rollback() + return fmt.Errorf("sqs: ensure queue deleted column: %w", err) + } + if err := ensureColumn(ctx, tx, "sqs_queues", "purged_at_ns INTEGER NOT NULL DEFAULT 0"); err != nil { + _ = tx.Rollback() + return fmt.Errorf("sqs: ensure queue purged column: %w", err) + } if err := ensureColumn(ctx, tx, "sqs_messages", "message_group_id TEXT NOT NULL DEFAULT ''"); err != nil { _ = tx.Rollback() return fmt.Errorf("sqs: ensure message group column: %w", err) @@ -238,7 +255,7 @@ func (r *SQLiteRepository) loadLocked() (domain.State, error) { state := domain.NewState() queueRows, err := r.db.QueryContext(ctx, ` - SELECT name, url, attributes_json, ordering_hint, dead_letter_queue, policy_json, created_at_ns, updated_at_ns + SELECT name, url, attributes_json, ordering_hint, dead_letter_queue, policy_json, created_at_ns, updated_at_ns, deleted_at_ns, purged_at_ns FROM sqs_queues ORDER BY name `) @@ -255,8 +272,10 @@ func (r *SQLiteRepository) loadLocked() (domain.State, error) { policy string createdAtNS int64 updatedAtNS int64 + deletedAtNS int64 + purgedAtNS int64 ) - if err := queueRows.Scan(&queue.Name, &queue.URL, &attributes, &ordering, &queue.Recovery.DeadLetterQueue, &policy, &createdAtNS, &updatedAtNS); err != nil { + if err := queueRows.Scan(&queue.Name, &queue.URL, &attributes, &ordering, &queue.Recovery.DeadLetterQueue, &policy, &createdAtNS, &updatedAtNS, &deletedAtNS, &purgedAtNS); err != nil { return domain.State{}, fmt.Errorf("sqs: scan queue: %w", err) } queue.Attributes, err = decodeStringMap(attributes) @@ -270,6 +289,8 @@ func (r *SQLiteRepository) loadLocked() (domain.State, error) { queue.OrderingHint = ordering queue.CreatedAt = unixNanoToTime(createdAtNS) queue.UpdatedAt = unixNanoToTime(updatedAtNS) + queue.DeletedAt = unixNanoToTime(deletedAtNS) + queue.PurgedAt = unixNanoToTime(purgedAtNS) state.Queues = append(state.Queues, queue) } if err := queueRows.Err(); err != nil { @@ -402,6 +423,67 @@ func (r *SQLiteRepository) loadLocked() (domain.State, error) { return domain.State{}, fmt.Errorf("sqs: iterate recovery metadata: %w", err) } + governanceRows, err := r.db.QueryContext(ctx, ` + SELECT queue_name, tags_json, permissions_json, move_tasks_json + FROM sqs_queue_governance + ORDER BY queue_name + `) + if err != nil { + return domain.State{}, fmt.Errorf("sqs: query governance: %w", err) + } + defer governanceRows.Close() + + for governanceRows.Next() { + var ( + queueName string + tagsJSON string + permissions string + moveTasks string + ) + if err := governanceRows.Scan(&queueName, &tagsJSON, &permissions, &moveTasks); err != nil { + return domain.State{}, fmt.Errorf("sqs: scan governance: %w", err) + } + if state.QueueTags == nil { + state.QueueTags = map[string]map[string]string{} + } + if state.QueuePermissions == nil { + state.QueuePermissions = map[string]map[string]domain.QueuePermission{} + } + if state.MoveTasks == nil { + state.MoveTasks = map[string]map[string]domain.MessageMoveTask{} + } + tags, err := decodeStringMap(tagsJSON) + if err != nil { + return domain.State{}, fmt.Errorf("sqs: decode governance tags: %w", err) + } + state.QueueTags[queueName] = tags + + var permissionsList []domain.QueuePermission + if err := json.Unmarshal([]byte(permissions), &permissionsList); err != nil { + return domain.State{}, fmt.Errorf("sqs: decode governance permissions: %w", err) + } + if len(permissionsList) > 0 { + state.QueuePermissions[queueName] = make(map[string]domain.QueuePermission, len(permissionsList)) + for _, permission := range permissionsList { + state.QueuePermissions[queueName][permission.Label] = permission + } + } + + var moveTaskList []domain.MessageMoveTask + if err := json.Unmarshal([]byte(moveTasks), &moveTaskList); err != nil { + return domain.State{}, fmt.Errorf("sqs: decode governance move tasks: %w", err) + } + if len(moveTaskList) > 0 { + state.MoveTasks[queueName] = make(map[string]domain.MessageMoveTask, len(moveTaskList)) + for _, task := range moveTaskList { + state.MoveTasks[queueName][task.TaskHandle] = task + } + } + } + if err := governanceRows.Err(); err != nil { + return domain.State{}, fmt.Errorf("sqs: iterate governance: %w", err) + } + return state, nil } @@ -419,6 +501,10 @@ func (r *SQLiteRepository) saveLocked(state domain.State) error { _ = tx.Rollback() return fmt.Errorf("sqs: clear recovery metadata: %w", err) } + if _, err := tx.ExecContext(ctx, `DELETE FROM sqs_queue_governance`); err != nil { + _ = tx.Rollback() + return fmt.Errorf("sqs: clear governance: %w", err) + } if _, err := tx.ExecContext(ctx, `DELETE FROM sqs_messages`); err != nil { _ = tx.Rollback() return fmt.Errorf("sqs: clear messages: %w", err) @@ -430,8 +516,8 @@ func (r *SQLiteRepository) saveLocked(state domain.State) error { queueStmt, err := tx.PrepareContext(ctx, ` INSERT INTO sqs_queues ( - name, url, attributes_json, ordering_hint, dead_letter_queue, policy_json, created_at_ns, updated_at_ns - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?) + name, url, attributes_json, ordering_hint, dead_letter_queue, policy_json, created_at_ns, updated_at_ns, deleted_at_ns, purged_at_ns + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) `) if err != nil { _ = tx.Rollback() @@ -459,12 +545,48 @@ func (r *SQLiteRepository) saveLocked(state domain.State) error { policy, timeToUnixNano(queue.CreatedAt), timeToUnixNano(queue.UpdatedAt), + timeToUnixNano(queue.DeletedAt), + timeToUnixNano(queue.PurgedAt), ); err != nil { _ = tx.Rollback() return fmt.Errorf("sqs: insert queue %q: %w", queue.Name, err) } } + governanceQueueNames := collectGovernanceQueueNames(normalized) + governanceStmt, err := tx.PrepareContext(ctx, ` + INSERT INTO sqs_queue_governance ( + queue_name, tags_json, permissions_json, move_tasks_json + ) VALUES (?, ?, ?, ?) + `) + if err != nil { + _ = tx.Rollback() + return fmt.Errorf("sqs: prepare governance insert: %w", err) + } + defer governanceStmt.Close() + + for _, queueName := range governanceQueueNames { + tags, err := encodeStringMap(normalized.QueueTags[queueName]) + if err != nil { + _ = tx.Rollback() + return fmt.Errorf("sqs: encode governance tags: %w", err) + } + permissionsJSON, err := json.Marshal(govPermissionSlice(normalized.QueuePermissions[queueName])) + if err != nil { + _ = tx.Rollback() + return fmt.Errorf("sqs: encode governance permissions: %w", err) + } + moveTasksJSON, err := json.Marshal(govMoveTaskSlice(normalized.MoveTasks[queueName])) + if err != nil { + _ = tx.Rollback() + return fmt.Errorf("sqs: encode governance move tasks: %w", err) + } + if _, err := governanceStmt.ExecContext(ctx, queueName, tags, string(permissionsJSON), string(moveTasksJSON)); err != nil { + _ = tx.Rollback() + return fmt.Errorf("sqs: insert governance %q: %w", queueName, err) + } + } + messageStmt, err := tx.PrepareContext(ctx, ` INSERT INTO sqs_messages ( queue_name, message_id, body, attributes_json, metadata_json, tags_json, receipt_keys_json, @@ -577,6 +699,57 @@ func (r *SQLiteRepository) saveLocked(state domain.State) error { return nil } +func collectGovernanceQueueNames(state domain.State) []string { + names := make(map[string]struct{}) + for name := range state.QueueTags { + names[name] = struct{}{} + } + for name := range state.QueuePermissions { + names[name] = struct{}{} + } + for name := range state.MoveTasks { + names[name] = struct{}{} + } + ordered := make([]string, 0, len(names)) + for name := range names { + ordered = append(ordered, name) + } + sort.Strings(ordered) + return ordered +} + +func govPermissionSlice(values map[string]domain.QueuePermission) []domain.QueuePermission { + if len(values) == 0 { + return []domain.QueuePermission{} + } + labels := make([]string, 0, len(values)) + for label := range values { + labels = append(labels, label) + } + sort.Strings(labels) + result := make([]domain.QueuePermission, 0, len(labels)) + for _, label := range labels { + result = append(result, values[label]) + } + return result +} + +func govMoveTaskSlice(values map[string]domain.MessageMoveTask) []domain.MessageMoveTask { + if len(values) == 0 { + return []domain.MessageMoveTask{} + } + handles := make([]string, 0, len(values)) + for handle := range values { + handles = append(handles, handle) + } + sort.Strings(handles) + result := make([]domain.MessageMoveTask, 0, len(handles)) + for _, handle := range handles { + result = append(result, values[handle]) + } + return result +} + func encodeStringMap(values map[string]string) (string, error) { if values == nil { return "{}", nil diff --git a/core/internal/resources/sqs/application/repository_sqlite_test.go b/core/internal/resources/sqs/application/repository_sqlite_test.go index 53aa923..fcc646d 100644 --- a/core/internal/resources/sqs/application/repository_sqlite_test.go +++ b/core/internal/resources/sqs/application/repository_sqlite_test.go @@ -81,6 +81,8 @@ func TestSQLiteRepositoryPersistsQueueAndMessageStateAcrossRestart(t *testing.T) }, CreatedAt: createdAt, UpdatedAt: createdAt.Add(time.Minute), + DeletedAt: createdAt.Add(2 * time.Minute), + PurgedAt: createdAt.Add(3 * time.Minute), }) state.Messages = append(state.Messages, domain.Message{ Queue: "queue-a", @@ -112,6 +114,27 @@ func TestSQLiteRepositoryPersistsQueueAndMessageStateAcrossRestart(t *testing.T) Message: "message-1", Detail: map[string]string{"reason": "retry"}, } + state.QueueTags["queue-a"] = map[string]string{ + "env": "dev", + } + state.QueuePermissions["queue-a"] = map[string]domain.QueuePermission{ + "label-a": { + Label: "label-a", + AWSAccountIDs: []string{"123456789012"}, + Actions: []string{"SendMessage"}, + }, + } + state.MoveTasks["queue-a"] = map[string]domain.MessageMoveTask{ + "task-1": { + TaskHandle: "task-1", + SourceQueue: "queue-a", + SourceArn: "arn:aws:sqs:us-east-1:123456789012:queue-a", + DestinationArn: "arn:aws:sqs:us-east-1:123456789012:queue-dlq", + MaxNumberOfMessagesPerSecond: 10, + ApproximateNumberOfMessagesMoved: 2, + Status: "RUNNING", + }, + } if err := repo.Save(state); err != nil { t.Fatalf("save state: %v", err) @@ -151,6 +174,12 @@ func TestSQLiteRepositoryPersistsQueueAndMessageStateAcrossRestart(t *testing.T) if got, want := queue.Recovery.DeadLetterQueue, "queue-dlq"; got != want { t.Fatalf("unexpected dead-letter queue after restart: got %q want %q", got, want) } + if got, want := queue.DeletedAt, createdAt.Add(2*time.Minute); !got.Equal(want) { + t.Fatalf("unexpected deleted_at after restart: got %v want %v", got, want) + } + if got, want := queue.PurgedAt, createdAt.Add(3*time.Minute); !got.Equal(want) { + t.Fatalf("unexpected purged_at after restart: got %v want %v", got, want) + } message := loaded.Messages[0] if got, want := message.Body, "payload"; got != want { @@ -195,6 +224,15 @@ func TestSQLiteRepositoryPersistsQueueAndMessageStateAcrossRestart(t *testing.T) if got, want := loaded.RecoveryMetadata["queue-a/message-1"].Detail["reason"], "retry"; got != want { t.Fatalf("unexpected recovery metadata after restart: got %q want %q", got, want) } + if got, want := loaded.QueueTags["queue-a"]["env"], "dev"; got != want { + t.Fatalf("unexpected queue tags after restart: got %q want %q", got, want) + } + if got, want := loaded.QueuePermissions["queue-a"]["label-a"].Actions[0], "SendMessage"; got != want { + t.Fatalf("unexpected queue permission after restart: got %q want %q", got, want) + } + if got, want := loaded.MoveTasks["queue-a"]["task-1"].Status, "RUNNING"; got != want { + t.Fatalf("unexpected move task status after restart: got %q want %q", got, want) + } statePath := filepath.Join(baseDir, "instances", "instance-a", "sqs", sqliteFileName) if _, err := os.Stat(statePath); err != nil { diff --git a/core/internal/resources/sqs/application/service.go b/core/internal/resources/sqs/application/service.go index 304a203..5cb6203 100644 --- a/core/internal/resources/sqs/application/service.go +++ b/core/internal/resources/sqs/application/service.go @@ -2,6 +2,9 @@ package application import ( "context" + "crypto/md5" + "encoding/hex" + "encoding/json" "errors" "fmt" "sort" @@ -10,7 +13,9 @@ import ( "sync" "time" + "github.com/google/uuid" "github.com/michasdev/mildstack/core/internal/application/orchestrator" + "github.com/michasdev/mildstack/core/internal/resources/awscontext" "github.com/michasdev/mildstack/core/internal/resources/sqs/contracts" "github.com/michasdev/mildstack/core/internal/resources/sqs/domain" "github.com/michasdev/mildstack/core/internal/resources/sqs/infrastructure" @@ -35,12 +40,16 @@ const ( maxLongPollWait = 20 * time.Second workerPollInterval = 50 * time.Millisecond leaseVisibilityTimeoutMetaKey = "visibility_timeout_seconds" + queueLifecycleCooldown = 60 * time.Second ) var ( errQueueNotFound = errors.New("sqs: queue not found") errReceiptHandleMismatch = errors.New("sqs: receipt handle does not match active lease") errInvalidVisibilityWindow = errors.New("sqs: visibility timeout must be non-negative") + errEmptyBatchRequest = errors.New("sqs: batch request is empty") + errTooManyBatchEntries = errors.New("sqs: batch request contains more than 10 entries") + errDuplicateBatchEntryIDs = errors.New("sqs: batch request contains duplicate entry IDs") ) func New() *Service { @@ -141,6 +150,573 @@ func (s *Service) Policy() orchestrator.EmulationPolicy { return s.policy.Clone() } +func (s *Service) QueueURL(queueName string) string { + return queueURLForAccount(queueName, "") +} + +func (s *Service) QueueARN(queueName string) string { + return queueARNForAccount(queueName, "") +} + +func (s *Service) CreateQueue(queueName string, attributes map[string]string) (domain.Queue, error) { + queueName = trimName(queueName) + if queueName == "" { + return domain.Queue{}, fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + now := s.clock.Now() + normalizedAttributes := cloneMap(attributes) + if normalizedAttributes == nil { + normalizedAttributes = map[string]string{} + } + recovery := queueRecoveryFromAttributes(normalizedAttributes) + + if index, queue, ok := s.queueRecordByNameLocked(queueName); ok { + if queue.DeletedAt.IsZero() { + if equalStringMaps(queue.Attributes, normalizedAttributes) { + return s.queueResponse(queueName, queue, normalizedAttributes), nil + } + return domain.Queue{}, fmt.Errorf("sqs: queue %q already exists with different attributes", queueName) + } + if now.Sub(queue.DeletedAt) < queueLifecycleCooldown { + return domain.Queue{}, fmt.Errorf("sqs: queue %q is still in delete cooldown", queueName) + } + + queue.URL = s.QueueURL(queueName) + queue.Attributes = normalizedAttributes + queue.Recovery = recovery + queue.OrderingHint = orderingHintFromAttributes(normalizedAttributes, queue.OrderingHint) + queue.CreatedAt = now + queue.UpdatedAt = now + queue.DeletedAt = time.Time{} + queue.PurgedAt = time.Time{} + s.state.Queues[index] = queue + if err := s.commitStateLocked(); err != nil { + return domain.Queue{}, err + } + return s.queueResponse(queueName, queue, normalizedAttributes), nil + } + + queue := domain.Queue{ + Name: queueName, + URL: s.QueueURL(queueName), + Attributes: normalizedAttributes, + Recovery: recovery, + OrderingHint: orderingHintFromAttributes(normalizedAttributes, ""), + CreatedAt: now, + UpdatedAt: now, + } + s.state.Queues = append(s.state.Queues, queue) + if err := s.commitStateLocked(); err != nil { + return domain.Queue{}, err + } + return s.queueResponse(queueName, queue, normalizedAttributes), nil +} + +func (s *Service) DeleteQueue(queueName string) error { + queueName = trimName(queueName) + if queueName == "" { + return fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + index, queue, ok := s.queueRecordByNameLocked(queueName) + if !ok { + return errQueueNotFound + } + now := s.clock.Now() + if !queue.DeletedAt.IsZero() { + if now.Sub(queue.DeletedAt) < queueLifecycleCooldown { + return fmt.Errorf("sqs: queue %q is still in delete cooldown", queueName) + } + return errQueueNotFound + } + + queue.DeletedAt = now + queue.UpdatedAt = now + queue.PurgedAt = time.Time{} + s.state.Queues[index] = queue + s.removeMessagesForQueueLocked(queueName) + return s.commitStateLocked() +} + +func (s *Service) GetQueueUrl(queueName, ownerAccountID string) (string, error) { + queueName = trimName(queueName) + if queueName == "" { + return "", fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + queue, ok := s.activeQueueByNameLocked(queueName) + if !ok || !ownerAccountMatches(ownerAccountID) { + return "", errQueueNotFound + } + if queue.URL == "" { + return queueURLForAccount(queueName, ownerAccountID), nil + } + return queue.URL, nil +} + +func (s *Service) ListQueues(queueNamePrefix string, maxResults int, nextToken, ownerAccountID string) ([]domain.Queue, string, error) { + queueNamePrefix = trimName(queueNamePrefix) + nextToken = trimName(nextToken) + ownerAccountID = trimName(ownerAccountID) + if maxResults < 0 { + maxResults = 0 + } + + s.mu.Lock() + defer s.mu.Unlock() + + if !ownerAccountMatches(ownerAccountID) { + return nil, "", errQueueNotFound + } + + activeQueues := make([]domain.Queue, 0, len(s.state.Queues)) + for _, queue := range s.state.ListQueues() { + if !queue.DeletedAt.IsZero() { + continue + } + if queueNamePrefix != "" && !strings.HasPrefix(queue.Name, queueNamePrefix) { + continue + } + if queue.URL == "" { + queue.URL = s.QueueURL(queue.Name) + } + activeQueues = append(activeQueues, queue) + } + + startIndex := 0 + if nextToken != "" { + startIndex = len(activeQueues) + for i, queue := range activeQueues { + if queue.Name > nextToken { + startIndex = i + break + } + if queue.Name == nextToken { + startIndex = i + 1 + } + } + } + if startIndex > len(activeQueues) { + startIndex = len(activeQueues) + } + + endIndex := len(activeQueues) + if maxResults > 0 && startIndex+maxResults < endIndex { + endIndex = startIndex + maxResults + } + + page := append([]domain.Queue(nil), activeQueues[startIndex:endIndex]...) + nextPageToken := "" + if endIndex < len(activeQueues) { + nextPageToken = activeQueues[endIndex-1].Name + } + + return page, nextPageToken, nil +} + +func (s *Service) PurgeQueue(queueName string) error { + queueName = trimName(queueName) + if queueName == "" { + return fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + index, queue, ok := s.queueRecordByNameLocked(queueName) + if !ok || !queue.DeletedAt.IsZero() { + return errQueueNotFound + } + now := s.clock.Now() + if !queue.PurgedAt.IsZero() && now.Sub(queue.PurgedAt) < queueLifecycleCooldown { + return fmt.Errorf("sqs: queue %q is still in purge cooldown", queueName) + } + + s.removeMessagesForQueueLocked(queueName) + queue.PurgedAt = now + queue.UpdatedAt = now + s.state.Queues[index] = queue + return s.commitStateLocked() +} + +func (s *Service) GetQueueAttributes(queueName string, attributeNames []string, ownerAccountID string) (contracts.QueueAttributesView, error) { + queueName = trimName(queueName) + ownerAccountID = trimName(ownerAccountID) + if queueName == "" { + return contracts.QueueAttributesView{}, fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + queue, ok := s.activeQueueByNameLocked(queueName) + if !ok || !ownerAccountMatches(ownerAccountID) { + return contracts.QueueAttributesView{}, errQueueNotFound + } + + attributes := selectQueueAttributes(queue.Attributes, attributeNames, s.QueueARN(queueName)) + if attributes == nil { + attributes = map[string]string{} + } + + return contracts.QueueAttributesView{ + QueueName: queueName, + QueueURL: queueURLForAccount(queueName, ownerAccountID), + QueueARN: queueARNForAccount(queueName, ownerAccountID), + Attributes: attributes, + }, nil +} + +func (s *Service) SetQueueAttributes(queueName string, attributes map[string]string) (contracts.QueueAttributesView, error) { + queueName = trimName(queueName) + if queueName == "" { + return contracts.QueueAttributesView{}, fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + index, queue, ok := s.queueRecordByNameLocked(queueName) + if !ok || !queue.DeletedAt.IsZero() { + return contracts.QueueAttributesView{}, errQueueNotFound + } + + normalizedAttributes := cloneMap(queue.Attributes) + if normalizedAttributes == nil { + normalizedAttributes = map[string]string{} + } + for key, value := range attributes { + normalizedAttributes[trimName(key)] = value + } + + queue.Attributes = normalizedAttributes + queue.Recovery = queueRecoveryFromAttributes(normalizedAttributes) + queue.OrderingHint = orderingHintFromAttributes(normalizedAttributes, queue.OrderingHint) + queue.UpdatedAt = s.clock.Now() + s.state.Queues[index] = queue + + if err := s.commitStateLocked(); err != nil { + return contracts.QueueAttributesView{}, err + } + + return contracts.QueueAttributesView{ + QueueName: queueName, + QueueURL: queueURLForAccount(queueName, ""), + QueueARN: queueARNForAccount(queueName, ""), + Attributes: cloneMap(normalizedAttributes), + }, nil +} + +func (s *Service) TagQueue(queueName string, tags map[string]string) error { + queueName = trimName(queueName) + if queueName == "" { + return fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + queue, ok := s.activeQueueByNameLocked(queueName) + if !ok { + return errQueueNotFound + } + + if s.state.QueueTags == nil { + s.state.QueueTags = map[string]map[string]string{} + } + current := cloneMap(s.state.QueueTags[queueName]) + if current == nil { + current = map[string]string{} + } + for key, value := range tags { + key = trimName(key) + if key == "" { + continue + } + current[key] = value + } + s.state.QueueTags[queue.Name] = current + queue.UpdatedAt = s.clock.Now() + if index, _, ok := s.queueRecordByNameLocked(queueName); ok { + s.state.Queues[index] = queue + } + return s.commitStateLocked() +} + +func (s *Service) UntagQueue(queueName string, tagKeys []string) error { + queueName = trimName(queueName) + if queueName == "" { + return fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + queue, ok := s.activeQueueByNameLocked(queueName) + if !ok { + return errQueueNotFound + } + + current := cloneMap(s.state.QueueTags[queueName]) + if current == nil { + current = map[string]string{} + } + for _, tagKey := range tagKeys { + tagKey = trimName(tagKey) + if tagKey == "" { + continue + } + delete(current, tagKey) + } + s.state.QueueTags[queue.Name] = current + queue.UpdatedAt = s.clock.Now() + if index, _, ok := s.queueRecordByNameLocked(queueName); ok { + s.state.Queues[index] = queue + } + return s.commitStateLocked() +} + +func (s *Service) AddPermission(queueName, label string, awsAccountIDs, actions []string) error { + queueName = trimName(queueName) + label = trimName(label) + if queueName == "" { + return fmt.Errorf("sqs: queue name is required") + } + if label == "" { + return fmt.Errorf("sqs: permission label is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + queue, ok := s.activeQueueByNameLocked(queueName) + if !ok { + return errQueueNotFound + } + + if s.state.QueuePermissions == nil { + s.state.QueuePermissions = map[string]map[string]domain.QueuePermission{} + } + permissions := s.state.QueuePermissions[queueName] + if permissions == nil { + permissions = map[string]domain.QueuePermission{} + } + + now := s.clock.Now() + permissions[label] = domain.QueuePermission{ + Label: label, + AWSAccountIDs: uniqueSortedTrimmedStrings(awsAccountIDs), + Actions: uniqueSortedTrimmedStrings(actions), + CreatedAt: now, + UpdatedAt: now, + } + s.state.QueuePermissions[queue.Name] = permissions + queue.UpdatedAt = now + if index, _, ok := s.queueRecordByNameLocked(queueName); ok { + s.state.Queues[index] = queue + } + return s.commitStateLocked() +} + +func (s *Service) RemovePermission(queueName, label string) error { + queueName = trimName(queueName) + label = trimName(label) + if queueName == "" { + return fmt.Errorf("sqs: queue name is required") + } + if label == "" { + return fmt.Errorf("sqs: permission label is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + queue, ok := s.activeQueueByNameLocked(queueName) + if !ok { + return errQueueNotFound + } + + permissions := cloneQueuePermissionMap(s.state.QueuePermissions[queueName]) + delete(permissions, label) + if s.state.QueuePermissions == nil { + s.state.QueuePermissions = map[string]map[string]domain.QueuePermission{} + } + s.state.QueuePermissions[queue.Name] = permissions + queue.UpdatedAt = s.clock.Now() + if index, _, ok := s.queueRecordByNameLocked(queueName); ok { + s.state.Queues[index] = queue + } + return s.commitStateLocked() +} + +func (s *Service) ListQueueTags(queueName string) (map[string]string, error) { + queueName = trimName(queueName) + if queueName == "" { + return map[string]string{}, fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + if _, ok := s.activeQueueByNameLocked(queueName); !ok { + return map[string]string{}, errQueueNotFound + } + return cloneMap(s.state.QueueTags[queueName]), nil +} + +func (s *Service) ListDeadLetterSourceQueues(queueName string) ([]string, error) { + queueName = trimName(queueName) + if queueName == "" { + return nil, fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + if _, ok := s.activeQueueByNameLocked(queueName); !ok { + return nil, errQueueNotFound + } + + sources := make([]string, 0) + for _, queue := range s.state.ListQueues() { + if queue.DeletedAt.IsZero() && queue.Recovery.DeadLetterQueue == queueName { + sources = append(sources, queue.Name) + } + } + sort.Strings(sources) + return sources, nil +} + +func (s *Service) StartMessageMoveTask(sourceArn, destinationArn string, maxNumberOfMessagesPerSecond int) (string, error) { + sourceArn = trimName(sourceArn) + destinationArn = trimName(destinationArn) + if sourceArn == "" { + return "", fmt.Errorf("sqs: source ARN is required") + } + if maxNumberOfMessagesPerSecond < 0 { + return "", fmt.Errorf("sqs: max number of messages per second must be non-negative") + } + + s.mu.Lock() + defer s.mu.Unlock() + + sourceQueue, ok := s.queueByARNLocked(sourceArn) + if !ok { + return "", errQueueNotFound + } + for _, task := range s.state.MoveTasks[sourceQueue.Name] { + if strings.EqualFold(task.Status, "RUNNING") { + return "", fmt.Errorf("sqs: a message move task is already running for %s", sourceQueue.Name) + } + } + + hasSourceQueue := false + for _, queue := range s.state.ListQueues() { + if queue.DeletedAt.IsZero() && queue.Recovery.DeadLetterQueue == sourceQueue.Name { + hasSourceQueue = true + break + } + } + if !hasSourceQueue { + return "", fmt.Errorf("sqs: source queue is not configured as a dead-letter queue") + } + + if s.state.MoveTasks == nil { + s.state.MoveTasks = map[string]map[string]domain.MessageMoveTask{} + } + tasks := s.state.MoveTasks[sourceQueue.Name] + if tasks == nil { + tasks = map[string]domain.MessageMoveTask{} + } + + now := s.clock.Now() + task := domain.MessageMoveTask{ + TaskHandle: sourceArn + "|" + uuid.NewString(), + SourceQueue: sourceQueue.Name, + SourceArn: sourceArn, + DestinationArn: destinationArn, + MaxNumberOfMessagesPerSecond: maxNumberOfMessagesPerSecond, + Status: "RUNNING", + StartedAt: now, + UpdatedAt: now, + } + tasks[task.TaskHandle] = task + s.state.MoveTasks[sourceQueue.Name] = tasks + if index, _, ok := s.queueRecordByNameLocked(sourceQueue.Name); ok { + sourceQueue.UpdatedAt = now + s.state.Queues[index] = sourceQueue + } + if err := s.commitStateLocked(); err != nil { + return "", err + } + return task.TaskHandle, nil +} + +func (s *Service) CancelMessageMoveTask(taskHandle string) (int64, error) { + taskHandle = trimName(taskHandle) + if taskHandle == "" { + return 0, fmt.Errorf("sqs: task handle is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + queueName, task, ok := s.findMessageMoveTaskLocked(taskHandle) + if !ok { + return 0, errQueueNotFound + } + now := s.clock.Now() + task.Status = "CANCELLED" + task.CancelledAt = now + task.UpdatedAt = now + if s.state.MoveTasks == nil { + s.state.MoveTasks = map[string]map[string]domain.MessageMoveTask{} + } + tasks := cloneMessageMoveTaskMap(s.state.MoveTasks[queueName]) + tasks[taskHandle] = task + s.state.MoveTasks[queueName] = tasks + if err := s.commitStateLocked(); err != nil { + return 0, err + } + return task.ApproximateNumberOfMessagesMoved, nil +} + +func (s *Service) ListMessageMoveTasks(queueName string) ([]domain.MessageMoveTask, error) { + queueName = trimName(queueName) + if queueName == "" { + return nil, fmt.Errorf("sqs: queue name is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + if _, ok := s.activeQueueByNameLocked(queueName); !ok { + return nil, errQueueNotFound + } + + tasks := make([]domain.MessageMoveTask, 0, len(s.state.MoveTasks[queueName])) + for _, task := range s.state.MoveTasks[queueName] { + tasks = append(tasks, task) + } + sort.SliceStable(tasks, func(i, j int) bool { + if tasks[i].StartedAt.Equal(tasks[j].StartedAt) { + return tasks[i].TaskHandle < tasks[j].TaskHandle + } + return tasks[i].StartedAt.After(tasks[j].StartedAt) + }) + return tasks, nil +} + func (s *Service) RegisterRoutes(registrar orchestrator.RouteRegistrar) error { for _, route := range infrastructure.Routes() { if err := registrar.Register(route); err != nil { @@ -205,103 +781,492 @@ func (s *Service) ReceiveMessage(queueName string, maxMessages int, waitTime tim return messages, nil } - remaining := deadline.Sub(now) - sleep := workerPollInterval - if remaining < sleep { - sleep = remaining + remaining := deadline.Sub(now) + sleep := workerPollInterval + if remaining < sleep { + sleep = remaining + } + if sleep <= 0 { + return messages, nil + } + s.clock.Sleep(sleep) + } +} + +func (s *Service) DeleteMessage(queueName string, receiptHandle string) error { + queueName = trimName(queueName) + receiptHandle = trimName(receiptHandle) + if queueName == "" { + return fmt.Errorf("sqs: queue name is required") + } + if receiptHandle == "" { + return fmt.Errorf("sqs: receipt handle is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + idx, ok := s.findMessageByReceiptLocked(queueName, receiptHandle) + if !ok { + return errReceiptHandleMismatch + } + + s.state.Messages = append(s.state.Messages[:idx], s.state.Messages[idx+1:]...) + return s.commitStateLocked() +} + +func (s *Service) ChangeMessageVisibility(queueName string, receiptHandle string, visibility time.Duration) error { + queueName = trimName(queueName) + receiptHandle = trimName(receiptHandle) + if queueName == "" { + return fmt.Errorf("sqs: queue name is required") + } + if receiptHandle == "" { + return fmt.Errorf("sqs: receipt handle is required") + } + if visibility < 0 { + return errInvalidVisibilityWindow + } + if visibility > 12*time.Hour { + visibility = 12 * time.Hour + } + + s.mu.Lock() + defer s.mu.Unlock() + + idx, ok := s.findMessageByReceiptLocked(queueName, receiptHandle) + if !ok { + return errReceiptHandleMismatch + } + + message := &s.state.Messages[idx] + if message.Metadata == nil { + message.Metadata = map[string]string{} + } + if visibility == 0 { + message.ReceivedAt = time.Time{} + message.AvailableAt = s.clock.Now() + message.Metadata[leaseVisibilityTimeoutMetaKey] = "0" + } else { + message.ReceivedAt = s.clock.Now() + message.Metadata[leaseVisibilityTimeoutMetaKey] = strconv.FormatInt(int64(visibility/time.Second), 10) + } + + return s.commitStateLocked() +} + +func (s *Service) SendMessage(queueName string, request contracts.SendMessageRequest) (contracts.SendMessageResult, error) { + queueName = trimName(queueName) + if queueName == "" { + return contracts.SendMessageResult{}, fmt.Errorf("sqs: queue name is required") + } + if trimName(request.MessageBody) == "" { + return contracts.SendMessageResult{}, fmt.Errorf("sqs: message body is required") + } + + s.mu.Lock() + defer s.mu.Unlock() + + queue, ok := s.activeQueueByNameLocked(queueName) + if !ok { + return contracts.SendMessageResult{}, errQueueNotFound + } + + message, result, err := s.enqueueMessageLocked(queueName, queue, request, "", "", 0, 1) + if err != nil { + return contracts.SendMessageResult{}, err + } + s.state.Messages = append(s.state.Messages, message) + if err := s.commitStateLocked(); err != nil { + return contracts.SendMessageResult{}, err + } + return result, nil +} + +func (s *Service) SendMessageBatch(queueName string, request contracts.SendMessageBatchRequest) (contracts.SendMessageBatchResult, error) { + queueName = trimName(queueName) + if queueName == "" { + return contracts.SendMessageBatchResult{}, fmt.Errorf("sqs: queue name is required") + } + if len(request.Entries) == 0 { + return contracts.SendMessageBatchResult{}, errEmptyBatchRequest + } + if len(request.Entries) > 10 { + return contracts.SendMessageBatchResult{}, errTooManyBatchEntries + } + if hasDuplicateBatchEntryIDsByID(request.Entries, func(entry contracts.SendMessageBatchRequestEntry) string { return entry.Id }) { + return contracts.SendMessageBatchResult{}, errDuplicateBatchEntryIDs + } + + s.mu.Lock() + defer s.mu.Unlock() + + queue, ok := s.activeQueueByNameLocked(queueName) + if !ok { + return contracts.SendMessageBatchResult{}, errQueueNotFound + } + + result := contracts.SendMessageBatchResult{ + Successful: make([]contracts.SendMessageBatchResultEntry, 0, len(request.Entries)), + Failed: make([]contracts.BatchResultErrorEntry, 0), + } + pending := make([]domain.Message, 0, len(request.Entries)) + for index, entry := range request.Entries { + entryResult, message, err := s.enqueueBatchMessageLocked(queueName, queue, entry, index, len(request.Entries)) + if err != nil { + result.Failed = append(result.Failed, batchFailureEntry(entry.Id, err.Error(), true)) + continue + } + pending = append(pending, message) + result.Successful = append(result.Successful, entryResult) + } + + if len(pending) > 0 { + s.state.Messages = append(s.state.Messages, pending...) + if err := s.commitStateLocked(); err != nil { + return contracts.SendMessageBatchResult{}, err + } + } + + return result, nil +} + +func (s *Service) DeleteMessageBatch(queueName string, request contracts.DeleteMessageBatchRequest) (contracts.DeleteMessageBatchResult, error) { + queueName = trimName(queueName) + if queueName == "" { + return contracts.DeleteMessageBatchResult{}, fmt.Errorf("sqs: queue name is required") + } + if len(request.Entries) == 0 { + return contracts.DeleteMessageBatchResult{}, errEmptyBatchRequest + } + if len(request.Entries) > 10 { + return contracts.DeleteMessageBatchResult{}, errTooManyBatchEntries + } + if hasDuplicateBatchEntryIDsByID(request.Entries, func(entry contracts.DeleteMessageBatchRequestEntry) string { return entry.Id }) { + return contracts.DeleteMessageBatchResult{}, errDuplicateBatchEntryIDs + } + + s.mu.Lock() + defer s.mu.Unlock() + + if _, ok := s.activeQueueByNameLocked(queueName); !ok { + return contracts.DeleteMessageBatchResult{}, errQueueNotFound + } + + result := contracts.DeleteMessageBatchResult{ + Successful: make([]contracts.DeleteMessageBatchResultEntry, 0, len(request.Entries)), + Failed: make([]contracts.BatchResultErrorEntry, 0), + } + deleted := false + for _, entry := range request.Entries { + id := trimName(entry.Id) + receiptHandle := trimName(entry.ReceiptHandle) + if id == "" || receiptHandle == "" { + result.Failed = append(result.Failed, batchFailureEntry(entry.Id, "receipt handle is required", true)) + continue + } + + idx, ok := s.findMessageByReceiptLocked(queueName, receiptHandle) + if !ok { + result.Failed = append(result.Failed, batchFailureEntry(entry.Id, "receipt handle is invalid", true)) + continue + } + + s.state.Messages = append(s.state.Messages[:idx], s.state.Messages[idx+1:]...) + deleted = true + result.Successful = append(result.Successful, contracts.DeleteMessageBatchResultEntry{Id: id}) + } + + if deleted { + if err := s.commitStateLocked(); err != nil { + return contracts.DeleteMessageBatchResult{}, err + } + } + + return result, nil +} + +func (s *Service) ChangeMessageVisibilityBatch(queueName string, request contracts.ChangeMessageVisibilityBatchRequest) (contracts.ChangeMessageVisibilityBatchResult, error) { + queueName = trimName(queueName) + if queueName == "" { + return contracts.ChangeMessageVisibilityBatchResult{}, fmt.Errorf("sqs: queue name is required") + } + if len(request.Entries) == 0 { + return contracts.ChangeMessageVisibilityBatchResult{}, errEmptyBatchRequest + } + if len(request.Entries) > 10 { + return contracts.ChangeMessageVisibilityBatchResult{}, errTooManyBatchEntries + } + if hasDuplicateBatchEntryIDsByID(request.Entries, func(entry contracts.ChangeMessageVisibilityBatchRequestEntry) string { return entry.Id }) { + return contracts.ChangeMessageVisibilityBatchResult{}, errDuplicateBatchEntryIDs + } + + s.mu.Lock() + defer s.mu.Unlock() + + if _, ok := s.activeQueueByNameLocked(queueName); !ok { + return contracts.ChangeMessageVisibilityBatchResult{}, errQueueNotFound + } + + result := contracts.ChangeMessageVisibilityBatchResult{ + Successful: make([]contracts.ChangeMessageVisibilityBatchResultEntry, 0, len(request.Entries)), + Failed: make([]contracts.BatchResultErrorEntry, 0), + } + changed := false + for _, entry := range request.Entries { + id := trimName(entry.Id) + receiptHandle := trimName(entry.ReceiptHandle) + if id == "" || receiptHandle == "" { + result.Failed = append(result.Failed, batchFailureEntry(entry.Id, "receipt handle is required", true)) + continue + } + + if entry.VisibilityTimeout < 0 { + result.Failed = append(result.Failed, batchFailureEntry(entry.Id, "visibility timeout must be non-negative", true)) + continue + } + + if err := s.changeMessageVisibilityLocked(queueName, receiptHandle, time.Duration(entry.VisibilityTimeout)*time.Second); err != nil { + result.Failed = append(result.Failed, batchFailureEntry(entry.Id, err.Error(), true)) + continue + } + + changed = true + result.Successful = append(result.Successful, contracts.ChangeMessageVisibilityBatchResultEntry{Id: id}) + } + + if changed { + if err := s.commitStateLocked(); err != nil { + return contracts.ChangeMessageVisibilityBatchResult{}, err + } + } + + return result, nil +} + +func (s *Service) commitStateLocked() error { + if s.repo != nil { + if err := s.repo.Save(s.state.Clone()); err != nil { + return fmt.Errorf("sqs: save repository: %w", err) + } + } + s.publishSnapshotLocked() + return nil +} + +func (s *Service) queueResponse(queueName string, queue domain.Queue, attributes map[string]string) domain.Queue { + queue.Name = trimName(queue.Name) + if queue.Name == "" { + queue.Name = trimName(queueName) + } + queue.URL = queueURLForAccount(queue.Name, "") + if len(attributes) == 0 && queue.Attributes == nil { + queue.Attributes = map[string]string{} + return queue + } + queue.Attributes = cloneMap(attributes) + if queue.Attributes == nil { + queue.Attributes = map[string]string{} + } + return queue +} + +func (s *Service) queueRecordByNameLocked(name string) (int, domain.Queue, bool) { + for idx, queue := range s.state.Queues { + if queue.Name == name { + return idx, queue, true + } + } + return -1, domain.Queue{}, false +} + +func (s *Service) queueByARNLocked(queueARN string) (domain.Queue, bool) { + queueARN = trimName(queueARN) + for _, queue := range s.state.Queues { + if queue.DeletedAt.IsZero() && queueARNForAccount(queue.Name, "") == queueARN { + return queue, true + } + } + return domain.Queue{}, false +} + +func (s *Service) findMessageMoveTaskLocked(taskHandle string) (string, domain.MessageMoveTask, bool) { + for queueName, tasks := range s.state.MoveTasks { + if task, ok := tasks[taskHandle]; ok { + return queueName, task, true } - if sleep <= 0 { - return messages, nil + } + return "", domain.MessageMoveTask{}, false +} + +func (s *Service) queueByNameLocked(name string) (domain.Queue, bool) { + queue, ok := s.activeQueueByNameLocked(name) + return queue, ok +} + +func (s *Service) activeQueueByNameLocked(name string) (domain.Queue, bool) { + for _, queue := range s.state.Queues { + if queue.Name == name && queue.DeletedAt.IsZero() { + return queue, true } - s.clock.Sleep(sleep) } + return domain.Queue{}, false } -func (s *Service) DeleteMessage(queueName string, receiptHandle string) error { - queueName = trimName(queueName) - receiptHandle = trimName(receiptHandle) - if queueName == "" { - return fmt.Errorf("sqs: queue name is required") +func (s *Service) removeMessagesForQueueLocked(queueName string) { + filtered := s.state.Messages[:0] + for _, message := range s.state.Messages { + if trimName(message.Queue) == queueName { + continue + } + filtered = append(filtered, message) } - if receiptHandle == "" { - return fmt.Errorf("sqs: receipt handle is required") + s.state.Messages = filtered + + if len(s.state.RecoveryMetadata) == 0 { + return } + for key, metadata := range s.state.RecoveryMetadata { + if metadata.Queue == queueName || strings.HasPrefix(key, queueName+"/") { + delete(s.state.RecoveryMetadata, key) + } + } +} - s.mu.Lock() - defer s.mu.Unlock() +func ownerAccountMatches(ownerAccountID string) bool { + ownerAccountID = trimName(ownerAccountID) + if ownerAccountID == "" { + return true + } + return ownerAccountID == awscontext.Default().AccountID +} - idx, ok := s.findMessageByReceiptLocked(queueName, receiptHandle) - if !ok { - return errReceiptHandleMismatch +func selectQueueAttributes(attributes map[string]string, attributeNames []string, queueARN string) map[string]string { + if len(attributes) == 0 && len(attributeNames) == 0 { + return map[string]string{"QueueArn": queueARN} } - s.state.Messages = append(s.state.Messages[:idx], s.state.Messages[idx+1:]...) - return s.commitStateLocked() + selected := map[string]string{} + includeAll := len(attributeNames) == 0 + for _, name := range attributeNames { + if strings.EqualFold(trimName(name), "All") { + includeAll = true + break + } + } + if includeAll { + for key, value := range attributes { + selected[key] = value + } + } else { + allowed := map[string]struct{}{} + for _, name := range attributeNames { + allowed[trimName(name)] = struct{}{} + } + for key, value := range attributes { + if _, ok := allowed[key]; ok { + selected[key] = value + } + } + } + selected["QueueArn"] = queueARN + return selected } -func (s *Service) ChangeMessageVisibility(queueName string, receiptHandle string, visibility time.Duration) error { - queueName = trimName(queueName) - receiptHandle = trimName(receiptHandle) - if queueName == "" { - return fmt.Errorf("sqs: queue name is required") +func orderingHintFromAttributes(attributes map[string]string, fallback string) string { + if strings.EqualFold(trimName(attributes["FifoQueue"]), "true") { + return "fifo" } - if receiptHandle == "" { - return fmt.Errorf("sqs: receipt handle is required") + if strings.EqualFold(fallback, "fifo") { + return "fifo" } - if visibility < 0 { - return errInvalidVisibilityWindow + return "standard" +} + +func equalStringMaps(left, right map[string]string) bool { + if len(left) != len(right) { + if len(left) == 0 && len(right) == 0 { + return true + } + return false } - if visibility > 12*time.Hour { - visibility = 12 * time.Hour + for key, leftValue := range left { + if right[key] != leftValue { + return false + } } + return true +} - s.mu.Lock() - defer s.mu.Unlock() - - idx, ok := s.findMessageByReceiptLocked(queueName, receiptHandle) - if !ok { - return errReceiptHandleMismatch +func uniqueSortedTrimmedStrings(values []string) []string { + if len(values) == 0 { + return []string{} } - message := &s.state.Messages[idx] - if message.Metadata == nil { - message.Metadata = map[string]string{} - } - if visibility == 0 { - message.ReceivedAt = time.Time{} - message.AvailableAt = s.clock.Now() - message.Metadata[leaseVisibilityTimeoutMetaKey] = "0" - } else { - message.ReceivedAt = s.clock.Now() - message.Metadata[leaseVisibilityTimeoutMetaKey] = strconv.FormatInt(int64(visibility/time.Second), 10) + seen := make(map[string]struct{}, len(values)) + for _, value := range values { + value = trimName(value) + if value == "" { + continue + } + seen[value] = struct{}{} } - return s.commitStateLocked() + ordered := make([]string, 0, len(seen)) + for value := range seen { + ordered = append(ordered, value) + } + sort.Strings(ordered) + return ordered } -func (s *Service) commitStateLocked() error { - if s.repo != nil { - if err := s.repo.Save(s.state.Clone()); err != nil { - return fmt.Errorf("sqs: save repository: %w", err) +func cloneQueuePermissionMap(values map[string]domain.QueuePermission) map[string]domain.QueuePermission { + if values == nil { + return map[string]domain.QueuePermission{} + } + + cloned := make(map[string]domain.QueuePermission, len(values)) + for label, permission := range values { + cloned[label] = domain.QueuePermission{ + Label: permission.Label, + AWSAccountIDs: append([]string(nil), permission.AWSAccountIDs...), + Actions: append([]string(nil), permission.Actions...), + CreatedAt: permission.CreatedAt, + UpdatedAt: permission.UpdatedAt, } } - s.publishSnapshotLocked() - return nil + return cloned } -func (s *Service) queueByNameLocked(name string) (domain.Queue, bool) { - for _, queue := range s.state.Queues { - if queue.Name == name { - return queue, true +func cloneMessageMoveTaskMap(values map[string]domain.MessageMoveTask) map[string]domain.MessageMoveTask { + if values == nil { + return map[string]domain.MessageMoveTask{} + } + + cloned := make(map[string]domain.MessageMoveTask, len(values)) + for handle, task := range values { + cloned[handle] = domain.MessageMoveTask{ + TaskHandle: task.TaskHandle, + SourceQueue: task.SourceQueue, + SourceArn: task.SourceArn, + DestinationArn: task.DestinationArn, + MaxNumberOfMessagesPerSecond: task.MaxNumberOfMessagesPerSecond, + ApproximateNumberOfMessagesMoved: task.ApproximateNumberOfMessagesMoved, + Status: task.Status, + StartedAt: task.StartedAt, + UpdatedAt: task.UpdatedAt, + CancelledAt: task.CancelledAt, } } - return domain.Queue{}, false + return cloned } func (s *Service) receiveReadyMessagesLocked(queueName string, maxMessages int, now time.Time) ([]domain.Message, error) { s.mu.Lock() defer s.mu.Unlock() + s.sweepDeadLettersLocked(now) + queue, ok := s.queueByNameLocked(queueName) if !ok { return nil, errQueueNotFound @@ -328,6 +1293,9 @@ func (s *Service) receiveReadyMessagesLocked(queueName string, maxMessages int, if message.Metadata == nil { message.Metadata = map[string]string{} } + if trimName(message.Metadata["approximate_first_receive_timestamp"]) == "" { + message.Metadata["approximate_first_receive_timestamp"] = strconv.FormatInt(now.UnixMilli(), 10) + } timeout := queueVisibilityTimeout(queue, *message) message.ReceivedAt = now message.Recovery.Attempts++ @@ -349,7 +1317,243 @@ func (s *Service) receiveReadyMessagesLocked(queueName string, maxMessages int, return selected, nil } +func (s *Service) enqueueMessageLocked(queueName string, queue domain.Queue, request contracts.SendMessageRequest, batchID string, batchEntryID string, batchEntryIndex int, batchEntryCount int) (domain.Message, contracts.SendMessageResult, error) { + messageGroupID := trimName(request.MessageGroupId) + if IsFIFOQueue(queue) && messageGroupID == "" { + return domain.Message{}, contracts.SendMessageResult{}, fmt.Errorf("sqs: message group id is required for fifo queues") + } + + now := s.clock.Now() + message := domain.Message{ + Queue: queueName, + MessageID: uuid.NewString(), + Body: request.MessageBody, + Attributes: messageAttributesToStrings(request.MessageAttributes), + Metadata: messageSystemAttributesToStrings(request.MessageSystemAttributes), + MessageGroupID: messageGroupID, + BatchID: trimName(batchID), + BatchEntryID: trimName(batchEntryID), + BatchEntryIndex: batchEntryIndex, + BatchEntryCount: batchEntryCount, + SentAt: now, + } + + effectiveDelaySeconds := request.DelaySeconds + if effectiveDelaySeconds <= 0 { + effectiveDelaySeconds = parseDelaySeconds(queue.Attributes["DelaySeconds"]) + } + if effectiveDelaySeconds > 0 { + message.AvailableAt = now.Add(time.Duration(effectiveDelaySeconds) * time.Second) + } + if message.Metadata == nil { + message.Metadata = map[string]string{} + } + if request.MessageDeduplicationId != "" { + message.Metadata["MessageDeduplicationId"] = request.MessageDeduplicationId + } + + if IsFIFOQueue(queue) { + message.SequenceNumber = s.nextSequenceNumberLocked(queueName, messageGroupID) + } + + result := contracts.SendMessageResult{ + MD5OfMessageBody: md5OfString(request.MessageBody), + MessageId: message.MessageID, + } + if message.SequenceNumber > 0 { + result.SequenceNumber = strconv.FormatInt(message.SequenceNumber, 10) + } + if len(request.MessageAttributes) > 0 { + result.MD5OfMessageAttributes = md5OfMap(message.Attributes) + } + if len(request.MessageSystemAttributes) > 0 { + result.MD5OfMessageSystemAttributes = md5OfMap(message.Metadata) + } + + return message, result, nil +} + +func (s *Service) enqueueBatchMessageLocked(queueName string, queue domain.Queue, entry contracts.SendMessageBatchRequestEntry, batchIndex int, batchCount int) (contracts.SendMessageBatchResultEntry, domain.Message, error) { + if trimName(entry.Id) == "" { + return contracts.SendMessageBatchResultEntry{}, domain.Message{}, fmt.Errorf("sqs: batch entry id is required") + } + if trimName(entry.MessageBody) == "" { + return contracts.SendMessageBatchResultEntry{}, domain.Message{}, fmt.Errorf("sqs: message body is required") + } + + message, result, err := s.enqueueMessageLocked(queueName, queue, contracts.SendMessageRequest{ + DelaySeconds: entry.DelaySeconds, + MessageAttributes: entry.MessageAttributes, + MessageBody: entry.MessageBody, + MessageDeduplicationId: entry.MessageDeduplicationId, + MessageGroupId: entry.MessageGroupId, + MessageSystemAttributes: entry.MessageSystemAttributes, + }, "", entry.Id, batchIndex, batchCount) + if err != nil { + return contracts.SendMessageBatchResultEntry{}, domain.Message{}, err + } + + return contracts.SendMessageBatchResultEntry{ + Id: trimName(entry.Id), + MD5OfMessageAttributes: result.MD5OfMessageAttributes, + MD5OfMessageBody: result.MD5OfMessageBody, + MD5OfMessageSystemAttributes: result.MD5OfMessageSystemAttributes, + MessageId: result.MessageId, + SequenceNumber: result.SequenceNumber, + }, message, nil +} + +func (s *Service) changeMessageVisibilityLocked(queueName, receiptHandle string, visibility time.Duration) error { + if visibility < 0 { + return errInvalidVisibilityWindow + } + idx, ok := s.findMessageByReceiptLocked(queueName, receiptHandle) + if !ok { + return errReceiptHandleMismatch + } + + message := &s.state.Messages[idx] + if message.Metadata == nil { + message.Metadata = map[string]string{} + } + if visibility == 0 { + message.ReceivedAt = time.Time{} + message.AvailableAt = s.clock.Now() + message.Metadata[leaseVisibilityTimeoutMetaKey] = "0" + return nil + } + + message.ReceivedAt = s.clock.Now() + message.Metadata[leaseVisibilityTimeoutMetaKey] = strconv.FormatInt(int64(visibility/time.Second), 10) + return nil +} + +func (s *Service) nextSequenceNumberLocked(queueName, groupID string) int64 { + var maxSequence int64 + for _, message := range s.state.Messages { + if trimName(message.Queue) != queueName { + continue + } + if trimName(groupID) != "" && trimName(message.MessageGroupID) != trimName(groupID) { + continue + } + if message.SequenceNumber > maxSequence { + maxSequence = message.SequenceNumber + } + } + return maxSequence + 1 +} + +func md5OfString(value string) string { + sum := md5.Sum([]byte(value)) + return hex.EncodeToString(sum[:]) +} + +func md5OfMap(values map[string]string) string { + if len(values) == 0 { + return "" + } + + keys := make([]string, 0, len(values)) + for key := range values { + keys = append(keys, key) + } + sort.Strings(keys) + + builder := strings.Builder{} + for _, key := range keys { + builder.WriteString(key) + builder.WriteString("=") + builder.WriteString(values[key]) + builder.WriteString(";") + } + return md5OfString(builder.String()) +} + +func messageAttributesToStrings(attributes map[string]contracts.MessageAttributeValue) map[string]string { + if len(attributes) == 0 { + return nil + } + + keys := make([]string, 0, len(attributes)) + for key := range attributes { + keys = append(keys, key) + } + sort.Strings(keys) + + values := make(map[string]string, len(keys)) + for _, key := range keys { + value := attributes[key] + switch { + case value.StringValue != "": + values[key] = value.StringValue + case len(value.BinaryValue) > 0: + values[key] = string(value.BinaryValue) + case value.DataType != "": + values[key] = value.DataType + default: + values[key] = "" + } + } + return values +} + +func messageSystemAttributesToStrings(attributes map[string]contracts.MessageAttributeValue) map[string]string { + if len(attributes) == 0 { + return nil + } + + keys := make([]string, 0, len(attributes)) + for key := range attributes { + keys = append(keys, key) + } + sort.Strings(keys) + + values := make(map[string]string, len(keys)) + for _, key := range keys { + value := attributes[key] + switch { + case value.StringValue != "": + values[key] = value.StringValue + case len(value.BinaryValue) > 0: + values[key] = string(value.BinaryValue) + case value.DataType != "": + values[key] = value.DataType + default: + values[key] = "" + } + } + return values +} + +func hasDuplicateBatchEntryIDsByID[T any](entries []T, getID func(T) string) bool { + seen := map[string]struct{}{} + for _, entry := range entries { + id := trimName(getID(entry)) + if id == "" { + continue + } + if _, ok := seen[id]; ok { + return true + } + seen[id] = struct{}{} + } + return false +} + +func batchFailureEntry(id, message string, senderFault bool) contracts.BatchResultErrorEntry { + return contracts.BatchResultErrorEntry{ + Code: "InvalidParameterValue", + Id: trimName(id), + Message: message, + SenderFault: senderFault, + } +} + func (s *Service) findMessageByReceiptLocked(queueName, receiptHandle string) (int, bool) { + if _, ok := s.activeQueueByNameLocked(queueName); !ok { + return -1, false + } for idx, message := range s.state.Messages { if !messageVisibleInQueue(message, queueName) { continue @@ -417,6 +1621,18 @@ func parseMessageVisibilityTimeout(raw string) time.Duration { return time.Duration(seconds) * time.Second } +func parseDelaySeconds(raw string) int { + raw = trimName(raw) + if raw == "" { + return 0 + } + seconds, err := strconv.Atoi(raw) + if err != nil || seconds < 0 { + return 0 + } + return seconds +} + func nextReceiptHandle(queueName string, message domain.Message) string { return fmt.Sprintf("%s/%s/%d", queueName, message.MessageID, len(message.ReceiptKeys)+1) } @@ -437,6 +1653,97 @@ func cloneMap(values map[string]string) map[string]string { return cloned } +func queueRecoveryFromAttributes(attributes map[string]string) domain.QueueRecovery { + recovery := domain.QueueRecovery{ + Policy: map[string]string{}, + } + if len(attributes) == 0 { + return recovery + } + + rawPolicy := trimName(attributes["RedrivePolicy"]) + if rawPolicy == "" { + return recovery + } + + var parsed map[string]any + if err := json.Unmarshal([]byte(rawPolicy), &parsed); err != nil { + recovery.Policy["raw"] = rawPolicy + return recovery + } + + for key, value := range parsed { + normalizedKey := camelToSnake(key) + recovery.Policy[normalizedKey] = fmt.Sprint(value) + } + if targetArn := trimName(recovery.Policy["dead_letter_target_arn"]); targetArn != "" { + if queueName, _, err := queueNameAndAccountFromARN(targetArn); err == nil { + recovery.DeadLetterQueue = queueName + } + } + return recovery +} + +func camelToSnake(value string) string { + if value == "" { + return "" + } + + var out strings.Builder + for i, r := range value { + if i > 0 && r >= 'A' && r <= 'Z' { + out.WriteByte('_') + } + out.WriteRune(r) + } + return strings.ToLower(out.String()) +} + +func queueNameAndAccountFromARN(raw string) (string, string, error) { + trimmed := trimName(raw) + if trimmed == "" { + return "", "", fmt.Errorf("sqs: arn is required") + } + + parts := strings.Split(trimmed, ":") + if len(parts) < 6 || parts[0] != "arn" { + return "", "", fmt.Errorf("sqs: invalid arn: %s", raw) + } + + accountID := trimName(parts[len(parts)-2]) + queueName := trimName(parts[len(parts)-1]) + if accountID == "" || queueName == "" { + return "", "", fmt.Errorf("sqs: invalid arn: %s", raw) + } + return queueName, accountID, nil +} + +func queueURLForAccount(queueName, ownerAccountID string) string { + queueName = trimName(queueName) + if queueName == "" { + return "" + } + + aws := awscontext.Default() + if ownerAccountID = trimName(ownerAccountID); ownerAccountID != "" { + aws = aws.WithAccountID(ownerAccountID) + } + return fmt.Sprintf("https://sqs.%s.amazonaws.com/%s/%s", aws.Region, aws.AccountID, queueName) +} + +func queueARNForAccount(queueName, ownerAccountID string) string { + queueName = trimName(queueName) + if queueName == "" { + return "" + } + + aws := awscontext.Default() + if ownerAccountID = trimName(ownerAccountID); ownerAccountID != "" { + aws = aws.WithAccountID(ownerAccountID) + } + return aws.ServiceARN("sqs", queueName) +} + func messageVisibleInQueue(message domain.Message, queueName string) bool { if message.Queue != queueName { return false diff --git a/core/internal/resources/sqs/application/service_test.go b/core/internal/resources/sqs/application/service_test.go index 309e247..dace812 100644 --- a/core/internal/resources/sqs/application/service_test.go +++ b/core/internal/resources/sqs/application/service_test.go @@ -2,12 +2,15 @@ package application import ( "context" + "crypto/md5" + "encoding/hex" "testing" "time" "github.com/michasdev/mildstack/core/internal/application/orchestrator" "github.com/michasdev/mildstack/core/internal/application/runtime" deliveryhttp "github.com/michasdev/mildstack/core/internal/delivery/http" + "github.com/michasdev/mildstack/core/internal/resources/sqs/contracts" "github.com/michasdev/mildstack/core/internal/resources/sqs/domain" ) @@ -81,6 +84,485 @@ func TestSQSServiceMetadataRoutesAndPolicy(t *testing.T) { assertRouteExists(t, entry.Routes, "DELETE", "/api/v1/runtime/services/sqs/queues/:queue/messages/:receiptHandle") } +func TestSQSServiceExposesQueueLifecycleAPI(t *testing.T) { + t.Helper() + + clock := newManualClock(time.Date(2026, time.April, 19, 12, 0, 0, 0, time.UTC)) + service := newServiceWithClock(domain.NewState(), nil, clock) + type lifecycleAPI interface { + QueueURL(string) string + QueueARN(string) string + CreateQueue(string, map[string]string) (domain.Queue, error) + DeleteQueue(string) error + GetQueueUrl(string, string) (string, error) + ListQueues(string, int, string, string) ([]domain.Queue, string, error) + PurgeQueue(string) error + GetQueueAttributes(string, []string, string) (contracts.QueueAttributesView, error) + SetQueueAttributes(string, map[string]string) (contracts.QueueAttributesView, error) + } + + if _, ok := any(service).(lifecycleAPI); !ok { + t.Fatal("expected service to expose queue lifecycle API") + } + + if got, want := service.QueueURL("orders"), "https://sqs.us-east-1.amazonaws.com/123456789012/orders"; got != want { + t.Fatalf("unexpected queue url helper: got %q want %q", got, want) + } + if got, want := service.QueueARN("orders"), "arn:aws:sqs:us-east-1:123456789012:orders"; got != want { + t.Fatalf("unexpected queue arn helper: got %q want %q", got, want) + } + + queue, err := service.CreateQueue("orders", map[string]string{ + "VisibilityTimeout": "30", + "RedrivePolicy": `{"deadLetterTargetArn":"arn:aws:sqs:us-east-1:123456789012:orders-dlq"}`, + }) + if err != nil { + t.Fatalf("create queue: %v", err) + } + if got, want := queue.URL, service.QueueURL("orders"); got != want { + t.Fatalf("unexpected queue url: got %q want %q", got, want) + } + if got, want := queue.Attributes["VisibilityTimeout"], "30"; got != want { + t.Fatalf("unexpected queue attribute: got %q want %q", got, want) + } + + sameQueue, err := service.CreateQueue("orders", map[string]string{ + "VisibilityTimeout": "30", + "RedrivePolicy": `{"deadLetterTargetArn":"arn:aws:sqs:us-east-1:123456789012:orders-dlq"}`, + }) + if err != nil { + t.Fatalf("idempotent create: %v", err) + } + if got, want := sameQueue.URL, queue.URL; got != want { + t.Fatalf("unexpected idempotent queue url: got %q want %q", got, want) + } + if _, err := service.CreateQueue("orders", map[string]string{"VisibilityTimeout": "45"}); err == nil { + t.Fatal("expected create with different attributes to fail") + } + + archiveQueue, err := service.CreateQueue("orders-archive", map[string]string{"VisibilityTimeout": "45"}) + if err != nil { + t.Fatalf("create archive queue: %v", err) + } + if got, want := archiveQueue.URL, service.QueueURL("orders-archive"); got != want { + t.Fatalf("unexpected archive queue url: got %q want %q", got, want) + } + + list, nextToken, err := service.ListQueues("ord", 1, "", "") + if err != nil { + t.Fatalf("list queues: %v", err) + } + if got, want := len(list), 1; got != want { + t.Fatalf("unexpected paged queue count: got %d want %d", got, want) + } + if got, want := list[0].Name, "orders"; got != want { + t.Fatalf("unexpected first page queue: got %q want %q", got, want) + } + if got, want := nextToken, "orders"; got != want { + t.Fatalf("unexpected next token: got %q want %q", got, want) + } + + nextPage, nextToken, err := service.ListQueues("ord", 10, nextToken, "") + if err != nil { + t.Fatalf("second page list queues: %v", err) + } + if got, want := len(nextPage), 1; got != want { + t.Fatalf("unexpected second page queue count: got %d want %d", got, want) + } + if got, want := nextPage[0].Name, "orders-archive"; got != want { + t.Fatalf("unexpected second page queue: got %q want %q", got, want) + } + if got, want := nextToken, ""; got != want { + t.Fatalf("unexpected terminal next token: got %q want %q", got, want) + } + + queueURL, err := service.GetQueueUrl("orders", "") + if err != nil { + t.Fatalf("get queue url: %v", err) + } + if got, want := queueURL, service.QueueURL("orders"); got != want { + t.Fatalf("unexpected get queue url result: got %q want %q", got, want) + } + + if _, err := service.SetQueueAttributes("orders", map[string]string{ + "VisibilityTimeout": "45", + "RedriveAllowPolicy": `{"redrivePermission":"byQueue"}`, + "RedrivePolicy": `{"deadLetterTargetArn":"arn:aws:sqs:us-east-1:123456789012:orders-dlq"}`, + "ContentBasedDeduplication": "true", + }); err != nil { + t.Fatalf("set queue attributes: %v", err) + } + + attrView, err := service.GetQueueAttributes("orders", []string{"All"}, "") + if err != nil { + t.Fatalf("get queue attributes: %v", err) + } + if got, want := attrView.QueueURL, service.QueueURL("orders"); got != want { + t.Fatalf("unexpected queue attribute url: got %q want %q", got, want) + } + if got, want := attrView.QueueARN, service.QueueARN("orders"); got != want { + t.Fatalf("unexpected queue attribute arn: got %q want %q", got, want) + } + if got, want := attrView.Attributes["VisibilityTimeout"], "45"; got != want { + t.Fatalf("unexpected queue attribute value: got %q want %q", got, want) + } + if got, want := attrView.Attributes["RedriveAllowPolicy"], `{"redrivePermission":"byQueue"}`; got != want { + t.Fatalf("unexpected opaque attribute value: got %q want %q", got, want) + } + if got, want := attrView.Attributes["QueueArn"], service.QueueARN("orders"); got != want { + t.Fatalf("unexpected queue arn attribute: got %q want %q", got, want) + } + + if err := service.DeleteQueue("orders"); err != nil { + t.Fatalf("delete queue: %v", err) + } + if _, err := service.GetQueueUrl("orders", ""); err == nil { + t.Fatal("expected deleted queue url lookup to fail") + } + if _, _, err := service.ListQueues("ord", 10, "", ""); err != nil { + t.Fatalf("list queues after delete: %v", err) + } + if _, err := service.CreateQueue("orders", map[string]string{"VisibilityTimeout": "30"}); err == nil { + t.Fatal("expected recreate during delete cooldown to fail") + } + + clock.Sleep(queueLifecycleCooldown + time.Second) + recreated, err := service.CreateQueue("orders", map[string]string{"VisibilityTimeout": "30"}) + if err != nil { + t.Fatalf("recreate after cooldown: %v", err) + } + if got, want := recreated.URL, service.QueueURL("orders"); got != want { + t.Fatalf("unexpected recreated queue url: got %q want %q", got, want) + } + + service.state.Messages = append(service.state.Messages, domain.Message{ + Queue: "orders-archive", + MessageID: "message-1", + Body: "payload", + }) + if err := service.PurgeQueue("orders-archive"); err != nil { + t.Fatalf("purge queue: %v", err) + } + if got, want := len(service.state.Messages), 0; got != want { + t.Fatalf("expected purge to delete queue messages, got %d", got) + } + if err := service.PurgeQueue("orders-archive"); err == nil { + t.Fatal("expected back-to-back purge to fail") + } +} + +func TestSQSServiceExposesMessageSurfaceSeams(t *testing.T) { + t.Helper() + + service := newService(domain.NewState(), nil) + type messageAPI interface { + ReceiveMessage(string, int, time.Duration) ([]domain.Message, error) + DeleteMessage(string, string) error + ChangeMessageVisibility(string, string, time.Duration) error + SendMessage(string, contracts.SendMessageRequest) (contracts.SendMessageResult, error) + SendMessageBatch(string, contracts.SendMessageBatchRequest) (contracts.SendMessageBatchResult, error) + DeleteMessageBatch(string, contracts.DeleteMessageBatchRequest) (contracts.DeleteMessageBatchResult, error) + ChangeMessageVisibilityBatch(string, contracts.ChangeMessageVisibilityBatchRequest) (contracts.ChangeMessageVisibilityBatchResult, error) + } + + if _, ok := any(service).(messageAPI); !ok { + t.Fatal("expected service to expose the message surface API") + } + + if _, err := service.CreateQueue("queue-a", nil); err != nil { + t.Fatalf("create queue: %v", err) + } + + sendResult, err := service.SendMessage("queue-a", contracts.SendMessageRequest{ + MessageBody: "payload", + QueueUrl: service.QueueURL("queue-a"), + }) + if err != nil { + t.Fatalf("send message: %v", err) + } + if sendResult.MessageId == "" { + t.Fatal("expected send message to return a message id") + } + expectedBodyDigest := md5.Sum([]byte("payload")) + if got, want := sendResult.MD5OfMessageBody, hex.EncodeToString(expectedBodyDigest[:]); got != want { + t.Fatalf("unexpected message digest: got %q want %q", got, want) + } + if got, want := len(service.state.Messages), 1; got != want { + t.Fatalf("unexpected stored message count: got %d want %d", got, want) + } + if got, want := service.state.Messages[0].Body, "payload"; got != want { + t.Fatalf("unexpected stored message body: got %q want %q", got, want) + } + + batchResult, err := service.SendMessageBatch("queue-a", contracts.SendMessageBatchRequest{ + QueueUrl: service.QueueURL("queue-a"), + Entries: []contracts.SendMessageBatchRequestEntry{ + {Id: "entry-1", MessageBody: "one"}, + {Id: "entry-2", MessageBody: "two"}, + {Id: "entry-3", MessageBody: ""}, + }, + }) + if err != nil { + t.Fatalf("send message batch: %v", err) + } + if got, want := len(batchResult.Successful), 2; got != want { + t.Fatalf("unexpected successful batch count: got %d want %d", got, want) + } + if got, want := len(batchResult.Failed), 1; got != want { + t.Fatalf("unexpected failed batch count: got %d want %d", got, want) + } + if got, want := batchResult.Successful[0].Id, "entry-1"; got != want { + t.Fatalf("unexpected first batch id: got %q want %q", got, want) + } + if got, want := batchResult.Failed[0].Id, "entry-3"; got != want { + t.Fatalf("unexpected failed batch id: got %q want %q", got, want) + } + if got, want := len(service.state.Messages), 3; got != want { + t.Fatalf("unexpected stored message count after batch send: got %d want %d", got, want) + } +} + +func TestSQSServiceSendMessageAppliesQueueDelayWhenMessageDelayMissing(t *testing.T) { + t.Helper() + + now := time.Date(2026, time.April, 19, 12, 0, 0, 0, time.UTC) + clock := newManualClock(now) + service := newServiceWithClock(domain.State{ + Service: "sqs", + Queues: []domain.Queue{ + { + Name: "queue-a", + Attributes: map[string]string{ + "DelaySeconds": "2", + }, + }, + }, + }, nil, clock) + + if _, err := service.SendMessage("queue-a", contracts.SendMessageRequest{ + MessageBody: "payload", + QueueUrl: service.QueueURL("queue-a"), + }); err != nil { + t.Fatalf("send message: %v", err) + } + + if got, want := service.state.Messages[0].AvailableAt, now.Add(2*time.Second); !got.Equal(want) { + t.Fatalf("unexpected available_at: got %v want %v", got, want) + } +} + +func TestSQSServiceExposesGovernanceAndRedriveSeams(t *testing.T) { + t.Helper() + + service := newService(domain.State{ + Service: "sqs", + Queues: []domain.Queue{ + { + Name: "queue-a", + Recovery: domain.QueueRecovery{ + DeadLetterQueue: "queue-dlq", + }, + }, + { + Name: "queue-dlq", + }, + }, + }, nil) + + type governanceAPI interface { + TagQueue(string, map[string]string) error + UntagQueue(string, []string) error + AddPermission(string, string, []string, []string) error + RemovePermission(string, string) error + ListQueueTags(string) (map[string]string, error) + ListDeadLetterSourceQueues(string) ([]string, error) + StartMessageMoveTask(string, string, int) (string, error) + CancelMessageMoveTask(string) (int64, error) + ListMessageMoveTasks(string) ([]domain.MessageMoveTask, error) + } + + if _, ok := any(service).(governanceAPI); !ok { + t.Fatal("expected service to expose the governance and redrive API") + } + + if got := service.state.QueueTags; len(got) != 0 { + t.Fatalf("expected empty queue tag map at startup, got %d entries", len(got)) + } + if got := service.state.QueuePermissions; len(got) != 0 { + t.Fatalf("expected empty permission map at startup, got %d entries", len(got)) + } + if got := service.state.MoveTasks; len(got) != 0 { + t.Fatalf("expected empty move-task map at startup, got %d entries", len(got)) + } + + tags, err := service.ListQueueTags("queue-a") + if err != nil { + t.Fatalf("list queue tags: %v", err) + } + if len(tags) != 0 { + t.Fatalf("expected no tags for fresh queue, got %d", len(tags)) + } + + sources, err := service.ListDeadLetterSourceQueues("queue-dlq") + if err != nil { + t.Fatalf("list dead-letter source queues: %v", err) + } + if got, want := len(sources), 1; got != want { + t.Fatalf("unexpected source queue count: got %d want %d", got, want) + } + if got, want := sources[0], "queue-a"; got != want { + t.Fatalf("unexpected source queue: got %q want %q", got, want) + } + + if err := service.TagQueue("queue-a", map[string]string{"env": "dev"}); err != nil { + t.Fatalf("tag queue: %v", err) + } + if err := service.AddPermission("queue-a", "label-a", []string{"123456789012"}, []string{"SendMessage"}); err != nil { + t.Fatalf("add permission: %v", err) + } + + tags, err = service.ListQueueTags("queue-a") + if err != nil { + t.Fatalf("list queue tags after tag: %v", err) + } + if got, want := tags["env"], "dev"; got != want { + t.Fatalf("unexpected queue tag value: got %q want %q", got, want) + } + if got, want := service.state.QueuePermissions["queue-a"]["label-a"].AWSAccountIDs[0], "123456789012"; got != want { + t.Fatalf("unexpected permission account after add: got %q want %q", got, want) + } + + handle, err := service.StartMessageMoveTask("arn:aws:sqs:us-east-1:123456789012:queue-dlq", "", 10) + if err != nil { + t.Fatalf("start message move task: %v", err) + } + tasks, err := service.ListMessageMoveTasks("queue-dlq") + if err != nil { + t.Fatalf("list message move tasks: %v", err) + } + if got, want := len(tasks), 1; got != want { + t.Fatalf("unexpected move task count: got %d want %d", got, want) + } + if got, want := tasks[0].TaskHandle, handle; got != want { + t.Fatalf("unexpected move task handle: got %q want %q", got, want) + } + if got, want := tasks[0].DestinationArn, ""; got != want { + t.Fatalf("unexpected destination arn: got %q want %q", got, want) + } + if got, want := tasks[0].Status, "RUNNING"; got != want { + t.Fatalf("unexpected move task status: got %q want %q", got, want) + } + moved, err := service.CancelMessageMoveTask(handle) + if err != nil { + t.Fatalf("cancel message move task: %v", err) + } + if got, want := moved, int64(0); got != want { + t.Fatalf("unexpected moved count after cancel: got %d want %d", got, want) + } + tasks, err = service.ListMessageMoveTasks("queue-dlq") + if err != nil { + t.Fatalf("list message move tasks after cancel: %v", err) + } + if got, want := tasks[0].Status, "CANCELLED"; got != want { + t.Fatalf("unexpected move task status after cancel: got %q want %q", got, want) + } + + if err := service.UntagQueue("queue-a", []string{"env"}); err != nil { + t.Fatalf("untag queue: %v", err) + } + tags, err = service.ListQueueTags("queue-a") + if err != nil { + t.Fatalf("list queue tags after untag: %v", err) + } + if got, want := len(tags), 0; got != want { + t.Fatalf("expected tag map to be empty after untag, got %d", got) + } + if err := service.RemovePermission("queue-a", "label-a"); err != nil { + t.Fatalf("remove permission: %v", err) + } + if got, want := len(service.state.QueuePermissions["queue-a"]), 0; got != want { + t.Fatalf("expected permission map to be empty after remove, got %d", got) + } +} + +func TestSQSServiceBatchMessageHelpersReturnPairedResults(t *testing.T) { + t.Helper() + + clock := newManualClock(time.Date(2026, time.April, 19, 12, 0, 0, 0, time.UTC)) + service := newServiceWithClock(domain.NewState(), nil, clock) + if _, err := service.CreateQueue("queue-a", map[string]string{"VisibilityTimeout": "30"}); err != nil { + t.Fatalf("create queue: %v", err) + } + + sendResult, err := service.SendMessageBatch("queue-a", contracts.SendMessageBatchRequest{ + QueueUrl: service.QueueURL("queue-a"), + Entries: []contracts.SendMessageBatchRequestEntry{ + {Id: "entry-1", MessageBody: "one"}, + {Id: "entry-2", MessageBody: "two"}, + }, + }) + if err != nil { + t.Fatalf("send batch: %v", err) + } + if got, want := len(sendResult.Successful), 2; got != want { + t.Fatalf("unexpected send batch success count: got %d want %d", got, want) + } + + messages, err := service.ReceiveMessage("queue-a", 2, 0) + if err != nil { + t.Fatalf("receive messages: %v", err) + } + if got, want := len(messages), 2; got != want { + t.Fatalf("unexpected receive count: got %d want %d", got, want) + } + + deleteResult, err := service.DeleteMessageBatch("queue-a", contracts.DeleteMessageBatchRequest{ + QueueUrl: service.QueueURL("queue-a"), + Entries: []contracts.DeleteMessageBatchRequestEntry{ + {Id: "delete-1", ReceiptHandle: CurrentReceiptHandle(messages[0])}, + {Id: "delete-2", ReceiptHandle: "missing"}, + }, + }) + if err != nil { + t.Fatalf("delete batch: %v", err) + } + if got, want := len(deleteResult.Successful), 1; got != want { + t.Fatalf("unexpected delete success count: got %d want %d", got, want) + } + if got, want := len(deleteResult.Failed), 1; got != want { + t.Fatalf("unexpected delete failure count: got %d want %d", got, want) + } + if got, want := deleteResult.Successful[0].Id, "delete-1"; got != want { + t.Fatalf("unexpected delete success id: got %q want %q", got, want) + } + if got, want := deleteResult.Failed[0].Id, "delete-2"; got != want { + t.Fatalf("unexpected delete failure id: got %q want %q", got, want) + } + + visibilityResult, err := service.ChangeMessageVisibilityBatch("queue-a", contracts.ChangeMessageVisibilityBatchRequest{ + QueueUrl: service.QueueURL("queue-a"), + Entries: []contracts.ChangeMessageVisibilityBatchRequestEntry{ + {Id: "vis-1", ReceiptHandle: CurrentReceiptHandle(messages[1]), VisibilityTimeout: 120}, + {Id: "vis-2", ReceiptHandle: "missing", VisibilityTimeout: 120}, + }, + }) + if err != nil { + t.Fatalf("change visibility batch: %v", err) + } + if got, want := len(visibilityResult.Successful), 1; got != want { + t.Fatalf("unexpected visibility success count: got %d want %d", got, want) + } + if got, want := len(visibilityResult.Failed), 1; got != want { + t.Fatalf("unexpected visibility failure count: got %d want %d", got, want) + } + if got, want := visibilityResult.Successful[0].Id, "vis-1"; got != want { + t.Fatalf("unexpected visibility success id: got %q want %q", got, want) + } + if got, want := visibilityResult.Failed[0].Id, "vis-2"; got != want { + t.Fatalf("unexpected visibility failure id: got %q want %q", got, want) + } +} + func TestSQSServiceAttachStateUsesNamespacedCopySafeSnapshot(t *testing.T) { t.Helper() @@ -107,6 +589,31 @@ func TestSQSServiceAttachStateUsesNamespacedCopySafeSnapshot(t *testing.T) { ReceiptKeys: []string{"r-1"}, }, }, + QueueTags: map[string]map[string]string{ + "queue-a": map[string]string{"env": "dev"}, + }, + QueuePermissions: map[string]map[string]domain.QueuePermission{ + "queue-a": map[string]domain.QueuePermission{ + "label-a": { + Label: "label-a", + AWSAccountIDs: []string{"123456789012"}, + Actions: []string{"SendMessage"}, + }, + }, + }, + MoveTasks: map[string]map[string]domain.MessageMoveTask{ + "queue-a": map[string]domain.MessageMoveTask{ + "task-1": { + TaskHandle: "task-1", + SourceQueue: "queue-a", + SourceArn: "arn:aws:sqs:us-east-1:123456789012:queue-a", + DestinationArn: "arn:aws:sqs:us-east-1:123456789012:queue-dlq", + MaxNumberOfMessagesPerSecond: 10, + ApproximateNumberOfMessagesMoved: 2, + Status: "RUNNING", + }, + }, + }, }, nil) hook := runtime.NewStateHook() @@ -135,6 +642,7 @@ func TestSQSServiceAttachStateUsesNamespacedCopySafeSnapshot(t *testing.T) { } messages[0].(map[string]any)["body"] = "mutated" messages[0].(map[string]any)["tags"].([]string)[0] = "mutated" + state["queue_tags"].(map[string]any)["queue-a"].(map[string]any)["env"] = "prod" if got, want := service.state.Queues[0].Name, "queue-a"; got != want { t.Fatalf("service queue name was aliased: got %q want %q", got, want) @@ -148,6 +656,9 @@ func TestSQSServiceAttachStateUsesNamespacedCopySafeSnapshot(t *testing.T) { if got, want := service.state.Messages[0].Tags[0], "alpha"; got != want { t.Fatalf("service message tags were aliased: got %q want %q", got, want) } + if got, want := service.state.QueueTags["queue-a"]["env"], "dev"; got != want { + t.Fatalf("service queue tags were aliased: got %q want %q", got, want) + } } func TestSQSServiceReceiveMessageRespectsStandardOrderAndFifoOrdering(t *testing.T) { @@ -268,6 +779,109 @@ func TestSQSServiceDeadLetterEligibilityUsesRecoveryPolicy(t *testing.T) { } } +func TestSQSServiceQueueAttributesPopulateDeadLetterRecovery(t *testing.T) { + t.Helper() + + service := newService(domain.NewState(), nil) + + if _, err := service.CreateQueue("queue-dlq", nil); err != nil { + t.Fatalf("create dlq: %v", err) + } + if _, err := service.CreateQueue("queue-a", map[string]string{ + "RedrivePolicy": `{"deadLetterTargetArn":"arn:aws:sqs:us-east-1:123456789012:queue-dlq","maxReceiveCount":"2"}`, + }); err != nil { + t.Fatalf("create source queue: %v", err) + } + + _, queue, ok := service.queueRecordByNameLocked("queue-a") + if !ok { + t.Fatal("expected source queue to exist") + } + if got, want := queue.Recovery.DeadLetterQueue, "queue-dlq"; got != want { + t.Fatalf("unexpected dead letter queue: got %q want %q", got, want) + } + if got, want := queue.Recovery.Policy["max_receive_count"], "2"; got != want { + t.Fatalf("unexpected max receive count policy: got %q want %q", got, want) + } + + if _, err := service.SetQueueAttributes("queue-a", map[string]string{ + "RedrivePolicy": `{"deadLetterTargetArn":"arn:aws:sqs:us-east-1:123456789012:queue-dlq","maxReceiveCount":"3"}`, + }); err != nil { + t.Fatalf("set queue attributes: %v", err) + } + + _, queue, ok = service.queueRecordByNameLocked("queue-a") + if !ok { + t.Fatal("expected source queue to exist after update") + } + if got, want := queue.Recovery.Policy["max_receive_count"], "3"; got != want { + t.Fatalf("unexpected updated max receive count policy: got %q want %q", got, want) + } + + sources, err := service.ListDeadLetterSourceQueues("queue-dlq") + if err != nil { + t.Fatalf("list dead letter source queues: %v", err) + } + if got, want := len(sources), 1; got != want { + t.Fatalf("unexpected source queue count: got %d want %d", got, want) + } + if got, want := sources[0], "queue-a"; got != want { + t.Fatalf("unexpected source queue name: got %q want %q", got, want) + } +} + +func TestSQSServiceReceiveMessageMovesDeadLetterEligibleMessagesBeforeDelivery(t *testing.T) { + t.Helper() + + now := time.Date(2026, time.April, 19, 12, 0, 0, 0, time.UTC) + clock := newManualClock(now) + service := newServiceWithClock(domain.State{ + Service: "sqs", + Queues: []domain.Queue{ + { + Name: "queue-a", + Recovery: domain.QueueRecovery{ + DeadLetterQueue: "queue-dlq", + Policy: map[string]string{ + "max_receive_count": "2", + }, + }, + }, + {Name: "queue-dlq"}, + }, + Messages: []domain.Message{ + { + Queue: "queue-a", + MessageID: "message-1", + Body: "payload", + ReceivedAt: now.Add(-31 * time.Second), + Recovery: domain.MessageRecovery{ + Attempts: 2, + }, + }, + }, + }, nil, clock) + + messages, err := service.ReceiveMessage("queue-a", 1, 0) + if err != nil { + t.Fatalf("receive from source queue: %v", err) + } + if got, want := len(messages), 0; got != want { + t.Fatalf("unexpected source queue message count: got %d want %d", got, want) + } + + dlqMessages, err := service.ReceiveMessage("queue-dlq", 1, 0) + if err != nil { + t.Fatalf("receive from dlq: %v", err) + } + if got, want := len(dlqMessages), 1; got != want { + t.Fatalf("unexpected dlq message count: got %d want %d", got, want) + } + if got, want := dlqMessages[0].MessageID, "message-1"; got != want { + t.Fatalf("unexpected dlq message id: got %q want %q", got, want) + } +} + func TestSQSServiceNewWithPersistenceLoadsRepositoryState(t *testing.T) { t.Helper() @@ -313,6 +927,25 @@ func TestSQSServiceNewWithPersistenceLoadsRepositoryState(t *testing.T) { ReceivedAt: time.Date(2026, time.April, 19, 12, 4, 0, 0, time.UTC), ReceiptKeys: []string{"r-1", "r-2"}, }) + state.QueueTags["queue-a"] = map[string]string{"env": "dev"} + state.QueuePermissions["queue-a"] = map[string]domain.QueuePermission{ + "label-a": { + Label: "label-a", + AWSAccountIDs: []string{"123456789012"}, + Actions: []string{"SendMessage"}, + }, + } + state.MoveTasks["queue-a"] = map[string]domain.MessageMoveTask{ + "task-1": { + TaskHandle: "task-1", + SourceQueue: "queue-a", + SourceArn: "arn:aws:sqs:us-east-1:123456789012:queue-a", + DestinationArn: "arn:aws:sqs:us-east-1:123456789012:queue-dlq", + MaxNumberOfMessagesPerSecond: 10, + ApproximateNumberOfMessagesMoved: 2, + Status: "RUNNING", + }, + } if err := repo.Save(state); err != nil { _ = repo.Close() t.Fatalf("save seeded state: %v", err) @@ -361,6 +994,15 @@ func TestSQSServiceNewWithPersistenceLoadsRepositoryState(t *testing.T) { if got, want := service.state.Messages[0].DeadLetterQueue, "queue-dlq"; got != want { t.Fatalf("unexpected dead letter queue after load: got %q want %q", got, want) } + if got, want := service.state.QueueTags["queue-a"]["env"], "dev"; got != want { + t.Fatalf("unexpected queue tags after load: got %q want %q", got, want) + } + if got, want := service.state.QueuePermissions["queue-a"]["label-a"].AWSAccountIDs[0], "123456789012"; got != want { + t.Fatalf("unexpected queue permission after load: got %q want %q", got, want) + } + if got, want := service.state.MoveTasks["queue-a"]["task-1"].Status, "RUNNING"; got != want { + t.Fatalf("unexpected move task after load: got %q want %q", got, want) + } } func TestSQSServiceStopClosesRepositoryIdempotently(t *testing.T) { diff --git a/core/internal/resources/sqs/contracts/catalog.go b/core/internal/resources/sqs/contracts/catalog.go index 3c08b9d..e0f19ba 100644 --- a/core/internal/resources/sqs/contracts/catalog.go +++ b/core/internal/resources/sqs/contracts/catalog.go @@ -13,16 +13,17 @@ type ActionSpec struct { Version string ReturnsQueueURL bool UsesQueueContext bool + MessageSurface bool } var catalog = []ActionSpec{ {Action: "AddPermission", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, {Action: "CancelMessageMoveTask", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, - {Action: "ChangeMessageVisibility", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, - {Action: "ChangeMessageVisibilityBatch", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, + {Action: "ChangeMessageVisibility", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true, MessageSurface: true}, + {Action: "ChangeMessageVisibilityBatch", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true, MessageSurface: true}, {Action: "CreateQueue", Scope: ScopeRoot, Version: "2012-11-05", ReturnsQueueURL: true}, - {Action: "DeleteMessage", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, - {Action: "DeleteMessageBatch", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, + {Action: "DeleteMessage", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true, MessageSurface: true}, + {Action: "DeleteMessageBatch", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true, MessageSurface: true}, {Action: "DeleteQueue", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, {Action: "GetQueueAttributes", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, {Action: "GetQueueUrl", Scope: ScopeRoot, Version: "2012-11-05", ReturnsQueueURL: true}, @@ -31,10 +32,10 @@ var catalog = []ActionSpec{ {Action: "ListQueues", Scope: ScopeRoot, Version: "2012-11-05", ReturnsQueueURL: true}, {Action: "ListQueueTags", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, {Action: "PurgeQueue", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, - {Action: "ReceiveMessage", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, + {Action: "ReceiveMessage", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true, MessageSurface: true}, {Action: "RemovePermission", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, - {Action: "SendMessage", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, - {Action: "SendMessageBatch", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, + {Action: "SendMessage", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true, MessageSurface: true}, + {Action: "SendMessageBatch", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true, MessageSurface: true}, {Action: "SetQueueAttributes", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, {Action: "StartMessageMoveTask", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, {Action: "TagQueue", Scope: ScopeQueue, Version: "2012-11-05", UsesQueueContext: true}, diff --git a/core/internal/resources/sqs/contracts/catalog_test.go b/core/internal/resources/sqs/contracts/catalog_test.go index b6b73c0..46190b4 100644 --- a/core/internal/resources/sqs/contracts/catalog_test.go +++ b/core/internal/resources/sqs/contracts/catalog_test.go @@ -93,6 +93,44 @@ func TestCatalogHasTransportScopeMetadata(t *testing.T) { } } +func TestCatalogMarksPhase39MessageSurfaceActions(t *testing.T) { + t.Helper() + + byAction := make(map[string]ActionSpec) + for _, spec := range Catalog() { + byAction[spec.Action] = spec + } + + messageActions := []string{ + "ChangeMessageVisibility", + "ChangeMessageVisibilityBatch", + "DeleteMessage", + "DeleteMessageBatch", + "ReceiveMessage", + "SendMessage", + "SendMessageBatch", + } + for _, action := range messageActions { + spec, ok := byAction[action] + if !ok { + t.Fatalf("expected action %q in catalog", action) + } + if !spec.MessageSurface { + t.Fatalf("expected action %s to be marked as message surface", action) + } + } + + for _, action := range []string{"AddPermission", "CreateQueue", "ListQueues", "TagQueue"} { + spec, ok := byAction[action] + if !ok { + t.Fatalf("expected action %q in catalog", action) + } + if spec.MessageSurface { + t.Fatalf("did not expect action %s to be marked as message surface", action) + } + } +} + func TestCatalogReturnsCopies(t *testing.T) { t.Helper() diff --git a/core/internal/resources/sqs/contracts/errors.go b/core/internal/resources/sqs/contracts/errors.go new file mode 100644 index 0000000..cbf622e --- /dev/null +++ b/core/internal/resources/sqs/contracts/errors.go @@ -0,0 +1,7 @@ +package contracts + +import "errors" + +// ErrSQSOperationDeferred marks queue operations that are routed but not yet +// behaviorally implemented in the current phase. +var ErrSQSOperationDeferred = errors.New("sqs: operation deferred") diff --git a/core/internal/resources/sqs/contracts/message.go b/core/internal/resources/sqs/contracts/message.go new file mode 100644 index 0000000..c93968e --- /dev/null +++ b/core/internal/resources/sqs/contracts/message.go @@ -0,0 +1,163 @@ +package contracts + +// MessageAttributeValue mirrors the AWS message attribute payload shape. +type MessageAttributeValue struct { + BinaryListValues [][]byte `json:"BinaryListValues,omitempty"` + BinaryValue []byte `json:"BinaryValue,omitempty"` + DataType string `json:"DataType,omitempty"` + StringListValues []string `json:"StringListValues,omitempty"` + StringValue string `json:"StringValue,omitempty"` +} + +// SendMessageRequest preserves the AWS field names required by the message +// write path. +type SendMessageRequest struct { + DelaySeconds int `json:"DelaySeconds,omitempty"` + MessageAttributes map[string]MessageAttributeValue `json:"MessageAttributes,omitempty"` + MessageBody string `json:"MessageBody"` + MessageDeduplicationId string `json:"MessageDeduplicationId,omitempty"` + MessageGroupId string `json:"MessageGroupId,omitempty"` + MessageSystemAttributes map[string]MessageAttributeValue `json:"MessageSystemAttributes,omitempty"` + QueueUrl string `json:"QueueUrl"` +} + +// SendMessageResult preserves the AWS response field names for SendMessage. +type SendMessageResult struct { + MD5OfMessageAttributes string `json:"MD5OfMessageAttributes,omitempty"` + MD5OfMessageBody string `json:"MD5OfMessageBody,omitempty"` + MD5OfMessageSystemAttributes string `json:"MD5OfMessageSystemAttributes,omitempty"` + MessageId string `json:"MessageId,omitempty"` + SequenceNumber string `json:"SequenceNumber,omitempty"` +} + +// SendMessageBatchRequestEntry mirrors the AWS batch message entry payload. +type SendMessageBatchRequestEntry struct { + DelaySeconds int `json:"DelaySeconds,omitempty"` + Id string `json:"Id"` + MessageAttributes map[string]MessageAttributeValue `json:"MessageAttributes,omitempty"` + MessageBody string `json:"MessageBody"` + MessageDeduplicationId string `json:"MessageDeduplicationId,omitempty"` + MessageGroupId string `json:"MessageGroupId,omitempty"` + MessageSystemAttributes map[string]MessageAttributeValue `json:"MessageSystemAttributes,omitempty"` +} + +// SendMessageBatchRequest preserves the AWS field names for batched sends. +type SendMessageBatchRequest struct { + Entries []SendMessageBatchRequestEntry `json:"Entries"` + QueueUrl string `json:"QueueUrl"` +} + +// SendMessageBatchResultEntry mirrors the AWS success payload for a batch send. +type SendMessageBatchResultEntry struct { + Id string `json:"Id,omitempty"` + MD5OfMessageAttributes string `json:"MD5OfMessageAttributes,omitempty"` + MD5OfMessageBody string `json:"MD5OfMessageBody,omitempty"` + MD5OfMessageSystemAttributes string `json:"MD5OfMessageSystemAttributes,omitempty"` + MessageId string `json:"MessageId,omitempty"` + SequenceNumber string `json:"SequenceNumber,omitempty"` +} + +// BatchResultErrorEntry preserves the AWS batch failure payload shape. +type BatchResultErrorEntry struct { + Code string `json:"Code,omitempty"` + Id string `json:"Id,omitempty"` + Message string `json:"Message,omitempty"` + SenderFault bool `json:"SenderFault,omitempty"` +} + +// SendMessageBatchResult preserves the AWS batch response field names. +type SendMessageBatchResult struct { + Failed []BatchResultErrorEntry `json:"Failed,omitempty"` + Successful []SendMessageBatchResultEntry `json:"Successful,omitempty"` +} + +// DeleteMessageRequest preserves the AWS delete payload. +type DeleteMessageRequest struct { + QueueUrl string `json:"QueueUrl"` + ReceiptHandle string `json:"ReceiptHandle"` +} + +// DeleteMessageBatchRequestEntry preserves the AWS batch delete entry shape. +type DeleteMessageBatchRequestEntry struct { + Id string `json:"Id"` + ReceiptHandle string `json:"ReceiptHandle"` +} + +// DeleteMessageBatchRequest preserves the AWS batch delete payload. +type DeleteMessageBatchRequest struct { + Entries []DeleteMessageBatchRequestEntry `json:"Entries"` + QueueUrl string `json:"QueueUrl"` +} + +// DeleteMessageBatchResultEntry mirrors the AWS delete batch success entry. +type DeleteMessageBatchResultEntry struct { + Id string `json:"Id,omitempty"` +} + +// DeleteMessageBatchResult preserves the AWS batch delete response shape. +type DeleteMessageBatchResult struct { + Failed []BatchResultErrorEntry `json:"Failed,omitempty"` + Successful []DeleteMessageBatchResultEntry `json:"Successful,omitempty"` +} + +// ChangeMessageVisibilityRequest preserves the AWS visibility payload. +type ChangeMessageVisibilityRequest struct { + QueueUrl string `json:"QueueUrl"` + ReceiptHandle string `json:"ReceiptHandle"` + VisibilityTimeout int `json:"VisibilityTimeout"` +} + +// ChangeMessageVisibilityBatchRequestEntry mirrors the AWS batch visibility +// entry payload. +type ChangeMessageVisibilityBatchRequestEntry struct { + Id string `json:"Id"` + ReceiptHandle string `json:"ReceiptHandle"` + VisibilityTimeout int `json:"VisibilityTimeout"` +} + +// ChangeMessageVisibilityBatchRequest preserves the AWS batch visibility payload. +type ChangeMessageVisibilityBatchRequest struct { + Entries []ChangeMessageVisibilityBatchRequestEntry `json:"Entries"` + QueueUrl string `json:"QueueUrl"` +} + +// ChangeMessageVisibilityBatchResultEntry mirrors the AWS batch visibility +// success entry. +type ChangeMessageVisibilityBatchResultEntry struct { + Id string `json:"Id,omitempty"` +} + +// ChangeMessageVisibilityBatchResult preserves the AWS batch visibility +// response shape. +type ChangeMessageVisibilityBatchResult struct { + Failed []BatchResultErrorEntry `json:"Failed,omitempty"` + Successful []ChangeMessageVisibilityBatchResultEntry `json:"Successful,omitempty"` +} + +// ReceiveMessageRequest preserves the AWS receive payload. +type ReceiveMessageRequest struct { + AttributeNames []string `json:"AttributeNames,omitempty"` + MaxNumberOfMessages int `json:"MaxNumberOfMessages,omitempty"` + MessageAttributeNames []string `json:"MessageAttributeNames,omitempty"` + MessageSystemAttributeNames []string `json:"MessageSystemAttributeNames,omitempty"` + QueueUrl string `json:"QueueUrl"` + ReceiveRequestAttemptId string `json:"ReceiveRequestAttemptId,omitempty"` + VisibilityTimeout int `json:"VisibilityTimeout,omitempty"` + WaitTimeSeconds int `json:"WaitTimeSeconds,omitempty"` +} + +// ReceivedMessage mirrors the AWS receive response message shape. +type ReceivedMessage struct { + Attributes map[string]string `json:"Attributes,omitempty"` + Body string `json:"Body,omitempty"` + MD5OfBody string `json:"MD5OfBody,omitempty"` + MD5OfMessageAttributes string `json:"MD5OfMessageAttributes,omitempty"` + MessageAttributes map[string]MessageAttributeValue `json:"MessageAttributes,omitempty"` + MessageId string `json:"MessageId,omitempty"` + ReceiptHandle string `json:"ReceiptHandle,omitempty"` +} + +// ReceiveMessageResult preserves the AWS receive response wrapper. +type ReceiveMessageResult struct { + Messages []ReceivedMessage `json:"Messages,omitempty"` +} diff --git a/core/internal/resources/sqs/contracts/message_test.go b/core/internal/resources/sqs/contracts/message_test.go new file mode 100644 index 0000000..a0a0571 --- /dev/null +++ b/core/internal/resources/sqs/contracts/message_test.go @@ -0,0 +1,120 @@ +package contracts + +import ( + "encoding/json" + "strings" + "testing" +) + +func TestMessageContractsPreserveAWSFieldNames(t *testing.T) { + t.Helper() + + request := SendMessageRequest{ + DelaySeconds: 5, + MessageAttributes: map[string]MessageAttributeValue{ + "trace": { + DataType: "String", + StringValue: "abc", + }, + }, + MessageBody: "payload", + MessageDeduplicationId: "dedupe-1", + MessageGroupId: "group-1", + MessageSystemAttributes: map[string]MessageAttributeValue{ + "AWSTraceHeader": { + DataType: "String", + StringValue: "Root=1-12345678-1234567890abcdef12345678", + }, + }, + QueueUrl: "https://sqs.us-east-1.amazonaws.com/123456789012/orders", + } + + data, err := json.Marshal(request) + if err != nil { + t.Fatalf("marshal send request: %v", err) + } + assertJSONContains(t, string(data), []string{ + `"DelaySeconds":5`, + `"MessageAttributes"`, + `"MessageBody":"payload"`, + `"MessageDeduplicationId":"dedupe-1"`, + `"MessageGroupId":"group-1"`, + `"MessageSystemAttributes"`, + `"QueueUrl":"https://sqs.us-east-1.amazonaws.com/123456789012/orders"`, + }) + + batch := SendMessageBatchRequest{ + Entries: []SendMessageBatchRequestEntry{ + { + Id: "entry-1", + MessageBody: "payload", + }, + }, + QueueUrl: request.QueueUrl, + } + data, err = json.Marshal(batch) + if err != nil { + t.Fatalf("marshal send batch request: %v", err) + } + assertJSONContains(t, string(data), []string{ + `"Entries"`, + `"Id":"entry-1"`, + `"MessageBody":"payload"`, + `"QueueUrl":"https://sqs.us-east-1.amazonaws.com/123456789012/orders"`, + }) + + sendResult := SendMessageResult{ + MD5OfMessageAttributes: "md5-attrs", + MD5OfMessageBody: "md5-body", + MD5OfMessageSystemAttributes: "md5-system", + MessageId: "message-1", + SequenceNumber: "42", + } + data, err = json.Marshal(sendResult) + if err != nil { + t.Fatalf("marshal send result: %v", err) + } + assertJSONContains(t, string(data), []string{ + `"MD5OfMessageAttributes":"md5-attrs"`, + `"MD5OfMessageBody":"md5-body"`, + `"MD5OfMessageSystemAttributes":"md5-system"`, + `"MessageId":"message-1"`, + `"SequenceNumber":"42"`, + }) + + receiveResult := ReceiveMessageResult{ + Messages: []ReceivedMessage{ + { + Attributes: map[string]string{ + "ApproximateReceiveCount": "1", + }, + Body: "payload", + MD5OfBody: "md5-body", + MD5OfMessageAttributes: "md5-attrs", + MessageId: "message-1", + ReceiptHandle: "receipt-1", + }, + }, + } + data, err = json.Marshal(receiveResult) + if err != nil { + t.Fatalf("marshal receive result: %v", err) + } + assertJSONContains(t, string(data), []string{ + `"Messages"`, + `"Body":"payload"`, + `"MD5OfBody":"md5-body"`, + `"MessageId":"message-1"`, + `"ReceiptHandle":"receipt-1"`, + }) +} + +func assertJSONContains(t *testing.T, json string, expected []string) { + t.Helper() + + for _, token := range expected { + if !strings.Contains(json, token) { + t.Fatalf("expected json to contain %q, got %s", token, json) + } + } +} diff --git a/core/internal/resources/sqs/contracts/queue.go b/core/internal/resources/sqs/contracts/queue.go new file mode 100644 index 0000000..3e39afd --- /dev/null +++ b/core/internal/resources/sqs/contracts/queue.go @@ -0,0 +1,10 @@ +package contracts + +// QueueAttributesView is the shared queue attribute response contract used by +// the application service and native transport scaffolding. +type QueueAttributesView struct { + QueueName string + QueueURL string + QueueARN string + Attributes map[string]string +} diff --git a/core/internal/resources/sqs/domain/state.go b/core/internal/resources/sqs/domain/state.go index 0e525a3..b13e0cc 100644 --- a/core/internal/resources/sqs/domain/state.go +++ b/core/internal/resources/sqs/domain/state.go @@ -13,6 +13,9 @@ type State struct { Queues []Queue Messages []Message RecoveryMetadata map[string]RecoveryMetadata + QueueTags map[string]map[string]string + QueuePermissions map[string]map[string]QueuePermission + MoveTasks map[string]map[string]MessageMoveTask } type Queue struct { @@ -23,6 +26,8 @@ type Queue struct { Recovery QueueRecovery CreatedAt time.Time UpdatedAt time.Time + DeletedAt time.Time + PurgedAt time.Time } type QueueRecovery struct { @@ -30,6 +35,27 @@ type QueueRecovery struct { Policy map[string]string } +type QueuePermission struct { + Label string + AWSAccountIDs []string + Actions []string + CreatedAt time.Time + UpdatedAt time.Time +} + +type MessageMoveTask struct { + TaskHandle string + SourceQueue string + SourceArn string + DestinationArn string + MaxNumberOfMessagesPerSecond int + ApproximateNumberOfMessagesMoved int64 + Status string + StartedAt time.Time + UpdatedAt time.Time + CancelledAt time.Time +} + type Message struct { Queue string MessageID string @@ -70,6 +96,9 @@ func NewState() State { Queues: []Queue{}, Messages: []Message{}, RecoveryMetadata: map[string]RecoveryMetadata{}, + QueueTags: map[string]map[string]string{}, + QueuePermissions: map[string]map[string]QueuePermission{}, + MoveTasks: map[string]map[string]MessageMoveTask{}, } } @@ -87,6 +116,8 @@ func (s State) Snapshot() map[string]any { }, "created_at": snapshotTime(queue.CreatedAt), "updated_at": snapshotTime(queue.UpdatedAt), + "deleted_at": snapshotTime(queue.DeletedAt), + "purged_at": snapshotTime(queue.PurgedAt), }) } @@ -134,11 +165,81 @@ func (s State) Snapshot() map[string]any { } } + queueTags := make(map[string]any, len(s.QueueTags)) + tagQueueNames := make([]string, 0, len(s.QueueTags)) + for queueName := range s.QueueTags { + tagQueueNames = append(tagQueueNames, queueName) + } + sort.Strings(tagQueueNames) + for _, queueName := range tagQueueNames { + queueTags[queueName] = cloneStringMapAny(s.QueueTags[queueName]) + } + + queuePermissions := make(map[string]any, len(s.QueuePermissions)) + permissionQueueNames := make([]string, 0, len(s.QueuePermissions)) + for queueName := range s.QueuePermissions { + permissionQueueNames = append(permissionQueueNames, queueName) + } + sort.Strings(permissionQueueNames) + for _, queueName := range permissionQueueNames { + labels := make([]string, 0, len(s.QueuePermissions[queueName])) + for label := range s.QueuePermissions[queueName] { + labels = append(labels, label) + } + sort.Strings(labels) + entries := make([]any, 0, len(labels)) + for _, label := range labels { + permission := s.QueuePermissions[queueName][label] + entries = append(entries, map[string]any{ + "label": permission.Label, + "aws_account_ids": append([]string(nil), permission.AWSAccountIDs...), + "actions": append([]string(nil), permission.Actions...), + "created_at": snapshotTime(permission.CreatedAt), + "updated_at": snapshotTime(permission.UpdatedAt), + }) + } + queuePermissions[queueName] = entries + } + + moveTasks := make(map[string]any, len(s.MoveTasks)) + moveQueueNames := make([]string, 0, len(s.MoveTasks)) + for queueName := range s.MoveTasks { + moveQueueNames = append(moveQueueNames, queueName) + } + sort.Strings(moveQueueNames) + for _, queueName := range moveQueueNames { + handles := make([]string, 0, len(s.MoveTasks[queueName])) + for handle := range s.MoveTasks[queueName] { + handles = append(handles, handle) + } + sort.Strings(handles) + entries := make([]any, 0, len(handles)) + for _, handle := range handles { + task := s.MoveTasks[queueName][handle] + entries = append(entries, map[string]any{ + "task_handle": task.TaskHandle, + "source_queue": task.SourceQueue, + "source_arn": task.SourceArn, + "destination_arn": task.DestinationArn, + "max_number_of_messages_per_second": task.MaxNumberOfMessagesPerSecond, + "approximate_number_of_messages_moved": task.ApproximateNumberOfMessagesMoved, + "status": task.Status, + "started_at": snapshotTime(task.StartedAt), + "updated_at": snapshotTime(task.UpdatedAt), + "cancelled_at": snapshotTime(task.CancelledAt), + }) + } + moveTasks[queueName] = entries + } + return map[string]any{ "service": s.Service, "queues": queues, "messages": messages, "recovery_metadata": recovery, + "queue_tags": queueTags, + "queue_permissions": queuePermissions, + "move_tasks": moveTasks, } } @@ -148,6 +249,9 @@ func (s State) Clone() State { Queues: make([]Queue, len(s.Queues)), Messages: make([]Message, len(s.Messages)), RecoveryMetadata: make(map[string]RecoveryMetadata, len(s.RecoveryMetadata)), + QueueTags: make(map[string]map[string]string, len(s.QueueTags)), + QueuePermissions: make(map[string]map[string]QueuePermission, len(s.QueuePermissions)), + MoveTasks: make(map[string]map[string]MessageMoveTask, len(s.MoveTasks)), } copy(cloned.Queues, s.Queues) for i := range cloned.Queues { @@ -169,6 +273,38 @@ func (s State) Clone() State { Detail: cloneStringMap(value.Detail), } } + for queueName, tags := range s.QueueTags { + cloned.QueueTags[queueName] = cloneStringMap(tags) + } + for queueName, permissions := range s.QueuePermissions { + cloned.QueuePermissions[queueName] = make(map[string]QueuePermission, len(permissions)) + for label, permission := range permissions { + cloned.QueuePermissions[queueName][label] = QueuePermission{ + Label: permission.Label, + AWSAccountIDs: append([]string(nil), permission.AWSAccountIDs...), + Actions: append([]string(nil), permission.Actions...), + CreatedAt: permission.CreatedAt, + UpdatedAt: permission.UpdatedAt, + } + } + } + for queueName, tasks := range s.MoveTasks { + cloned.MoveTasks[queueName] = make(map[string]MessageMoveTask, len(tasks)) + for handle, task := range tasks { + cloned.MoveTasks[queueName][handle] = MessageMoveTask{ + TaskHandle: task.TaskHandle, + SourceQueue: task.SourceQueue, + SourceArn: task.SourceArn, + DestinationArn: task.DestinationArn, + MaxNumberOfMessagesPerSecond: task.MaxNumberOfMessagesPerSecond, + ApproximateNumberOfMessagesMoved: task.ApproximateNumberOfMessagesMoved, + Status: task.Status, + StartedAt: task.StartedAt, + UpdatedAt: task.UpdatedAt, + CancelledAt: task.CancelledAt, + } + } + } return cloned } diff --git a/core/internal/resources/sqs/domain/state_test.go b/core/internal/resources/sqs/domain/state_test.go index 9cb4949..dc53322 100644 --- a/core/internal/resources/sqs/domain/state_test.go +++ b/core/internal/resources/sqs/domain/state_test.go @@ -24,6 +24,8 @@ func TestStateSnapshotCopiesLiveData(t *testing.T) { }, CreatedAt: time.Date(2026, time.April, 19, 10, 0, 0, 0, time.UTC), UpdatedAt: time.Date(2026, time.April, 19, 10, 1, 0, 0, time.UTC), + DeletedAt: time.Date(2026, time.April, 19, 10, 2, 0, 0, time.UTC), + PurgedAt: time.Date(2026, time.April, 19, 10, 3, 0, 0, time.UTC), }) state.Messages = append(state.Messages, Message{ Queue: "queue-a", @@ -85,6 +87,12 @@ func TestStateSnapshotCopiesLiveData(t *testing.T) { if got, want := state.Queues[0].OrderingHint, "fifo"; got != want { t.Fatalf("queue ordering hint was aliased: got %q want %q", got, want) } + if got, want := state.Queues[0].DeletedAt, time.Date(2026, time.April, 19, 10, 2, 0, 0, time.UTC); !got.Equal(want) { + t.Fatalf("queue deleted at was aliased: got %v want %v", got, want) + } + if got, want := state.Queues[0].PurgedAt, time.Date(2026, time.April, 19, 10, 3, 0, 0, time.UTC); !got.Equal(want) { + t.Fatalf("queue purged at was aliased: got %v want %v", got, want) + } if got, want := state.Messages[0].Body, "payload"; got != want { t.Fatalf("message body was aliased: got %q want %q", got, want) } @@ -115,6 +123,8 @@ func TestStateCloneReturnsDeepCopy(t *testing.T) { "DelaySeconds": "0", }, OrderingHint: "standard", + DeletedAt: time.Date(2026, time.April, 19, 11, 0, 0, 0, time.UTC), + PurgedAt: time.Date(2026, time.April, 19, 11, 1, 0, 0, time.UTC), Recovery: QueueRecovery{ Policy: map[string]string{"enabled": "true"}, }, @@ -165,6 +175,12 @@ func TestStateCloneReturnsDeepCopy(t *testing.T) { if got, want := state.Queues[0].Recovery.Policy["enabled"], "true"; got != want { t.Fatalf("queue policy was shared with clone: got %q want %q", got, want) } + if got, want := state.Queues[0].DeletedAt, time.Date(2026, time.April, 19, 11, 0, 0, 0, time.UTC); !got.Equal(want) { + t.Fatalf("queue deleted at was shared with clone: got %v want %v", got, want) + } + if got, want := state.Queues[0].PurgedAt, time.Date(2026, time.April, 19, 11, 1, 0, 0, time.UTC); !got.Equal(want) { + t.Fatalf("queue purged at was shared with clone: got %v want %v", got, want) + } if got, want := state.Messages[0].Tags[0], "alpha"; got != want { t.Fatalf("message tags were shared with clone: got %q want %q", got, want) } @@ -202,6 +218,8 @@ func TestStateKeepsRoomForRecoveryMetadataAndAttributes(t *testing.T) { "RedrivePolicy": "present", }, OrderingHint: "fifo", + DeletedAt: time.Date(2026, time.April, 19, 12, 0, 0, 0, time.UTC), + PurgedAt: time.Date(2026, time.April, 19, 12, 1, 0, 0, time.UTC), Recovery: QueueRecovery{ DeadLetterQueue: "queue-dlq", }, @@ -218,6 +236,12 @@ func TestStateKeepsRoomForRecoveryMetadataAndAttributes(t *testing.T) { if got, want := state.Queues[0].OrderingHint, "fifo"; got != want { t.Fatalf("unexpected queue ordering hint: got %q want %q", got, want) } + if got, want := state.Queues[0].DeletedAt, time.Date(2026, time.April, 19, 12, 0, 0, 0, time.UTC); !got.Equal(want) { + t.Fatalf("unexpected queue deleted at: got %v want %v", got, want) + } + if got, want := state.Queues[0].PurgedAt, time.Date(2026, time.April, 19, 12, 1, 0, 0, time.UTC); !got.Equal(want) { + t.Fatalf("unexpected queue purged at: got %v want %v", got, want) + } if got, want := state.RecoveryMetadata["queue-a"].Detail["state"], "ready"; got != want { t.Fatalf("unexpected recovery detail: got %q want %q", got, want) } diff --git a/package-lock.json b/package-lock.json new file mode 100644 index 0000000..66661b9 --- /dev/null +++ b/package-lock.json @@ -0,0 +1,6 @@ +{ + "name": "mildstack", + "lockfileVersion": 3, + "requires": true, + "packages": {} +} diff --git a/package.json b/package.json new file mode 100644 index 0000000..0967ef4 --- /dev/null +++ b/package.json @@ -0,0 +1 @@ +{}