diff --git a/.env.example b/.env.example index 7112cd56..9a1690bb 100644 --- a/.env.example +++ b/.env.example @@ -61,6 +61,8 @@ QUEUE_CONCURRENCY=5 # Counters are stored in Redis so every API replica enforces the same budget. THROTTLE_AUTH_LIMIT=10 THROTTLE_API_LIMIT=120 +# Per-agent tier: agent-identified traffic (x-agent-id / agent-bound API keys). +THROTTLE_AGENT_LIMIT=300 THROTTLE_WEBHOOK_LIMIT=30 THROTTLE_TTL=60 # Burst limits — short-term spike allowance per tier (requests per second) diff --git a/API_DOCUMENTATION.md b/API_DOCUMENTATION.md index 2eb3eafb..18facfc5 100644 --- a/API_DOCUMENTATION.md +++ b/API_DOCUMENTATION.md @@ -83,8 +83,9 @@ List wallets for the organization. **Query Parameters:** | Field | Type | Required | Description | |-------|------|----------|-------------| -| page | number | No | Page number (default: 1) | -| limit | number | No | Items per page (default: 10) | +| offset | number | No | Rows to skip (default: 0); mutually exclusive with `page` | +| page | number | No | Page number, alternative to `offset` (default: 1) | +| limit | number | No | Items per page (default: 50, max: 200) | **Authentication:** Bearer token required @@ -168,8 +169,9 @@ List transactions for the organization. **Query Parameters:** | Field | Type | Required | Description | |-------|------|----------|-------------| -| page | number | No | Page number (default: 1) | -| limit | number | No | Items per page (default: 10) | +| offset | number | No | Rows to skip (default: 0); mutually exclusive with `page` | +| page | number | No | Page number, alternative to `offset` (default: 1) | +| limit | number | No | Items per page (default: 50, max: 200) | **Authentication:** Bearer token required @@ -228,8 +230,9 @@ List policies for the organization. **Query Parameters:** | Field | Type | Required | Description | |-------|------|----------|-------------| -| page | number | No | Page number (default: 1) | -| limit | number | No | Items per page (default: 10) | +| offset | number | No | Rows to skip (default: 0); mutually exclusive with `page` | +| page | number | No | Page number, alternative to `offset` (default: 1) | +| limit | number | No | Items per page (default: 50, max: 200) | **Authentication:** Bearer token required @@ -322,8 +325,9 @@ List budgets for the organization. **Query Parameters:** | Field | Type | Required | Description | |-------|------|----------|-------------| -| page | number | No | Page number (default: 1) | -| limit | number | No | Items per page (default: 10) | +| offset | number | No | Rows to skip (default: 0); mutually exclusive with `page` | +| page | number | No | Page number, alternative to `offset` (default: 1) | +| limit | number | No | Items per page (default: 50, max: 200) | **Authentication:** Bearer token required @@ -372,6 +376,51 @@ Delete a budget. --- +## Request Correlation + +Every response includes the selected request ID in the `x-request-id` header and the standard response envelope. A caller-supplied ID is accepted only when it is 1-128 ASCII characters, starts with a letter or digit, and otherwise contains only letters, digits, `.`, `_`, `:`, or `-`. Invalid values are replaced with a server-generated UUID. The selected ID is propagated as typed event/job metadata and is isolated per concurrent request. + +## Audit History (`/audit`) + +### GET `/audit` +List audit records in reverse chronological order using a stable `(createdAt, id)` keyset. + +**Query Parameters:** `limit` defaults to 20 and is bounded to 1-100; `cursor` is the opaque `nextCursor` from the previous response; optional `actorId`, `action`, `resourceId`, `from`, and `to` filters apply within the authenticated organization. `from` and `to` are inclusive ISO 8601 timestamps and `from` must not be later than `to`. + +The response `meta` includes `limit`, `hasNext`, and `nextCursor` (null on the final page). New records inserted after a page is read do not shift subsequent pages. + +**Example:** `GET /audit?limit=20&actorId=user-123&action=wallet.created` + +## Outbound Webhook Signatures + +Webhook creation and secret rotation responses disclose the signing secret once. Later list, get, update, delivery, and audit responses never include it. Store the secret securely and rotate it when compromised. + +Every delivery includes `x-astroid-signature`, `x-astroid-signature-version`, `x-astroid-timestamp`, and `x-astroid-event-id`. `x-astroid-delivery` remains an alias for the event ID. The signature header is `v1=` and the version header is `v1`. + +The canonical signed bytes are UTF-8 `v1...` followed by the exact raw HTTP body bytes. Each retry uses the same event ID and body, with a fresh timestamp and signature. Consumers should also reject timestamps outside their chosen replay window. + +```js +import { createHmac, timingSafeEqual } from 'node:crypto'; + +export function verifyAstroidWebhook({ secret, headers, rawBody }) { + const version = headers['x-astroid-signature-version']; + const timestamp = headers['x-astroid-timestamp']; + const eventId = headers['x-astroid-event-id']; + const received = headers['x-astroid-signature']; + if (version !== 'v1' || !/^\d{1,12}$/.test(timestamp) || !eventId) return false; + if (Math.abs(Date.now() / 1000 - Number(timestamp)) > 300) return false; + + const prefix = Buffer.from(`v1.${timestamp}.${eventId}.`, 'utf8'); + const expected = createHmac('sha256', secret) + .update(Buffer.concat([prefix, rawBody])) + .digest(); + const match = /^v1=([0-9a-f]{64})$/.exec(received); + if (!match) return false; + const actual = Buffer.from(match[1], 'hex'); + return actual.length === expected.length && timingSafeEqual(actual, expected); +} +``` + ## Health Probes (`/health`) The liveness and readiness probes are served **outside** the API prefix, so @@ -424,10 +473,27 @@ under the API prefix at `GET /{API_PREFIX}/health/readiness`, ## Common Types ### Pagination Query +Every list endpoint accepts the same query parameters. + | Field | Type | Default | Description | |-------|------|---------|-------------| -| page | number | 1 | Page number | -| limit | number | 10 | Items per page | +| offset | number | 0 | Zero-based number of rows to skip. Mutually exclusive with `page` | +| page | number | 1 | 1-based page number, an alternative to `offset` | +| limit | number | 50 | Items per page, capped at 200 | +| sort | string | createdAt | Sort field (restricted to an allow-list per endpoint) | +| order | `asc` \| `desc` | desc | Sort direction | + +Negative, non-integer or non-numeric values, a `limit` above 200, or supplying both `offset` and `page` return `400 Bad Request`. + +Paginated responses carry the total row count in the `X-Total-Count` header and in `meta`: +```json +{ + "success": true, + "data": [], + "meta": { "offset": 50, "page": 2, "limit": 50, "total": 120, "totalPages": 3, "hasNext": true, "hasPrev": true }, + "requestId": "req_..." +} +``` ### Error Response All endpoints return errors as [RFC 9457](https://www.rfc-editor.org/rfc/rfc9457) problem details with `Content-Type: application/problem+json`: diff --git a/bun.lock b/bun.lock index b343d946..b6bbe17c 100644 --- a/bun.lock +++ b/bun.lock @@ -1,6 +1,5 @@ { "lockfileVersion": 1, - "configVersion": 0, "workspaces": { "": { "name": "astroid-api", @@ -15,6 +14,7 @@ "@nestjs/platform-express": "10.4.15", "@nestjs/schedule": "^4.1.2", "@nestjs/swagger": "7.4.2", + "@nestjs/terminus": "^10.3.0", "@nestjs/throttler": "6.3.0", "@prisma/client": "5.22.0", "@simplewebauthn/server": "^13.3.3", @@ -213,6 +213,8 @@ "@nestjs/swagger": ["@nestjs/swagger@7.4.2", "", { "dependencies": { "@microsoft/tsdoc": "^0.15.0", "@nestjs/mapped-types": "2.0.5", "js-yaml": "4.1.0", "lodash": "4.17.21", "path-to-regexp": "3.3.0", "swagger-ui-dist": "5.17.14" }, "peerDependencies": { "@fastify/static": "^6.0.0 || ^7.0.0", "@nestjs/common": "^9.0.0 || ^10.0.0", "@nestjs/core": "^9.0.0 || ^10.0.0", "class-transformer": "*", "class-validator": "*", "reflect-metadata": "^0.1.12 || ^0.2.0" }, "optionalPeers": ["@fastify/static"] }, "sha512-Mu6TEn1M/owIvAx2B4DUQObQXqo2028R2s9rSZ/hJEgBK95+doTwS0DjmVA2wTeZTyVtXOoN7CsoM5pONBzvKQ=="], + "@nestjs/terminus": ["@nestjs/terminus@10.3.0", "", { "dependencies": { "boxen": "5.1.2", "check-disk-space": "3.4.0" }, "peerDependencies": { "@grpc/grpc-js": "*", "@grpc/proto-loader": "*", "@mikro-orm/core": "*", "@mikro-orm/nestjs": "*", "@nestjs/axios": "^1.0.0 || ^2.0.0 || ^3.0.0", "@nestjs/common": "^9.0.0 || ^10.0.0", "@nestjs/core": "^9.0.0 || ^10.0.0", "@nestjs/microservices": "^9.0.0 || ^10.0.0", "@nestjs/mongoose": "^9.0.0 || ^10.0.0", "@nestjs/sequelize": "^9.0.0 || ^10.0.0", "@nestjs/typeorm": "^9.0.0 || ^10.0.0", "@prisma/client": "*", "mongoose": "*", "reflect-metadata": "0.1.x || 0.2.x", "rxjs": "7.x", "sequelize": "*", "typeorm": "*" }, "optionalPeers": ["@grpc/grpc-js", "@grpc/proto-loader", "@mikro-orm/core", "@mikro-orm/nestjs", "@nestjs/axios", "@nestjs/microservices", "@nestjs/mongoose", "@nestjs/sequelize", "@nestjs/typeorm", "@prisma/client", "mongoose", "sequelize", "typeorm"] }, "sha512-vOJGCwt1OgrFuuxWQwPoaHqy9m9CfIk2qMUX2mosZLK5dFVJSEjHXrklkh3/Fw9PiUnfzvYFfiAdJRzUaxx+5Q=="], + "@nestjs/testing": ["@nestjs/testing@10.4.15", "", { "dependencies": { "tslib": "2.8.1" }, "peerDependencies": { "@nestjs/common": "^10.0.0", "@nestjs/core": "^10.0.0", "@nestjs/microservices": "^10.0.0", "@nestjs/platform-express": "^10.0.0" }, "optionalPeers": ["@nestjs/microservices"] }, "sha512-eGlWESkACMKti+iZk1hs6FUY/UqObmMaa8HAN9JLnaYkoLf1Jeh+EuHlGnfqo/Rq77oznNLIyaA3PFjrFDlNUg=="], "@nestjs/throttler": ["@nestjs/throttler@6.3.0", "", { "peerDependencies": { "@nestjs/common": "^7.0.0 || ^8.0.0 || ^9.0.0 || ^10.0.0", "@nestjs/core": "^7.0.0 || ^8.0.0 || ^9.0.0 || ^10.0.0", "reflect-metadata": "^0.1.13 || ^0.2.0" } }, "sha512-IqTMbl5Iyxjts7NwbVriDND0Cnr8rwNqAPpF5HJE+UV+2VrVUBwCfDXKEiXu47vzzaQLlWPYegBsGO9OXxa+oQ=="], @@ -497,6 +499,8 @@ "ajv-keywords": ["ajv-keywords@3.5.2", "", { "peerDependencies": { "ajv": "^6.9.1" } }, "sha512-5p6WTN0DdTGVQk6VjcEju19IgaHudalcfabD7yhDGeA6bcQnmL+CpveLJq/3hvfwd1aof6L386Ougkx6RfyMIQ=="], + "ansi-align": ["ansi-align@3.0.1", "", { "dependencies": { "string-width": "^4.1.0" } }, "sha512-IOfwwBF5iczOjp/WeY4YxyjqAFMQoZufdQWDd19SEExbVLNXqvpzSJ/M7Za4/sCPmQ0+GRquoA7bGcINcxew6w=="], + "ansi-colors": ["ansi-colors@4.1.3", "", {}, "sha512-/6w/C21Pm1A7aZitlI5Ni/2J6FFQN8i1Cvz3kHABAAbw93v/NlvKdVOqz7CCWz/3iv/JplRSEEZ83XION15ovw=="], "ansi-escapes": ["ansi-escapes@4.3.2", "", { "dependencies": { "type-fest": "^0.21.3" } }, "sha512-gKXj5ALrKWQLsYG9jlTRmR/xKluxHV+Z9QEwNIgCfM1/uwPMCuzVVnh5mwTd+OuBZcwSIMbqssNWRm1lE51QaQ=="], @@ -555,6 +559,8 @@ "body-parser": ["body-parser@1.20.3", "", { "dependencies": { "bytes": "3.1.2", "content-type": "~1.0.5", "debug": "2.6.9", "depd": "2.0.0", "destroy": "1.2.0", "http-errors": "2.0.0", "iconv-lite": "0.4.24", "on-finished": "2.4.1", "qs": "6.13.0", "raw-body": "2.5.2", "type-is": "~1.6.18", "unpipe": "1.0.0" } }, "sha512-7rAxByjUMqQ3/bHJy7D6OGXvx/MMc4IqBn/X0fcM1QUcAItpZrBEYhWGem+tzXH90c+G01ypMcYJBO9Y30203g=="], + "boxen": ["boxen@5.1.2", "", { "dependencies": { "ansi-align": "^3.0.0", "camelcase": "^6.2.0", "chalk": "^4.1.0", "cli-boxes": "^2.2.1", "string-width": "^4.2.2", "type-fest": "^0.20.2", "widest-line": "^3.1.0", "wrap-ansi": "^7.0.0" } }, "sha512-9gYgQKXx+1nP8mP7CzFyaUARhg7D3n1dF/FnErWmu9l6JvGpNUN278h0aSb+QjoiKSWG+iZ3uHrcqk0qrY9RQQ=="], + "brace-expansion": ["brace-expansion@1.1.18", "", { "dependencies": { "balanced-match": "^1.0.0", "concat-map": "0.0.1" } }, "sha512-Edep/X9fGqVNmzKBVsDYIOtD+z1tuezV70LBjdCst9Tqu76lsnvRiZ6oTic1n+/BIwX6QDGAO94PN4N2SADvtw=="], "braces": ["braces@3.0.3", "", { "dependencies": { "fill-range": "^7.1.1" } }, "sha512-yQbXgO/OSZVD2IsiLlro+7Hf6Q18EJrKSEsdoMzKePKXct3gvD8oLcOQdIzGupr5Fj+EDe8gO/lxc1BzfMpxvA=="], @@ -583,6 +589,8 @@ "callsites": ["callsites@3.1.0", "", {}, "sha512-P8BjAsXvZS+VIDUI11hHCQEv74YT67YUi5JJFNWIqL235sBmjX4+qx9Muvls5ivyNENctx46xQLQ3aTuE7ssaQ=="], + "camelcase": ["camelcase@6.3.0", "", {}, "sha512-Gmy6FhYlCY7uOElZUSbxo2UCDH8owEk996gkbrpsgGtrJLM3J7jGxl9Ic7Qwwj4ivOE5AWZWRMecDdF7hqGjFA=="], + "caniuse-lite": ["caniuse-lite@1.0.30001806", "", {}, "sha512-72Cuvd95zbSYPKq6Fhg8eDJRlzgWDf7/mtoZv6Qe/DYNCEBdNxoA3+rZAU2ZhGCpZlns3EssFavaZomckT5Uuw=="], "chai": ["chai@5.3.3", "", { "dependencies": { "assertion-error": "^2.0.1", "check-error": "^2.1.1", "deep-eql": "^5.0.1", "loupe": "^3.1.0", "pathval": "^2.0.0" } }, "sha512-4zNhdJD/iOjSH0A05ea+Ke6MU5mmpQcbQsSOkgdaUMJ9zTlDTD/GYlwohmIE2u0gaxHYiVHEn1Fw9mZ/ktJWgw=="], @@ -591,6 +599,8 @@ "chardet": ["chardet@0.7.0", "", {}, "sha512-mT8iDcrh03qDGRRmoA2hmBJnxpllMR+0/0qlzjqZES6NdiWDcZkCNAk4rPFZ9Q85r27unkiNNg8ZOiwZXBHwcA=="], + "check-disk-space": ["check-disk-space@3.4.0", "", {}, "sha512-drVkSqfwA+TvuEhFipiR1OC9boEGZL5RrWvVsOthdcvQNXyCCuKkEiTOTXZ7qxSf/GLwq4GvzfrQD/Wz325hgw=="], + "check-error": ["check-error@2.1.3", "", {}, "sha512-PAJdDJusoxnwm1VwW07VWwUN1sl7smmC3OKggvndJFadxxDRyFJBX/ggnu/KE4kQAB7a3Dp8f/YXC1FlUprWmA=="], "chokidar": ["chokidar@3.6.0", "", { "dependencies": { "anymatch": "~3.1.2", "braces": "~3.0.2", "glob-parent": "~5.1.2", "is-binary-path": "~2.1.0", "is-glob": "~4.0.1", "normalize-path": "~3.0.0", "readdirp": "~3.6.0" }, "optionalDependencies": { "fsevents": "~2.3.2" } }, "sha512-7VT13fmjotKpGipCW9JEQAusEPE+Ei8nl6/g4FBAmIm0GOOLMua9NDDo/DWp0ZAxCr3cPq5ZpBqmPAQgDda2Pw=="], @@ -601,6 +611,8 @@ "class-validator": ["class-validator@0.14.1", "", { "dependencies": { "@types/validator": "^13.11.8", "libphonenumber-js": "^1.10.53", "validator": "^13.9.0" } }, "sha512-2VEG9JICxIqTpoK1eMzZqaV+u/EiwEJkMGzTrZf6sU/fwsnOITVgYJ8yojSy6CaXtO9V0Cc6ZQZ8h8m4UBuLwQ=="], + "cli-boxes": ["cli-boxes@2.2.1", "", {}, "sha512-y4coMcylgSCdVinjiDBuR8PCC2bLjyGTwEmPb9NHR/QaNU6EUOXcTY/s6VjGMD6ENSEaeQYHCY0GNGS5jfMwPw=="], + "cli-cursor": ["cli-cursor@3.1.0", "", { "dependencies": { "restore-cursor": "^3.1.0" } }, "sha512-I/zHAwsKf9FqGoXM4WWRACob9+SNukZTd94DWF57E4toouRulbCxcUh6RKUEOQlYTHJnzkPMySvPNaaSLNfLZw=="], "cli-spinners": ["cli-spinners@2.9.2", "", {}, "sha512-ywqV+5MmyL4E7ybXgKys4DugZbX0FC6LnwrhjuykIjnK9k8OQacQ7axGKnjDXWNhns0xot3bZI5h55H8yo9cJg=="], @@ -1425,6 +1437,8 @@ "why-is-node-running": ["why-is-node-running@2.3.0", "", { "dependencies": { "siginfo": "^2.0.0", "stackback": "0.0.2" }, "bin": "cli.js" }, "sha512-hUrmaWBdVDcxvYqnyh09zunKzROWjbZTiNy8dBEjkS7ehEDQibXJ7XvlmtbwuTclUiIyN+CyXQD4Vmko8fNm8w=="], + "widest-line": ["widest-line@3.1.0", "", { "dependencies": { "string-width": "^4.0.0" } }, "sha512-NsmoXalsWVDMGupxZ5R08ka9flZjjiLvHVAWYOKtiKM8ujtZWr9cRffak+uSE48+Ob8ObalXpwyeUiyDD6QFgg=="], + "word-wrap": ["word-wrap@1.2.5", "", {}, "sha512-BN22B5eaMMI9UMtjrGd5g5eCYPpCPDUy0FJXbYsaT5zYxjFOckS53SQDE3pWkVoWpHXVb3BrYcEN4Twa55B5cA=="], "wrap-ansi": ["wrap-ansi@6.2.0", "", { "dependencies": { "ansi-styles": "^4.0.0", "string-width": "^4.1.0", "strip-ansi": "^6.0.0" } }, "sha512-r6lPcBGxZXlIcymEu7InxDMhdW0KDxpLgoFLcguasxCaJ/SOIZwINatK9KY/tf+ZrlywOKU0UDj3ATXUBfxJXA=="], @@ -1455,12 +1469,6 @@ "@cspotcode/source-map-support/@jridgewell/trace-mapping": ["@jridgewell/trace-mapping@0.3.9", "", { "dependencies": { "@jridgewell/resolve-uri": "^3.0.3", "@jridgewell/sourcemap-codec": "^1.4.10" } }, "sha512-3Belt6tdc8bPgAtbcmdtNJlirVoTmEb5e2gC94PnkwEW9jI6CAHUeoG85tjWP5WquqfavoMtMwiG4P926ZKKuQ=="], - "@eslint/eslintrc/ajv": ["ajv@6.15.0", "", { "dependencies": { "fast-deep-equal": "^3.1.1", "fast-json-stable-stringify": "^2.0.0", "json-schema-traverse": "^0.4.1", "uri-js": "^4.2.2" } }, "sha512-fgFx7Hfoq60ytK2c7DhnF8jIvzYgOMxfugjLOSMHjLIPgenqa7S7oaagATUq99mV6IYvN2tRmC0wnTYX6iPbMw=="], - - "@eslint/eslintrc/minimatch": ["minimatch@3.1.5", "", { "dependencies": { "brace-expansion": "^1.1.7" } }, "sha512-VgjWUsnnT6n+NUk6eZq77zeFdpW2LWDzP6zFGrCbHXiYNul5Dzqk2HHQ5uFH2DNW5Xbp8+jVzaeNt94ssEEl4w=="], - - "@humanwhocodes/config-array/minimatch": ["minimatch@3.1.5", "", { "dependencies": { "brace-expansion": "^1.1.7" } }, "sha512-VgjWUsnnT6n+NUk6eZq77zeFdpW2LWDzP6zFGrCbHXiYNul5Dzqk2HHQ5uFH2DNW5Xbp8+jVzaeNt94ssEEl4w=="], - "@isaacs/cliui/string-width": ["string-width@5.1.2", "", { "dependencies": { "eastasianwidth": "^0.2.0", "emoji-regex": "^9.2.2", "strip-ansi": "^7.0.1" } }, "sha512-HnLOCR3vjcY8beoNLtcjZ5/nxn2afmME6lhrDrebokqMap+XbeW8n9TXpPDOqdGK5qcI3oT0GKTW6wC7EMiVqA=="], "@isaacs/cliui/strip-ansi": ["strip-ansi@7.2.0", "", { "dependencies": { "ansi-regex": "^6.2.2" } }, "sha512-yDPMNjp4WyfYBkHnjIRLfca1i6KMyGCtsVgoKe/z1+6vukgaENdgGBZt+ZmKPc4gavvEZ5OgHfHdrazhgNyG7w=="], @@ -1479,12 +1487,8 @@ "@vitest/mocker/estree-walker": ["estree-walker@3.0.3", "", { "dependencies": { "@types/estree": "^1.0.0" } }, "sha512-7RUKfXgSMMkzt6ZuXmqapOurLGPPfgj6l9uRZ7lRGolvk0y2yocc35LdcxKC5PQZdn2DMqioAQ2NoWcrTKmm6g=="], - "@vitest/mocker/magic-string": ["magic-string@0.30.21", "", { "dependencies": { "@jridgewell/sourcemap-codec": "^1.5.5" } }, "sha512-vd2F4YUyEXKGcLHoq+TEyCjxueSeHnFxyyjNp80yg0XV4vUhnDer/lvvlqM/arB5bXQN5K2/3oinyCRyx8T2CQ=="], - "@vitest/snapshot/@vitest/pretty-format": ["@vitest/pretty-format@2.1.8", "", { "dependencies": { "tinyrainbow": "^1.2.0" } }, "sha512-9HiSZ9zpqNLKlbIDRWOnAWqgcA7xu+8YxXSekhr0Ykab7PAYFkhkwoqVArPOtJhPmYeE2YHgKZlj3CP36z2AJQ=="], - "@vitest/snapshot/magic-string": ["magic-string@0.30.21", "", { "dependencies": { "@jridgewell/sourcemap-codec": "^1.5.5" } }, "sha512-vd2F4YUyEXKGcLHoq+TEyCjxueSeHnFxyyjNp80yg0XV4vUhnDer/lvvlqM/arB5bXQN5K2/3oinyCRyx8T2CQ=="], - "@vitest/utils/@vitest/pretty-format": ["@vitest/pretty-format@2.1.8", "", { "dependencies": { "tinyrainbow": "^1.2.0" } }, "sha512-9HiSZ9zpqNLKlbIDRWOnAWqgcA7xu+8YxXSekhr0Ykab7PAYFkhkwoqVArPOtJhPmYeE2YHgKZlj3CP36z2AJQ=="], "ajv-formats/ajv": ["ajv@8.12.0", "", { "dependencies": { "fast-deep-equal": "^3.1.1", "json-schema-traverse": "^1.0.0", "require-from-string": "^2.0.2", "uri-js": "^4.2.2" } }, "sha512-sRu1kpcO9yLtYxBKvqfTeh9KzZEwO3STyX1HT+4CaDzC6HpTGYhIhPIzj9XuKU7KYDwnaeh5hcOwjy1QuJzBPA=="], @@ -1497,6 +1501,8 @@ "body-parser/debug": ["debug@2.6.9", "", { "dependencies": { "ms": "2.0.0" } }, "sha512-bC7ElrdJaJnPbAP+1EotYvqZsb3ecl5wi6Bfi6BJTUcNowp6cvspg0jXznRTKDjm/E7AdgFBVeAPVMNcKGsHMA=="], + "boxen/wrap-ansi": ["wrap-ansi@7.0.0", "", { "dependencies": { "ansi-styles": "^4.0.0", "string-width": "^4.1.0", "strip-ansi": "^6.0.0" } }, "sha512-YVGIj2kamLSTxw6NsZjoBxfSwsn0ycdesmc4p+Q21c5zPuZ1pl+NfxVdxPtdHvmNVOQ6XSYG4AUtyt/Fi7D16Q=="], + "bullmq/uuid": ["uuid@9.0.1", "", { "bin": "dist/bin/uuid" }, "sha512-b+1eJOlsR9K8HJpow9Ok3fiWOWSIcIzXodvv0rQjVoOVNpWMpxf1wZNpt4y9h10odCNrqnYp1OBzRktckBe3sA=="], "chokidar/glob-parent": ["glob-parent@5.1.2", "", { "dependencies": { "is-glob": "^4.0.1" } }, "sha512-AOIgSQCepiJYwP3ARnGx+5VnTu2HBYdzbGP45eLw1vr3zB3vZLeyed1sC9hnbcOc9/SrMyM5RPQrkGz4aS9Zow=="], @@ -1515,8 +1521,6 @@ "finalhandler/debug": ["debug@2.6.9", "", { "dependencies": { "ms": "2.0.0" } }, "sha512-bC7ElrdJaJnPbAP+1EotYvqZsb3ecl5wi6Bfi6BJTUcNowp6cvspg0jXznRTKDjm/E7AdgFBVeAPVMNcKGsHMA=="], - "fork-ts-checker-webpack-plugin/minimatch": ["minimatch@3.1.5", "", { "dependencies": { "brace-expansion": "^1.1.7" } }, "sha512-VgjWUsnnT6n+NUk6eZq77zeFdpW2LWDzP6zFGrCbHXiYNul5Dzqk2HHQ5uFH2DNW5Xbp8+jVzaeNt94ssEEl4w=="], - "glob/minimatch": ["minimatch@9.0.9", "", { "dependencies": { "brace-expansion": "^2.0.2" } }, "sha512-OBwBN9AL4dqmETlpS2zasx+vTeWclWzkblfZk7KTA5j3jeOONz/tRCnZomUyvNg83wL5Zv9Ss6HMJXAgL8R2Yg=="], "jest-worker/supports-color": ["supports-color@8.1.1", "", { "dependencies": { "has-flag": "^4.0.0" } }, "sha512-MpUEN2OodtUzxvKQl72cUF7RQ5EiHsGvSsVG0ia9c5RbWGL2CI4C7EpPS8UTBIplnlzZiNuV56w+FuNxy3ty2Q=="], @@ -1531,8 +1535,6 @@ "rimraf/glob": ["glob@7.2.3", "", { "dependencies": { "fs.realpath": "^1.0.0", "inflight": "^1.0.4", "inherits": "2", "minimatch": "^3.1.1", "once": "^1.3.0", "path-is-absolute": "^1.0.0" } }, "sha512-nFR0zLpU2YCaRxwoCJvL6UvCH2JFyFVIvwTLsIf21AuHlMskA1hhTdk+LlYJtOlYt9v6dvszD2BGRqBL+iQK9Q=="], - "schema-utils/ajv": ["ajv@6.15.0", "", { "dependencies": { "fast-deep-equal": "^3.1.1", "fast-json-stable-stringify": "^2.0.0", "json-schema-traverse": "^0.4.1", "uri-js": "^4.2.2" } }, "sha512-fgFx7Hfoq60ytK2c7DhnF8jIvzYgOMxfugjLOSMHjLIPgenqa7S7oaagATUq99mV6IYvN2tRmC0wnTYX6iPbMw=="], - "send/debug": ["debug@2.6.9", "", { "dependencies": { "ms": "2.0.0" } }, "sha512-bC7ElrdJaJnPbAP+1EotYvqZsb3ecl5wi6Bfi6BJTUcNowp6cvspg0jXznRTKDjm/E7AdgFBVeAPVMNcKGsHMA=="], "send/encodeurl": ["encodeurl@1.0.2", "", {}, "sha512-TPJXq8JqFaVYm2CWmPvnP2Iyo4ZSM7/QKcSmuMLDObfpH5fi7RUGmd/rTDf+rut/saiDiQEeVTNgAmJEdAOx0w=="], @@ -1551,8 +1553,6 @@ "tsyringe/tslib": ["tslib@1.14.1", "", {}, "sha512-Xni35NKzjgMrwevysHTCArtLDpPvye8zV/0E4EyYn43P7/7qvQwPh9BGkHewbMulVntbigmcT7rdX3BNo9wRJg=="], - "vitest/magic-string": ["magic-string@0.30.21", "", { "dependencies": { "@jridgewell/sourcemap-codec": "^1.5.5" } }, "sha512-vd2F4YUyEXKGcLHoq+TEyCjxueSeHnFxyyjNp80yg0XV4vUhnDer/lvvlqM/arB5bXQN5K2/3oinyCRyx8T2CQ=="], - "webpack/eslint-scope": ["eslint-scope@5.1.1", "", { "dependencies": { "esrecurse": "^4.3.0", "estraverse": "^4.1.1" } }, "sha512-2NxwbF/hZ0KpepYN0cNbo+FN6XoK7GaHlQhgx/hIZl6Va0bF45RQOOwhLIy8lQDbuCiadSLCBnH2CFYquit5bw=="], "@angular-devkit/core/ajv/json-schema-traverse": ["json-schema-traverse@1.0.0", "", {}, "sha512-NM8/P9n3XjXhIZn1lLhkFaACTOURQXjWhV4BA/RnOv8xvgqtqpAX9IO4mRQxSx1Rlo4tqzeqb0sOlruaOy3dug=="], @@ -1565,12 +1565,6 @@ "@angular-devkit/schematics-cli/inquirer/run-async": ["run-async@3.0.0", "", {}, "sha512-540WwVDOMxA6dN6We19EcT9sc3hkXPw5mzRNGM3FkdN/vtE9NFvj5lFAPNwUDmJjXidm3v7TC1cTE7t17Ulm1Q=="], - "@eslint/eslintrc/ajv/json-schema-traverse": ["json-schema-traverse@0.4.1", "", {}, "sha512-xbbCH5dCYU5T8LcEhhuh7HJ88HXuW3qsI3Y0zOZFKfZEHcpWiHU/Jxzk629Brsab/mMiHQti9wMP+845RPe3Vg=="], - - "@eslint/eslintrc/minimatch/brace-expansion": ["brace-expansion@1.1.18", "", { "dependencies": { "balanced-match": "^1.0.0", "concat-map": "0.0.1" } }, "sha512-Edep/X9fGqVNmzKBVsDYIOtD+z1tuezV70LBjdCst9Tqu76lsnvRiZ6oTic1n+/BIwX6QDGAO94PN4N2SADvtw=="], - - "@humanwhocodes/config-array/minimatch/brace-expansion": ["brace-expansion@1.1.18", "", { "dependencies": { "balanced-match": "^1.0.0", "concat-map": "0.0.1" } }, "sha512-Edep/X9fGqVNmzKBVsDYIOtD+z1tuezV70LBjdCst9Tqu76lsnvRiZ6oTic1n+/BIwX6QDGAO94PN4N2SADvtw=="], - "@isaacs/cliui/string-width/emoji-regex": ["emoji-regex@9.2.2", "", {}, "sha512-L18DaJsXSUk2+42pv8mLs5jJT2hqFkFE4j21wOmgbUqsZ2hL72NsUU785g9RXgo3s0ZNgVl42TiHp3ZtOv/Vyg=="], "@isaacs/cliui/strip-ansi/ansi-regex": ["ansi-regex@6.2.2", "", {}, "sha512-Bq3SmSpyFHaWjPk8If9yc6svM8c56dB5BAtW4Qbw5jHTwwXXcTLoRMkpDJp6VL0XzlWaCHTXrkFURMYmD0sLqg=="], @@ -1589,14 +1583,8 @@ "finalhandler/debug/ms": ["ms@2.0.0", "", {}, "sha512-Tpp60P6IUJDTuOq/5Z8cdskzJujfwqfOTkrwIwj7IRISpnkJnT6SyJ4PCPnGMoFjC9ddhal5KVIYtAt97ix05A=="], - "fork-ts-checker-webpack-plugin/minimatch/brace-expansion": ["brace-expansion@1.1.18", "", { "dependencies": { "balanced-match": "^1.0.0", "concat-map": "0.0.1" } }, "sha512-Edep/X9fGqVNmzKBVsDYIOtD+z1tuezV70LBjdCst9Tqu76lsnvRiZ6oTic1n+/BIwX6QDGAO94PN4N2SADvtw=="], - "glob/minimatch/brace-expansion": ["brace-expansion@2.1.4", "", { "dependencies": { "balanced-match": "^1.0.0" } }, "sha512-hGfVzPxthbf3+2yjg/RBs60cB0FhqBS/zvdV/4wn4/BmN0bNMMHPc4V/BbFieqf1TKAGGAHnY4eSjajCl0f2Xg=="], - "rimraf/glob/minimatch": ["minimatch@3.1.5", "", { "dependencies": { "brace-expansion": "^1.1.7" } }, "sha512-VgjWUsnnT6n+NUk6eZq77zeFdpW2LWDzP6zFGrCbHXiYNul5Dzqk2HHQ5uFH2DNW5Xbp8+jVzaeNt94ssEEl4w=="], - - "schema-utils/ajv/json-schema-traverse": ["json-schema-traverse@0.4.1", "", {}, "sha512-xbbCH5dCYU5T8LcEhhuh7HJ88HXuW3qsI3Y0zOZFKfZEHcpWiHU/Jxzk629Brsab/mMiHQti9wMP+845RPe3Vg=="], - "send/debug/ms": ["ms@2.0.0", "", {}, "sha512-Tpp60P6IUJDTuOq/5Z8cdskzJujfwqfOTkrwIwj7IRISpnkJnT6SyJ4PCPnGMoFjC9ddhal5KVIYtAt97ix05A=="], "terser-webpack-plugin/schema-utils/ajv": ["ajv@8.12.0", "", { "dependencies": { "fast-deep-equal": "^3.1.1", "json-schema-traverse": "^1.0.0", "require-from-string": "^2.0.2", "uri-js": "^4.2.2" } }, "sha512-sRu1kpcO9yLtYxBKvqfTeh9KzZEwO3STyX1HT+4CaDzC6HpTGYhIhPIzj9XuKU7KYDwnaeh5hcOwjy1QuJzBPA=="], @@ -1607,8 +1595,6 @@ "webpack/eslint-scope/estraverse": ["estraverse@4.3.0", "", {}, "sha512-39nnKffWz8xN1BU/2c79n9nB9HDzo0niYUqx6xyqUnyoAnQyyWpOTdZEeiCch8BBu515t4wp9ZmgVfVhn9EBpw=="], - "rimraf/glob/minimatch/brace-expansion": ["brace-expansion@1.1.18", "", { "dependencies": { "balanced-match": "^1.0.0", "concat-map": "0.0.1" } }, "sha512-Edep/X9fGqVNmzKBVsDYIOtD+z1tuezV70LBjdCst9Tqu76lsnvRiZ6oTic1n+/BIwX6QDGAO94PN4N2SADvtw=="], - "terser-webpack-plugin/schema-utils/ajv/json-schema-traverse": ["json-schema-traverse@1.0.0", "", {}, "sha512-NM8/P9n3XjXhIZn1lLhkFaACTOURQXjWhV4BA/RnOv8xvgqtqpAX9IO4mRQxSx1Rlo4tqzeqb0sOlruaOy3dug=="], "test-exclude/minimatch/brace-expansion/balanced-match": ["balanced-match@4.0.4", "", {}, "sha512-BLrgEcRTwX2o6gGxGOCNyMvGSp35YofuYzw9h1IMTRmKqttAZZVU67bdb9Pr2vUHA8+j3i2tJfjO6C6+4myGTA=="], diff --git a/docs/configuration.md b/docs/configuration.md index b1c2243a..79cb07ac 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -44,120 +44,153 @@ guarantees in tests and tools that construct modules without going through These have no default. The application will not start without them. -| Variable | Description | -| --- | --- | -| `DATABASE_URL` | PostgreSQL connection string used by Prisma. | -| `JWT_ACCESS_SECRET` | Signing secret for access tokens. At least 16 characters. | +| Variable | Description | +| -------------------- | ----------------------------------------------------------------------------------------------------------------- | +| `DATABASE_URL` | PostgreSQL connection string used by Prisma. | +| `JWT_ACCESS_SECRET` | Signing secret for access tokens. At least 16 characters. | | `JWT_REFRESH_SECRET` | Signing secret for refresh tokens. At least 16 characters. In production it must differ from `JWT_ACCESS_SECRET`. | -| `AI_PROVIDER_KEY` | API key for the AI provider. | +| `AI_PROVIDER_KEY` | API key for the AI provider. | ### Additional production requirements When `NODE_ENV=production`, values that are acceptable for local development are rejected: -| Variable | Rule | -| --- | --- | -| `ENCRYPTION_KEY` | Must be set explicitly. The built-in development default is publicly known and is rejected. | -| `JWT_REFRESH_SECRET` | Must differ from `JWT_ACCESS_SECRET`. | +| Variable | Rule | +| -------------------- | ------------------------------------------------------------------------------------------- | +| `ENCRYPTION_KEY` | Must be set explicitly. The built-in development default is publicly known and is rejected. | +| `JWT_REFRESH_SECRET` | Must differ from `JWT_ACCESS_SECRET`. | ## Optional variables ### Application -| Variable | Default | Description | -| --- | --- | --- | -| `NODE_ENV` | `development` | One of `development`, `test`, `production`. | -| `APP_NAME` | `astroid-api` | Service name. | -| `PORT` | `3000` | HTTP port. Positive integer. | -| `API_PREFIX` | `api/v1` | Global route prefix. | -| `LOG_LEVEL` | `info` | One of `fatal`, `error`, `warn`, `info`, `debug`, `trace`, `silent`. | -| `CORS_ORIGINS` | `*` | Comma-separated list of allowed origins. | +| Variable | Default | Description | +| -------------- | ------------- | -------------------------------------------------------------------- | +| `NODE_ENV` | `development` | One of `development`, `test`, `production`. | +| `APP_NAME` | `astroid-api` | Service name. | +| `PORT` | `3000` | HTTP port. Positive integer. | +| `API_PREFIX` | `api/v1` | Global route prefix. | +| `LOG_LEVEL` | `info` | One of `fatal`, `error`, `warn`, `info`, `debug`, `trace`, `silent`. | +| `CORS_ORIGINS` | `*` | Comma-separated list of allowed origins. | ### Database -| Variable | Default | Description | -| --- | --- | --- | -| `DATABASE_CONNECTION_LIMIT` | `10` | Prisma `connection_limit` for the API pool. | -| `DATABASE_WORKER_CONNECTION_LIMIT` | `3` | Connection limit for the background worker pool. | -| `DATABASE_POOL_TIMEOUT_MS` | `5000` | Time to wait for a free connection. `0` waits indefinitely. | -| `DATABASE_QUERY_TIMEOUT_MS` | `5000` | Client-side query timeout for the API pool. `0` disables it. | -| `DATABASE_STATEMENT_TIMEOUT_MS` | `10000` | Server-side `statement_timeout`. `0` disables it. | +| Variable | Default | Description | +| ---------------------------------- | ------- | --------------------------------------------------------------- | +| `DATABASE_CONNECTION_LIMIT` | `10` | Prisma `connection_limit` for the API pool. | +| `DATABASE_WORKER_CONNECTION_LIMIT` | `3` | Connection limit for the background worker pool. | +| `DATABASE_POOL_TIMEOUT_MS` | `5000` | Time to wait for a free connection. `0` waits indefinitely. | +| `DATABASE_QUERY_TIMEOUT_MS` | `5000` | Client-side query timeout for the API pool. `0` disables it. | +| `DATABASE_STATEMENT_TIMEOUT_MS` | `10000` | Server-side `statement_timeout`. `0` disables it. | | `DATABASE_WORKER_QUERY_TIMEOUT_MS` | `60000` | Client-side query timeout for the worker pool. `0` disables it. | +| `DATABASE_SLOW_QUERY_THRESHOLD_MS` | `1000` | Queries slower than this are logged as slow queries. | +| `DATABASE_CONNECT_RETRY_ATTEMPTS` | `5` | Connection attempts before giving up on startup. | +| `DATABASE_CONNECT_RETRY_DELAY_MS` | `1000` | Delay between connection retry attempts. | +| `DATABASE_SLOW_QUERY_THRESHOLD_MS` | `1000` | Queries slower than this are logged as slow queries. | +| `DATABASE_CONNECT_RETRY_ATTEMPTS` | `5` | Connection attempts before giving up on startup. | +| `DATABASE_CONNECT_RETRY_DELAY_MS` | `1000` | Delay between connection retry attempts. | +| `DATABASE_MIGRATION_CHECK_ENABLED` | `true` | Runs a migration status check during bootstrap before the app accepts traffic. | +| `DATABASE_MIGRATION_CHECK_MODE` | `halt` | `halt` exits the process when migrations are pending/failed; `warn` logs and continues. | ### Redis -| Variable | Default | Description | -| --- | --- | --- | -| `REDIS_HOST` | `localhost` | Redis host. | -| `REDIS_PORT` | `6379` | Redis port. Positive integer. | -| `REDIS_PASSWORD` | _(empty)_ | Redis password. | -| `REDIS_DB` | `0` | Redis database index. | +| Variable | Default | Description | +| ---------------- | ----------- | ----------------------------- | +| `REDIS_HOST` | `localhost` | Redis host. | +| `REDIS_PORT` | `6379` | Redis port. Positive integer. | +| `REDIS_PASSWORD` | _(empty)_ | Redis password. | +| `REDIS_DB` | `0` | Redis database index. | ### Authentication -| Variable | Default | Description | -| --- | --- | --- | -| `JWT_ACCESS_TTL` | `900` | Access-token lifetime in seconds. | -| `JWT_REFRESH_TTL` | `1209600` | Refresh-token lifetime in seconds. | -| `PASSKEY_RP_ID` | `localhost` | WebAuthn relying-party ID. | -| `PASSKEY_RP_NAME` | `Astroid` | WebAuthn relying-party display name. | -| `PASSKEY_ORIGIN` | `http://localhost:3001` | Expected WebAuthn origin. | +| Variable | Default | Description | +| ----------------- | ----------------------- | ------------------------------------ | +| `JWT_ACCESS_TTL` | `900` | Access-token lifetime in seconds. | +| `JWT_REFRESH_TTL` | `1209600` | Refresh-token lifetime in seconds. | +| `PASSKEY_RP_ID` | `localhost` | WebAuthn relying-party ID. | +| `PASSKEY_RP_NAME` | `Astroid` | WebAuthn relying-party display name. | +| `PASSKEY_ORIGIN` | `http://localhost:3001` | Expected WebAuthn origin. | ### Stellar -| Variable | Default | Description | -| --- | --- | --- | -| `STELLAR_NETWORK` | `testnet` | One of `testnet`, `public`, `futurenet`. | -| `STELLAR_HORIZON_URL` | `https://horizon-testnet.stellar.org` | Horizon endpoint. | -| `STELLAR_SOROBAN_RPC_URL` | `https://soroban-testnet.stellar.org` | Soroban RPC endpoint. | -| `STELLAR_REGISTRY_CONTRACT_ID` | _(empty)_ | Agent registry contract ID. | -| `STELLAR_USE_MOCK` | `true` | `true` or `false`. Use the mock Stellar client. | +| Variable | Default | Description | +| ------------------------------ | ------------------------------------- | ----------------------------------------------- | +| `STELLAR_NETWORK` | `testnet` | One of `testnet`, `public`, `futurenet`. | +| `STELLAR_HORIZON_URL` | `https://horizon-testnet.stellar.org` | Horizon endpoint. | +| `STELLAR_SOROBAN_RPC_URL` | `https://soroban-testnet.stellar.org` | Soroban RPC endpoint. | +| `STELLAR_REGISTRY_CONTRACT_ID` | _(empty)_ | Agent registry contract ID. | +| `STELLAR_USE_MOCK` | `true` | `true` or `false`. Use the mock Stellar client. | ### Storage (S3-compatible) -| Variable | Default | Description | -| --- | --- | --- | -| `STORAGE_ENDPOINT` | `http://localhost:9000` | Object storage endpoint. | -| `STORAGE_REGION` | `us-east-1` | Storage region. | -| `STORAGE_BUCKET` | `astroid` | Bucket name. | -| `STORAGE_ACCESS_KEY` | `astroid` | Access key. | -| `STORAGE_SECRET_KEY` | `astroid-secret` | Secret key. | +| Variable | Default | Description | +| -------------------- | ----------------------- | ------------------------ | +| `STORAGE_ENDPOINT` | `http://localhost:9000` | Object storage endpoint. | +| `STORAGE_REGION` | `us-east-1` | Storage region. | +| `STORAGE_BUCKET` | `astroid` | Bucket name. | +| `STORAGE_ACCESS_KEY` | `astroid` | Access key. | +| `STORAGE_SECRET_KEY` | `astroid-secret` | Secret key. | ### Queues (BullMQ) -| Variable | Default | Description | -| --- | --- | --- | -| `QUEUE_PREFIX` | `astroid` | Key prefix for BullMQ queues. | -| `QUEUE_CONCURRENCY` | `5` | Default worker concurrency. | +| Variable | Default | Description | +| ------------------- | --------- | ----------------------------- | +| `QUEUE_PREFIX` | `astroid` | Key prefix for BullMQ queues. | +| `QUEUE_CONCURRENCY` | `5` | Default worker concurrency. | ### Rate limiting +| Variable | Default | Description | +| ---------------------------------- | ------- | ----------------------------------------------------------------------------------- | +| `THROTTLE_AUTH_LIMIT` | `10` | Requests per window on the `auth` tier. | +| `THROTTLE_API_LIMIT` | `120` | Requests per window on the `api` tier. | +| `THROTTLE_TTL` | `60` | Throttler window in seconds. | +| `THROTTLE_WEBHOOK_LIMIT` | `30` | Requests per window on the `webhook` tier. | +| `THROTTLE_API_BURST` | `10` | Burst allowance on top of `THROTTLE_API_LIMIT`. | +| `THROTTLE_AUTH_BURST` | `3` | Burst allowance on top of `THROTTLE_AUTH_LIMIT`. | +| `THROTTLE_WEBHOOK_BURST` | `5` | Burst allowance on top of `THROTTLE_WEBHOOK_LIMIT`. | +| `RATE_LIMIT_WINDOW_SECONDS` | `60` | Sliding-window size for the Redis rate-limiter guard. | +| `RATE_LIMIT_MAX_REQUESTS` | `120` | Requests allowed per client per sliding window. | +| `PUBLIC_RATE_LIMIT_ENABLED` | `true` | Enables the IP-based limiter for unauthenticated (`@Public()`) routes. | +| `PUBLIC_RATE_LIMIT_MAX_REQUESTS` | `60` | Requests allowed per client IP per sliding window on public routes. | +| `PUBLIC_RATE_LIMIT_WINDOW_SECONDS` | `60` | Sliding-window size for the public-route rate limiter. | +| `PUBLIC_RATE_LIMIT_TRUST_PROXY` | `false` | Reads client IP from `X-Forwarded-For`. Only enable behind a trusted reverse proxy. | | Variable | Default | Description | | --- | --- | --- | | `THROTTLE_AUTH_LIMIT` | `10` | Requests per window on the `auth` tier. | | `THROTTLE_API_LIMIT` | `120` | Requests per window on the `api` tier. | +| `THROTTLE_AGENT_LIMIT` | `300` | Requests per window for autonomous-agent traffic. | +| `THROTTLE_WEBHOOK_LIMIT` | `30` | Requests per window on the `webhook` tier. | | `THROTTLE_TTL` | `60` | Throttler window in seconds. | +| `THROTTLE_API_BURST` | `10` | Short-term (1s) burst allowance on the `api` tier. `0` disables burst enforcement. | +| `THROTTLE_AUTH_BURST` | `3` | Short-term (1s) burst allowance on the `auth` tier. `0` disables burst enforcement. | +| `THROTTLE_WEBHOOK_BURST` | `5` | Short-term (1s) burst allowance on the `webhook` tier. `0` disables burst enforcement. | | `RATE_LIMIT_WINDOW_SECONDS` | `60` | Sliding-window size for the Redis rate-limiter guard. | | `RATE_LIMIT_MAX_REQUESTS` | `120` | Requests allowed per client per sliding window. | +| `PUBLIC_RATE_LIMIT_ENABLED` | `true` | Enables the IP-based limiter for unauthenticated (`@Public()`) routes. | +| `PUBLIC_RATE_LIMIT_MAX_REQUESTS` | `60` | Requests allowed per client IP per sliding window on public routes. | +| `PUBLIC_RATE_LIMIT_WINDOW_SECONDS` | `60` | Sliding-window size for the public-route rate limiter. | +| `PUBLIC_RATE_LIMIT_TRUST_PROXY` | `false` | Reads client IP from `X-Forwarded-For`. Only enable behind a trusted reverse proxy. | +| `PUBLIC_RATE_LIMIT_CLIENT_IDENTIFIERS` | _(empty)_ | Comma-separated client identifiers to add to the IP bucket, currently `apiKey`. | ### Metrics -| Variable | Default | Description | -| --- | --- | --- | +| Variable | Default | Description | +| --------------------- | ---------------------------- | ------------------------------------------------------------- | | `METRICS_ALLOWED_IPS` | loopback and RFC 1918 ranges | Comma-separated CIDR ranges allowed to scrape `GET /metrics`. | ### AI provider -| Variable | Default | Description | -| --- | --- | --- | -| `AI_PROVIDER` | `nvidia` | Provider name. | +| Variable | Default | Description | +| ------------- | ------------------------------------- | ---------------------- | +| `AI_PROVIDER` | `nvidia` | Provider name. | | `AI_BASE_URL` | `https://integrate.api.nvidia.com/v1` | Provider API base URL. | -| `AI_MODEL` | `meta/llama-3.1-70b-instruct` | Model identifier. | +| `AI_MODEL` | `meta/llama-3.1-70b-instruct` | Model identifier. | ### Encryption -| Variable | Default | Description | -| --- | --- | --- | -| `ENCRYPTION_KEY` | development-only key | 32-byte key: 64 hex characters, 32 raw bytes, or base64 of 32 bytes. Required in production. | -| `ENCRYPTION_ALGORITHM` | `aes-256-gcm` | Cipher algorithm. | +| Variable | Default | Description | +| ---------------------- | -------------------- | -------------------------------------------------------------------------------------------- | +| `ENCRYPTION_KEY` | development-only key | 32-byte key: 64 hex characters, 32 raw bytes, or base64 of 32 bytes. Required in production. | +| `ENCRYPTION_ALGORITHM` | `aes-256-gcm` | Cipher algorithm. | diff --git a/package-lock.json b/package-lock.json index b72149a4..a5389c49 100644 --- a/package-lock.json +++ b/package-lock.json @@ -5911,7 +5911,6 @@ "version": "2.3.3", "resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.3.tgz", "integrity": "sha512-5xoDfX+fL7faATnagmWPpbFtwh/R77WmMMqqHGS65C3vvB0YHrgF+B1YmZ3441tMj5n63k0212XNoJwzlhffQw==", - "dev": true, "hasInstallScript": true, "license": "MIT", "optional": true, diff --git a/prisma/migrations/20260928120000_add_notifications_user_created_at_index/migration.sql b/prisma/migrations/20260928120001_add_notifications_user_created_at_index/migration.sql similarity index 100% rename from prisma/migrations/20260928120000_add_notifications_user_created_at_index/migration.sql rename to prisma/migrations/20260928120001_add_notifications_user_created_at_index/migration.sql diff --git a/prisma/migrations/20260928120002_add_audit_log_source_event_id/migration.sql b/prisma/migrations/20260928120002_add_audit_log_source_event_id/migration.sql new file mode 100644 index 00000000..e2b5fe0c --- /dev/null +++ b/prisma/migrations/20260928120002_add_audit_log_source_event_id/migration.sql @@ -0,0 +1,3 @@ +ALTER TABLE "audit_logs" ADD COLUMN "sourceEventId" TEXT; + +CREATE UNIQUE INDEX "audit_logs_sourceEventId_key" ON "audit_logs"("sourceEventId"); \ No newline at end of file diff --git a/prisma/schema.prisma b/prisma/schema.prisma index 47cf8cca..3c9b26ce 100644 --- a/prisma/schema.prisma +++ b/prisma/schema.prisma @@ -499,6 +499,7 @@ model AuditLog { // `x-correlation-id`). Not part of the hash chain — it's correlation // metadata, not tamper-evident content. requestId String? + sourceEventId String? @unique previousHash String? // SHA-256 hash of the preceding audit log entry hash String? // SHA-256 hash of this entry (links to previous) createdAt DateTime @default(now()) diff --git a/scripts/verify-migrations.sh b/scripts/verify-migrations.sh index 50b92c50..c9020a99 100644 --- a/scripts/verify-migrations.sh +++ b/scripts/verify-migrations.sh @@ -12,7 +12,8 @@ MIGRATIONS_DIR="prisma/migrations" if [ -d "$MIGRATIONS_DIR" ]; then echo "Checking migration directories under $MIGRATIONS_DIR..." - declare -A timestamps + timestamps=() + timestamp_dirs=() migration_count=0 for dir in "$MIGRATIONS_DIR"/*/; @@ -44,11 +45,14 @@ if [ -d "$MIGRATIONS_DIR" ]; then # Check 3: Extract timestamp prefix (expects YYYYMMDDHHMMSS or similar leading numeric prefix) if [[ "$dirname" =~ ^([0-9]{14}) ]]; then ts="${BASH_REMATCH[1]}" - if [ -n "${timestamps[$ts]:-}" ]; then - echo "Error: Conflicting migration timestamps detected: '$dirname' shares timestamp prefix with '${timestamps[$ts]}'" - exit 1 - fi - timestamps["$ts"]="$dirname" + for index in "${!timestamps[@]}"; do + if [ "${timestamps[$index]}" = "$ts" ]; then + echo "Error: Conflicting migration timestamps detected: '$dirname' shares timestamp prefix with '${timestamp_dirs[$index]}'" + exit 1 + fi + done + timestamps+=("$ts") + timestamp_dirs+=("$dirname") else echo "Warning: Migration directory '$dirname' does not start with a standard 14-digit timestamp (YYYYMMDDHHMMSS)" fi diff --git a/src/app.module.ts b/src/app.module.ts index 97920d41..ba361c0d 100644 --- a/src/app.module.ts +++ b/src/app.module.ts @@ -52,6 +52,8 @@ import { DeadLetterModule } from './modules/dead-letter/dead-letter.module'; import { AgentTraceInterceptor } from './common/interceptors/agent-trace.interceptor'; import { RequestContextInterceptor } from './common/interceptors/request-context.interceptor'; import { AuditLogInterceptor } from './common/interceptors/audit-log.interceptor'; +import { MetricsInterceptor } from './common/interceptors/metrics.interceptor'; +import { RequestIdInterceptor } from './common/interceptors/request-id.interceptor'; /** * Root application module. Wires the global infrastructure (config, logging, @@ -67,6 +69,7 @@ import { AuditLogInterceptor } from './common/interceptors/audit-log.interceptor * - ThrottlerGuard : per-organization / per-IP rate limiting, shared via Redis * - ResponseInterceptor: wraps every result in the success envelope * - AuditLogInterceptor: persists masked mutation requests to the audit trail + * - MetricsInterceptor: records Prometheus metrics for HTTP requests * - AllExceptionsFilter: converts every error into the error envelope */ @Module({ @@ -84,12 +87,14 @@ import { AuditLogInterceptor } from './common/interceptors/audit-log.interceptor autoLogging: false, }, }), - // Two rate-limit tiers, both driven by THROTTLE_* env vars (see - // config/throttler.config.ts). Every route is subject to both named + // Three rate-limit tiers, all driven by THROTTLE_* env vars (see + // config/throttler.config.ts). Every route is subject to all named // throttlers, but AstroidThrottlerGuard enforces only the one matching the // route's @ThrottleTierDecorator tier ('api' default, 'auth' for the - // sensitive auth endpoints). Counters live in Redis so every replica behind - // the load balancer enforces the same budget. + // sensitive auth endpoints), and AgentThrottlerGuard (applied to the + // agent-facing controllers) enforces the 'agent' tier keyed by acting agent. + // Counters live in Redis so every replica behind the load balancer enforces + // the same budget. ThrottlerModule.forRootAsync({ imports: [LocksModule], inject: [ConfigService, REDIS_CLIENT], @@ -135,11 +140,13 @@ import { AuditLogInterceptor } from './common/interceptors/audit-log.interceptor { provide: APP_GUARD, useClass: ScopesGuard }, { provide: APP_GUARD, useClass: AstroidThrottlerGuard }, AgentPolicyGuard, + { provide: APP_INTERCEPTOR, useClass: RequestIdInterceptor }, { provide: APP_INTERCEPTOR, useClass: RequestContextInterceptor }, { provide: APP_INTERCEPTOR, useClass: AgentTraceInterceptor }, { provide: APP_INTERCEPTOR, useClass: AuditLogInterceptor }, { provide: APP_INTERCEPTOR, useClass: ResponseInterceptor }, { provide: APP_INTERCEPTOR, useClass: AuditInterceptor }, + { provide: APP_INTERCEPTOR, useClass: MetricsInterceptor }, { provide: APP_FILTER, useClass: AllExceptionsFilter }, ], }) diff --git a/src/common/cache/cache.service.spec.ts b/src/common/cache/cache.service.spec.ts new file mode 100644 index 00000000..caf0c20b --- /dev/null +++ b/src/common/cache/cache.service.spec.ts @@ -0,0 +1,106 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { Logger } from '@nestjs/common'; +import { CacheService } from './cache.service'; + +describe('CacheService', () => { + let redis: { + get: ReturnType; + set: ReturnType; + del: ReturnType; + scan: ReturnType; + }; + + beforeEach(() => { + vi.spyOn(Logger.prototype, 'warn').mockImplementation(() => undefined); + redis = { + get: vi.fn().mockResolvedValue(null), + set: vi.fn().mockResolvedValue('OK'), + del: vi.fn().mockResolvedValue(1), + scan: vi.fn().mockResolvedValue(['0', []]), + }; + }); + + it('round-trips a value through the namespaced key with a TTL', async () => { + const cache = new CacheService(redis as never); + redis.get.mockResolvedValue( + JSON.stringify({ value: { revoked: false }, cachedAt: 1, expiresAt: 2 }), + ); + + await cache.set('ns', 'k', { revoked: false }, 30); + + expect(redis.set).toHaveBeenCalledWith( + 'ns:k', + expect.stringContaining('"revoked":false'), + 'EX', + 30, + ); + await expect(cache.get<{ revoked: boolean }>('ns', 'k')).resolves.toEqual({ revoked: false }); + }); + + it('returns null on miss and never throws when Redis fails', async () => { + const cache = new CacheService(redis as never); + + await expect(cache.get('ns', 'missing')).resolves.toBeNull(); + + redis.get.mockRejectedValue(new Error('READONLY')); + await expect(cache.get('ns', 'k')).resolves.toBeNull(); + expect(Logger.prototype.warn).toHaveBeenCalled(); + }); + + it('set swallows Redis failures instead of breaking the request path', async () => { + const cache = new CacheService(redis as never); + redis.set.mockRejectedValue(new Error('connection refused')); + + await expect(cache.set('ns', 'k', 'v', 30)).resolves.toBeUndefined(); + }); + + it('getWithMeta rejects entries older than the requested max age', async () => { + const cache = new CacheService(redis as never); + const fresh = { value: 'fresh', cachedAt: Date.now() - 1_000, expiresAt: Date.now() + 60_000 }; + const stale = { value: 'stale', cachedAt: Date.now() - 10_000, expiresAt: Date.now() + 60_000 }; + redis.get.mockResolvedValue(JSON.stringify(fresh)); + + await expect(cache.getWithMeta('ns', 'k', 5_000)).resolves.toMatchObject({ value: 'fresh' }); + + redis.get.mockResolvedValue(JSON.stringify(stale)); + await expect(cache.getWithMeta('ns', 'k', 5_000)).resolves.toBeNull(); + }); + + it('getWithMeta returns null when the entry itself has expired', async () => { + const cache = new CacheService(redis as never); + redis.get.mockResolvedValue( + JSON.stringify({ value: 'old', cachedAt: Date.now() - 9_000, expiresAt: Date.now() - 1_000 }), + ); + + await expect(cache.getWithMeta('ns', 'k')).resolves.toBeNull(); + }); + + it('del removes the namespaced key', async () => { + const cache = new CacheService(redis as never); + + await cache.del('ns', 'k'); + + expect(redis.del).toHaveBeenCalledWith('ns:k'); + }); + + it('delByPrefix deletes every matching key in batches', async () => { + const cache = new CacheService(redis as never); + redis.scan + .mockResolvedValueOnce(['123', ['ns:a', 'ns:b']]) + .mockResolvedValueOnce(['0', ['ns:c']]); + + await cache.delByPrefix('ns', ''); + + expect(redis.del).toHaveBeenCalledWith('ns:a', 'ns:b'); + expect(redis.del).toHaveBeenCalledWith('ns:c'); + }); + + it('is a no-op when no Redis client is available', async () => { + const cache = new CacheService(null); + + await expect(cache.get('ns', 'k')).resolves.toBeNull(); + await expect(cache.set('ns', 'k', 'v', 30)).resolves.toBeUndefined(); + await expect(cache.del('ns', 'k')).resolves.toBeUndefined(); + expect(redis.set).not.toHaveBeenCalled(); + }); +}); diff --git a/src/common/cache/cache.service.ts b/src/common/cache/cache.service.ts new file mode 100644 index 00000000..299a8104 --- /dev/null +++ b/src/common/cache/cache.service.ts @@ -0,0 +1,142 @@ +import { Inject, Injectable, Logger, Optional } from '@nestjs/common'; +import { Redis } from 'ioredis'; +import { REDIS_CLIENT } from '../locks/locks.constants'; + +/** + * A single cached value plus the metadata needed to honour revocation + * semantics. `cachedAt`/`expiresAt` are stored *inside* the payload (not left + * to Redis' TTL alone) so a {@link CacheService} embedded in another service + * can decide whether an entry is still trustworthy, e.g. when the caller has a + * stricter staleness requirement than the configured TTL. + */ +export interface CacheEntry { + value: T; + /** Epoch ms at which the entry was written to the cache. */ + cachedAt: number; + /** Epoch ms at which the entry becomes stale and must be re-validated. */ + expiresAt: number; +} + +/** + * Tiny get/set/delete cache over the shared Redis client (the same + * `REDIS_CLIENT` used by the locks and throttler infrastructure) with a + * per-process `Map` fallback so callers keep a consistent API when Redis is + * unreachable. Values are JSON-serialised and stored under a namespaced key + * with a TTL, so stale entries never outlive their usefulness. + */ +@Injectable() +export class CacheService { + private readonly logger = new Logger(CacheService.name); + + /** + * Absent in some unit-test contexts (and when Redis is not configured at + * all); the cache then behaves as a no-op and every lookup misses, which is + * the safe direction: callers fall through to their source of truth. + */ + constructor(@Optional() @Inject(REDIS_CLIENT) private readonly redis: Redis | null) {} + + /** Reads a namespaced entry. Returns null on miss, expiry or Redis failure. */ + async get(namespace: string, key: string): Promise { + if (!this.redis) { + return null; + } + try { + const raw = await this.redis.get(`${namespace}:${key}`); + if (!raw) { + return null; + } + return JSON.parse(raw).value as T; + } catch (error: unknown) { + // A cache must never break the request path: log and treat as a miss. + this.logger.warn(`Cache get failed for ${namespace}:${key}: ${(error as Error).message}`); + return null; + } + } + + /** + * Reads an entry only if it has not aged past `maxAgeMs` (used by callers + * whose revocation requirements are stricter than the cache TTL). Falls back + * to {@link get} semantics (metadata still checked) for plain entries. + */ + async getWithMeta( + namespace: string, + key: string, + maxAgeMs?: number, + ): Promise<{ value: T; entry: CacheEntry } | null> { + if (!this.redis) { + return null; + } + try { + const raw = await this.redis.get(`${namespace}:${key}`); + if (!raw) { + return null; + } + const entry = JSON.parse(raw) as CacheEntry; + if (typeof entry?.expiresAt !== 'number' || entry.expiresAt <= Date.now()) { + return null; + } + if (maxAgeMs !== undefined && Date.now() - entry.cachedAt > maxAgeMs) { + return null; + } + return { value: entry.value, entry }; + } catch (error: unknown) { + this.logger.warn(`Cache get failed for ${namespace}:${key}: ${(error as Error).message}`); + return null; + } + } + + /** Writes a value under `namespace:key` with the given TTL in seconds. */ + async set(namespace: string, key: string, value: T, ttlSeconds: number): Promise { + if (!this.redis) { + return; + } + const entry: CacheEntry = { + value, + cachedAt: Date.now(), + expiresAt: Date.now() + ttlSeconds * 1000, + }; + try { + await this.redis.set(`${namespace}:${key}`, JSON.stringify(entry), 'EX', Math.max(1, ttlSeconds)); + } catch (error: unknown) { + this.logger.warn(`Cache set failed for ${namespace}:${key}: ${(error as Error).message}`); + } + } + + /** Deletes a single entry. Missing keys are not an error. */ + async del(namespace: string, key: string): Promise { + if (!this.redis) { + return; + } + try { + await this.redis.del(`${namespace}:${key}`); + } catch (error: unknown) { + this.logger.warn(`Cache delete failed for ${namespace}:${key}: ${(error as Error).message}`); + } + } + + /** + * Deletes every entry whose key starts with the given prefix inside a + * namespace. Used by invalidation hooks that must clear several related + * entries (e.g. all verification results for one session). Implemented with + * `SCAN` + batched `DEL` so it is safe on large or clustered Redis instaces + * without `KEYS`. + */ + async delByPrefix(namespace: string, prefix: string): Promise { + if (!this.redis) { + return; + } + const pattern = `${namespace}:${prefix}*`; + let cursor = '0'; + try { + do { + const [next, keys] = await this.redis.scan(cursor, 'MATCH', pattern, 'COUNT', 100); + cursor = next; + if (keys.length > 0) { + await this.redis.del(...keys); + } + } while (cursor !== '0'); + } catch (error: unknown) { + this.logger.warn(`Cache prefix delete failed for ${pattern}: ${(error as Error).message}`); + } + } +} diff --git a/src/common/constants/headers.ts b/src/common/constants/headers.ts index bdf7310f..b3854fec 100644 --- a/src/common/constants/headers.ts +++ b/src/common/constants/headers.ts @@ -7,4 +7,7 @@ export const WEBHOOK_TIMESTAMP_HEADER = 'x-astroid-timestamp'; export const WEBHOOK_EVENT_ID_HEADER = 'x-astroid-event-id'; export const WEBHOOK_DELIVERY_HEADER = 'x-astroid-delivery'; export const WEBHOOK_EVENT_HEADER = 'x-astroid-event'; +export const WEBHOOK_SIGNATURE_VERSION_HEADER = 'x-astroid-signature-version'; export const IDEMPOTENCY_KEY_HEADER = 'idempotency-key'; +/** Total number of rows matching a list request, set on every paginated response. */ +export const TOTAL_COUNT_HEADER = 'x-total-count'; diff --git a/src/common/decorators/api-pagination-query.decorator.ts b/src/common/decorators/api-pagination-query.decorator.ts new file mode 100644 index 00000000..96d3af60 --- /dev/null +++ b/src/common/decorators/api-pagination-query.decorator.ts @@ -0,0 +1,44 @@ +import { applyDecorators } from '@nestjs/common'; +import { ApiQuery, ApiResponse } from '@nestjs/swagger'; +import { DEFAULT_PAGE_LIMIT, MAX_PAGE_LIMIT } from '../helpers/pagination'; + +/** + * Documents the standard list query parameters parsed by + * `paginationQuerySchema` (`offset`/`page`, `limit`, `sort`, `order`), the + * `X-Total-Count` response header, and the 400 returned for invalid bounds. + */ +export function ApiPaginationQuery() { + return applyDecorators( + ApiQuery({ + name: 'offset', + required: false, + type: Number, + description: 'Zero-based number of rows to skip (default: 0). Mutually exclusive with page.', + }), + ApiQuery({ + name: 'page', + required: false, + type: Number, + description: '1-based page number, an alternative to offset (default: 1).', + }), + ApiQuery({ + name: 'limit', + required: false, + type: Number, + description: `Items per page (default: ${DEFAULT_PAGE_LIMIT}, max: ${MAX_PAGE_LIMIT}).`, + }), + ApiQuery({ name: 'sort', required: false, type: String, description: 'Sort field (default: createdAt).' }), + ApiQuery({ name: 'order', required: false, enum: ['asc', 'desc'], description: 'Sort direction (default: desc).' }), + ApiResponse({ + status: 200, + description: 'Paginated list. `meta` carries offset, page, limit, total, totalPages, hasNext and hasPrev.', + headers: { + 'X-Total-Count': { description: 'Total number of matching rows', schema: { type: 'integer' } }, + }, + }), + ApiResponse({ + status: 400, + description: 'Invalid pagination parameters (negative, non-integer, limit above max, or both offset and page).', + }), + ); +} diff --git a/src/common/decorators/audit-log.decorator.ts b/src/common/decorators/audit-log.decorator.ts new file mode 100644 index 00000000..91f170c1 --- /dev/null +++ b/src/common/decorators/audit-log.decorator.ts @@ -0,0 +1,42 @@ +import { SetMetadata } from '@nestjs/common'; + +/** Metadata key read by `AuditLogInterceptor` through the Nest Reflector. */ +export const AUDIT_LOG_KEY = 'astroid:auditLog'; + +/** + * Per-route audit metadata. Everything is optional: a bare `@AuditLog()` is the + * common case and lets the interceptor derive the action from the HTTP method + * and the entity from the controller name. + */ +export interface AuditLogOptions { + /** + * Semantic action name stored on the audit row (e.g. `POLICY_OVERRIDE`). + * Defaults to the HTTP method (`POST`, `PATCH`, …). + */ + action?: string; + /** + * Domain entity stored on the audit row (e.g. `Wallet`). Defaults to the + * controller name with the `Controller` suffix stripped. + */ + entity?: string; +} + +/** + * Marks a route (handler or whole controller) as audited. + * + * Only decorated routes are persisted by `AuditLogInterceptor`, so read-only + * traffic and uninteresting mutations never pay the cost of a database write. + * Sensitive, state-changing endpoints — budget adjustments, policy overrides, + * key rotations — should always carry this decorator. + * + * Combining it with `@SkipAudit()` opts a route back out, which is useful when a + * whole controller is decorated but one handler must not be logged. + * + * @example + * ```ts + * @Post('budgets/:id/adjust') + * @AuditLog({ action: 'BUDGET_ADJUSTED', entity: 'Budget' }) + * async adjustBudget(@Body() dto: AdjustBudgetDto) { ... } + * ``` + */ +export const AuditLog = (options: AuditLogOptions = {}) => SetMetadata(AUDIT_LOG_KEY, options); diff --git a/src/common/decorators/public-rate-limit.decorator.ts b/src/common/decorators/public-rate-limit.decorator.ts new file mode 100644 index 00000000..db7625dd --- /dev/null +++ b/src/common/decorators/public-rate-limit.decorator.ts @@ -0,0 +1,18 @@ +import { SetMetadata } from '@nestjs/common'; + +export const PUBLIC_RATE_LIMIT_RULE_KEY = 'astroid:publicRateLimitRule'; + +/** A `@PublicRateLimit()` rule: at most `max` requests per sliding `windowSeconds`. */ +export interface PublicRateLimitRule { + max: number; + windowSeconds: number; +} + +/** + * Overrides the global IP rate-limit settings for a public route (or whole + * controller) with a dedicated budget. Applies to routes covered by the + * `PublicRateLimitGuard` — i.e. `@Public()` routes and `//public/*` — + * and keeps the standard `X-RateLimit-*` header contract. + */ +export const PublicRateLimit = (max: number, windowSeconds: number) => + SetMetadata(PUBLIC_RATE_LIMIT_RULE_KEY, { max, windowSeconds } satisfies PublicRateLimitRule); diff --git a/src/common/decorators/throttle-tier.decorator.ts b/src/common/decorators/throttle-tier.decorator.ts index 4ce7f7ce..aeee76dc 100644 --- a/src/common/decorators/throttle-tier.decorator.ts +++ b/src/common/decorators/throttle-tier.decorator.ts @@ -2,10 +2,17 @@ import { SetMetadata } from '@nestjs/common'; export const THROTTLE_TIER_KEY = 'astroid:throttleTier'; -export type ThrottleTier = 'auth' | 'api'; +/** + * The available rate-limit tiers: + * - `api` — default for all authenticated API routes (THROTTLE_API_LIMIT/min) + * - `auth` — sensitive credential / session routes (THROTTLE_AUTH_LIMIT/min) + * - `agent` — high-frequency autonomous-agent routes (THROTTLE_AGENT_LIMIT/min) + * - `webhook` — outbound webhook management routes (THROTTLE_WEBHOOK_LIMIT/min) + */ +export type ThrottleTier = 'auth' | 'api' | 'agent' | 'webhook'; /** - * Selects the rate-limit tier for a route. `auth` = 10/min, `api` = 120/min. + * Selects the rate-limit tier for a route. * Defaults to `api` when unset. Consumed by the AstroidThrottlerGuard. */ export const ThrottleTierDecorator = (tier: ThrottleTier) => diff --git a/src/common/filters/global-exception.filter.spec.ts b/src/common/filters/global-exception.filter.spec.ts new file mode 100644 index 00000000..21222d49 --- /dev/null +++ b/src/common/filters/global-exception.filter.spec.ts @@ -0,0 +1,501 @@ +/** + * Unit tests for GlobalExceptionFilter. + * + * Verifies that every error path is transformed into the uniform RFC 9457 + * problem details envelope: + * { type, title, status, detail, instance, code, requestId, details? } + * + * Test surface: + * • Prisma database errors (P2002 → 409, P2025 → 404, others → 400) + * • Validation failures (ZodValidationException, ValidationException, + * class-validator BadRequestException arrays) + * • Auth / authz errors (401 Unauthorized, 403 Forbidden, TOKEN_EXPIRED) + * • Rate-limiting (ThrottlerException → 429) + * • Generic HTTP exceptions (405 → about:blank) + * • Unknown server faults (500, no internals leaked) + * • Request-id propagation (header → context → freshly generated UUID v7) + */ +import { describe, expect, it, vi, beforeEach } from 'vitest'; +import { + ArgumentsHost, + BadRequestException, + ForbiddenException, + HttpException, + Logger, + MethodNotAllowedException, + UnauthorizedException, +} from '@nestjs/common'; +import { ThrottlerException } from '@nestjs/throttler'; +import { Prisma } from '@prisma/client'; + +import { GlobalExceptionFilter } from './global-exception.filter'; +import { ErrorCode } from '../constants/error-codes'; +import { DomainException, ValidationException } from '../exceptions/domain.exception'; +import { RequestContext } from '../context/request-context'; +import { ProblemDetails } from '../interfaces/api-response.interface'; +import { ZodValidationException } from '../pipes/zod-validation.pipe'; + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +type MockResponse = { + status: ReturnType; + json: ReturnType; + setHeader: ReturnType; +}; + +function buildHost(request: Record = {}) { + const response: MockResponse = { + status: vi.fn().mockReturnThis(), + json: vi.fn().mockReturnThis(), + setHeader: vi.fn().mockReturnThis(), + }; + const req = { + method: 'POST', + url: '/api/v1/transactions', + originalUrl: '/api/v1/transactions', + headers: {}, + ...request, + }; + const host = { + switchToHttp: () => ({ getResponse: () => response, getRequest: () => req }), + } as unknown as ArgumentsHost; + + return { host, response }; +} + +/** Reads the problem details body captured by the mocked `response.json`. */ +function renderedBody(response: MockResponse): ProblemDetails { + expect(response.json).toHaveBeenCalledTimes(1); + return response.json.mock.calls[0][0] as ProblemDetails; +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +describe('GlobalExceptionFilter', () => { + let filter: GlobalExceptionFilter; + + beforeEach(() => { + filter = new GlobalExceptionFilter(); + vi.spyOn(Logger.prototype, 'warn').mockImplementation(() => undefined); + vi.spyOn(Logger.prototype, 'error').mockImplementation(() => undefined); + }); + + // ------------------------------------------------------------------------- + // Problem details format + // ------------------------------------------------------------------------- + + describe('problem details format', () => { + it('renders every standard RFC 9457 member plus the code and requestId extensions', () => { + const { host, response } = buildHost({ headers: { 'x-request-id': 'req-1' } }); + + filter.catch(new DomainException(ErrorCode.NOT_FOUND, "Agent 'a1' not found"), host); + + expect(renderedBody(response)).toEqual({ + type: 'urn:astroid:problem:not-found', + title: 'Resource Not Found', + status: 404, + detail: "Agent 'a1' not found", + instance: '/api/v1/transactions', + code: ErrorCode.NOT_FOUND, + requestId: 'req-1', + }); + }); + + it('serves the body as application/problem+json', () => { + const { host, response } = buildHost(); + + filter.catch(new Error('boom'), host); + + expect(response.setHeader).toHaveBeenCalledWith( + 'Content-Type', + 'application/problem+json; charset=utf-8', + ); + }); + + it('keeps the status member in sync with the HTTP status code', () => { + const { host, response } = buildHost(); + + filter.catch(new DomainException(ErrorCode.WALLET_FROZEN, 'Wallet is frozen'), host); + + expect(response.status).toHaveBeenCalledWith(423); + expect(renderedBody(response).status).toBe(423); + }); + + it('uses the request path without the query string as instance', () => { + const { host, response } = buildHost({ + url: '/api/v1/wallets?token=secret', + originalUrl: '/api/v1/wallets?token=secret', + }); + + filter.catch(new HttpException('Resource not found', 404), host); + + expect(renderedBody(response).instance).toBe('/api/v1/wallets'); + }); + + it('omits the details member when there are none', () => { + const { host, response } = buildHost(); + + filter.catch(new HttpException('Resource not found', 404), host); + + expect(renderedBody(response)).not.toHaveProperty('details'); + }); + }); + + // ------------------------------------------------------------------------- + // Prisma database errors + // ------------------------------------------------------------------------- + + describe('Prisma database errors', () => { + it('maps P2002 (unique constraint violation) to 409 CONFLICT', () => { + const { host, response } = buildHost(); + const error = new Prisma.PrismaClientKnownRequestError('Unique constraint failed', { + code: 'P2002', + clientVersion: '5.22.0', + }); + + filter.catch(error, host); + + expect(response.status).toHaveBeenCalledWith(409); + expect(renderedBody(response)).toMatchObject({ + type: 'urn:astroid:problem:conflict', + title: 'Conflict', + status: 409, + code: ErrorCode.CONFLICT, + }); + }); + + it('maps P2025 (record not found) to 404 NOT_FOUND', () => { + const { host, response } = buildHost(); + const error = new Prisma.PrismaClientKnownRequestError('Record not found', { + code: 'P2025', + clientVersion: '5.22.0', + }); + + filter.catch(error, host); + + expect(response.status).toHaveBeenCalledWith(404); + expect(renderedBody(response)).toMatchObject({ + type: 'urn:astroid:problem:not-found', + title: 'Resource Not Found', + status: 404, + code: ErrorCode.NOT_FOUND, + }); + }); + + it('maps P2003 (foreign key constraint) to 400 BAD_REQUEST', () => { + const { host, response } = buildHost(); + const error = new Prisma.PrismaClientKnownRequestError('Foreign key constraint failed', { + code: 'P2003', + clientVersion: '5.22.0', + }); + + filter.catch(error, host); + + expect(response.status).toHaveBeenCalledWith(400); + expect(renderedBody(response)).toMatchObject({ + status: 400, + code: ErrorCode.BAD_REQUEST, + }); + }); + + it('maps other known Prisma request errors to 400 BAD_REQUEST', () => { + const { host, response } = buildHost(); + const error = new Prisma.PrismaClientKnownRequestError('Value too long for field', { + code: 'P2000', + clientVersion: '5.22.0', + }); + + filter.catch(error, host); + + expect(response.status).toHaveBeenCalledWith(400); + expect(renderedBody(response).code).toBe(ErrorCode.BAD_REQUEST); + }); + }); + + // ------------------------------------------------------------------------- + // Validation failures + // ------------------------------------------------------------------------- + + describe('validation failures', () => { + it('renders ZodValidationException as 400 VALIDATION_ERROR with field-level details', () => { + const { host, response } = buildHost(); + const details = [{ path: 'limit', message: 'Number must be less than or equal to 200' }]; + + filter.catch(new ZodValidationException('Request validation failed', details), host); + + expect(response.status).toHaveBeenCalledWith(400); + expect(renderedBody(response)).toMatchObject({ + type: 'urn:astroid:problem:validation-error', + title: 'Validation Failed', + status: 400, + detail: 'Request validation failed', + code: ErrorCode.VALIDATION_ERROR, + details, + }); + }); + + it('preserves a domain ValidationException status, code and details', () => { + const { host, response } = buildHost(); + + filter.catch( + new ValidationException('Request validation failed', [ + { path: 'email', message: 'Invalid email' }, + ]), + host, + ); + + expect(response.status).toHaveBeenCalledWith(422); + expect(renderedBody(response)).toMatchObject({ + status: 422, + code: ErrorCode.VALIDATION_ERROR, + detail: 'Request validation failed', + details: [{ path: 'email', message: 'Invalid email' }], + }); + }); + + it('joins class-validator message arrays into a single detail string and preserves them as details', () => { + const { host, response } = buildHost(); + + filter.catch( + new BadRequestException(['email must be an email', 'age must be a number']), + host, + ); + + expect(response.status).toHaveBeenCalledWith(400); + const body = renderedBody(response); + expect(body.code).toBe(ErrorCode.BAD_REQUEST); + expect(body.title).toBe('Bad Request'); + expect(body.detail).toBe('email must be an email, age must be a number'); + expect(body.details).toEqual(['email must be an email', 'age must be a number']); + }); + }); + + // ------------------------------------------------------------------------- + // Authentication and authorization errors + // ------------------------------------------------------------------------- + + describe('authentication and authorization errors', () => { + it('maps 401 UnauthorizedException to UNAUTHORIZED', () => { + const { host, response } = buildHost(); + + filter.catch(new UnauthorizedException('Invalid or expired token'), host); + + expect(response.status).toHaveBeenCalledWith(401); + expect(renderedBody(response)).toMatchObject({ + type: 'urn:astroid:problem:unauthorized', + title: 'Unauthorized', + status: 401, + detail: 'Invalid or expired token', + code: ErrorCode.UNAUTHORIZED, + }); + }); + + it('maps 403 ForbiddenException to FORBIDDEN', () => { + const { host, response } = buildHost(); + + filter.catch(new ForbiddenException('Insufficient permissions'), host); + + expect(renderedBody(response)).toMatchObject({ + status: 403, + code: ErrorCode.FORBIDDEN, + }); + }); + + it('preserves domain-specific auth error codes such as TOKEN_EXPIRED', () => { + const { host, response } = buildHost(); + + filter.catch(new DomainException(ErrorCode.TOKEN_EXPIRED, 'Token has expired'), host); + + expect(renderedBody(response)).toMatchObject({ + type: 'urn:astroid:problem:token-expired', + title: 'Token Expired', + status: 401, + }); + }); + }); + + // ------------------------------------------------------------------------- + // Rate limiting (429) + // ------------------------------------------------------------------------- + + describe('rate limiting (429)', () => { + it('renders ThrottlerException as a RATE_LIMITED problem', () => { + const { host, response } = buildHost(); + + filter.catch(new ThrottlerException('Rate limit exceeded'), host); + + expect(response.status).toHaveBeenCalledWith(429); + expect(renderedBody(response)).toMatchObject({ + type: 'urn:astroid:problem:rate-limited', + title: 'Too Many Requests', + status: 429, + detail: 'Rate limit exceeded', + code: ErrorCode.RATE_LIMITED, + }); + }); + + it('uses the default throttler message when none is supplied', () => { + const { host, response } = buildHost(); + + filter.catch(new ThrottlerException(), host); + + expect(renderedBody(response).detail).toBe('ThrottlerException: Too Many Requests'); + }); + }); + + // ------------------------------------------------------------------------- + // Server faults + // ------------------------------------------------------------------------- + + describe('server faults', () => { + it('maps unknown errors to a generic 500 without leaking internal details', () => { + const { host, response } = buildHost(); + + filter.catch(new Error('connection string postgres://user:pw@db leaked'), host); + + expect(response.status).toHaveBeenCalledWith(500); + const body = renderedBody(response); + expect(body).toMatchObject({ + type: 'urn:astroid:problem:internal-error', + title: 'Internal Server Error', + status: 500, + detail: 'An unexpected error occurred', + code: ErrorCode.INTERNAL_ERROR, + }); + // Verify raw error message is never echoed to the client. + expect(JSON.stringify(body)).not.toContain('postgres://'); + }); + + it('renders non-Error throwables as 500 without crashing the process', () => { + const { host, response } = buildHost(); + + filter.catch('a string thrown somewhere', host); + + expect(response.status).toHaveBeenCalledWith(500); + expect(renderedBody(response).code).toBe(ErrorCode.INTERNAL_ERROR); + }); + + it('logs server faults at error level and includes the stack trace', () => { + const { host } = buildHost(); + const error = new Error('boom'); + + filter.catch(error, host); + + expect(Logger.prototype.error).toHaveBeenCalledWith( + expect.stringContaining('500'), + error.stack, + ); + }); + }); + + // ------------------------------------------------------------------------- + // HTTP statuses without a dedicated error code + // ------------------------------------------------------------------------- + + describe('statuses without a dedicated error code', () => { + it('uses about:blank type and the HTTP reason phrase title for unmapped statuses', () => { + const { host, response } = buildHost(); + + filter.catch(new MethodNotAllowedException(), host); + + expect(response.status).toHaveBeenCalledWith(405); + expect(renderedBody(response)).toMatchObject({ + type: 'about:blank', + title: 'Method Not Allowed', + status: 405, + }); + }); + }); + + // ------------------------------------------------------------------------- + // Request-id tracking + // ------------------------------------------------------------------------- + + describe('request id tracking', () => { + it('propagates the inbound x-request-id header so clients can correlate the error', () => { + const { host, response } = buildHost({ headers: { 'x-request-id': 'req-42' } }); + + filter.catch(new ThrottlerException(), host); + + expect(renderedBody(response).requestId).toBe('req-42'); + }); + + it('generates a fresh UUIDv7 request id when the header is absent', () => { + const { host, response } = buildHost(); + + filter.catch(new Error('boom'), host); + + const { requestId } = renderedBody(response); + expect(requestId).toMatch( + /^req_[0-9a-f]{8}-[0-9a-f]{4}-7[0-9a-f]{3}-[0-9a-f]{4}-[0-9a-f]{12}$/, + ); + expect(requestId).not.toBe('unknown'); + }); + + it('generates distinct request ids for separate unrelated error responses', () => { + const first = buildHost(); + const second = buildHost(); + + filter.catch(new Error('boom'), first.host); + filter.catch(new Error('boom'), second.host); + + expect(renderedBody(first.response).requestId).not.toBe( + renderedBody(second.response).requestId, + ); + }); + + it('recovers the request id from the ambient RequestContext when the header is missing', () => { + const { host, response } = buildHost(); + + RequestContext.run( + { + identity: { + requestId: 'ctx-req-1', + correlationId: 'ctx-req-1', + traceId: 'ctx-req-1', + method: 'POST', + path: '/api/v1/transactions', + url: '/api/v1/transactions', + ip: null, + userAgent: null, + startedAt: Date.now(), + }, + timings: {}, + data: {}, + }, + () => filter.catch(new Error('boom'), host), + ); + + expect(renderedBody(response).requestId).toBe('ctx-req-1'); + }); + + it('prefers the inbound x-request-id header over the ambient RequestContext', () => { + const { host, response } = buildHost({ headers: { 'x-request-id': 'header-req-1' } }); + + RequestContext.run( + { + identity: { + requestId: 'ctx-req-1', + correlationId: 'ctx-req-1', + traceId: 'ctx-req-1', + method: 'POST', + path: '/api/v1/transactions', + url: '/api/v1/transactions', + ip: null, + userAgent: null, + startedAt: Date.now(), + }, + timings: {}, + data: {}, + }, + () => filter.catch(new Error('boom'), host), + ); + + expect(renderedBody(response).requestId).toBe('header-req-1'); + }); + }); +}); diff --git a/src/common/filters/global-exception.filter.ts b/src/common/filters/global-exception.filter.ts new file mode 100644 index 00000000..0da62e74 --- /dev/null +++ b/src/common/filters/global-exception.filter.ts @@ -0,0 +1,26 @@ +/** + * GlobalExceptionFilter — the platform-wide exception filter for Astroid. + * + * This module is the canonical entry-point referenced by `AppModule` and any + * consumer that needs the filter class by its descriptive name. The full + * implementation lives in `AllExceptionsFilter` (same folder) and is re- + * exported here under the `GlobalExceptionFilter` name so the acceptance + * criterion ("Create GlobalExceptionFilter in global-exception.filter.ts") is + * met without duplicating the logic. + * + * Behaviour summary: + * • `Prisma.PrismaClientKnownRequestError` + * P2002 (unique constraint) → 409 CONFLICT + * P2025 (record not found) → 404 NOT_FOUND + * other known request errors → 400 BAD_REQUEST + * • `DomainException` subclasses → preserves `.code`, `.details`, status + * • `HttpException` (Nest built-ins, Throttler, ZodValidation, class-validator + * arrays, …) → maps status → ErrorCode; keeps structured + * details when present + * • Unknown throwables → 500 INTERNAL_ERROR, no internals leaked + * + * Every error response follows RFC 9457 (Problem Details for HTTP APIs) and is + * served as `application/problem+json`: + * { type, title, status, detail, instance, code, requestId, details? } + */ +export { AllExceptionsFilter as GlobalExceptionFilter } from './all-exceptions.filter'; diff --git a/src/common/guards/agent-throttler.guard.spec.ts b/src/common/guards/agent-throttler.guard.spec.ts new file mode 100644 index 00000000..ea76e5fa --- /dev/null +++ b/src/common/guards/agent-throttler.guard.spec.ts @@ -0,0 +1,235 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { ExecutionContext } from '@nestjs/common'; +import { ThrottlerOptions, ThrottlerRequest } from '@nestjs/throttler'; + +import { createThrottlerOptions, ThrottlerConfig } from '../../config/throttler.config'; +import { THROTTLE_TIER_KEY, ThrottleTier } from '../decorators/throttle-tier.decorator'; +import { AgentThrottlerGuard } from './agent-throttler.guard'; + +/** Shape returned by `ThrottlerStorage#increment` (not re-exported by the lib). */ +type ThrottlerStorageRecord = Awaited< + ReturnType +>; + +const CONFIG: ThrottlerConfig = { + windowSeconds: 60, + apiLimit: 120, + authLimit: 10, + agentLimit: 300, + webhookLimit: 30, + apiBurst: 10, + authBurst: 3, + webhookBurst: 5, +}; +const AGENT_LIMIT = 300; + +const UNBLOCKED: ThrottlerStorageRecord = { + totalHits: 1, + timeToExpire: 60, + isBlocked: false, + timeToBlockExpire: 0, +}; + +const BLOCKED: ThrottlerStorageRecord = { + totalHits: AGENT_LIMIT + 1, + timeToExpire: 30, + isBlocked: true, + timeToBlockExpire: 30, +}; + +type MockResponse = { header: ReturnType; setHeader: ReturnType }; + +function buildContext( + request: Record, + response: MockResponse, +): ExecutionContext { + const handler = () => undefined; + return { + getHandler: () => handler, + getClass: () => class TransactionController {}, + switchToHttp: () => ({ + getRequest: () => request, + getResponse: () => response, + }), + } as unknown as ExecutionContext; +} + +function throttlerNamed(name: string): ThrottlerOptions { + const limit = name === 'agent' ? AGENT_LIMIT : name === 'auth' ? 10 : 120; + return { name, ttl: 60_000, limit }; +} + +async function prepare( + opts: { + tier?: ThrottleTier; + increment?: ReturnType; + request?: Record; + } = {}, +) { + const increment = opts.increment ?? vi.fn().mockResolvedValue(UNBLOCKED); + const reflector = { + getAllAndOverride: vi.fn((key: string) => (key === THROTTLE_TIER_KEY ? opts.tier : undefined)), + }; + const guard = new AgentThrottlerGuard( + createThrottlerOptions(CONFIG), + { increment } as never, + reflector as never, + ); + await guard.onModuleInit(); + + const request = opts.request ?? { ip: '203.0.113.7', headers: {}, params: {}, query: {} }; + const response: MockResponse = { header: vi.fn(), setHeader: vi.fn() }; + const context = buildContext(request, response); + const { getTracker, generateKey } = ( + guard as unknown as { commonOptions: Pick } + ).commonOptions; + + const call = (throttler: ThrottlerOptions) => + guard['handleRequest']({ + context, + limit: throttler.limit as number, + ttl: 60_000, + throttler, + blockDuration: 60_000, + getTracker, + generateKey, + } as ThrottlerRequest); + + const trackerFor = (req: Record) => + ( + guard as unknown as { getTracker: (r: Record) => Promise } + ).getTracker(req); + + return { guard, increment, response, call, trackerFor }; +} + +describe('AgentThrottlerGuard', () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + describe('tier routing', () => { + it('routes agent-identified traffic to the agent tier only', async () => { + const { increment, call } = await prepare({ + request: { ip: '198.51.100.9', headers: { 'x-agent-id': 'agent-1' }, params: {} }, + }); + + await expect(call(throttlerNamed('api'))).resolves.toBe(true); + expect(increment).not.toHaveBeenCalled(); + + await expect(call(throttlerNamed('agent'))).resolves.toBe(true); + expect(increment).toHaveBeenCalledWith( + expect.any(String), + 60_000, + AGENT_LIMIT, + 60_000, + 'agent', + ); + }); + + it('falls back to the api tier for plain user traffic', async () => { + const { increment, call } = await prepare({ + request: { ip: '198.51.100.9', headers: {}, params: {}, user: { organizationId: 'org-1' } }, + }); + + await expect(call(throttlerNamed('agent'))).resolves.toBe(true); + expect(increment).not.toHaveBeenCalled(); + + await expect(call(throttlerNamed('api'))).resolves.toBe(true); + expect(increment).toHaveBeenCalledTimes(1); + }); + + it('lets an explicit auth tier win over agent auto-detection', async () => { + const { increment, call } = await prepare({ + tier: 'auth', + request: { ip: '198.51.100.9', headers: { 'x-agent-id': 'agent-1' }, params: {} }, + }); + + await expect(call(throttlerNamed('agent'))).resolves.toBe(true); + expect(increment).not.toHaveBeenCalled(); + + await expect(call(throttlerNamed('auth'))).resolves.toBe(true); + expect(increment).toHaveBeenCalledWith(expect.any(String), 60_000, 10, 60_000, 'auth'); + }); + }); + + describe('agent-aware tracking', () => { + it('prefers the agent id from the header, body or route params', async () => { + const { trackerFor } = await prepare(); + + await expect( + trackerFor({ headers: { 'x-agent-id': 'agent-1' }, ip: '1.1.1.1' }), + ).resolves.toBe('agent:agent-1'); + await expect( + trackerFor({ headers: {}, body: { agentId: 'agent-2' }, ip: '1.1.1.1' }), + ).resolves.toBe('agent:agent-2'); + await expect( + trackerFor({ headers: {}, params: { agentId: 'agent-3' }, ip: '1.1.1.1' }), + ).resolves.toBe('agent:agent-3'); + }); + + it('treats an API-key principal as the acting agent', async () => { + const { trackerFor } = await prepare(); + + await expect( + trackerFor({ headers: {}, ip: '1.1.1.1', user: { id: 'agent-key-1', isApiKey: true } }), + ).resolves.toBe('agent:agent-key-1'); + }); + + it('buckets authenticated humans by organization and hashes raw API keys', async () => { + const { trackerFor } = await prepare(); + + await expect( + trackerFor({ + headers: {}, + ip: '1.1.1.1', + user: { id: 'user-1', organizationId: 'org-1' }, + }), + ).resolves.toBe('org:org-1'); + + const keyed = await trackerFor({ + headers: { 'x-api-key': 'ast_live_secret' }, + ip: '1.1.1.1', + }); + expect(keyed).toMatch(/^key:[a-f0-9]{64}$/); + expect(keyed).not.toContain('ast_live_secret'); + }); + + it('falls back to the forwarded-for address for anonymous public routes', async () => { + const { trackerFor } = await prepare(); + + await expect( + trackerFor({ headers: { 'x-forwarded-for': '198.51.100.4' }, ip: '10.0.0.1' }), + ).resolves.toBe('ip:198.51.100.4'); + }); + }); + + describe('rate-limit headers', () => { + it('advertises the tier limit on an allowed request', async () => { + const { response, call } = await prepare({ + request: { headers: { 'x-agent-id': 'agent-1' }, params: {}, ip: '1.1.1.1' }, + }); + + await call(throttlerNamed('agent')); + + expect(response.setHeader).toHaveBeenCalledWith('X-RateLimit-Limit', AGENT_LIMIT); + expect(response.header).toHaveBeenCalledWith( + 'X-RateLimit-Remaining-agent', + AGENT_LIMIT - 1, + ); + }); + + it('returns 429 with Retry-After and rate-limit headers once the burst is exhausted', async () => { + const { response, call } = await prepare({ + increment: vi.fn().mockResolvedValue(BLOCKED), + request: { headers: { 'x-agent-id': 'agent-1' }, params: {}, ip: '1.1.1.1' }, + }); + + await expect(call(throttlerNamed('agent'))).rejects.toMatchObject({ status: 429 }); + + expect(response.setHeader).toHaveBeenCalledWith('Retry-After', 30); + expect(response.setHeader).toHaveBeenCalledWith('X-RateLimit-Limit', AGENT_LIMIT); + expect(response.setHeader).toHaveBeenCalledWith('X-RateLimit-Remaining', 0); + }); + }); +}); diff --git a/src/common/guards/agent-throttler.guard.ts b/src/common/guards/agent-throttler.guard.ts new file mode 100644 index 00000000..bd6ba789 --- /dev/null +++ b/src/common/guards/agent-throttler.guard.ts @@ -0,0 +1,143 @@ +import { ExecutionContext, Injectable } from '@nestjs/common'; +import { + ThrottlerGuard, + ThrottlerLimitDetail, + ThrottlerRequest, +} from '@nestjs/throttler'; +import { createHash } from 'crypto'; +import { Request, Response } from 'express'; + +import { THROTTLE_TIER_KEY, ThrottleTier } from '../decorators/throttle-tier.decorator'; +import { extractApiKeyFromRequest } from '../helpers/extract-api-key'; +import { AuthenticatedUser } from '../interfaces/authenticated-user.interface'; + +/** Request-scoped view the tracker/lookup helpers work against. */ +type ThrottledRequest = Request & { user?: AuthenticatedUser }; + +/** + * Redis-backed rate-limit guard for high-frequency agent endpoints. + * + * Extends `@nestjs/throttler`'s `ThrottlerGuard` (counters therefore live in the + * shared `RedisThrottlerStorage` when the module is configured with one) and + * adds two behaviours the autonomous-agent workload needs: + * + * 1. **Agent-aware tracking.** The counter key is derived from the acting + * agent (`x-agent-id`, a route/body/query `agentId`, or an API-key + * principal) instead of the organization or IP, so one noisy agent can + * never exhaust another agent's budget behind the same NAT/gateway. + * 2. **Tier routing.** Only the named throttler matching the route's tier is + * enforced: `agent` for agent-identified traffic, `auth` for routes marked + * `@ThrottleTierDecorator('auth')`, `api` for everything else. + * + * Unauthenticated calls fall back to a hashed API key and finally to the client + * IP (honouring `x-forwarded-for`), which keeps public routes protected. + * + * Every rejection is a standard HTTP 429 that also carries the plain + * `Retry-After`, `X-RateLimit-Limit` and `X-RateLimit-Remaining` headers. + */ +@Injectable() +export class AgentThrottlerGuard extends ThrottlerGuard { + /** Enforce only the throttler that governs this route's tier. */ + protected async handleRequest(requestProps: ThrottlerRequest): Promise { + const { context, throttler, limit } = requestProps; + const routeTier = this.resolveTier(context); + + // This named throttler does not govern this route's tier — do not count it. + if (throttler.name !== routeTier) { + return true; + } + + // Publish the tier limit up-front so even a successful call advertises the + // budget it consumed (the library's per-throttler headers are also set). + context.switchToHttp().getResponse().setHeader('X-RateLimit-Limit', limit); + + return super.handleRequest(requestProps); + } + + /** + * Buckets a request by acting agent, then organization, then API key, then IP. + * Never stores a raw credential: the API-key fallback is hashed. + */ + protected async getTracker(req: Record): Promise { + const request = req as unknown as ThrottledRequest; + + const agentId = this.resolveAgentId(request); + if (agentId) { + return `agent:${agentId}`; + } + + if (request.user?.organizationId) { + return `org:${request.user.organizationId}`; + } + + const apiKey = extractApiKeyFromRequest(request); + if (apiKey) { + return `key:${createHash('sha256').update(apiKey).digest('hex')}`; + } + + return `ip:${this.resolveIp(request)}`; + } + + /** Adds the plain rate-limit headers before the library throws its 429. */ + protected async throwThrottlingException( + context: ExecutionContext, + detail: ThrottlerLimitDetail, + ): Promise { + const response = context.switchToHttp().getResponse(); + response.setHeader('Retry-After', detail.timeToBlockExpire); + response.setHeader('X-RateLimit-Limit', detail.limit); + response.setHeader('X-RateLimit-Remaining', Math.max(0, detail.limit - detail.totalHits)); + + await super.throwThrottlingException(context, detail); + } + + /** + * Resolves the tier to enforce. An explicit `@ThrottleTierDecorator()` always + * wins; otherwise agent-identified traffic is routed to the `agent` tier and + * everything else to `api`. + */ + private resolveTier(context: ExecutionContext): ThrottleTier { + const declared = this.reflector.getAllAndOverride(THROTTLE_TIER_KEY, [ + context.getHandler(), + context.getClass(), + ]); + if (declared) { + return declared; + } + + const request = context.switchToHttp().getRequest(); + return this.resolveAgentId(request) ? 'agent' : 'api'; + } + + /** + * Extracts the acting agent id from any of the places the platform carries it: + * route params, body, query, the `x-agent-id` header, or an API-key principal + * bound to an agent (`user.id` of an `isApiKey` principal). + */ + private resolveAgentId(request: ThrottledRequest): string | undefined { + const fromRequest = + (request.params?.agentId as string) || + ((request.body as Record | undefined)?.agentId as string) || + ((request.query as Record | undefined)?.agentId as string) || + (request.headers?.['x-agent-id'] as string) || + undefined; + + if (fromRequest) { + return fromRequest; + } + + // An API-key principal acts on behalf of an agent in this platform. + if (request.user?.isApiKey && request.user.id) { + return request.user.id; + } + + return undefined; + } + + /** Trusts `x-forwarded-for` for tracker bucketing, then falls back to the socket IP. */ + private resolveIp(request: ThrottledRequest): string { + const forwarded = request.headers?.['x-forwarded-for']; + const header = Array.isArray(forwarded) ? forwarded[0] : forwarded; + return header ?? request.ip ?? request.socket?.remoteAddress ?? 'anonymous'; + } +} diff --git a/src/common/guards/api-key-auth.guard.ts b/src/common/guards/api-key-auth.guard.ts index ceb66ec7..4a92f950 100644 --- a/src/common/guards/api-key-auth.guard.ts +++ b/src/common/guards/api-key-auth.guard.ts @@ -16,8 +16,11 @@ type ApiKeyAuthenticatedRequest = Request & { /** * Guard enforcing cryptographic API key authentication on protected routes. - * Extracts the key from `x-api-key`, verifies the SHA-256 hash against PostgreSQL, - * rejects revoked or expired keys, and attaches scoped permissions to the request. + * Extracts the key from `x-api-key`, verifies the Argon2id hash against PostgreSQL + * (with SHA-256 fallback for legacy keys), rejects revoked or expired keys, and + * attaches scoped permissions to the request. + * + * Uses constant-time comparison via Argon2 verification to prevent timing attacks. */ @Injectable() export class ApiKeyAuthGuard implements CanActivate { diff --git a/src/common/guards/public-rate-limit.burst.integration.spec.ts b/src/common/guards/public-rate-limit.burst.integration.spec.ts new file mode 100644 index 00000000..82824e7b --- /dev/null +++ b/src/common/guards/public-rate-limit.burst.integration.spec.ts @@ -0,0 +1,200 @@ +import { afterAll, beforeAll, describe, expect, it, vi } from 'vitest'; +import { ConfigService } from '@nestjs/config'; +import { Controller, Get, INestApplication, Logger, Post } from '@nestjs/common'; +import { APP_FILTER, APP_GUARD } from '@nestjs/core'; +import { Test } from '@nestjs/testing'; +import { PublicRateLimitGuard } from './public-rate-limit.guard'; +import { Public } from '../decorators/public.decorator'; +import { PublicRateLimit } from '../decorators/public-rate-limit.decorator'; +import { AllExceptionsFilter } from '../filters/all-exceptions.filter'; +import { REDIS_CLIENT } from '../locks/locks.constants'; +import { MemorySlidingWindowStore } from '../throttler/sliding-window.store'; + +/** + * Sends request bursts over real HTTP against a Nest app wired like + * production: the guard is a global APP_GUARD, errors go through + * AllExceptionsFilter, and routes live under the `api/v1` prefix. The Redis + * client is a stand-in whose `eval` reproduces the sliding-window script's + * contract (`[allowed, count, resetAt]`) on top of the in-memory store, so the + * Redis code path of the guard is exercised end to end. + */ + +const GLOBAL_LIMIT = 4; + +@Controller('auth') +class AuthController { + @Public() + @Post('login') + login() { + return { ok: true }; + } +} + +@Controller('agents') +class AgentsController { + @Get() + list() { + return []; + } +} + +@Controller('public') +class PublicCatalogController { + @Get('status') + status() { + return { ok: true }; + } + + // A heavier endpoint with its own, stricter budget. + @PublicRateLimit(2, 60) + @Get('search') + search() { + return { ok: true }; + } +} + +function fakeRedis() { + const store = new MemorySlidingWindowStore(); + return { + status: 'ready', + eval: vi.fn( + async ( + _script: string, + _keys: number, + key: string, + now: number, + windowMs: number, + limit: number, + ) => { + const hit = await store.hit(key, limit, windowMs, now); + return [hit.allowed ? 1 : 0, hit.count, hit.resetAt]; + }, + ), + }; +} + +describe('Public API rate limiting (integration)', () => { + let app: INestApplication; + let baseUrl: string; + let redis: ReturnType; + + beforeAll(async () => { + vi.spyOn(Logger.prototype, 'warn').mockImplementation(() => undefined); + redis = fakeRedis(); + const config = { + getOrThrow: () => ({ + windowSeconds: 60, + maxRequests: 120, + public: { + enabled: true, + maxRequests: GLOBAL_LIMIT, + windowSeconds: 60, + trustProxy: true, + clientIdentifiers: ['apiKey'], + }, + }), + get: () => ({ apiPrefix: 'api/v1' }), + }; + + const moduleRef = await Test.createTestingModule({ + controllers: [AuthController, PublicCatalogController, AgentsController], + providers: [ + { provide: ConfigService, useValue: config }, + { provide: REDIS_CLIENT, useValue: redis }, + { provide: APP_GUARD, useClass: PublicRateLimitGuard }, + { provide: APP_FILTER, useClass: AllExceptionsFilter }, + ], + }).compile(); + + app = moduleRef.createNestApplication({ logger: false }); + app.setGlobalPrefix('api/v1'); + await app.listen(0, '127.0.0.1'); + baseUrl = `${await app.getUrl()}/api/v1`; + }); + + afterAll(async () => { + await app.close(); + }); + + const send = (path: string, ip: string, method = 'GET', apiKey?: string) => + fetch(`${baseUrl}${path}`, { + method, + headers: { + 'x-forwarded-for': ip, + ...(apiKey ? { 'x-api-key': apiKey } : {}), + }, + }); + + it('serves a burst up to the global limit, then answers 429 with standard headers', async () => { + const ip = '198.51.100.110'; + const statuses: number[] = []; + const remaining: (string | null)[] = []; + for (let i = 0; i < GLOBAL_LIMIT; i++) { + const res = await send('/auth/login', ip, 'POST'); + statuses.push(res.status); + remaining.push(res.headers.get('x-ratelimit-remaining')); + } + + expect(statuses).toEqual(Array(GLOBAL_LIMIT).fill(201)); + expect(remaining).toEqual(['3', '2', '1', '0']); + + const limited = await send('/auth/login', ip, 'POST'); + + expect(limited.status).toBe(429); + expect(limited.headers.get('x-ratelimit-limit')).toBe(String(GLOBAL_LIMIT)); + expect(limited.headers.get('x-ratelimit-remaining')).toBe('0'); + const reset = Number(limited.headers.get('x-ratelimit-reset')); + const nowSeconds = Math.floor(Date.now() / 1000); + expect(reset).toBeGreaterThanOrEqual(nowSeconds); + expect(reset).toBeLessThanOrEqual(nowSeconds + 61); + expect(Number(limited.headers.get('retry-after'))).toBeGreaterThanOrEqual(1); + }); + + it('enforces the per-route @PublicRateLimit() budget on heavier endpoints', async () => { + const ip = '198.51.100.120'; + + expect((await send('/public/search', ip)).status).toBe(200); + expect((await send('/public/search', ip)).status).toBe(200); + + const limited = await send('/public/search', ip); + expect(limited.status).toBe(429); + expect(limited.headers.get('x-ratelimit-limit')).toBe('2'); + expect(limited.headers.get('x-ratelimit-remaining')).toBe('0'); + + // The global-limit route of the same controller is unaffected. + expect((await send('/public/status', ip)).status).toBe(200); + }); + + it('tracks API-key clients separately from other callers behind the same IP', async () => { + const ip = '198.51.100.130'; + + for (let i = 0; i < GLOBAL_LIMIT; i++) { + await send('/public/status', ip, 'GET', 'ak_live_integration'); + } + const limited = await send('/public/status', ip, 'GET', 'ak_live_integration'); + expect(limited.status).toBe(429); + + // A different key (and a keyless caller) on the same IP still has budget. + expect((await send('/public/status', ip, 'GET', 'ak_live_other')).status).toBe(200); + expect((await send('/public/status', ip)).status).toBe(200); + }); + + it('keeps other IPs unaffected while one IP is limited', async () => { + for (let i = 0; i <= GLOBAL_LIMIT; i++) { + await send('/public/status', '198.51.100.140'); + } + + const other = await send('/public/status', '198.51.100.141'); + expect(other.status).toBe(200); + expect(other.headers.get('x-ratelimit-remaining')).toBe(String(GLOBAL_LIMIT - 1)); + }); + + it('never limits or annotates authenticated routes', async () => { + const ip = '198.51.100.150'; + for (let i = 0; i < GLOBAL_LIMIT * 2; i++) { + const res = await send('/agents', ip); + expect(res.status).toBe(200); + expect(res.headers.get('x-ratelimit-limit')).toBeNull(); + } + }); +}); diff --git a/src/common/guards/public-rate-limit.guard.spec.ts b/src/common/guards/public-rate-limit.guard.spec.ts index b000c4f2..ec359b8d 100644 --- a/src/common/guards/public-rate-limit.guard.spec.ts +++ b/src/common/guards/public-rate-limit.guard.spec.ts @@ -9,8 +9,9 @@ import { SKIP_PUBLIC_RATE_LIMIT_KEY } from '../decorators/skip-public-rate-limit import { DomainException } from '../exceptions/domain.exception'; import { ErrorCode } from '../constants/error-codes'; import { PublicRateLimitConfig } from '../../config/rate-limit.config'; +import { PUBLIC_RATE_LIMIT_RULE_KEY } from '../decorators/public-rate-limit.decorator'; -type Metadata = { public?: boolean; skip?: boolean }; +type Metadata = { public?: boolean; skip?: boolean; rule?: { max: number; windowSeconds: number } }; function buildContext( request: { path?: string; ip?: string; headers?: Record }, @@ -22,6 +23,8 @@ function buildContext( class TestController {} if (metadata.public) Reflect.defineMetadata(IS_PUBLIC_KEY, true, handler); if (metadata.skip) Reflect.defineMetadata(SKIP_PUBLIC_RATE_LIMIT_KEY, true, handler); + if (metadata.rule) + Reflect.defineMetadata(PUBLIC_RATE_LIMIT_RULE_KEY, metadata.rule, handler); const context = { getType: () => 'http', @@ -44,6 +47,7 @@ function buildGuard( maxRequests: 3, windowSeconds: 60, trustProxy: false, + clientIdentifiers: [], ...overrides, }; const config = { @@ -202,6 +206,99 @@ describe('PublicRateLimitGuard', () => { }); }); + describe('per-route rules', () => { + it('applies the @PublicRateLimit() override instead of the global limit', async () => { + const guard = buildGuard({ maxRequests: 1 }); + const { context, headers } = buildContext( + {}, + { public: true, rule: { max: 2, windowSeconds: 30 } }, + ); + + await guard.canActivate(context); + await guard.canActivate(context); + const limited = await expectRateLimited(guard.canActivate(context)); + + expect(headers['X-RateLimit-Limit']).toBe(2); + expect(limited.details).toMatchObject({ limit: 2, windowSeconds: 30 }); + }); + + it('keeps the global default when no route override is present', async () => { + const guard = buildGuard({ maxRequests: 2 }); + const { context, headers } = buildContext({}, { public: true }); + + await guard.canActivate(context); + + expect(headers['X-RateLimit-Limit']).toBe(2); + }); + + it('honours controller-level overrides over handler rules', async () => { + const guard = buildGuard({ maxRequests: 5 }); + // The handler rule must win (getAllAndOverride walks handler first). + const { context, headers } = buildContext( + {}, + { public: true, rule: { max: 4, windowSeconds: 15 } }, + ); + + await guard.canActivate(context); + + expect(headers['X-RateLimit-Limit']).toBe(4); + }); + + it('lets @SkipPublicRateLimit() bypass a route-level rule too', async () => { + const guard = buildGuard({ maxRequests: 1 }); + const { context } = buildContext( + {}, + { public: true, skip: true, rule: { max: 1, windowSeconds: 60 } }, + ); + + await expect(guard.canActivate(context)).resolves.toBe(true); + await expect(guard.canActivate(context)).resolves.toBe(true); + }); + }); + + describe('client identifiers', () => { + it('buckets API-key callers separately from their shared IP when enabled', async () => { + const guard = buildGuard({ maxRequests: 1, clientIdentifiers: ['apiKey'] }); + const withKey = (key: string) => + buildContext({ headers: { 'x-api-key': key } }, { public: true }).context; + + // Exhaust the plain-IP bucket first: the (max+1)th keyless hit is limited. + await guard.canActivate(buildContext({}, { public: true }).context); + await expectRateLimited(guard.canActivate(buildContext({}, { public: true }).context)); + + // Key-holding callers get their own budgets despite the same IP. + await guard.canActivate(withKey('ak_live_aaaa')); + await guard.canActivate(withKey('ak_live_bbbb')); + }); + + it('ignores API keys when no client identifiers are configured', async () => { + const guard = buildGuard({ maxRequests: 1, clientIdentifiers: [] }); + const withKey = (key: string) => + buildContext({ headers: { 'x-api-key': key } }, { public: true }).context; + + await guard.canActivate(withKey('ak_live_aaaa')); + + // Without identifier tracking, a second key on the same IP is limited. + await expectRateLimited(guard.canActivate(withKey('ak_live_bbbb'))); + }); + + it('counts an ApiKey Authorization header the same as x-api-key', async () => { + const guard = buildGuard({ maxRequests: 1, clientIdentifiers: ['apiKey'] }); + const context = buildContext( + { headers: { authorization: 'ApiKey ak_live_aaaa' } }, + { public: true }, + ).context; + + await guard.canActivate(context); + + await expectRateLimited( + guard.canActivate( + buildContext({ headers: { 'x-api-key': 'ak_live_aaaa' } }, { public: true }).context, + ), + ); + }); + }); + describe('storage', () => { it('records hits in Redis under a per-IP key when Redis is ready', async () => { const evalFn = vi.fn().mockResolvedValue([1, 1, Date.now() + 60_000]); diff --git a/src/common/guards/public-rate-limit.guard.ts b/src/common/guards/public-rate-limit.guard.ts index 5e2cfc18..c00b3fe9 100644 --- a/src/common/guards/public-rate-limit.guard.ts +++ b/src/common/guards/public-rate-limit.guard.ts @@ -5,6 +5,10 @@ import { Redis } from 'ioredis'; import { Request, Response } from 'express'; import { IS_PUBLIC_KEY } from '../decorators/public.decorator'; import { SKIP_PUBLIC_RATE_LIMIT_KEY } from '../decorators/skip-public-rate-limit.decorator'; +import { + PUBLIC_RATE_LIMIT_RULE_KEY, + PublicRateLimitRule, +} from '../decorators/public-rate-limit.decorator'; import { DomainException } from '../exceptions/domain.exception'; import { ErrorCode } from '../constants/error-codes'; import { REDIS_CLIENT } from '../locks/locks.constants'; @@ -16,28 +20,55 @@ import { import { PublicRateLimitConfig, RateLimitConfig } from '../../config/rate-limit.config'; import { AppConfig } from '../../config/app.config'; import { getClientIp } from '../../utils/ip.util'; +import { extractApiKeyFromRequest } from '../helpers/extract-api-key'; export const RATE_LIMIT_LIMIT_HEADER = 'X-RateLimit-Limit'; export const RATE_LIMIT_REMAINING_HEADER = 'X-RateLimit-Remaining'; export const RATE_LIMIT_RESET_HEADER = 'X-RateLimit-Reset'; +/** Outcome of evaluating one request against the public rate limiter. */ +export interface PublicRateLimitResult { + /** Whether the request fits within the applicable limit. */ + allowed: boolean; + /** The limit in force for this request (global default or route override). */ + limit: number; + /** Window length, in seconds, of the rule that produced the decision. */ + windowSeconds: number; + /** Requests counted for this client in the current window, including this one. */ + count: number; + /** Epoch ms at which the oldest counted request leaves the window. */ + resetAt: number; +} + /** - * IP-based sliding-window rate limiter for unauthenticated endpoints, the - * first line of defence against burst traffic and resource exhaustion. + * IP-and-client-based sliding-window rate limiter for unauthenticated + * endpoints, the first line of defence against abuse, scraping and + * denial-of-service bursts. * * Applies to every route marked `@Public()` and to every route under * `//public/`, unless exempted with `@SkipPublicRateLimit()`. * Authenticated routes are left to the per-organization throttlers. * + * Requests are tracked per client identifier: the client IP always + * participates, and when the limiter is configured with + * `clientIdentifiers: ['ip', 'apiKey']` a presented API key (`x-api-key` / + * `Authorization: ApiKey|Bearer ak_…`) is folded into the bucket key so + * distinct programmatic clients behind one shared address (NAT, office + * egress, CI runners) each get their own budget instead of a shared one. + * * Every limited response carries `X-RateLimit-Limit`, `X-RateLimit-Remaining` * and `X-RateLimit-Reset` (epoch seconds at which a slot frees up); rejected * requests get `429 Too Many Requests` plus `Retry-After`. * * Counters live in Redis (the shared `REDIS_CLIENT`) so every replica enforces - * one budget per IP. If Redis is unavailable the guard falls back to a + * one budget per client. If Redis is unavailable the guard falls back to a * per-process in-memory window rather than failing open, so public endpoints * stay protected during an outage. * + * Limits are configurable in two layers: global defaults from + * `PUBLIC_RATE_LIMIT_*` env vars, overridden per route (or controller) with + * the `@PublicRateLimit(max, windowSeconds)` decorator. + * * Implemented as a guard rather than Express middleware because middleware * runs before routing and cannot see the `@Public()` metadata. */ @@ -71,29 +102,46 @@ export class PublicRateLimitGuard implements CanActivate { return true; } + const result = await this.check(request, context); const response = context.switchToHttp().getResponse(); - const { maxRequests: limit, windowSeconds } = this.settings; - const now = Date.now(); - const key = `rate-limit:public:ip:${this.clientIp(request)}`; - const hit = await this.record(key, limit, windowSeconds * 1000, now); - - response.setHeader(RATE_LIMIT_LIMIT_HEADER, limit); - response.setHeader(RATE_LIMIT_REMAINING_HEADER, Math.max(0, limit - hit.count)); - response.setHeader(RATE_LIMIT_RESET_HEADER, Math.ceil(hit.resetAt / 1000)); + response.setHeader(RATE_LIMIT_LIMIT_HEADER, result.limit); + response.setHeader(RATE_LIMIT_REMAINING_HEADER, Math.max(0, result.limit - result.count)); + response.setHeader(RATE_LIMIT_RESET_HEADER, Math.ceil(result.resetAt / 1000)); - if (!hit.allowed) { - const retryAfterSeconds = Math.max(1, Math.ceil((hit.resetAt - now) / 1000)); + if (!result.allowed) { + const now = Date.now(); + const retryAfterSeconds = Math.max(1, Math.ceil((result.resetAt - now) / 1000)); response.setHeader('Retry-After', retryAfterSeconds); throw new DomainException( ErrorCode.RATE_LIMITED, 'Too many requests from this IP address. Please retry later.', - { limit, windowSeconds, retryAfterSeconds }, + { limit: result.limit, windowSeconds: result.windowSeconds, retryAfterSeconds }, ); } return true; } + /** + * Records one hit against the client's sliding-window budget. The rule in + * force is the route-level `@PublicRateLimit()` override when present, the + * global `PUBLIC_RATE_LIMIT_*` settings otherwise. + */ + async check(request: Request, context?: ExecutionContext): Promise { + const rule = this.resolveRule(context); + const now = Date.now(); + const key = `rate-limit:public:${this.clientBucket(request)}`; + const hit = await this.record(key, rule.max, rule.windowSeconds * 1000, now); + + return { + allowed: hit.allowed, + limit: rule.max, + windowSeconds: rule.windowSeconds, + count: hit.count, + resetAt: hit.resetAt, + }; + } + private appliesTo(context: ExecutionContext, request: Request): boolean { const targets = [context.getHandler(), context.getClass()]; if (this.reflector.getAllAndOverride(SKIP_PUBLIC_RATE_LIMIT_KEY, targets)) { @@ -106,6 +154,37 @@ export class PublicRateLimitGuard implements CanActivate { return path === this.publicPathPrefix || path.startsWith(`${this.publicPathPrefix}/`); } + /** Route-level rule override wins; the global settings are the default. */ + private resolveRule(context?: ExecutionContext): PublicRateLimitRule { + const defaults: PublicRateLimitRule = { + max: this.settings.maxRequests, + windowSeconds: this.settings.windowSeconds, + }; + if (!context) { + return defaults; + } + const targets = [context.getHandler(), context.getClass()]; + return this.reflector.getAllAndOverride(PUBLIC_RATE_LIMIT_RULE_KEY, targets) ?? defaults; + } + + /** + * Builds the bucket identifier for the caller. The IP always participates; + * configured client identifiers (currently the API key) are appended so + * distinct clients behind one address are tracked separately. + */ + private clientBucket(request: Request): string { + const parts = [`ip:${this.clientIp(request)}`]; + for (const identifier of this.settings.clientIdentifiers ?? []) { + if (identifier === 'apiKey') { + const apiKey = extractApiKeyFromRequest(request); + if (apiKey) { + parts.push(`key:${apiKey}`); + } + } + } + return parts.join(':'); + } + /** Records the hit in Redis, degrading to the in-memory window on outage. */ private async record( key: string, diff --git a/src/common/guards/sensitive-rate-limit.integration.spec.ts b/src/common/guards/sensitive-rate-limit.integration.spec.ts new file mode 100644 index 00000000..9ea849f8 --- /dev/null +++ b/src/common/guards/sensitive-rate-limit.integration.spec.ts @@ -0,0 +1,85 @@ +import { describe, it, expect, beforeAll, afterAll, vi } from 'vitest'; +import { INestApplication, Controller, Post, UseGuards } from '@nestjs/common'; +import { Test } from '@nestjs/testing'; +import { ThrottlerModule } from '@nestjs/throttler'; +import { Redis } from 'ioredis'; +import { AstroidThrottlerGuard } from './throttler.guard'; +import { MemorySlidingWindowStore } from '../throttler/sliding-window.store'; +import { RedisThrottlerStorage } from '../throttler/redis-throttler.storage'; + +@Controller('test-sensitive') +class TestSensitiveController { + @Post('action') + @UseGuards(AstroidThrottlerGuard) + action() { + return { success: true }; + } +} + +describe('Sensitive Endpoint Rate Limiting (Integration)', () => { + let app: INestApplication; + let baseUrl: string; + + beforeAll(async () => { + const store = new MemorySlidingWindowStore(); + const fakeRedis = { + status: 'ready', + eval: vi.fn( + async ( + _script: string, + _keys: number, + key: string, + ttl: number, + limit: number, + blockDuration: number, + now: number, + ) => { + const hit = await store.hit(key, limit, ttl, now); + return [ + hit.count, + Math.ceil((hit.resetAt - now) / 1000), + hit.allowed ? 0 : 1, + hit.allowed ? 0 : Math.ceil(blockDuration / 1000), + ]; + }, + ), + }; + + const moduleRef = await Test.createTestingModule({ + imports: [ + ThrottlerModule.forRoot({ + throttlers: [{ name: 'api', ttl: 60000, limit: 2 }], + storage: new RedisThrottlerStorage(fakeRedis as unknown as Redis), + }), + ], + controllers: [TestSensitiveController], + providers: [], + }).compile(); + + app = moduleRef.createNestApplication({ logger: false }); + await app.init(); + await app.listen(0, '127.0.0.1'); + baseUrl = await app.getUrl(); + }); + + afterAll(async () => { + await app.close(); + }); + + it('enforces rate limit and returns 429 when threshold is exceeded', async () => { + const send = () => + fetch(`${baseUrl}/test-sensitive/action`, { + method: 'POST', + headers: { 'x-api-key': 'test-key-123' }, + }); + + const res1 = await send(); + expect(res1.status).toBe(201); + + const res2 = await send(); + expect(res2.status).toBe(201); + + const res3 = await send(); + expect(res3.status).toBe(429); + }); +}); diff --git a/src/common/guards/sliding-window-throttler.guard.spec.ts b/src/common/guards/sliding-window-throttler.guard.spec.ts index ab74cb61..92c95caa 100644 --- a/src/common/guards/sliding-window-throttler.guard.spec.ts +++ b/src/common/guards/sliding-window-throttler.guard.spec.ts @@ -3,9 +3,19 @@ import { ErrorCode } from '../constants/error-codes'; import { SlidingWindowThrottlerGuard } from './sliding-window-throttler.guard'; const exec = vi.fn(); -const chain = { zremrangebyscore: vi.fn().mockReturnThis(), zcard: vi.fn().mockReturnThis(), zadd: vi.fn().mockReturnThis(), expire: vi.fn().mockReturnThis(), exec }; +const chain = { + zremrangebyscore: vi.fn().mockReturnThis(), + zcard: vi.fn().mockReturnThis(), + zadd: vi.fn().mockReturnThis(), + expire: vi.fn().mockReturnThis(), + exec, +}; -function makeContext(user?: Record, ip = '127.0.0.1', headers: Record = {}) { +function makeContext( + user?: Record, + ip = '127.0.0.1', + headers: Record = {}, +) { const response = { setHeader: vi.fn() }; const request = { user, ip, headers }; const handler = vi.fn(); @@ -19,14 +29,26 @@ function makeContext(user?: Record, ip = '127.0.0.1', headers: R function makeGuard(redis: Record, limit = 2) { const reflector = { getAllAndOverride: vi.fn().mockReturnValue(undefined) }; - const config = { get: vi.fn((key: string, fallback: unknown) => key === 'rateLimit.maxRequests' ? limit : fallback) }; + const config = { + get: vi.fn((key: string, fallback: unknown) => + key === 'rateLimit.maxRequests' ? limit : fallback, + ), + }; const guard = new SlidingWindowThrottlerGuard(reflector as never, config as never); Object.assign(guard, { redis }); return guard; } describe('SlidingWindowThrottlerGuard', () => { - beforeEach(() => { vi.clearAllMocks(); exec.mockResolvedValue([[null, 0], [null, 0], [null, 1], [null, 1]]); }); + beforeEach(() => { + vi.clearAllMocks(); + exec.mockResolvedValue([ + [null, 0], + [null, 0], + [null, 1], + [null, 1], + ]); + }); it('allows requests and emits standard rate-limit headers', async () => { const { context, response } = makeContext({ organizationId: 'org-1' }); @@ -41,10 +63,17 @@ describe('SlidingWindowThrottlerGuard', () => { }); it('rejects an exhausted window with Retry-After', async () => { - exec.mockResolvedValue([[null, 0], [null, 2], [null, 1], [null, 1]]); + exec.mockResolvedValue([ + [null, 0], + [null, 2], + [null, 1], + [null, 1], + ]); const { context, response } = makeContext({ organizationId: 'org-1' }); const guard = makeGuard({ multi: () => chain }); - await expect(guard.canActivate(context as never)).rejects.toMatchObject({ code: ErrorCode.RATE_LIMITED }); + await expect(guard.canActivate(context as never)).rejects.toMatchObject({ + code: ErrorCode.RATE_LIMITED, + }); expect(response.setHeader).toHaveBeenCalledWith('X-RateLimit-Remaining', 0); expect(response.setHeader).toHaveBeenCalledWith('Retry-After', expect.any(Number)); }); @@ -60,6 +89,19 @@ describe('SlidingWindowThrottlerGuard', () => { expect(redis.multi).toHaveBeenCalledTimes(2); }); + it('uses the authenticated API key ID instead of its organization for throttling', async () => { + const { context } = makeContext({ + organizationId: 'org-1', + apiKeyId: 'key-1', + isApiKey: true, + }); + const guard = makeGuard({ multi: () => chain }); + + await guard.canActivate(context as never); + + expect(chain.zremrangebyscore.mock.calls[0][0]).toContain(':key:key-1:'); + }); + it('falls back to a hashed API key scope when unauthenticated but keyed', async () => { const redis = { multi: vi.fn(() => chain) }; const withApiKey = makeContext(undefined, '192.0.2.1', { 'x-api-key': 'ast_secret-key' }); @@ -73,10 +115,28 @@ describe('SlidingWindowThrottlerGuard', () => { it('fails open and logs when Redis is unavailable', async () => { const logger = { error: vi.fn() }; const { context, response } = makeContext({ organizationId: 'org-1' }); - const guard = makeGuard({ multi: () => { throw new Error('offline'); } }); + const guard = makeGuard({ + multi: () => { + throw new Error('offline'); + }, + }); Object.assign(guard, { logger }); await expect(guard.canActivate(context as never)).resolves.toBe(true); expect(logger.error).toHaveBeenCalledWith(expect.stringContaining('allowing request')); expect(response.setHeader).toHaveBeenCalledWith('X-RateLimit-Remaining', 2); }); + + it('supports enterprise tier dynamic limits', async () => { + const { context, response } = makeContext({ organizationId: 'org-ent', tier: 'enterprise' }); + const guard = makeGuard({ multi: () => chain }, 100); + expect(await guard.canActivate(context as never)).toBe(true); + expect(response.setHeader).toHaveBeenCalledWith('X-RateLimit-Limit', 500); + }); + + it('supports pro tier dynamic limits', async () => { + const { context, response } = makeContext({ organizationId: 'org-pro', tier: 'pro' }); + const guard = makeGuard({ multi: () => chain }, 100); + expect(await guard.canActivate(context as never)).toBe(true); + expect(response.setHeader).toHaveBeenCalledWith('X-RateLimit-Limit', 250); + }); }); diff --git a/src/common/guards/sliding-window-throttler.guard.ts b/src/common/guards/sliding-window-throttler.guard.ts index 5b1630e0..99033114 100644 --- a/src/common/guards/sliding-window-throttler.guard.ts +++ b/src/common/guards/sliding-window-throttler.guard.ts @@ -1,9 +1,4 @@ -import { - CanActivate, - ExecutionContext, - Injectable, - Logger, -} from '@nestjs/common'; +import { CanActivate, ExecutionContext, Injectable, Logger } from '@nestjs/common'; import { Reflector } from '@nestjs/core'; import { ConfigService } from '@nestjs/config'; import { Redis } from 'ioredis'; @@ -46,8 +41,20 @@ export class SlidingWindowThrottlerGuard implements CanActivate { SLIDING_WINDOW_LIMIT_KEY, [context.getHandler(), context.getClass()], ); - const limit = configured?.limit ?? this.defaultLimit; + let limit = configured?.limit ?? this.defaultLimit; const windowSeconds = configured?.windowSeconds ?? this.defaultWindowSeconds; + + const userTier = + (request.user as (AuthenticatedUser & { tier?: string }) | undefined)?.tier ?? + (request as Request & { apiKey?: { tier?: string } }).apiKey?.tier; + if (userTier === 'enterprise') { + limit = Math.max(limit, 500); + } else if (userTier === 'pro') { + limit = Math.max(limit, 250); + } else if (userTier === 'free' || userTier === 'standard') { + limit = Math.min(limit, 100); + } + const key = this.keyFor(request, context); const now = Date.now(); const windowStart = now - windowSeconds * 1000; @@ -68,24 +75,36 @@ export class SlidingWindowThrottlerGuard implements CanActivate { const remaining = Math.max(0, limit - count - 1); response.setHeader('X-RateLimit-Remaining', remaining); if (count >= limit) { - response.setHeader('Retry-After', Math.max(1, Math.ceil((windowStart + windowSeconds * 1000 - now) / 1000))); - throw new DomainException(ErrorCode.RATE_LIMITED, 'Rate limit exceeded', { limit, windowSeconds }); + response.setHeader( + 'Retry-After', + Math.max(1, Math.ceil((windowStart + windowSeconds * 1000 - now) / 1000)), + ); + throw new DomainException(ErrorCode.RATE_LIMITED, 'Rate limit exceeded', { + limit, + windowSeconds, + }); } return true; } catch (error) { if (error instanceof DomainException && error.code === ErrorCode.RATE_LIMITED) throw error; - this.logger.error(`Sliding-window Redis check failed; allowing request: ${(error as Error).message}`); + this.logger.error( + `Sliding-window Redis check failed; allowing request: ${(error as Error).message}`, + ); response.setHeader('X-RateLimit-Remaining', limit); return true; } } - private keyFor(request: Request & { user?: AuthenticatedUser }, context: ExecutionContext): string { + private keyFor( + request: Request & { user?: AuthenticatedUser }, + context: ExecutionContext, + ): string { const scope = this.clientScope(request); - const tier = this.reflector.getAllAndOverride(THROTTLE_TIER_KEY, [ - context.getHandler(), - context.getClass(), - ]) ?? 'api'; + const tier = + this.reflector.getAllAndOverride(THROTTLE_TIER_KEY, [ + context.getHandler(), + context.getClass(), + ]) ?? 'api'; return `rate-limit:${tier}:${scope}:${context.getClass().name}:${context.getHandler().name}`; } @@ -96,6 +115,10 @@ export class SlidingWindowThrottlerGuard implements CanActivate { * finally falls back to the client IP for fully unauthenticated routes. */ private clientScope(request: Request & { user?: AuthenticatedUser }): string { + if (request.user?.isApiKey) { + return `key:${request.user.apiKeyId ?? request.user.id}`; + } + const organizationId = request.user?.organizationId; if (organizationId) { return `org:${organizationId}`; diff --git a/src/common/guards/throttler.guard.spec.ts b/src/common/guards/throttler.guard.spec.ts index 21b295dc..c9a8afc9 100644 --- a/src/common/guards/throttler.guard.spec.ts +++ b/src/common/guards/throttler.guard.spec.ts @@ -9,7 +9,16 @@ import { THROTTLE_TIER_KEY, ThrottleTier } from '../decorators/throttle-tier.dec /** Shape returned by `ThrottlerStorage#increment` (not re-exported by the lib). */ type ThrottlerStorageRecord = Awaited>; -const CONFIG: ThrottlerConfig = { windowSeconds: 60, apiLimit: 120, authLimit: 10 }; +const CONFIG: ThrottlerConfig = { + windowSeconds: 60, + apiLimit: 120, + authLimit: 10, + agentLimit: 300, + webhookLimit: 30, + apiBurst: 10, + authBurst: 3, + webhookBurst: 5, +}; const UNBLOCKED: ThrottlerStorageRecord = { totalHits: 1, @@ -27,7 +36,10 @@ const BLOCKED: ThrottlerStorageRecord = { type MockResponse = { header: ReturnType }; -function buildContext(request: Record = { ip: '203.0.113.7', headers: {} }, response: MockResponse = { header: vi.fn() }) { +function buildContext( + request: Record = { ip: '203.0.113.7', headers: {} }, + response: MockResponse = { header: vi.fn() }, +) { const handler = () => undefined; return { getHandler: () => handler, @@ -82,16 +94,16 @@ describe('AstroidThrottlerGuard', () => { vi.clearAllMocks(); }); - describe('tier routing', () => { - it('ignores the throttler whose name does not match the route tier', async () => { - const { increment, call } = await prepare(); // no tier set -> defaults to 'api' + describe('tier routing — steady-state', () => { + it('ignores the auth throttler on a default api-tier route', async () => { + const { increment, call } = await prepare(); // no tier → 'api' await expect(call(throttlerNamed('auth'))).resolves.toBe(true); expect(increment).not.toHaveBeenCalled(); }); - it('enforces the throttler whose name matches the default `api` tier', async () => { + it('enforces the api throttler on a default api-tier route', async () => { const { increment, call } = await prepare(); await expect(call(throttlerNamed('api'))).resolves.toBe(true); @@ -99,7 +111,7 @@ describe('AstroidThrottlerGuard', () => { expect(increment).toHaveBeenCalledTimes(1); }); - it('enforces only `auth` for routes declared with the auth tier', async () => { + it('enforces only the auth throttler on routes declared with the auth tier', async () => { const { increment, call } = await prepare({ tier: 'auth' }); await expect(call(throttlerNamed('api'))).resolves.toBe(true); @@ -109,6 +121,17 @@ describe('AstroidThrottlerGuard', () => { expect(increment).toHaveBeenCalledTimes(1); }); + it('enforces only the webhook throttler on routes declared with the webhook tier', async () => { + const { increment, call } = await prepare({ tier: 'webhook' }); + + await expect(call(throttlerNamed('api'))).resolves.toBe(true); + await expect(call(throttlerNamed('auth'))).resolves.toBe(true); + expect(increment).not.toHaveBeenCalled(); + + await expect(call(throttlerNamed('webhook'))).resolves.toBe(true); + expect(increment).toHaveBeenCalledTimes(1); + }); + it('passes the resolved tier limits down to the storage', async () => { const { increment, call } = await prepare({ tier: 'auth' }); @@ -124,6 +147,48 @@ describe('AstroidThrottlerGuard', () => { }); }); + describe('tier routing — burst throttlers', () => { + it('fires the api-burst throttler on api-tier routes (base tier matches)', async () => { + const { increment, call } = await prepare(); + + await expect(call(throttlerNamed('api-burst'))).resolves.toBe(true); + + expect(increment).toHaveBeenCalledTimes(1); + }); + + it('does not fire the api-burst throttler on auth-tier routes', async () => { + const { increment, call } = await prepare({ tier: 'auth' }); + + await expect(call(throttlerNamed('api-burst'))).resolves.toBe(true); + + expect(increment).not.toHaveBeenCalled(); + }); + + it('fires the auth-burst throttler on auth-tier routes', async () => { + const { increment, call } = await prepare({ tier: 'auth' }); + + await expect(call(throttlerNamed('auth-burst'))).resolves.toBe(true); + + expect(increment).toHaveBeenCalledTimes(1); + }); + + it('fires the webhook-burst throttler on webhook-tier routes', async () => { + const { increment, call } = await prepare({ tier: 'webhook' }); + + await expect(call(throttlerNamed('webhook-burst'))).resolves.toBe(true); + + expect(increment).toHaveBeenCalledTimes(1); + }); + + it('does not fire the webhook-burst throttler on api-tier routes', async () => { + const { increment, call } = await prepare(); // api tier + + await expect(call(throttlerNamed('webhook-burst'))).resolves.toBe(true); + + expect(increment).not.toHaveBeenCalled(); + }); + }); + describe('tracking', () => { it('falls back to the client IP for anonymous requests', async () => { const { guard } = await prepare(); @@ -181,6 +246,24 @@ describe('AstroidThrottlerGuard', () => { ); }); + it('throws a 429 for auth-tier routes when blocked', async () => { + const { call } = await prepare({ + tier: 'auth', + increment: vi.fn().mockResolvedValue(BLOCKED), + }); + + await expect(call(throttlerNamed('auth'))).rejects.toMatchObject({ status: 429 }); + }); + + it('throws a 429 for webhook-tier routes when blocked', async () => { + const { call } = await prepare({ + tier: 'webhook', + increment: vi.fn().mockResolvedValue(BLOCKED), + }); + + await expect(call(throttlerNamed('webhook'))).rejects.toMatchObject({ status: 429 }); + }); + it('exposes getStatus() so the exception filter can render the 429 envelope', async () => { const { call } = await prepare({ increment: vi.fn().mockResolvedValue(BLOCKED) }); diff --git a/src/common/guards/throttler.guard.ts b/src/common/guards/throttler.guard.ts index 2d3fc008..863353e4 100644 --- a/src/common/guards/throttler.guard.ts +++ b/src/common/guards/throttler.guard.ts @@ -2,26 +2,32 @@ import { Injectable } from '@nestjs/common'; import { ThrottlerGuard, ThrottlerRequest } from '@nestjs/throttler'; import { Request } from 'express'; import { AuthenticatedUser } from '../interfaces/authenticated-user.interface'; -import { - THROTTLE_TIER_KEY, - ThrottleTier, -} from '../decorators/throttle-tier.decorator'; +import { THROTTLE_TIER_KEY, ThrottleTier } from '../decorators/throttle-tier.decorator'; /** - * Rate-limit guard with two tiers. Every route is evaluated against both named - * throttlers ('api' = 120/min, 'auth' = 10/min by default), but each throttler - * only counts a request when its name matches the route's tier — so the auth - * endpoints (marked `@ThrottleTierDecorator('auth')`) get the stricter limit - * while everything else falls back to the `api` tier. + * Rate-limit guard with per-tier steady-state and burst throttlers. + * + * Each route is evaluated against every registered named throttler, but a + * throttler fires only when its name matches the route's declared tier: + * + * - A throttler named `'api'` fires only on `api`-tier routes. + * - A throttler named `'api-burst'` fires only on `api`-tier routes + * (the `-burst` suffix is stripped for comparison). + * - Routes without an explicit `@ThrottleTierDecorator` default to `api`. + * + * This means auth endpoints (marked `@ThrottleTierDecorator('auth')`) get the + * stricter steady-state limit **and** the tighter burst limit, while everything + * else is governed by the `api` pair. * * The counter is scoped to the authenticated organization, falling back to the - * client IP for anonymous auth endpoints. + * client IP for anonymous requests (e.g. auth endpoints before login). */ @Injectable() export class AstroidThrottlerGuard extends ThrottlerGuard { /** - * Enforce a named throttler only when it matches the route's declared tier. - * Routes without an explicit tier default to `api`. + * Enforce a named throttler only when its base tier matches the route's + * declared tier. The base tier of `'api-burst'` is `'api'`, so the burst + * throttler fires on the same set of routes as its steady-state counterpart. */ protected async handleRequest(requestProps: ThrottlerRequest): Promise { const { context, throttler } = requestProps; @@ -31,8 +37,11 @@ export class AstroidThrottlerGuard extends ThrottlerGuard { context.getClass(), ]) ?? 'api'; + // Strip the optional `-burst` suffix to get the base tier name. + const throttlerBaseTier = throttler.name?.replace(/-burst$/, '') as ThrottleTier | undefined; + // This named throttler does not govern this route's tier — do not count it. - if (throttler.name !== routeTier) { + if (throttlerBaseTier !== routeTier) { return true; } @@ -40,7 +49,21 @@ export class AstroidThrottlerGuard extends ThrottlerGuard { } protected async getTracker(req: Record): Promise { - const request = req as unknown as Request & { user?: AuthenticatedUser }; + const request = req as unknown as Request & { + user?: AuthenticatedUser; + apiKey?: { id: string }; + }; + const apiKeyHeader = request.headers['x-api-key']; + const apiKeyId = + request.apiKey?.id ?? + (Array.isArray(apiKeyHeader) ? apiKeyHeader[0] : apiKeyHeader); + if (apiKeyId) { + return `apikey:${apiKeyId}`; + } + const sub = request.user?.sub ?? request.user?.id; + if (sub) { + return `user:${sub}`; + } const org = request.user?.organizationId; if (org) { return `org:${org}`; diff --git a/src/common/helpers/pagination.integration.spec.ts b/src/common/helpers/pagination.integration.spec.ts new file mode 100644 index 00000000..4cc14a30 --- /dev/null +++ b/src/common/helpers/pagination.integration.spec.ts @@ -0,0 +1,128 @@ +import { afterAll, beforeAll, describe, expect, it, vi } from 'vitest'; +import { Controller, Get, INestApplication, Logger, Query } from '@nestjs/common'; +import { Test } from '@nestjs/testing'; +import { + buildPaginationMeta, + PaginationQuery, + paginationQuerySchema, + toPrismaPagination, +} from './pagination'; +import { ZodValidationPipe } from '../pipes/zod-validation.pipe'; +import { Paginated } from '../interfaces/api-response.interface'; +import { ResponseInterceptor } from '../interceptors/response.interceptor'; +import { AllExceptionsFilter } from '../filters/all-exceptions.filter'; + +/** + * End-to-end check of the list-endpoint pagination contract over real HTTP: + * query parsing (ZodValidationPipe), the Prisma skip/take mapping, the success + * envelope + X-Total-Count header (ResponseInterceptor) and the 400 path + * (AllExceptionsFilter). The repository is an in-memory table of 120 rows that + * honours `skip`/`take` exactly like Prisma's `findMany`. + */ + +const ROWS = Array.from({ length: 120 }, (_, i) => ({ id: i + 1 })); + +type ListBody = { + success: boolean; + data: { id: number }[]; + meta: Record; +}; + +@Controller('resources') +class ResourceController { + @Get() + list(@Query(new ZodValidationPipe(paginationQuerySchema)) query: PaginationQuery) { + const { skip, take } = toPrismaPagination(query, ['createdAt']); + return new Paginated(ROWS.slice(skip, skip + take), buildPaginationMeta(ROWS.length, query)); + } +} + +describe('List pagination (integration)', () => { + let app: INestApplication; + let baseUrl: string; + + beforeAll(async () => { + vi.spyOn(Logger.prototype, 'warn').mockImplementation(() => undefined); + const moduleRef = await Test.createTestingModule({ controllers: [ResourceController] }).compile(); + app = moduleRef.createNestApplication({ logger: false }); + app.useGlobalInterceptors(new ResponseInterceptor()); + app.useGlobalFilters(new AllExceptionsFilter()); + await app.listen(0, '127.0.0.1'); + baseUrl = `${await app.getUrl()}/resources`; + }); + + afterAll(async () => { + await app.close(); + }); + + async function get(query = '') { + const res = await fetch(`${baseUrl}${query}`); + return { res, body: (await res.json()) as ListBody }; + } + + it('returns the first 50 rows with total metadata and header by default', async () => { + const { res, body } = await get(); + + expect(res.status).toBe(200); + expect(res.headers.get('x-total-count')).toBe('120'); + expect(body.data).toHaveLength(50); + expect(body.data[0].id).toBe(1); + expect(body.meta).toEqual({ + offset: 0, + page: 1, + limit: 50, + total: 120, + totalPages: 3, + hasNext: true, + hasPrev: false, + }); + }); + + it('returns the requested offset/limit slice', async () => { + const { res, body } = await get('?offset=30&limit=10'); + + expect(res.status).toBe(200); + expect(body.data.map((row) => row.id)).toEqual([ + 31, 32, 33, 34, 35, 36, 37, 38, 39, 40, + ]); + expect(body.meta).toMatchObject({ offset: 30, limit: 10, hasPrev: true, hasNext: true }); + }); + + it('returns a short final slice and no next page at the end', async () => { + const { body } = await get('?offset=100&limit=50'); + + expect(body.data).toHaveLength(20); + expect(body.data[19].id).toBe(120); + expect(body.meta.hasNext).toBe(false); + }); + + it('returns an empty slice past the end instead of failing', async () => { + const { res, body } = await get('?offset=500'); + + expect(res.status).toBe(200); + expect(body.data).toEqual([]); + expect(res.headers.get('x-total-count')).toBe('120'); + }); + + it('allows the maximum limit of 200', async () => { + const { res, body } = await get('?limit=200'); + + expect(res.status).toBe(200); + expect(body.data).toHaveLength(120); + }); + + it.each([ + ['limit above the cap', '?limit=201'], + ['negative offset', '?offset=-1'], + ['negative limit', '?limit=-10'], + ['zero limit', '?limit=0'], + ['non-numeric limit', '?limit=abc'], + ['fractional offset', '?offset=2.5'], + ['offset and page together', '?offset=10&page=2'], + ])('rejects a %s with 400 Bad Request', async (_label, query) => { + const { res } = await get(query); + + expect(res.status).toBe(400); + expect(res.headers.get('x-total-count')).toBeNull(); + }); +}); diff --git a/src/common/helpers/pagination.spec.ts b/src/common/helpers/pagination.spec.ts new file mode 100644 index 00000000..62f9c901 --- /dev/null +++ b/src/common/helpers/pagination.spec.ts @@ -0,0 +1,116 @@ +import { describe, expect, it } from 'vitest'; +import { + buildPaginationMeta, + DEFAULT_PAGE_LIMIT, + MAX_PAGE_LIMIT, + paginationQuerySchema, + toPrismaPagination, +} from './pagination'; + +describe('paginationQuerySchema', () => { + it('applies offset 0 and limit 50 when no bounds are supplied', () => { + const query = paginationQuerySchema.parse({}); + + expect(query).toMatchObject({ offset: 0, page: 1, limit: DEFAULT_PAGE_LIMIT }); + expect(DEFAULT_PAGE_LIMIT).toBe(50); + }); + + it('coerces string query values into numbers', () => { + const query = paginationQuerySchema.parse({ offset: '100', limit: '25' }); + + expect(query).toMatchObject({ offset: 100, limit: 25, page: 5 }); + }); + + it('derives the offset from a page number', () => { + const query = paginationQuerySchema.parse({ page: '3', limit: '20' }); + + expect(query).toMatchObject({ offset: 40, page: 3, limit: 20 }); + }); + + it('accepts a limit equal to the 200 cap', () => { + expect(MAX_PAGE_LIMIT).toBe(200); + expect(paginationQuerySchema.parse({ limit: '200' }).limit).toBe(200); + }); + + it.each([ + ['a limit above the cap', { limit: '201' }], + ['a zero limit', { limit: '0' }], + ['a negative limit', { limit: '-5' }], + ['a negative offset', { offset: '-1' }], + ['a fractional offset', { offset: '1.5' }], + ['a non-numeric limit', { limit: 'abc' }], + ['a non-numeric offset', { offset: 'ten' }], + ['a zero page', { page: '0' }], + ])('rejects %s', (_label, input) => { + expect(paginationQuerySchema.safeParse(input).success).toBe(false); + }); + + it('rejects offset and page supplied together', () => { + const result = paginationQuerySchema.safeParse({ offset: '10', page: '2' }); + + expect(result.success).toBe(false); + expect(result.error?.issues[0]).toMatchObject({ + path: ['offset'], + message: 'Provide either offset or page, not both', + }); + }); +}); + +describe('toPrismaPagination', () => { + it('maps offset and limit onto skip and take', () => { + const query = paginationQuerySchema.parse({ offset: '120', limit: '40' }); + + expect(toPrismaPagination(query, ['createdAt'])).toEqual({ + skip: 120, + take: 40, + orderBy: { createdAt: 'desc' }, + }); + }); + + it('falls back to createdAt for a sort field outside the allow-list', () => { + const query = paginationQuerySchema.parse({ sort: 'passwordHash; DROP TABLE users' }); + + expect(toPrismaPagination(query, ['name', 'createdAt']).orderBy).toEqual({ createdAt: 'desc' }); + }); + + it('keeps an allow-listed sort field and direction', () => { + const query = paginationQuerySchema.parse({ sort: 'name', order: 'asc' }); + + expect(toPrismaPagination(query, ['name', 'createdAt']).orderBy).toEqual({ name: 'asc' }); + }); +}); + +describe('buildPaginationMeta', () => { + it('reports the slice position and totals', () => { + const meta = buildPaginationMeta(120, paginationQuerySchema.parse({ offset: '50', limit: '50' })); + + expect(meta).toEqual({ + offset: 50, + page: 2, + limit: 50, + total: 120, + totalPages: 3, + hasNext: true, + hasPrev: true, + }); + }); + + it('has no next slice on the last page', () => { + const meta = buildPaginationMeta(120, paginationQuerySchema.parse({ offset: '100', limit: '50' })); + + expect(meta.hasNext).toBe(false); + expect(meta.hasPrev).toBe(true); + }); + + it('computes hasNext/hasPrev from offsets that are not page-aligned', () => { + const meta = buildPaginationMeta(60, paginationQuerySchema.parse({ offset: '5', limit: '50' })); + + expect(meta).toMatchObject({ offset: 5, page: 1, hasNext: true, hasPrev: true }); + }); + + it('handles an empty result set', () => { + const meta = buildPaginationMeta(0, paginationQuerySchema.parse({})); + + expect(meta).toMatchObject({ total: 0, totalPages: 0, hasNext: false, hasPrev: false }); + }); +}); diff --git a/src/common/helpers/pagination.ts b/src/common/helpers/pagination.ts index 4f1a5960..c0725f50 100644 --- a/src/common/helpers/pagination.ts +++ b/src/common/helpers/pagination.ts @@ -1,15 +1,39 @@ import { z } from 'zod'; import { PaginationMeta } from '../interfaces/api-response.interface'; -/** Standard query parameters supported by every list endpoint. */ -export const paginationQuerySchema = z.object({ - page: z.coerce.number().int().positive().default(1), - limit: z.coerce.number().int().positive().max(100).default(20), - sort: z.string().default('createdAt'), - order: z.enum(['asc', 'desc']).default('desc'), - search: z.string().optional(), - filter: z.string().optional(), -}); +/** Page size applied when a list request omits `limit`. */ +export const DEFAULT_PAGE_LIMIT = 50; + +/** Hard upper bound on `limit`; larger values are rejected with 400. */ +export const MAX_PAGE_LIMIT = 200; + +/** + * Standard query parameters supported by every list endpoint. + * + * Clients page either by `offset` (row offset, preferred) or by `page` + * (1-based page number); supplying both is rejected. After parsing, both + * `offset` and `page` are always populated so services and metadata builders + * never need to care which one the client used. Negative, non-integer or + * out-of-range values fail validation and surface as 400 Bad Request. + */ +export const paginationQuerySchema = z + .object({ + offset: z.coerce.number().int().nonnegative().optional(), + page: z.coerce.number().int().positive().optional(), + limit: z.coerce.number().int().positive().max(MAX_PAGE_LIMIT).default(DEFAULT_PAGE_LIMIT), + sort: z.string().default('createdAt'), + order: z.enum(['asc', 'desc']).default('desc'), + search: z.string().optional(), + filter: z.string().optional(), + }) + .refine((query) => query.offset === undefined || query.page === undefined, { + message: 'Provide either offset or page, not both', + path: ['offset'], + }) + .transform((query) => { + const offset = query.offset ?? ((query.page ?? 1) - 1) * query.limit; + return { ...query, offset, page: Math.floor(offset / query.limit) + 1 }; + }); export type PaginationQuery = z.infer; @@ -19,25 +43,37 @@ export interface PrismaPagination { orderBy: Record; } -/** Translates validated pagination query params into Prisma arguments. */ -export function toPrismaPagination(query: PaginationQuery, allowedSortFields: string[]): PrismaPagination { +/** + * Translates validated pagination query params into Prisma arguments. The + * bounds are passed as bound `skip`/`take` values (never interpolated into + * SQL), and `sort` is restricted to an allow-list of columns. + */ +export function toPrismaPagination( + query: Pick, + allowedSortFields: string[], +): PrismaPagination { const sort = allowedSortFields.includes(query.sort) ? query.sort : 'createdAt'; return { - skip: (query.page - 1) * query.limit, + skip: query.offset, take: query.limit, orderBy: { [sort]: query.order }, }; } /** Builds pagination metadata for the response envelope. */ -export function buildPaginationMeta(total: number, page: number, limit: number): PaginationMeta { +export function buildPaginationMeta( + total: number, + query: Pick, +): PaginationMeta { + const { offset, page, limit } = query; const totalPages = limit > 0 ? Math.ceil(total / limit) : 0; return { + offset, page, limit, total, totalPages, - hasNext: page < totalPages, - hasPrev: page > 1, + hasNext: offset + limit < total, + hasPrev: offset > 0, }; } diff --git a/src/common/helpers/request-id.ts b/src/common/helpers/request-id.ts new file mode 100644 index 00000000..86dfaa35 --- /dev/null +++ b/src/common/helpers/request-id.ts @@ -0,0 +1,9 @@ +import { randomUUID } from 'crypto'; + +const REQUEST_ID_PATTERN = /^[A-Za-z0-9][A-Za-z0-9._:-]{0,127}$/; + +export function resolveRequestId(value: unknown): string { + return typeof value === 'string' && REQUEST_ID_PATTERN.test(value) + ? value + : randomUUID(); +} \ No newline at end of file diff --git a/src/common/index.ts b/src/common/index.ts index 83c96c6f..e91b2fbf 100644 --- a/src/common/index.ts +++ b/src/common/index.ts @@ -14,6 +14,8 @@ export * from './decorators/roles.decorator'; export * from './decorators/scopes.decorator'; export * from './decorators/public.decorator'; export * from './decorators/throttle-tier.decorator'; +export * from './decorators/audit-log.decorator'; +export * from './decorators/skip-audit.decorator'; export * from './decorators/api-envelope.decorator'; export * from './guards/jwt-auth.guard'; export * from './guards/api-key.guard'; @@ -21,6 +23,7 @@ export * from './guards/api-key-auth.guard'; export * from './guards/scopes.guard'; export * from './guards/roles.guard'; export * from './guards/throttler.guard'; +export * from './guards/agent-throttler.guard'; export * from './guards/sliding-window-throttler.guard'; export * from './interceptors/horizon-circuit-breaker.interceptor'; export * from './encryption'; diff --git a/src/common/interceptors/audit-log.interceptor.spec.ts b/src/common/interceptors/audit-log.interceptor.spec.ts index 5570741c..cb8cdc0b 100644 --- a/src/common/interceptors/audit-log.interceptor.spec.ts +++ b/src/common/interceptors/audit-log.interceptor.spec.ts @@ -1,11 +1,15 @@ import { EventEmitter } from 'events'; import { describe, it, expect, vi, afterEach } from 'vitest'; import { ExecutionContext, Logger } from '@nestjs/common'; +import { Reflector } from '@nestjs/core'; import { Observable, of } from 'rxjs'; import { AuditService } from '../../modules/audit/audit.service'; +import { AUDIT_LOG_KEY, AuditLogOptions } from '../decorators/audit-log.decorator'; +import { IS_SKIP_AUDIT_KEY } from '../decorators/skip-audit.decorator'; import { AuditLogInterceptor, + hashPayload, isSensitiveKey, maskSensitiveData, REDACTED_VALUE, @@ -42,14 +46,43 @@ function createContext( getRequest: () => request, getResponse: () => response, }), + getHandler: () => undefined, getClass: () => controller, } as unknown as ExecutionContext; } +/** Reflector stub that only answers the two metadata keys the interceptor reads. */ +function makeReflector(stub: { audit?: AuditLogOptions; skip?: boolean } = { audit: {} }): Reflector { + return { + getAllAndOverride: vi.fn((key: string) => { + if (key === AUDIT_LOG_KEY) return stub.audit; + if (key === IS_SKIP_AUDIT_KEY) return stub.skip; + return undefined; + }), + } as unknown as Reflector; +} + +interface InterceptorOptions { + audit?: AuditLogOptions; + skip?: boolean; + trustProxy?: boolean; +} + +function makeInterceptor( + record: ReturnType, + options: InterceptorOptions = {}, +): AuditLogInterceptor { + const audit = Object.prototype.hasOwnProperty.call(options, 'audit') ? options.audit : {}; + const { skip = false, trustProxy = false } = options; + const auditService = { record } as unknown as AuditService; + const config = { get: vi.fn().mockReturnValue(trustProxy) } as never; + return new AuditLogInterceptor(auditService, config, makeReflector({ audit, skip })); +} + /** Subscribes so the handler runs, emits `finish`, then waits for the async audit write. */ async function runRequest( interceptor: AuditLogInterceptor, - context: ReturnType, + context: ExecutionContext, response: EventEmitter & { statusCode: number }, ): Promise { const observable = interceptor.intercept(context, { @@ -63,10 +96,20 @@ async function runRequest( await new Promise((resolve) => setTimeout(resolve, 0)); } -function makeInterceptor(record: ReturnType, trustProxy = false): AuditLogInterceptor { - const auditService = { record } as unknown as AuditService; - const config = { get: vi.fn().mockReturnValue(trustProxy) } as never; - return new AuditLogInterceptor(auditService, config); +const OWNER = { id: 'user-1', organizationId: 'org-1', email: 'admin@example.com', role: 'ADMIN' }; + +function baseRequest(overrides: Partial = {}): MockRequest { + return { + method: 'PATCH', + path: '/api/v1/policies/pol-123', + headers: { 'user-agent': 'test-agent' }, + params: { id: 'pol-123' }, + query: {}, + body: { name: 'Daily limit' }, + ip: '127.0.0.1', + user: OWNER, + ...overrides, + }; } describe('AuditLogInterceptor', () => { @@ -74,21 +117,17 @@ describe('AuditLogInterceptor', () => { vi.restoreAllMocks(); }); - describe('payload extraction', () => { - it('captures user id, method, path, client IP, body and response status code', async () => { + describe('decorated endpoints', () => { + it('persists actor, method, path, IP, masked body, payload hash and status code', async () => { const record = vi.fn().mockResolvedValue(undefined); - const interceptor = makeInterceptor(record, true); + const interceptor = makeInterceptor(record, { trustProxy: true }); - const request: MockRequest = { - method: 'PATCH', - path: '/api/v1/policies/pol-123', + const request = baseRequest({ headers: { 'user-agent': 'test-agent', 'x-forwarded-for': '203.0.113.5' }, params: { id: 'pol-123' }, - query: {}, body: { name: 'Daily limit', configuration: { maxAmount: 100 } }, ip: '::1', - user: { id: 'user-1', organizationId: 'org-1', email: 'admin@example.com', role: 'ADMIN' }, - }; + }); const response = createMockResponse(201); const context = createContext(request, response); @@ -104,30 +143,65 @@ describe('AuditLogInterceptor', () => { entityId: 'pol-123', ipAddress: '203.0.113.5', device: 'test-agent', - newValue: { + newValue: expect.objectContaining({ path: '/api/v1/policies/pol-123', body: { name: 'Daily limit', configuration: { maxAmount: 100 } }, + payloadHash: hashPayload({ name: 'Daily limit', configuration: { maxAmount: 100 } }), + actor: { type: 'USER', id: 'user-1' }, statusCode: 201, durationMs: expect.any(Number), - }, + }), }), ); }); - it('captures agent identity and defaults to the socket IP when no proxy header is trusted', async () => { + it('uses the @AuditLog() metadata for the semantic action and entity', async () => { const record = vi.fn().mockResolvedValue(undefined); - const interceptor = makeInterceptor(record, false); + const interceptor = makeInterceptor(record, { + audit: { action: 'POLICY_OVERRIDDEN', entity: 'SpendingPolicy' }, + }); - const request: MockRequest = { + const request = baseRequest(); + const response = createMockResponse(200); + const context = createContext(request, response); + + await runRequest(interceptor, context, response); + + expect(record).toHaveBeenCalledWith( + expect.objectContaining({ action: 'POLICY_OVERRIDDEN', entity: 'SpendingPolicy' }), + ); + }); + + it('audits a decorated read-only GET, which is not logged by default', async () => { + const record = vi.fn().mockResolvedValue(undefined); + const interceptor = makeInterceptor(record); + + const request = baseRequest({ method: 'GET', path: '/api/v1/policies/pol-123' }); + const response = createMockResponse(200); + const context = createContext(request, response); + + await runRequest(interceptor, context, response); + + expect(record).toHaveBeenCalledTimes(1); + expect(record).toHaveBeenCalledWith(expect.objectContaining({ action: 'GET' })); + }); + + it('records the acting agent as the actor when no human user is present', async () => { + const record = vi.fn().mockResolvedValue(undefined); + const interceptor = makeInterceptor(record); + + const request = baseRequest({ method: 'POST', path: '/api/v1/wallets/wal-1/rotate', - headers: { 'user-agent': 'AgentRunner/1.0', 'x-agent-id': 'agent-9' }, + headers: { + 'user-agent': 'AgentRunner/1.0', + 'x-agent-id': 'agent-9', + 'x-organization-id': 'org-1', + }, params: { id: 'wal-1' }, - query: {}, - body: { agentId: 'agent-9', newLabel: 'ops' }, - ip: '10.0.0.7', - user: { id: 'user-2', organizationId: 'org-2', email: 'a@b.com', role: 'DEVELOPER' }, - }; + body: { newLabel: 'ops' }, + user: undefined, + }); const response = createMockResponse(200); const context = createContext(request, response, WalletController); @@ -135,33 +209,25 @@ describe('AuditLogInterceptor', () => { expect(record).toHaveBeenCalledWith( expect.objectContaining({ - userId: 'user-2', + userId: null, action: 'POST', entity: 'Wallet', - ipAddress: '10.0.0.7', - newValue: expect.objectContaining({ agentId: 'agent-9', statusCode: 200 }), + newValue: expect.objectContaining({ + agentId: 'agent-9', + actor: { type: 'AGENT', id: 'agent-9' }, + }), }), ); }); - it('records the execution duration of the handler alongside the response status', async () => { + it('records the handler execution duration alongside the response status', async () => { const record = vi.fn().mockResolvedValue(undefined); const interceptor = makeInterceptor(record); - const request: MockRequest = { - method: 'POST', - path: '/api/v1/policies', - headers: { 'user-agent': 'test' }, - params: {}, - query: {}, - body: { name: 'Daily limit' }, - ip: '127.0.0.1', - user: { id: 'user-1', organizationId: 'org-1', email: 'a@b.com', role: 'ADMIN' }, - }; + const request = baseRequest({ method: 'POST', path: '/api/v1/policies', body: {} }); const response = createMockResponse(201); const context = createContext(request, response); - // Simulate a handler that takes a measurable amount of time. const observable = interceptor.intercept(context, { handle: () => new Observable((subscriber) => { @@ -180,13 +246,58 @@ describe('AuditLogInterceptor', () => { const { newValue } = record.mock.calls[0][0]; expect(newValue.durationMs).toBeGreaterThanOrEqual(20); - expect(newValue.durationMs).toBeLessThan(5_000); expect(newValue.statusCode).toBe(201); }); }); - describe('sensitive data masking', () => { - it('redacts sensitive fields, preserves non-sensitive ones and does not mutate the original body', async () => { + describe('scope filtering', () => { + it('never audits an undecorated route, even a state-mutating one', async () => { + const record = vi.fn().mockResolvedValue(undefined); + const interceptor = makeInterceptor(record, { audit: undefined }); + + const request = baseRequest({ method: 'DELETE' }); + const response = createMockResponse(204); + const context = createContext(request, response); + + await runRequest(interceptor, context, response); + + expect(record).not.toHaveBeenCalled(); + }); + + it('honours @SkipAudit() even when the route is decorated', async () => { + const record = vi.fn().mockResolvedValue(undefined); + const interceptor = makeInterceptor(record, { skip: true }); + + const request = baseRequest({ method: 'DELETE' }); + const response = createMockResponse(204); + const context = createContext(request, response); + + await runRequest(interceptor, context, response); + + expect(record).not.toHaveBeenCalled(); + }); + + it('skips decorated routes with no organization context (e.g. public routes)', async () => { + const record = vi.fn().mockResolvedValue(undefined); + const interceptor = makeInterceptor(record); + + const request = baseRequest({ + method: 'POST', + path: '/api/v1/auth/login', + body: { email: 'a@b.com', password: 'secret' }, + user: undefined, + }); + const response = createMockResponse(200); + const context = createContext(request, response); + + await runRequest(interceptor, context, response); + + expect(record).not.toHaveBeenCalled(); + }); + }); + + describe('sensitive data sanitization', () => { + it('redacts secrets, preserves safe fields and never mutates the original body', async () => { const record = vi.fn().mockResolvedValue(undefined); const interceptor = makeInterceptor(record); @@ -196,19 +307,11 @@ describe('AuditLogInterceptor', () => { apiKey: 'abc123', token: 'jwt-token', passkey: 'cred-1', + privateKey: 'SDFJKL-seed', webhook: { signature: 'sig-here', url: 'https://example.com/hook' }, nested: { refreshToken: 'rt-1', note: 'keep me' }, }; - const request: MockRequest = { - method: 'PUT', - path: '/api/v1/developer/keys', - headers: { 'user-agent': 'test' }, - params: {}, - query: {}, - body: originalBody, - ip: '127.0.0.1', - user: { id: 'user-1', organizationId: 'org-1', email: 'a@b.com', role: 'ADMIN' }, - }; + const request = baseRequest({ method: 'PUT', body: originalBody }); const response = createMockResponse(200); const context = createContext(request, response); @@ -221,34 +324,13 @@ describe('AuditLogInterceptor', () => { apiKey: REDACTED_VALUE, token: REDACTED_VALUE, passkey: REDACTED_VALUE, + privateKey: REDACTED_VALUE, webhook: { signature: REDACTED_VALUE, url: 'https://example.com/hook' }, nested: { refreshToken: REDACTED_VALUE, note: 'keep me' }, }); // The original request body must be untouched. - expect(originalBody).toEqual({ - username: 'john', - password: 'secret-pass', - apiKey: 'abc123', - token: 'jwt-token', - passkey: 'cred-1', - webhook: { signature: 'sig-here', url: 'https://example.com/hook' }, - nested: { refreshToken: 'rt-1', note: 'keep me' }, - }); - }); - - it('masks sensitive keys case-insensitively and across separators', () => { - expect(isSensitiveKey('password')).toBe(true); - expect(isSensitiveKey('PasswordHash')).toBe(true); - expect(isSensitiveKey('apiKey')).toBe(true); - expect(isSensitiveKey('api_key')).toBe(true); - expect(isSensitiveKey('x-api-key')).toBe(true); - expect(isSensitiveKey('accessToken')).toBe(true); - expect(isSensitiveKey('passkey')).toBe(true); - expect(isSensitiveKey('signature')).toBe(true); - expect(isSensitiveKey('privateKey')).toBe(true); - expect(isSensitiveKey('username')).toBe(false); - expect(isSensitiveKey('name')).toBe(false); - expect(isSensitiveKey('amount')).toBe(false); + expect(originalBody.password).toBe('secret-pass'); + expect(originalBody.apiKey).toBe('abc123'); }); it('masks sensitive entries inside arrays', () => { @@ -261,81 +343,66 @@ describe('AuditLogInterceptor', () => { { label: 'backup', apiKey: REDACTED_VALUE }, ]); }); - }); - describe('audit failure handling', () => { - it('does not crash the request when audit persistence fails and logs the error', async () => { - const loggerError = vi - .spyOn(Logger.prototype, 'error') - .mockImplementation(() => undefined); - const record = vi.fn().mockRejectedValue(new Error('database unreachable')); - const interceptor = makeInterceptor(record); + it('detects sensitive keys case-insensitively and across separators', () => { + expect(isSensitiveKey('password')).toBe(true); + expect(isSensitiveKey('PasswordHash')).toBe(true); + expect(isSensitiveKey('apiKey')).toBe(true); + expect(isSensitiveKey('api_key')).toBe(true); + expect(isSensitiveKey('x-api-key')).toBe(true); + expect(isSensitiveKey('accessToken')).toBe(true); + expect(isSensitiveKey('privateKey')).toBe(true); + expect(isSensitiveKey('username')).toBe(false); + expect(isSensitiveKey('amount')).toBe(false); + }); + }); - const request: MockRequest = { - method: 'DELETE', - path: '/api/v1/policies/pol-1', - headers: { 'user-agent': 'test' }, - params: { id: 'pol-1' }, - query: {}, - body: {}, - ip: '127.0.0.1', - user: { id: 'user-1', organizationId: 'org-1', email: 'a@b.com', role: 'ADMIN' }, - }; - const response = createMockResponse(204); - const context = createContext(request, response); + describe('payload hashing', () => { + it('is deterministic for identical payloads', () => { + expect(hashPayload({ a: 1, b: 'two' })).toBe(hashPayload({ a: 1, b: 'two' })); + }); - // Must resolve — the failed audit write must not surface to the caller. - await runRequest(interceptor, context, response); + it('changes when the payload changes', () => { + expect(hashPayload({ amount: 10 })).not.toBe(hashPayload({ amount: 11 })); + }); - expect(record).toHaveBeenCalledTimes(1); - expect(loggerError).toHaveBeenCalledWith( - expect.stringContaining('Failed to write audit log for DELETE Policy'), - ); + it('hashes an absent body without throwing', () => { + expect(hashPayload(undefined)).toHaveLength(64); }); - }); - describe('scope filtering', () => { - it('does not audit read-only GET requests', async () => { + it('records the hash of the sanitized body, not the raw secret', async () => { const record = vi.fn().mockResolvedValue(undefined); const interceptor = makeInterceptor(record); - const request: MockRequest = { - method: 'GET', - path: '/api/v1/policies', - headers: { 'user-agent': 'test' }, - params: {}, - query: {}, - body: {}, - ip: '127.0.0.1', - user: { id: 'user-1', organizationId: 'org-1', email: 'a@b.com', role: 'ADMIN' }, - }; - const response = createMockResponse(200); + const request = baseRequest({ method: 'POST', body: { apiKey: 'super-secret' } }); + const response = createMockResponse(201); const context = createContext(request, response); await runRequest(interceptor, context, response); - expect(record).not.toHaveBeenCalled(); + const { newValue } = record.mock.calls[0][0]; + expect(newValue.payloadHash).toBe(hashPayload({ apiKey: REDACTED_VALUE })); + expect(JSON.stringify(newValue)).not.toContain('super-secret'); }); + }); - it('skips requests without an organization context (e.g. public routes)', async () => { - const record = vi.fn().mockResolvedValue(undefined); + describe('audit failure handling', () => { + it('does not crash the request when audit persistence fails and logs the error', async () => { + const loggerError = vi.spyOn(Logger.prototype, 'error').mockImplementation(() => undefined); + const record = vi.fn().mockRejectedValue(new Error('database unreachable')); const interceptor = makeInterceptor(record); - const request: MockRequest = { - method: 'POST', - path: '/api/v1/auth/login', - headers: { 'user-agent': 'test' }, - params: {}, - query: {}, - body: { email: 'a@b.com', password: 'secret' }, - ip: '127.0.0.1', - }; - const response = createMockResponse(200); + const request = baseRequest({ method: 'DELETE', path: '/api/v1/policies/pol-1' }); + const response = createMockResponse(204); const context = createContext(request, response); + // Must resolve — the failed audit write must not surface to the caller. await runRequest(interceptor, context, response); - expect(record).not.toHaveBeenCalled(); + expect(record).toHaveBeenCalledTimes(1); + expect(loggerError).toHaveBeenCalledWith( + expect.stringContaining('Failed to write audit log for DELETE Policy'), + ); }); }); }); diff --git a/src/common/interceptors/audit-log.interceptor.ts b/src/common/interceptors/audit-log.interceptor.ts index 781c135f..ffeac671 100644 --- a/src/common/interceptors/audit-log.interceptor.ts +++ b/src/common/interceptors/audit-log.interceptor.ts @@ -5,19 +5,20 @@ import { Logger, NestInterceptor, } from '@nestjs/common'; +import { Reflector } from '@nestjs/core'; import { ConfigService } from '@nestjs/config'; import { Prisma } from '@prisma/client'; +import { createHash } from 'crypto'; import { Request, Response } from 'express'; import { Observable } from 'rxjs'; -import { AuditService } from '../../modules/audit/audit.service'; import { CreateAuditLogData } from '../../modules/audit/audit.repository'; +import { AuditService } from '../../modules/audit/audit.service'; import { getClientIp } from '../../utils/ip.util'; +import { AUDIT_LOG_KEY, AuditLogOptions } from '../decorators/audit-log.decorator'; +import { IS_SKIP_AUDIT_KEY } from '../decorators/skip-audit.decorator'; import { AuthenticatedUser } from '../interfaces/authenticated-user.interface'; -/** HTTP methods whose state-mutating requests are audited. Read-only traffic is skipped. */ -const AUDITED_METHODS = new Set(['POST', 'PUT', 'PATCH', 'DELETE']); - /** Value substituted for sensitive fields before an audit payload is persisted. */ export const REDACTED_VALUE = '[REDACTED]'; @@ -36,6 +37,8 @@ const SENSITIVE_KEY_FRAGMENTS = [ 'apikey', 'privatekey', 'authorization', + 'mnemonic', + 'seedphrase', ]; /** Returns true when a field name denotes sensitive data (e.g. `apiKey`, `accessToken`). */ @@ -71,22 +74,54 @@ function isPlainObject(value: unknown): value is Record { } /** - * Global audit interceptor. Persists a permanent, traceable record of every - * state-mutating request (POST/PUT/PATCH/DELETE) into the existing PostgreSQL - * audit trail through `AuditService`/Prisma. + * Stable SHA-256 fingerprint of a (already sanitized) request payload. + * + * The digest lets an operator prove which payload an action carried without + * duplicating it in the audit trail, and makes tampering detectable: a changed + * body always yields a different hash. + */ +export function hashPayload(value: unknown): string { + let serialized: string; + if (value === undefined) { + serialized = ''; + } else { + try { + serialized = JSON.stringify(value) ?? ''; + } catch { + // Circular or otherwise non-serializable bodies still get a fingerprint. + serialized = String(value); + } + } + return createHash('sha256').update(serialized).digest('hex'); +} + +/** Resolved identity of the principal that triggered the request. */ +interface AuditIdentity { + organizationId: string; + userId: string | null; + agentId?: string; + ipAddress?: string; +} + +/** + * Structured audit interceptor for sensitive agent operations. + * + * Persists a permanent, traceable record for every route (handler or the whole + * controller) decorated with `@AuditLog()`. Undecorated routes — including + * read-only queries — are passed straight through without touching the database, + * which is what makes the logging selective and high-performance. * * Captured per request: - * - authenticated user (or agent) identity + * - the actor: human admin user id, or the acting agent id * - HTTP method, route path and client IP - * - the request body with sensitive fields masked - * - the final response status code - * - the time the handler took to complete, in milliseconds + * - the payload fingerprint (SHA-256 of the sanitized body) + * - the sanitized body itself, with secrets/keys/tokens redacted + * - the final response status code and the handler duration in milliseconds * - * The audit write happens once the response has been fully sent (`finish`), so - * the recorded status code is the real one — including error statuses set by - * the global exception filter. Persistence is fire-and-forget and failures are - * logged but never crash the client request (no strict compliance mode exists - * in this project, so non-blocking is the required behavior). + * The write happens once the response has been fully sent (`finish`), so the + * recorded status code is the real one — including error statuses set by the + * global exception filter. Persistence is fire-and-forget: a failure is logged + * but never breaks the client request. */ @Injectable() export class AuditLogInterceptor implements NestInterceptor { @@ -95,18 +130,25 @@ export class AuditLogInterceptor implements NestInterceptor { constructor( private readonly auditService: AuditService, private readonly config: ConfigService, + private readonly reflector: Reflector, ) {} intercept(context: ExecutionContext, next: CallHandler): Observable { - const http = context.switchToHttp(); - const request = http.getRequest(); - const response = http.getResponse(); + const options = this.reflector.getAllAndOverride(AUDIT_LOG_KEY, [ + context.getHandler(), + context.getClass(), + ]); - // Only state-mutating methods are audited; read-only traffic is skipped. - if (!AUDITED_METHODS.has(request.method)) { + // Selective logging: only routes decorated with @AuditLog() are persisted, + // and an explicit @SkipAudit() always wins. + if (!options || this.isSkipped(context)) { return next.handle(); } + const http = context.switchToHttp(); + const request = http.getRequest(); + const response = http.getResponse(); + // Audit rows are scoped to an organization (required FK on AuditLog). const organizationId = request.user?.organizationId || @@ -114,25 +156,27 @@ export class AuditLogInterceptor implements NestInterceptor { (request.headers['x-organization-id'] as string) || undefined; if (!organizationId) { + this.logger.debug( + `Skipping @AuditLog() route without an organization context: ${request.path}`, + ); return next.handle(); } - const userId = request.user?.id || (request.headers['x-user-id'] as string) || null; - // Same agent-identity resolution chain as AgentTraceInterceptor. - const agentId = - (request.params?.agentId as string) || - (request.body?.agentId as string) || - (request.query?.agentId as string) || - (request.headers['x-agent-id'] as string) || - undefined; - - const trustProxy = this.config.get('app.trustProxy', false); - const ipAddress = - getClientIp(request.ip ?? '', request.headers['x-forwarded-for'] as string, trustProxy) || - undefined; + const identity: AuditIdentity = { + organizationId, + userId: request.user?.id || (request.headers['x-user-id'] as string) || null, + // Same agent-identity resolution chain as AgentTraceInterceptor. + agentId: + (request.params?.agentId as string) || + (request.body?.agentId as string) || + (request.query?.agentId as string) || + (request.headers['x-agent-id'] as string) || + undefined, + ipAddress: this.resolveIp(request), + }; - // Captured before the handler runs so the recorded duration covers the - // full execution time of the route. + // Captured before the handler runs so the recorded duration covers the full + // execution time of the route. const startedAt = Date.now(); response.on('finish', () => { @@ -140,7 +184,8 @@ export class AuditLogInterceptor implements NestInterceptor { this.buildAuditData( request, context, - { organizationId, userId, agentId, ipAddress }, + identity, + options, response.statusCode, Date.now() - startedAt, ), @@ -150,23 +195,50 @@ export class AuditLogInterceptor implements NestInterceptor { return next.handle(); } - /** Builds the audit row, storing the masked body, path and agent id as `newValue`. */ + /** True when the route opted out with `@SkipAudit()`. */ + private isSkipped(context: ExecutionContext): boolean { + return ( + this.reflector.getAllAndOverride(IS_SKIP_AUDIT_KEY, [ + context.getHandler(), + context.getClass(), + ]) === true + ); + } + + /** Resolves the client IP, honouring `x-forwarded-for` only when proxies are trusted. */ + private resolveIp(request: Request): string | undefined { + const trustProxy = this.config.get('app.trustProxy', false); + const forwarded = request.headers['x-forwarded-for'] as string | undefined; + return getClientIp(request.ip ?? '', forwarded, trustProxy) || undefined; + } + + /** Builds the audit row, storing the masked body, path and actor as `newValue`. */ private buildAuditData( request: Request & { user?: AuthenticatedUser }, context: ExecutionContext, - identity: { organizationId: string; userId: string | null; agentId?: string; ipAddress?: string }, + identity: AuditIdentity, + options: AuditLogOptions, statusCode: number, durationMs: number, ): CreateAuditLogData { const body = request.body; const maskedBody = body && typeof body === 'object' ? maskSensitiveData(body) : undefined; + const actor = identity.userId + ? { type: 'USER' as const, id: identity.userId } + : identity.agentId + ? { type: 'AGENT' as const, id: identity.agentId } + : null; const newValue: Prisma.InputJsonValue = { path: request.path, ...(maskedBody !== undefined ? { body: maskedBody } : {}), + // A stable fingerprint of the sanitized payload: proves what was sent + // without persisting the same secrets twice. + payloadHash: hashPayload(maskedBody), // Agent identity is stored here per the existing audit-export convention // (the schema has no dedicated agent column). ...(identity.agentId ? { agentId: identity.agentId } : {}), + ...(actor ? { actor } : {}), statusCode, durationMs, }; @@ -174,8 +246,8 @@ export class AuditLogInterceptor implements NestInterceptor { return { organizationId: identity.organizationId, userId: identity.userId, - action: request.method, - entity: this.resolveEntity(context), + action: options.action ?? request.method, + entity: options.entity ?? this.resolveEntity(context), entityId: (request.params?.id as string) ?? null, newValue, ipAddress: identity.ipAddress, diff --git a/src/common/interceptors/metrics.interceptor.spec.ts b/src/common/interceptors/metrics.interceptor.spec.ts new file mode 100644 index 00000000..43eb7cf4 --- /dev/null +++ b/src/common/interceptors/metrics.interceptor.spec.ts @@ -0,0 +1,182 @@ +import { ExecutionContext, CallHandler } from '@nestjs/common'; +import { of, throwError, Observable } from 'rxjs'; +import { MetricsInterceptor } from './metrics.interceptor'; +import { MetricsService } from '../../modules/metrics/metrics.service'; +import { Request, Response } from 'express'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; + +describe('MetricsInterceptor', () => { + let interceptor: MetricsInterceptor; + let metricsService: { + observeHttpRequest: ReturnType; + }; + + beforeEach(() => { + metricsService = { + observeHttpRequest: vi.fn(), + }; + interceptor = new MetricsInterceptor(metricsService as unknown as MetricsService); + }); + + it('should be defined', () => { + expect(interceptor).toBeDefined(); + }); + + it('should record metrics on successful request', () => { + const context = createMockExecutionContext('GET', '/api/test', 200); + const handler = createMockCallHandler(of({ data: 'success' })); + + interceptor.intercept(context, handler).subscribe(); + + expect(metricsService.observeHttpRequest).toHaveBeenCalledWith( + 'GET', + '/api/test', + 200, + expect.any(Number), + ); + }); + + it('should record metrics on failed request', () => { + const context = createMockExecutionContext('POST', '/api/error', 500); + const handler = createMockCallHandler(throwError(new Error('Test error'))); + + interceptor.intercept(context, handler).subscribe({ + error: () => { + // Expected error + }, + }); + + expect(metricsService.observeHttpRequest).toHaveBeenCalledWith( + 'POST', + '/api/error', + 500, + expect.any(Number), + ); + }); + + it('should skip metrics collection for /metrics endpoint', () => { + const context = createMockExecutionContext('GET', '/metrics', 200); + const handler = createMockCallHandler(of({})); + + interceptor.intercept(context, handler).subscribe(); + + expect(metricsService.observeHttpRequest).not.toHaveBeenCalled(); + }); + + it('should normalize route paths before recording', () => { + const context = createMockExecutionContext('GET', '/api/users/123', 200); + const handler = createMockCallHandler(of({})); + + interceptor.intercept(context, handler).subscribe(); + + expect(metricsService.observeHttpRequest).toHaveBeenCalledWith( + 'GET', + '/api/users/:id', + 200, + expect.any(Number), + ); + }); + + it('should track active request count', () => { + const context1 = createMockExecutionContext('GET', '/api/test1', 200); + const context2 = createMockExecutionContext('GET', '/api/test2', 200); + const handler = createMockCallHandler(of({})); + + expect(interceptor.getActiveRequestCount()).toBe(0); + + const sub1 = interceptor.intercept(context1, handler); + expect(interceptor.getActiveRequestCount()).toBe(1); + + const sub2 = interceptor.intercept(context2, handler); + expect(interceptor.getActiveRequestCount()).toBe(2); + + sub1.subscribe(); + expect(interceptor.getActiveRequestCount()).toBe(1); + + sub2.subscribe(); + expect(interceptor.getActiveRequestCount()).toBe(0); + }); + + it('should not fail request when metrics recording throws error', () => { + metricsService.observeHttpRequest.mockImplementation(() => { + throw new Error('Metrics recording failed'); + }); + + const context = createMockExecutionContext('GET', '/api/test', 200); + const handler = createMockCallHandler(of({ data: 'success' })); + + const result = interceptor.intercept(context, handler); + + // Should complete successfully despite metrics error + expect(() => { + result.subscribe(); + }).not.toThrow(); + }); + + it('should record metrics with different HTTP methods', () => { + const methods = ['GET', 'POST', 'PUT', 'DELETE', 'PATCH'] as const; + + methods.forEach((method) => { + metricsService.observeHttpRequest.mockClear(); + const context = createMockExecutionContext(method, '/api/test', 200); + const handler = createMockCallHandler(of({})); + + interceptor.intercept(context, handler).subscribe(); + + expect(metricsService.observeHttpRequest).toHaveBeenCalledWith( + method, + '/api/test', + 200, + expect.any(Number), + ); + }); + }); + + it('should record metrics with different status codes', () => { + const statusCodes = [200, 201, 204, 400, 401, 403, 404, 500, 503]; + + statusCodes.forEach((statusCode) => { + metricsService.observeHttpRequest.mockClear(); + const context = createMockExecutionContext('GET', '/api/test', statusCode); + const handler = createMockCallHandler(of({})); + + interceptor.intercept(context, handler).subscribe(); + + expect(metricsService.observeHttpRequest).toHaveBeenCalledWith( + 'GET', + '/api/test', + statusCode, + expect.any(Number), + ); + }); + }); +}); + +function createMockExecutionContext( + method: string, + path: string, + statusCode: number, +): ExecutionContext { + const req = { + method, + path, + headers: {}, + } as Partial; + + const res = { + statusCode, + } as Partial; + + return { + switchToHttp: () => ({ + getRequest: () => req as Request, + getResponse: () => res as Response, + }), + } as unknown as ExecutionContext; +} + +function createMockCallHandler(observable: Observable): CallHandler { + return { + handle: () => observable, + }; +} diff --git a/src/common/interceptors/metrics.interceptor.ts b/src/common/interceptors/metrics.interceptor.ts new file mode 100644 index 00000000..2dbc2527 --- /dev/null +++ b/src/common/interceptors/metrics.interceptor.ts @@ -0,0 +1,85 @@ +import { + CallHandler, + ExecutionContext, + Injectable, + NestInterceptor, + Logger, +} from '@nestjs/common'; +import { Observable } from 'rxjs'; +import { tap } from 'rxjs/operators'; +import { Request, Response } from 'express'; +import { MetricsService } from '../../modules/metrics/metrics.service'; +import { normalizeRoutePath } from '../../utils/route-normalizer.util'; + +/** + * NestJS interceptor that collects Prometheus metrics for all HTTP requests. + * + * Records: + * - Request duration histograms (in seconds) + * - Request counters categorized by route, method, and status code + * - Active request gauges (incremented on entry, decremented on completion) + * + * The /metrics endpoint itself is excluded to prevent self-instrumentation. + */ +@Injectable() +export class MetricsInterceptor implements NestInterceptor { + private readonly logger = new Logger(MetricsInterceptor.name); + private activeRequests = 0; + + constructor(private readonly metricsService: MetricsService) {} + + intercept(context: ExecutionContext, next: CallHandler): Observable { + const http = context.switchToHttp(); + const req = http.getRequest(); + const res = http.getResponse(); + + // Skip metrics collection for the /metrics endpoint + if (req.path === '/metrics') { + return next.handle(); + } + + const startTime = process.hrtime.bigint(); + this.activeRequests++; + + return next.handle().pipe( + tap({ + next: () => { + this.recordMetrics(req, res, startTime); + }, + error: () => { + this.recordMetrics(req, res, startTime); + }, + finalize: () => { + this.activeRequests--; + }, + }), + ); + } + + private recordMetrics(req: Request, res: Response, startTime: bigint): void { + try { + const durationSeconds = Number(process.hrtime.bigint() - startTime) / 1e9; + const route = normalizeRoutePath(req.path); + + this.metricsService.observeHttpRequest( + req.method, + route, + res.statusCode, + durationSeconds, + ); + } catch (error) { + // Log metric recording errors but don't fail the request + this.logger.error( + `Failed to record metrics for ${req.method} ${req.path}: ${(error as Error).message}`, + ); + } + } + + /** + * Returns the current number of active requests being processed. + * This can be used for monitoring system load. + */ + getActiveRequestCount(): number { + return this.activeRequests; + } +} diff --git a/src/common/interceptors/request-context.interceptor.spec.ts b/src/common/interceptors/request-context.interceptor.spec.ts index c1f267c4..bebf0997 100644 --- a/src/common/interceptors/request-context.interceptor.spec.ts +++ b/src/common/interceptors/request-context.interceptor.spec.ts @@ -163,4 +163,33 @@ describe('RequestContextInterceptor', () => { expect(capturedAgent).toBe('agent-9'); expect(capturedAuthMethod).toBe('service'); }); + + it('keeps request IDs isolated across concurrent async contexts', async () => { + const makeContext = (requestId: string) => ({ + identity: { + requestId, + correlationId: requestId, + traceId: requestId, + method: 'GET', + path: '/', + url: '/', + ip: null, + userAgent: null, + startedAt: Date.now(), + }, + timings: {}, + data: {}, + }); + + const results = await Promise.all( + ['req-concurrent-a', 'req-concurrent-b'].map((requestId) => + RequestContext.run(makeContext(requestId), async () => { + await new Promise((resolve) => setImmediate(resolve)); + return RequestContext.getRequestId(); + }), + ), + ); + + expect(results).toEqual(['req-concurrent-a', 'req-concurrent-b']); + }); }); diff --git a/src/common/interceptors/request-context.interceptor.ts b/src/common/interceptors/request-context.interceptor.ts index 2d4cee30..71254f27 100644 --- a/src/common/interceptors/request-context.interceptor.ts +++ b/src/common/interceptors/request-context.interceptor.ts @@ -6,7 +6,6 @@ import { } from '@nestjs/common'; import { Observable } from 'rxjs'; import { Request } from 'express'; -import { v7 as uuidv7 } from 'uuid'; import { RequestContext, RequestContextData, @@ -17,6 +16,7 @@ import { CORRELATION_ID_HEADER, REQUEST_ID_HEADER, } from '../constants/headers'; +import { resolveRequestId } from '../helpers/request-id'; /** * Seeds the structured request context (see {@link RequestContext}) at the very @@ -49,10 +49,9 @@ export class RequestContextInterceptor implements NestInterceptor { } private seed(req: Request & { user?: AuthenticatedUser }): RequestContextData { - const requestId = - RequestContext.getRequestId() ?? - (req.headers[REQUEST_ID_HEADER] as string | undefined) ?? - `req_${uuidv7()}`; + const requestId = resolveRequestId( + RequestContext.getRequestId() ?? req.headers[REQUEST_ID_HEADER], + ); const traceId = (req.headers[CORRELATION_ID_HEADER] as string | undefined) ?? diff --git a/src/common/interceptors/request-id.interceptor.spec.ts b/src/common/interceptors/request-id.interceptor.spec.ts new file mode 100644 index 00000000..5f187f5b --- /dev/null +++ b/src/common/interceptors/request-id.interceptor.spec.ts @@ -0,0 +1,293 @@ +import { describe, it, expect, vi, beforeEach } from 'vitest'; +import { ExecutionContext, CallHandler, Logger } from '@nestjs/common'; +import { of, throwError } from 'rxjs'; +import { RequestIdInterceptor } from './request-id.interceptor'; +import { REQUEST_ID_HEADER } from '../constants/headers'; + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +/** + * Builds a minimal mock ExecutionContext for HTTP requests. Callers supply only + * the fields relevant to their test case. + */ +function buildContext(options: { + incomingRequestId?: string; + method?: string; + path?: string; +}): { + context: ExecutionContext; + requestHeaders: Record; + responseHeaders: Record; + requestRef: { id?: string; headers: Record; method: string; path: string }; +} { + const requestHeaders: Record = {}; + if (options.incomingRequestId !== undefined) { + requestHeaders[REQUEST_ID_HEADER] = options.incomingRequestId; + } + + const responseHeaders: Record = {}; + + const requestRef = { + id: undefined as string | undefined, + headers: requestHeaders, + method: options.method ?? 'GET', + path: options.path ?? '/api/v1/test', + }; + + const context = { + switchToHttp: () => ({ + getRequest: () => requestRef, + getResponse: () => ({ + setHeader: (name: string, value: string) => { + responseHeaders[name] = value; + }, + statusCode: 200, + }), + }), + } as unknown as ExecutionContext; + + return { context, requestHeaders, responseHeaders, requestRef }; +} + +/** + * Executes the interceptor and resolves once the observable completes or errors. + */ +function run( + interceptor: RequestIdInterceptor, + context: ExecutionContext, + callHandler: CallHandler, +): Promise { + return new Promise((resolve, reject) => { + interceptor.intercept(context, callHandler).subscribe({ + next: (val) => resolve(val), + error: (err) => reject(err), + }); + }); +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +describe('RequestIdInterceptor', () => { + let interceptor: RequestIdInterceptor; + + beforeEach(() => { + interceptor = new RequestIdInterceptor(); + // Silence logger output during tests — we assert on behaviour, not log lines. + vi.spyOn(Logger.prototype, 'log').mockImplementation(() => undefined); + vi.spyOn(Logger.prototype, 'warn').mockImplementation(() => undefined); + }); + + // ── Header preservation ───────────────────────────────────────────────── + + it('should preserve an incoming X-Request-ID header', async () => { + const { context, requestHeaders, responseHeaders, requestRef } = buildContext({ + incomingRequestId: 'client-provided-id-123', + }); + + const callHandler: CallHandler = { handle: () => of(null) }; + await run(interceptor, context, callHandler); + + // Header kept on the request + expect(requestHeaders[REQUEST_ID_HEADER]).toBe('client-provided-id-123'); + // Echoed on the response + expect(responseHeaders[REQUEST_ID_HEADER]).toBe('client-provided-id-123'); + // Attached to request.id + expect(requestRef.id).toBe('client-provided-id-123'); + }); + + it('should preserve a UUID-format X-Request-ID header unchanged', async () => { + const uuid = '550e8400-e29b-41d4-a716-446655440000'; + const { context, requestHeaders, responseHeaders } = buildContext({ + incomingRequestId: uuid, + }); + + const callHandler: CallHandler = { handle: () => of(null) }; + await run(interceptor, context, callHandler); + + expect(requestHeaders[REQUEST_ID_HEADER]).toBe(uuid); + expect(responseHeaders[REQUEST_ID_HEADER]).toBe(uuid); + }); + + // ── Automatic ID generation ────────────────────────────────────────────── + + it('should generate a UUID when no X-Request-ID header is present', async () => { + const { context, requestHeaders, responseHeaders, requestRef } = buildContext({}); + + const callHandler: CallHandler = { handle: () => of(null) }; + await run(interceptor, context, callHandler); + + const generated = requestHeaders[REQUEST_ID_HEADER]; + expect(generated).toBeDefined(); + expect(typeof generated).toBe('string'); + // crypto.randomUUID() produces the standard 8-4-4-4-12 format + expect(generated).toMatch( + /^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$/i, + ); + expect(responseHeaders[REQUEST_ID_HEADER]).toBe(generated); + expect(requestRef.id).toBe(generated); + }); + + it('should generate a UUID when the X-Request-ID header is an empty string', async () => { + const { context, requestHeaders } = buildContext({ incomingRequestId: '' }); + + const callHandler: CallHandler = { handle: () => of(null) }; + await run(interceptor, context, callHandler); + + const generated = requestHeaders[REQUEST_ID_HEADER]; + expect(generated).toBeDefined(); + expect(generated).not.toBe(''); + expect(generated).toMatch( + /^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$/i, + ); + }); + + it('should generate a UUID when the X-Request-ID header is whitespace only', async () => { + const { context, requestHeaders } = buildContext({ incomingRequestId: ' ' }); + + const callHandler: CallHandler = { handle: () => of(null) }; + await run(interceptor, context, callHandler); + + const generated = requestHeaders[REQUEST_ID_HEADER]; + expect(generated).toBeDefined(); + expect(generated?.trim()).not.toBe(''); + expect(generated).toMatch( + /^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$/i, + ); + }); + + it('replaces request IDs containing unsupported characters or exceeding 128 characters', async () => { + for (const incomingRequestId of ['bad id', 'bad\nid', 'x'.repeat(129)]) { + const { context, requestHeaders, responseHeaders } = buildContext({ incomingRequestId }); + await run(interceptor, context, { handle: () => of(null) }); + + expect(requestHeaders[REQUEST_ID_HEADER]).not.toBe(incomingRequestId); + expect(requestHeaders[REQUEST_ID_HEADER]).toMatch( + /^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$/i, + ); + expect(responseHeaders[REQUEST_ID_HEADER]).toBe(requestHeaders[REQUEST_ID_HEADER]); + } + }); + + it('should generate unique IDs for each request', async () => { + const { context: ctx1 } = buildContext({}); + const { context: ctx2 } = buildContext({}); + + let id1: string | undefined; + let id2: string | undefined; + + const handler1: CallHandler = { + handle: () => { + id1 = (ctx1.switchToHttp().getRequest() as { headers: Record }).headers[REQUEST_ID_HEADER]; + return of(null); + }, + }; + const handler2: CallHandler = { + handle: () => { + id2 = (ctx2.switchToHttp().getRequest() as { headers: Record }).headers[REQUEST_ID_HEADER]; + return of(null); + }, + }; + + await run(interceptor, ctx1, handler1); + await run(interceptor, ctx2, handler2); + + expect(id1).toBeDefined(); + expect(id2).toBeDefined(); + expect(id1).not.toBe(id2); + }); + + // ── request.id attachment ──────────────────────────────────────────────── + + it('should attach the request id to request.id for Express compatibility', async () => { + const { context, requestRef } = buildContext({ incomingRequestId: 'express-compat-id' }); + const callHandler: CallHandler = { handle: () => of(null) }; + await run(interceptor, context, callHandler); + + expect(requestRef.id).toBe('express-compat-id'); + }); + + // ── Response header ────────────────────────────────────────────────────── + + it('should set X-Request-ID on the response even when the handler throws', async () => { + const { context, responseHeaders } = buildContext({ incomingRequestId: 'error-case-id' }); + + const callHandler: CallHandler = { + handle: () => throwError(() => new Error('handler error')), + }; + + await run(interceptor, context, callHandler).catch(() => { + // Expected — we just want to inspect the response headers. + }); + + // Response header must be set before handle() is called (synchronous). + expect(responseHeaders[REQUEST_ID_HEADER]).toBe('error-case-id'); + }); + + // ── Structured logging ─────────────────────────────────────────────────── + + it('should emit a structured log on request entry', async () => { + const { context } = buildContext({ + incomingRequestId: 'log-test-id', + method: 'POST', + path: '/api/v1/agents', + }); + + const callHandler: CallHandler = { handle: () => of(null) }; + await run(interceptor, context, callHandler); + + expect(Logger.prototype.log).toHaveBeenCalledWith( + expect.objectContaining({ + message: 'Request received', + requestId: 'log-test-id', + method: 'POST', + path: '/api/v1/agents', + }), + ); + }); + + it('should emit a structured log on successful response completion', async () => { + const { context } = buildContext({ + incomingRequestId: 'log-complete-id', + method: 'GET', + path: '/api/v1/wallets', + }); + + const callHandler: CallHandler = { handle: () => of(null) }; + await run(interceptor, context, callHandler); + + expect(Logger.prototype.log).toHaveBeenCalledWith( + expect.objectContaining({ + message: 'Request completed', + requestId: 'log-complete-id', + method: 'GET', + path: '/api/v1/wallets', + }), + ); + }); + + it('should emit a warn log when the handler errors', async () => { + const { context } = buildContext({ + incomingRequestId: 'log-error-id', + method: 'DELETE', + path: '/api/v1/agents/1', + }); + + const callHandler: CallHandler = { + handle: () => throwError(() => new Error('something went wrong')), + }; + + await run(interceptor, context, callHandler).catch(() => undefined); + + expect(Logger.prototype.warn).toHaveBeenCalledWith( + expect.objectContaining({ + message: 'Request errored', + requestId: 'log-error-id', + error: 'something went wrong', + }), + ); + }); +}); diff --git a/src/common/interceptors/request-id.interceptor.ts b/src/common/interceptors/request-id.interceptor.ts new file mode 100644 index 00000000..b2c7b475 --- /dev/null +++ b/src/common/interceptors/request-id.interceptor.ts @@ -0,0 +1,95 @@ +import { + CallHandler, + ExecutionContext, + Injectable, + Logger, + NestInterceptor, +} from '@nestjs/common'; +import { Request, Response } from 'express'; +import { Observable } from 'rxjs'; +import { tap } from 'rxjs/operators'; +import { REQUEST_ID_HEADER } from '../constants/headers'; +import { resolveRequestId } from '../helpers/request-id'; + +/** + * Global interceptor that ensures every HTTP request carries a stable, + * cryptographically-secure request identifier throughout its full lifecycle. + * + * Execution order (runs first among APP_INTERCEPTORs): + * 1. Reads the existing `X-Request-ID` header forwarded by the client or an + * upstream proxy (e.g. a load balancer, API gateway). + * 2. Falls back to `crypto.randomUUID()` when the header is absent or empty. + * 3. Normalises the resolved ID by writing it back onto `request.headers` so + * that downstream interceptors (RequestContextInterceptor, + * AgentTraceInterceptor, ResponseInterceptor) and `pino-http`'s `genReqId` + * all see a consistent value. + * 4. Attaches the ID to `request.id` for compatibility with frameworks and + * middleware that read the Express `id` property. + * 5. Sets the `X-Request-ID` response header so clients and debugging tools + * can correlate a response with the originating request. + * 6. Emits a structured log entry on request start and on response completion, + * carrying `{ requestId, method, path }` for end-to-end distributed + * tracing across controllers, services and background jobs. + * + * This interceptor intentionally performs no async work and injects no services + * so it can be instantiated as a plain class without a DI container (important + * for unit tests and for being wired as the very first APP_INTERCEPTOR). + */ +@Injectable() +export class RequestIdInterceptor implements NestInterceptor { + private readonly logger = new Logger(RequestIdInterceptor.name); + + intercept(context: ExecutionContext, next: CallHandler): Observable { + const http = context.switchToHttp(); + const request = http.getRequest(); + const response = http.getResponse(); + + // 1. Preserve a valid incoming header; replace missing or invalid values. + const incoming = request.headers[REQUEST_ID_HEADER] as string | undefined; + const requestId = resolveRequestId(incoming); + + // 2. Normalise — stamp the resolved ID back onto the request headers so + // every downstream consumer reads the same value regardless of whether + // the client supplied one. + request.headers[REQUEST_ID_HEADER] = requestId; + + // 3. Attach to `request.id` for Express-ecosystem compatibility. + request.id = requestId; + + // 4. Echo onto the response immediately (before the handler runs) so the + // header is present even when the handler throws synchronously. + response.setHeader(REQUEST_ID_HEADER, requestId); + + // 5. Structured log on request entry. + this.logger.log({ + message: 'Request received', + requestId, + method: request.method, + path: request.path, + }); + + return next.handle().pipe( + // 6. Structured log on response completion (success and error alike). + tap({ + next: () => { + this.logger.log({ + message: 'Request completed', + requestId, + method: request.method, + path: request.path, + statusCode: response.statusCode, + }); + }, + error: (err: unknown) => { + this.logger.warn({ + message: 'Request errored', + requestId, + method: request.method, + path: request.path, + error: err instanceof Error ? err.message : String(err), + }); + }, + }), + ); + } +} diff --git a/src/common/interceptors/response.interceptor.spec.ts b/src/common/interceptors/response.interceptor.spec.ts index de75091b..a5e3e39a 100644 --- a/src/common/interceptors/response.interceptor.spec.ts +++ b/src/common/interceptors/response.interceptor.spec.ts @@ -18,6 +18,9 @@ describe('ResponseInterceptor', () => { getRequest: () => ({ headers: requestId ? { [REQUEST_ID_HEADER]: requestId } : {}, }), + getResponse: () => ({ + setHeader: () => undefined, + }), }), } as unknown as ExecutionContext; }; @@ -60,7 +63,7 @@ describe('ResponseInterceptor', () => { it('extracts items and meta from Paginated responses', async () => { const paginated = new Paginated( [{ id: '1' }, { id: '2' }], - { total: 2, page: 1, limit: 10, totalPages: 1, hasNext: false, hasPrev: false }, + { total: 2, page: 1, limit: 10, offset: 0, totalPages: 1, hasNext: false, hasPrev: false }, ); const context = createMockContext('test-request-id'); @@ -72,7 +75,7 @@ describe('ResponseInterceptor', () => { expect(result).toEqual({ success: true, data: [{ id: '1' }, { id: '2' }], - meta: { total: 2, page: 1, limit: 10, totalPages: 1, hasNext: false, hasPrev: false }, + meta: { total: 2, page: 1, limit: 10, offset: 0, totalPages: 1, hasNext: false, hasPrev: false }, requestId: 'test-request-id', }); }); diff --git a/src/common/interceptors/response.interceptor.ts b/src/common/interceptors/response.interceptor.ts index 0efb7e92..b3a5a7bb 100644 --- a/src/common/interceptors/response.interceptor.ts +++ b/src/common/interceptors/response.interceptor.ts @@ -1,17 +1,19 @@ import { CallHandler, ExecutionContext, Injectable, NestInterceptor } from '@nestjs/common'; -import { Request } from 'express'; +import { Request, Response } from 'express'; import { Observable } from 'rxjs'; import { map } from 'rxjs/operators'; import { ApiSuccessResponse, + CursorPaginated, Paginated, } from '../interfaces/api-response.interface'; -import { REQUEST_ID_HEADER } from '../constants/headers'; +import { REQUEST_ID_HEADER, TOTAL_COUNT_HEADER } from '../constants/headers'; /** * Wraps every successful controller return value in the canonical success - * envelope. If a handler returns a `Paginated`, its items become `data` and - * its pagination info becomes `meta`. + * envelope. If a handler returns a `Paginated`, its items become `data`, its + * pagination info becomes `meta`, and the total row count is also exposed via + * the `X-Total-Count` header for clients that page from headers. */ @Injectable() export class ResponseInterceptor implements NestInterceptor> { @@ -19,12 +21,17 @@ export class ResponseInterceptor implements NestInterceptor, ): Observable> { - const request = context.switchToHttp().getRequest(); + const http = context.switchToHttp(); + const request = http.getRequest(); const requestId = (request.headers[REQUEST_ID_HEADER] as string) ?? 'unknown'; return next.handle().pipe( map((payload): ApiSuccessResponse => { if (payload instanceof Paginated) { + http.getResponse().setHeader(TOTAL_COUNT_HEADER, String(payload.meta.total)); + return { success: true, data: payload.items, meta: payload.meta, requestId }; + } + if (payload instanceof CursorPaginated) { return { success: true, data: payload.items, meta: payload.meta, requestId }; } return { success: true, data: payload ?? null, meta: {}, requestId }; diff --git a/src/common/interfaces/api-response.interface.ts b/src/common/interfaces/api-response.interface.ts index 9c5b42f0..57400c0e 100644 --- a/src/common/interfaces/api-response.interface.ts +++ b/src/common/interfaces/api-response.interface.ts @@ -10,6 +10,8 @@ export interface ApiMeta { } export interface PaginationMeta extends ApiMeta { + /** Zero-based row offset of the first item in `data`. */ + offset: number; page: number; limit: number; total: number; @@ -18,6 +20,12 @@ export interface PaginationMeta extends ApiMeta { hasPrev: boolean; } +export interface CursorPaginationMeta extends ApiMeta { + limit: number; + hasNext: boolean; + nextCursor: string | null; +} + export interface ApiSuccessResponse { success: true; data: T; @@ -61,3 +69,10 @@ export class Paginated { public readonly meta: PaginationMeta, ) {} } + +export class CursorPaginated { + constructor( + public readonly items: T[], + public readonly meta: CursorPaginationMeta, + ) {} +} diff --git a/src/common/interfaces/authenticated-user.interface.ts b/src/common/interfaces/authenticated-user.interface.ts index 12e65167..23d1f02b 100644 --- a/src/common/interfaces/authenticated-user.interface.ts +++ b/src/common/interfaces/authenticated-user.interface.ts @@ -3,11 +3,14 @@ import { UserRole } from '@prisma/client'; /** The authenticated principal attached to each request by JWT or API key strategies. */ export interface AuthenticatedUser { id: string; + sub?: string; organizationId: string; email?: string; role: UserRole; + tier?: string; sessionId?: string; apiKeyId?: string; + createdById?: string | null; scopes?: string[]; permissions?: string[]; isApiKey?: boolean; @@ -17,6 +20,7 @@ export interface AuthenticatedUser { export interface AuthenticatedApiKey { id: string; keyId: string; + apiKeyId?: string; organizationId: string; createdById?: string | null; name: string; diff --git a/src/common/locks/redis-lock.util.spec.ts b/src/common/locks/redis-lock.util.spec.ts index f13dbbcd..66ad0c28 100644 --- a/src/common/locks/redis-lock.util.spec.ts +++ b/src/common/locks/redis-lock.util.spec.ts @@ -140,6 +140,43 @@ describe('RedisLock', () => { await expect(lock.withLock('agent-1', fn, 5000, 2, 0)).rejects.toThrow('boom'); expect(redis.eval).toHaveBeenCalledTimes(1); }); + + it('allows only one concurrent handler for the same resource key', async () => { + let held = false; + let finishHandler!: () => void; + let signalEntered!: () => void; + const handlerGate = new Promise((resolve) => { + finishHandler = resolve; + }); + const handlerEntered = new Promise((resolve) => { + signalEntered = resolve; + }); + redis.set.mockImplementation(async () => { + if (held) return null; + held = true; + return 'OK'; + }); + redis.eval.mockImplementation(async () => { + held = false; + return 1; + }); + const handler = vi.fn(async () => { + signalEntered(); + await handlerGate; + }); + + const first = lock.withLock('wallet:1', handler); + await handlerEntered; + + await expect(lock.withLock('wallet:1', handler)).rejects.toBeInstanceOf( + LockNotAcquiredException, + ); + expect(handler).toHaveBeenCalledTimes(1); + + finishHandler(); + await first; + expect(held).toBe(false); + }); }); it('disconnects the shared client on module destroy', () => { diff --git a/src/config/database.config.ts b/src/config/database.config.ts index f441d0d8..434a0cec 100644 --- a/src/config/database.config.ts +++ b/src/config/database.config.ts @@ -28,6 +28,8 @@ export type DatabaseConfig = { slowQueryThresholdMs: number; connectionRetryAttempts: number; connectionRetryDelayMs: number; + migrationCheckEnabled: boolean; + migrationCheckMode: 'halt' | 'warn'; }; export const databaseConfig = registerAs('database', (): DatabaseConfig => { @@ -43,5 +45,7 @@ export const databaseConfig = registerAs('database', (): DatabaseConfig => { slowQueryThresholdMs: env.DATABASE_SLOW_QUERY_THRESHOLD_MS, connectionRetryAttempts: env.DATABASE_CONNECT_RETRY_ATTEMPTS, connectionRetryDelayMs: env.DATABASE_CONNECT_RETRY_DELAY_MS, + migrationCheckEnabled: env.DATABASE_MIGRATION_CHECK_ENABLED, + migrationCheckMode: env.DATABASE_MIGRATION_CHECK_MODE, }; }); diff --git a/src/config/env.validation.ts b/src/config/env.validation.ts index 9f38b578..aa1da734 100644 --- a/src/config/env.validation.ts +++ b/src/config/env.validation.ts @@ -39,6 +39,8 @@ export const databaseEnvSchema = z.object({ DATABASE_SLOW_QUERY_THRESHOLD_MS: z.coerce.number().int().nonnegative().default(1000), DATABASE_CONNECT_RETRY_ATTEMPTS: z.coerce.number().int().positive().max(10).default(5), DATABASE_CONNECT_RETRY_DELAY_MS: z.coerce.number().int().positive().max(60000).default(1000), + DATABASE_MIGRATION_CHECK_ENABLED: z.coerce.boolean().default(true), + DATABASE_MIGRATION_CHECK_MODE: z.enum(['halt', 'warn']).default('halt'), }); export const redisEnvSchema = z.object({ @@ -85,7 +87,16 @@ export const queueEnvSchema = z.object({ export const throttleEnvSchema = z.object({ THROTTLE_AUTH_LIMIT: z.coerce.number().int().positive().default(10), THROTTLE_API_LIMIT: z.coerce.number().int().positive().default(120), + /** Requests allowed per window for traffic identified as an autonomous agent. */ + THROTTLE_AGENT_LIMIT: z.coerce.number().int().positive().default(300), + THROTTLE_WEBHOOK_LIMIT: z.coerce.number().int().positive().default(30), THROTTLE_TTL: z.coerce.number().int().positive().default(60), + // Short-term burst allowance per tier (requests per second). A burst window + // is intentionally kept very short (1 s) so spikes don't exhaust the full + // steady-state quota. Set to 0 to disable burst enforcement. + THROTTLE_API_BURST: z.coerce.number().int().nonnegative().default(10), + THROTTLE_AUTH_BURST: z.coerce.number().int().nonnegative().default(3), + THROTTLE_WEBHOOK_BURST: z.coerce.number().int().nonnegative().default(5), }); export const rateLimitEnvSchema = z.object({ @@ -165,18 +176,21 @@ export const encryptionEnvSchema = z.object({ * Production additionally rejects insecure-but-valid values that are fine for * local development. */ -export const environmentSchema = appEnvSchema - .merge(databaseEnvSchema) - .merge(redisEnvSchema) - .merge(authEnvSchema) - .merge(stellarEnvSchema) - .merge(storageEnvSchema) - .merge(queueEnvSchema) - .merge(throttleEnvSchema) - .merge(rateLimitEnvSchema) - .merge(metricsEnvSchema) - .merge(aiEnvSchema) - .merge(encryptionEnvSchema) +export const environmentSchema = z + .object({ + ...appEnvSchema.shape, + ...databaseEnvSchema.shape, + ...redisEnvSchema.shape, + ...authEnvSchema.shape, + ...stellarEnvSchema.shape, + ...storageEnvSchema.shape, + ...queueEnvSchema.shape, + ...throttleEnvSchema.shape, + ...rateLimitEnvSchema.shape, + ...metricsEnvSchema.shape, + ...aiEnvSchema.shape, + ...encryptionEnvSchema.shape, + }) .superRefine((env, ctx) => { if (env.NODE_ENV !== 'production') { return; diff --git a/src/config/rate-limit.config.ts b/src/config/rate-limit.config.ts index 75ec9893..d48f2d3a 100644 --- a/src/config/rate-limit.config.ts +++ b/src/config/rate-limit.config.ts @@ -1,12 +1,21 @@ import { registerAs } from '@nestjs/config'; import { rateLimitEnvSchema, validateEnv } from './env.validation'; +/** Client identifiers that can participate in the public rate-limit bucket. */ +export type PublicRateLimitIdentifier = 'ip' | 'apiKey'; + /** Settings for the IP-based limiter applied to unauthenticated routes. */ export type PublicRateLimitConfig = { enabled: boolean; maxRequests: number; windowSeconds: number; trustProxy: boolean; + /** + * Identifiers folded into the bucket key. The client IP always + * participates; 'apiKey' additionally separates key-holding clients + * behind a shared address. Order defines bucket-key composition. + */ + clientIdentifiers: PublicRateLimitIdentifier[]; }; export type RateLimitConfig = { @@ -15,6 +24,26 @@ export type RateLimitConfig = { public: PublicRateLimitConfig; }; +/** + * Parses the `PUBLIC_RATE_LIMIT_CLIENT_IDENTIFIERS` list. Unknown entries are + * ignored so a typo cannot break startup; the IP always participates anyway. + * + * Read outside the Zod environment schema (same pattern as + * `BALANCE_CACHE_TTL`) so the schema file keeps a fixed set of keys. + */ +function parseClientIdentifiers(raw: string | undefined): PublicRateLimitIdentifier[] { + if (!raw) { + return []; + } + const known: PublicRateLimitIdentifier[] = ['ip', 'apiKey']; + return raw + .split(',') + .map((entry) => entry.trim()) + .filter((entry): entry is PublicRateLimitIdentifier => + known.includes(entry as PublicRateLimitIdentifier), + ); +} + /** * Config for the Redis-backed sliding-window rate limiters: the per-route * `SlidingWindowThrottlerGuard` (top-level fields) and the IP-based @@ -30,6 +59,7 @@ export const rateLimitConfig = registerAs('rateLimit', (): RateLimitConfig => { maxRequests: env.PUBLIC_RATE_LIMIT_MAX_REQUESTS, windowSeconds: env.PUBLIC_RATE_LIMIT_WINDOW_SECONDS, trustProxy: env.PUBLIC_RATE_LIMIT_TRUST_PROXY, + clientIdentifiers: parseClientIdentifiers(process.env.PUBLIC_RATE_LIMIT_CLIENT_IDENTIFIERS), }, }; }); diff --git a/src/config/throttler.config.spec.ts b/src/config/throttler.config.spec.ts index 8c0f65fb..e629d42e 100644 --- a/src/config/throttler.config.spec.ts +++ b/src/config/throttler.config.spec.ts @@ -14,11 +14,18 @@ describe('throttlerConfig', () => { delete process.env.THROTTLE_TTL; delete process.env.THROTTLE_API_LIMIT; delete process.env.THROTTLE_AUTH_LIMIT; + delete process.env.THROTTLE_AGENT_LIMIT; + delete process.env.THROTTLE_WEBHOOK_LIMIT; expect(throttlerConfig() as ThrottlerConfig).toEqual({ windowSeconds: 60, apiLimit: 120, authLimit: 10, + agentLimit: 300, + webhookLimit: 30, + apiBurst: 10, + authBurst: 3, + webhookBurst: 5, }); }); @@ -26,11 +33,21 @@ describe('throttlerConfig', () => { process.env.THROTTLE_TTL = '30'; process.env.THROTTLE_API_LIMIT = '500'; process.env.THROTTLE_AUTH_LIMIT = '5'; + process.env.THROTTLE_AGENT_LIMIT = '900'; + process.env.THROTTLE_WEBHOOK_LIMIT = '60'; + process.env.THROTTLE_API_BURST = '20'; + process.env.THROTTLE_AUTH_BURST = '2'; + process.env.THROTTLE_WEBHOOK_BURST = '8'; expect(throttlerConfig() as ThrottlerConfig).toEqual({ windowSeconds: 30, apiLimit: 500, authLimit: 5, + agentLimit: 900, + webhookLimit: 60, + apiBurst: 20, + authBurst: 2, + webhookBurst: 8, }); }); @@ -39,30 +56,90 @@ describe('throttlerConfig', () => { expect(() => throttlerConfig()).toThrow(/THROTTLE_TTL/); }); + + it('accepts zero burst values to disable burst enforcement', () => { + process.env.THROTTLE_API_BURST = '0'; + process.env.THROTTLE_AUTH_BURST = '0'; + process.env.THROTTLE_WEBHOOK_BURST = '0'; + + const config = throttlerConfig() as ThrottlerConfig; + + expect(config.apiBurst).toBe(0); + expect(config.authBurst).toBe(0); + expect(config.webhookBurst).toBe(0); + }); }); describe('createThrottlerOptions', () => { - const config: ThrottlerConfig = { windowSeconds: 60, apiLimit: 120, authLimit: 10 }; - - it('exposes exactly two named tiers so AstroidThrottlerGuard can route by tier', () => { + const config: ThrottlerConfig = { + windowSeconds: 60, + apiLimit: 120, + authLimit: 10, + agentLimit: 300, + webhookLimit: 30, + apiBurst: 10, + authBurst: 3, + webhookBurst: 5, + }; + + it('exposes four steady-state tiers so guards can route by tier', () => { const options = createThrottlerOptions(config); expect(Array.isArray(options)).toBe(false); - expect(options.throttlers.map((throttler) => throttler.name)).toEqual(['api', 'auth']); + expect(options.throttlers.filter((t) => !t.name?.endsWith('-burst')).map((t) => t.name)).toEqual([ + 'api', + 'auth', + 'agent', + 'webhook', + ]); }); it('converts the configured window from seconds to the milliseconds @nestjs/throttler expects', () => { const options = createThrottlerOptions({ ...config, windowSeconds: 30 }); + const steadyState = options.throttlers.filter((t) => !t.name?.endsWith('-burst')); - expect(options.throttlers[0].ttl).toBe(30_000); - expect(options.throttlers[1].ttl).toBe(30_000); + expect(steadyState[0].ttl).toBe(30_000); + expect(steadyState[1].ttl).toBe(30_000); + expect(steadyState[2].ttl).toBe(30_000); + expect(steadyState[3].ttl).toBe(30_000); }); - it('applies the stricter limit to the auth tier only', () => { + it('applies tier-specific limits to api, auth, agent and webhook', () => { const options = createThrottlerOptions(config); expect(options.throttlers.find((t) => t.name === 'api')?.limit).toBe(120); expect(options.throttlers.find((t) => t.name === 'auth')?.limit).toBe(10); + expect(options.throttlers.find((t) => t.name === 'agent')?.limit).toBe(300); + expect(options.throttlers.find((t) => t.name === 'webhook')?.limit).toBe(30); + }); + + it('registers burst throttlers with a 1-second TTL for non-zero burst values', () => { + const options = createThrottlerOptions(config); + + const apiBurst = options.throttlers.find((t) => t.name === 'api-burst'); + const authBurst = options.throttlers.find((t) => t.name === 'auth-burst'); + const webhookBurst = options.throttlers.find((t) => t.name === 'webhook-burst'); + + expect(apiBurst).toBeDefined(); + expect(apiBurst?.ttl).toBe(1_000); + expect(apiBurst?.limit).toBe(10); + + expect(authBurst).toBeDefined(); + expect(authBurst?.ttl).toBe(1_000); + expect(authBurst?.limit).toBe(3); + + expect(webhookBurst).toBeDefined(); + expect(webhookBurst?.ttl).toBe(1_000); + expect(webhookBurst?.limit).toBe(5); + }); + + it('omits burst throttlers when burst limits are zero', () => { + const noBurstConfig: ThrottlerConfig = { ...config, apiBurst: 0, authBurst: 0, webhookBurst: 0 }; + const options = createThrottlerOptions(noBurstConfig); + + expect(options.throttlers.find((t) => t.name === 'api-burst')).toBeUndefined(); + expect(options.throttlers.find((t) => t.name === 'auth-burst')).toBeUndefined(); + expect(options.throttlers.find((t) => t.name === 'webhook-burst')).toBeUndefined(); }); it('attaches the shared Redis storage, without which counters stay in-process', () => { diff --git a/src/config/throttler.config.ts b/src/config/throttler.config.ts index 53a4740e..1568d96e 100644 --- a/src/config/throttler.config.ts +++ b/src/config/throttler.config.ts @@ -9,20 +9,31 @@ import { throttleEnvSchema, validateEnv } from './env.validation'; export type TieredThrottlerOptions = Exclude; export type ThrottlerConfig = { - /** Fixed-window length in seconds, shared by every tier. */ + /** Fixed-window length in seconds, shared by every steady-state tier. */ windowSeconds: number; /** Requests allowed per window on the public `api` tier. */ apiLimit: number; /** Requests allowed per window on the sensitive `auth` tier. */ authLimit: number; + /** Requests allowed per window for a single autonomous agent. */ + agentLimit: number; + /** Requests allowed per window on the `webhook` management tier. */ + webhookLimit: number; + /** + * Burst throttlers — each applies a 1-second window with a per-tier + * maximum so single-second spikes don't consume the full steady-state quota. + * A value of 0 disables burst enforcement for that tier. + */ + apiBurst: number; + authBurst: number; + webhookBurst: number; }; /** * Rate-limit configuration, driven by the `THROTTLE_*` environment variables. * - * Historically these values lived under the `queue` namespace even though - * BullMQ never read them — they only ever configured `@nestjs/throttler`. The - * dedicated `throttler` namespace makes the ownership explicit. + * The dedicated `throttler` namespace makes the ownership of these variables + * explicit (they previously lived ambiguously under `queue`). */ export const throttlerConfig = registerAs('throttler', (): ThrottlerConfig => { const env = validateEnv(throttleEnvSchema, process.env); @@ -30,13 +41,27 @@ export const throttlerConfig = registerAs('throttler', (): ThrottlerConfig => { windowSeconds: env.THROTTLE_TTL, apiLimit: env.THROTTLE_API_LIMIT, authLimit: env.THROTTLE_AUTH_LIMIT, + agentLimit: env.THROTTLE_AGENT_LIMIT, + webhookLimit: env.THROTTLE_WEBHOOK_LIMIT, + apiBurst: env.THROTTLE_API_BURST, + authBurst: env.THROTTLE_AUTH_BURST, + webhookBurst: env.THROTTLE_WEBHOOK_BURST, }; }); /** - * Builds the two tiered throttlers consumed by `AstroidThrottlerGuard`: - * - `api` — every route that does not declare a tier explicitly - * - `auth` — routes marked with `@ThrottleTierDecorator('auth')` + * Builds the named throttlers consumed by `AstroidThrottlerGuard`: + * + * Steady-state tiers (TTL = `windowSeconds`): + * - `api` — every route that does not declare a tier explicitly + * - `auth` — routes marked with `@ThrottleTierDecorator('auth')` + * - `agent` — high-frequency routes enforced by `AgentThrottlerGuard` + * - `webhook` — routes marked with `@ThrottleTierDecorator('webhook')` + * + * Burst tiers (TTL = 1 second), only registered when the burst limit > 0: + * - `api-burst` — short-term spike guard for `api` routes + * - `auth-burst` — short-term spike guard for `auth` routes + * - `webhook-burst` — short-term spike guard for `webhook` routes * * The options must be returned in the object form (not the bare array) so the * shared Redis {@link ThrottlerStorage} can be attached: `@nestjs/throttler` @@ -50,12 +75,29 @@ export function createThrottlerOptions( storage?: ThrottlerStorage, ): TieredThrottlerOptions { const ttl = config.windowSeconds * 1000; + const burstTtl = 1_000; // 1 second burst window + + const throttlers: ThrottlerOptions[] = [ + // ── Steady-state tiers ────────────────────────────────────────────────── + { name: 'api', ttl, limit: config.apiLimit }, + { name: 'auth', ttl, limit: config.authLimit }, + { name: 'agent', ttl, limit: config.agentLimit }, + { name: 'webhook', ttl, limit: config.webhookLimit }, + ]; + + // ── Burst tiers — only wired when burst > 0 ───────────────────────────── + if (config.apiBurst > 0) { + throttlers.push({ name: 'api-burst', ttl: burstTtl, limit: config.apiBurst }); + } + if (config.authBurst > 0) { + throttlers.push({ name: 'auth-burst', ttl: burstTtl, limit: config.authBurst }); + } + if (config.webhookBurst > 0) { + throttlers.push({ name: 'webhook-burst', ttl: burstTtl, limit: config.webhookBurst }); + } return { ...(storage ? { storage } : {}), - throttlers: [ - { name: 'api', ttl, limit: config.apiLimit }, - { name: 'auth', ttl, limit: config.authLimit }, - ], + throttlers, }; } diff --git a/src/database/prisma.service.spec.ts b/src/database/prisma.service.spec.ts index 0cffade0..7a29a0ba 100644 --- a/src/database/prisma.service.spec.ts +++ b/src/database/prisma.service.spec.ts @@ -300,3 +300,41 @@ describe('PrismaService', () => { expect(checkMigrationStatusMock).not.toHaveBeenCalled(); }); }); + +describe('getPoolStats aggregation logic', () => { + // PrismaService.getPoolStats aggregates pg_stat_activity rows fetched via + // $queryRawUnsafe. The mock PrismaClient above replaces `this` on + // construction (a constructor returning an object shadows the derived + // instance per JS semantics), so PrismaService's own prototype methods + // aren't reachable through it — this exercises the same aggregation logic + // directly against a stub client instead, mirroring what getPoolStats does. + async function aggregate( + rows: { state: string | null; wait_event_type: string | null; count: bigint }[], + ): Promise<{ active: number; idle: number; waiting: number }> { + let active = 0; + let idle = 0; + let waiting = 0; + for (const row of rows) { + const count = Number(row.count); + if (row.wait_event_type === 'Lock') { + waiting += count; + } else if (row.state === 'active') { + active += count; + } else if (row.state?.startsWith('idle')) { + idle += count; + } + } + return { active, idle, waiting }; + } + + it('aggregates pg_stat_activity rows into active/idle/waiting counts', async () => { + const stats = await aggregate([ + { state: 'active', wait_event_type: null, count: 2n }, + { state: 'idle', wait_event_type: null, count: 5n }, + { state: 'idle in transaction', wait_event_type: null, count: 1n }, + { state: 'active', wait_event_type: 'Lock', count: 3n }, + ]); + + expect(stats).toEqual({ active: 2, idle: 6, waiting: 3 }); + }); +}); diff --git a/src/database/prisma.service.ts b/src/database/prisma.service.ts index 6235b6db..40d2df43 100644 --- a/src/database/prisma.service.ts +++ b/src/database/prisma.service.ts @@ -169,6 +169,46 @@ export class PrismaService extends PrismaClient implements OnModuleInit, OnModul await this.workerClient.$disconnect(); } + /** + * Reads live connection counts for this database from Postgres' + * `pg_stat_activity`. Prisma's Rust query engine doesn't expose pool + * internals (active/idle/waiting) through the Node client, so this is the + * only accurate source for those numbers — used by MetricsService to + * publish `db_pool_connections`. + */ + async getPoolStats(): Promise<{ active: number; idle: number; waiting: number }> { + try { + const rows = await this.$queryRawUnsafe< + { state: string | null; wait_event_type: string | null; count: bigint }[] + >( + `SELECT state, wait_event_type, count(*) AS count + FROM pg_stat_activity + WHERE datname = current_database() + GROUP BY state, wait_event_type`, + ); + + let active = 0; + let idle = 0; + let waiting = 0; + + for (const row of rows) { + const count = Number(row.count); + if (row.wait_event_type === 'Lock') { + waiting += count; + } else if (row.state === 'active') { + active += count; + } else if (row.state?.startsWith('idle')) { + idle += count; + } + } + + return { active, idle, waiting }; + } catch (error) { + this.logger.warn(`Failed to read pool stats from pg_stat_activity: ${(error as Error).message}`); + return { active: 0, idle: 0, waiting: 0 }; + } + } + /** Registers a Nest shutdown hook so the process closes the pool cleanly. */ async enableShutdownHooks(app: INestApplication): Promise { process.on('beforeExit', () => { diff --git a/src/events/domain-event.types.ts b/src/events/domain-event.types.ts index 6ad1c162..f23d9c8f 100644 --- a/src/events/domain-event.types.ts +++ b/src/events/domain-event.types.ts @@ -6,16 +6,27 @@ import { DomainEventNameType } from './event-names'; * immutable ledger entry and to fan out to webhooks. */ export interface DomainEventEnvelope> { + eventId?: string; name: DomainEventNameType; organizationId?: string; aggregateType: string; aggregateId?: string; actorId?: string; + requestId?: string; correlationId?: string; + metadata?: DomainEventMetadata; payload: TPayload; occurredAt: Date; } +export const DOMAIN_EVENT_ENVELOPE = 'astroid.domain_event'; + +export interface DomainEventMetadata { + requestId: string; + correlationId: string; + traceId?: string; +} + // Base payload types that extend Record for flexibility export interface OrganizationRegisteredPayload extends Record { organizationId: string; diff --git a/src/events/event-bus.service.spec.ts b/src/events/event-bus.service.spec.ts new file mode 100644 index 00000000..0be92c57 --- /dev/null +++ b/src/events/event-bus.service.spec.ts @@ -0,0 +1,69 @@ +import { describe, expect, it, vi } from 'vitest'; +import { RequestContext } from '../common/context/request-context'; +import { EventBusService } from './event-bus.service'; +import { PrismaService } from '../database/prisma.service'; +import { TypedEventEmitter } from './typed-event-emitter.service'; + +function context(requestId: string) { + return { + identity: { + requestId, + correlationId: `corr-${requestId}`, + traceId: `trace-${requestId}`, + method: 'POST', + path: '/wallets', + url: '/wallets', + ip: null, + userAgent: null, + startedAt: Date.now(), + }, + timings: {}, + data: {}, + }; +} + +describe('EventBusService correlation metadata', () => { + it('passes request identity as typed metadata without changing the event payload', async () => { + const emit = vi.fn(); + const emitEnvelope = vi.fn(); + const service = new EventBusService( + { domainEvent: { create: vi.fn() } } as unknown as PrismaService, + { emit, emitEnvelope } as unknown as TypedEventEmitter, + ); + const payload = { walletId: 'wallet-1' }; + + await RequestContext.run(context('req-123'), () => + service.emit('wallet.created', payload, { aggregateType: 'Wallet', persist: false }), + ); + + expect(emit).toHaveBeenCalledWith('wallet.created', payload, { + requestId: 'req-123', + correlationId: 'corr-req-123', + traceId: 'trace-req-123', + }); + expect(emitEnvelope).toHaveBeenCalledWith(expect.objectContaining({ + eventId: expect.any(String), + requestId: 'req-123', + correlationId: 'corr-req-123', + payload, + })); + }); + + it('generates request correlation metadata for internal events', async () => { + const emit = vi.fn(); + const emitEnvelope = vi.fn(); + const service = new EventBusService( + { domainEvent: { create: vi.fn() } } as unknown as PrismaService, + { emit, emitEnvelope } as unknown as TypedEventEmitter, + ); + + await service.emit('wallet.created', { walletId: 'wallet-1' }, { + aggregateType: 'Wallet', + persist: false, + }); + + const metadata = emit.mock.calls[0][2]; + expect(metadata.requestId).toMatch(/^[0-9a-f-]{36}$/i); + expect(metadata.correlationId).toBe(metadata.requestId); + }); +}); \ No newline at end of file diff --git a/src/events/event-bus.service.ts b/src/events/event-bus.service.ts index 0577a88b..c0933b92 100644 --- a/src/events/event-bus.service.ts +++ b/src/events/event-bus.service.ts @@ -1,14 +1,18 @@ import { Injectable, Logger } from '@nestjs/common'; +import { randomUUID } from 'crypto'; import { PrismaService } from '../database/prisma.service'; import { DomainEventNameType } from './event-names'; import { DomainEventEnvelope } from './domain-event.types'; import { TypedEventEmitter, DomainEventMap } from './typed-event-emitter.service'; +import { RequestContext } from '../common/context/request-context'; +import { resolveRequestId } from '../common/helpers/request-id'; export interface EmitOptions { organizationId?: string; aggregateType: string; aggregateId?: string; actorId?: string; + requestId?: string; correlationId?: string; /** When false, the event is broadcast but NOT written to the ledger. */ persist?: boolean; @@ -38,13 +42,23 @@ export class EventBusService { payload: DomainEventMap[K], options: EmitOptions, ): Promise { + const requestId = options.requestId ?? RequestContext.getRequestId() ?? resolveRequestId(undefined); + const correlationId = options.correlationId ?? RequestContext.getCorrelationId() ?? requestId; + const metadata = { + requestId, + correlationId, + traceId: RequestContext.getTraceId() ?? correlationId, + }; const envelope: DomainEventEnvelope> = { + eventId: randomUUID(), name: name as unknown as DomainEventNameType, organizationId: options.organizationId, aggregateType: options.aggregateType, aggregateId: options.aggregateId, actorId: options.actorId, - correlationId: options.correlationId, + requestId, + correlationId, + metadata, payload: payload as unknown as Record, occurredAt: new Date(), }; @@ -55,13 +69,15 @@ export class EventBusService { // Broadcast synchronously in-process using typed emitter for type safety. // Subscribers isolate their own errors. - this.typedEmitter.emit(name, payload); + this.typedEmitter.emit(name, payload, metadata); + this.typedEmitter.emitEnvelope(envelope); } private async persist(envelope: DomainEventEnvelope): Promise { try { await this.prisma.domainEvent.create({ data: { + id: envelope.eventId, organizationId: envelope.organizationId ?? null, name: envelope.name, aggregateType: envelope.aggregateType, diff --git a/src/events/event-names.ts b/src/events/event-names.ts index 6837244d..30bdb80d 100644 --- a/src/events/event-names.ts +++ b/src/events/event-names.ts @@ -59,6 +59,7 @@ export const DomainEventName = { // Risk RiskEvaluated: 'risk.evaluated', RiskAlert: 'risk.alert', + TransactionRiskScoringRequested: 'transaction.risk_scoring_requested', // Notification / audit NotificationCreated: 'notification.created', diff --git a/src/events/typed-event-emitter.service.spec.ts b/src/events/typed-event-emitter.service.spec.ts index 2af003e3..a2475e16 100644 --- a/src/events/typed-event-emitter.service.spec.ts +++ b/src/events/typed-event-emitter.service.spec.ts @@ -40,6 +40,17 @@ describe('TypedEventEmitter', () => { expect(result).toBe(false); }); + it('forwards typed metadata as a separate event argument', () => { + const handler = vi.fn(); + const payload: DomainEventMap['wallet.created'] = { walletId: 'wallet-123' }; + const metadata = { requestId: 'req-1', correlationId: 'corr-1' }; + eventEmitter.on('wallet.created', handler); + + typedEmitter.emit('wallet.created', payload, metadata); + + expect(handler).toHaveBeenCalledWith(payload, metadata); + }); + it('enforces type safety at compile time', () => { const payload: DomainEventMap['agent.registered'] = { agentId: 'agent-123', diff --git a/src/events/typed-event-emitter.service.ts b/src/events/typed-event-emitter.service.ts index 5a9155a0..94d1b11f 100644 --- a/src/events/typed-event-emitter.service.ts +++ b/src/events/typed-event-emitter.service.ts @@ -1,6 +1,8 @@ import { Injectable } from '@nestjs/common'; import { EventEmitter2 } from '@nestjs/event-emitter'; import * as PayloadTypes from './domain-event.types'; +import { DomainEventMetadata } from './domain-event.types'; +import { DOMAIN_EVENT_ENVELOPE, DomainEventEnvelope } from './domain-event.types'; /** * Type-safe mapping of event names to their payload types. @@ -82,8 +84,15 @@ export class TypedEventEmitter { emit( event: K, payload: DomainEventMap[K], + metadata?: DomainEventMetadata, ): boolean { - return this.emitter.emit(event as string, payload); + return metadata + ? this.emitter.emit(event as string, payload, metadata) + : this.emitter.emit(event as string, payload); + } + + emitEnvelope(envelope: DomainEventEnvelope): boolean { + return this.emitter.emit(DOMAIN_EVENT_ENVELOPE, envelope); } /** @@ -93,7 +102,7 @@ export class TypedEventEmitter { */ on( event: K, - handler: (payload: DomainEventMap[K]) => void | Promise, + handler: (payload: DomainEventMap[K], metadata?: DomainEventMetadata) => void | Promise, ): this { this.emitter.on(event as string, handler); return this; @@ -106,7 +115,7 @@ export class TypedEventEmitter { */ once( event: K, - handler: (payload: DomainEventMap[K]) => void | Promise, + handler: (payload: DomainEventMap[K], metadata?: DomainEventMetadata) => void | Promise, ): this { this.emitter.once(event as string, handler); return this; diff --git a/src/main.ts b/src/main.ts index 805ab05e..b27b3036 100644 --- a/src/main.ts +++ b/src/main.ts @@ -8,7 +8,9 @@ import { Request, Response, NextFunction } from 'express'; import { AppModule } from './app.module'; import { PrismaService } from './database/prisma.service'; import { AppConfig } from './config/app.config'; +import { TOTAL_COUNT_HEADER } from './common/constants/headers'; import { assertValidEnvironment, EnvironmentValidationError } from './config/env.validation'; +import { DatabaseConfig } from './config/database.config'; async function bootstrap() { // Fail fast on missing or malformed configuration, before any module is @@ -20,6 +22,23 @@ async function bootstrap() { const app = await NestFactory.create(AppModule, { bufferLogs: true }); const config = app.get(ConfigService); const appConfig = config.getOrThrow('app'); + const databaseConfig = config.getOrThrow('database'); + const prisma = app.get(PrismaService); + + // Startup migration check: refuse to accept traffic against a database + // whose schema hasn't caught up with prisma/migrations (mode 'halt'), or + // log a warning and continue (mode 'warn'). Reuses the same check + // PrismaService.onModuleInit already ran (and logged) on connect. + if (databaseConfig.migrationCheckEnabled) { + const logger = app.get(PinoLogger); + const result = await prisma.validateMigrations(); + + if (!result.upToDate && databaseConfig.migrationCheckMode === 'halt') { + logger.error(result.message, 'MigrationCheck'); + await app.close(); + throw new Error(`Migration check failed: ${result.message}`); + } + } // Structured logging (nestjs-pino) app.useLogger(app.get(PinoLogger)); @@ -56,8 +75,13 @@ async function bootstrap() { next(); }); - // CORS - app.enableCors({ origin: appConfig.corsOrigins, credentials: true }); + // CORS. X-Total-Count is exposed so browser clients can read the total row + // count of paginated list responses. + app.enableCors({ + origin: appConfig.corsOrigins, + credentials: true, + exposedHeaders: [TOTAL_COUNT_HEADER], + }); // Global validation pipe (transforms + validates DTOs) app.useGlobalPipes( @@ -99,7 +123,6 @@ async function bootstrap() { } // Prisma shutdown hook - const prisma = app.get(PrismaService); await prisma.enableShutdownHooks(app); await app.listen(appConfig.port); diff --git a/src/middleware/request-id.middleware.ts b/src/middleware/request-id.middleware.ts index 8b6687cf..53e1f359 100644 --- a/src/middleware/request-id.middleware.ts +++ b/src/middleware/request-id.middleware.ts @@ -1,10 +1,10 @@ import { Injectable, NestMiddleware } from '@nestjs/common'; import { NextFunction, Request, Response } from 'express'; -import { v7 as uuidv7 } from 'uuid'; import { CORRELATION_ID_HEADER, REQUEST_ID_HEADER, } from '../common/constants/headers'; +import { resolveRequestId } from '../common/helpers/request-id'; /** * Ensures every request carries a stable `x-request-id` (generating one when @@ -14,8 +14,8 @@ import { @Injectable() export class RequestIdMiddleware implements NestMiddleware { use(req: Request, res: Response, next: NextFunction): void { - const existing = req.headers[REQUEST_ID_HEADER] as string | undefined; - const requestId = existing && existing.length > 0 ? existing : `req_${uuidv7()}`; + const existing = req.headers[REQUEST_ID_HEADER]; + const requestId = resolveRequestId(existing); req.headers[REQUEST_ID_HEADER] = requestId; const correlation = req.headers[CORRELATION_ID_HEADER] as string | undefined; diff --git a/src/modules/agents/agent.controller.ts b/src/modules/agents/agent.controller.ts index c05248b8..2ea260ea 100644 --- a/src/modules/agents/agent.controller.ts +++ b/src/modules/agents/agent.controller.ts @@ -6,7 +6,6 @@ import { ApiResponse, ApiParam, ApiBody, - ApiQuery, } from '@nestjs/swagger'; import { AgentStatus, UserRole } from '@prisma/client'; import { AgentService } from './agent.service'; @@ -28,6 +27,7 @@ import { UseAgentLock } from '../../common/locks/agent-lock.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { AuthenticatedUser } from '../../common/interfaces/authenticated-user.interface'; import { PaginationQuery, paginationQuerySchema } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; import { ApiEnvelope } from '../../common/decorators/api-envelope.decorator'; import { SlidingWindowThrottlerGuard, @@ -46,8 +46,7 @@ export class AgentController { summary: 'List agents', description: 'Returns a paginated list of agents for the current organization.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiEnvelope(AgentResponseDto as never, { isArray: true }) @ApiResponse({ status: 200, description: 'Paginated list of agents' }) @ApiResponse({ status: 401, description: 'Not authenticated' }) diff --git a/src/modules/agents/agent.service.spec.ts b/src/modules/agents/agent.service.spec.ts index 5cb121c1..e4577bb4 100644 --- a/src/modules/agents/agent.service.spec.ts +++ b/src/modules/agents/agent.service.spec.ts @@ -250,6 +250,7 @@ describe('AgentService', () => { }); const result = await service.list(orgId, { + offset: 0, page: 1, limit: 10, sort: 'createdAt', diff --git a/src/modules/agents/agent.service.ts b/src/modules/agents/agent.service.ts index 701aa320..31815f67 100644 --- a/src/modules/agents/agent.service.ts +++ b/src/modules/agents/agent.service.ts @@ -94,7 +94,7 @@ export class AgentService { const pagination = toPrismaPagination(query, SORTABLE); const { items, total } = await this.repository.findManyAndCount(where, pagination); const decryptedItems = items.map((agent) => this.decryptAgent(agent)); - return new Paginated(decryptedItems, buildPaginationMeta(total, query.page, query.limit)); + return new Paginated(decryptedItems, buildPaginationMeta(total, query)); } async getOrThrow(organizationId: string, id: string): Promise { diff --git a/src/modules/analytics/analytics.repository.ts b/src/modules/analytics/analytics.repository.ts index e8cdac79..c7003377 100644 --- a/src/modules/analytics/analytics.repository.ts +++ b/src/modules/analytics/analytics.repository.ts @@ -7,49 +7,54 @@ import { PrismaService } from '../../database/prisma.service'; export class AnalyticsRepository { constructor(private readonly prisma: PrismaService) {} - countAgents(organizationId: string) { - return this.prisma.agent.count({ where: { organizationId, deletedAt: null } }); - } - - countWallets(organizationId: string) { - return this.prisma.wallet.count({ where: { organizationId, deletedAt: null } }); - } - - countPendingProposals(organizationId: string) { - return this.prisma.proposal.count({ where: { organizationId, status: 'PENDING' } }); - } - - aggregateSpend(organizationId: string, since?: Date) { - const where: Prisma.TransactionWhereInput = { + /** + * Fetches the dashboard overview's counts and spend aggregates in a single + * batched roundtrip (was 5 separate queries) via `$transaction([...])`, then + * fetches the two status/risk-band distributions in parallel. Prisma's + * `groupBy` return type doesn't infer correctly inside a `$transaction` + * array, so those two stay outside the batch as concurrent queries. + */ + async overview(organizationId: string, since30d: Date) { + const completedWhere: Prisma.TransactionWhereInput = { organizationId, status: TransactionStatus.COMPLETED, deletedAt: null, }; - if (since) { - where.createdAt = { gte: since }; - } - return this.prisma.transaction.aggregate({ - where, - _sum: { amount: true }, - _count: { _all: true }, - _avg: { riskScore: true }, - }); - } - groupByStatus(organizationId: string) { - return this.prisma.transaction.groupBy({ - by: ['status'], - where: { organizationId, deletedAt: null }, - _count: { _all: true }, - }); - } + const [[agents, wallets, pendingProposals, allTime, last30d], byStatus, byRisk] = + await Promise.all([ + this.prisma.$transaction([ + this.prisma.agent.count({ where: { organizationId, deletedAt: null } }), + this.prisma.wallet.count({ where: { organizationId, deletedAt: null } }), + this.prisma.proposal.count({ where: { organizationId, status: 'PENDING' } }), + this.prisma.transaction.aggregate({ + where: completedWhere, + _sum: { amount: true }, + _count: { _all: true }, + _avg: { riskScore: true }, + }), + this.prisma.transaction.aggregate({ + where: { ...completedWhere, createdAt: { gte: since30d } }, + _sum: { amount: true }, + _count: { _all: true }, + _avg: { riskScore: true }, + }), + ]), + this.prisma.transaction.groupBy({ + by: ['status'], + where: { organizationId, deletedAt: null }, + orderBy: { status: 'asc' }, + _count: { _all: true }, + }), + this.prisma.transaction.groupBy({ + by: ['riskBand'], + where: { organizationId, deletedAt: null }, + orderBy: { riskBand: 'asc' }, + _count: { _all: true }, + }), + ]); - groupByRiskBand(organizationId: string) { - return this.prisma.transaction.groupBy({ - by: ['riskBand'], - where: { organizationId, deletedAt: null }, - _count: { _all: true }, - }); + return { agents, wallets, pendingProposals, allTime, last30d, byStatus, byRisk }; } spendByAgent(organizationId: string) { diff --git a/src/modules/analytics/analytics.service.spec.ts b/src/modules/analytics/analytics.service.spec.ts new file mode 100644 index 00000000..6c936f57 --- /dev/null +++ b/src/modules/analytics/analytics.service.spec.ts @@ -0,0 +1,80 @@ +import { describe, it, expect, vi, beforeEach } from 'vitest'; +import { AnalyticsService } from './analytics.service'; +import { AnalyticsRepository } from './analytics.repository'; + +describe('AnalyticsService', () => { + let service: AnalyticsService; + let repository: { overview: ReturnType; spendByAgent: ReturnType }; + + beforeEach(() => { + repository = { + overview: vi.fn(), + spendByAgent: vi.fn(), + }; + service = new AnalyticsService(repository as unknown as AnalyticsRepository); + }); + + describe('overview', () => { + it('fetches every card in a single batched repository call', async () => { + repository.overview.mockResolvedValue({ + agents: 3, + wallets: 2, + pendingProposals: 1, + allTime: { _sum: { amount: 100 }, _count: { _all: 10 }, _avg: { riskScore: 42 } }, + last30d: { _sum: { amount: 50 }, _count: { _all: 5 }, _avg: { riskScore: 20 } }, + byStatus: [{ status: 'COMPLETED', _count: { _all: 8 } }], + byRisk: [{ riskBand: 'LOW', _count: { _all: 6 } }], + }); + + const result = await service.overview('org-1'); + + expect(repository.overview).toHaveBeenCalledTimes(1); + expect(repository.overview).toHaveBeenCalledWith('org-1', expect.any(Date)); + expect(result.counts).toEqual({ + agents: 3, + wallets: 2, + pendingProposals: 1, + transactions: 10, + }); + expect(result.spend.allTime).toBe('100'); + expect(result.spend.last30Days).toBe('50'); + expect(result.spend.averageRiskScore).toBe(42); + expect(result.transactionsByStatus).toEqual([{ status: 'COMPLETED', count: 8 }]); + expect(result.transactionsByRiskBand).toEqual([{ riskBand: 'LOW', count: 6 }]); + }); + + it('defaults spend to zero when there is no transaction history', async () => { + repository.overview.mockResolvedValue({ + agents: 0, + wallets: 0, + pendingProposals: 0, + allTime: { _sum: { amount: null }, _count: { _all: 0 }, _avg: { riskScore: null } }, + last30d: { _sum: { amount: null }, _count: { _all: 0 }, _avg: { riskScore: null } }, + byStatus: [], + byRisk: [], + }); + + const result = await service.overview('org-empty'); + + expect(result.spend.allTime).toBe('0'); + expect(result.spend.last30Days).toBe('0'); + expect(result.spend.averageRiskScore).toBe(0); + }); + }); + + describe('spendByAgent', () => { + it('maps repository rows to the response shape, preserving repository order', async () => { + repository.spendByAgent.mockResolvedValue([ + { agentId: 'a2', _sum: { amount: 100 }, _count: { _all: 2 } }, + { agentId: 'a1', _sum: { amount: 10 }, _count: { _all: 1 } }, + ]); + + const result = await service.spendByAgent('org-1'); + + expect(result).toEqual([ + { agentId: 'a2', totalSpent: '100', transactionCount: 2 }, + { agentId: 'a1', totalSpent: '10', transactionCount: 1 }, + ]); + }); + }); +}); diff --git a/src/modules/analytics/analytics.service.ts b/src/modules/analytics/analytics.service.ts index db349d95..a8fee952 100644 --- a/src/modules/analytics/analytics.service.ts +++ b/src/modules/analytics/analytics.service.ts @@ -13,16 +13,8 @@ export class AnalyticsService { /** High-level overview cards for the dashboard home. */ async overview(organizationId: string) { const since30d = new Date(Date.now() - 30 * 86_400_000); - const [agents, wallets, pendingProposals, allTime, last30d, byStatus, byRisk] = - await Promise.all([ - this.repository.countAgents(organizationId), - this.repository.countWallets(organizationId), - this.repository.countPendingProposals(organizationId), - this.repository.aggregateSpend(organizationId), - this.repository.aggregateSpend(organizationId, since30d), - this.repository.groupByStatus(organizationId), - this.repository.groupByRiskBand(organizationId), - ]); + const { agents, wallets, pendingProposals, allTime, last30d, byStatus, byRisk } = + await this.repository.overview(organizationId, since30d); return { counts: { diff --git a/src/modules/approvals/approval.controller.ts b/src/modules/approvals/approval.controller.ts index 06ce1c32..a7031dd3 100644 --- a/src/modules/approvals/approval.controller.ts +++ b/src/modules/approvals/approval.controller.ts @@ -24,6 +24,7 @@ import { AuditAction } from '../../common/decorators/audit-action.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { AuthenticatedUser } from '../../common/interfaces/authenticated-user.interface'; import { PaginationQuery, paginationQuerySchema } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; @ApiTags('approvals') @ApiBearerAuth('access-token') @@ -38,8 +39,7 @@ export class ApprovalController { 'Returns a paginated list of approval proposals for the current organization. ' + 'Supports filtering by status.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiQuery({ name: 'status', required: false, enum: ['PENDING', 'APPROVED', 'REJECTED', 'EXPIRED'], description: 'Filter by proposal status' }) @ApiResponse({ status: 200, description: 'Paginated list of proposals' }) @ApiResponse({ status: 401, description: 'Not authenticated' }) diff --git a/src/modules/approvals/approval.service.ts b/src/modules/approvals/approval.service.ts index efee59c8..f5ce934b 100644 --- a/src/modules/approvals/approval.service.ts +++ b/src/modules/approvals/approval.service.ts @@ -41,7 +41,7 @@ export class ApprovalService { } const pagination = toPrismaPagination(query, SORTABLE); const { items, total } = await this.repository.findManyAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); + return new Paginated(items, buildPaginationMeta(total, query)); } async getOrThrow(organizationId: string, id: string) { diff --git a/src/modules/audit/audit-cursor.spec.ts b/src/modules/audit/audit-cursor.spec.ts new file mode 100644 index 00000000..000a71ca --- /dev/null +++ b/src/modules/audit/audit-cursor.spec.ts @@ -0,0 +1,42 @@ +import { describe, expect, it } from 'vitest'; +import { decodeAuditCursor, encodeAuditCursor, isValidAuditCursor } from './audit-cursor'; +import { auditListQuerySchema } from './audit-list.dto'; + +describe('audit cursor', () => { + it('round-trips the stable timestamp and ID tuple', () => { + const cursor = { createdAt: new Date('2026-09-29T12:30:00.000Z'), id: 'audit-123' }; + expect(decodeAuditCursor(encodeAuditCursor(cursor))).toEqual(cursor); + }); + + it.each(['', 'not-a-cursor', 'e30', Buffer.from('{"v":2}').toString('base64url')])( + 'rejects malformed cursor %s', + (cursor) => { + expect(isValidAuditCursor(cursor)).toBe(false); + expect(auditListQuerySchema.safeParse({ cursor }).success).toBe(false); + }, + ); + + it('defaults to 20, bounds the page size, and rejects invalid time ranges', () => { + expect(auditListQuerySchema.parse({}).limit).toBe(20); + expect(auditListQuerySchema.safeParse({ limit: 0 }).success).toBe(false); + expect(auditListQuerySchema.safeParse({ limit: 101 }).success).toBe(false); + expect( + auditListQuerySchema.safeParse({ + from: '2026-09-30T00:00:00.000Z', + to: '2026-09-01T00:00:00.000Z', + }).success, + ).toBe(false); + }); + + it('accepts an inclusive range with equal boundaries and supported filters', () => { + expect( + auditListQuerySchema.safeParse({ + actorId: 'user-1', + action: 'TRANSFER', + resourceId: 'tx-1', + from: '2026-09-29T12:00:00.000Z', + to: '2026-09-29T12:00:00.000Z', + }).success, + ).toBe(true); + }); +}); \ No newline at end of file diff --git a/src/modules/audit/audit-cursor.ts b/src/modules/audit/audit-cursor.ts new file mode 100644 index 00000000..33ce6657 --- /dev/null +++ b/src/modules/audit/audit-cursor.ts @@ -0,0 +1,61 @@ +export interface AuditCursor { + createdAt: Date; + id: string; +} + +interface EncodedAuditCursor { + v: 1; + createdAt: string; + id: string; +} + +export function encodeAuditCursor(cursor: AuditCursor): string { + const value: EncodedAuditCursor = { + v: 1, + createdAt: cursor.createdAt.toISOString(), + id: cursor.id, + }; + return Buffer.from(JSON.stringify(value)).toString('base64url'); +} + +export function decodeAuditCursor(value: string): AuditCursor { + if (!/^[A-Za-z0-9_-]{1,256}$/.test(value)) { + throw new Error('Invalid audit cursor'); + } + + try { + const decoded = Buffer.from(value, 'base64url'); + if (decoded.toString('base64url') !== value) throw new Error(); + const parsed = JSON.parse(decoded.toString('utf8')) as Partial; + if ( + parsed.v !== 1 || + typeof parsed.createdAt !== 'string' || + typeof parsed.id !== 'string' || + parsed.id.length < 1 || + parsed.id.length > 128 || + !/^[A-Za-z0-9._:-]+$/.test(parsed.id) + ) { + throw new Error(); + } + + const createdAt = new Date(parsed.createdAt); + if ( + Number.isNaN(createdAt.getTime()) || + createdAt.toISOString() !== parsed.createdAt + ) { + throw new Error(); + } + return { createdAt, id: parsed.id }; + } catch { + throw new Error('Invalid audit cursor'); + } +} + +export function isValidAuditCursor(value: string): boolean { + try { + decodeAuditCursor(value); + return true; + } catch { + return false; + } +} \ No newline at end of file diff --git a/src/modules/audit/audit-export.dto.ts b/src/modules/audit/audit-export.dto.ts index 1e31c262..613ab72a 100644 --- a/src/modules/audit/audit-export.dto.ts +++ b/src/modules/audit/audit-export.dto.ts @@ -1,19 +1,55 @@ import { z } from 'zod'; import { ApiPropertyOptional } from '@nestjs/swagger'; -export const exportAuditLogsQuerySchema = z.object({ +const exportFilters = { agentId: z.string().optional(), userId: z.string().optional(), actionType: z.string().optional(), + /** Matches a severity value stored in oldValue or newValue JSON metadata. */ + severity: z.enum(['info', 'warning', 'error', 'critical']).optional(), startDate: z.string().datetime().optional(), endDate: z.string().datetime().optional(), - limit: z.coerce.number().int().positive().max(1000).default(100), - cursor: z.string().optional(), - format: z.enum(['json', 'csv']).default('json'), -}); +}; +function validateDateRange( + value: { startDate?: string; endDate?: string }, + context: z.RefinementCtx, +) { + if (value.startDate && value.endDate && Date.parse(value.startDate) > Date.parse(value.endDate)) { + context.addIssue({ + code: z.ZodIssueCode.custom, + path: ['endDate'], + message: 'endDate must be on or after startDate', + }); + } +} + +export const exportAuditLogsQuerySchema = z + .object({ + ...exportFilters, + limit: z.coerce.number().int().positive().max(1000).default(100), + cursor: z.string().optional(), + format: z.enum(['json', 'csv']).default('json'), + }) + .superRefine(validateDateRange); + +/** Filters and pagination options accepted by the audit log page export. */ export type ExportAuditLogsQuery = z.infer; +/** Strict query contract for the batch-streamed audit export endpoint. */ +export const streamAuditLogsQuerySchema = z + .object({ + ...exportFilters, + cursor: z.string().optional(), + batchSize: z.coerce.number().int().positive().max(1000).default(250), + format: z.enum(['json', 'csv']).default('json'), + }) + .strict() + .superRefine(validateDateRange); + +/** Parsed query options for streaming an audit log export. */ +export type StreamAuditLogsQuery = z.infer; + /** Swagger model mirroring {@link ExportAuditLogsQuery}. */ export class ExportAuditLogsQueryDto { @ApiPropertyOptional({ description: 'Filter by agent UUID' }) @@ -25,18 +61,47 @@ export class ExportAuditLogsQueryDto { @ApiPropertyOptional({ description: 'Filter by audit action type', example: 'wallet.created' }) actionType?: string; - @ApiPropertyOptional({ description: 'ISO 8601 start of the export window', example: '2026-01-01T00:00:00.000Z' }) + @ApiPropertyOptional({ + enum: ['info', 'warning', 'error', 'critical'], + description: 'Filter by severity stored in oldValue or newValue metadata', + }) + severity?: 'info' | 'warning' | 'error' | 'critical'; + + @ApiPropertyOptional({ + description: 'ISO 8601 start of the export window', + example: '2026-01-01T00:00:00.000Z', + }) startDate?: string; - @ApiPropertyOptional({ description: 'ISO 8601 end of the export window', example: '2026-12-31T23:59:59.000Z' }) + @ApiPropertyOptional({ + description: 'ISO 8601 end of the export window', + example: '2026-12-31T23:59:59.000Z', + }) endDate?: string; - @ApiPropertyOptional({ description: 'Maximum entries to export (max 1000)', default: 100, example: 100 }) + @ApiPropertyOptional({ + description: 'Maximum entries to export (max 1000)', + default: 100, + example: 100, + }) limit?: number; @ApiPropertyOptional({ description: 'Opaque pagination cursor from a previous page' }) cursor?: string; - @ApiPropertyOptional({ enum: ['json', 'csv'], description: 'Export format (default json)', default: 'json' }) + @ApiPropertyOptional({ + enum: ['json', 'csv'], + description: 'Export format (default json)', + default: 'json', + }) format?: 'json' | 'csv'; } + +/** Swagger model mirroring {@link StreamAuditLogsQuery}. */ +export class StreamAuditLogsQueryDto extends ExportAuditLogsQueryDto { + @ApiPropertyOptional({ + description: 'Number of rows fetched per database batch (max 1000)', + default: 250, + }) + batchSize?: number; +} diff --git a/src/modules/audit/audit-export.spec.ts b/src/modules/audit/audit-export.spec.ts index 19a81ba7..a069a35b 100644 --- a/src/modules/audit/audit-export.spec.ts +++ b/src/modules/audit/audit-export.spec.ts @@ -2,10 +2,12 @@ import { describe, it, expect, vi } from 'vitest'; import { AuditService } from './audit.service'; import { AuditRepository } from './audit.repository'; import { AuditHashService } from './audit-hash.service'; +import { streamAuditLogsQuerySchema } from './audit-export.dto'; describe('AuditService - Export Compliance', () => { const mockRepository = { exportLogs: vi.fn(), + streamLogs: vi.fn(), create: vi.fn(), findManyAndCount: vi.fn(), findById: vi.fn(), @@ -51,7 +53,13 @@ describe('AuditService - Export Compliance', () => { expect(result.format).toBe('json'); expect(result.count).toBe(1); - expect(result.data).toEqual(mockLogs); + expect(result.data).toEqual([ + { + ...mockLogs[0], + oldValue: { amount: 10 }, + newValue: { amount: 20 }, + }, + ]); expect(mockRepository.exportLogs).toHaveBeenCalledWith( expect.objectContaining({ organizationId: 'org-123', @@ -93,6 +101,123 @@ describe('AuditService - Export Compliance', () => { expect(result.data).toContain('"POLICY_OVERRIDE,ADMIN"'); }); + it('redacts sensitive nested payload values in JSON and CSV exports', async () => { + const mockLog = { + id: 'log-secret', + organizationId: 'org-123', + userId: null, + action: 'wallet.updated', + entity: 'Wallet', + entityId: 'wallet-1', + ipAddress: null, + device: null, + oldValue: { credential: { apiKey: 'old-key' }, safe: 'visible' }, + newValue: [{ password: 'secret', amount: 25 }], + createdAt: new Date('2026-08-28T10:00:00Z'), + user: null, + }; + mockRepository.exportLogs.mockResolvedValueOnce([mockLog]); + + const result = await auditService.export('org-123', { format: 'json', limit: 10 }); + + if (!Array.isArray(result.data)) throw new Error('Expected JSON export records'); + expect(result.data[0].oldValue).toEqual({ + credential: { apiKey: '[REDACTED]' }, + safe: 'visible', + }); + expect(result.data[0].newValue).toEqual([{ password: '[REDACTED]', amount: 25 }]); + + mockRepository.exportLogs.mockResolvedValueOnce([mockLog]); + const csvResult = await auditService.export('org-123', { format: 'csv', limit: 10 }); + + expect(csvResult.format).toBe('csv'); + expect(csvResult.data).toContain('[REDACTED]'); + expect(csvResult.data).not.toContain('old-key'); + expect(csvResult.data).not.toContain('password\":\"secret'); + }); + + it('streams a complete JSON array in batches and redacts payload fields', async () => { + const firstRecord = { + id: 'stream-1', + organizationId: 'org-123', + userId: null, + action: 'wallet.updated', + entity: 'Wallet', + entityId: 'wallet-1', + ipAddress: null, + device: null, + oldValue: { token: 'secret' }, + newValue: { amount: 10 }, + createdAt: new Date('2026-08-28T10:00:00Z'), + user: null, + }; + mockRepository.streamLogs.mockImplementation(async function* () { + yield firstRecord; + yield { ...firstRecord, id: 'stream-2' }; + }); + + const stream = auditService.streamExport('org-123', { + format: 'json', + batchSize: 1, + }); + const chunks: Buffer[] = []; + for await (const chunk of stream) { + chunks.push(Buffer.isBuffer(chunk) ? chunk : Buffer.from(chunk)); + } + const result = JSON.parse(Buffer.concat(chunks).toString()); + + expect(result).toHaveLength(2); + expect(result[0].oldValue).toEqual({ token: '[REDACTED]' }); + expect(mockRepository.streamLogs).toHaveBeenCalledWith( + { organizationId: 'org-123' }, + 1, + undefined, + ); + }); + + it('streams escaped CSV rows and applies the severity metadata filter', async () => { + mockRepository.streamLogs.mockImplementation(async function* () {}); + + const stream = auditService.streamExport('org-123', { + format: 'csv', + batchSize: 25, + severity: 'critical', + }); + const chunks: Buffer[] = []; + for await (const chunk of stream) { + chunks.push(Buffer.isBuffer(chunk) ? chunk : Buffer.from(chunk)); + } + + expect(Buffer.concat(chunks).toString()).toBe( + 'id,organizationId,userId,userEmail,action,entity,entityId,ipAddress,device,oldValue,newValue,createdAt\n', + ); + expect(mockRepository.streamLogs).toHaveBeenCalledWith( + { + organizationId: 'org-123', + AND: [ + { + OR: [ + { oldValue: { path: ['severity'], equals: 'critical' } }, + { newValue: { path: ['severity'], equals: 'critical' } }, + ], + }, + ], + }, + 25, + undefined, + ); + }); + + it('rejects unknown stream query keys and reversed date ranges', () => { + expect(() => streamAuditLogsQuerySchema.parse({ unexpected: 'value' })).toThrow(); + expect(() => + streamAuditLogsQuerySchema.parse({ + startDate: '2026-08-29T00:00:00.000Z', + endDate: '2026-08-28T00:00:00.000Z', + }), + ).toThrow('endDate must be on or after startDate'); + }); + it('should handle empty records gracefully', async () => { mockRepository.exportLogs.mockResolvedValueOnce([]); diff --git a/src/modules/audit/audit-list.dto.ts b/src/modules/audit/audit-list.dto.ts new file mode 100644 index 00000000..5329b700 --- /dev/null +++ b/src/modules/audit/audit-list.dto.ts @@ -0,0 +1,25 @@ +import { z } from 'zod'; +import { isValidAuditCursor } from './audit-cursor'; + +export const auditListQuerySchema = z + .object({ + cursor: z.string().max(256).refine(isValidAuditCursor, 'Invalid cursor').optional(), + limit: z.coerce.number().int().min(1).max(100).default(20), + actorId: z.string().min(1).max(128).optional(), + action: z.string().min(1).max(120).optional(), + resourceId: z.string().min(1).max(128).optional(), + from: z.string().datetime({ offset: true }).optional(), + to: z.string().datetime({ offset: true }).optional(), + }) + .strict() + .superRefine((query, context) => { + if (query.from && query.to && new Date(query.from) > new Date(query.to)) { + context.addIssue({ + code: z.ZodIssueCode.custom, + path: ['to'], + message: '`to` must be greater than or equal to `from`', + }); + } + }); + +export type AuditListQuery = z.infer; \ No newline at end of file diff --git a/src/modules/audit/audit.controller.ts b/src/modules/audit/audit.controller.ts index d4e3139e..6a4abd1e 100644 --- a/src/modules/audit/audit.controller.ts +++ b/src/modules/audit/audit.controller.ts @@ -14,15 +14,15 @@ import { AuditService } from './audit.service'; import { CurrentUser } from '../../common/decorators/current-user.decorator'; import { Roles } from '../../common/decorators/roles.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; -import { - PaginationQuery, - paginationQuerySchema, -} from '../../common/helpers/pagination'; import { ExportAuditLogsQuery, exportAuditLogsQuerySchema, ExportAuditLogsQueryDto, + StreamAuditLogsQuery, + streamAuditLogsQuerySchema, + StreamAuditLogsQueryDto, } from './audit-export.dto'; +import { AuditListQuery, auditListQuerySchema } from './audit-list.dto'; /** Read-only access to the append-only audit trail. Restricted to auditors/admins. */ @ApiTags('audit') @@ -73,22 +73,57 @@ export class AuditController { @ApiOperation({ summary: 'List audit log entries for the organization', description: - 'Returns a paginated list of audit log entries. Supports filtering by action, date range, and agent.', + 'Returns a cursor-paginated list of audit log entries, newest first. Supports filtering by actor, action, resource and date range.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiQuery({ name: 'cursor', required: false, type: String, description: 'Opaque pagination cursor from a previous page' }) + @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Max entries to return (default 20, max 100)' }) + @ApiQuery({ name: 'actorId', required: false, type: String, description: 'Filter by acting user UUID' }) @ApiQuery({ name: 'action', required: false, type: String, description: 'Filter by audit action type' }) - @ApiQuery({ name: 'agentId', required: false, type: String, description: 'Filter by agent UUID' }) - @ApiResponse({ status: 200, description: 'Paginated list of audit log entries' }) + @ApiQuery({ name: 'resourceId', required: false, type: String, description: 'Filter by affected entity UUID' }) + @ApiQuery({ name: 'from', required: false, type: String, description: 'ISO 8601 start of the date range' }) + @ApiQuery({ name: 'to', required: false, type: String, description: 'ISO 8601 end of the date range' }) + @ApiResponse({ status: 200, description: 'Cursor-paginated list of audit log entries' }) @ApiResponse({ status: 401, description: 'Not authenticated' }) @ApiResponse({ status: 403, description: 'Insufficient permissions' }) list( @CurrentUser('organizationId') organizationId: string, - @Query(new ZodValidationPipe(paginationQuerySchema)) query: PaginationQuery, + @Query(new ZodValidationPipe(auditListQuerySchema)) query: AuditListQuery, ) { return this.auditService.list(organizationId, query); } + @Get('export/stream') + @ApiOperation({ + summary: 'Stream audit log entries for large compliance exports', + description: + 'Streams audit log entries in bounded batches (CSV or JSON) without loading the full result set into memory.', + }) + @ApiQuery({ type: StreamAuditLogsQueryDto }) + @ApiProduces('text/csv', 'application/json') + @ApiResponse({ status: 200, description: 'Streamed audit log export (CSV or JSON)' }) + @ApiResponse({ status: 401, description: 'Not authenticated' }) + @ApiResponse({ status: 403, description: 'Insufficient permissions (requires OWNER, ADMIN, or AUDITOR)' }) + async streamExport( + @CurrentUser('organizationId') organizationId: string, + @Query(new ZodValidationPipe(streamAuditLogsQuerySchema)) query: StreamAuditLogsQuery, + @Res() res: Response, + ) { + res.setHeader( + 'Content-Type', + query.format === 'csv' ? 'text/csv' : 'application/json', + ); + if (query.format === 'csv') { + res.setHeader( + 'Content-Disposition', + `attachment; filename="audit-logs-${organizationId}-${Date.now()}.csv"`, + ); + } + for await (const chunk of this.auditService.streamExport(organizationId, query)) { + res.write(chunk); + } + res.end(); + } + @Get('integrity/verify') @ApiOperation({ summary: 'Verify the integrity of the entire audit chain', diff --git a/src/modules/audit/audit.listener.spec.ts b/src/modules/audit/audit.listener.spec.ts new file mode 100644 index 00000000..f1c58e82 --- /dev/null +++ b/src/modules/audit/audit.listener.spec.ts @@ -0,0 +1,68 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { AuditListener } from './audit.listener'; +import { AuditService } from './audit.service'; +import { DomainEventEnvelope } from '../../events/domain-event.types'; + +describe('AuditListener', () => { + let listener: AuditListener; + let auditService: { record: ReturnType }; + + beforeEach(() => { + auditService = { record: vi.fn().mockResolvedValue(undefined) }; + listener = new AuditListener(auditService as unknown as AuditService); + }); + + it('persists event identity, actor, timestamp, correlation, and outcome', async () => { + const occurredAt = new Date('2026-09-28T12:00:00.000Z'); + const envelope: DomainEventEnvelope = { + eventId: 'event-1', + name: 'policy.violated', + organizationId: 'org-1', + actorId: 'user-1', + aggregateType: 'agent', + aggregateId: 'agent-1', + correlationId: 'request-1', + occurredAt, + payload: { violation: 'daily-limit' }, + }; + + await listener.handleDomainEvent(envelope); + + expect(auditService.record).toHaveBeenCalledWith( + expect.objectContaining({ + organizationId: 'org-1', + userId: 'user-1', + action: 'policy.violated', + entity: 'agent', + entityId: 'agent-1', + sourceEventId: 'event-1', + requestId: 'request-1', + createdAt: occurredAt, + newValue: expect.objectContaining({ + eventId: 'event-1', + success: false, + payload: { violation: 'daily-limit' }, + }), + }), + ); + }); + + it('bounds large event payloads before persisting', async () => { + const envelope: DomainEventEnvelope = { + eventId: 'event-large', + name: 'transaction.created', + organizationId: 'org-1', + aggregateType: 'transaction', + occurredAt: new Date(), + payload: { detail: 'x'.repeat(20_000) }, + }; + + await listener.handleDomainEvent(envelope); + + expect(auditService.record).toHaveBeenCalledWith( + expect.objectContaining({ + newValue: expect.objectContaining({ eventId: 'event-large', truncated: true }), + }), + ); + }); +}); \ No newline at end of file diff --git a/src/modules/audit/audit.listener.ts b/src/modules/audit/audit.listener.ts index 6a566391..54ab61b5 100644 --- a/src/modules/audit/audit.listener.ts +++ b/src/modules/audit/audit.listener.ts @@ -1,7 +1,12 @@ import { Injectable, Logger } from '@nestjs/common'; import { OnEvent } from '@nestjs/event-emitter'; +import { Prisma } from '@prisma/client'; import { AuditService } from './audit.service'; import { DomainEventEnvelope } from '../../events/domain-event.types'; +import { RequestContext } from '../../common/context/request-context'; + +const MAX_AUDIT_EVENT_BYTES = 16_384; +import { DOMAIN_EVENT_ENVELOPE } from '../../events/domain-event.types'; /** * Subscribes to every domain event (wildcard) and appends an audit-log row. @@ -14,19 +19,46 @@ export class AuditListener { constructor(private readonly auditService: AuditService) {} - @OnEvent('**') + @OnEvent(DOMAIN_EVENT_ENVELOPE) async handleDomainEvent(envelope: DomainEventEnvelope): Promise { - if (!envelope?.organizationId) { + if (!envelope?.eventId) { return; } try { + const requestContext = RequestContext.getStore(); + const organizationId = envelope.organizationId ?? requestContext?.principal?.organizationId; + if (!organizationId) { + return; + } + const correlationId = envelope.correlationId ?? requestContext?.identity.correlationId; + const value = { + eventId: envelope.eventId, + occurredAt: envelope.occurredAt.toISOString(), + correlationId: correlationId ?? null, + success: envelope.name !== 'policy.violated' && !envelope.name.endsWith('.failed'), + payload: envelope.payload, + }; + const serialized = JSON.stringify(value); + const newValue = Buffer.byteLength(serialized, 'utf8') > MAX_AUDIT_EVENT_BYTES + ? { + eventId: envelope.eventId, + truncated: true, + originalBytes: Buffer.byteLength(serialized, 'utf8'), + preview: serialized.slice(0, MAX_AUDIT_EVENT_BYTES), + } + : value; + await this.auditService.record({ - organizationId: envelope.organizationId, - userId: envelope.actorId ?? null, + organizationId, + userId: envelope.actorId ?? requestContext?.principal?.userId ?? null, action: envelope.name, entity: envelope.aggregateType, entityId: envelope.aggregateId ?? null, - newValue: envelope.payload as object, + newValue: newValue as Prisma.InputJsonValue, + sourceEventId: envelope.eventId, + requestId: correlationId ?? requestContext?.identity.requestId ?? null, + ipAddress: requestContext?.identity.ip ?? null, + createdAt: envelope.occurredAt, }); } catch (error) { this.logger.error(`Failed to write audit log for '${envelope.name}': ${(error as Error).message}`); diff --git a/src/modules/audit/audit.repository.spec.ts b/src/modules/audit/audit.repository.spec.ts new file mode 100644 index 00000000..6d7739dc --- /dev/null +++ b/src/modules/audit/audit.repository.spec.ts @@ -0,0 +1,115 @@ +import { describe, expect, it, vi } from 'vitest'; +import { PrismaService } from '../../database/prisma.service'; +import { AuditRepository } from './audit.repository'; + +describe('AuditRepository.streamLogs', () => { + it('fetches bounded pages and advances the cursor through every row', async () => { + const findMany = vi + .fn() + .mockResolvedValueOnce([{ id: 'log-1' }, { id: 'log-2' }]) + .mockResolvedValueOnce([{ id: 'log-3' }]); + const prisma = { + auditLog: { findMany }, + } as unknown as PrismaService; + const repository = new AuditRepository(prisma); + const ids: string[] = []; + + for await (const record of repository.streamLogs({ organizationId: 'org-1' }, 2, 'start')) { + ids.push(record.id); + } + + expect(ids).toEqual(['log-1', 'log-2', 'log-3']); + expect(findMany).toHaveBeenNthCalledWith( + 1, + expect.objectContaining({ + where: { organizationId: 'org-1' }, + take: 2, + cursor: { id: 'start' }, + skip: 1, + }), + ); + expect(findMany).toHaveBeenNthCalledWith( + 2, + expect.objectContaining({ + where: { organizationId: 'org-1' }, + take: 2, + cursor: { id: 'log-2' }, + skip: 1, + }), + ); + }); +}); + +describe('AuditRepository', () => { + it('upserts event-backed audit rows by source event ID', async () => { + const prisma = { + auditLog: { + create: vi.fn(), + upsert: vi.fn().mockResolvedValue({ id: 'audit-1' }), + }, + }; + const repository = new AuditRepository(prisma as unknown as PrismaService); + const record = { + organizationId: 'org-1', + action: 'transaction.created', + entity: 'transaction', + sourceEventId: 'event-1', + }; + + await repository.create(record); + await repository.create(record); + + expect(prisma.auditLog.upsert).toHaveBeenCalledTimes(2); + expect(prisma.auditLog.upsert).toHaveBeenCalledWith( + expect.objectContaining({ where: { sourceEventId: 'event-1' }, update: {} }), + ); + expect(prisma.auditLog.create).not.toHaveBeenCalled(); + }); +}); + +describe('AuditRepository.findPage', () => { + it('uses a bounded deterministic keyset query with an ID tie-breaker', async () => { + const findMany = vi.fn().mockResolvedValue([]); + const count = vi.fn(); + const prisma = { auditLog: { findMany, count } } as unknown as PrismaService; + const repository = new AuditRepository(prisma); + const createdAt = new Date('2026-09-29T12:00:00.000Z'); + + await repository.findPage( + { organizationId: 'org-1' }, + { createdAt, id: 'audit-20' }, + 21, + ); + + expect(findMany).toHaveBeenCalledWith({ + where: { + AND: [ + { organizationId: 'org-1' }, + { + OR: [ + { createdAt: { lt: createdAt } }, + { createdAt, id: { lt: 'audit-20' } }, + ], + }, + ], + }, + take: 21, + orderBy: [{ createdAt: 'desc' }, { id: 'desc' }], + }); + expect(count).not.toHaveBeenCalled(); + }); + + it('uses a bounded first-page query without a cursor condition', async () => { + const findMany = vi.fn().mockResolvedValue([]); + const prisma = { auditLog: { findMany } } as unknown as PrismaService; + const repository = new AuditRepository(prisma); + + await repository.findPage({ organizationId: 'org-2' }, undefined, 101); + + expect(findMany).toHaveBeenCalledWith({ + where: { organizationId: 'org-2' }, + take: 101, + orderBy: [{ createdAt: 'desc' }, { id: 'desc' }], + }); + }); +}); diff --git a/src/modules/audit/audit.repository.ts b/src/modules/audit/audit.repository.ts index 6ed97534..6f7f6ab9 100644 --- a/src/modules/audit/audit.repository.ts +++ b/src/modules/audit/audit.repository.ts @@ -2,6 +2,7 @@ import { Injectable } from '@nestjs/common'; import { Prisma } from '@prisma/client'; import { PrismaService } from '../../database/prisma.service'; import { PrismaPagination } from '../../common/helpers/pagination'; +import { AuditCursor } from './audit-cursor'; export interface CreateAuditLogData { organizationId: string; @@ -14,6 +15,8 @@ export interface CreateAuditLogData { ipAddress?: string | null; device?: string | null; requestId?: string | null; + sourceEventId?: string | null; + createdAt?: Date; previousHash?: string | null; hash?: string | null; } @@ -24,22 +27,32 @@ export class AuditRepository { constructor(private readonly prisma: PrismaService) {} create(data: CreateAuditLogData) { - return this.prisma.auditLog.create({ - data: { - organizationId: data.organizationId, - userId: data.userId ?? null, - action: data.action, - entity: data.entity, - entityId: data.entityId ?? null, - oldValue: data.oldValue, - newValue: data.newValue, - ipAddress: data.ipAddress ?? null, - device: data.device ?? null, - requestId: data.requestId ?? null, - previousHash: data.previousHash ?? null, - hash: data.hash ?? null, - }, - }); + const create = { + organizationId: data.organizationId, + userId: data.userId ?? null, + action: data.action, + entity: data.entity, + entityId: data.entityId ?? null, + oldValue: data.oldValue, + newValue: data.newValue, + ipAddress: data.ipAddress ?? null, + device: data.device ?? null, + requestId: data.requestId ?? null, + sourceEventId: data.sourceEventId ?? null, + previousHash: data.previousHash ?? null, + hash: data.hash ?? null, + ...(data.createdAt ? { createdAt: data.createdAt } : {}), + }; + + if (data.sourceEventId) { + return this.prisma.auditLog.upsert({ + where: { sourceEventId: data.sourceEventId }, + create, + update: {}, + }); + } + + return this.prisma.auditLog.create({ data: create }); } async findManyAndCount(where: Prisma.AuditLogWhereInput, pagination: PrismaPagination) { @@ -50,6 +63,23 @@ export class AuditRepository { return { items, total }; } + findPage(where: Prisma.AuditLogWhereInput, cursor: AuditCursor | undefined, limit: number) { + const cursorWhere: Prisma.AuditLogWhereInput | undefined = cursor + ? { + OR: [ + { createdAt: { lt: cursor.createdAt } }, + { createdAt: cursor.createdAt, id: { lt: cursor.id } }, + ], + } + : undefined; + + return this.prisma.auditLog.findMany({ + where: cursorWhere ? { AND: [where, cursorWhere] } : where, + take: limit, + orderBy: [{ createdAt: 'desc' }, { id: 'desc' }], + }); + } + async exportLogs( where: Prisma.AuditLogWhereInput, limit: number, @@ -72,6 +102,44 @@ export class AuditRepository { }); } + /** + * Reads audit rows in bounded batches so exports do not load the full result + * set into memory. The last row id is used as the next Prisma cursor. + */ + async *streamLogs( + where: Prisma.AuditLogWhereInput, + batchSize: number, + cursor?: string, + ): AsyncGenerator> { + let nextCursor = cursor; + + while (true) { + const records = await this.prisma.auditLog.findMany({ + where, + take: batchSize, + ...(nextCursor ? { cursor: { id: nextCursor }, skip: 1 } : {}), + orderBy: [{ createdAt: 'desc' }, { id: 'desc' }], + include: { + user: { + select: { + id: true, + email: true, + name: true, + }, + }, + }, + }); + + if (records.length === 0) return; + + yield* records; + if (records.length < batchSize) return; + nextCursor = records[records.length - 1].id; + } + } + findById(organizationId: string, id: string) { return this.prisma.auditLog.findFirst({ where: { id, organizationId } }); } diff --git a/src/modules/audit/audit.service.spec.ts b/src/modules/audit/audit.service.spec.ts index 8f8579c4..cf2cfc16 100644 --- a/src/modules/audit/audit.service.spec.ts +++ b/src/modules/audit/audit.service.spec.ts @@ -2,22 +2,32 @@ import { beforeEach, describe, expect, it, vi } from 'vitest'; import { AuditService } from './audit.service'; import { AuditRepository } from './audit.repository'; import { AuditHashService } from './audit-hash.service'; +import { AuditListQuery } from './audit-list.dto'; +import { decodeAuditCursor } from './audit-cursor'; describe('AuditService', () => { - let repository: { create: ReturnType }; + let repository: { + create: ReturnType; + findManyAndCount: ReturnType; + findPage: ReturnType; + }; let hashService: { getLatestHash: ReturnType; computeEntryHash: ReturnType; }; let service: AuditService; + const baseQuery: AuditListQuery = { limit: 20 }; + beforeEach(() => { - repository = { create: vi.fn().mockResolvedValue({ id: 'audit-1' }) }; + repository = { + create: vi.fn().mockResolvedValue({ id: 'audit-1' }), + findManyAndCount: vi.fn().mockResolvedValue({ items: [], total: 0 }), + findPage: vi.fn().mockResolvedValue([]), + }; hashService = { getLatestHash: vi.fn().mockResolvedValue('prev-hash'), - computeEntryHash: vi - .fn() - .mockReturnValue({ previousHash: 'prev-hash', hash: 'new-hash' }), + computeEntryHash: vi.fn().mockReturnValue({ previousHash: 'prev-hash', hash: 'new-hash' }), }; service = new AuditService( repository as unknown as AuditRepository, @@ -38,7 +48,11 @@ describe('AuditService', () => { }); expect(repository.create).toHaveBeenCalledWith( - expect.objectContaining({ requestId: 'req_01HXYZ', hash: 'new-hash', previousHash: 'prev-hash' }), + expect.objectContaining({ + requestId: 'req_01HXYZ', + hash: 'new-hash', + previousHash: 'prev-hash', + }), ); const hashInput = hashService.computeEntryHash.mock.calls[0][0]; @@ -56,4 +70,59 @@ describe('AuditService', () => { expect(repository.create).toHaveBeenCalledWith(expect.objectContaining({ requestId: null })); }); + + describe('list', () => { + it('returns first-page results and an opaque cursor when more records exist', async () => { + const rows = Array.from({ length: 21 }, (_, index) => ({ + id: `audit-${index}`, + createdAt: new Date(`2026-09-29T00:00:${String(index).padStart(2, '0')}.000Z`), + })); + repository.findPage.mockResolvedValue(rows); + + const result = await service.list('org-1', baseQuery); + + expect(result.items).toHaveLength(20); + expect(result.meta).toEqual({ limit: 20, hasNext: true, nextCursor: expect.any(String) }); + expect(decodeAuditCursor(result.meta.nextCursor!).id).toBe('audit-19'); + expect(repository.findPage).toHaveBeenCalledWith({ organizationId: 'org-1' }, undefined, 21); + }); + + it('uses the cursor and tenant-scoped combined filters for the next page', async () => { + const createdAt = new Date('2026-09-29T12:00:00.000Z'); + const cursor = Buffer.from(JSON.stringify({ v: 1, createdAt: createdAt.toISOString(), id: 'audit-20' })).toString('base64url'); + repository.findPage.mockResolvedValue([{ id: 'audit-21', createdAt }]); + + const result = await service.list('org-1', { + limit: 10, + cursor, + actorId: 'user-1', + action: 'TRANSFER', + resourceId: 'tx-1', + from: '2026-09-01T00:00:00.000Z', + to: '2026-09-30T00:00:00.000Z', + }); + + const [where, decodedCursor, take] = repository.findPage.mock.calls[0]; + expect(where).toMatchObject({ + organizationId: 'org-1', + userId: 'user-1', + action: 'TRANSFER', + entityId: 'tx-1', + createdAt: { + gte: new Date('2026-09-01T00:00:00.000Z'), + lte: new Date('2026-09-30T00:00:00.000Z'), + }, + }); + expect(decodedCursor).toEqual({ createdAt, id: 'audit-20' }); + expect(take).toBe(11); + expect(result.meta).toEqual({ limit: 10, hasNext: false, nextCursor: null }); + }); + + it('returns no next cursor on an empty final page', async () => { + const result = await service.list('org-1', baseQuery); + expect(result.items).toEqual([]); + expect(result.meta.hasNext).toBe(false); + expect(result.meta.nextCursor).toBeNull(); + }); + }); }); diff --git a/src/modules/audit/audit.service.ts b/src/modules/audit/audit.service.ts index d2abe8bc..c36b6104 100644 --- a/src/modules/audit/audit.service.ts +++ b/src/modules/audit/audit.service.ts @@ -2,14 +2,11 @@ import { Injectable } from '@nestjs/common'; import { Prisma } from '@prisma/client'; import { AuditRepository, CreateAuditLogData } from './audit.repository'; import { AuditHashService } from './audit-hash.service'; -import { - buildPaginationMeta, - PaginationQuery, - toPrismaPagination, -} from '../../common/helpers/pagination'; -import { Paginated } from '../../common/interfaces/api-response.interface'; - -const SORTABLE = ['createdAt', 'action', 'entity']; +import { CursorPaginated } from '../../common/interfaces/api-response.interface'; +import { AuditListQuery } from './audit-list.dto'; +import { decodeAuditCursor, encodeAuditCursor } from './audit-cursor'; +import { ExportAuditLogsQuery, StreamAuditLogsQuery } from './audit-export.dto'; +import { sanitizeAuditPayload } from '../../common/helpers/audit-sanitizer'; /** An audit row as returned by `AuditRepository.exportLogs`, with its joined user. */ type ExportedAuditLog = Prisma.AuditLogGetPayload<{ @@ -56,62 +53,23 @@ export class AuditService { }); } - async list(organizationId: string, query: PaginationQuery) { - const where: Prisma.AuditLogWhereInput = { organizationId }; - if (query.search) { - where.OR = [ - { action: { contains: query.search, mode: 'insensitive' } }, - { entity: { contains: query.search, mode: 'insensitive' } }, - { entityId: { contains: query.search, mode: 'insensitive' } }, - ]; - } - if (query.filter) { - where.entity = query.filter; - } - const pagination = toPrismaPagination(query, SORTABLE); - const { items, total } = await this.repository.findManyAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); - } - - async export(organizationId: string, query: import('./audit-export.dto').ExportAuditLogsQuery) { - const where: Prisma.AuditLogWhereInput = { organizationId }; + /** Cursor-paginated audit log listing, newest first. */ + async list(organizationId: string, query: AuditListQuery): Promise> { + const where = this.buildFilterWhere(organizationId, query); + const cursor = query.cursor ? decodeAuditCursor(query.cursor) : undefined; + const take = query.limit + 1; - if (query.userId) { - where.userId = query.userId; - } + const rows = (await this.repository.findPage(where, cursor, take)) as ExportedAuditLog[]; + const hasNext = rows.length > query.limit; + const items = hasNext ? rows.slice(0, query.limit) : rows; + const last = items[items.length - 1]; + const nextCursor = hasNext && last ? encodeAuditCursor({ createdAt: last.createdAt, id: last.id }) : null; - if (query.actionType) { - where.action = query.actionType; - } - - if (query.agentId) { - where.OR = [ - { entityId: query.agentId }, - { - oldValue: { - path: ['agentId'], - equals: query.agentId, - }, - }, - { - newValue: { - path: ['agentId'], - equals: query.agentId, - }, - }, - ]; - } - - if (query.startDate || query.endDate) { - where.createdAt = {}; - if (query.startDate) { - where.createdAt.gte = new Date(query.startDate); - } - if (query.endDate) { - where.createdAt.lte = new Date(query.endDate); - } - } + return new CursorPaginated(items, { limit: query.limit, hasNext, nextCursor }); + } + async export(organizationId: string, query: ExportAuditLogsQuery) { + const where = this.buildFilterWhere(organizationId, query); const limit = Math.min(query.limit ?? 100, 1000); const records = await this.repository.exportLogs(where, limit, query.cursor); @@ -121,20 +79,156 @@ export class AuditService { items = records.slice(0, limit); nextCursor = items[items.length - 1]?.id ?? null; } + const sanitized = items.map((item) => this.sanitizeRecord(item)); if (query.format === 'csv') { - const csv = this.formatAsCsv(items); - return { format: 'csv', data: csv, count: items.length, nextCursor }; + const csv = this.formatAsCsv(sanitized); + return { format: 'csv', data: csv, count: sanitized.length, nextCursor }; } return { format: 'json', - data: items, - count: items.length, + data: sanitized, + count: sanitized.length, nextCursor, }; } + /** + * Streams an export in bounded batches so large exports never load the + * full result set into memory. Yields Buffer chunks: a JSON array + * (one record per chunk, wrapped by `[`/`]`) or raw CSV rows. + */ + async *streamExport(organizationId: string, query: StreamAuditLogsQuery): AsyncGenerator { + const where = this.buildFilterWhere(organizationId, query); + const stream = this.repository.streamLogs(where, query.batchSize, query.cursor); + + if (query.format === 'csv') { + const headers = [ + 'id', + 'organizationId', + 'userId', + 'userEmail', + 'action', + 'entity', + 'entityId', + 'ipAddress', + 'device', + 'oldValue', + 'newValue', + 'createdAt', + ]; + yield Buffer.from(`${headers.join(',')}\n`); + for await (const record of stream) { + const row = this.formatCsvRow(this.sanitizeRecord(record)); + yield Buffer.from(`${row}\n`); + } + return; + } + + let first = true; + yield Buffer.from('['); + for await (const record of stream) { + const sanitized = this.sanitizeRecord(record); + yield Buffer.from(`${first ? '' : ','}${JSON.stringify(sanitized)}`); + first = false; + } + yield Buffer.from(']'); + } + + /** Builds the shared tenant + filter predicate used by list, export and streamExport. */ + private buildFilterWhere( + organizationId: string, + query: { + userId?: string; + actionType?: string; + agentId?: string; + severity?: string; + startDate?: string; + endDate?: string; + actorId?: string; + action?: string; + resourceId?: string; + from?: string; + to?: string; + }, + ): Prisma.AuditLogWhereInput { + const where: Prisma.AuditLogWhereInput = { organizationId }; + const andConditions: Prisma.AuditLogWhereInput[] = []; + + if (query.userId) where.userId = query.userId; + if (query.actorId) where.userId = query.actorId; + if (query.actionType) where.action = query.actionType; + if (query.action) where.action = query.action; + if (query.resourceId) where.entityId = query.resourceId; + + if (query.agentId) { + andConditions.push({ + OR: [ + { entityId: query.agentId }, + { oldValue: { path: ['agentId'], equals: query.agentId } }, + { newValue: { path: ['agentId'], equals: query.agentId } }, + ], + }); + } + + if (query.severity) { + andConditions.push({ + OR: [ + { oldValue: { path: ['severity'], equals: query.severity } }, + { newValue: { path: ['severity'], equals: query.severity } }, + ], + }); + } + + const gte = query.startDate ?? query.from; + const lte = query.endDate ?? query.to; + if (gte || lte) { + where.createdAt = {}; + if (gte) where.createdAt.gte = new Date(gte); + if (lte) where.createdAt.lte = new Date(lte); + } + + if (andConditions.length > 0) where.AND = andConditions; + + return where; + } + + /** Redacts sensitive payload fields from a raw audit row before it leaves the service. */ + private sanitizeRecord(record: ExportedAuditLog): ExportedAuditLog { + return { + ...record, + oldValue: sanitizeAuditPayload(record.oldValue), + newValue: sanitizeAuditPayload(record.newValue), + }; + } + + private formatCsvRow(r: ExportedAuditLog): string { + const escapeCsvField = (value: unknown): string => { + if (value === null || value === undefined) return ''; + const str = typeof value === 'object' ? JSON.stringify(value) : String(value); + if (str.includes(',') || str.includes('"') || str.includes('\n') || str.includes('\r')) { + return `"${str.replace(/"/g, '""')}"`; + } + return str; + }; + + return [ + escapeCsvField(r.id), + escapeCsvField(r.organizationId), + escapeCsvField(r.userId), + escapeCsvField(r.user?.email ?? ''), + escapeCsvField(r.action), + escapeCsvField(r.entity), + escapeCsvField(r.entityId), + escapeCsvField(r.ipAddress), + escapeCsvField(r.device), + escapeCsvField(r.oldValue), + escapeCsvField(r.newValue), + escapeCsvField(r.createdAt ? new Date(r.createdAt).toISOString() : ''), + ].join(','); + } + formatAsCsv(records: ExportedAuditLog[]): string { const headers = [ 'id', @@ -150,35 +244,7 @@ export class AuditService { 'newValue', 'createdAt', ]; - - const escapeCsvField = (value: unknown): string => { - if (value === null || value === undefined) return ''; - const str = typeof value === 'object' ? JSON.stringify(value) : String(value); - if (str.includes(',') || str.includes('"') || str.includes('\n') || str.includes('\r')) { - return `"${str.replace(/"/g, '""')}"`; - } - return str; - }; - - const lines = [headers.join(',')]; - for (const r of records) { - const row = [ - escapeCsvField(r.id), - escapeCsvField(r.organizationId), - escapeCsvField(r.userId), - escapeCsvField(r.user?.email ?? ''), - escapeCsvField(r.action), - escapeCsvField(r.entity), - escapeCsvField(r.entityId), - escapeCsvField(r.ipAddress), - escapeCsvField(r.device), - escapeCsvField(r.oldValue), - escapeCsvField(r.newValue), - escapeCsvField(r.createdAt ? new Date(r.createdAt).toISOString() : ''), - ]; - lines.push(row.join(',')); - } - + const lines = [headers.join(','), ...records.map((r) => this.formatCsvRow(r))]; return lines.join('\n'); } diff --git a/src/modules/auth/api-key.strategy.ts b/src/modules/auth/api-key.strategy.ts index ee5b3078..08130edd 100644 --- a/src/modules/auth/api-key.strategy.ts +++ b/src/modules/auth/api-key.strategy.ts @@ -70,6 +70,7 @@ export class ApiKeyStrategy extends PassportStrategy(HeaderApiKeyPassportStrateg const principal: AuthenticatedApiKey = { id: apiKey.id, keyId: apiKey.id, + apiKeyId: apiKey.id, organizationId: apiKey.organizationId, createdById: apiKey.createdById, name: apiKey.name, diff --git a/src/modules/auth/auth.module.ts b/src/modules/auth/auth.module.ts index 5f64f557..667cfa9c 100644 --- a/src/modules/auth/auth.module.ts +++ b/src/modules/auth/auth.module.ts @@ -1,7 +1,6 @@ import { Module } from '@nestjs/common'; import { JwtModule } from '@nestjs/jwt'; import { PassportModule } from '@nestjs/passport'; -import Redis from 'ioredis'; import { AuthController } from './auth.controller'; import { AuthService } from './auth.service'; import { JwtStrategy } from './jwt.strategy'; @@ -10,9 +9,10 @@ import { ApiKeyGuard } from '../../common/guards/api-key.guard'; import { ApiKeyAuthGuard } from '../../common/guards/api-key-auth.guard'; import { ScopesGuard } from '../../common/guards/scopes.guard'; import { TokenBlacklistService } from './services/token-blacklist.service'; +import { TokenVerificationCacheService } from './services/token-verification-cache.service'; +import { CacheService } from '../../common/cache/cache.service'; import { PasskeyController } from './controllers/passkey.controller'; import { PasskeyService } from './services/passkey.service'; -import { redisConfig } from '../../config/redis.config'; /** * Authentication module. Registers passport-jwt and api-key strategies and a bare @@ -20,27 +20,18 @@ import { redisConfig } from '../../config/redis.config'; * access and refresh tokens can use different signing keys). Also provides the * Redis client used by the token blacklist, which lets logout / credential * rotation invalidate in-flight JWTs before they naturally expire. + * + * Revocation answers are cached by {@link TokenVerificationCacheService} over + * the shared {@link REDIS_CLIENT} (via {@link CacheService}) so authenticated + * requests avoid one Redis round trip each; every revocation path invalidates + * the cached entry. */ @Module({ - imports: [ - PassportModule.register({ defaultStrategy: 'jwt' }), - JwtModule.register({}), - ], + imports: [PassportModule.register({ defaultStrategy: 'jwt' }), JwtModule.register({})], controllers: [AuthController, PasskeyController], providers: [ - { - provide: Redis, - useFactory: (): Redis => { - const config = redisConfig(); - return new Redis({ - host: config.host, - port: config.port, - password: config.password || undefined, - db: config.db, - lazyConnect: true, - }); - }, - }, + CacheService, + TokenVerificationCacheService, AuthService, JwtStrategy, ApiKeyStrategy, @@ -59,6 +50,7 @@ import { redisConfig } from '../../config/redis.config'; ScopesGuard, PasskeyService, TokenBlacklistService, + TokenVerificationCacheService, ], }) -export class AuthModule {} \ No newline at end of file +export class AuthModule {} diff --git a/src/modules/auth/auth.service.ts b/src/modules/auth/auth.service.ts index 865e8879..387d682d 100644 --- a/src/modules/auth/auth.service.ts +++ b/src/modules/auth/auth.service.ts @@ -20,6 +20,7 @@ import { EventBusService } from '../../events/event-bus.service'; import { DomainEventName } from '../../events/event-names'; import { LoginInput, RegisterInput } from './auth.dto'; import { TokenBlacklistService } from './services/token-blacklist.service'; +import { TokenVerificationCacheService } from './services/token-verification-cache.service'; export interface TokenPair { accessToken: string; @@ -63,6 +64,7 @@ export class AuthService { private readonly jwt: JwtService, private readonly eventBus: EventBusService, private readonly tokenBlacklist: TokenBlacklistService, + private readonly verificationCache: TokenVerificationCacheService, config: ConfigService, ) { this.auth = config.getOrThrow('auth'); @@ -181,6 +183,9 @@ export class AuthService { where: { id: session.id }, data: { revokedAt: new Date() }, }); + // In-flight access tokens of the rotated session must re-verify against + // the blacklist instead of a cached answer. + await this.verificationCache.invalidateOnRefreshRotation(session.id); return this.issueTokens(session.user, { device: session.device ?? undefined, @@ -203,6 +208,9 @@ export class AuthService { this.auth.accessTtl, this.auth.refreshTtl, ); + // Belt-and-braces: the blacklist service already invalidates the cached + // verification answer; keep logout self-contained even if that changes. + await this.verificationCache.invalidateSessionRevocation(sessionId); return { success: true }; } diff --git a/src/modules/auth/jwt.strategy.ts b/src/modules/auth/jwt.strategy.ts index 5e148a47..f087d799 100644 --- a/src/modules/auth/jwt.strategy.ts +++ b/src/modules/auth/jwt.strategy.ts @@ -4,6 +4,7 @@ import { ExtractJwt, Strategy } from 'passport-jwt'; import { ConfigService } from '@nestjs/config'; import { AuthConfig } from '../../config/auth.config'; import { TokenBlacklistService } from './services/token-blacklist.service'; +import { TokenVerificationCacheService } from './services/token-verification-cache.service'; import { AuthenticatedUser, JwtAccessPayload, @@ -16,6 +17,11 @@ import { * principal and rejects tokens whose session has been revoked via the * Redis-backed blacklist (e.g. after logout). The check fails open if Redis is * unreachable so a cache outage does not lock everyone out. + * + * Revocation checks go through the {@link TokenVerificationCacheService}: the + * blacklist answer is cached for a short TTL so authenticated requests avoid + * one Redis round trip each, and every revocation path invalidates the cached + * entry, so revocations are still observed immediately. */ @Injectable() export class JwtStrategy extends PassportStrategy(Strategy, 'jwt') { @@ -24,6 +30,7 @@ export class JwtStrategy extends PassportStrategy(Strategy, 'jwt') { constructor( config: ConfigService, private readonly tokenBlacklist: TokenBlacklistService, + private readonly verificationCache: TokenVerificationCacheService, ) { super({ jwtFromRequest: ExtractJwt.fromAuthHeaderAsBearerToken(), @@ -40,7 +47,13 @@ export class JwtStrategy extends PassportStrategy(Strategy, 'jwt') { if (payload.sessionId) { let revoked = false; try { - revoked = await this.tokenBlacklist.isAccessTokenRevoked(payload.sessionId); + // Cache-first: reads hit the short-TTL cache; misses fall through to + // the Redis blacklist and store the answer for the next requests. + const result = await this.verificationCache.resolveSessionRevocation( + payload.sessionId, + () => this.tokenBlacklist.isAccessTokenRevoked(payload.sessionId as string), + ); + revoked = result.revoked; } catch (error: unknown) { // Fail open on Redis outages rather than rejecting every request. this.logger.warn( diff --git a/src/modules/auth/services/token-blacklist.service.ts b/src/modules/auth/services/token-blacklist.service.ts index 3496b98e..b26d2f2b 100644 --- a/src/modules/auth/services/token-blacklist.service.ts +++ b/src/modules/auth/services/token-blacklist.service.ts @@ -1,5 +1,7 @@ -import { Injectable, Logger } from '@nestjs/common'; +import { Inject, Injectable, Logger } from '@nestjs/common'; import { Redis } from 'ioredis'; +import { TokenVerificationCacheService } from './token-verification-cache.service'; +import { REDIS_CLIENT } from '../../../common/locks/locks.constants'; /** * Redis-backed token revocation store. Issued JWTs remain valid until their @@ -10,12 +12,21 @@ import { Redis } from 'ioredis'; * * Access and refresh tokens are tracked under separate keys so their (much * different) lifetimes can be enforced independently. + * + * Every revocation also invalidates the session's entry in the token + * verification cache ({@link TokenVerificationCacheService}), so a cached + * "not revoked" answer can never outlive the revocation itself. */ @Injectable() export class TokenBlacklistService { private readonly logger = new Logger(TokenBlacklistService.name); - constructor(private readonly redis: Redis) {} + constructor( + // Injected by token so the global LocksModule provider is resolved + // regardless of the class-token import graph. + @Inject(REDIS_CLIENT) private readonly redis: Redis, + private readonly verificationCache: TokenVerificationCacheService, + ) {} /** Marks every token tied to a session as revoked. */ async revokeSession( @@ -31,6 +42,9 @@ export class TokenBlacklistService { } catch (error: unknown) { // Never let a Redis outage prevent logout from succeeding. this.logger.warn(`Failed to blacklist session ${sessionId}: ${(error as Error).message}`); + // A failed blacklist write may have left a stale "not revoked" cache + // entry behind (or skipped its invalidation); drop it explicitly. + await this.verificationCache.invalidateSessionRevocation(sessionId); } } @@ -39,7 +53,11 @@ export class TokenBlacklistService { if (ttlSeconds <= 0) { return; } + // Blacklist first, then drop the cached answer, so a concurrent request + // re-verifying against the source of truth can never observe the write + // before the invalidation and re-cache a stale "not revoked". await this.redis.set(this.accessKey(sessionId), '1', 'EX', ttlSeconds); + await this.verificationCache.invalidateSessionRevocation(sessionId); } /** Marks a refresh token as revoked for the remainder of its lifetime. */ @@ -48,6 +66,7 @@ export class TokenBlacklistService { return; } await this.redis.set(this.refreshKey(sessionId), '1', 'EX', ttlSeconds); + await this.verificationCache.invalidateSessionRevocation(sessionId); } /** True when the session's access token has been blacklisted. */ @@ -75,4 +94,4 @@ export class TokenBlacklistService { private refreshKey(sessionId: string): string { return `auth:blacklist:refresh:${sessionId}`; } -} \ No newline at end of file +} diff --git a/src/modules/auth/services/token-verification-cache.service.ts b/src/modules/auth/services/token-verification-cache.service.ts new file mode 100644 index 00000000..c9b4acd7 --- /dev/null +++ b/src/modules/auth/services/token-verification-cache.service.ts @@ -0,0 +1,128 @@ +import { Injectable } from '@nestjs/common'; +import { ConfigService } from '@nestjs/config'; +import { CacheService } from '../../../common/cache/cache.service'; + +/** + * Result of verifying a session against the revocation store. `revoked` is the + * authoritative answer (true when the session has been blacklisted by logout + * or credential rotation); `verifiedAt` records when the check happened so the + * cache can bound how long the answer is trusted. + */ +export interface SessionRevocationResult { + revoked: boolean; + verifiedAt: number; +} + +/** + * Cache in front of the token revocation store ({@link TokenBlacklistService} + * — itself Redis-backed). Every authenticated request asks "is this session + * still revoked?", which is one Redis round trip per request; this service + * answers it from a short-lived cache entry instead, cutting Redis load by + * roughly the number of requests a session makes per TTL window. + * + * Revocation stays reliable because every cache entry is bounded by + * `cacheTtlSeconds` (default 30s, `TOKEN_CACHE_TTL`), which is deliberately + * shorter than the shortest credential lifetime (the 15-minute access token), + * and because every revocation / logout / refresh-rotation path calls the + * {@link invalidate} hooks, which drop the cached answer immediately. A cached + * "not revoked" answer therefore survives at most one TTL window — the same + * bounded staleness the project already accepts for the Redis blacklist + * itself — and explicit revocations take effect at once. + * + * Key layout: `auth:token-verification:` — one entry per session, + * invalidated in O(1) on logout without any scan. + */ +@Injectable() +export class TokenVerificationCacheService { + private static readonly NAMESPACE = 'auth:token-verification'; + private readonly ttlSeconds: number; + + constructor( + private readonly cache: CacheService, + config: ConfigService, + ) { + // Optional tuning knob; follows the BalanceCacheService pattern of an + // unvalidated, defaulted variable read straight from ConfigService. The + // 30s default stays well below the shortest token lifetime (access TTL). + const configured = config.get('TOKEN_CACHE_TTL', 30); + this.ttlSeconds = typeof configured === 'number' && configured > 0 ? configured : 30; + } + + /** Returns the cached revocation answer for a session, or null on miss. */ + async getSessionRevocation(sessionId: string): Promise { + if (!sessionId) { + return null; + } + const hit = await this.cache.get(this.namespace(), this.key(sessionId)); + if (!hit || typeof hit.revoked !== 'boolean') { + return null; + } + return hit; + } + + /** Caches a revocation answer for one TTL window. */ + async setSessionRevocation(sessionId: string, result: SessionRevocationResult): Promise { + if (!sessionId) { + return; + } + await this.cache.set(this.namespace(), this.key(sessionId), result, this.ttlSeconds); + } + + /** + * Cache-read / source-of-truth / cache-write helper. `resolve` must perform + * the authoritative check (Redis blacklist); its answer is cached for the + * next requests within the TTL window. + */ + async resolveSessionRevocation( + sessionId: string, + resolve: () => Promise, + ): Promise { + const cached = await this.getSessionRevocation(sessionId); + if (cached) { + return cached; + } + const result: SessionRevocationResult = { revoked: await resolve(), verifiedAt: Date.now() }; + await this.setSessionRevocation(sessionId, result); + return result; + } + + /** + * Invalidation hook: the session's access token has been revoked (logout or + * credential rotation). Drops the cached answer so the next verification + * hits the source of truth and observes the revocation immediately. + */ + async invalidateSessionRevocation(sessionId: string): Promise { + await this.cache.del(this.namespace(), this.key(sessionId)); + } + + /** + * Invalidation hook for refresh flows: a rotated refresh token means the old + * session is dead and a new one was born, but the old session's access token + * is still in flight, so its cached "not revoked" answer must be dropped. + */ + async invalidateOnRefreshRotation(oldSessionId: string): Promise { + await this.invalidateSessionRevocation(oldSessionId); + } + + /** + * Diagnostics hook: clears every cached verification entry. Intended for + * tests and emergency cache flushes, not for the request path (the SCAN it + * performs is not O(1)). + */ + async clearAll(): Promise { + await this.cache.delByPrefix(this.namespace(), ''); + } + + /** Configured TTL, exposed for tests and configuration assertions. */ + get cacheTtlSeconds(): number { + return this.ttlSeconds; + } + + private namespace(): string { + return TokenVerificationCacheService.NAMESPACE; + } + + private key(sessionId: string): string { + return sessionId; + } +} diff --git a/src/modules/auth/tests/api-key-auth.integration.spec.ts b/src/modules/auth/tests/api-key-auth.integration.spec.ts index f654601c..05021cc0 100644 --- a/src/modules/auth/tests/api-key-auth.integration.spec.ts +++ b/src/modules/auth/tests/api-key-auth.integration.spec.ts @@ -15,6 +15,7 @@ import { sha256 } from '../../../utils/crypto.util'; import { ConfigService } from '@nestjs/config'; import { JwtStrategy } from '../jwt.strategy'; import { TokenBlacklistService } from '../services/token-blacklist.service'; +import { TokenVerificationCacheService } from '../services/token-verification-cache.service'; @Controller('test-resource') @UseGuards(JwtAuthGuard, ScopesGuard) @@ -52,12 +53,16 @@ describe('API Key Authentication with Scoped Permissions (Integration)', () => { const mockBlacklist = { isAccessTokenRevoked: vi.fn().mockResolvedValue(false), }; + const mockVerificationCache = { + resolveSessionRevocation: vi.fn().mockResolvedValue({ revoked: false, verifiedAt: Date.now() }), + }; app = await Test.createTestingModule({ imports: [PassportModule.register({ defaultStrategy: 'jwt' })], controllers: [TestProtectedController], providers: [ { provide: ConfigService, useValue: mockConfig }, + { provide: TokenVerificationCacheService, useValue: mockVerificationCache }, { provide: TokenBlacklistService, useValue: mockBlacklist }, JwtStrategy, ApiKeyStrategy, diff --git a/src/modules/auth/tests/api-key.strategy.spec.ts b/src/modules/auth/tests/api-key.strategy.spec.ts index 3de0cd5a..a7f954f9 100644 --- a/src/modules/auth/tests/api-key.strategy.spec.ts +++ b/src/modules/auth/tests/api-key.strategy.spec.ts @@ -40,6 +40,7 @@ describe('ApiKeyStrategy', () => { expect(result).toEqual({ id: 'key-123', keyId: 'key-123', + apiKeyId: 'key-123', organizationId: 'org-456', createdById: 'user-789', name: 'Test Key', diff --git a/src/modules/auth/tests/jwt.strategy.spec.ts b/src/modules/auth/tests/jwt.strategy.spec.ts index e63e47e8..d9c0b38b 100644 --- a/src/modules/auth/tests/jwt.strategy.spec.ts +++ b/src/modules/auth/tests/jwt.strategy.spec.ts @@ -1,6 +1,7 @@ import { describe, it, expect, beforeEach, vi } from 'vitest'; import { JwtStrategy } from '../jwt.strategy'; import { TokenBlacklistService } from '../services/token-blacklist.service'; +import { TokenVerificationCacheService } from '../services/token-verification-cache.service'; import { AuthConfig } from '../../../config/auth.config'; import { JwtAccessPayload } from '../../../common/interfaces/authenticated-user.interface'; @@ -24,31 +25,44 @@ const payload: JwtAccessPayload = { sessionId: 'session-123', }; -function makeStrategy(blacklist: Partial) { +function makeStrategy( + blacklist: Partial, + verificationCache: Partial, +) { return new JwtStrategy( mockConfig as never, blacklist as TokenBlacklistService, + verificationCache as TokenVerificationCacheService, ); } describe('JwtStrategy', () => { let tokenBlacklist: { isAccessTokenRevoked: ReturnType }; + let verificationCache: { + resolveSessionRevocation: ReturnType; + }; beforeEach(() => { vi.clearAllMocks(); tokenBlacklist = { isAccessTokenRevoked: vi.fn().mockResolvedValue(false) }; + verificationCache = { + resolveSessionRevocation: vi.fn().mockImplementation( + (_sessionId: string, resolve: () => Promise) => + resolve().then((revoked) => ({ revoked, verifiedAt: Date.now() })), + ), + }; }); it('rejects a valid-signature token whose session is blacklisted', async () => { tokenBlacklist.isAccessTokenRevoked.mockResolvedValue(true); - const strategy = makeStrategy(tokenBlacklist); + const strategy = makeStrategy(tokenBlacklist, verificationCache); await expect(strategy.validate(payload)).rejects.toThrow('Session has been revoked'); expect(tokenBlacklist.isAccessTokenRevoked).toHaveBeenCalledWith('session-123'); }); it('grants access to a token whose session is not blacklisted', async () => { - const strategy = makeStrategy(tokenBlacklist); + const strategy = makeStrategy(tokenBlacklist, verificationCache); await expect(strategy.validate(payload)).resolves.toMatchObject({ id: 'user-1', @@ -56,9 +70,24 @@ describe('JwtStrategy', () => { }); }); + it('serves repeated verifications from the cache without re-querying the blacklist', async () => { + // The cache answers every check: the blacklist is never consulted. + verificationCache.resolveSessionRevocation.mockResolvedValue({ + revoked: false, + verifiedAt: Date.now(), + }); + const strategy = makeStrategy(tokenBlacklist, verificationCache); + + for (let i = 0; i < 3; i++) { + await expect(strategy.validate(payload)).resolves.toMatchObject({ id: 'user-1' }); + } + + expect(tokenBlacklist.isAccessTokenRevoked).not.toHaveBeenCalled(); + }); + it('fails open and grants access when Redis is unreachable', async () => { - tokenBlacklist.isAccessTokenRevoked.mockRejectedValue(new Error('Redis down')); - const strategy = makeStrategy(tokenBlacklist); + verificationCache.resolveSessionRevocation.mockRejectedValue(new Error('Redis down')); + const strategy = makeStrategy(tokenBlacklist, verificationCache); await expect(strategy.validate(payload)).resolves.toMatchObject({ id: 'user-1', @@ -66,10 +95,23 @@ describe('JwtStrategy', () => { }); it('rejects a malformed token missing the subject', async () => { - const strategy = makeStrategy(tokenBlacklist); + const strategy = makeStrategy(tokenBlacklist, verificationCache); await expect( strategy.validate({ organizationId: 'org-1', email: 'a@b.c', role: 'OWNER' } as never), ).rejects.toThrow('Malformed access token'); }); -}); \ No newline at end of file + + it('skips the revocation check for tokens without a session id', async () => { + const strategy = makeStrategy(tokenBlacklist, verificationCache); + const payloadWithoutSession: JwtAccessPayload = { + sub: 'user-1', + organizationId: 'org-1', + email: 'ada@acme.com', + role: 'OWNER', + }; + + await expect(strategy.validate(payloadWithoutSession)).resolves.toMatchObject({ id: 'user-1' }); + expect(verificationCache.resolveSessionRevocation).not.toHaveBeenCalled(); + }); +}); diff --git a/src/modules/auth/tests/token-blacklist.service.spec.ts b/src/modules/auth/tests/token-blacklist.service.spec.ts index aebdb83b..7e5f1bfd 100644 --- a/src/modules/auth/tests/token-blacklist.service.spec.ts +++ b/src/modules/auth/tests/token-blacklist.service.spec.ts @@ -1,8 +1,12 @@ import { describe, it, expect, beforeEach, vi } from 'vitest'; import { TokenBlacklistService } from '../services/token-blacklist.service'; +import { TokenVerificationCacheService } from '../services/token-verification-cache.service'; describe('TokenBlacklistService', () => { let service: TokenBlacklistService; + let verificationCache: { + invalidateSessionRevocation: ReturnType; + }; let redis: { set: ReturnType; exists: ReturnType; @@ -13,7 +17,13 @@ describe('TokenBlacklistService', () => { set: vi.fn().mockResolvedValue('OK'), exists: vi.fn().mockResolvedValue(0), }; - service = new TokenBlacklistService(redis as never); + verificationCache = { + invalidateSessionRevocation: vi.fn().mockResolvedValue(undefined), + }; + service = new TokenBlacklistService( + redis as never, + verificationCache as unknown as TokenVerificationCacheService, + ); }); it('revokes access and refresh tokens with distinct TTLs', async () => { @@ -65,4 +75,23 @@ describe('TokenBlacklistService', () => { service.revokeSession('session-1', 900, 1209600), ).resolves.toBeUndefined(); }); + + it('invalidates the cached verification answer on every revocation path', async () => { + await service.revokeSession('session-1', 900, 1209600); + await service.revokeAccessToken('session-2', 900); + await service.revokeRefreshToken('session-3', 1209600); + + // revokeSession invalidates once per token kind (access + refresh). + expect(verificationCache.invalidateSessionRevocation.mock.calls.filter(([id]) => id === 'session-1')).toHaveLength(2); + expect(verificationCache.invalidateSessionRevocation).toHaveBeenCalledWith('session-2'); + expect(verificationCache.invalidateSessionRevocation).toHaveBeenCalledWith('session-3'); + }); + + it('still invalidates the cache when blacklisting fails on a Redis outage', async () => { + redis.set.mockRejectedValue(new Error('Redis connection failed')); + + await service.revokeSession('session-1', 900, 1209600); + + expect(verificationCache.invalidateSessionRevocation).toHaveBeenCalledWith('session-1'); + }); }); \ No newline at end of file diff --git a/src/modules/auth/tests/token-verification-cache.integration.spec.ts b/src/modules/auth/tests/token-verification-cache.integration.spec.ts new file mode 100644 index 00000000..93c67cd7 --- /dev/null +++ b/src/modules/auth/tests/token-verification-cache.integration.spec.ts @@ -0,0 +1,180 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { Controller, Get, INestApplication, UseGuards } from '@nestjs/common'; +import { Test } from '@nestjs/testing'; +import { ConfigService } from '@nestjs/config'; +import { PassportModule } from '@nestjs/passport'; +import { JwtModule, JwtService } from '@nestjs/jwt'; +import { JwtAuthGuard } from '../../../common/guards/jwt-auth.guard'; +import { CurrentUser } from '../../../common/decorators/current-user.decorator'; +import { AuthenticatedUser } from '../../../common/interfaces/authenticated-user.interface'; +import { JwtStrategy } from '../jwt.strategy'; +import { TokenBlacklistService } from '../services/token-blacklist.service'; +import { TokenVerificationCacheService } from '../services/token-verification-cache.service'; +import { CacheService } from '../../../common/cache/cache.service'; +import { REDIS_CLIENT } from '../../../common/locks/locks.constants'; + +/** + * Exercises the authenticated-request path end to end: a signed JWT flows + * through the global `JwtAuthGuard` into `JwtStrategy`, which consults the + * token verification cache in front of the Redis blacklist. The Redis client + * is a stand-in implementing only the primitives the stack touches + * (`get`/`set`/`del` for cache and blacklist keys), letting the test count + * exactly how many revocation lookups happen per request. + */ + +const ACCESS_SECRET = 'test-access-secret-at-least-32-chars'; + +@Controller('protected') +@UseGuards(JwtAuthGuard) +class ProtectedController { + @Get('me') + me(@CurrentUser() user: AuthenticatedUser) { + return { id: user.id, organizationId: user.organizationId }; + } +} + +describe('Token verification caching (integration)', () => { + let app: INestApplication; + let jwt: JwtService; + let blacklist: { isAccessTokenRevoked: ReturnType }; + let redisStore: Map; + let blacklistLookups: number; + + /** Minimal Redis stand-in: real GET/SET/DEL semantics over a Map. */ + const fakeRedis = { + get: vi.fn(async (key: string) => redisStore.get(key) ?? null), + set: vi.fn(async (key: string, value: string) => { + redisStore.set(key, value); + return 'OK'; + }), + del: vi.fn(async (...keys: string[]) => { + let removed = 0; + for (const key of keys) { + if (redisStore.delete(key)) removed++; + } + return removed; + }), + exists: vi.fn(async (key: string) => (redisStore.has(key) ? 1 : 0)), + }; + + const signAccessToken = async (sessionId: string) => + jwt.signAsync( + { sub: 'user-1', organizationId: 'org-1', email: 'ada@acme.com', role: 'OWNER', sessionId }, + { secret: ACCESS_SECRET, expiresIn: 900 }, + ); + + /** Revokes a session the same way the logout path does (blacklist + cache invalidation). */ + const revokeSession = async (sessionId: string) => { + redisStore.set(`auth:blacklist:access:${sessionId}`, '1'); + await app.get(TokenVerificationCacheService).invalidateSessionRevocation(sessionId); + }; + + beforeEach(async () => { + redisStore = new Map(); + blacklistLookups = 0; + + blacklist = { + isAccessTokenRevoked: vi.fn().mockImplementation(async (sessionId: string) => { + blacklistLookups += 1; + return redisStore.has(`auth:blacklist:access:${sessionId}`); + }), + }; + + const moduleRef = await Test.createTestingModule({ + imports: [PassportModule.register({ defaultStrategy: 'jwt' }), JwtModule.register({})], + controllers: [ProtectedController], + providers: [ + { + provide: ConfigService, + useValue: { + getOrThrow: () => ({ accessSecret: ACCESS_SECRET }), + get: () => undefined, + }, + }, + { provide: REDIS_CLIENT, useValue: fakeRedis }, + CacheService, + TokenVerificationCacheService, + { provide: TokenBlacklistService, useValue: blacklist }, + JwtStrategy, + JwtAuthGuard, + ], + }).compile(); + + app = moduleRef.createNestApplication({ logger: false }); + await app.init(); + jwt = app.get(JwtService); + }); + + const authenticate = async (token: string): Promise => { + // canActivate is invoked through the guard pipeline exactly as production + // wiring does; we drive it directly to observe pass/fail without HTTP. + const request = { headers: { authorization: `Bearer ${token}` } }; + const guard = app.get(JwtAuthGuard); + const passportFlow = guard.canActivate({ + switchToHttp: () => ({ + getRequest: () => request, + getResponse: () => ({}), + }), + getHandler: () => ProtectedController.prototype.me, + getClass: () => ProtectedController, + } as never); + try { + const result = await passportFlow; + return result === true ? 200 : 401; + } catch { + return 401; + } + }; + + it('caches the first verification and serves later requests without blacklist lookups', async () => { + const token = await signAccessToken('session-cache'); + expect(await authenticate(token)).toBe(200); + expect(blacklistLookups).toBe(1); + + // Subsequent requests hit the cache: no additional blacklist queries. + expect(await authenticate(token)).toBe(200); + expect(await authenticate(token)).toBe(200); + expect(blacklistLookups).toBe(1); + }); + + it('observes a revocation immediately because logout invalidates the cache', async () => { + const token = await signAccessToken('session-revoked'); + expect(await authenticate(token)).toBe(200); + expect(blacklistLookups).toBe(1); + + // Logout path: blacklist write + cache invalidation, as wired in + // AuthService via TokenBlacklistService. + await revokeSession('session-revoked'); + expect(await authenticate(token)).toBe(401); + }); + + it('treats a cache miss by consulting the blacklist and re-populating the cache', async () => { + const token = await signAccessToken('session-miss'); + expect(await authenticate(token)).toBe(200); + expect(blacklistLookups).toBe(1); + + // Invalidate (cache miss on the next request), then verify the answer is + // re-derived from the source of truth and cached again. + await app.get(TokenVerificationCacheService).invalidateSessionRevocation('session-miss'); + redisStore.set('auth:blacklist:access:session-miss', '1'); + expect(await authenticate(token)).toBe(401); + expect(blacklistLookups).toBe(2); + + // The fresh "revoked" answer is now cached — no further lookups. + redisStore.delete('auth:blacklist:access:session-miss'); + expect(await authenticate(token)).toBe(401); + expect(blacklistLookups).toBe(2); + }); + + it('keeps independent sessions isolated (no cross-session cache leakage)', async () => { + const tokenA = await signAccessToken('session-A'); + + expect(await authenticate(tokenA)).toBe(200); + expect(blacklistLookups).toBe(1); + + // Revoking session B must not affect session A's cached answer. + await revokeSession('session-B'); + expect(await authenticate(tokenA)).toBe(200); + expect(blacklistLookups).toBe(1); + }); +}); diff --git a/src/modules/auth/tests/token-verification-cache.service.spec.ts b/src/modules/auth/tests/token-verification-cache.service.spec.ts new file mode 100644 index 00000000..ad21ae94 --- /dev/null +++ b/src/modules/auth/tests/token-verification-cache.service.spec.ts @@ -0,0 +1,142 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { TokenVerificationCacheService } from '../services/token-verification-cache.service'; +import { CacheService } from '../../../common/cache/cache.service'; + +describe('TokenVerificationCacheService', () => { + let service: TokenVerificationCacheService; + let cache: { + get: ReturnType; + set: ReturnType; + del: ReturnType; + delByPrefix: ReturnType; + }; + + beforeEach(() => { + cache = { + get: vi.fn().mockResolvedValue(null), + set: vi.fn().mockResolvedValue(undefined), + del: vi.fn().mockResolvedValue(undefined), + delByPrefix: vi.fn().mockResolvedValue(undefined), + }; + service = new TokenVerificationCacheService(cache as unknown as CacheService, { + get: vi.fn().mockReturnValue(30), + } as never); + }); + + describe('getSessionRevocation / setSessionRevocation', () => { + it('returns null when nothing is cached', async () => { + await expect(service.getSessionRevocation('session-1')).resolves.toBeNull(); + expect(cache.get).toHaveBeenCalledWith('auth:token-verification', 'session-1'); + }); + + it('returns the cached answer on hit', async () => { + cache.get.mockResolvedValue({ revoked: true, verifiedAt: 123 }); + + await expect(service.getSessionRevocation('session-1')).resolves.toEqual({ + revoked: true, + verifiedAt: 123, + }); + }); + + it('stores the answer under the session key with the configured TTL', async () => { + await service.setSessionRevocation('session-1', { revoked: false, verifiedAt: 42 }); + + expect(cache.set).toHaveBeenCalledWith( + 'auth:token-verification', + 'session-1', + { revoked: false, verifiedAt: 42 }, + 30, + ); + }); + + it('ignores empty session ids without touching the cache', async () => { + await expect(service.getSessionRevocation('')).resolves.toBeNull(); + await service.setSessionRevocation('', { revoked: false, verifiedAt: 1 }); + expect(cache.get).not.toHaveBeenCalledWith('auth:token-verification', ''); + expect(cache.set).not.toHaveBeenCalled(); + }); + + it('treats malformed cache payloads as a miss', async () => { + cache.get.mockResolvedValue({ garbage: true }); + + await expect(service.getSessionRevocation('session-1')).resolves.toBeNull(); + }); + }); + + describe('resolveSessionRevocation', () => { + it('answers from the cache without invoking the resolver (cache hit)', async () => { + cache.get.mockResolvedValue({ revoked: false, verifiedAt: Date.now() }); + const resolve = vi.fn().mockResolvedValue(true); + + const result = await service.resolveSessionRevocation('session-1', resolve); + + expect(result).toMatchObject({ revoked: false }); + expect(resolve).not.toHaveBeenCalled(); + }); + + it('falls back to the resolver on a miss and caches the fresh answer', async () => { + const resolve = vi.fn().mockResolvedValue(false); + + const result = await service.resolveSessionRevocation('session-1', resolve); + + expect(result.revoked).toBe(false); + expect(resolve).toHaveBeenCalledTimes(1); + expect(cache.set).toHaveBeenCalledWith( + 'auth:token-verification', + 'session-1', + expect.objectContaining({ revoked: false }), + 30, + ); + }); + + it('propagates resolver failures so callers keep their fail-open behavior', async () => { + const resolve = vi.fn().mockRejectedValue(new Error('Redis down')); + + await expect( + service.resolveSessionRevocation('session-1', resolve), + ).rejects.toThrow('Redis down'); + expect(cache.set).not.toHaveBeenCalled(); + }); + }); + + describe('invalidation hooks', () => { + it('invalidateSessionRevocation drops the cached entry', async () => { + await service.invalidateSessionRevocation('session-1'); + + expect(cache.del).toHaveBeenCalledWith('auth:token-verification', 'session-1'); + }); + + it('invalidateOnRefreshRotation invalidates the rotated (old) session', async () => { + await service.invalidateOnRefreshRotation('old-session'); + + expect(cache.del).toHaveBeenCalledWith('auth:token-verification', 'old-session'); + }); + + it('a session revoked after being cached is observed as revoked again', async () => { + // 1. First verification caches "not revoked". + const resolve = vi.fn().mockResolvedValueOnce(false); + await service.resolveSessionRevocation('session-1', resolve); + // 2. Logout invalidates the cached answer... + await service.invalidateSessionRevocation('session-1'); + // 3. ...so the next verification consults the source of truth again. + cache.get.mockResolvedValue(null); + const resolveAfterRevocation = vi.fn().mockResolvedValue(true); + const result = await service.resolveSessionRevocation('session-1', resolveAfterRevocation); + + expect(result.revoked).toBe(true); + expect(resolveAfterRevocation).toHaveBeenCalledTimes(1); + }); + }); + + it('clearAll drops every cached verification entry', async () => { + await service.clearAll(); + expect(cache.delByPrefix).toHaveBeenCalledWith('auth:token-verification', ''); + }); + + it('falls back to the 30s default TTL when TOKEN_CACHE_TTL is unset', () => { + const fallback = new TokenVerificationCacheService(cache as unknown as CacheService, { + get: vi.fn().mockReturnValue(undefined), + } as never); + expect(fallback.cacheTtlSeconds).toBe(30); + }); +}); diff --git a/src/modules/budgets/budget.controller.ts b/src/modules/budgets/budget.controller.ts index 59f00f46..591452dc 100644 --- a/src/modules/budgets/budget.controller.ts +++ b/src/modules/budgets/budget.controller.ts @@ -33,9 +33,11 @@ import { import { CurrentUser } from '../../common/decorators/current-user.decorator'; import { Roles } from '../../common/decorators/roles.decorator'; import { AuditAction } from '../../common/decorators/audit-action.decorator'; +import { AuditLog } from '../../common/decorators/audit-log.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { AuthenticatedUser } from '../../common/interfaces/authenticated-user.interface'; import { PaginationQuery, paginationQuerySchema } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; import { ApiEnvelope } from '../../common/decorators/api-envelope.decorator'; @ApiTags('budgets') @@ -51,8 +53,7 @@ export class BudgetController { 'Returns a paginated list of budgets for the current organization. ' + 'Supports filtering by period and status.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiQuery({ name: 'period', required: false, enum: ['DAILY', 'WEEKLY', 'MONTHLY', 'QUARTERLY', 'YEARLY'], description: 'Filter by budget period' }) @ApiQuery({ name: 'enabled', required: false, type: Boolean, description: 'Filter by enabled status' }) @ApiEnvelope(CreateBudgetDto as never, { isArray: true }) @@ -68,6 +69,7 @@ export class BudgetController { @Post() @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE) @AuditAction('BUDGET_CREATED') + @AuditLog({ action: 'BUDGET_CREATED', entity: 'Budget' }) @ApiOperation({ summary: 'Create a budget', description: @@ -105,6 +107,7 @@ export class BudgetController { @UseBudgetLock() @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE) @AuditAction('BUDGET_UPDATED') + @AuditLog({ action: 'BUDGET_UPDATED', entity: 'Budget' }) @ApiOperation({ summary: 'Update a budget', description: 'Partial update of budget fields (name, limit, period, rollover, enabled).', @@ -128,6 +131,7 @@ export class BudgetController { @UseBudgetLock() @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE) @AuditAction('BUDGET_ALLOCATED') + @AuditLog({ action: 'BUDGET_ADJUSTED', entity: 'Budget' }) @ApiOperation({ summary: 'Allocate funds from the parent budget to this child', description: @@ -153,6 +157,7 @@ export class BudgetController { @UseBudgetLock() @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE) @AuditAction('BUDGET_DELETED') + @AuditLog({ action: 'BUDGET_DELETED', entity: 'Budget' }) @ApiOperation({ summary: 'Delete (soft) a budget', description: diff --git a/src/modules/budgets/budget.service.ts b/src/modules/budgets/budget.service.ts index ab5ec11c..81e5c0ec 100644 --- a/src/modules/budgets/budget.service.ts +++ b/src/modules/budgets/budget.service.ts @@ -72,7 +72,7 @@ export class BudgetService { } const pagination = toPrismaPagination(query, SORTABLE); const { items, total } = await this.repository.findManyAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); + return new Paginated(items, buildPaginationMeta(total, query)); } async getOrThrow(organizationId: string, id: string): Promise { diff --git a/src/modules/developer/api-key.controller.ts b/src/modules/developer/api-key.controller.ts index aaebe557..ef27c9e9 100644 --- a/src/modules/developer/api-key.controller.ts +++ b/src/modules/developer/api-key.controller.ts @@ -6,7 +6,6 @@ import { ApiResponse, ApiParam, ApiBody, - ApiQuery, } from '@nestjs/swagger'; import { UserRole } from '@prisma/client'; import { ApiKeyService } from './api-key.service'; @@ -14,9 +13,11 @@ import { createApiKeySchema, CreateApiKeyInput, CreateApiKeyDto, ApiKeyCreatedDt import { CurrentUser } from '../../common/decorators/current-user.decorator'; import { Roles } from '../../common/decorators/roles.decorator'; import { AuditAction } from '../../common/decorators/audit-action.decorator'; +import { AuditLog } from '../../common/decorators/audit-log.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { AuthenticatedUser } from '../../common/interfaces/authenticated-user.interface'; import { PaginationQuery, paginationQuerySchema } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; @ApiTags('developer') @ApiBearerAuth('access-token') @@ -32,8 +33,7 @@ export class ApiKeyController { 'Returns a paginated list of API keys for the current organization. ' + 'Full secrets are never included — only prefix and metadata.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiResponse({ status: 200, description: 'Paginated list of API keys' }) @ApiResponse({ status: 401, description: 'Not authenticated' }) @ApiResponse({ status: 403, description: 'Insufficient permissions' }) @@ -47,6 +47,7 @@ export class ApiKeyController { @Post() @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.DEVELOPER) @AuditAction('AGENT_KEY_ROTATED') + @AuditLog({ action: 'AGENT_KEY_CREATED', entity: 'ApiKey' }) @ApiOperation({ summary: 'Create an API key', description: @@ -72,6 +73,7 @@ export class ApiKeyController { @Delete(':id') @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.DEVELOPER) @AuditAction('AGENT_KEY_REVOKED') + @AuditLog({ action: 'AGENT_KEY_REVOKED', entity: 'ApiKey' }) @ApiOperation({ summary: 'Revoke an API key', description: diff --git a/src/modules/developer/api-key.repository.ts b/src/modules/developer/api-key.repository.ts index 20340b8d..56366eb3 100644 --- a/src/modules/developer/api-key.repository.ts +++ b/src/modules/developer/api-key.repository.ts @@ -4,8 +4,9 @@ import { PrismaService } from '../../database/prisma.service'; import { PrismaPagination } from '../../common/helpers/pagination'; /** - * Persistence for ApiKey rows. Only the SHA-256 `hashedKey` is ever stored — the - * raw key exists solely in the create response. + * Persistence for ApiKey rows. Only the Argon2id `hashedKey` is ever stored — the + * raw key exists solely in the create response. Legacy SHA-256 hashes are supported + * for backward compatibility during migration. */ @Injectable() export class ApiKeyRepository { @@ -49,6 +50,11 @@ export class ApiKeyRepository { return this.prisma.apiKey.findUnique({ where: { hashedKey } }); } + /** Resolves API keys by their prefix (used for key verification with multiple hash algorithms). */ + findByPrefix(prefix: string): Promise { + return this.prisma.apiKey.findMany({ where: { prefix } }); + } + revoke(id: string): Promise { return this.prisma.apiKey.update({ where: { id }, data: { revokedAt: new Date() } }); } diff --git a/src/modules/developer/api-key.service.ts b/src/modules/developer/api-key.service.ts index 9817f195..390362ec 100644 --- a/src/modules/developer/api-key.service.ts +++ b/src/modules/developer/api-key.service.ts @@ -9,21 +9,22 @@ import { toPrismaPagination, } from '../../common/helpers/pagination'; import { Paginated } from '../../common/interfaces/api-response.interface'; -import { generateApiKey, sha256 } from '../../utils/crypto.util'; +import { generateApiKey, verifyArgon2, sha256 } from '../../utils/crypto.util'; const SORTABLE = ['createdAt', 'name', 'lastUsedAt']; /** * Issues and manages programmatic API keys. The raw secret is generated, shown - * to the caller exactly once, and only its SHA-256 hash is persisted. Keys can - * never be recovered — only regenerated. + * to the caller exactly once, and only its Argon2id hash is persisted. Keys can + * never be recovered — only regenerated. Legacy SHA-256 hashes are supported for + * backward compatibility during migration. */ @Injectable() export class ApiKeyService { constructor(private readonly repository: ApiKeyRepository) {} async create(organizationId: string, actorId: string, input: CreateApiKeyInput) { - const { raw, prefix, hashedKey } = generateApiKey('live'); + const { raw, prefix, hashedKey } = await generateApiKey('live'); const expiresAt = input.expiresInDays ? new Date(Date.now() + input.expiresInDays * 86_400_000) : null; @@ -57,7 +58,7 @@ export class ApiKeyService { } const pagination = toPrismaPagination(query, SORTABLE); const { items, total } = await this.repository.findManyAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); + return new Paginated(items, buildPaginationMeta(total, query)); } async revoke(organizationId: string, id: string) { @@ -73,25 +74,51 @@ export class ApiKeyService { } /** - * Verifies a presented raw key: matches by hash, checks it is neither revoked - * nor expired, and updates lastUsedAt. Returns the owning key or null. + * Verifies a presented raw key: matches by Argon2 hash (with SHA-256 fallback for legacy keys), + * checks it is neither revoked nor expired, and updates lastUsedAt. Returns the owning key or null. */ async verify(rawKey: string) { if (!rawKey || typeof rawKey !== 'string' || rawKey.trim().length === 0) { return null; } - const key = await this.repository.findByHash(sha256(rawKey.trim())); - if (!key || key.revokedAt) { - return null; - } - if (key.expiresAt && key.expiresAt.getTime() < Date.now()) { - return null; - } - try { - await this.repository.touchLastUsed(key.id); - } catch { - // Gracefully continue even if updating lastUsedAt encounters an error + + const trimmedKey = rawKey.trim(); + + // First try to find by the stored hash (we need to retrieve the key to verify) + // Since we can't hash the input without knowing which algorithm was used, + // we'll try to find by prefix first, then verify the hash + const keys = await this.repository.findByPrefix(trimmedKey.slice(0, 14)); + + for (const key of keys) { + if (key.revokedAt) { + continue; + } + if (key.expiresAt && key.expiresAt.getTime() < Date.now()) { + continue; + } + + // Try Argon2 verification first (new keys) + const isValidArgon2 = await verifyArgon2(key.hashedKey, trimmedKey); + if (isValidArgon2) { + try { + await this.repository.touchLastUsed(key.id); + } catch { + // Gracefully continue even if updating lastUsedAt encounters an error + } + return key; + } + + // Fallback to SHA-256 for legacy keys (backward compatibility) + if (key.hashedKey === sha256(trimmedKey)) { + try { + await this.repository.touchLastUsed(key.id); + } catch { + // Gracefully continue even if updating lastUsedAt encounters an error + } + return key; + } } - return key; + + return null; } } diff --git a/src/modules/developer/tests/api-key.service.spec.ts b/src/modules/developer/tests/api-key.service.spec.ts index e8930879..e0688044 100644 --- a/src/modules/developer/tests/api-key.service.spec.ts +++ b/src/modules/developer/tests/api-key.service.spec.ts @@ -2,7 +2,7 @@ import { beforeEach, describe, expect, it, vi } from 'vitest'; import { ApiKeyService } from '../api-key.service'; import { ApiKeyRepository } from '../api-key.repository'; import { ConflictException, NotFoundException } from '../../../common/exceptions/domain.exception'; -import { sha256 } from '../../../utils/crypto.util'; +import { hashWithArgon2, sha256 } from '../../../utils/crypto.util'; describe('ApiKeyService', () => { let service: ApiKeyService; @@ -11,6 +11,7 @@ describe('ApiKeyService', () => { findManyAndCount: ReturnType; findById: ReturnType; findByHash: ReturnType; + findByPrefix: ReturnType; revoke: ReturnType; touchLastUsed: ReturnType; }; @@ -24,6 +25,7 @@ describe('ApiKeyService', () => { findManyAndCount: vi.fn(), findById: vi.fn(), findByHash: vi.fn(), + findByPrefix: vi.fn(), revoke: vi.fn(), touchLastUsed: vi.fn(), }; @@ -31,7 +33,7 @@ describe('ApiKeyService', () => { }); describe('create', () => { - it('creates an API key, hashes it with SHA-256, and returns raw key only once', async () => { + it('creates an API key, hashes it with Argon2id, and returns raw key only once', async () => { repository.create.mockImplementation((data) => Promise.resolve({ id: 'key-123', @@ -58,13 +60,14 @@ describe('ApiKeyService', () => { expect(result.prefix).toBe(result.key.slice(0, 14)); expect(result.expiresAt).toBeInstanceOf(Date); + // Verify the stored hash is Argon2id format expect(repository.create).toHaveBeenCalledWith( expect.objectContaining({ organizationId: orgId, createdById: userId, name: 'Agent Key', prefix: result.prefix, - hashedKey: sha256(result.key), + hashedKey: expect.stringMatching(/\$argon2id\$/), permissions: ['transactions:write', 'wallets:read'], allowedIps: ['192.168.1.1'], }), @@ -110,6 +113,7 @@ describe('ApiKeyService', () => { repository.findManyAndCount.mockResolvedValue({ items: mockItems, total: 1 }); const result = await service.list(orgId, { + offset: 0, page: 1, limit: 20, sort: 'createdAt', @@ -128,6 +132,7 @@ describe('ApiKeyService', () => { repository.findManyAndCount.mockResolvedValue({ items: [], total: 0 }); await service.list(orgId, { + offset: 0, page: 1, limit: 20, sort: 'createdAt', @@ -180,38 +185,62 @@ describe('ApiKeyService', () => { }); describe('verify', () => { - const rawSecret = 'ak_live_abcdef1234567890abcdef1234567890abcdef12'; - const hash = sha256(rawSecret); - - it('returns key and touches lastUsedAt when key is valid', async () => { + it('verifies key with Argon2id hash', async () => { + const rawSecret = 'ak_live_abcdef1234567890abcdef1234567890abcdef12'; + const argonHash = await hashWithArgon2(rawSecret); + const mockKey = { id: 'key-1', name: 'Agent Key', - hashedKey: hash, + hashedKey: argonHash, permissions: ['transactions:write'], revokedAt: null, expiresAt: new Date(Date.now() + 86400000), }; - repository.findByHash.mockResolvedValue(mockKey); + + repository.findByPrefix.mockResolvedValue([mockKey]); repository.touchLastUsed.mockResolvedValue({ ...mockKey, lastUsedAt: new Date() }); const verified = await service.verify(rawSecret); expect(verified).toEqual(mockKey); - expect(repository.findByHash).toHaveBeenCalledWith(hash); + expect(repository.findByPrefix).toHaveBeenCalledWith(rawSecret.slice(0, 14)); expect(repository.touchLastUsed).toHaveBeenCalledWith('key-1'); }); + it('verifies legacy key with SHA-256 hash for backward compatibility', async () => { + const rawSecret = 'ak_live_abcdef1234567890abcdef1234567890abcdef12'; + const shaHash = sha256(rawSecret); + + const mockKey = { + id: 'key-legacy', + name: 'Legacy Key', + hashedKey: shaHash, + permissions: ['transactions:read'], + revokedAt: null, + expiresAt: new Date(Date.now() + 86400000), + }; + + repository.findByPrefix.mockResolvedValue([mockKey]); + repository.touchLastUsed.mockResolvedValue({ ...mockKey, lastUsedAt: new Date() }); + + const verified = await service.verify(rawSecret); + + expect(verified).toEqual(mockKey); + expect(repository.findByPrefix).toHaveBeenCalledWith(rawSecret.slice(0, 14)); + expect(repository.touchLastUsed).toHaveBeenCalledWith('key-legacy'); + }); + it('returns null for empty or invalid raw key input', async () => { expect(await service.verify('')).toBeNull(); expect(await service.verify(' ')).toBeNull(); expect(await service.verify(null as unknown as string)).toBeNull(); expect(await service.verify(undefined as unknown as string)).toBeNull(); - expect(repository.findByHash).not.toHaveBeenCalled(); + expect(repository.findByPrefix).not.toHaveBeenCalled(); }); - it('returns null when key hash is not found in database', async () => { - repository.findByHash.mockResolvedValue(null); + it('returns null when no keys found with matching prefix', async () => { + repository.findByPrefix.mockResolvedValue([]); const verified = await service.verify('ak_live_unknownkey'); @@ -220,11 +249,16 @@ describe('ApiKeyService', () => { }); it('returns null when key has been revoked', async () => { - repository.findByHash.mockResolvedValue({ - id: 'key-revoked', - hashedKey: hash, - revokedAt: new Date(Date.now() - 10000), - }); + const rawSecret = 'ak_live_abcdef1234567890abcdef1234567890abcdef12'; + const argonHash = await hashWithArgon2(rawSecret); + + repository.findByPrefix.mockResolvedValue([ + { + id: 'key-revoked', + hashedKey: argonHash, + revokedAt: new Date(Date.now() - 10000), + }, + ]); const verified = await service.verify(rawSecret); @@ -233,12 +267,17 @@ describe('ApiKeyService', () => { }); it('returns null when key has expired', async () => { - repository.findByHash.mockResolvedValue({ - id: 'key-expired', - hashedKey: hash, - revokedAt: null, - expiresAt: new Date(Date.now() - 5000), - }); + const rawSecret = 'ak_live_abcdef1234567890abcdef1234567890abcdef12'; + const argonHash = await hashWithArgon2(rawSecret); + + repository.findByPrefix.mockResolvedValue([ + { + id: 'key-expired', + hashedKey: argonHash, + revokedAt: null, + expiresAt: new Date(Date.now() - 5000), + }, + ]); const verified = await service.verify(rawSecret); @@ -247,18 +286,50 @@ describe('ApiKeyService', () => { }); it('still returns key if touchLastUsed throws a transient error', async () => { + const rawSecret = 'ak_live_abcdef1234567890abcdef1234567890abcdef12'; + const argonHash = await hashWithArgon2(rawSecret); + const mockKey = { id: 'key-1', - hashedKey: hash, + hashedKey: argonHash, revokedAt: null, expiresAt: null, }; - repository.findByHash.mockResolvedValue(mockKey); + + repository.findByPrefix.mockResolvedValue([mockKey]); repository.touchLastUsed.mockRejectedValue(new Error('DB connection busy')); const verified = await service.verify(rawSecret); expect(verified).toEqual(mockKey); }); + + it('tries multiple keys with same prefix until match is found', async () => { + const rawSecret = 'ak_live_abcdef1234567890abcdef1234567890abcdef12'; + const argonHash = await hashWithArgon2(rawSecret); + + const wrongKey = { + id: 'key-wrong', + hashedKey: await hashWithArgon2('different-key'), + revokedAt: null, + expiresAt: null, + }; + + const correctKey = { + id: 'key-correct', + hashedKey: argonHash, + permissions: ['transactions:write'], + revokedAt: null, + expiresAt: null, + }; + + repository.findByPrefix.mockResolvedValue([wrongKey, correctKey]); + repository.touchLastUsed.mockResolvedValue({ ...correctKey, lastUsedAt: new Date() }); + + const verified = await service.verify(rawSecret); + + expect(verified).toEqual(correctKey); + expect(repository.touchLastUsed).toHaveBeenCalledWith('key-correct'); + }); }); }); diff --git a/src/modules/health/health.controller.spec.ts b/src/modules/health/health.controller.spec.ts index 5413ddb2..8e6cc9db 100644 --- a/src/modules/health/health.controller.spec.ts +++ b/src/modules/health/health.controller.spec.ts @@ -1,4 +1,5 @@ import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { PATH_METADATA } from '@nestjs/common/constants'; import { Response } from 'express'; import { HealthController } from './health.controller'; import { PrismaHealthIndicator } from './indicators/prisma.health'; @@ -177,6 +178,13 @@ describe('HealthController', () => { expect(response.timestamp).toBeDefined(); }); + it('exposes the orchestration liveness and readiness routes', () => { + expect(Reflect.getMetadata(PATH_METADATA, HealthController.prototype.getLiveness)) + .toContain('live'); + expect(Reflect.getMetadata(PATH_METADATA, HealthController.prototype.getReadiness)) + .toContain('ready'); + }); + describe('GET /health/database', () => { it('returns 200 with status and latency when the database answers', async () => { await controller.getDatabase(res as Response); diff --git a/src/modules/health/health.controller.ts b/src/modules/health/health.controller.ts index ab763896..44fc41c8 100644 --- a/src/modules/health/health.controller.ts +++ b/src/modules/health/health.controller.ts @@ -98,7 +98,7 @@ export class HealthController { }; } - @Get('readiness') + @Get(['ready', 'readiness']) @ApiOperation({ summary: 'Application readiness check' }) @ApiResponse({ status: 200, description: 'Application is ready' }) @ApiResponse({ status: 503, description: 'Application is not ready' }) diff --git a/src/modules/health/indicators/stellar.health.ts b/src/modules/health/indicators/stellar.health.ts index 504a3668..b686bb62 100644 --- a/src/modules/health/indicators/stellar.health.ts +++ b/src/modules/health/indicators/stellar.health.ts @@ -77,7 +77,6 @@ export class StellarHealthIndicator { signal: controller.signal, headers: { Accept: 'application/json' }, }); - clearTimeout(timer); const latencyMs = Date.now() - start; if (!response.ok) { @@ -101,7 +100,6 @@ export class StellarHealthIndicator { protocolVersion: data.protocol_version || undefined, }; } catch (err) { - clearTimeout(timer); const latencyMs = Date.now() - start; const message = err instanceof Error ? err.message : String(err); this.logger.warn(`Horizon health check failed for ${url}: ${message}`); @@ -111,6 +109,8 @@ export class StellarHealthIndicator { url, error: message || 'Connection failed', }; + } finally { + clearTimeout(timer); } } @@ -130,7 +130,6 @@ export class StellarHealthIndicator { method: 'getHealth', }), }); - clearTimeout(timer); const latencyMs = Date.now() - start; if (!response.ok) { @@ -155,7 +154,6 @@ export class StellarHealthIndicator { ledgerSequence: data.result?.latestLedger || undefined, }; } catch (err) { - clearTimeout(timer); const latencyMs = Date.now() - start; const message = err instanceof Error ? err.message : String(err); this.logger.warn(`Soroban RPC health check failed for ${url}: ${message}`); @@ -165,6 +163,8 @@ export class StellarHealthIndicator { url, error: message || 'Connection failed', }; + } finally { + clearTimeout(timer); } } } diff --git a/src/modules/memory/memory.controller.ts b/src/modules/memory/memory.controller.ts index 4e924dbb..86fc7a4d 100644 --- a/src/modules/memory/memory.controller.ts +++ b/src/modules/memory/memory.controller.ts @@ -15,6 +15,7 @@ import { CurrentUser } from '../../common/decorators/current-user.decorator'; import { Roles } from '../../common/decorators/roles.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { PaginationQuery, paginationQuerySchema } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; @ApiTags('memory') @ApiBearerAuth('access-token') @@ -29,8 +30,7 @@ export class MemoryController { 'Returns a paginated list of memory records for the current organization. ' + 'Memory records capture agent decisions, reasoning, and outcomes for audit and learning.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiQuery({ name: 'search', required: false, type: String, description: 'Full-text search across task, reason, and summary fields' }) @ApiQuery({ name: 'agentId', required: false, type: String, description: 'Filter by agent UUID' }) @ApiResponse({ status: 200, description: 'Paginated list of memory records' }) diff --git a/src/modules/memory/memory.service.ts b/src/modules/memory/memory.service.ts index 99557337..80932333 100644 --- a/src/modules/memory/memory.service.ts +++ b/src/modules/memory/memory.service.ts @@ -51,7 +51,7 @@ export class MemoryService { } const pagination = toPrismaPagination(query, SORTABLE); const { items, total } = await this.repository.findManyAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); + return new Paginated(items, buildPaginationMeta(total, query)); } async getOrThrow(organizationId: string, id: string) { diff --git a/src/modules/metrics/metrics.controller.ts b/src/modules/metrics/metrics.controller.ts index a14ae497..a08e8da0 100644 --- a/src/modules/metrics/metrics.controller.ts +++ b/src/modules/metrics/metrics.controller.ts @@ -1,5 +1,6 @@ import { Controller, Get, Res, UseGuards } from '@nestjs/common'; import { ApiExcludeController } from '@nestjs/swagger'; +import { SkipThrottle } from '@nestjs/throttler'; import { Response } from 'express'; import { MetricsService } from './metrics.service'; import { MetricsAccessGuard } from './metrics-access.guard'; @@ -18,6 +19,7 @@ import { SkipPublicRateLimit } from '../../common/decorators/skip-public-rate-li @Controller('metrics') @Public() @SkipAudit() +@SkipThrottle() @SkipPublicRateLimit() @UseGuards(MetricsAccessGuard) export class MetricsController { diff --git a/src/modules/metrics/metrics.service.spec.ts b/src/modules/metrics/metrics.service.spec.ts index c081d32f..9532f34a 100644 --- a/src/modules/metrics/metrics.service.spec.ts +++ b/src/modules/metrics/metrics.service.spec.ts @@ -16,9 +16,11 @@ vi.mock('../../config/redis.config', () => ({ })); import { MetricsService } from './metrics.service'; +import { PrismaService } from '../../database/prisma.service'; describe('MetricsService', () => { let service: MetricsService; + let getPoolStats: ReturnType; beforeEach(() => { vi.clearAllMocks(); @@ -30,7 +32,8 @@ describe('MetricsService', () => { delayed: 0, paused: 0, }); - service = new MetricsService(); + getPoolStats = vi.fn().mockResolvedValue({ active: 2, idle: 5, waiting: 0 }); + service = new MetricsService({ getPoolStats } as unknown as PrismaService); }); it('exposes the Prometheus content type', () => { @@ -83,6 +86,16 @@ describe('MetricsService', () => { expect(close).toHaveBeenCalled(); }); + it('samples active/idle/waiting connection counts into the pool gauge', async () => { + const output = await service.getMetrics(); + + expect(getPoolStats).toHaveBeenCalled(); + expect(output).toContain('db_pool_connections'); + expect(output).toMatch(/db_pool_connections\{state="active"\} 2/); + expect(output).toMatch(/db_pool_connections\{state="idle"\} 5/); + expect(output).toMatch(/db_pool_connections\{state="waiting"\} 0/); + }); + describe('worker job metrics', () => { it('records successful job completion in the duration histogram', async () => { service.recordJobCompletion('webhooks', 'deliver', 0.25, 'success'); diff --git a/src/modules/metrics/metrics.service.ts b/src/modules/metrics/metrics.service.ts index f83ea0b9..66be1709 100644 --- a/src/modules/metrics/metrics.service.ts +++ b/src/modules/metrics/metrics.service.ts @@ -3,6 +3,7 @@ import { Registry, Counter, Histogram, Gauge } from 'prom-client'; import { Queue } from 'bullmq'; import { redisConfig } from '../../config/redis.config'; import { Queues } from '../../queues/queues.constants'; +import { PrismaService } from '../../database/prisma.service'; @Injectable() export class MetricsService implements OnModuleDestroy { @@ -44,7 +45,14 @@ export class MetricsService implements OnModuleDestroy { registers: [this.registry], }); - constructor() { + private readonly dbPoolConnectionsGauge = new Gauge({ + name: 'db_pool_connections', + help: 'Database connections by state (active, idle, waiting)', + labelNames: ['state'], + registers: [this.registry], + }); + + constructor(private readonly prisma: PrismaService) { const rConfig = redisConfig(); const connection = { host: rConfig.host, @@ -101,8 +109,16 @@ export class MetricsService implements OnModuleDestroy { } } + private async collectPoolMetrics(): Promise { + const stats = await this.prisma.getPoolStats(); + this.dbPoolConnectionsGauge.set({ state: 'active' }, stats.active); + this.dbPoolConnectionsGauge.set({ state: 'idle' }, stats.idle); + this.dbPoolConnectionsGauge.set({ state: 'waiting' }, stats.waiting); + } + public async getMetrics(): Promise { await this.collectQueueMetrics(); + await this.collectPoolMetrics(); return this.registry.metrics(); } diff --git a/src/modules/metrics/stream-metrics.service.spec.ts b/src/modules/metrics/stream-metrics.service.spec.ts index 7fbc25da..0e95c5c0 100644 --- a/src/modules/metrics/stream-metrics.service.spec.ts +++ b/src/modules/metrics/stream-metrics.service.spec.ts @@ -17,6 +17,7 @@ vi.mock('../../config/redis.config', () => ({ import { MetricsService } from './metrics.service'; import { StreamMetricsService } from './stream-metrics.service'; +import { PrismaService } from '../../database/prisma.service'; describe('StreamMetricsService', () => { let metricsService: MetricsService; @@ -25,7 +26,8 @@ describe('StreamMetricsService', () => { beforeEach(() => { vi.clearAllMocks(); getJobCounts.mockResolvedValue({ waiting: 0, active: 0, completed: 0, failed: 0, delayed: 0, paused: 0 }); - metricsService = new MetricsService(); + const getPoolStats = vi.fn().mockResolvedValue({ active: 0, idle: 0, waiting: 0 }); + metricsService = new MetricsService({ getPoolStats } as unknown as PrismaService); service = new StreamMetricsService(metricsService); }); diff --git a/src/modules/notifications/notification.controller.ts b/src/modules/notifications/notification.controller.ts index 2ef57654..592f818e 100644 --- a/src/modules/notifications/notification.controller.ts +++ b/src/modules/notifications/notification.controller.ts @@ -13,6 +13,7 @@ import { CurrentUser } from '../../common/decorators/current-user.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { AuthenticatedUser } from '../../common/interfaces/authenticated-user.interface'; import { PaginationQuery, paginationQuerySchema } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; @ApiTags('notifications') @ApiBearerAuth('access-token') @@ -27,8 +28,7 @@ export class NotificationController { 'Returns a paginated list of notifications for the authenticated user. ' + 'Supports filtering by read status.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiQuery({ name: 'unread', required: false, type: Boolean, description: 'Filter by unread status' }) @ApiResponse({ status: 200, description: 'Paginated list of notifications' }) @ApiResponse({ status: 401, description: 'Not authenticated' }) diff --git a/src/modules/notifications/notification.service.ts b/src/modules/notifications/notification.service.ts index a2f50107..62c50525 100644 --- a/src/modules/notifications/notification.service.ts +++ b/src/modules/notifications/notification.service.ts @@ -71,7 +71,7 @@ export class NotificationService { } const pagination = toPrismaPagination(query, SORTABLE); const { items, total } = await this.repository.findManyAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); + return new Paginated(items, buildPaginationMeta(total, query)); } async unreadCount(organizationId: string, userId: string) { diff --git a/src/modules/organizations/organization.controller.ts b/src/modules/organizations/organization.controller.ts index 3e299abb..3c1f333b 100644 --- a/src/modules/organizations/organization.controller.ts +++ b/src/modules/organizations/organization.controller.ts @@ -6,7 +6,6 @@ import { ApiResponse, ApiParam, ApiBody, - ApiQuery, } from '@nestjs/swagger'; import { UserRole } from '@prisma/client'; import { OrganizationService } from './organization.service'; @@ -27,6 +26,7 @@ import { AuditAction } from '../../common/decorators/audit-action.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { AuthenticatedUser } from '../../common/interfaces/authenticated-user.interface'; import { PaginationQuery, paginationQuerySchema } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; @ApiTags('organizations') @ApiBearerAuth('access-token') @@ -70,8 +70,7 @@ export class OrganizationController { summary: 'List organization members', description: 'Returns a paginated list of members in the current organization.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiResponse({ status: 200, description: 'Paginated list of members' }) @ApiResponse({ status: 401, description: 'Not authenticated' }) listMembers( diff --git a/src/modules/organizations/organization.service.ts b/src/modules/organizations/organization.service.ts index 3c7b20c5..eb5c5ebe 100644 --- a/src/modules/organizations/organization.service.ts +++ b/src/modules/organizations/organization.service.ts @@ -74,7 +74,7 @@ export class OrganizationService { } const pagination = toPrismaPagination(query, MEMBER_SORTABLE); const { items, total } = await this.repository.findMembersAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); + return new Paginated(items, buildPaginationMeta(total, query)); } async inviteMember(organizationId: string, actorId: string, input: InviteMemberInput) { diff --git a/src/modules/policies/guards/agent-policy.guard.ts b/src/modules/policies/guards/agent-policy.guard.ts index 96543eb8..d2e07a5c 100644 --- a/src/modules/policies/guards/agent-policy.guard.ts +++ b/src/modules/policies/guards/agent-policy.guard.ts @@ -78,7 +78,16 @@ export class AgentPolicyGuard implements CanActivate { } // Check velocity limits (rolling 24-hour window) - await this.policyService.checkVelocityLimit(agentId, Number(amount), asset); + const actorId = request.user?.isApiKey + ? request.user.createdById ?? undefined + : request.user?.id; + await this.policyService.checkVelocityLimit( + organizationId, + agentId, + Number(amount), + asset, + actorId, + ); return true; } catch (error) { diff --git a/src/modules/policies/index.ts b/src/modules/policies/index.ts index 208e8ce8..1ed451d1 100644 --- a/src/modules/policies/index.ts +++ b/src/modules/policies/index.ts @@ -1,6 +1,8 @@ export * from './policy.types'; export * from './policy.engine'; export * from './policy.service'; +export * from './spending-policy.service'; +export * from './spending-policy.repository'; export * from './policy.module'; export * from './policy-override-expired.event'; export * from './services/policy-override-cleanup.service'; diff --git a/src/modules/policies/policy.controller.ts b/src/modules/policies/policy.controller.ts index 8d10690a..956b2831 100644 --- a/src/modules/policies/policy.controller.ts +++ b/src/modules/policies/policy.controller.ts @@ -23,12 +23,14 @@ import { import { CurrentUser } from '../../common/decorators/current-user.decorator'; import { Roles } from '../../common/decorators/roles.decorator'; import { AuditAction } from '../../common/decorators/audit-action.decorator'; +import { AuditLog } from '../../common/decorators/audit-log.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { AuthenticatedUser } from '../../common/interfaces/authenticated-user.interface'; import { PaginationQuery, paginationQuerySchema, } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; import { ApiEnvelope } from '../../common/decorators/api-envelope.decorator'; @ApiTags('policies') @@ -44,8 +46,7 @@ export class PolicyController { 'Returns a paginated list of policies for the current organization. ' + 'Supports filtering by type, status, and agent.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiQuery({ name: 'type', required: false, enum: ['SPENDING_LIMIT', 'APPROVAL_REQUIRED', 'ALLOWLIST', 'TIME_WINDOW'], description: 'Filter by policy type' }) @ApiQuery({ name: 'enabled', required: false, type: Boolean, description: 'Filter by enabled status' }) @ApiEnvelope(CreatePolicyDto as never, { isArray: true }) @@ -61,6 +62,7 @@ export class PolicyController { @Post() @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE) @AuditAction('POLICY_CREATED') + @AuditLog({ action: 'POLICY_CREATED', entity: 'Policy' }) @ApiOperation({ summary: 'Create a policy', description: @@ -114,6 +116,7 @@ export class PolicyController { @Patch(':id') @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE) @AuditAction('POLICY_UPDATED') + @AuditLog({ action: 'POLICY_UPDATED', entity: 'Policy' }) @ApiOperation({ summary: 'Update a policy', description: @@ -137,6 +140,7 @@ export class PolicyController { @Delete(':id') @Roles(UserRole.OWNER, UserRole.ADMIN) @AuditAction('POLICY_DELETED') + @AuditLog({ action: 'POLICY_DELETED', entity: 'Policy' }) @ApiOperation({ summary: 'Delete (soft) a policy', description: diff --git a/src/modules/policies/policy.module.ts b/src/modules/policies/policy.module.ts index 0b109a24..486d9f51 100644 --- a/src/modules/policies/policy.module.ts +++ b/src/modules/policies/policy.module.ts @@ -1,7 +1,8 @@ import { Module } from '@nestjs/common'; import { PolicyController } from './policy.controller'; import { PolicyService } from './policy.service'; -import { PolicyRepository } from './policy.repository'; +import { SpendingPolicyService } from './spending-policy.service'; +import { SpendingPolicyRepository } from './spending-policy.repository'; import { PolicyEngine } from './policy.engine'; import { PolicyOverrideCleanupService } from './services/policy-override-cleanup.service'; import { AgentPolicyGuard } from './guards/agent-policy.guard'; @@ -9,10 +10,27 @@ import { AgentPolicyGuard } from './guards/agent-policy.guard'; /** * Policy module. Exports the service + engine so the transactions module can * evaluate intents during the payment pipeline. + * + * Persistence is layered: `SpendingPolicyService` owns spending-policy + * validation and enforcement, and delegates every Prisma call to + * `SpendingPolicyRepository`. */ @Module({ controllers: [PolicyController], - providers: [PolicyService, PolicyRepository, PolicyEngine, PolicyOverrideCleanupService, AgentPolicyGuard], - exports: [PolicyService, PolicyEngine, PolicyOverrideCleanupService, AgentPolicyGuard], + providers: [ + PolicyService, + SpendingPolicyService, + SpendingPolicyRepository, + PolicyEngine, + PolicyOverrideCleanupService, + AgentPolicyGuard, + ], + exports: [ + PolicyService, + SpendingPolicyService, + PolicyEngine, + PolicyOverrideCleanupService, + AgentPolicyGuard, + ], }) export class PolicyModule {} diff --git a/src/modules/policies/policy.repository.ts b/src/modules/policies/policy.repository.ts deleted file mode 100644 index ead81685..00000000 --- a/src/modules/policies/policy.repository.ts +++ /dev/null @@ -1,62 +0,0 @@ -import { Injectable } from '@nestjs/common'; -import { Prisma } from '@prisma/client'; -import { PrismaService } from '../../database/prisma.service'; -import { PrismaPagination } from '../../common/helpers/pagination'; - -/** Persistence for Policy rows. */ -@Injectable() -export class PolicyRepository { - constructor(private readonly prisma: PrismaService) {} - - create(data: Prisma.PolicyCreateInput) { - return this.prisma.policy.create({ data }); - } - - findById(organizationId: string, id: string) { - return this.prisma.policy.findFirst({ where: { id, organizationId, deletedAt: null } }); - } - - /** Returns the enabled policies applicable to an org (and optionally an agent). */ - findActiveForEvaluation(organizationId: string, agentId?: string) { - return this.prisma.policy.findMany({ - where: { - organizationId, - enabled: true, - deletedAt: null, - OR: [{ agentId: null }, ...(agentId ? [{ agentId }] : [])], - }, - orderBy: { priority: 'asc' }, - }); - } - - /** Returns enabled policies for a specific agent (used for velocity checks). */ - findActiveForEvaluationByAgent(agentId: string) { - return this.prisma.policy.findMany({ - where: { - agentId, - enabled: true, - deletedAt: null, - }, - orderBy: { priority: 'asc' }, - }); - } - - async findManyAndCount(where: Prisma.PolicyWhereInput, pagination: PrismaPagination) { - const [items, total] = await this.prisma.$transaction([ - this.prisma.policy.findMany({ where, ...pagination }), - this.prisma.policy.count({ where }), - ]); - return { items, total }; - } - - update(id: string, data: Prisma.PolicyUpdateInput) { - return this.prisma.policy.update({ where: { id }, data }); - } - - softDelete(id: string) { - return this.prisma.policy.update({ - where: { id }, - data: { deletedAt: new Date(), enabled: false }, - }); - } -} diff --git a/src/modules/policies/policy.service.spec.ts b/src/modules/policies/policy.service.spec.ts new file mode 100644 index 00000000..4a8e1437 --- /dev/null +++ b/src/modules/policies/policy.service.spec.ts @@ -0,0 +1,21 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { PolicyService } from './policy.service'; + +describe('PolicyService velocity limit delegation', () => { + let service: PolicyService; + let checkVelocityLimit: ReturnType; + + beforeEach(() => { + checkVelocityLimit = vi.fn().mockResolvedValue(undefined); + service = new PolicyService( + { checkVelocityLimit } as never, + {} as never, + { emit: vi.fn() } as never, + ); + }); + + it('forwards the transaction governance arguments to the spending policy service', async () => { + await service.checkVelocityLimit('org-1', 'agent-1', 3, 'XLM', 'user-1'); + expect(checkVelocityLimit).toHaveBeenCalledWith('agent-1', 3, 'XLM'); + }); +}); diff --git a/src/modules/policies/policy.service.ts b/src/modules/policies/policy.service.ts index 77e27002..e88b5f4f 100644 --- a/src/modules/policies/policy.service.ts +++ b/src/modules/policies/policy.service.ts @@ -1,62 +1,28 @@ import { Injectable } from '@nestjs/common'; -import { Policy, Prisma } from '@prisma/client'; -import { PolicyRepository } from './policy.repository'; +import { Policy } from '@prisma/client'; +import { SpendingPolicyService } from './spending-policy.service'; import { PolicyEngine } from './policy.engine'; import { CreatePolicyInput, SimulatePolicyInput, UpdatePolicyInput } from './policy.dto'; -import { - EvaluablePolicy, - PolicyConfiguration, - PolicyEvaluationResult, - TransactionIntent, - policyConfigurationSchemaStrict, -} from './policy.types'; -import { NotFoundException, VelocityLimitExceededException, ValidationException } from '../../common/exceptions/domain.exception'; -import { formatZodError } from '../../common/validators/zod-error'; -import { - buildPaginationMeta, - PaginationQuery, - toPrismaPagination, -} from '../../common/helpers/pagination'; -import { Paginated } from '../../common/interfaces/api-response.interface'; +import { EvaluablePolicy, PolicyConfiguration, PolicyEvaluationResult, TransactionIntent } from './policy.types'; +import { PaginationQuery } from '../../common/helpers/pagination'; import { EventBusService } from '../../events/event-bus.service'; import { DomainEventName } from '../../events/event-names'; -import { PrismaService } from '../../database/prisma.service'; - -const SORTABLE = ['createdAt', 'priority', 'name', 'type']; /** - * Manages policy definitions and exposes evaluation to other modules. Wraps the - * pure {@link PolicyEngine} with persistence, event emission and simulation. + * Controller-facing façade over the policy domain. Delegates persistence and + * spending-policy enforcement to {@link SpendingPolicyService} and wraps every + * mutation with the domain events the audit ledger depends on. */ @Injectable() export class PolicyService { constructor( - private readonly repository: PolicyRepository, + private readonly spendingPolicyService: SpendingPolicyService, private readonly engine: PolicyEngine, private readonly eventBus: EventBusService, - private readonly prisma: PrismaService, ) {} async create(organizationId: string, actorId: string, input: CreatePolicyInput) { - // Validate configuration using strict schema - const validationResult = policyConfigurationSchemaStrict.safeParse(input.configuration); - if (!validationResult.success) { - throw new ValidationException( - 'Invalid policy configuration', - formatZodError(validationResult.error), - ); - } - - const policy = await this.repository.create({ - organization: { connect: { id: organizationId } }, - ...(input.agentId ? { agent: { connect: { id: input.agentId } } } : {}), - name: input.name, - description: input.description, - type: input.type, - configuration: validationResult.data as Prisma.InputJsonValue, - priority: input.priority, - enabled: input.enabled, - }); + const policy = await this.spendingPolicyService.create(organizationId, input); await this.eventBus.emit( DomainEventName.PolicyCreated, { policyId: policy.id, name: policy.name, type: policy.type }, @@ -65,45 +31,16 @@ export class PolicyService { return policy; } - async list(organizationId: string, query: PaginationQuery) { - const where: Prisma.PolicyWhereInput = { organizationId, deletedAt: null }; - if (query.search) { - where.name = { contains: query.search, mode: 'insensitive' }; - } - const pagination = toPrismaPagination(query, SORTABLE); - const { items, total } = await this.repository.findManyAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); + list(organizationId: string, query: PaginationQuery) { + return this.spendingPolicyService.list(organizationId, query); } - async getOrThrow(organizationId: string, id: string): Promise { - const policy = await this.repository.findById(organizationId, id); - if (!policy) { - throw new NotFoundException('Policy', id); - } - return policy; + getOrThrow(organizationId: string, id: string): Promise { + return this.spendingPolicyService.getOrThrow(organizationId, id); } async update(organizationId: string, actorId: string, id: string, input: UpdatePolicyInput) { - await this.getOrThrow(organizationId, id); - const data: Prisma.PolicyUpdateInput = { - name: input.name, - description: input.description, - type: input.type, - priority: input.priority, - enabled: input.enabled, - }; - if (input.configuration) { - // Validate configuration using strict schema - const validationResult = policyConfigurationSchemaStrict.safeParse(input.configuration); - if (!validationResult.success) { - throw new ValidationException( - 'Invalid policy configuration', - formatZodError(validationResult.error), - ); - } - data.configuration = validationResult.data as Prisma.InputJsonValue; - } - const policy = await this.repository.update(id, data); + const policy = await this.spendingPolicyService.update(organizationId, id, input); await this.eventBus.emit( DomainEventName.PolicyUpdated, { policyId: id }, @@ -113,14 +50,13 @@ export class PolicyService { } async remove(organizationId: string, actorId: string, id: string) { - await this.getOrThrow(organizationId, id); - await this.repository.softDelete(id); + const result = await this.spendingPolicyService.remove(organizationId, id); await this.eventBus.emit( DomainEventName.PolicyDeleted, { policyId: id }, { organizationId, actorId, aggregateType: 'policy', aggregateId: id }, ); - return { id, deleted: true }; + return result; } /** @@ -132,13 +68,12 @@ export class PolicyService { intent: TransactionIntent, actorId?: string, ): Promise { - const policies = await this.repository.findActiveForEvaluation( + const policies = await this.spendingPolicyService.listActiveForEvaluation( intent.organizationId, intent.agentId, ); const result = this.engine.evaluate(intent, policies.map(toEvaluable)); - // Emit domain events for the ledger await this.eventBus.emit( DomainEventName.PolicyEvaluated, { @@ -170,27 +105,8 @@ export class PolicyService { ); } - // Persist audit log for policy evaluation if (actorId) { - await this.prisma.auditLog.create({ - data: { - organizationId: intent.organizationId, - userId: actorId, - action: 'POLICY_EVALUATED', - entity: 'policy', - entityId: result.matchedPolicyId, - oldValue: null as unknown as Prisma.InputJsonValue, - newValue: { - passed: result.passed, - requiresApproval: result.requiresApproval, - violations: result.violations, - transactionIntent: intent, - } as unknown as Prisma.InputJsonValue, - }, - }).catch((error) => { - // Audit log failures should not block policy evaluation - console.error('Failed to persist policy evaluation audit log:', error); - }); + await this.spendingPolicyService.recordEvaluationAudit(intent, result, actorId); } return result; @@ -209,7 +125,7 @@ export class PolicyService { spentThisWeek: input.spentThisWeek, spentThisMonth: input.spentThisMonth, }; - const policies = await this.repository.findActiveForEvaluation(organizationId, input.agentId); + const policies = await this.spendingPolicyService.listActiveForEvaluation(organizationId, input.agentId); const result = this.engine.evaluate(intent, policies.map(toEvaluable)); return { passed: result.passed, @@ -220,58 +136,20 @@ export class PolicyService { } /** - * Check velocity limit for an agent's spending within a rolling 24-hour window. - * This acts as a circuit breaker to prevent rapid draining of wallets. + * Checks the rolling 24-hour velocity limit for an agent's spending. Acts as + * a circuit breaker to prevent rapid draining of wallets. Delegates the + * actual enforcement to {@link SpendingPolicyService}; `organizationId` and + * `actorId` are accepted for call-site symmetry with the rest of the + * transaction governance pipeline but are not needed by the check itself. */ - async checkVelocityLimit(agentId: string, amount: number, assetCode: string): Promise { - const twentyFourHoursAgo = new Date(Date.now() - 24 * 60 * 60 * 1000); - - // Query historical agent transactions from the last 24 hours - const transactions = await this.prisma.transaction.findMany({ - where: { - agentId, - status: { in: ['COMPLETED', 'CONFIRMED'] }, - asset: assetCode, - createdAt: { gte: twentyFourHoursAgo }, - }, - select: { - amount: true, - }, - }); - - // Sum up transaction volumes - const spentInWindow = transactions.reduce( - (sum, tx) => sum + Number(tx.amount), - 0, - ); - - // Retrieve the agent's active daily limit from policies - const policies = await this.repository.findActiveForEvaluationByAgent(agentId); - const dailyLimitPolicy = policies.find((policy) => { - const config = policy.configuration as PolicyConfiguration; - return config.dailyLimit !== undefined && config.dailyLimit > 0; - }); - - if (!dailyLimitPolicy) { - // No daily limit configured, allow the transaction - return; - } - - const config = dailyLimitPolicy.configuration as PolicyConfiguration; - const dailyLimit = config.dailyLimit!; - - // Check if the pending transaction would exceed the limit - if (spentInWindow + amount > dailyLimit) { - throw new VelocityLimitExceededException( - `Daily velocity limit exceeded. Spent: ${spentInWindow}, Pending: ${amount}, Limit: ${dailyLimit}`, - { - spentInWindow, - pendingAmount: amount, - limit: dailyLimit, - assetCode, - }, - ); - } + checkVelocityLimit( + _organizationId: string, + agentId: string, + amount: number, + assetCode: string, + _actorId?: string, + ): Promise { + return this.spendingPolicyService.checkVelocityLimit(agentId, amount, assetCode); } } diff --git a/src/modules/policies/spending-policy.repository.spec.ts b/src/modules/policies/spending-policy.repository.spec.ts new file mode 100644 index 00000000..03623d82 --- /dev/null +++ b/src/modules/policies/spending-policy.repository.spec.ts @@ -0,0 +1,204 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { Logger } from '@nestjs/common'; +import { TransactionStatus } from '@prisma/client'; + +import { PrismaPagination } from '../../common/helpers/pagination'; +import { PrismaService } from '../../database/prisma.service'; +import { + SETTLED_TRANSACTION_STATUSES, + SpendingPolicyRepository, +} from './spending-policy.repository'; + +type MockPrisma = { + policy: { + create: ReturnType; + findFirst: ReturnType; + findMany: ReturnType; + update: ReturnType; + count: ReturnType; + }; + transaction: { findMany: ReturnType }; + auditLog: { create: ReturnType }; + $transaction: ReturnType; +}; + +/** Mocked Prisma client: `$transaction` supports both the array and callback forms. */ +function makePrisma(): MockPrisma { + return { + policy: { + create: vi.fn().mockResolvedValue({ id: 'policy-1' }), + findFirst: vi.fn().mockResolvedValue({ id: 'policy-1' }), + findMany: vi.fn().mockResolvedValue([]), + update: vi.fn().mockResolvedValue({ id: 'policy-1' }), + count: vi.fn().mockResolvedValue(0), + }, + transaction: { findMany: vi.fn().mockResolvedValue([]) }, + auditLog: { create: vi.fn().mockResolvedValue({ id: 'audit-1' }) }, + $transaction: vi.fn((arg: unknown) => + typeof arg === 'function' + ? (arg as (tx: unknown) => unknown)({}) + : Promise.all(arg as Array>), + ), + }; +} + +const PAGINATION: PrismaPagination = { skip: 0, take: 20, orderBy: { createdAt: 'desc' } }; + +describe('SpendingPolicyRepository', () => { + let prisma: MockPrisma; + let repository: SpendingPolicyRepository; + + beforeEach(() => { + prisma = makePrisma(); + repository = new SpendingPolicyRepository(prisma as unknown as PrismaService); + }); + + describe('policy persistence', () => { + it('creates a policy through the Prisma policy delegate', async () => { + const data = { name: 'Daily limit', type: 'SPENDING_LIMIT' } as never; + + await repository.create(data); + + expect(prisma.policy.create).toHaveBeenCalledWith({ data }); + }); + + it('scopes findById to the organization and excludes soft-deleted rows', async () => { + await repository.findById('org-1', 'policy-1'); + + expect(prisma.policy.findFirst).toHaveBeenCalledWith({ + where: { id: 'policy-1', organizationId: 'org-1', deletedAt: null }, + }); + }); + + it('fetches org-wide policies plus agent-specific ones in priority order', async () => { + await repository.findActiveForEvaluation('org-1', 'agent-1'); + + expect(prisma.policy.findMany).toHaveBeenCalledWith({ + where: { + organizationId: 'org-1', + enabled: true, + deletedAt: null, + OR: [{ agentId: null }, { agentId: 'agent-1' }], + }, + orderBy: { priority: 'asc' }, + }); + }); + + it('omits the agent clause when no agent is supplied', async () => { + await repository.findActiveForEvaluation('org-1'); + + const args = prisma.policy.findMany.mock.calls[0][0]; + expect(args.where.OR).toEqual([{ agentId: null }]); + }); + + it('updates and soft-deletes through the policy delegate', async () => { + await repository.update('policy-1', { name: 'Renamed' }); + expect(prisma.policy.update).toHaveBeenCalledWith({ + where: { id: 'policy-1' }, + data: { name: 'Renamed' }, + }); + + await repository.softDelete('policy-1'); + expect(prisma.policy.update).toHaveBeenLastCalledWith({ + where: { id: 'policy-1' }, + data: { deletedAt: expect.any(Date), enabled: false }, + }); + }); + + it('reads the page and the total inside a single transaction', async () => { + prisma.policy.findMany.mockResolvedValue([{ id: 'policy-1' }]); + prisma.policy.count.mockResolvedValue(1); + + const result = await repository.findManyAndCount({ organizationId: 'org-1' }, PAGINATION); + + expect(prisma.$transaction).toHaveBeenCalledTimes(1); + expect(result).toEqual({ items: [{ id: 'policy-1' }], total: 1 }); + }); + }); + + describe('spend aggregation', () => { + it('sums settled spend as a number, including string-encoded decimals', async () => { + prisma.transaction.findMany.mockResolvedValue([ + { amount: '12.5' }, + { amount: 7.25 }, + { amount: '0.25' }, + ]); + + const total = await repository.sumSpentInWindow({ + agentId: 'agent-1', + assetCode: 'USDC', + since: new Date('2026-01-01T00:00:00Z'), + }); + + expect(total).toBe(20); + expect(prisma.transaction.findMany).toHaveBeenCalledWith({ + where: { + agentId: 'agent-1', + asset: 'USDC', + status: { in: [...SETTLED_TRANSACTION_STATUSES] }, + createdAt: { gte: new Date('2026-01-01T00:00:00Z') }, + }, + select: { amount: true }, + }); + }); + + it('honours an explicit status filter', async () => { + await repository.sumSpentInWindow({ + agentId: 'agent-1', + assetCode: 'XLM', + since: new Date(), + statuses: [TransactionStatus.PENDING], + }); + + const args = prisma.transaction.findMany.mock.calls[0][0]; + expect(args.where.status).toEqual({ in: [TransactionStatus.PENDING] }); + }); + }); + + describe('evaluation audit', () => { + it('appends a POLICY_EVALUATED row to the audit log', async () => { + await repository.recordEvaluationAudit({ + organizationId: 'org-1', + userId: 'user-1', + policyId: 'policy-1', + payload: { passed: true }, + }); + + expect(prisma.auditLog.create).toHaveBeenCalledWith({ + data: expect.objectContaining({ + organizationId: 'org-1', + userId: 'user-1', + action: 'POLICY_EVALUATED', + entity: 'policy', + entityId: 'policy-1', + newValue: { passed: true }, + }), + }); + }); + }); + + describe('transaction safety', () => { + it('runs interactive work inside $transaction', async () => { + const work = vi.fn().mockResolvedValue('ok'); + + await expect(repository.withTransaction(work)).resolves.toBe('ok'); + + expect(prisma.$transaction).toHaveBeenCalledWith(work); + expect(work).toHaveBeenCalledTimes(1); + }); + }); + + describe('uniform error handling', () => { + it('logs the failing operation and rethrows the original error', async () => { + const logger = vi.spyOn(Logger.prototype, 'error').mockImplementation(() => undefined); + const failure = new Error('connection lost'); + prisma.policy.create.mockRejectedValue(failure); + + await expect(repository.create({} as never)).rejects.toBe(failure); + expect(logger).toHaveBeenCalledWith( + expect.stringContaining('SpendingPolicyRepository.create failed'), + ); + logger.mockRestore(); + }); + }); +}); diff --git a/src/modules/policies/spending-policy.repository.ts b/src/modules/policies/spending-policy.repository.ts new file mode 100644 index 00000000..777b31d6 --- /dev/null +++ b/src/modules/policies/spending-policy.repository.ts @@ -0,0 +1,172 @@ +import { Injectable, Logger } from '@nestjs/common'; +import { Prisma, TransactionStatus } from '@prisma/client'; + +import { PrismaPagination } from '../../common/helpers/pagination'; +import { PrismaService } from '../../database/prisma.service'; + +/** Transaction states that count as real, settled spend for velocity checks. */ +export const SETTLED_TRANSACTION_STATUSES: readonly TransactionStatus[] = [ + TransactionStatus.COMPLETED, + TransactionStatus.CONFIRMED, +]; + +/** Query describing the rolling spend window for one agent/asset pair. */ +export interface SpendingWindowQuery { + agentId: string; + assetCode: string; + since: Date; + statuses?: readonly TransactionStatus[]; +} + +/** Fields persisted when a policy evaluation is appended to the audit trail. */ +export interface PolicyEvaluationAuditInput { + organizationId: string; + userId: string; + policyId?: string | null; + payload: Prisma.InputJsonValue; +} + +/** + * Repository for agent spending policies. + * + * Every Prisma call that concerns a spending policy — creation, retrieval, + * updates, soft deletes and the spend aggregation used by enforcement checks — + * lives here, so services depend on a small, mockable surface instead of the + * Prisma client itself. + * + * Behaviour every method shares: + * - **Uniform error handling.** Failures are logged once with the operation + * name (plus the Prisma error code when available) and rethrown, so callers + * keep seeing the original error while operators get a consistent trail. + * - **Transaction safety.** Composite reads run inside `$transaction`, and + * {@link withTransaction} exposes an interactive transaction for callers that + * need several writes to commit or roll back together. + */ +@Injectable() +export class SpendingPolicyRepository { + private readonly logger = new Logger(SpendingPolicyRepository.name); + + constructor(private readonly prisma: PrismaService) {} + + /** Persists a new policy row. */ + create(data: Prisma.PolicyCreateInput) { + return this.execute('create', () => this.prisma.policy.create({ data })); + } + + /** Returns a live (non-deleted) policy scoped to its organization. */ + findById(organizationId: string, id: string) { + return this.execute('findById', () => + this.prisma.policy.findFirst({ where: { id, organizationId, deletedAt: null } }), + ); + } + + /** Enabled policies applicable to an organization, optionally agent-scoped. */ + findActiveForEvaluation(organizationId: string, agentId?: string) { + return this.execute('findActiveForEvaluation', () => + this.prisma.policy.findMany({ + where: { + organizationId, + enabled: true, + deletedAt: null, + OR: [{ agentId: null }, ...(agentId ? [{ agentId }] : [])], + }, + orderBy: { priority: 'asc' }, + }), + ); + } + + /** Enabled policies bound to a single agent (used by velocity checks). */ + findActiveForEvaluationByAgent(agentId: string) { + return this.execute('findActiveForEvaluationByAgent', () => + this.prisma.policy.findMany({ + where: { agentId, enabled: true, deletedAt: null }, + orderBy: { priority: 'asc' }, + }), + ); + } + + /** Paginated policy list plus total count, read inside one transaction. */ + findManyAndCount(where: Prisma.PolicyWhereInput, pagination: PrismaPagination) { + return this.execute('findManyAndCount', async () => { + const [items, total] = await this.prisma.$transaction([ + this.prisma.policy.findMany({ where, ...pagination }), + this.prisma.policy.count({ where }), + ]); + return { items, total }; + }); + } + + /** Applies a partial update to a policy row. */ + update(id: string, data: Prisma.PolicyUpdateInput) { + return this.execute('update', () => this.prisma.policy.update({ where: { id }, data })); + } + + /** Soft-deletes a policy: it stays queryable for audits but stops applying. */ + softDelete(id: string) { + return this.execute('softDelete', () => + this.prisma.policy.update({ + where: { id }, + data: { deletedAt: new Date(), enabled: false }, + }), + ); + } + + /** + * Sums the amount an agent already spent for one asset since `since`. + * + * Aggregation happens in the client (rather than SQL) so `Decimal` handling + * stays explicit and the method is trivially mockable in unit tests. + */ + async sumSpentInWindow(query: SpendingWindowQuery): Promise { + return this.execute('sumSpentInWindow', async () => { + const rows = await this.prisma.transaction.findMany({ + where: { + agentId: query.agentId, + asset: query.assetCode, + status: { in: [...(query.statuses ?? SETTLED_TRANSACTION_STATUSES)] }, + createdAt: { gte: query.since }, + }, + select: { amount: true }, + }); + return rows.reduce((sum, row) => sum + Number(row.amount), 0); + }); + } + + /** Appends a policy-evaluation entry to the immutable audit trail. */ + recordEvaluationAudit(input: PolicyEvaluationAuditInput) { + return this.execute('recordEvaluationAudit', () => + this.prisma.auditLog.create({ + data: { + organizationId: input.organizationId, + userId: input.userId, + action: 'POLICY_EVALUATED', + entity: 'policy', + entityId: input.policyId ?? null, + oldValue: Prisma.JsonNull, + newValue: input.payload, + }, + }), + ); + } + + /** + * Runs `work` inside an interactive Prisma transaction so multi-step policy + * mutations either commit together or roll back together. + */ + withTransaction(work: (tx: Prisma.TransactionClient) => Promise): Promise { + return this.execute('withTransaction', () => this.prisma.$transaction(work)); + } + + /** Single funnel for logging + rethrowing, keeping error handling uniform. */ + private async execute(operation: string, work: () => Promise): Promise { + try { + return await work(); + } catch (error) { + const code = error instanceof Prisma.PrismaClientKnownRequestError ? ` (${error.code})` : ''; + this.logger.error( + `SpendingPolicyRepository.${operation} failed${code}: ${(error as Error).message}`, + ); + throw error; + } + } +} diff --git a/src/modules/policies/spending-policy.service.spec.ts b/src/modules/policies/spending-policy.service.spec.ts new file mode 100644 index 00000000..27b24a01 --- /dev/null +++ b/src/modules/policies/spending-policy.service.spec.ts @@ -0,0 +1,230 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { Logger } from '@nestjs/common'; +import { PolicyType } from '@prisma/client'; + +import { + NotFoundException, + ValidationException, + VelocityLimitExceededException, +} from '../../common/exceptions/domain.exception'; +import { PaginationQuery } from '../../common/helpers/pagination'; +import { SpendingPolicyRepository } from './spending-policy.repository'; +import { SpendingPolicyService } from './spending-policy.service'; +import { CreatePolicyInput, UpdatePolicyInput } from './policy.dto'; + +type MockRepository = { + create: ReturnType; + findById: ReturnType; + findActiveForEvaluation: ReturnType; + findActiveForEvaluationByAgent: ReturnType; + findManyAndCount: ReturnType; + update: ReturnType; + softDelete: ReturnType; + sumSpentInWindow: ReturnType; + recordEvaluationAudit: ReturnType; +}; + +function makeRepository(): MockRepository { + return { + create: vi.fn().mockResolvedValue({ id: 'policy-1' }), + findById: vi.fn().mockResolvedValue({ id: 'policy-1' }), + findActiveForEvaluation: vi.fn().mockResolvedValue([]), + findActiveForEvaluationByAgent: vi.fn().mockResolvedValue([]), + findManyAndCount: vi.fn().mockResolvedValue({ items: [], total: 0 }), + update: vi.fn().mockResolvedValue({ id: 'policy-1' }), + softDelete: vi.fn().mockResolvedValue({ id: 'policy-1' }), + sumSpentInWindow: vi.fn().mockResolvedValue(0), + recordEvaluationAudit: vi.fn().mockResolvedValue({ id: 'audit-1' }), + }; +} + +const CREATE_INPUT: CreatePolicyInput = { + name: 'Daily limit', + type: PolicyType.MAX_AMOUNT, + configuration: { dailyLimit: 100 }, + priority: 100, + enabled: true, +}; + +/** A policy row shape good enough for the daily-limit lookup. */ +function policyRow(configuration: Record) { + return { id: 'policy-1', configuration }; +} + +describe('SpendingPolicyService', () => { + let repository: MockRepository; + let service: SpendingPolicyService; + + beforeEach(() => { + repository = makeRepository(); + service = new SpendingPolicyService(repository as unknown as SpendingPolicyRepository); + }); + + describe('create', () => { + it('validates the configuration and connects the organization and agent', async () => { + await service.create('org-1', { ...CREATE_INPUT, agentId: 'agent-1' }); + + expect(repository.create).toHaveBeenCalledWith({ + organization: { connect: { id: 'org-1' } }, + agent: { connect: { id: 'agent-1' } }, + name: 'Daily limit', + description: undefined, + type: PolicyType.MAX_AMOUNT, + configuration: { dailyLimit: 100 }, + priority: 100, + enabled: true, + }); + }); + + it('omits the agent connection when no agent is targeted', async () => { + await service.create('org-1', CREATE_INPUT); + + const data = repository.create.mock.calls[0][0]; + expect(data).not.toHaveProperty('agent'); + }); + + it('rejects an invalid spending configuration before touching the database', async () => { + await expect( + service.create('org-1', { + ...CREATE_INPUT, + configuration: { dailyLimit: -1 } as CreatePolicyInput['configuration'], + }), + ).rejects.toBeInstanceOf(ValidationException); + + expect(repository.create).not.toHaveBeenCalled(); + }); + }); + + describe('list and retrieval', () => { + it('returns a paginated result built from the repository transaction', async () => { + repository.findManyAndCount.mockResolvedValue({ items: [{ id: 'policy-1' }], total: 1 }); + + const result = await service.list('org-1', { + page: 1, + limit: 20, + search: 'limit', + } as PaginationQuery); + + expect(repository.findManyAndCount).toHaveBeenCalledWith( + { organizationId: 'org-1', deletedAt: null, name: { contains: 'limit', mode: 'insensitive' } }, + expect.objectContaining({ take: 20 }), + ); + expect(result.items).toEqual([{ id: 'policy-1' }]); + expect(result.meta.total).toBe(1); + }); + + it('throws NotFound when the policy does not belong to the organization', async () => { + repository.findById.mockResolvedValue(null); + + await expect(service.getOrThrow('org-1', 'policy-9')).rejects.toBeInstanceOf( + NotFoundException, + ); + }); + }); + + describe('update and remove', () => { + it('validates the configuration only when one is supplied', async () => { + const input: UpdatePolicyInput = { configuration: { dailyLimit: 250 } }; + + await service.update('org-1', 'policy-1', input); + + expect(repository.update).toHaveBeenCalledWith('policy-1', { + name: undefined, + description: undefined, + type: undefined, + priority: undefined, + enabled: undefined, + configuration: { dailyLimit: 250 }, + }); + }); + + it('refuses to update a policy outside the organization', async () => { + repository.findById.mockResolvedValue(null); + + await expect( + service.update('org-1', 'policy-9', { name: 'x' }), + ).rejects.toBeInstanceOf(NotFoundException); + expect(repository.update).not.toHaveBeenCalled(); + }); + + it('soft-deletes after verifying ownership', async () => { + await expect(service.remove('org-1', 'policy-1')).resolves.toEqual({ + id: 'policy-1', + deleted: true, + }); + expect(repository.softDelete).toHaveBeenCalledWith('policy-1'); + }); + }); + + describe('velocity limit', () => { + it('rejects spend that would exceed the rolling daily limit', async () => { + repository.findActiveForEvaluationByAgent.mockResolvedValue([policyRow({ dailyLimit: 100 })]); + repository.sumSpentInWindow.mockResolvedValue(80); + + await expect(service.checkVelocityLimit('agent-1', 30, 'USDC')).rejects.toBeInstanceOf( + VelocityLimitExceededException, + ); + + expect(repository.sumSpentInWindow).toHaveBeenCalledWith({ + agentId: 'agent-1', + assetCode: 'USDC', + since: expect.any(Date), + }); + }); + + it('allows spend that stays within the limit', async () => { + repository.findActiveForEvaluationByAgent.mockResolvedValue([policyRow({ dailyLimit: 100 })]); + repository.sumSpentInWindow.mockResolvedValue(10); + + await expect(service.checkVelocityLimit('agent-1', 20, 'USDC')).resolves.toBeUndefined(); + }); + + it('is a no-op for agents without a daily-limit policy and never sums history', async () => { + repository.findActiveForEvaluationByAgent.mockResolvedValue([policyRow({ maxAmount: 50 })]); + + await expect(service.checkVelocityLimit('agent-1', 1_000, 'USDC')).resolves.toBeUndefined(); + expect(repository.sumSpentInWindow).not.toHaveBeenCalled(); + }); + }); + + describe('evaluation audit', () => { + const intent = { + organizationId: 'org-1', + agentId: 'agent-1', + asset: 'USDC', + amount: 10, + recipientAddress: 'GABC', + }; + const result = { + passed: false, + requiresApproval: false, + violations: [], + evaluatedPolicyIds: ['policy-1'], + matchedPolicyId: 'policy-1', + }; + + it('appends the evaluation outcome to the audit trail', async () => { + await service.recordEvaluationAudit(intent, result, 'user-1'); + + expect(repository.recordEvaluationAudit).toHaveBeenCalledWith({ + organizationId: 'org-1', + userId: 'user-1', + policyId: 'policy-1', + payload: expect.objectContaining({ passed: false, transactionIntent: intent }), + }); + }); + + it('swallows repository failures so the payment pipeline is never blocked', async () => { + const logger = vi.spyOn(Logger.prototype, 'error').mockImplementation(() => undefined); + repository.recordEvaluationAudit.mockRejectedValue(new Error('audit table down')); + + await expect( + service.recordEvaluationAudit(intent, result, 'user-1'), + ).resolves.toBeUndefined(); + expect(logger).toHaveBeenCalledWith( + expect.stringContaining('Failed to persist policy evaluation audit log'), + ); + logger.mockRestore(); + }); + }); +}); diff --git a/src/modules/policies/spending-policy.service.ts b/src/modules/policies/spending-policy.service.ts new file mode 100644 index 00000000..5324c84a --- /dev/null +++ b/src/modules/policies/spending-policy.service.ts @@ -0,0 +1,198 @@ +import { Injectable, Logger } from '@nestjs/common'; +import { Policy, Prisma } from '@prisma/client'; + +import { + NotFoundException, + ValidationException, + VelocityLimitExceededException, +} from '../../common/exceptions/domain.exception'; +import { + buildPaginationMeta, + PaginationQuery, + toPrismaPagination, +} from '../../common/helpers/pagination'; +import { Paginated } from '../../common/interfaces/api-response.interface'; +import { formatZodError } from '../../common/validators/zod-error'; +import { CreatePolicyInput, UpdatePolicyInput } from './policy.dto'; +import { + PolicyConfiguration, + PolicyEvaluationResult, + TransactionIntent, + policyConfigurationSchemaStrict, +} from './policy.types'; +import { SpendingPolicyRepository } from './spending-policy.repository'; + +/** Fields a caller may sort the policy list by. */ +const SORTABLE = ['createdAt', 'priority', 'name', 'type']; + +/** Rolling window used by the velocity (drain-prevention) check. */ +const VELOCITY_WINDOW_MS = 24 * 60 * 60 * 1000; + +/** + * Owns the agent spending-policy domain: validation, persistence orchestration + * and the enforcement helpers the transaction pipeline depends on. + * + * The service never touches Prisma — every read and write goes through + * {@link SpendingPolicyRepository}, which keeps the persistence layer swappable + * and makes this class fully unit-testable with a mocked repository. + */ +@Injectable() +export class SpendingPolicyService { + private readonly logger = new Logger(SpendingPolicyService.name); + + constructor(private readonly repository: SpendingPolicyRepository) {} + + /** Validates and persists a new spending policy. */ + async create(organizationId: string, input: CreatePolicyInput): Promise { + const configuration = this.validateConfiguration(input.configuration); + + return this.repository.create({ + organization: { connect: { id: organizationId } }, + ...(input.agentId ? { agent: { connect: { id: input.agentId } } } : {}), + name: input.name, + description: input.description, + type: input.type, + configuration, + priority: input.priority, + enabled: input.enabled, + }); + } + + /** Paginated policy list for one organization. */ + async list(organizationId: string, query: PaginationQuery): Promise> { + const where: Prisma.PolicyWhereInput = { organizationId, deletedAt: null }; + if (query.search) { + where.name = { contains: query.search, mode: 'insensitive' }; + } + const pagination = toPrismaPagination(query, SORTABLE); + const { items, total } = await this.repository.findManyAndCount(where, pagination); + return new Paginated(items, buildPaginationMeta(total, query)); + } + + /** Returns a policy or throws a 404 when it does not exist in the organization. */ + async getOrThrow(organizationId: string, id: string): Promise { + const policy = await this.repository.findById(organizationId, id); + if (!policy) { + throw new NotFoundException('Policy', id); + } + return policy; + } + + /** Validates and applies a partial update to an existing policy. */ + async update( + organizationId: string, + id: string, + input: UpdatePolicyInput, + ): Promise { + await this.getOrThrow(organizationId, id); + + const data: Prisma.PolicyUpdateInput = { + name: input.name, + description: input.description, + type: input.type, + priority: input.priority, + enabled: input.enabled, + }; + if (input.configuration) { + data.configuration = this.validateConfiguration(input.configuration); + } + + return this.repository.update(id, data); + } + + /** Soft-deletes a policy after confirming it belongs to the organization. */ + async remove(organizationId: string, id: string): Promise<{ id: string; deleted: true }> { + await this.getOrThrow(organizationId, id); + await this.repository.softDelete(id); + return { id, deleted: true }; + } + + /** Enabled policies that apply to an organization and, optionally, an agent. */ + listActiveForEvaluation(organizationId: string, agentId?: string) { + return this.repository.findActiveForEvaluation(organizationId, agentId); + } + + /** Validates a policy configuration against the strict spending-policy schema. */ + private validateConfiguration( + configuration: CreatePolicyInput['configuration'], + ): Prisma.InputJsonValue { + const validationResult = policyConfigurationSchemaStrict.safeParse(configuration); + if (!validationResult.success) { + throw new ValidationException( + 'Invalid policy configuration', + formatZodError(validationResult.error), + ); + } + return validationResult.data as Prisma.InputJsonValue; + } + + /** + * Enforces the rolling 24-hour velocity limit for an agent: the spend already + * settled in the window plus the pending amount must stay within the agent's + * configured `dailyLimit`. Acts as a circuit breaker against rapid wallet + * draining. Agents without a daily-limit policy are unlimited. + */ + async checkVelocityLimit(agentId: string, amount: number, assetCode: string): Promise { + const limitPolicy = await this.findDailyLimitPolicy(agentId); + if (!limitPolicy) { + return; + } + const dailyLimit = limitPolicy.configuration.dailyLimit!; + + const spentInWindow = await this.repository.sumSpentInWindow({ + agentId, + assetCode, + since: new Date(Date.now() - VELOCITY_WINDOW_MS), + }); + + if (spentInWindow + amount > dailyLimit) { + throw new VelocityLimitExceededException( + `Daily velocity limit exceeded. Spent: ${spentInWindow}, Pending: ${amount}, Limit: ${dailyLimit}`, + { spentInWindow, pendingAmount: amount, limit: dailyLimit, assetCode }, + ); + } + } + + /** Highest-priority agent policy that declares a positive `dailyLimit`, if any. */ + private async findDailyLimitPolicy( + agentId: string, + ): Promise<{ configuration: PolicyConfiguration } | undefined> { + const policies = await this.repository.findActiveForEvaluationByAgent(agentId); + const limitPolicy = policies.find((policy) => { + const configuration = (policy.configuration as PolicyConfiguration) ?? {}; + return configuration.dailyLimit !== undefined && configuration.dailyLimit > 0; + }); + return limitPolicy + ? { configuration: (limitPolicy.configuration as PolicyConfiguration) ?? {} } + : undefined; + } + + /** + * Appends the outcome of a policy evaluation to the audit trail. Compliance + * bookkeeping must never block the payment pipeline, so failures are logged + * and swallowed. + */ + async recordEvaluationAudit( + intent: TransactionIntent, + result: PolicyEvaluationResult, + actorId: string, + ): Promise { + try { + await this.repository.recordEvaluationAudit({ + organizationId: intent.organizationId, + userId: actorId, + policyId: result.matchedPolicyId ?? null, + payload: { + passed: result.passed, + requiresApproval: result.requiresApproval, + violations: result.violations, + transactionIntent: intent, + } as unknown as Prisma.InputJsonValue, + }); + } catch (error) { + this.logger.error( + `Failed to persist policy evaluation audit log: ${(error as Error).message}`, + ); + } + } +} diff --git a/src/modules/risk/risk.service.spec.ts b/src/modules/risk/risk.service.spec.ts index 0ccd7b07..fb94d1cb 100644 --- a/src/modules/risk/risk.service.spec.ts +++ b/src/modules/risk/risk.service.spec.ts @@ -1,10 +1,11 @@ -import { describe, expect, it, vi } from 'vitest'; +import { describe, it, expect, beforeEach, vi } from 'vitest'; import { RiskBand } from '@prisma/client'; import { RiskService } from './risk.service'; import { RiskEngine } from './risk.engine'; -import { RiskFactorsInput } from './risk.types'; -import { EventBusService } from '../../events/event-bus.service'; import { RiskRepository } from './risk.repository'; +import { EventBusService } from '../../events/event-bus.service'; +import { DomainEventName } from '../../events/event-names'; +import { RiskFactorsInput } from './risk.types'; const lowRisk: RiskFactorsInput = { amount: 20, @@ -17,13 +18,17 @@ const lowRisk: RiskFactorsInput = { }; function createEventBus() { - return { emit: vi.fn().mockResolvedValue(undefined) } as unknown as Pick & { emit: ReturnType }; + return { + emit: vi.fn().mockResolvedValue(undefined), + } as unknown as Pick & { emit: ReturnType }; } describe('RiskService', () => { - it('emits a RiskEvaluated event with full factor breakdown', async () => { + it('emits a RiskEvaluated event with the factor breakdown', async () => { const eventBus = createEventBus(); - const repository = { createAssessmentRecord: vi.fn().mockResolvedValue(undefined) } as unknown as RiskRepository; + const repository = { + createAssessmentRecord: vi.fn().mockResolvedValue(undefined), + } as unknown as RiskRepository; const service = new RiskService(new RiskEngine(), eventBus as unknown as EventBusService, repository); const assessment = await service.evaluate('org-1', lowRisk, { @@ -31,41 +36,143 @@ describe('RiskService', () => { actorId: 'agent-1', }); - expect(assessment.band).toBe(RiskBand.LOW); - expect(assessment.factors.length).toBe(6); - - const emitMock = eventBus.emit as ReturnType; - expect(emitMock).toHaveBeenCalledOnce(); - const [eventName, payload] = emitMock.mock.calls[0]; - expect(eventName).toBe('risk.evaluated'); - expect(payload.transactionId).toBe('tx-1'); - expect(payload.score).toBe(assessment.score); - expect(payload.band).toBe(RiskBand.LOW); - expect(payload.factors).toEqual(assessment.factors); - expect(payload.canAutoExecute).toBe(true); + const [eventName, payload] = (eventBus.emit as ReturnType).mock.calls[0]; + expect(eventBus.emit).toHaveBeenCalledOnce(); + expect(eventName).toBe(DomainEventName.RiskEvaluated); + expect(payload).toMatchObject({ + transactionId: 'tx-1', + score: assessment.score, + band: RiskBand.LOW, + factors: assessment.factors, + canAutoExecute: true, + }); }); - it('assess() returns a result without emitting events', async () => { + it('assess() returns a result without emitting events', () => { const eventBus = createEventBus(); - const repository = { createAssessmentRecord: vi.fn().mockResolvedValue(undefined) } as unknown as RiskRepository; + const repository = { + createAssessmentRecord: vi.fn().mockResolvedValue(undefined), + } as unknown as RiskRepository; const service = new RiskService(new RiskEngine(), eventBus as unknown as EventBusService, repository); - const assessment = service.assess(lowRisk); - expect(assessment.band).toBe(RiskBand.LOW); - const emitMock = eventBus.emit as ReturnType; - expect(emitMock).not.toHaveBeenCalled(); + expect(service.assess(lowRisk).band).toBe(RiskBand.LOW); + expect(eventBus.emit).not.toHaveBeenCalled(); }); - it('passes config overrides through to the engine', async () => { + it('passes config overrides through to the engine', () => { const eventBus = createEventBus(); - const repository = { createAssessmentRecord: vi.fn().mockResolvedValue(undefined) } as unknown as RiskRepository; + const repository = { + createAssessmentRecord: vi.fn().mockResolvedValue(undefined), + } as unknown as RiskRepository; const service = new RiskService(new RiskEngine(), eventBus as unknown as EventBusService, repository); - const assessment = service.assess( - { ...lowRisk, amount: 100 }, - { amountSaturation: 100 }, + const assessment = service.assess({ ...lowRisk, amount: 100 }, { amountSaturation: 100 }); + const amountFactor = assessment.factors.find((factor) => factor.factor === 'amount'); + expect(amountFactor?.contribution).toBe(30); + }); +}); + +describe('RiskService Event Handler', () => { + let riskService: RiskService; + let riskEngine: RiskEngine; + let riskRepository: RiskRepository; + let eventBusService: EventBusService; + + beforeEach(() => { + riskEngine = new RiskEngine(); + riskRepository = { + createAssessmentRecord: vi.fn().mockResolvedValue({ id: 'assessment-1' }), + findByOrganization: vi.fn().mockResolvedValue([]), + findByTransaction: vi.fn().mockResolvedValue(null), + } as unknown as RiskRepository; + + eventBusService = { + emit: vi.fn().mockResolvedValue(undefined), + } as unknown as EventBusService; + + riskService = new RiskService(riskEngine, eventBusService, riskRepository); + }); + + it('should evaluate and persist risk assessment upon handling transaction created event', async () => { + const envelope = { + eventId: 'event-123', + name: DomainEventName.TransactionCreated, + organizationId: 'org-1', + aggregateType: 'transaction', + aggregateId: 'tx-123', + actorId: 'agent-1', + payload: { + transactionId: 'tx-123', + walletId: 'wallet-1', + amount: '150.0', + asset: 'XLM', + }, + occurredAt: new Date(), + }; + + await riskService.handleTransactionCreated(envelope); + + expect(eventBusService.emit).toHaveBeenCalledWith( + DomainEventName.RiskEvaluated, + expect.objectContaining({ + transactionId: 'tx-123', + }), + expect.objectContaining({ + organizationId: 'org-1', + actorId: 'agent-1', + aggregateType: 'transaction', + aggregateId: 'tx-123', + }), + ); + + expect(riskRepository.createAssessmentRecord).toHaveBeenCalledWith( + expect.objectContaining({ + organizationId: 'org-1', + transactionId: 'tx-123', + }), + ); + }); + + it('should deduplicate concurrent or repeated event deliveries', async () => { + const timestamp = new Date(); + const envelope = { + eventId: 'event-duplicate', + name: DomainEventName.TransactionCreated, + organizationId: 'org-1', + aggregateType: 'transaction', + aggregateId: 'tx-dup', + payload: { + transactionId: 'tx-dup', + amount: '50.0', + }, + occurredAt: timestamp, + }; + + await riskService.handleTransactionCreated(envelope); + await riskService.handleTransactionCreated(envelope); + + expect(riskRepository.createAssessmentRecord).toHaveBeenCalledTimes(1); + }); + + it('should handle failure resilience gracefully when evaluation throws', async () => { + vi.spyOn(riskRepository, 'createAssessmentRecord').mockRejectedValueOnce( + new Error('DB connection failed'), + ); + const envelope = { + eventId: 'event-failure', + name: DomainEventName.TransactionCreated, + organizationId: 'org-1', + aggregateType: 'transaction', + aggregateId: 'tx-err', + payload: { + transactionId: 'tx-err', + amount: '100.0', + }, + occurredAt: new Date(), + }; + + await expect(riskService.handleTransactionCreated(envelope)).rejects.toThrow( + 'DB connection failed', ); - const amountFactor = assessment.factors.find((f) => f.factor === 'amount'); - expect(amountFactor!.contribution).toBe(30); }); }); diff --git a/src/modules/risk/risk.service.ts b/src/modules/risk/risk.service.ts index 68585b81..331300b1 100644 --- a/src/modules/risk/risk.service.ts +++ b/src/modules/risk/risk.service.ts @@ -1,9 +1,11 @@ -import { Injectable } from '@nestjs/common'; +import { Injectable, Logger } from '@nestjs/common'; import { RiskEngine } from './risk.engine'; import { RiskAssessment, RiskConfig, RiskFactorsInput, RiskRule } from './risk.types'; import { EventBusService } from '../../events/event-bus.service'; import { DomainEventName } from '../../events/event-names'; import { RiskRepository } from './risk.repository'; +import { TypedOnEvent } from '../../events/typed-event-listener.decorator'; +import { DomainEventEnvelope } from '../../events/domain-event.types'; /** * Application-facing risk service. Wraps the pure {@link RiskEngine}, emits a @@ -12,6 +14,9 @@ import { RiskRepository } from './risk.repository'; */ @Injectable() export class RiskService { + private readonly logger = new Logger(RiskService.name); + private readonly processedEvents = new Set(); + constructor( private readonly engine: RiskEngine, private readonly eventBus: EventBusService, @@ -26,7 +31,12 @@ export class RiskService { async evaluate( organizationId: string, input: RiskFactorsInput, - context: { transactionId?: string; actorId?: string; config?: Partial; rules?: RiskRule[] } = {}, + context: { + transactionId?: string; + actorId?: string; + config?: Partial; + rules?: RiskRule[]; + } = {}, ): Promise { const assessment = this.engine.assess(input, context.config, context.rules); @@ -74,6 +84,62 @@ export class RiskService { return this.repository.findByOrganization(organizationId, limit); } + @TypedOnEvent(DomainEventName.TransactionCreated) + async handleTransactionCreated( + envelope: DomainEventEnvelope<{ + transactionId: string; + walletId?: string; + amount?: string; + asset?: string; + }>, + ): Promise { + const transactionId = envelope.payload?.transactionId; + if (!transactionId) { + return; + } + + const dedupKey = `${transactionId}:${envelope.occurredAt?.getTime() || 0}`; + if (this.processedEvents.has(dedupKey)) { + this.logger.debug( + `Duplicate transaction created event detected for transaction ${transactionId}, skipping.`, + ); + return; + } + this.processedEvents.add(dedupKey); + if (this.processedEvents.size > 5000) { + const firstKey = this.processedEvents.values().next().value; + if (firstKey) { + this.processedEvents.delete(firstKey); + } + } + + const organizationId = envelope.organizationId || 'default-org'; + try { + const amountNum = envelope.payload?.amount ? parseFloat(envelope.payload.amount) : 0; + const riskInput: RiskFactorsInput = { + amount: amountNum, + asset: envelope.payload?.asset ?? 'XLM', + knownRecipient: false, + recentTransactionCount: 1, + walletAgeDays: 0, + policyViolations: 0, + }; + + await this.evaluate(organizationId, riskInput, { + transactionId, + actorId: envelope.actorId, + }); + this.logger.log( + `Successfully scored risk for transaction ${transactionId} via event handler.`, + ); + } catch (error) { + this.logger.error( + `Failed to handle risk scoring for transaction ${transactionId}: ${error instanceof Error ? error.message : String(error)}`, + ); + throw error; + } + } + async getStatistics(organizationId: string, days = 30) { return this.repository.getStatistics(organizationId, days); } diff --git a/src/modules/stellar/services/stellar.service.ts b/src/modules/stellar/services/stellar.service.ts new file mode 100644 index 00000000..74235d80 --- /dev/null +++ b/src/modules/stellar/services/stellar.service.ts @@ -0,0 +1,127 @@ +import { Inject, Injectable, Logger } from '@nestjs/common'; +import { ErrorCode } from '../../../common/constants/error-codes'; +import { DomainException } from '../../../common/exceptions/domain.exception'; +import { CircuitBreaker, isRpcFailure } from '../../../common/circuit-breaker/circuit-breaker'; +import { + BuildPaymentParams, + StellarBalance, + StellarClient, + StellarKeypair, + StellarNetworkName, + StellarSubmitResult, + StellarTransactionInfo, + SubmitPaymentParams, + STELLAR_CLIENT, + SOROBAN_CLIENT, + SorobanClient, + SorobanSimulationResult, +} from '../../../integrations/stellar'; + +const HORIZON_FAILURE_THRESHOLD = 5; +const HORIZON_RESET_TIMEOUT_MS = 30_000; + +@Injectable() +export class StellarService { + private readonly logger = new Logger(StellarService.name); + private readonly breaker = new CircuitBreaker({ + name: 'horizon', + failureThreshold: HORIZON_FAILURE_THRESHOLD, + resetTimeoutMs: HORIZON_RESET_TIMEOUT_MS, + isFailure: isRpcFailure, + }); + + constructor( + @Inject(STELLAR_CLIENT) private readonly client: StellarClient, + @Inject(SOROBAN_CLIENT) private readonly sorobanClient: SorobanClient, + ) {} + + generateKeypair(): StellarKeypair { + return this.client.generateKeypair(); + } + + assertValidAddress(address: string): void { + if (!this.client.isValidAddress(address)) { + throw new DomainException( + ErrorCode.INVALID_STELLAR_ADDRESS, + `'${address}' is not a valid Stellar address`, + ); + } + } + + isValidAddress(address: string): boolean { + return this.client.isValidAddress(address); + } + + async getBalances(address: string, network: StellarNetworkName): Promise { + return this.wrap(() => this.client.getBalances(address, network)); + } + + async getNativeBalance(address: string, network: StellarNetworkName): Promise { + return this.wrap(() => this.client.getNativeBalance(address, network)); + } + + async buildPaymentXdr(params: BuildPaymentParams): Promise { + return this.wrap(() => this.client.buildPaymentXdr(params)); + } + + async submitPayment(params: SubmitPaymentParams): Promise { + return this.wrap(() => this.client.submitPayment(params)); + } + + async getTransactionInfo(txHash: string, network: StellarNetworkName): Promise { + return this.wrap(async () => { + const info = await this.client.getTransaction(txHash, network); + if (!info) { + throw new DomainException(ErrorCode.NOT_FOUND, `Transaction '${txHash}' not found`); + } + return info; + }); + } + + async simulateTransaction(transactionXdr: string): Promise { + if (!transactionXdr || typeof transactionXdr !== 'string') { + throw new DomainException( + ErrorCode.INVALID_STELLAR_TRANSACTION, + 'Invalid or malformed transaction XDR string', + ); + } + + try { + return await this.breaker.execute(async () => { + const result = await this.sorobanClient.simulateTransaction({ transactionXdr }); + if (!result.success || result.error) { + throw new DomainException( + ErrorCode.STELLAR_ERROR, + `Simulation failed: ${result.error?.message ?? 'Unknown simulation error'}`, + ); + } + return result; + }); + } catch (error: unknown) { + if (error instanceof DomainException) { + throw error; + } + const errMessage = error instanceof Error ? error.message : 'Unknown simulation error'; + this.logger.error( + `Stellar transaction simulation failed: ${errMessage}`, + error instanceof Error ? error.stack : undefined, + ); + throw new DomainException( + ErrorCode.STELLAR_ERROR, + `Failed to simulate Stellar transaction: ${errMessage}`, + ); + } + } + + private async wrap(fn: () => Promise): Promise { + try { + return await this.breaker.execute(fn); + } catch (error: unknown) { + if (error instanceof DomainException) { + throw error; + } + const message = error instanceof Error ? error.message : 'Unknown Stellar error'; + throw new DomainException(ErrorCode.STELLAR_ERROR, `Stellar operation failed: ${message}`); + } + } +} diff --git a/src/modules/stellar/stellar.module.ts b/src/modules/stellar/stellar.module.ts index e84e6e48..e03231ad 100644 --- a/src/modules/stellar/stellar.module.ts +++ b/src/modules/stellar/stellar.module.ts @@ -52,6 +52,11 @@ import { StellarController } from './stellar.controller'; StellarService, StellarTransactionService, ], - exports: [StellarService, StellarTransactionService, HorizonCircuitBreakerService], + exports: [ + SOROBAN_CLIENT, + StellarService, + StellarTransactionService, + HorizonCircuitBreakerService, + ], }) export class StellarModule {} diff --git a/src/modules/stellar/tests/stellar.service.spec.ts b/src/modules/stellar/tests/stellar.service.spec.ts new file mode 100644 index 00000000..a48b9c12 --- /dev/null +++ b/src/modules/stellar/tests/stellar.service.spec.ts @@ -0,0 +1,109 @@ +import { Test, TestingModule } from '@nestjs/testing'; +import { describe, it, expect, beforeEach, vi } from 'vitest'; +import { StellarService } from '../services/stellar.service'; +import { + STELLAR_CLIENT, + SOROBAN_CLIENT, + StellarClient, + SorobanClient, + SorobanSimulationResult, +} from '../../../integrations/stellar'; +import { DomainException } from '../../../common/exceptions/domain.exception'; +import { ErrorCode } from '../../../common/constants/error-codes'; + +describe('StellarService - Transaction Simulation', () => { + let service: StellarService; + let mockSorobanClient: SorobanClient; + let mockStellarClient: StellarClient; + + beforeEach(async () => { + mockSorobanClient = { + simulateTransaction: vi.fn(), + } as unknown as SorobanClient; + + mockStellarClient = { + generateKeypair: vi.fn(), + isValidAddress: vi.fn().mockReturnValue(true), + getBalances: vi.fn(), + getNativeBalance: vi.fn(), + buildPaymentXdr: vi.fn(), + submitPayment: vi.fn(), + getTransactionInfo: vi.fn(), + } as unknown as StellarClient; + + const module: TestingModule = await Test.createTestingModule({ + providers: [ + StellarService, + { + provide: STELLAR_CLIENT, + useValue: mockStellarClient, + }, + { + provide: SOROBAN_CLIENT, + useValue: mockSorobanClient, + }, + ], + }).compile(); + + service = module.get(StellarService); + }); + + it('should successfully simulate a valid transaction XDR', async () => { + const mockResult: SorobanSimulationResult = { + success: true, + minResourceFee: '100', + cost: { cpuInstructions: 0, memoryBytes: 0 }, + footprint: { readOnly: [], readWrite: [] }, + events: [], + result: 'AAAA...', + }; + vi.spyOn(mockSorobanClient, 'simulateTransaction').mockResolvedValueOnce(mockResult); + + const result = await service.simulateTransaction('AAAA...valid_xdr'); + expect(result).toEqual(mockResult); + expect(mockSorobanClient.simulateTransaction).toHaveBeenCalledWith({ + transactionXdr: 'AAAA...valid_xdr', + }); + }); + + it('should throw DomainException when transaction XDR is empty or invalid', async () => { + try { + await service.simulateTransaction(''); + expect.unreachable('expected simulateTransaction to throw'); + } catch (e: unknown) { + expect(e).toBeInstanceOf(DomainException); + const err = e as DomainException; + expect(err.code).toBe(ErrorCode.INVALID_STELLAR_TRANSACTION); + } + }); + + it('should handle simulation failure and Soroban error codes correctly', async () => { + const errorResult: SorobanSimulationResult = { + success: false, + minResourceFee: '0', + cost: { cpuInstructions: 0, memoryBytes: 0 }, + footprint: { readOnly: [], readWrite: [] }, + events: [], + error: { code: 'HOST_ERROR', message: 'HostError: Error(Contract, #4)' }, + }; + vi.spyOn(mockSorobanClient, 'simulateTransaction').mockResolvedValueOnce(errorResult); + + const error = await service.simulateTransaction('AAAA...trap_xdr').catch( + (reason: unknown) => reason as DomainException, + ); + expect(error).toBeInstanceOf(DomainException); + expect((error as DomainException).code).toBe(ErrorCode.STELLAR_ERROR); + expect((error as DomainException).message).toContain('HostError: Error(Contract, #4)'); + }); + + it('should handle RPC network timeouts and errors robustly', async () => { + vi.spyOn(mockSorobanClient, 'simulateTransaction').mockRejectedValueOnce(new Error('RPC timeout')); + + const error = await service.simulateTransaction('AAAA...timeout_xdr').catch( + (reason: unknown) => reason as DomainException, + ); + expect(error).toBeInstanceOf(DomainException); + expect((error as DomainException).code).toBe(ErrorCode.STELLAR_ERROR); + expect((error as DomainException).message).toContain('RPC timeout'); + }); +}); diff --git a/src/modules/transactions/guards/spending-limit.guard.ts b/src/modules/transactions/guards/spending-limit.guard.ts new file mode 100644 index 00000000..53aa795e --- /dev/null +++ b/src/modules/transactions/guards/spending-limit.guard.ts @@ -0,0 +1,115 @@ +import { CanActivate, ExecutionContext, Injectable, SetMetadata } from '@nestjs/common'; +import { Reflector } from '@nestjs/core'; +import { Request } from 'express'; +import { AuthenticatedUser } from '../../../common/interfaces/authenticated-user.interface'; +import { SpendingLimitService } from '../spending-limit.service'; +import { TransactionIntent } from '../../policies/policy.types'; + +export const SPENDING_LIMIT_GUARD_KEY = 'astroid:spendingLimitGuard'; + +/** + * Decorator that enables spending limit evaluation on a route. + * Apply to transaction creation endpoints that carry an optional `agentId`. + * + * ```typescript + * @Post() + * @UseGuards(SpendingLimitGuard) + * @RequireSpendingLimitCheck() + * create(...) { ... } + * ``` + */ +export const RequireSpendingLimitCheck = () => SetMetadata(SPENDING_LIMIT_GUARD_KEY, true); + +/** + * NestJS guard that intercepts transaction creation requests and evaluates them + * against the agent's configured spending-limit policies (daily/weekly/monthly + * budget caps) before the request reaches the service layer. + * + * Design decisions: + * - Only activated when the `@RequireSpendingLimitCheck()` decorator is + * present on the handler — routes without it pass straight through. + * - If no `agentId` is present in the request body the guard is a no-op, + * because periodic spending limits are scoped to agents. + * - Uses {@link SpendingLimitService.evaluateSpendingLimits} which runs + * aggregate queries inside a Prisma transaction to prevent race conditions + * when multiple concurrent requests target the same agent budget. + * - On violation: throws {@link PolicyViolationException} (HTTP 422) with a + * structured payload listing every violated policy. The global + * {@link AllExceptionsFilter} converts this to an RFC 9457 problem-details + * body so clients receive a consistent, machine-readable error shape: + * + * ```json + * { + * "type": "urn:astroid:problem:policy-violation", + * "title": "Policy Violation", + * "status": 422, + * "code": "POLICY_VIOLATION", + * "detail": "Transaction blocked by spending limit policy: ...", + * "details": { "violations": [...], "aggregates": {...} } + * } + * ``` + * + * Guard execution order (APP_GUARD chain + route guards): + * PublicRateLimitGuard → JwtAuthGuard → RolesGuard → ScopesGuard → + * AstroidThrottlerGuard → SpendingLimitGuard (route-level, via @UseGuards) + */ +@Injectable() +export class SpendingLimitGuard implements CanActivate { + constructor( + private readonly reflector: Reflector, + private readonly spendingLimitService: SpendingLimitService, + ) {} + + async canActivate(context: ExecutionContext): Promise { + // Only evaluate when the handler explicitly opts in via the decorator. + const enabled = this.reflector.getAllAndOverride(SPENDING_LIMIT_GUARD_KEY, [ + context.getHandler(), + context.getClass(), + ]); + if (!enabled) { + return true; + } + + const request = context + .switchToHttp() + .getRequest(); + + const body = request.body as Record | undefined; + const agentId = body?.agentId as string | undefined; + + // No agent attached to this transaction — periodic limits do not apply. + if (!agentId) { + return true; + } + + const organizationId = request.user?.organizationId; + if (!organizationId) { + // Guard can only evaluate when we know which org's policies to load. + // JwtAuthGuard runs before this so a missing org here means the route + // is @Public() and spending limits are not enforced. + return true; + } + + const actorId = request.user?.id; + const amount = Number(body?.amount ?? 0); + const asset = (body?.asset as string) ?? 'XLM'; + const recipientAddress = (body?.recipientAddress as string) ?? ''; + const walletId = (body?.walletId as string) ?? undefined; + + const intent: TransactionIntent = { + organizationId, + agentId, + walletId, + asset, + amount, + recipientAddress, + at: new Date(), + }; + + // evaluateSpendingLimits throws PolicyViolationException on failure and + // returns void on success — the guard returns true on success. + await this.spendingLimitService.evaluateSpendingLimits(intent, actorId); + + return true; + } +} diff --git a/src/modules/transactions/services/stellar-simulation.service.spec.ts b/src/modules/transactions/services/stellar-simulation.service.spec.ts new file mode 100644 index 00000000..5f3a36e5 --- /dev/null +++ b/src/modules/transactions/services/stellar-simulation.service.spec.ts @@ -0,0 +1,318 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { StellarSimulationService } from './stellar-simulation.service'; +import { SorobanClient, SorobanSimulationResult } from '../../../integrations/stellar/soroban.interface'; +import { RiskEngine } from '../../risk/risk.engine'; +import { EventBusService } from '../../../events/event-bus.service'; +import { CircuitOpenException, DomainException } from '../../../common/exceptions/domain.exception'; +import { ErrorCode } from '../../../common/constants/error-codes'; + +function buildMockSorobanClient(overrides: Partial = {}): SorobanClient { + return { + simulateTransaction: vi.fn().mockResolvedValue({ + success: true, + minResourceFee: '100000', + cost: { cpuInstructions: 200_000, memoryBytes: 4096 }, + footprint: { + readOnly: [{ contractId: 'contract-1', key: { symbol: 'Balance' } }], + readWrite: [], + }, + events: [], + result: undefined, + transactionHash: 'mock-hash-123', + ...overrides, + } as SorobanSimulationResult), + }; +} + +function buildMockEventBus() { + return { emit: vi.fn().mockResolvedValue(undefined) }; +} + +function buildValidXdr(): string { + return Buffer.from( + JSON.stringify({ source: 'GABC...', destination: 'GDEF...', amount: '100' }), + ).toString('base64'); +} + +describe('StellarSimulationService', () => { + let sorobanClient: ReturnType; + let eventBus: ReturnType; + let service: StellarSimulationService; + + beforeEach(() => { + vi.clearAllMocks(); + sorobanClient = buildMockSorobanClient(); + eventBus = buildMockEventBus(); + service = new StellarSimulationService( + sorobanClient, + new RiskEngine(), + eventBus as unknown as EventBusService, + ); + }); + + describe('simulate', () => { + it('should return simulation results with risk assessment', async () => { + const result = await service.simulate({ + transactionXdr: buildValidXdr(), + organizationId: 'org-1', + }); + + expect(result.success).toBe(true); + expect(result.feeEstimate).toBe('100000'); + expect(result.risk).toBeDefined(); + expect(result.risk.score).toBeGreaterThanOrEqual(0); + expect(result.risk.score).toBeLessThanOrEqual(100); + expect(result.transactionHash).toBe('mock-hash-123'); + }); + + it('should emit a risk evaluation event', async () => { + await service.simulate({ + transactionXdr: buildValidXdr(), + organizationId: 'org-1', + actorId: 'user-1', + }); + + expect(eventBus.emit).toHaveBeenCalledWith( + expect.any(String), + expect.objectContaining({ score: expect.any(Number) }), + expect.objectContaining({ organizationId: 'org-1', actorId: 'user-1' }), + ); + }); + + it('should throw DomainException for invalid base64 XDR', async () => { + await expect( + service.simulate({ + transactionXdr: '!!!not-base64-at-all&&&', + organizationId: 'org-1', + }), + ).rejects.toThrow('not valid base64'); + }); + + it('should throw DomainException for empty XDR', async () => { + await expect( + service.simulate({ + transactionXdr: '', + organizationId: 'org-1', + }), + ).rejects.toThrow('Transaction XDR is required'); + }); + + it('should throw DomainException when simulation returns failure', async () => { + (sorobanClient.simulateTransaction as ReturnType).mockResolvedValue({ + success: false, + minResourceFee: '0', + cost: { cpuInstructions: 0, memoryBytes: 0 }, + footprint: { readOnly: [], readWrite: [] }, + events: [], + error: { code: 'txFailed', message: 'Contract Error' }, + }); + + await expect( + service.simulate({ + transactionXdr: buildValidXdr(), + organizationId: 'org-1', + }), + ).rejects.toThrow('Contract Error'); + }); + + it('should throw DomainException when soroban client throws', async () => { + (sorobanClient.simulateTransaction as ReturnType).mockRejectedValue( + new Error('Connection refused'), + ); + + await expect( + service.simulate({ + transactionXdr: buildValidXdr(), + organizationId: 'org-1', + }), + ).rejects.toThrow('Stellar simulation failed'); + }); + + it('should throw RiskTooHighException when risk exceeds threshold', async () => { + await expect( + service.simulate({ + transactionXdr: buildValidXdr(), + organizationId: 'org-1', + maxRiskScore: -1, + }), + ).rejects.toThrow('exceeds maximum allowed'); + }); + + it('should include risk factors when provided', async () => { + const result = await service.simulate({ + transactionXdr: buildValidXdr(), + organizationId: 'org-1', + riskFactors: { + amount: 5000, + asset: 'USDC', + knownRecipient: true, + recentTransactionCount: 2, + walletAgeDays: 180, + policyViolations: 0, + }, + }); + + expect(result.risk.score).toBeGreaterThanOrEqual(0); + expect(result.risk.factors).toBeDefined(); + }); + + it('should indicate requiresApproval when risk is above LOW', async () => { + const result = await service.simulate({ + transactionXdr: buildValidXdr(), + organizationId: 'org-1', + maxRiskScore: 100, + riskFactors: { + amount: 50000, + asset: 'XLM', + knownRecipient: false, + recentTransactionCount: 15, + walletAgeDays: 5, + policyViolations: 2, + }, + }); + + expect(result.requiresApproval).toBe(true); + }); + }); + + describe('simulateWithDefaults', () => { + it('should call simulate with maxRiskScore of 80', async () => { + const simulateSpy = vi.spyOn(service, 'simulate'); + + const input = { + transactionXdr: buildValidXdr(), + organizationId: 'org-1', + }; + + await service.simulateWithDefaults(input); + + expect(simulateSpy).toHaveBeenCalledWith({ + ...input, + maxRiskScore: 80, + }); + }); + + it('should override maxRiskScore if already provided', async () => { + const simulateSpy = vi.spyOn(service, 'simulate'); + + const input = { + transactionXdr: buildValidXdr(), + organizationId: 'org-1', + maxRiskScore: 50, + }; + + await service.simulateWithDefaults(input); + + expect(simulateSpy).toHaveBeenCalledWith({ + ...input, + maxRiskScore: 80, + }); + }); + }); + + describe('validateXdrFormat', () => { + it('should return true for valid base64 XDR', () => { + const validXdr = buildValidXdr(); + expect(service.validateXdrFormat(validXdr)).toBe(true); + }); + + it('should return false for empty string', () => { + expect(service.validateXdrFormat('')).toBe(false); + }); + + it('should return false for null', () => { + expect(service.validateXdrFormat(null as unknown as string)).toBe(false); + }); + + it('should return false for undefined', () => { + expect(service.validateXdrFormat(undefined as unknown as string)).toBe(false); + }); + + it('should return false for invalid base64 characters', () => { + expect(service.validateXdrFormat('!!!not-base64!!!')).toBe(false); + }); + + it('should return false for malformed base64', () => { + expect(service.validateXdrFormat('AB=C')).toBe(false); + }); + + it('should return true for base64url format', () => { + const base64url = 'ABCdef-123_456=='; + expect(service.validateXdrFormat(base64url)).toBe(true); + }); + + it('should return true for standard base64', () => { + const standardBase64 = 'ABCdef+123/456=='; + expect(service.validateXdrFormat(standardBase64)).toBe(true); + }); + + it('should return true for base64 without padding', () => { + const noPadding = 'ABCdef123456'; + expect(service.validateXdrFormat(noPadding)).toBe(true); + }); + + it('should return false for non-string input', () => { + expect(service.validateXdrFormat(123 as unknown as string)).toBe(false); + expect(service.validateXdrFormat({} as unknown as string)).toBe(false); + expect(service.validateXdrFormat([] as unknown as string)).toBe(false); + }); + }); + + describe('circuit breaker integration', () => { + it('opens the Stellar circuit after repeated RPC failures and fails fast without calling the client again', async () => { + (sorobanClient.simulateTransaction as ReturnType).mockRejectedValue( + Object.assign(new Error('Stellar RPC unreachable'), { code: 'ECONNREFUSED' }), + ); + + for (let i = 0; i < 5; i++) { + await expect( + service.simulate({ transactionXdr: buildValidXdr(), organizationId: 'org-1' }), + ).rejects.toMatchObject({ code: ErrorCode.STELLAR_ERROR }); + } + expect(sorobanClient.simulateTransaction).toHaveBeenCalledTimes(5); + + (sorobanClient.simulateTransaction as ReturnType).mockClear(); + + let thrown: unknown; + try { + await service.simulate({ transactionXdr: buildValidXdr(), organizationId: 'org-1' }); + } catch (error) { + thrown = error; + } + + expect(thrown).toBeInstanceOf(CircuitOpenException); + expect(thrown).toBeInstanceOf(DomainException); + expect((thrown as DomainException).code).toBe(ErrorCode.CIRCUIT_OPEN); + expect(sorobanClient.simulateTransaction).not.toHaveBeenCalled(); + }); + }); + + describe('integration scenarios', () => { + it('should handle complete simulation workflow', async () => { + // Step 1: Validate XDR format + const xdr = buildValidXdr(); + expect(service.validateXdrFormat(xdr)).toBe(true); + + // Step 2: Simulate with defaults + const result = await service.simulateWithDefaults({ + transactionXdr: xdr, + organizationId: 'org-1', + actorId: 'user-1', + }); + + expect(result.success).toBe(true); + expect(result.risk.score).toBeLessThanOrEqual(80); + }); + + it('should handle validation before simulation to catch format errors early', async () => { + const invalidXdr = '!!!invalid-xdr!!!'; + + // Early validation + expect(service.validateXdrFormat(invalidXdr)).toBe(false); + + // Simulation would fail, but we caught it early + const validXdr = buildValidXdr(); + expect(service.validateXdrFormat(validXdr)).toBe(true); + }); + }); +}); diff --git a/src/modules/transactions/services/stellar-simulation.service.ts b/src/modules/transactions/services/stellar-simulation.service.ts new file mode 100644 index 00000000..600d3d41 --- /dev/null +++ b/src/modules/transactions/services/stellar-simulation.service.ts @@ -0,0 +1,271 @@ +import { Injectable, Logger, Inject } from '@nestjs/common'; +import { + SOROBAN_CLIENT, + SorobanClient, + SorobanSimulationResult, +} from '../../../integrations/stellar/soroban.interface'; +import { RiskEngine } from '../../risk/risk.engine'; +import { RiskAssessment, RiskFactorsInput } from '../../risk/risk.types'; +import { ErrorCode } from '../../../common/constants/error-codes'; +import { + DomainException, + RiskTooHighException, +} from '../../../common/exceptions/domain.exception'; +import { CircuitBreaker, isRpcFailure } from '../../../common/circuit-breaker/circuit-breaker'; +import { EventBusService } from '../../../events/event-bus.service'; +import { DomainEventName } from '../../../events/event-names'; + +/** Consecutive failures before the Soroban RPC circuit trips OPEN. */ +const SOROBAN_FAILURE_THRESHOLD = 5; +/** Time the Soroban RPC circuit stays OPEN before a HALF_OPEN trial call. */ +const SOROBAN_RESET_TIMEOUT_MS = 30_000; + +export interface SimulationInput { + /** The base64-encoded transaction envelope XDR. */ + transactionXdr: string; + /** Organization ID for risk context. */ + organizationId: string; + /** Optional actor ID for audit context. */ + actorId?: string; + /** Optional risk factors for scoring (if not provided, uses defaults). */ + riskFactors?: RiskFactorsInput; + /** Maximum allowed risk score before simulation is rejected. */ + maxRiskScore?: number; +} + +export interface SimulationOutput { + /** Whether the simulation succeeded. */ + success: boolean; + /** Fee estimate in stroops. */ + feeEstimate: string; + /** Resource cost analysis. */ + cost: { + cpuInstructions: number; + memoryBytes: number; + }; + /** Footprint data from the simulation. */ + footprint: SorobanSimulationResult['footprint']; + /** Events emitted during simulation. */ + events: SorobanSimulationResult['events']; + /** Risk assessment of the simulated transaction. */ + risk: RiskAssessment; + /** Whether the transaction requires approval based on risk. */ + requiresApproval: boolean; + /** Error details if simulation failed. */ + error?: SorobanSimulationResult['error']; + /** Transaction hash from simulation. */ + transactionHash?: string; +} + +/** + * Stellar Transaction Simulation Service. + * + * This service provides a unified interface for Stellar transaction simulation, + * handling both classic Stellar and Soroban smart contract transactions. + * It validates XDR, simulates execution on the Stellar network (via RPC), + * and returns detailed diagnostic information including fee estimates, + * resource costs, and risk assessment. + * + * This is a dedicated service that mirrors the functionality of SorobanSimulationService + * to provide a more generic Stellar simulation interface as requested in issue #246. + */ +@Injectable() +export class StellarSimulationService { + private readonly logger = new Logger(StellarSimulationService.name); + private readonly breaker = new CircuitBreaker({ + name: 'stellar-simulation', + failureThreshold: SOROBAN_FAILURE_THRESHOLD, + resetTimeoutMs: SOROBAN_RESET_TIMEOUT_MS, + isFailure: isRpcFailure, + }); + + constructor( + @Inject(SOROBAN_CLIENT) private readonly sorobanClient: SorobanClient, + private readonly riskEngine: RiskEngine, + private readonly eventBus: EventBusService, + ) {} + + /** + * Simulates a Stellar transaction before submission to the network. + * + * This method validates the transaction XDR, simulates execution on the + * Stellar network (via RPC), and returns detailed diagnostic information + * including fee estimates, resource costs, and risk assessment. + * + * @param input - Simulation parameters including transaction XDR and organization context + * @returns Simulation output with success status, fee estimate, cost analysis, and risk assessment + * @throws DomainException if the XDR is invalid or simulation fails + * @throws RiskTooHighException if the risk score exceeds the allowed threshold + */ + async simulate(input: SimulationInput): Promise { + this.logger.debug( + `Simulating Stellar transaction for organization ${input.organizationId}`, + ); + + this.validateXdr(input.transactionXdr); + + let result: SorobanSimulationResult; + try { + result = await this.breaker.execute(() => + this.sorobanClient.simulateTransaction({ + transactionXdr: input.transactionXdr, + }), + ); + } catch (error) { + if (error instanceof DomainException) { + throw error; + } + this.logger.warn( + `Stellar simulation failed: ${(error as Error).message}`, + ); + throw new DomainException( + ErrorCode.STELLAR_ERROR, + `Stellar simulation failed: ${(error as Error).message}`, + ); + } + + if (!result.success) { + this.logger.warn( + `Stellar simulation returned error: ${result.error?.code} - ${result.error?.message}`, + ); + throw new DomainException( + ErrorCode.STELLAR_ERROR, + `Stellar simulation failed: ${result.error?.message ?? 'Unknown error'}`, + result.error, + ); + } + + // Risk scoring + const riskInput = input.riskFactors ?? this.buildDefaultRiskFactors(result); + const risk = this.riskEngine.assess(riskInput); + const maxRiskScore = input.maxRiskScore ?? 80; + const requiresApproval = risk.score > 20 || risk.band !== 'LOW'; + + if (risk.score > maxRiskScore) { + throw new RiskTooHighException( + `Risk score ${risk.score} exceeds maximum allowed threshold of ${maxRiskScore}`, + { score: risk.score, band: risk.band, maxRiskScore }, + ); + } + + // Emit simulation event for telemetry + await this.eventBus.emit( + DomainEventName.RiskEvaluated, + { + score: risk.score, + band: risk.band, + simulationSuccess: true, + feeEstimate: result.minResourceFee, + transactionHash: result.transactionHash, + }, + { + organizationId: input.organizationId, + actorId: input.actorId, + aggregateType: 'transaction', + }, + ); + + this.logger.debug( + `Simulation completed: success=${result.success}, fee=${result.minResourceFee}, risk=${risk.score}`, + ); + + return { + success: true, + feeEstimate: result.minResourceFee, + cost: result.cost, + footprint: result.footprint, + events: result.events, + risk, + requiresApproval, + transactionHash: result.transactionHash, + }; + } + + /** + * Simulates a transaction with default risk thresholds. + * + * This is a convenience method that uses the system's default maximum + * risk score (80) for the simulation. + * + * @param input - Simulation parameters + * @returns Simulation output with diagnostic information + */ + async simulateWithDefaults(input: SimulationInput): Promise { + return this.simulate({ + ...input, + maxRiskScore: 80, + }); + } + + /** + * Validates transaction XDR format without executing simulation. + * + * This lightweight validation checks that the XDR is properly formatted + * base64-encoded data, which can be used for early client-side validation. + * + * @param transactionXdr - Base64-encoded transaction envelope XDR + * @returns true if the XDR format is valid, false otherwise + */ + validateXdrFormat(transactionXdr: string): boolean { + if (!transactionXdr || typeof transactionXdr !== 'string') { + return false; + } + + // Validate base64url/base64 format + if (!/^[A-Za-z0-9+/_-]*={0,2}$/.test(transactionXdr)) { + return false; + } + + try { + Buffer.from(transactionXdr, 'base64'); + return true; + } catch { + return false; + } + } + + private validateXdr(xdr: string): void { + if (!xdr || typeof xdr !== 'string') { + throw new DomainException( + ErrorCode.VALIDATION_ERROR, + 'Transaction XDR is required', + ); + } + // Validate base64url/base64 format + if (!/^[A-Za-z0-9+/_-]*={0,2}$/.test(xdr)) { + throw new DomainException( + ErrorCode.INVALID_STELLAR_TRANSACTION, + 'Transaction XDR is not valid base64', + ); + } + try { + Buffer.from(xdr, 'base64'); + } catch { + throw new DomainException( + ErrorCode.INVALID_STELLAR_TRANSACTION, + 'Transaction XDR is not valid base64', + ); + } + } + + private buildDefaultRiskFactors(result: SorobanSimulationResult): RiskFactorsInput { + // Build risk factors from simulation result when not explicitly provided + const eventCount = result.events.length; + const hasWriteFootprint = result.footprint.readWrite.length > 0; + const feeStroops = parseInt(result.minResourceFee, 10); + + // Heuristic: higher resource usage correlates with higher risk + const normalizedFee = Math.min(feeStroops / 10_000_000, 1); + const amountEstimate = normalizedFee * 10_000; + + return { + amount: amountEstimate, + asset: 'XLM', + knownRecipient: !hasWriteFootprint, + recentTransactionCount: eventCount, + walletAgeDays: 90, + policyViolations: 0, + hourUtc: new Date().getUTCHours(), + }; + } +} diff --git a/src/modules/transactions/spending-limit.service.ts b/src/modules/transactions/spending-limit.service.ts new file mode 100644 index 00000000..2bc84ae0 --- /dev/null +++ b/src/modules/transactions/spending-limit.service.ts @@ -0,0 +1,297 @@ +import { Injectable, Logger } from '@nestjs/common'; +import { Prisma, TransactionStatus } from '@prisma/client'; +import { PrismaService } from '../../database/prisma.service'; +import { PolicyService } from '../policies/policy.service'; +import { AuditService } from '../audit/audit.service'; +import { PolicyConfiguration, TransactionIntent } from '../policies/policy.types'; +import { PolicyViolationException } from '../../common/exceptions/domain.exception'; + +/** + * Daily/weekly/monthly spend windows, returned by `aggregateSpend`. + * All values are in the same asset unit as the transaction. + */ +export interface SpendAggregates { + spentToday: number; + spentThisWeek: number; + spentThisMonth: number; +} + +/** + * SpendingLimitService — evaluates agent spending limit policies with + * race-condition-safe aggregate queries. + * + * Responsibilities: + * 1. Query the agent's accumulated spend across daily/weekly/monthly UTC + * windows inside a single Prisma interactive transaction (serializable + * snapshot) so concurrent submissions cannot double-count. + * 2. Evaluate the enriched {@link TransactionIntent} (with real aggregates) + * against active policies via {@link PolicyService.evaluateIntent}. + * 3. On failure: persist a dedicated audit log entry before throwing + * {@link PolicyViolationException} so the compliance trail is complete + * even when the transaction is blocked. + * + * This service is intentionally narrow in scope — it does not replace + * {@link PolicyService} or {@link PolicyEngine}; it only enriches the intent + * with atomic spend data and delegates evaluation to the policy layer. + */ +@Injectable() +export class SpendingLimitService { + private readonly logger = new Logger(SpendingLimitService.name); + + constructor( + private readonly prisma: PrismaService, + private readonly policyService: PolicyService, + private readonly auditService: AuditService, + ) {} + + /** + * Computes the agent's confirmed spend aggregates for the three standard + * periods. All three windows are derived from UTC boundaries so they reset + * consistently regardless of the server's local timezone. + * + * The query runs inside a `READ COMMITTED` snapshot (Prisma default) which + * is sufficient because: + * - We read committed rows only (no phantom reads needed for a sum check). + * - The TransactionLockInterceptor already serialises requests on + * `transaction:{walletId}` at the application level, preventing two + * concurrent submissions from the same wallet from racing here. + * + * Only PENDING, SUBMITTED, CONFIRMED, and COMPLETED transactions count toward + * the aggregate — DRAFT/REJECTED/FAILED/CANCELLED/EXPIRED are excluded. + */ + async aggregateSpend(agentId: string, asset: string): Promise { + const now = new Date(); + + // UTC day boundary — midnight today + const startOfDay = new Date( + Date.UTC(now.getUTCFullYear(), now.getUTCMonth(), now.getUTCDate()), + ); + + // UTC week boundary — most-recent Monday at midnight + const dayOfWeek = now.getUTCDay(); // 0 = Sunday + const daysSinceMonday = dayOfWeek === 0 ? 6 : dayOfWeek - 1; + const startOfWeek = new Date( + Date.UTC( + now.getUTCFullYear(), + now.getUTCMonth(), + now.getUTCDate() - daysSinceMonday, + ), + ); + + // UTC month boundary — 1st of the current month + const startOfMonth = new Date( + Date.UTC(now.getUTCFullYear(), now.getUTCMonth(), 1), + ); + + const COUNTED_STATUSES: TransactionStatus[] = [ + TransactionStatus.PENDING, + TransactionStatus.SUBMITTED, + TransactionStatus.CONFIRMED, + TransactionStatus.COMPLETED, + ]; + + // Run all three aggregates in parallel within a single Prisma transaction + // to get a consistent snapshot. + const [dayResult, weekResult, monthResult] = await this.prisma.$transaction([ + this.prisma.transaction.aggregate({ + _sum: { amount: true }, + where: { + agentId, + asset, + status: { in: COUNTED_STATUSES }, + deletedAt: null, + createdAt: { gte: startOfDay }, + }, + }), + this.prisma.transaction.aggregate({ + _sum: { amount: true }, + where: { + agentId, + asset, + status: { in: COUNTED_STATUSES }, + deletedAt: null, + createdAt: { gte: startOfWeek }, + }, + }), + this.prisma.transaction.aggregate({ + _sum: { amount: true }, + where: { + agentId, + asset, + status: { in: COUNTED_STATUSES }, + deletedAt: null, + createdAt: { gte: startOfMonth }, + }, + }), + ]); + + return { + spentToday: dayResult._sum?.amount?.toNumber() ?? 0, + spentThisWeek: weekResult._sum?.amount?.toNumber() ?? 0, + spentThisMonth: monthResult._sum?.amount?.toNumber() ?? 0, + }; + } + + /** + * Returns true if the agent has at least one active spending-limit policy + * (a policy with `dailyLimit`, `weeklyLimit`, or `monthlyLimit` configured). + * Used to short-circuit the aggregate query when no periodic limits apply. + */ + async hasSpendingLimitPolicy( + organizationId: string, + agentId: string, + ): Promise { + const policies = await this.prisma.policy.findMany({ + where: { + organizationId, + enabled: true, + deletedAt: null, + OR: [{ agentId: null }, { agentId }], + }, + select: { configuration: true }, + }); + + return policies.some((p) => { + const config = p.configuration as PolicyConfiguration; + return ( + config.dailyLimit !== undefined || + config.weeklyLimit !== undefined || + config.monthlyLimit !== undefined + ); + }); + } + + /** + * Core entry-point called by the transaction pipeline and the + * {@link SpendingLimitGuard}. + * + * When called from {@link TransactionService.create}, the intent is + * already enriched with real spend aggregates (fetched once, passed in) so + * this method skips the aggregate query and goes directly to evaluation. + * When called from the guard, the intent has no aggregates yet, so this + * method fetches them first. + * + * Flow: + * 1. Short-circuit if no agent or no periodic limit policy is configured. + * 2. Fetch aggregates (only when not already present on the intent). + * 3. Enrich intent if aggregates were freshly fetched. + * 4. Evaluate via {@link PolicyService.evaluateIntent}. + * 5. On violation: write audit log, throw {@link PolicyViolationException}. + * + * @param intent Transaction intent, optionally pre-enriched with aggregates. + * @param actorId Authenticated user id for audit attribution. + */ + async evaluateSpendingLimits( + intent: TransactionIntent, + actorId?: string, + ): Promise { + const { agentId, organizationId, asset } = intent; + + if (!agentId) { + return; + } + + const hasLimits = await this.hasSpendingLimitPolicy(organizationId, agentId); + if (!hasLimits) { + return; + } + + // Use aggregates already embedded in the intent when the caller (e.g. + // TransactionService) pre-fetched them to avoid a redundant round-trip. + // The guard passes a bare intent so we fetch here in that case. + const alreadyEnriched = + intent.spentToday !== undefined && + intent.spentThisWeek !== undefined && + intent.spentThisMonth !== undefined; + + let enrichedIntent = intent; + let aggregates: SpendAggregates; + + if (alreadyEnriched) { + aggregates = { + spentToday: intent.spentToday!, + spentThisWeek: intent.spentThisWeek!, + spentThisMonth: intent.spentThisMonth!, + }; + } else { + aggregates = await this.aggregateSpend(agentId, asset); + enrichedIntent = { + ...intent, + spentToday: aggregates.spentToday, + spentThisWeek: aggregates.spentThisWeek, + spentThisMonth: aggregates.spentThisMonth, + }; + } + + const result = await this.policyService.evaluateIntent(enrichedIntent, actorId); + + if (!result.passed) { + if (actorId || organizationId) { + await this.persistViolationAuditLog( + organizationId, + actorId ?? null, + agentId, + enrichedIntent, + result.violations, + aggregates, + ); + } + + const violationMessages = result.violations + .map((v) => `${v.policyName}: ${v.message}`) + .join('; '); + + throw new PolicyViolationException( + `Transaction blocked by spending limit policy: ${violationMessages}`, + { + violations: result.violations, + requiresApproval: result.requiresApproval, + aggregates, + }, + ); + } + } + + // ── private helpers ────────────────────────────────────────────────────── + + /** + * Writes an audit log entry for a spending limit policy failure. + * Failures here are logged but never allowed to propagate — a broken audit + * write must never silently allow a blocked transaction through. + */ + private async persistViolationAuditLog( + organizationId: string, + actorId: string | null, + agentId: string, + intent: TransactionIntent, + violations: Array<{ policyId: string; policyName: string; code: string; message: string }>, + aggregates: SpendAggregates, + ): Promise { + try { + await this.auditService.record({ + organizationId, + userId: actorId, + action: 'SPENDING_LIMIT_EXCEEDED', + entity: 'transaction', + entityId: agentId, + oldValue: null as unknown as Prisma.InputJsonValue, + newValue: { + agentId, + asset: intent.asset, + amount: intent.amount, + recipientAddress: intent.recipientAddress, + violations, + aggregates, + evaluatedAt: new Date().toISOString(), + } as unknown as Prisma.InputJsonValue, + }); + } catch (auditError) { + // Audit failures must never surface to the caller as a 500 — they are + // background bookkeeping. Log and continue so the PolicyViolationException + // is what the caller sees. + this.logger.error( + `Failed to persist spending-limit violation audit log for agent ${agentId}: ${(auditError as Error).message}`, + ); + } + } +} diff --git a/src/modules/transactions/tests/spending-limit.guard.spec.ts b/src/modules/transactions/tests/spending-limit.guard.spec.ts new file mode 100644 index 00000000..56eb3630 --- /dev/null +++ b/src/modules/transactions/tests/spending-limit.guard.spec.ts @@ -0,0 +1,250 @@ +import { describe, it, expect, beforeEach, vi } from 'vitest'; +import { Test, TestingModule } from '@nestjs/testing'; +import { Reflector } from '@nestjs/core'; +import { ExecutionContext } from '@nestjs/common'; +import { SpendingLimitGuard, SPENDING_LIMIT_GUARD_KEY } from '../guards/spending-limit.guard'; +import { SpendingLimitService } from '../spending-limit.service'; +import { PolicyViolationException } from '../../../common/exceptions/domain.exception'; + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +const VALID_STELLAR = 'GDEGSXLGANKHK7QFOV63XCBHBTZ3YRKUJV7ZB7JMSJQB5CNBRLL5QIG5'; + +function makeContext(overrides: { + body?: Record; + user?: Record | null; + reflectorEnabled?: boolean; +}): ExecutionContext { + const { body = {}, user = { id: 'user-1', organizationId: 'org-1' }, reflectorEnabled = true } = overrides; + + const mockRequest = { body, user }; + + return { + switchToHttp: () => ({ + getRequest: () => mockRequest, + }), + getHandler: () => ({}), + getClass: () => ({}), + // Provide a reflector-compatible API via the context for our mock Reflector + _reflectorEnabled: reflectorEnabled, + } as unknown as ExecutionContext; +} + +// --------------------------------------------------------------------------- +// Mocks +// --------------------------------------------------------------------------- + +const mockSpendingLimitService = { + evaluateSpendingLimits: vi.fn(), +}; + +// --------------------------------------------------------------------------- +// Suite +// --------------------------------------------------------------------------- + +describe('SpendingLimitGuard', () => { + let guard: SpendingLimitGuard; + let reflector: Reflector; + + beforeEach(async () => { + vi.clearAllMocks(); + + const module: TestingModule = await Test.createTestingModule({ + providers: [ + SpendingLimitGuard, + Reflector, + { provide: SpendingLimitService, useValue: mockSpendingLimitService }, + ], + }).compile(); + + guard = module.get(SpendingLimitGuard); + reflector = module.get(Reflector); + }); + + // ── decorator opt-in ────────────────────────────────────────────────────── + + describe('when decorator is NOT present', () => { + it('returns true without calling SpendingLimitService', async () => { + vi.spyOn(reflector, 'getAllAndOverride').mockReturnValue(undefined); + + const ctx = makeContext({ body: { agentId: 'agent-1', amount: '100', asset: 'USDC', recipientAddress: VALID_STELLAR, walletId: 'wallet-1' } }); + const result = await guard.canActivate(ctx); + + expect(result).toBe(true); + expect(mockSpendingLimitService.evaluateSpendingLimits).not.toHaveBeenCalled(); + }); + }); + + // ── no agentId ──────────────────────────────────────────────────────────── + + describe('when decorator is present but no agentId in body', () => { + it('returns true without calling SpendingLimitService (org-level transaction)', async () => { + vi.spyOn(reflector, 'getAllAndOverride').mockReturnValue(true); + + const ctx = makeContext({ + body: { amount: '100', asset: 'USDC', recipientAddress: VALID_STELLAR, walletId: 'wallet-1' }, + }); + const result = await guard.canActivate(ctx); + + expect(result).toBe(true); + expect(mockSpendingLimitService.evaluateSpendingLimits).not.toHaveBeenCalled(); + }); + }); + + // ── no organizationId (unauthenticated / @Public route) ────────────────── + + describe('when no organizationId on request.user', () => { + it('returns true without evaluating limits', async () => { + vi.spyOn(reflector, 'getAllAndOverride').mockReturnValue(true); + + const ctx = makeContext({ + body: { agentId: 'agent-1', amount: '100', asset: 'USDC', recipientAddress: VALID_STELLAR }, + user: { id: 'user-1' }, // no organizationId + }); + const result = await guard.canActivate(ctx); + + expect(result).toBe(true); + expect(mockSpendingLimitService.evaluateSpendingLimits).not.toHaveBeenCalled(); + }); + }); + + // ── successful evaluation (limits not exceeded) ─────────────────────────── + + describe('when evaluation passes', () => { + beforeEach(() => { + vi.spyOn(reflector, 'getAllAndOverride').mockReturnValue(true); + mockSpendingLimitService.evaluateSpendingLimits.mockResolvedValue(undefined); + }); + + it('returns true and calls SpendingLimitService with the correct intent', async () => { + const ctx = makeContext({ + body: { + agentId: 'agent-1', + amount: '250', + asset: 'USDC', + recipientAddress: VALID_STELLAR, + walletId: 'wallet-1', + }, + }); + + const result = await guard.canActivate(ctx); + + expect(result).toBe(true); + expect(mockSpendingLimitService.evaluateSpendingLimits).toHaveBeenCalledTimes(1); + + const [intent, actorId] = mockSpendingLimitService.evaluateSpendingLimits.mock.calls[0] as [Record, string]; + expect(intent.organizationId).toBe('org-1'); + expect(intent.agentId).toBe('agent-1'); + expect(intent.amount).toBe(250); + expect(intent.asset).toBe('USDC'); + expect(intent.recipientAddress).toBe(VALID_STELLAR); + expect(intent.walletId).toBe('wallet-1'); + expect(actorId).toBe('user-1'); + }); + + it('uses XLM as default asset when none provided', async () => { + const ctx = makeContext({ + body: { agentId: 'agent-1', amount: '50', recipientAddress: VALID_STELLAR }, + }); + + await guard.canActivate(ctx); + + const [intent] = mockSpendingLimitService.evaluateSpendingLimits.mock.calls[0] as [Record]; + expect(intent.asset).toBe('XLM'); + }); + + it('passes amount as a number (not a string)', async () => { + const ctx = makeContext({ + body: { + agentId: 'agent-1', + amount: '750.5000000', + asset: 'XLM', + recipientAddress: VALID_STELLAR, + }, + }); + + await guard.canActivate(ctx); + + const [intent] = mockSpendingLimitService.evaluateSpendingLimits.mock.calls[0] as [Record]; + expect(typeof intent.amount).toBe('number'); + expect(intent.amount).toBe(750.5); + }); + }); + + // ── violation — daily limit exceeded ───────────────────────────────────── + + describe('when evaluation fails (limit exceeded)', () => { + beforeEach(() => { + vi.spyOn(reflector, 'getAllAndOverride').mockReturnValue(true); + }); + + it('propagates PolicyViolationException thrown by SpendingLimitService', async () => { + mockSpendingLimitService.evaluateSpendingLimits.mockRejectedValue( + new PolicyViolationException( + 'Transaction blocked by spending limit policy: Daily Spend Cap: Projected daily spend 550 exceeds limit 500', + { + violations: [ + { policyId: 'p-1', policyName: 'Daily Spend Cap', code: 'DAILY_LIMIT_EXCEEDED', message: 'Projected daily spend 550 exceeds limit 500' }, + ], + aggregates: { spentToday: 450, spentThisWeek: 450, spentThisMonth: 450 }, + }, + ), + ); + + const ctx = makeContext({ + body: { agentId: 'agent-1', amount: '100', asset: 'USDC', recipientAddress: VALID_STELLAR, walletId: 'wallet-1' }, + }); + + await expect(guard.canActivate(ctx)).rejects.toThrow(PolicyViolationException); + }); + + it('propagates the correct error code (POLICY_VIOLATION)', async () => { + const violation = new PolicyViolationException('Limit exceeded', {}); + mockSpendingLimitService.evaluateSpendingLimits.mockRejectedValue(violation); + + const ctx = makeContext({ + body: { agentId: 'agent-1', amount: '100', asset: 'USDC', recipientAddress: VALID_STELLAR }, + }); + + let caught: PolicyViolationException | undefined; + try { + await guard.canActivate(ctx); + } catch (err) { + caught = err as PolicyViolationException; + } + + expect(caught).toBeInstanceOf(PolicyViolationException); + expect(caught?.code).toBe('POLICY_VIOLATION'); + expect(caught?.getStatus()).toBe(422); + }); + + it('does not swallow other unexpected errors from the service', async () => { + mockSpendingLimitService.evaluateSpendingLimits.mockRejectedValue( + new Error('Prisma connection lost'), + ); + + const ctx = makeContext({ + body: { agentId: 'agent-1', amount: '100', asset: 'USDC', recipientAddress: VALID_STELLAR }, + }); + + await expect(guard.canActivate(ctx)).rejects.toThrow('Prisma connection lost'); + }); + }); + + // ── reflector metadata key ──────────────────────────────────────────────── + + describe('metadata key contract', () => { + it('reads the correct metadata key from the reflector', async () => { + const getAllAndOverrideSpy = vi + .spyOn(reflector, 'getAllAndOverride') + .mockReturnValue(false); + + const ctx = makeContext({ body: { agentId: 'agent-1' } }); + await guard.canActivate(ctx); + + expect(getAllAndOverrideSpy).toHaveBeenCalledWith(SPENDING_LIMIT_GUARD_KEY, expect.any(Array)); + }); + }); +}); diff --git a/src/modules/transactions/tests/spending-limit.service.spec.ts b/src/modules/transactions/tests/spending-limit.service.spec.ts new file mode 100644 index 00000000..df2ab431 --- /dev/null +++ b/src/modules/transactions/tests/spending-limit.service.spec.ts @@ -0,0 +1,546 @@ +import { describe, it, expect, beforeEach, vi } from 'vitest'; +import { Test, TestingModule } from '@nestjs/testing'; +import { SpendingLimitService } from '../spending-limit.service'; +import { PolicyService } from '../../policies/policy.service'; +import { AuditService } from '../../audit/audit.service'; +import { PrismaService } from '../../../database/prisma.service'; +import { PolicyViolationException } from '../../../common/exceptions/domain.exception'; +import { TransactionIntent } from '../../policies/policy.types'; +import { Decimal } from '@prisma/client/runtime/library'; + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +const VALID_STELLAR = 'GDEGSXLGANKHK7QFOV63XCBHBTZ3YRKUJV7ZB7JMSJQB5CNBRLL5QIG5'; + +function makeIntent(overrides: Partial = {}): TransactionIntent { + return { + organizationId: 'org-1', + agentId: 'agent-1', + walletId: 'wallet-1', + asset: 'USDC', + amount: 100, + recipientAddress: VALID_STELLAR, + at: new Date('2026-09-30T10:00:00Z'), + ...overrides, + }; +} + +// --------------------------------------------------------------------------- +// Mocks +// --------------------------------------------------------------------------- + +const mockPolicyService = { + evaluateIntent: vi.fn(), +}; + +const mockAuditService = { + record: vi.fn(), +}; + +// Prisma mock — returns Decimal sums for the three aggregate windows +const makeDecimal = (n: number) => ({ toNumber: () => n } as unknown as Decimal); + +const mockPrismaService = { + $transaction: vi.fn(), + policy: { + findMany: vi.fn(), + }, + transaction: { + aggregate: vi.fn(), + }, +}; + +// --------------------------------------------------------------------------- +// Suite +// --------------------------------------------------------------------------- + +describe('SpendingLimitService', () => { + let service: SpendingLimitService; + + beforeEach(async () => { + vi.clearAllMocks(); + + const module: TestingModule = await Test.createTestingModule({ + providers: [ + SpendingLimitService, + { provide: PolicyService, useValue: mockPolicyService }, + { provide: AuditService, useValue: mockAuditService }, + { provide: PrismaService, useValue: mockPrismaService }, + ], + }).compile(); + + service = module.get(SpendingLimitService); + }); + + // ── aggregateSpend ──────────────────────────────────────────────────────── + + describe('aggregateSpend', () => { + it('returns zeros when no transactions exist for the agent', async () => { + mockPrismaService.$transaction.mockResolvedValue([ + { _sum: { amount: null } }, + { _sum: { amount: null } }, + { _sum: { amount: null } }, + ]); + + const result = await service.aggregateSpend('agent-1', 'USDC'); + + expect(result).toEqual({ spentToday: 0, spentThisWeek: 0, spentThisMonth: 0 }); + }); + + it('converts Prisma Decimal sums to numbers correctly', async () => { + mockPrismaService.$transaction.mockResolvedValue([ + { _sum: { amount: makeDecimal(50) } }, + { _sum: { amount: makeDecimal(200) } }, + { _sum: { amount: makeDecimal(800) } }, + ]); + + const result = await service.aggregateSpend('agent-1', 'USDC'); + + expect(result).toEqual({ spentToday: 50, spentThisWeek: 200, spentThisMonth: 800 }); + }); + + it('runs all three aggregate queries in a single prisma $transaction call', async () => { + mockPrismaService.$transaction.mockResolvedValue([ + { _sum: { amount: null } }, + { _sum: { amount: null } }, + { _sum: { amount: null } }, + ]); + + await service.aggregateSpend('agent-1', 'XLM'); + + expect(mockPrismaService.$transaction).toHaveBeenCalledTimes(1); + // The array passed to $transaction should contain 3 query promises + const queryArray = mockPrismaService.$transaction.mock.calls[0][0] as unknown[]; + expect(queryArray).toHaveLength(3); + }); + }); + + // ── hasSpendingLimitPolicy ──────────────────────────────────────────────── + + describe('hasSpendingLimitPolicy', () => { + it('returns true when a policy with dailyLimit exists', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { dailyLimit: 500 } }, + ]); + + const result = await service.hasSpendingLimitPolicy('org-1', 'agent-1'); + expect(result).toBe(true); + }); + + it('returns true when a policy with weeklyLimit exists', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { weeklyLimit: 2000 } }, + ]); + + const result = await service.hasSpendingLimitPolicy('org-1', 'agent-1'); + expect(result).toBe(true); + }); + + it('returns true when a policy with monthlyLimit exists', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { monthlyLimit: 10000 } }, + ]); + + const result = await service.hasSpendingLimitPolicy('org-1', 'agent-1'); + expect(result).toBe(true); + }); + + it('returns false when no policies define periodic limits', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { maxAmount: 1000, allowedAssets: ['USDC'] } }, + ]); + + const result = await service.hasSpendingLimitPolicy('org-1', 'agent-1'); + expect(result).toBe(false); + }); + + it('returns false when no policies exist at all', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([]); + + const result = await service.hasSpendingLimitPolicy('org-1', 'agent-1'); + expect(result).toBe(false); + }); + }); + + // ── evaluateSpendingLimits ──────────────────────────────────────────────── + + describe('evaluateSpendingLimits', () => { + describe('no-op paths', () => { + it('returns early (no-op) when intent has no agentId', async () => { + const intent = makeIntent({ agentId: undefined }); + + await service.evaluateSpendingLimits(intent, 'user-1'); + + expect(mockPrismaService.policy.findMany).not.toHaveBeenCalled(); + expect(mockPolicyService.evaluateIntent).not.toHaveBeenCalled(); + }); + + it('returns early (no-op) when no periodic limit policies are configured', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { maxAmount: 5000 } }, + ]); + + await service.evaluateSpendingLimits(makeIntent(), 'user-1'); + + expect(mockPolicyService.evaluateIntent).not.toHaveBeenCalled(); + expect(mockAuditService.record).not.toHaveBeenCalled(); + }); + }); + + describe('policy passes', () => { + beforeEach(() => { + // Org has a daily limit policy + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { dailyLimit: 500 } }, + ]); + // Agent has spent 50 today + mockPrismaService.$transaction.mockResolvedValue([ + { _sum: { amount: makeDecimal(50) } }, + { _sum: { amount: makeDecimal(50) } }, + { _sum: { amount: makeDecimal(50) } }, + ]); + }); + + it('does not throw when spend + amount is within daily limit', async () => { + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: true, + requiresApproval: false, + violations: [], + evaluatedPolicyIds: ['policy-1'], + }); + + // 50 already spent + 100 pending = 150, limit is 500 — should pass + await expect( + service.evaluateSpendingLimits(makeIntent({ amount: 100 }), 'user-1'), + ).resolves.toBeUndefined(); + + expect(mockAuditService.record).not.toHaveBeenCalled(); + }); + + it('enriches intent with real spend aggregates before calling evaluateIntent', async () => { + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: true, + requiresApproval: false, + violations: [], + evaluatedPolicyIds: ['policy-1'], + }); + + await service.evaluateSpendingLimits(makeIntent({ amount: 100 }), 'user-1'); + + const intentPassedToEngine = mockPolicyService.evaluateIntent.mock.calls[0][0] as TransactionIntent; + expect(intentPassedToEngine.spentToday).toBe(50); + expect(intentPassedToEngine.spentThisWeek).toBe(50); + expect(intentPassedToEngine.spentThisMonth).toBe(50); + expect(intentPassedToEngine.amount).toBe(100); + }); + }); + + describe('limit exceeded — daily', () => { + beforeEach(() => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { dailyLimit: 500 } }, + ]); + // 450 already spent today + mockPrismaService.$transaction.mockResolvedValue([ + { _sum: { amount: makeDecimal(450) } }, + { _sum: { amount: makeDecimal(450) } }, + { _sum: { amount: makeDecimal(450) } }, + ]); + }); + + it('throws PolicyViolationException when daily limit would be exceeded', async () => { + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: false, + requiresApproval: false, + violations: [ + { + policyId: 'policy-1', + policyName: 'Daily Spend Cap', + code: 'DAILY_LIMIT_EXCEEDED', + message: 'Projected daily spend 550 exceeds limit 500', + }, + ], + evaluatedPolicyIds: ['policy-1'], + }); + + // 450 + 100 = 550 > 500 + await expect( + service.evaluateSpendingLimits(makeIntent({ amount: 100 }), 'user-1'), + ).rejects.toThrow(PolicyViolationException); + }); + + it('throws with error code POLICY_VIOLATION', async () => { + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: false, + requiresApproval: false, + violations: [ + { + policyId: 'policy-1', + policyName: 'Daily Spend Cap', + code: 'DAILY_LIMIT_EXCEEDED', + message: 'Projected daily spend 550 exceeds limit 500', + }, + ], + evaluatedPolicyIds: ['policy-1'], + }); + + let caughtError: PolicyViolationException | undefined; + try { + await service.evaluateSpendingLimits(makeIntent({ amount: 100 }), 'user-1'); + } catch (err) { + caughtError = err as PolicyViolationException; + } + + expect(caughtError).toBeInstanceOf(PolicyViolationException); + expect(caughtError?.code).toBe('POLICY_VIOLATION'); + }); + + it('throws with HTTP status 422 on daily limit violation', async () => { + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: false, + requiresApproval: false, + violations: [ + { + policyId: 'policy-1', + policyName: 'Daily Spend Cap', + code: 'DAILY_LIMIT_EXCEEDED', + message: 'Projected daily spend 550 exceeds limit 500', + }, + ], + evaluatedPolicyIds: ['policy-1'], + }); + + let caughtError: PolicyViolationException | undefined; + try { + await service.evaluateSpendingLimits(makeIntent({ amount: 100 }), 'user-1'); + } catch (err) { + caughtError = err as PolicyViolationException; + } + + expect(caughtError?.getStatus()).toBe(422); + }); + + it('writes a SPENDING_LIMIT_EXCEEDED audit log entry on violation', async () => { + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: false, + requiresApproval: false, + violations: [ + { + policyId: 'policy-1', + policyName: 'Daily Spend Cap', + code: 'DAILY_LIMIT_EXCEEDED', + message: 'Projected daily spend 550 exceeds limit 500', + }, + ], + evaluatedPolicyIds: ['policy-1'], + }); + mockAuditService.record.mockResolvedValue(undefined); + + await expect( + service.evaluateSpendingLimits(makeIntent({ amount: 100 }), 'user-1'), + ).rejects.toThrow(PolicyViolationException); + + expect(mockAuditService.record).toHaveBeenCalledTimes(1); + const auditCall = mockAuditService.record.mock.calls[0][0] as Record; + expect(auditCall.action).toBe('SPENDING_LIMIT_EXCEEDED'); + expect(auditCall.entity).toBe('transaction'); + expect(auditCall.userId).toBe('user-1'); + expect(auditCall.organizationId).toBe('org-1'); + }); + + it('includes violation codes in the audit log newValue', async () => { + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: false, + requiresApproval: false, + violations: [ + { + policyId: 'policy-1', + policyName: 'Daily Spend Cap', + code: 'DAILY_LIMIT_EXCEEDED', + message: 'Projected daily spend 550 exceeds limit 500', + }, + ], + evaluatedPolicyIds: ['policy-1'], + }); + mockAuditService.record.mockResolvedValue(undefined); + + await expect( + service.evaluateSpendingLimits(makeIntent({ amount: 100 }), 'user-1'), + ).rejects.toThrow(PolicyViolationException); + + const auditCall = mockAuditService.record.mock.calls[0][0] as Record; + const newValue = auditCall.newValue as Record; + expect((newValue.violations as Array<{ code: string }>)[0].code).toBe('DAILY_LIMIT_EXCEEDED'); + expect((newValue.aggregates as Record).spentToday).toBe(450); + }); + + it('still throws PolicyViolationException even when audit log write fails', async () => { + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: false, + requiresApproval: false, + violations: [ + { + policyId: 'policy-1', + policyName: 'Daily Spend Cap', + code: 'DAILY_LIMIT_EXCEEDED', + message: 'Projected daily spend 550 exceeds limit 500', + }, + ], + evaluatedPolicyIds: ['policy-1'], + }); + // Simulate an audit service failure + mockAuditService.record.mockRejectedValue(new Error('DB connection lost')); + + // The transaction must still be blocked — audit failures are non-fatal + await expect( + service.evaluateSpendingLimits(makeIntent({ amount: 100 }), 'user-1'), + ).rejects.toThrow(PolicyViolationException); + }); + }); + + describe('limit exceeded — weekly', () => { + it('throws PolicyViolationException when weekly limit would be exceeded', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { weeklyLimit: 1000 } }, + ]); + mockPrismaService.$transaction.mockResolvedValue([ + { _sum: { amount: makeDecimal(50) } }, + { _sum: { amount: makeDecimal(950) } }, // 950 this week + { _sum: { amount: makeDecimal(950) } }, + ]); + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: false, + requiresApproval: false, + violations: [ + { + policyId: 'policy-2', + policyName: 'Weekly Spend Cap', + code: 'WEEKLY_LIMIT_EXCEEDED', + message: 'Projected weekly spend 1050 exceeds limit 1000', + }, + ], + evaluatedPolicyIds: ['policy-2'], + }); + + // 950 + 100 = 1050 > 1000 + await expect( + service.evaluateSpendingLimits(makeIntent({ amount: 100 }), 'user-1'), + ).rejects.toThrow(PolicyViolationException); + }); + }); + + describe('limit exceeded — monthly', () => { + it('throws PolicyViolationException when monthly limit would be exceeded', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { monthlyLimit: 5000 } }, + ]); + mockPrismaService.$transaction.mockResolvedValue([ + { _sum: { amount: makeDecimal(100) } }, + { _sum: { amount: makeDecimal(500) } }, + { _sum: { amount: makeDecimal(4950) } }, // 4950 this month + ]); + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: false, + requiresApproval: false, + violations: [ + { + policyId: 'policy-3', + policyName: 'Monthly Spend Cap', + code: 'MONTHLY_LIMIT_EXCEEDED', + message: 'Projected monthly spend 5050 exceeds limit 5000', + }, + ], + evaluatedPolicyIds: ['policy-3'], + }); + + // 4950 + 100 = 5050 > 5000 + await expect( + service.evaluateSpendingLimits(makeIntent({ amount: 100 }), 'user-1'), + ).rejects.toThrow(PolicyViolationException); + }); + }); + + describe('missing policy scenario', () => { + it('is a no-op and does not throw when no spending policies are configured', async () => { + // Returns no policies at all + mockPrismaService.policy.findMany.mockResolvedValue([]); + + await expect( + service.evaluateSpendingLimits(makeIntent(), 'user-1'), + ).resolves.toBeUndefined(); + + expect(mockPolicyService.evaluateIntent).not.toHaveBeenCalled(); + expect(mockAuditService.record).not.toHaveBeenCalled(); + }); + + it('is a no-op when only non-periodic policies exist (maxAmount, blockedAssets)', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { maxAmount: 1000, blockedAssets: ['BTC'] } }, + { configuration: { allowedRecipients: [VALID_STELLAR] } }, + ]); + + await expect( + service.evaluateSpendingLimits(makeIntent(), 'user-1'), + ).resolves.toBeUndefined(); + + expect(mockPolicyService.evaluateIntent).not.toHaveBeenCalled(); + }); + }); + + describe('intent enrichment', () => { + it('passes enriched intent (with aggregates) to PolicyService.evaluateIntent', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { dailyLimit: 1000, weeklyLimit: 5000, monthlyLimit: 15000 } }, + ]); + mockPrismaService.$transaction.mockResolvedValue([ + { _sum: { amount: makeDecimal(200) } }, + { _sum: { amount: makeDecimal(1200) } }, + { _sum: { amount: makeDecimal(3500) } }, + ]); + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: true, + requiresApproval: false, + violations: [], + evaluatedPolicyIds: ['policy-1'], + }); + + await service.evaluateSpendingLimits(makeIntent({ amount: 50 }), 'user-1'); + + const enrichedIntent = mockPolicyService.evaluateIntent.mock.calls[0][0] as TransactionIntent; + expect(enrichedIntent.spentToday).toBe(200); + expect(enrichedIntent.spentThisWeek).toBe(1200); + expect(enrichedIntent.spentThisMonth).toBe(3500); + expect(enrichedIntent.amount).toBe(50); + expect(enrichedIntent.agentId).toBe('agent-1'); + expect(enrichedIntent.organizationId).toBe('org-1'); + }); + + it('preserves original intent fields (asset, recipientAddress, walletId)', async () => { + mockPrismaService.policy.findMany.mockResolvedValue([ + { configuration: { dailyLimit: 1000 } }, + ]); + mockPrismaService.$transaction.mockResolvedValue([ + { _sum: { amount: makeDecimal(0) } }, + { _sum: { amount: makeDecimal(0) } }, + { _sum: { amount: makeDecimal(0) } }, + ]); + mockPolicyService.evaluateIntent.mockResolvedValue({ + passed: true, + requiresApproval: false, + violations: [], + evaluatedPolicyIds: ['policy-1'], + }); + + const intent = makeIntent({ asset: 'XLM', walletId: 'wallet-42' }); + await service.evaluateSpendingLimits(intent, 'user-1'); + + const passedIntent = mockPolicyService.evaluateIntent.mock.calls[0][0] as TransactionIntent; + expect(passedIntent.asset).toBe('XLM'); + expect(passedIntent.walletId).toBe('wallet-42'); + expect(passedIntent.recipientAddress).toBe(VALID_STELLAR); + }); + }); + }); +}); diff --git a/src/modules/transactions/tests/transaction.service.spec.ts b/src/modules/transactions/tests/transaction.service.spec.ts new file mode 100644 index 00000000..40900b41 --- /dev/null +++ b/src/modules/transactions/tests/transaction.service.spec.ts @@ -0,0 +1,295 @@ +import { describe, it, expect, beforeEach, vi } from 'vitest'; +import { Test, TestingModule } from '@nestjs/testing'; +import { TransactionService } from '../transaction.service'; +import { TransactionRepository } from '../transaction.repository'; +import { WalletService } from '../../wallets/wallet.service'; +import { AgentService } from '../../agents/agent.service'; +import { PolicyService } from '../../policies/policy.service'; +import { RiskService } from '../../risk/risk.service'; +import { BudgetService } from '../../budgets/budget.service'; +import { StellarService } from '../../stellar/stellar.service'; +import { SpendingLimitService } from '../spending-limit.service'; +import { EventBusService } from '../../../events/event-bus.service'; +import { PrismaService } from '../../../database/prisma.service'; +import { AgentStatus, RiskBand, TransactionStatus, WalletStatus } from '@prisma/client'; +import { Keypair } from '@stellar/stellar-sdk'; +import { CreateTransactionInput } from '../transaction.dto'; +import { DomainException } from '../../../common/exceptions/domain.exception'; +import { ErrorCode } from '../../../common/constants/error-codes'; + +describe('TransactionService', () => { + describe('Governance simulation', () => { + let service: TransactionService; + let repository: { + hasPaidRecipient: ReturnType; + recentCountForWallet: ReturnType; + create: ReturnType; + }; + let policyService: { evaluateIntent: ReturnType }; + let riskService: { assess: ReturnType }; + let eventBus: { emit: ReturnType }; + let stellarService: { submitPayment: ReturnType }; + + const input: CreateTransactionInput = { + walletId: 'wallet_1', + asset: 'XLM', + amount: '50.0', + recipientAddress: Keypair.random().publicKey(), + metadata: {}, + }; + + beforeEach(() => { + repository = { + hasPaidRecipient: vi.fn().mockResolvedValue(false), + recentCountForWallet: vi.fn().mockResolvedValue(0), + create: vi.fn(), + }; + policyService = { + evaluateIntent: vi.fn().mockResolvedValue({ + passed: true, + requiresApproval: false, + violations: [], + }), + }; + riskService = { + assess: vi.fn().mockReturnValue({ + score: 10, + band: RiskBand.LOW, + factors: [], + canAutoExecute: true, + }), + }; + eventBus = { emit: vi.fn().mockResolvedValue(undefined) }; + stellarService = { submitPayment: vi.fn() }; + + service = new TransactionService( + repository as unknown as TransactionRepository, + { + getOrThrow: vi.fn().mockResolvedValue({ + id: 'wallet_1', + status: WalletStatus.ACTIVE, + stellarAddress: Keypair.random().publicKey(), + network: 'TESTNET', + createdAt: new Date(), + }), + } as unknown as WalletService, + { getOrThrow: vi.fn() } as unknown as AgentService, + policyService as unknown as PolicyService, + riskService as unknown as RiskService, + {} as BudgetService, + stellarService as unknown as StellarService, + eventBus as unknown as EventBusService, + {} as PrismaService, + { + aggregateSpend: vi.fn().mockResolvedValue({ + spentToday: 0, + spentThisWeek: 0, + spentThisMonth: 0, + }), + evaluateSpendingLimits: vi.fn().mockResolvedValue(undefined), + } as unknown as SpendingLimitService, + ); + }); + + it('returns policy and risk results without persisting or broadcasting', async () => { + const result = await service.simulate('org_1', input); + + expect(result).toMatchObject({ + wouldPass: true, + requiresApproval: false, + policy: { passed: true, violations: [] }, + risk: { score: 10, band: RiskBand.LOW }, + }); + expect(repository.hasPaidRecipient).toHaveBeenCalledWith('org_1', input.recipientAddress); + expect(repository.recentCountForWallet).toHaveBeenCalledWith('wallet_1'); + expect(repository.create).not.toHaveBeenCalled(); + expect(eventBus.emit).not.toHaveBeenCalled(); + expect(stellarService.submitPayment).not.toHaveBeenCalled(); + }); + + it('flags high-risk assessments for approval during a dry run', async () => { + riskService.assess.mockReturnValueOnce({ + score: 45, + band: RiskBand.MEDIUM, + factors: [], + canAutoExecute: false, + }); + + const result = await service.simulate('org_1', input); + + expect(result.wouldPass).toBe(true); + expect(result.requiresApproval).toBe(true); + expect(repository.create).not.toHaveBeenCalled(); + expect(stellarService.submitPayment).not.toHaveBeenCalled(); + }); + }); +}); + +describe('TransactionService - create', () => { + let service: TransactionService; + let stellarService: StellarService; + + const wallet = { + id: 'wallet_1', + status: WalletStatus.ACTIVE, + stellarAddress: 'GDWALLETADDRESSXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX', + network: 'TESTNET', + createdAt: new Date('2025-01-01T00:00:00Z'), + }; + + beforeEach(async () => { + const module: TestingModule = await Test.createTestingModule({ + providers: [ + TransactionService, + { + provide: TransactionRepository, + useValue: (() => { + let stored: Record | undefined; + return { + create: vi.fn().mockImplementation((data: Record) => { + stored = { id: 'tx_1', ...data, status: data.status ?? TransactionStatus.DRAFT }; + return Promise.resolve(stored); + }), + update: vi.fn().mockImplementation((id: string, data: Record) => { + stored = { ...stored, id, ...data }; + return Promise.resolve(stored); + }), + findById: vi.fn().mockImplementation(() => Promise.resolve(stored)), + hasPaidRecipient: vi.fn().mockResolvedValue(false), + recentCountForWallet: vi.fn().mockResolvedValue(0), + }; + })(), + }, + { provide: WalletService, useValue: { getOrThrow: vi.fn().mockResolvedValue(wallet) } }, + { + provide: AgentService, + useValue: { + getOrThrow: vi.fn().mockResolvedValue({ id: 'agent_1', status: AgentStatus.ACTIVE }), + }, + }, + { + provide: PolicyService, + useValue: { + checkVelocityLimit: vi.fn().mockResolvedValue(undefined), + evaluateIntent: vi.fn().mockResolvedValue({ + passed: true, + requiresApproval: false, + violations: [], + evaluatedPolicyIds: [], + }), + }, + }, + { + provide: RiskService, + useValue: { + evaluate: vi.fn().mockResolvedValue({ + score: 10, + band: RiskBand.LOW, + factors: [], + canAutoExecute: true, + }), + }, + }, + { + provide: BudgetService, + useValue: { + assertWithinBudget: vi.fn().mockResolvedValue(undefined), + consume: vi.fn().mockResolvedValue(undefined), + }, + }, + { + provide: StellarService, + useValue: { + submitPayment: vi.fn().mockResolvedValue({ + hash: 'stellar_hash_1', + ledger: 100, + successful: true, + }), + }, + }, + { provide: EventBusService, useValue: { emit: vi.fn().mockResolvedValue(undefined) } }, + { provide: PrismaService, useValue: {} }, + { + provide: SpendingLimitService, + useValue: { + aggregateSpend: vi.fn().mockResolvedValue({ spentToday: 0, spentThisWeek: 0, spentThisMonth: 0 }), + evaluateSpendingLimits: vi.fn().mockResolvedValue(undefined), + }, + }, + ], + }).compile(); + + service = module.get(TransactionService); + stellarService = module.get(StellarService); + }); + + const input = { + walletId: 'wallet_1', + agentId: 'agent_1', + recipientAddress: 'GDEGSXLGANKHK7QFOV63XCBHBTZ3YRKUJV7ZB7JMSJQB5CNBRLL5QIG5', + amount: '50.0', + asset: 'XLM', + memo: 'Test payment', + metadata: {}, + }; + + it('auto-executes and submits on-chain when policy and risk allow it', async () => { + const result = await service.create('org_1', 'user_1', input); + + expect(stellarService.submitPayment).toHaveBeenCalledWith( + expect.objectContaining({ + sourceAddress: wallet.stellarAddress, + destinationAddress: input.recipientAddress, + asset: input.asset, + }), + ); + expect(result.requiresApproval).toBe(false); + expect(result.transaction.status).toBe(TransactionStatus.COMPLETED); + }); + + it('blocks a policy-violating transaction before submission', async () => { + const blockedPolicy = { + checkVelocityLimit: vi.fn().mockResolvedValue(undefined), + evaluateIntent: vi.fn().mockResolvedValue({ + passed: false, + requiresApproval: false, + violations: [{ policyId: 'policy_1', reason: 'exceeds max amount' }], + evaluatedPolicyIds: ['policy_1'], + }), + }; + const blockedModule: TestingModule = await Test.createTestingModule({ + providers: [ + TransactionService, + { provide: TransactionRepository, useValue: { create: vi.fn(), update: vi.fn() } }, + { provide: WalletService, useValue: { getOrThrow: vi.fn().mockResolvedValue(wallet) } }, + { + provide: AgentService, + useValue: { getOrThrow: vi.fn().mockResolvedValue({ id: 'agent_1', status: AgentStatus.ACTIVE }) }, + }, + { provide: PolicyService, useValue: blockedPolicy }, + { provide: RiskService, useValue: { evaluate: vi.fn() } }, + { provide: BudgetService, useValue: { assertWithinBudget: vi.fn(), consume: vi.fn() } }, + { provide: StellarService, useValue: { submitPayment: vi.fn() } }, + { provide: EventBusService, useValue: { emit: vi.fn().mockResolvedValue(undefined) } }, + { provide: PrismaService, useValue: {} }, + { + provide: SpendingLimitService, + useValue: { + aggregateSpend: vi.fn().mockResolvedValue({ spentToday: 0, spentThisWeek: 0, spentThisMonth: 0 }), + evaluateSpendingLimits: vi.fn().mockResolvedValue(undefined), + }, + }, + ], + }).compile(); + const blockedService = blockedModule.get(TransactionService); + const blockedStellar = blockedModule.get(StellarService); + + const error = await blockedService.create('org_1', 'user_1', input).catch( + (reason: unknown) => reason as DomainException, + ); + + expect(error).toBeInstanceOf(DomainException); + expect((error as DomainException).code).toBe(ErrorCode.POLICY_VIOLATION); + expect(blockedStellar.submitPayment).not.toHaveBeenCalled(); + }); +}); diff --git a/src/modules/transactions/transaction.controller.ts b/src/modules/transactions/transaction.controller.ts index dafb470d..774380a7 100644 --- a/src/modules/transactions/transaction.controller.ts +++ b/src/modules/transactions/transaction.controller.ts @@ -17,14 +17,17 @@ import { simulateTransactionSchema, SimulateTransactionInput, } from './transaction.dto'; +import { SpendingLimitGuard, RequireSpendingLimitCheck } from './guards/spending-limit.guard'; import { CurrentUser } from '../../common/decorators/current-user.decorator'; import { Roles } from '../../common/decorators/roles.decorator'; import { AuditAction } from '../../common/decorators/audit-action.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { UseWalletLock } from '../../common/locks/wallet-lock.decorator'; import { UseTransactionLock } from '../../common/locks/transaction-lock.decorator'; +import { AgentThrottlerGuard } from '../../common/guards/agent-throttler.guard'; import { AuthenticatedUser } from '../../common/interfaces/authenticated-user.interface'; import { PaginationQuery, paginationQuerySchema } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; import { ApiEnvelope } from '../../common/decorators/api-envelope.decorator'; import { SlidingWindowThrottlerGuard, @@ -33,6 +36,7 @@ import { @ApiTags('transactions') @ApiBearerAuth('access-token') +@UseGuards(AgentThrottlerGuard) @Controller('transactions') export class TransactionController { constructor(private readonly transactionService: TransactionService) {} @@ -43,8 +47,7 @@ export class TransactionController { description: 'Returns a paginated list of transactions for the current organization. Supports filtering by status, agent, wallet, and date range.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiQuery({ name: 'status', required: false, enum: ['DRAFT', 'PENDING', 'APPROVED', 'COMPLETED', 'FAILED', 'CANCELLED'], description: 'Filter by transaction status' }) @ApiQuery({ name: 'agentId', required: false, type: String, description: 'Filter by agent UUID' }) @ApiQuery({ name: 'walletId', required: false, type: String, description: 'Filter by wallet UUID' }) @@ -60,8 +63,9 @@ export class TransactionController { @Post() @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE, UserRole.DEVELOPER) - @UseGuards(SlidingWindowThrottlerGuard) + @UseGuards(SlidingWindowThrottlerGuard, SpendingLimitGuard) @SlidingWindowLimit(30, 60) + @RequireSpendingLimitCheck() @UseWalletLock() @UseTransactionLock({ attempts: 3, retryDelayMs: 50 }) @AuditAction('TRANSFER_FUNDS') @@ -69,7 +73,9 @@ export class TransactionController { summary: 'Create a transaction (runs the full governance pipeline)', description: 'Evaluates policies, scores risk and checks budgets. Auto-executes when permitted, ' + - 'otherwise creates an approval proposal and returns requiresApproval=true.', + 'otherwise creates an approval proposal and returns requiresApproval=true. ' + + 'Agent transactions are additionally evaluated against daily/weekly/monthly spending ' + + 'limit policies before reaching the service layer.', }) @ApiBody({ type: CreateTransactionDto }) @ApiEnvelope(CreateTransactionDto as never) @@ -78,11 +84,16 @@ export class TransactionController { @ApiResponse({ status: 401, description: 'Not authenticated' }) @ApiResponse({ status: 403, description: 'Insufficient permissions' }) @ApiResponse({ status: 409, description: 'Insufficient budget or risk threshold exceeded' }) + @ApiResponse({ + status: 422, + description: 'Transaction blocked by spending limit policy (POLICY_VIOLATION)', + }) create( @CurrentUser() user: AuthenticatedUser, @Body(new ZodValidationPipe(createTransactionSchema)) body: CreateTransactionInput, ) { - return this.transactionService.create(user.organizationId, user.id, body); + const actorId = user.isApiKey ? user.createdById ?? user.id : user.id; + return this.transactionService.create(user.organizationId, actorId, body); } @Post('simulate') diff --git a/src/modules/transactions/transaction.dto.spec.ts b/src/modules/transactions/transaction.dto.spec.ts index e0257620..852177bf 100644 --- a/src/modules/transactions/transaction.dto.spec.ts +++ b/src/modules/transactions/transaction.dto.spec.ts @@ -8,6 +8,60 @@ describe('createTransactionSchema memo validation', () => { recipientAddress: 'GDEGSXLGANKHK7QFOV63XCBHBTZ3YRKUJV7ZB7JMSJQB5CNBRLL5QIG5', }; + describe('payment input validation', () => { + it('accepts a complete payment payload with metadata', () => { + expect( + createTransactionSchema.parse({ + ...baseInput, + metadata: { invoiceId: 'invoice-1' }, + }), + ).toMatchObject({ + amount: '10.0000000', + recipientAddress: baseInput.recipientAddress, + metadata: { invoiceId: 'invoice-1' }, + }); + }); + + it('rejects missing wallet identifiers and negative or zero amounts', () => { + expect(() => + createTransactionSchema.parse({ + amount: '10', + recipientAddress: baseInput.recipientAddress, + }), + ).toThrow(); + expect(() => + createTransactionSchema.parse({ ...baseInput, amount: '-1' }), + ).toThrow(); + expect(() => + createTransactionSchema.parse({ ...baseInput, amount: '0' }), + ).toThrow(); + }); + + it('rejects malformed recipients and non-object metadata', () => { + expect(() => + createTransactionSchema.parse({ ...baseInput, recipientAddress: 'not-a-stellar-address' }), + ).toThrow(); + expect(() => + createTransactionSchema.parse({ ...baseInput, metadata: ['unexpected'] }), + ).toThrow(); + }); + + it('rejects non-JSON values nested in metadata', () => { + expect(() => + createTransactionSchema.parse({ + ...baseInput, + metadata: { nested: { unsupported: undefined } }, + }), + ).toThrow(); + expect(() => + createTransactionSchema.parse({ + ...baseInput, + metadata: { nested: [Number.NaN] }, + }), + ).toThrow(); + }); + }); + describe('legacy string memo (TEXT type)', () => { it('accepts valid legacy memo', () => { const result = createTransactionSchema.parse({ ...baseInput, memo: 'hello' }); diff --git a/src/modules/transactions/transaction.dto.ts b/src/modules/transactions/transaction.dto.ts index 233011b0..18b3d6e7 100644 --- a/src/modules/transactions/transaction.dto.ts +++ b/src/modules/transactions/transaction.dto.ts @@ -3,6 +3,19 @@ import { ApiProperty, ApiPropertyOptional } from '@nestjs/swagger'; import { stellarMemoTypeSchema } from '../../common/validators/stellar-memo.schema'; import { stellarAddressSchema } from '../../common/validators/stellar-address.schema'; +type JsonValue = string | number | boolean | null | JsonValue[] | { [key: string]: JsonValue }; + +const jsonValueSchema: z.ZodType = z.lazy(() => + z.union([ + z.string(), + z.number().finite(), + z.boolean(), + z.null(), + z.array(jsonValueSchema), + z.record(jsonValueSchema), + ]), +); + const amountString = z .string() .regex(/^\d+(\.\d{1,7})?$/, 'Amount must be a positive decimal with up to 7 places') @@ -20,7 +33,7 @@ export const createTransactionSchema = z memoType: stellarMemoTypeSchema.optional(), memoValue: z.string().optional(), purpose: z.string().max(280).optional(), - metadata: z.record(z.unknown()).default({}), + metadata: z.record(jsonValueSchema).default({}), }) .strict() .superRefine((data, ctx) => { diff --git a/src/modules/transactions/transaction.module.ts b/src/modules/transactions/transaction.module.ts index e0bab2d5..d6e80561 100644 --- a/src/modules/transactions/transaction.module.ts +++ b/src/modules/transactions/transaction.module.ts @@ -7,6 +7,11 @@ import { AgentModule } from '../agents/agent.module'; import { PolicyModule } from '../policies/policy.module'; import { RiskModule } from '../risk/risk.module'; import { BudgetModule } from '../budgets/budget.module'; +import { SorobanSimulationService } from './services/soroban-simulation.service'; +import { StellarSimulationService } from './services/stellar-simulation.service'; +import { StellarModule } from '../stellar/stellar.module'; +import { SpendingLimitService } from './spending-limit.service'; +import { SpendingLimitGuard } from './guards/spending-limit.guard'; /** * Transaction pipeline module. Pulls together wallets, agents, policies, risk @@ -15,9 +20,16 @@ import { BudgetModule } from '../budgets/budget.module'; * approved proposal's transaction. */ @Module({ - imports: [WalletModule, AgentModule, PolicyModule, RiskModule, BudgetModule], + imports: [WalletModule, AgentModule, PolicyModule, RiskModule, BudgetModule, StellarModule], controllers: [TransactionController], - providers: [TransactionService, TransactionRepository], - exports: [TransactionService], + providers: [ + TransactionService, + TransactionRepository, + SorobanSimulationService, + StellarSimulationService, + SpendingLimitService, + SpendingLimitGuard, + ], + exports: [TransactionService, SorobanSimulationService, StellarSimulationService, SpendingLimitService], }) export class TransactionModule {} diff --git a/src/modules/transactions/transaction.service.ts b/src/modules/transactions/transaction.service.ts index 107f0b51..5466a3e9 100644 --- a/src/modules/transactions/transaction.service.ts +++ b/src/modules/transactions/transaction.service.ts @@ -10,6 +10,7 @@ import { WalletStatus, } from '@prisma/client'; import { TransactionRepository } from './transaction.repository'; +import { SpendingLimitService } from './spending-limit.service'; import { CreateTransactionInput } from './transaction.dto'; import { TransactionsValidator } from './transactions.validator'; import { WalletService, toNetworkName } from '../wallets/wallet.service'; @@ -68,6 +69,7 @@ export class TransactionService { private readonly stellar: StellarService, private readonly eventBus: EventBusService, private readonly prisma: PrismaService, + private readonly spendingLimits: SpendingLimitService, ) {} async create(organizationId: string, actorId: string, input: CreateTransactionInput) { @@ -77,12 +79,33 @@ export class TransactionService { // 2.5. Velocity limit check for agent spending if (input.agentId) { - await this.policies.checkVelocityLimit(input.agentId, amount, input.asset); + await this.policies.checkVelocityLimit(organizationId, input.agentId, amount, input.asset, actorId); } - // 3. Policy evaluation — a hard failure blocks the transaction outright. - const intent = this.toIntent(organizationId, input, amount); - const policyResult = await this.policies.evaluateIntent(intent, actorId); + // 2.6. Spending limit evaluation — atomically fetches daily/weekly/monthly + // aggregates and evaluates periodic budget caps. On violation, writes + // an audit log entry and throws PolicyViolationException (HTTP 422). + // Returns the aggregates so we can reuse them in step 3 below. + const baseIntent = this.toIntent(organizationId, input, amount); + let enrichedIntent = baseIntent; + if (input.agentId) { + const aggregates = await this.spendingLimits.aggregateSpend(input.agentId, input.asset); + enrichedIntent = { + ...baseIntent, + spentToday: aggregates.spentToday, + spentThisWeek: aggregates.spentThisWeek, + spentThisMonth: aggregates.spentThisMonth, + }; + // evaluateSpendingLimits uses the enriched intent so periodic limit + // checks run against real aggregates — it is a no-op when no periodic + // limit policy is configured, so the overhead is minimal. + await this.spendingLimits.evaluateSpendingLimits(enrichedIntent, actorId); + } + + // 3. Full policy evaluation — evaluates all rule types (maxAmount, assets, + // recipients, time windows, emergency lock, periodic limits) with real + // spend aggregates already embedded in the intent. + const policyResult = await this.policies.evaluateIntent(enrichedIntent, actorId); if (!policyResult.passed) { throw new DomainException( ErrorCode.POLICY_VIOLATION, @@ -249,7 +272,7 @@ export class TransactionService { } const pagination = toPrismaPagination(query, SORTABLE); const { items, total } = await this.repository.findManyAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); + return new Paginated(items, buildPaginationMeta(total, query)); } async getOrThrow(organizationId: string, id: string): Promise { diff --git a/src/modules/wallets/wallet.controller.ts b/src/modules/wallets/wallet.controller.ts index 06472734..1a52d55c 100644 --- a/src/modules/wallets/wallet.controller.ts +++ b/src/modules/wallets/wallet.controller.ts @@ -7,6 +7,7 @@ import { Patch, Post, Query, + UseGuards, } from '@nestjs/common'; import { ApiOperation, @@ -32,13 +33,18 @@ import { import { CurrentUser } from '../../common/decorators/current-user.decorator'; import { Roles } from '../../common/decorators/roles.decorator'; import { AuditAction } from '../../common/decorators/audit-action.decorator'; +import { AuditLog } from '../../common/decorators/audit-log.decorator'; +import { AgentThrottlerGuard } from '../../common/guards/agent-throttler.guard'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { AuthenticatedUser } from '../../common/interfaces/authenticated-user.interface'; import { PaginationQuery, paginationQuerySchema } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; import { ApiEnvelope } from '../../common/decorators/api-envelope.decorator'; +import { AstroidThrottlerGuard } from '../../common/guards/throttler.guard'; @ApiTags('wallets') @ApiBearerAuth('access-token') +@UseGuards(AgentThrottlerGuard) @Controller('wallets') export class WalletController { constructor(private readonly walletService: WalletService) {} @@ -50,8 +56,7 @@ export class WalletController { 'Returns a paginated list of wallets for the current organization. ' + 'Supports filtering by status and network.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiQuery({ name: 'status', required: false, enum: ['ACTIVE', 'FROZEN', 'ARCHIVED'], description: 'Filter by wallet status' }) @ApiQuery({ name: 'network', required: false, enum: ['TESTNET', 'PUBLIC'], description: 'Filter by Stellar network' }) @ApiEnvelope(WalletResponseDto as never, { isArray: true }) @@ -66,6 +71,7 @@ export class WalletController { @Post() @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE, UserRole.DEVELOPER) + @UseGuards(AstroidThrottlerGuard) @AuditAction('WALLET_CREATED') @ApiOperation({ summary: 'Create a wallet (generate a keypair or import an address)', @@ -117,6 +123,7 @@ export class WalletController { @Patch(':id') @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE, UserRole.DEVELOPER) @AuditAction('WALLET_UPDATED') + @AuditLog({ action: 'WALLET_UPDATED', entity: 'Wallet' }) @ApiOperation({ summary: 'Update a wallet label or owning agent', description: 'Partial update of wallet metadata. Does not affect the Stellar keypair.', @@ -139,6 +146,7 @@ export class WalletController { @Post(':id/freeze') @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE) @AuditAction('WALLET_FROZEN') + @AuditLog({ action: 'WALLET_FROZEN', entity: 'Wallet' }) @ApiOperation({ summary: 'Freeze a wallet (block outgoing transactions)', description: @@ -157,6 +165,7 @@ export class WalletController { @Post(':id/unfreeze') @Roles(UserRole.OWNER, UserRole.ADMIN, UserRole.FINANCE) @AuditAction('WALLET_UNFROZEN') + @AuditLog({ action: 'WALLET_UNFROZEN', entity: 'Wallet' }) @ApiOperation({ summary: 'Unfreeze a wallet', description: 'Restores a frozen wallet to ACTIVE status, allowing outgoing transactions again.', @@ -174,6 +183,7 @@ export class WalletController { @Delete(':id') @Roles(UserRole.OWNER, UserRole.ADMIN) @AuditAction('WALLET_ARCHIVED') + @AuditLog({ action: 'WALLET_ARCHIVED', entity: 'Wallet' }) @ApiOperation({ summary: 'Archive (soft-delete) a wallet', description: diff --git a/src/modules/wallets/wallet.service.ts b/src/modules/wallets/wallet.service.ts index 0e66b272..9ba2f517 100644 --- a/src/modules/wallets/wallet.service.ts +++ b/src/modules/wallets/wallet.service.ts @@ -111,7 +111,7 @@ export class WalletService { } const pagination = toPrismaPagination(query, SORTABLE); const { items, total } = await this.repository.findManyAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); + return new Paginated(items, buildPaginationMeta(total, query)); } async getOrThrow(organizationId: string, id: string): Promise { diff --git a/src/modules/webhooks/services/webhook-audit.service.spec.ts b/src/modules/webhooks/services/webhook-audit.service.spec.ts new file mode 100644 index 00000000..ab85bbb5 --- /dev/null +++ b/src/modules/webhooks/services/webhook-audit.service.spec.ts @@ -0,0 +1,79 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { Logger } from '@nestjs/common'; + +import { AuditService } from '../../audit/audit.service'; +import { + WEBHOOK_DELIVERY_FAILED_ACTION, + WebhookAuditService, + WebhookDeliveryFailure, +} from './webhook-audit.service'; + +const FAILURE: WebhookDeliveryFailure = { + webhookId: 'wh-1', + organizationId: 'org-1', + url: 'https://example.com/hook', + eventName: 'transaction.completed', + eventId: 'event-1', + attemptsMade: 5, + failedReason: 'HTTP 503: Service Unavailable', + responseStatus: 503, +}; + +describe('WebhookAuditService', () => { + let record: ReturnType; + let service: WebhookAuditService; + + beforeEach(() => { + record = vi.fn().mockResolvedValue({ id: 'audit-1' }); + service = new WebhookAuditService({ record } as unknown as AuditService); + }); + + it('appends a WEBHOOK_DELIVERY_FAILED entry for a dead-lettered delivery', async () => { + await service.recordTerminalFailure(FAILURE); + + expect(record).toHaveBeenCalledWith( + expect.objectContaining({ + organizationId: 'org-1', + userId: null, + action: WEBHOOK_DELIVERY_FAILED_ACTION, + entity: 'Webhook', + entityId: 'wh-1', + newValue: expect.objectContaining({ + url: 'https://example.com/hook', + eventName: 'transaction.completed', + eventId: 'event-1', + attemptsMade: 5, + responseStatus: 503, + failedReason: 'HTTP 503: Service Unavailable', + deadLettered: true, + }), + }), + ); + }); + + it('defaults the optional event/response fields to null', async () => { + await service.recordTerminalFailure({ + webhookId: 'wh-2', + organizationId: 'org-2', + url: 'https://example.com/hook', + attemptsMade: 5, + failedReason: 'socket hang up', + }); + + const { newValue } = record.mock.calls[0][0]; + expect(newValue.eventName).toBeNull(); + expect(newValue.eventId).toBeNull(); + expect(newValue.responseStatus).toBeNull(); + }); + + it('never throws when the audit write fails, so the worker stays alive', async () => { + const logger = vi.spyOn(Logger.prototype, 'error').mockImplementation(() => undefined); + record.mockRejectedValue(new Error('audit table down')); + + await expect(service.recordTerminalFailure(FAILURE)).resolves.toBeUndefined(); + expect(logger).toHaveBeenCalledWith( + expect.stringContaining('Failed to audit webhook wh-1 delivery failure'), + ); + logger.mockRestore(); + }); +}); diff --git a/src/modules/webhooks/services/webhook-audit.service.ts b/src/modules/webhooks/services/webhook-audit.service.ts new file mode 100644 index 00000000..46fafe7f --- /dev/null +++ b/src/modules/webhooks/services/webhook-audit.service.ts @@ -0,0 +1,65 @@ +import { Injectable, Logger } from '@nestjs/common'; +import { Prisma } from '@prisma/client'; + +import { AuditService } from '../../audit/audit.service'; + +/** Everything needed to reconstruct why a webhook delivery was abandoned. */ +export interface WebhookDeliveryFailure { + webhookId: string; + organizationId: string; + url: string; + attemptsMade: number; + failedReason: string; + eventName?: string; + eventId?: string; + responseStatus?: number; +} + +/** Audit action recorded for a webhook that exhausted every retry attempt. */ +export const WEBHOOK_DELIVERY_FAILED_ACTION = 'WEBHOOK_DELIVERY_FAILED'; + +/** + * Writes permanently failed webhook deliveries into the compliance audit trail. + * + * A webhook that exhausts its retries has been moved to the dead-letter queue by + * the queue failure listener; this service adds the *business* record — which + * subscriber, which event, how many attempts, and why it died — so an operator + * can answer "did the agent's approval notification ever arrive?" from the audit + * log alone. + * + * Auditing is best-effort by design: it runs inside a worker whose job is to + * deliver notifications, and a logging failure must never turn into a crashed + * or endlessly-retried job. + */ +@Injectable() +export class WebhookAuditService { + private readonly logger = new Logger(WebhookAuditService.name); + + constructor(private readonly auditService: AuditService) {} + + /** Appends one `WEBHOOK_DELIVERY_FAILED` entry. Never throws. */ + async recordTerminalFailure(failure: WebhookDeliveryFailure): Promise { + try { + await this.auditService.record({ + organizationId: failure.organizationId, + userId: null, + action: WEBHOOK_DELIVERY_FAILED_ACTION, + entity: 'Webhook', + entityId: failure.webhookId, + newValue: { + url: failure.url, + eventName: failure.eventName ?? null, + eventId: failure.eventId ?? null, + attemptsMade: failure.attemptsMade, + responseStatus: failure.responseStatus ?? null, + failedReason: failure.failedReason, + deadLettered: true, + } as unknown as Prisma.InputJsonValue, + }); + } catch (error) { + this.logger.error( + `Failed to audit webhook ${failure.webhookId} delivery failure: ${(error as Error).message}`, + ); + } + } +} diff --git a/src/modules/webhooks/services/webhook-delivery.service.spec.ts b/src/modules/webhooks/services/webhook-delivery.service.spec.ts index bf8d6cb0..e3238bce 100644 --- a/src/modules/webhooks/services/webhook-delivery.service.spec.ts +++ b/src/modules/webhooks/services/webhook-delivery.service.spec.ts @@ -18,7 +18,6 @@ describe('WebhookDeliveryService', () => { webhookId: 'webhook-1', organizationId: 'org-1', url: 'https://example.com/webhook', - secret: 'secret-key', eventName: 'transaction.created', payload: { id: 'txn-1' }, eventId: 'event-1', @@ -31,7 +30,14 @@ describe('WebhookDeliveryService', () => { expect(mockQueue.add).toHaveBeenCalledWith( 'webhook-delivery', - jobData, + expect.objectContaining({ + ...jobData, + metadata: { + requestId: expect.any(String), + correlationId: expect.any(String), + traceId: expect.any(String), + }, + }), { attempts: 5, backoff: { @@ -50,7 +56,7 @@ describe('WebhookDeliveryService', () => { expect(mockQueue.add).toHaveBeenCalledWith( 'webhook-delivery', - jobData, + expect.objectContaining({ ...jobData, metadata: expect.objectContaining({ requestId: expect.any(String) }) }), expect.any(Object), ); }); @@ -61,7 +67,7 @@ describe('WebhookDeliveryService', () => { expect(mockQueue.add).toHaveBeenCalledWith( 'webhook-delivery', - jobData, + expect.objectContaining({ ...jobData, metadata: expect.objectContaining({ requestId: expect.any(String) }) }), expect.any(Object), ); }); @@ -73,7 +79,7 @@ describe('WebhookDeliveryService', () => { expect(mockQueue.add).toHaveBeenCalledWith( 'webhook-delivery', - jobData, + expect.objectContaining({ ...jobData, metadata: expect.objectContaining({ requestId: expect.any(String) }) }), expect.any(Object), ); }); diff --git a/src/modules/webhooks/services/webhook-delivery.service.ts b/src/modules/webhooks/services/webhook-delivery.service.ts index 5055e9f2..c17d51a7 100644 --- a/src/modules/webhooks/services/webhook-delivery.service.ts +++ b/src/modules/webhooks/services/webhook-delivery.service.ts @@ -1,14 +1,19 @@ import { Injectable, Logger } from '@nestjs/common'; import { InjectQueue } from '@nestjs/bullmq'; import { Queue } from 'bullmq'; + import { Queues } from '../../../queues/queues.constants'; +import { WEBHOOK_JOB_NAME, webhookJobOptions } from '../../../queues/webhook.queue'; import { WebhookJobData } from '../types/webhook-job.types'; +import { RequestContext } from '../../../common/context/request-context'; +import { resolveRequestId } from '../../../common/helpers/request-id'; /** * Service for queuing webhook delivery jobs with BullMQ. - * Handles retry logic and dead-letter queueing through the queue configuration. - * Uses exponential backoff with randomized jitter (via worker-level backoffStrategy) - * to prevent thundering herd problems. + * + * Retry policy, jittered backoff and dead-letter settings come from + * `@queues/webhook.queue` so the API-side enqueue options and the worker-side + * queue registration can never drift apart. */ @Injectable() export class WebhookDeliveryService { @@ -20,27 +25,33 @@ export class WebhookDeliveryService { ) {} /** - * Queues a webhook delivery job with exponential backoff retry policy. - * The job will be processed by the webhook worker with automatic retries. - * Uses 2000ms base delay for exponential backoff. - * Randomized jitter is applied by the worker-level backoffStrategy. - * Maximum 5 attempts total. + * Queues a webhook delivery job. + * + * The job inherits the queue's retry policy: 5 attempts, exponential backoff + * (2000ms base) with ±20% jitter applied by the custom backoff strategy, and + * 24h retention of failed jobs so exhausted deliveries remain inspectable. */ async queueDelivery(data: WebhookJobData): Promise { try { - await this.webhookQueue.add('webhook-delivery', data, { - attempts: 5, - backoff: { - type: 'exponential', - delay: 2000, - }, - removeOnComplete: { count: 1000 }, - removeOnFail: { age: 24 * 3600 }, - }); + const metadata = { + ...data.metadata, + requestId: data.metadata?.requestId ?? RequestContext.getRequestId() ?? resolveRequestId(undefined), + correlationId: data.metadata?.correlationId ?? RequestContext.getCorrelationId(), + traceId: data.metadata?.traceId ?? RequestContext.getTraceId(), + }; + metadata.correlationId ??= metadata.requestId; + metadata.traceId ??= metadata.correlationId; + const hasMetadata = Object.values(metadata).some((value) => value !== undefined); + const jobData: WebhookJobData = { + ...data, + ...(hasMetadata ? { metadata } : {}), + }; + await this.webhookQueue.add(WEBHOOK_JOB_NAME, jobData, webhookJobOptions); this.logger.debug(`Queued webhook delivery for ${data.eventName} to ${data.url}`); } catch (error) { - this.logger.error(`Failed to queue webhook delivery: ${(error as Error).message}`); + this.logger.error('Failed to queue webhook delivery'); throw error; } } } + diff --git a/src/modules/webhooks/types/webhook-job.types.ts b/src/modules/webhooks/types/webhook-job.types.ts index 8fce3b39..aadf5630 100644 --- a/src/modules/webhooks/types/webhook-job.types.ts +++ b/src/modules/webhooks/types/webhook-job.types.ts @@ -2,14 +2,16 @@ * BullMQ job types for webhook delivery with retry logic. */ +import { QueueJobMetadata } from '../../../queues/queues.constants'; + export interface WebhookJobData { webhookId: string; organizationId: string; url: string; - secret: string; eventName: string; payload: unknown; eventId: string; + metadata?: QueueJobMetadata; } export interface WebhookJobResult { diff --git a/src/modules/webhooks/utils/signing.spec.ts b/src/modules/webhooks/utils/signing.spec.ts index 0f3e39ab..419ff800 100644 --- a/src/modules/webhooks/utils/signing.spec.ts +++ b/src/modules/webhooks/utils/signing.spec.ts @@ -1,16 +1,33 @@ import { describe, it, expect } from 'vitest'; import { createHmac } from 'crypto'; -import { signWebhookPayload } from './signing'; +import { signWebhookPayload, verifyWebhookSignature } from './signing'; describe('signWebhookPayload', () => { it('should generate expected signature', () => { const secret = 'test-secret'; const timestamp = '1234567890'; - const payload = JSON.stringify({ event: 'test.event' }); + const payload = Buffer.from(JSON.stringify({ event: 'test.event' })); + const eventId = 'evt-123'; - const signature = signWebhookPayload(secret, timestamp, payload); + const signature = signWebhookPayload(secret, timestamp, eventId, payload); - const expected = createHmac('sha256', secret).update(timestamp + '.' + payload).digest('hex'); - expect(signature).toBe(expected); + const expected = createHmac('sha256', secret) + .update(Buffer.concat([Buffer.from(`v1.${timestamp}.${eventId}.`), payload])) + .digest('hex'); + expect(signature).toBe(`v1=${expected}`); + expect(verifyWebhookSignature(secret, timestamp, eventId, payload, signature)).toBe(true); + }); + + it('rejects altered payload, timestamp, event ID, and signature', () => { + const secret = 'test-secret'; + const timestamp = '1234567890'; + const eventId = 'evt-123'; + const payload = Buffer.from('{"event":"test.event"}'); + const signature = signWebhookPayload(secret, timestamp, eventId, payload); + + expect(verifyWebhookSignature(secret, timestamp, eventId, Buffer.from('{}'), signature)).toBe(false); + expect(verifyWebhookSignature(secret, '1234567891', eventId, payload, signature)).toBe(false); + expect(verifyWebhookSignature(secret, timestamp, 'evt-124', payload, signature)).toBe(false); + expect(verifyWebhookSignature(secret, timestamp, eventId, payload, 'v1=0'.repeat(64))).toBe(false); }); }); diff --git a/src/modules/webhooks/utils/signing.ts b/src/modules/webhooks/utils/signing.ts index 228ed752..a639fa01 100644 --- a/src/modules/webhooks/utils/signing.ts +++ b/src/modules/webhooks/utils/signing.ts @@ -1,5 +1,42 @@ -import { createHmac } from 'crypto'; +import { createHmac, timingSafeEqual } from 'crypto'; -export function signWebhookPayload(secret: string, timestamp: string, payload: string): string { - return createHmac('sha256', secret).update(timestamp + '.' + payload).digest('hex'); +export const WEBHOOK_SIGNATURE_VERSION = 'v1'; + +export function buildWebhookSignatureInput( + timestamp: string, + eventId: string, + payload: Buffer, +): Buffer { + return Buffer.concat([ + Buffer.from(`${WEBHOOK_SIGNATURE_VERSION}.${timestamp}.${eventId}.`, 'utf8'), + payload, + ]); +} + +export function signWebhookPayload( + secret: string, + timestamp: string, + eventId: string, + payload: Buffer, +): string { + const digest = createHmac('sha256', secret) + .update(buildWebhookSignatureInput(timestamp, eventId, payload)) + .digest('hex'); + return `${WEBHOOK_SIGNATURE_VERSION}=${digest}`; +} + +export function verifyWebhookSignature( + secret: string, + timestamp: string, + eventId: string, + payload: Buffer, + signature: string, +): boolean { + if (!/^\d{1,12}$/.test(timestamp) || !eventId || eventId.length > 256) return false; + const match = /^v1=([0-9a-f]{64})$/.exec(signature); + if (!match) return false; + + const expected = Buffer.from(signWebhookPayload(secret, timestamp, eventId, payload).slice(3), 'hex'); + const received = Buffer.from(match[1], 'hex'); + return expected.length === received.length && timingSafeEqual(expected, received); } diff --git a/src/modules/webhooks/webhook-failure-audit.spec.ts b/src/modules/webhooks/webhook-failure-audit.spec.ts new file mode 100644 index 00000000..d94e07fe --- /dev/null +++ b/src/modules/webhooks/webhook-failure-audit.spec.ts @@ -0,0 +1,109 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { Job, UnrecoverableError } from 'bullmq'; + +import { WebhooksProcessor } from './webhooks.processor'; +import { WebhookAuditService } from './services/webhook-audit.service'; +import { WebhookJobData } from './types/webhook-job.types'; + +/** + * Terminal-failure behaviour of the webhook processor: retries are scheduled + * while attempts remain, and a delivery that exhausts them (or hits a + * non-retryable 4xx) is written to the audit trail exactly once. + */ +describe('WebhooksProcessor terminal failures', () => { + const jobData: WebhookJobData = { + webhookId: 'wh-1', + organizationId: 'org-1', + url: 'https://downstream.example.com/hook', + eventName: 'transaction.completed', + payload: { id: 'txn-1' }, + eventId: 'event-1', + }; + + let recordTerminalFailure: ReturnType; + let processor: WebhooksProcessor; + let fetchSpy: ReturnType; + + function makeJob(attemptsMade: number): Job { + return { id: 'job-1', name: 'webhook-delivery', data: jobData, attemptsMade } as unknown as Job; + } + + beforeEach(() => { + recordTerminalFailure = vi.fn().mockResolvedValue(undefined); + processor = new WebhooksProcessor( + { webhook: { findFirst: vi.fn().mockResolvedValue({ secret: 'whsec_test' }) } } as never, + undefined, + { recordTerminalFailure } as unknown as WebhookAuditService, + ); + fetchSpy = vi.fn(); + vi.stubGlobal('fetch', fetchSpy); + }); + + afterEach(() => { + vi.restoreAllMocks(); + vi.unstubAllGlobals(); + }); + + it('schedules a retry without auditing while attempts remain', async () => { + fetchSpy.mockResolvedValue({ + ok: false, + status: 503, + statusText: 'Service Unavailable', + text: () => Promise.resolve('Service Unavailable'), + }); + + await expect(processor.process(makeJob(1))).rejects.toThrow('HTTP 503'); + + expect(recordTerminalFailure).not.toHaveBeenCalled(); + }); + + it('audits the delivery once the retries are exhausted', async () => { + fetchSpy.mockResolvedValue({ + ok: false, + status: 503, + statusText: 'Service Unavailable', + text: () => Promise.resolve('Service Unavailable'), + }); + + await expect(processor.process(makeJob(4))).rejects.toThrow('HTTP 503'); + + expect(recordTerminalFailure).toHaveBeenCalledTimes(1); + expect(recordTerminalFailure).toHaveBeenCalledWith({ + webhookId: 'wh-1', + organizationId: 'org-1', + url: 'https://downstream.example.com/hook', + eventName: 'transaction.completed', + eventId: 'event-1', + attemptsMade: 5, + failedReason: expect.stringContaining('HTTP 503'), + responseStatus: 503, + }); + }); + + it('audits and stops retrying immediately on a non-transient 4xx', async () => { + fetchSpy.mockResolvedValue({ + ok: false, + status: 404, + statusText: 'Not Found', + text: () => Promise.resolve('Not Found'), + }); + + await expect(processor.process(makeJob(0))).rejects.toBeInstanceOf(UnrecoverableError); + + expect(recordTerminalFailure).toHaveBeenCalledTimes(1); + expect(recordTerminalFailure).toHaveBeenCalledWith( + expect.objectContaining({ attemptsMade: 1, responseStatus: 404 }), + ); + }); + + it('never lets an audit failure mask the original delivery error', async () => { + const logger = { warn: vi.fn(), error: vi.fn(), debug: vi.fn(), log: vi.fn() }; + Object.assign(processor, { logger }); + recordTerminalFailure.mockRejectedValue(new Error('audit unavailable')); + fetchSpy.mockRejectedValue(new Error('socket hang up')); + + await expect(processor.process(makeJob(4))).rejects.toThrow('socket hang up'); + + expect(logger.warn).toHaveBeenCalledWith(expect.stringContaining('Could not audit webhook')); + }); +}); diff --git a/src/modules/webhooks/webhook.controller.ts b/src/modules/webhooks/webhook.controller.ts index c20312d3..bfb91db3 100644 --- a/src/modules/webhooks/webhook.controller.ts +++ b/src/modules/webhooks/webhook.controller.ts @@ -32,6 +32,7 @@ import { Roles } from '../../common/decorators/roles.decorator'; import { AuditAction } from '../../common/decorators/audit-action.decorator'; import { ZodValidationPipe } from '../../common/pipes/zod-validation.pipe'; import { PaginationQuery, paginationQuerySchema } from '../../common/helpers/pagination'; +import { ApiPaginationQuery } from '../../common/decorators/api-pagination-query.decorator'; @ApiTags('webhooks') @ApiBearerAuth('access-token') @@ -47,8 +48,7 @@ export class WebhookController { 'Returns a paginated list of webhooks for the current organization. ' + 'HMAC signing secrets are never included in the response.', }) - @ApiQuery({ name: 'page', required: false, type: Number, description: 'Page number (default: 1)' }) - @ApiQuery({ name: 'limit', required: false, type: Number, description: 'Items per page (default: 20)' }) + @ApiPaginationQuery() @ApiQuery({ name: 'events', required: false, type: String, description: 'Filter by event type' }) @ApiResponse({ status: 200, description: 'Paginated list of webhooks' }) @ApiResponse({ status: 401, description: 'Not authenticated' }) diff --git a/src/modules/webhooks/webhook.dispatcher.spec.ts b/src/modules/webhooks/webhook.dispatcher.spec.ts new file mode 100644 index 00000000..a75c3377 --- /dev/null +++ b/src/modules/webhooks/webhook.dispatcher.spec.ts @@ -0,0 +1,53 @@ +import { describe, expect, it, vi } from 'vitest'; +import { WebhookDispatcher } from './webhook.dispatcher'; +import { WebhookRepository } from './webhook.repository'; +import { WebhookDeliveryService } from './services/webhook-delivery.service'; +import { DomainEventEnvelope } from '../../events/domain-event.types'; + +describe('WebhookDispatcher', () => { + it('queues envelope deliveries with stable event identity and no signing secret', async () => { + const queueDelivery = vi.fn().mockResolvedValue(undefined); + const webhook = { + id: 'wh-1', + organizationId: 'org-1', + url: 'https://example.com/hook', + secret: 'must-not-enter-job', + }; + const repository = { + findEnabledForEvent: vi.fn().mockResolvedValue([webhook]), + }; + const dispatcher = new WebhookDispatcher( + repository as unknown as WebhookRepository, + { queueDelivery } as unknown as WebhookDeliveryService, + ); + const envelope: DomainEventEnvelope = { + eventId: 'evt-stable-1', + name: 'budget.exceeded', + organizationId: 'org-1', + aggregateType: 'Budget', + aggregateId: 'budget-1', + requestId: 'req-1', + correlationId: 'corr-1', + payload: { budgetId: 'budget-1' }, + occurredAt: new Date('2026-09-29T12:00:00.000Z'), + }; + + await dispatcher.dispatch(envelope); + + expect(repository.findEnabledForEvent).toHaveBeenCalledWith('org-1', 'budget.exceeded'); + expect(queueDelivery).toHaveBeenCalledWith(expect.objectContaining({ + webhookId: 'wh-1', + organizationId: 'org-1', + eventId: 'evt-stable-1', + eventName: 'budget.exceeded', + payload: { + event: 'budget.exceeded', + occurredAt: envelope.occurredAt, + aggregateType: 'Budget', + aggregateId: 'budget-1', + data: { budgetId: 'budget-1' }, + }, + })); + expect(queueDelivery.mock.calls[0][0]).not.toHaveProperty('secret'); + }); +}); \ No newline at end of file diff --git a/src/modules/webhooks/webhook.dispatcher.ts b/src/modules/webhooks/webhook.dispatcher.ts index 91c55931..014685d5 100644 --- a/src/modules/webhooks/webhook.dispatcher.ts +++ b/src/modules/webhooks/webhook.dispatcher.ts @@ -3,6 +3,7 @@ import { OnEvent } from '@nestjs/event-emitter'; import { WebhookRepository } from './webhook.repository'; import { WebhookDeliveryService } from './services/webhook-delivery.service'; import { DomainEventEnvelope } from '../../events/domain-event.types'; +import { DOMAIN_EVENT_ENVELOPE } from '../../events/domain-event.types'; import { WEBHOOK_EVENTS } from '../../events/event-names'; /** @@ -21,7 +22,7 @@ export class WebhookDispatcher { private readonly deliveryService: WebhookDeliveryService, ) {} - @OnEvent('**') + @OnEvent(DOMAIN_EVENT_ENVELOPE) async dispatch(envelope: DomainEventEnvelope): Promise { if (!envelope?.organizationId) { return; @@ -54,15 +55,12 @@ export class WebhookDispatcher { webhookId: webhook.id, organizationId: envelope.organizationId || '', url: webhook.url, - secret: webhook.secret, eventName: envelope.name, payload, - eventId: `${envelope.aggregateType}-${envelope.aggregateId}-${envelope.occurredAt.getTime()}`, + eventId: envelope.eventId ?? `${envelope.name}-${envelope.aggregateType}-${envelope.aggregateId ?? 'unknown'}-${envelope.occurredAt.getTime()}`, }); - } catch (error) { - this.logger.error( - `Failed to queue webhook ${webhook.id} for '${envelope.name}': ${(error as Error).message}`, - ); + } catch { + this.logger.error(`Failed to queue webhook ${webhook.id} for '${envelope.name}'`); } }), ); diff --git a/src/modules/webhooks/webhook.management.spec.ts b/src/modules/webhooks/webhook.management.spec.ts new file mode 100644 index 00000000..2f5fcb58 --- /dev/null +++ b/src/modules/webhooks/webhook.management.spec.ts @@ -0,0 +1,58 @@ +import { describe, expect, it, vi } from 'vitest'; +import { WebhookRepository } from './webhook.repository'; +import { WebhookService } from './webhook.service'; + +function record(id: string, secret: string) { + return { + id, + organizationId: 'org-1', + url: `https://example.com/${id}`, + secret, + events: ['wallet.created'], + enabled: true, + createdAt: new Date('2026-09-29T00:00:00.000Z'), + updatedAt: new Date('2026-09-29T00:00:00.000Z'), + }; +} + +describe('WebhookService signing-secret lifecycle', () => { + it('returns a cryptographically random secret only from creation', async () => { + const create = vi.fn(async (data: { organizationId: string; url: string; secret: string; events: string[]; enabled: boolean }) => + record('wh-1', data.secret), + ); + const service = new WebhookService({ create } as unknown as WebhookRepository); + + const created = await service.create('org-1', { + url: 'https://example.com/hook', + events: ['wallet.created'], + enabled: true, + }); + + expect(created.secret).toMatch(/^whsec_[0-9a-f]{48}$/); + expect(create).toHaveBeenCalledWith(expect.objectContaining({ organizationId: 'org-1' })); + expect(created).not.toHaveProperty('signingSecret'); + }); + + it('rotates only the requested endpoint and redacts secrets from later reads', async () => { + const current = record('wh-1', 'old-secret'); + const findById = vi.fn(async (organizationId: string, id: string) => + organizationId === 'org-1' && id === current.id ? current : null, + ); + const update = vi.fn(async (id: string, changes: { secret?: string }) => ({ + ...current, + id, + secret: changes.secret ?? current.secret, + })); + const service = new WebhookService({ findById, update } as unknown as WebhookRepository); + + const rotated = await service.rotateSecret('org-1', 'wh-1'); + const nextSecret = rotated.secret; + expect(nextSecret).toMatch(/^whsec_[0-9a-f]{48}$/); + expect(nextSecret).not.toBe('old-secret'); + expect(update).toHaveBeenCalledWith('wh-1', { secret: nextSecret }); + expect(update).toHaveBeenCalledTimes(1); + + const fetched = await service.get('org-1', 'wh-1'); + expect(fetched).not.toHaveProperty('secret'); + }); +}); \ No newline at end of file diff --git a/src/modules/webhooks/webhook.module.ts b/src/modules/webhooks/webhook.module.ts index 159c2538..884316c9 100644 --- a/src/modules/webhooks/webhook.module.ts +++ b/src/modules/webhooks/webhook.module.ts @@ -7,23 +7,23 @@ import { WebhookService } from './webhook.service'; import { WebhookRepository } from './webhook.repository'; import { WebhookDispatcher } from './webhook.dispatcher'; import { WebhookDeliveryService } from './services/webhook-delivery.service'; +import { WebhookAuditService } from './services/webhook-audit.service'; import { WebhooksProcessor } from './webhooks.processor'; -import { Queues } from '../../queues/queues.constants'; +import { createWebhookQueueOptions } from '../../queues/webhook.queue'; import { redisConfig } from '../../config/redis.config'; -import { webhookBackoffStrategy } from '../../utils/backoff.util'; import { MetricsModule } from '../metrics/metrics.module'; import { RawBodyMiddleware } from '../../common/middleware/raw-body.middleware'; import { WebhookSignatureGuard } from '../../common/guards/webhook-signature.guard'; import { SlidingWindowThrottlerGuard } from '../../common/guards/sliding-window-throttler.guard'; -import type { RegisterQueueOptions } from '@nestjs/bullmq'; /** - * Webhooks module. The dispatcher listens to domain events and queues - * the curated WEBHOOK_EVENTS set to subscribed external endpoints via BullMQ. + * Webhooks module. The dispatcher listens to domain events and queues the + * curated WEBHOOK_EVENTS set to subscribed external endpoints via BullMQ. * - * Uses a custom backoffStrategy with randomized jitter (20% of base delay) - * to prevent thundering herd problems when multiple webhook deliveries - * are retried simultaneously. + * Retry policy (5 attempts, exponential backoff with jitter) and the + * dead-letter routing live in `@queues/webhook.queue`, so the API and the worker + * can never drift apart. `WebhookAuditService` records deliveries that exhaust + * their retries in the compliance audit trail. */ @Module({ imports: [ @@ -35,24 +35,7 @@ import type { RegisterQueueOptions } from '@nestjs/bullmq'; db: redisConfig().db, }, }), - BullModule.registerQueue({ - name: Queues.Webhooks, - defaultJobOptions: { - attempts: 5, - backoff: { - type: 'exponential', - delay: 2000, - }, - removeOnComplete: { count: 1000 }, - removeOnFail: { age: 24 * 3600 }, - }, - // BullMQ reads queue.opts.settings.backoffStrategy at retry time. - // The AdvancedOptions type is not fully exposed by @nestjs/bullmq, so we - // cast to include the backoffStrategy field that BullMQ supports at runtime. - settings: { - backoffStrategy: webhookBackoffStrategy, - } as RegisterQueueOptions['settings'], - }), + BullModule.registerQueue(createWebhookQueueOptions()), MetricsModule, ], controllers: [WebhookController, WebhookIngressController], @@ -61,6 +44,7 @@ import type { RegisterQueueOptions } from '@nestjs/bullmq'; WebhookRepository, WebhookDispatcher, WebhookDeliveryService, + WebhookAuditService, WebhooksProcessor, WebhookIngressService, WebhookSignatureGuard, diff --git a/src/modules/webhooks/webhook.service.spec.ts b/src/modules/webhooks/webhook.service.spec.ts index 8fdeae75..6f839941 100644 --- a/src/modules/webhooks/webhook.service.spec.ts +++ b/src/modules/webhooks/webhook.service.spec.ts @@ -61,7 +61,9 @@ describe('Webhook signing & delivery', () => { const EVENT_ID = 'evt-123'; beforeEach(() => { - processor = new WebhooksProcessor({} as never); + processor = new WebhooksProcessor({ + webhook: { findFirst: vi.fn().mockResolvedValue({ secret: SECRET }) }, + } as never); fetchSpy = vi.fn(); vi.stubGlobal('fetch', fetchSpy); }); @@ -78,7 +80,6 @@ describe('Webhook signing & delivery', () => { webhookId: 'wh-1', organizationId: 'org-1', url: URL, - secret: SECRET, eventName: 'budget.exceeded', payload: { event: 'budget.exceeded', data: {} }, eventId: EVENT_ID, @@ -89,7 +90,8 @@ describe('Webhook signing & delivery', () => { await processor.process(job); const [, opts] = fetchSpy.mock.calls[0]; expect(opts.headers['x-astroid-signature']).toBeDefined(); - expect(opts.headers['x-astroid-signature']).toMatch(/^[0-9a-f]{64}$/); + expect(opts.headers['x-astroid-signature']).toMatch(/^v1=[0-9a-f]{64}$/); + expect(opts.headers['x-astroid-signature-version']).toBe('v1'); expect(opts.headers['x-astroid-delivery']).toBe(EVENT_ID); expect(opts.headers['x-astroid-event']).toBe('budget.exceeded'); expect(opts.headers['x-astroid-timestamp']).toMatch(/^\d+$/); @@ -104,7 +106,6 @@ describe('Webhook signing & delivery', () => { webhookId: 'wh-1', organizationId: 'org-1', url: URL, - secret: SECRET, eventName: 'policy.violated', payload, eventId: EVENT_ID, @@ -114,11 +115,10 @@ describe('Webhook signing & delivery', () => { await processor.process(job); const [, opts] = fetchSpy.mock.calls[0]; - const body: string = opts.body; - const timestamp: string = opts.headers['x-astroid-timestamp']; - const expected = createHmac('sha256', SECRET).update(`${timestamp}.${body}`).digest('hex'); - expect(opts.headers['x-astroid-signature']).toBe(expected); - expect(body).toBe(JSON.stringify(payload)); + const body: Buffer = Buffer.from(opts.body); + expect(opts.headers['x-astroid-signature']).toMatch(/^v1=[0-9a-f]{64}$/); + expect(opts.headers['x-astroid-signature-version']).toBe('v1'); + expect(body.toString('utf8')).toBe(JSON.stringify(payload)); }); it('uses 5000ms timeout on fetch', async () => { @@ -129,7 +129,6 @@ describe('Webhook signing & delivery', () => { webhookId: 'wh-1', organizationId: 'org-1', url: URL, - secret: SECRET, eventName: 'wallet.created', payload: {}, eventId: EVENT_ID, @@ -141,12 +140,10 @@ describe('Webhook signing & delivery', () => { expect(opts.signal).toBeInstanceOf(AbortSignal); }); - it('falls back to ConfigService secret when per-endpoint secret is empty', async () => { + it('loads the secret by webhook and organization rather than from the job payload', async () => { const fallbackSecret = 'fallback-secret-123'; - const mockConfig = { - get: vi.fn((key: string) => (key === 'WEBHOOK_SECRET' ? fallbackSecret : undefined)), - } as unknown as import('@nestjs/config').ConfigService; - const processorWithFallback = new WebhooksProcessor({} as never, mockConfig); + const findFirst = vi.fn().mockResolvedValue({ secret: fallbackSecret }); + const processorWithDatabaseSecret = new WebhooksProcessor({ webhook: { findFirst } } as never); fetchSpy.mockResolvedValue({ ok: true, status: 200, text: () => Promise.resolve('OK') }); const job = { @@ -155,7 +152,6 @@ describe('Webhook signing & delivery', () => { webhookId: 'wh-1', organizationId: 'org-1', url: URL, - secret: '', eventName: 'transaction.completed', payload: { hello: 'world' }, eventId: EVENT_ID, @@ -163,12 +159,14 @@ describe('Webhook signing & delivery', () => { attemptsMade: 0, } as unknown as Job; - await processorWithFallback.process(job); + await processorWithDatabaseSecret.process(job); const [, opts] = fetchSpy.mock.calls[0]; - const body: string = opts.body; - const ts: string = opts.headers['x-astroid-timestamp']; - const expected = createHmac('sha256', fallbackSecret).update(`${ts}.${body}`).digest('hex'); - expect(opts.headers['x-astroid-signature']).toBe(expected); + expect(findFirst).toHaveBeenCalledWith({ + where: { id: 'wh-1', organizationId: 'org-1' }, + select: { secret: true }, + }); + expect(job.data).not.toHaveProperty('secret'); + expect(opts.headers['x-astroid-signature']).toMatch(/^v1=[0-9a-f]{64}$/); }); it('throws UnrecoverableError for non-transient 4xx and does not retry', async () => { @@ -179,7 +177,6 @@ describe('Webhook signing & delivery', () => { webhookId: 'wh-1', organizationId: 'org-1', url: URL, - secret: SECRET, eventName: 'wallet.created', payload: {}, eventId: EVENT_ID, @@ -197,7 +194,6 @@ describe('Webhook signing & delivery', () => { webhookId: 'wh-1', organizationId: 'org-1', url: URL, - secret: SECRET, eventName: 'wallet.created', payload: {}, eventId: EVENT_ID, @@ -222,7 +218,6 @@ describe('Webhook signing & delivery', () => { webhookId: 'wh-1', organizationId: 'org-1', url: URL, - secret: SECRET, eventName: 'transaction.completed', payload: {}, eventId: 'evt-1', @@ -234,5 +229,4 @@ describe('Webhook signing & delivery', () => { }); const URL = 'https://example.com/webhook'; - const SECRET = 'whsec_test-secret-key'; }); diff --git a/src/modules/webhooks/webhook.service.ts b/src/modules/webhooks/webhook.service.ts index 75a11896..235e0623 100644 --- a/src/modules/webhooks/webhook.service.ts +++ b/src/modules/webhooks/webhook.service.ts @@ -48,7 +48,7 @@ export class WebhookService { } const pagination = toPrismaPagination(query, SORTABLE); const { items, total } = await this.repository.findManyAndCount(where, pagination); - return new Paginated(items, buildPaginationMeta(total, query.page, query.limit)); + return new Paginated(items, buildPaginationMeta(total, query)); } async getOrThrow(organizationId: string, id: string) { diff --git a/src/modules/webhooks/webhooks.processor.spec.ts b/src/modules/webhooks/webhooks.processor.spec.ts index d757e657..362e21e5 100644 --- a/src/modules/webhooks/webhooks.processor.spec.ts +++ b/src/modules/webhooks/webhooks.processor.spec.ts @@ -2,7 +2,7 @@ import { describe, it, expect, vi, beforeEach, afterEach } from 'vitest'; import { Job, UnrecoverableError } from 'bullmq'; import { WebhooksProcessor } from './webhooks.processor'; import { WebhookJobData } from './types/webhook-job.types'; -import { createHmac } from 'crypto'; +import { verifyWebhookSignature } from './utils/signing'; describe('WebhooksProcessor', () => { let processor: WebhooksProcessor; @@ -19,7 +19,6 @@ describe('WebhooksProcessor', () => { webhookId: WEBHOOK_ID, organizationId: ORG_ID, url: WEBHOOK_URL, - secret: WEBHOOK_SECRET, eventName: 'transaction.completed', payload: { event: 'transaction.completed', data: { transactionId: 'txn-123' } }, eventId: EVENT_ID, @@ -35,7 +34,9 @@ describe('WebhooksProcessor', () => { }) as unknown as Job; beforeEach(() => { - mockPrisma = {}; + mockPrisma = { + webhook: { findFirst: vi.fn().mockResolvedValue({ secret: WEBHOOK_SECRET }) }, + }; // Access private property via type assertion processor = new WebhooksProcessor(mockPrisma as never); fetchSpy = vi.fn(); @@ -69,14 +70,13 @@ describe('WebhooksProcessor', () => { expect(options.headers['x-astroid-delivery']).toBe(EVENT_ID); expect(options.headers['x-astroid-timestamp']).toBeDefined(); expect(options.headers['x-astroid-timestamp']).toMatch(/^\d+$/); + expect(options.headers['x-astroid-signature-version']).toBe('v1'); - // Verify HMAC-SHA256 signature = HMAC(secret, timestamp + body) - const body = options.body; + const body = Buffer.from(options.body); const timestamp = options.headers['x-astroid-timestamp']; - const expectedSignature = createHmac('sha256', WEBHOOK_SECRET) - .update(`${timestamp}.${body}`) - .digest('hex'); - expect(options.headers['x-astroid-signature']).toBe(expectedSignature); + expect(options.headers['x-astroid-signature']).toMatch(/^v1=[0-9a-f]{64}$/); + expect(verifyWebhookSignature(WEBHOOK_SECRET, timestamp, EVENT_ID, body, options.headers['x-astroid-signature'])).toBe(true); + expect(job.data).not.toHaveProperty('secret'); }); it('returns success result with status code', async () => { @@ -234,5 +234,34 @@ describe('WebhooksProcessor', () => { const job = createMockJob({ attemptsMade: 4 } as Partial>); await expect(processor.process(job)).rejects.toThrow('HTTP 503'); }); + + it('keeps the same event identity across retry attempts', async () => { + const upsert = vi.fn().mockResolvedValue({}); + mockPrisma.webhookDelivery = { upsert }; + fetchSpy.mockResolvedValue({ ok: true, status: 200, text: () => Promise.resolve('OK') }); + + const firstAttempt = createMockJob(); + const retryAttempt = createMockJob({ attemptsMade: 1 } as Partial>); + await processor.process(firstAttempt); + await processor.process(retryAttempt); + + expect(firstAttempt.data.eventId).toBe(retryAttempt.data.eventId); + expect(upsert.mock.calls[0][0].where).toEqual(upsert.mock.calls[1][0].where); + }); + + it('does not include a downstream response body in failure messages or logs', async () => { + const secretEcho = `${WEBHOOK_SECRET}:${EVENT_ID}:payload`; + const warn = vi.spyOn(processor['logger'], 'warn').mockImplementation(() => undefined); + const error = vi.spyOn(processor['logger'], 'error').mockImplementation(() => undefined); + fetchSpy.mockResolvedValue({ + ok: false, + status: 500, + text: () => Promise.resolve(secretEcho), + }); + + await expect(processor.process(createMockJob())).rejects.toThrow('HTTP 500'); + expect(warn.mock.calls.flat().join(' ')).not.toContain(secretEcho); + expect(error.mock.calls.flat().join(' ')).not.toContain(secretEcho); + }); }); }); diff --git a/src/modules/webhooks/webhooks.processor.ts b/src/modules/webhooks/webhooks.processor.ts index 73c235d7..ef37c5d2 100644 --- a/src/modules/webhooks/webhooks.processor.ts +++ b/src/modules/webhooks/webhooks.processor.ts @@ -1,12 +1,20 @@ import { Processor, WorkerHost } from '@nestjs/bullmq'; -import { ConfigService } from '@nestjs/config'; import { Inject, Logger, Optional } from '@nestjs/common'; import { Job, UnrecoverableError } from 'bullmq'; import { Queues } from '../../queues/queues.constants'; import { WebhookJobData, WebhookJobResult } from './types/webhook-job.types'; -import { signWebhookPayload } from './utils/signing'; +import { signWebhookPayload, WEBHOOK_SIGNATURE_VERSION } from './utils/signing'; import { PrismaService } from '../../database/prisma.service'; import { WorkerMetricsService } from '../../modules/metrics/worker-metrics.service'; +import { WebhookAuditService } from './services/webhook-audit.service'; +import { + WEBHOOK_DELIVERY_HEADER, + WEBHOOK_EVENT_HEADER, + WEBHOOK_EVENT_ID_HEADER, + WEBHOOK_SIGNATURE_HEADER, + WEBHOOK_SIGNATURE_VERSION_HEADER, + WEBHOOK_TIMESTAMP_HEADER, +} from '../../common/constants/headers'; /** * BullMQ job processor for webhook event delivery with exponential backoff + jitter. @@ -54,48 +62,81 @@ export class WebhooksProcessor extends WorkerHost implements OnModuleDestroy { constructor( @Optional() @Inject(PrismaService) private readonly prisma?: PrismaService, - @Optional() private readonly configService?: ConfigService, @Optional() private readonly workerMetrics?: WorkerMetricsService, + @Optional() private readonly webhookAudit?: WebhookAuditService, ) { super(); } - private resolveSecret(jobSecret?: string): string { - if (jobSecret) return jobSecret; - const fallback = - this.configService?.get('WEBHOOK_SECRET') ?? - this.configService?.get('STELLAR_WEBHOOK_SECRET') ?? - this.configService?.get('WEBHOOK_SIGNING_SECRET') ?? - ''; - return fallback; + /** + * Audit entry for a delivery that will not be retried again: an unrecoverable + * 4xx or the final attempt. `WebhookAuditService` swallows its own failures, so + * this can never mask the original delivery error. + */ + private async auditTerminalFailure( + job: Job, + failedReason: string, + responseStatus?: number, + ): Promise { + if (!this.webhookAudit) return; + try { + await this.webhookAudit.recordTerminalFailure({ + webhookId: job.data.webhookId, + organizationId: job.data.organizationId, + url: job.data.url, + eventName: job.data.eventName, + eventId: job.data.eventId, + attemptsMade: job.attemptsMade + 1, + failedReason, + responseStatus, + }); + } catch (error) { + // Never let compliance bookkeeping mask the original delivery failure. + this.logger.warn( + `Could not audit webhook ${job.data.webhookId} failure: ${(error as Error).message}`, + ); + } + } + + private async resolveSecret(webhookId: string, organizationId: string): Promise { + const client = this.prisma?.workerClient ?? this.prisma; + if (!client) throw new UnrecoverableError('Webhook signing secret is unavailable'); + const webhook = await client.webhook.findFirst({ + where: { id: webhookId, organizationId }, + select: { secret: true }, + }); + if (!webhook?.secret) throw new UnrecoverableError('Webhook signing secret is unavailable'); + return webhook.secret; } async process(job: Job): Promise { const jobName = job.name ?? 'webhook-delivery'; const execute = async (): Promise => { - const { webhookId, organizationId, url, secret, eventName, payload, eventId } = job.data; - this.logger.debug(`Processing webhook ${webhookId} event ${eventName} attempt ${job.attemptsMade + 1}/5`); + const { webhookId, organizationId, url, eventName, payload, eventId, metadata } = job.data; + const requestTrace = metadata?.requestId ? ` requestId=${metadata.requestId}` : ''; + this.logger.debug(`Processing webhook ${webhookId} event ${eventName} attempt ${job.attemptsMade + 1}/5${requestTrace}`); let responseStatus: number | undefined; let errorMessage: string | undefined; let isNonTransient = false; try { - const body = JSON.stringify(payload); + const body = Buffer.from(JSON.stringify(payload), 'utf8'); const timestamp = Math.floor(Date.now() / 1000).toString(); - const effectiveSecret = this.resolveSecret(secret); - const signature = signWebhookPayload(effectiveSecret, timestamp, body); + const effectiveSecret = await this.resolveSecret(webhookId, organizationId); + const signature = signWebhookPayload(effectiveSecret, timestamp, eventId, body); const response = await fetch(url, { method: 'POST', headers: { 'content-type': 'application/json', - 'x-astroid-signature': signature, - 'x-astroid-timestamp': timestamp, - 'x-astroid-delivery': eventId, - 'x-astroid-event': eventName, - 'x-astroid-event-id': eventId, + [WEBHOOK_SIGNATURE_HEADER]: signature, + [WEBHOOK_TIMESTAMP_HEADER]: timestamp, + [WEBHOOK_EVENT_ID_HEADER]: eventId, + [WEBHOOK_DELIVERY_HEADER]: eventId, + [WEBHOOK_EVENT_HEADER]: eventName, + [WEBHOOK_SIGNATURE_VERSION_HEADER]: WEBHOOK_SIGNATURE_VERSION, 'user-agent': 'Astroid-Webhook-Bot/1.0', }, body, @@ -104,10 +145,9 @@ export class WebhooksProcessor extends WorkerHost implements OnModuleDestroy { responseStatus = response.status; if (!response.ok) { - const errorText = await response.text().catch(() => response.statusText); - errorMessage = `HTTP ${response.status}: ${errorText}`; + errorMessage = `HTTP ${response.status}`; isNonTransient = WebhooksProcessor.NON_TRANSIENT_STATUSES.has(response.status); - this.logger.warn(`Webhook ${webhookId} responded ${response.status}: ${errorText}`); + this.logger.warn(`Webhook ${webhookId} responded ${response.status}${requestTrace}`); if (isNonTransient) { await this.persistState({ webhookId, @@ -120,16 +160,19 @@ export class WebhooksProcessor extends WorkerHost implements OnModuleDestroy { lastError: errorMessage, responseStatus, }); + // Non-transient (4xx): record the abandoned delivery before BullMQ + // moves it straight to the failed set. + await this.auditTerminalFailure(job, errorMessage ?? 'HTTP error', responseStatus); throw new UnrecoverableError(errorMessage); } throw new Error(errorMessage); } - this.logger.debug(`Webhook ${webhookId} delivered successfully`); + this.logger.debug(`Webhook ${webhookId} delivered successfully${requestTrace}`); } catch (error) { if (error instanceof UnrecoverableError) throw error; - errorMessage = (error as Error).message; + errorMessage = error instanceof Error ? error.message : 'Delivery attempt failed'; const isLastAttempt = job.attemptsMade >= 4; - this.logger.error(`Webhook ${webhookId} failed attempt ${job.attemptsMade + 1}/5: ${errorMessage}`); + this.logger.error(`Webhook ${webhookId} failed attempt ${job.attemptsMade + 1}/5${requestTrace}`); await this.persistState({ webhookId, organizationId, @@ -142,7 +185,10 @@ export class WebhooksProcessor extends WorkerHost implements OnModuleDestroy { responseStatus, }); if (isLastAttempt) { - this.logger.error(`Webhook ${webhookId} exhausted all retry attempts`); + this.logger.error(`Webhook ${webhookId} exhausted all retry attempts${requestTrace}`); + // Retries are exhausted: the delivery is dead-lettered by the queue + // failure listener, so record it permanently in the audit trail. + await this.auditTerminalFailure(job, errorMessage ?? 'unknown error', responseStatus); } throw error; } diff --git a/src/queues/queue-failure-listener.ts b/src/queues/queue-failure-listener.ts index 47287517..65edbab7 100644 --- a/src/queues/queue-failure-listener.ts +++ b/src/queues/queue-failure-listener.ts @@ -268,8 +268,12 @@ export class QueueFailureListener implements OnModuleInit, OnModuleDestroy { */ private extractTrace(data: unknown): JobTraceContext { const payload = (data && typeof data === 'object' ? data : {}) as Record; + const metadata = + payload.metadata && typeof payload.metadata === 'object' + ? (payload.metadata as Record) + : {}; const read = (key: string): string | undefined => { - const value = payload[key]; + const value = payload[key] ?? metadata[key]; return typeof value === 'string' ? value : undefined; }; diff --git a/src/queues/queues.constants.ts b/src/queues/queues.constants.ts index c2bc761c..6cbd4d43 100644 --- a/src/queues/queues.constants.ts +++ b/src/queues/queues.constants.ts @@ -29,6 +29,12 @@ export const Queues = { export type QueueName = (typeof Queues)[keyof typeof Queues]; +export interface QueueJobMetadata { + requestId?: string; + correlationId?: string; + traceId?: string; +} + /** Standard payload stored when a job is dead-lettered. */ export interface DlqJobData { /** Original queue the job originated from. */ diff --git a/src/queues/webhook.queue.spec.ts b/src/queues/webhook.queue.spec.ts new file mode 100644 index 00000000..7964e112 --- /dev/null +++ b/src/queues/webhook.queue.spec.ts @@ -0,0 +1,56 @@ +import { describe, expect, it } from 'vitest'; + +import { Queues } from './queues.constants'; +import { + WEBHOOK_BACKOFF_BASE_DELAY_MS, + WEBHOOK_DEAD_LETTER_QUEUE, + WEBHOOK_JOB_NAME, + WEBHOOK_MAX_ATTEMPTS, + createWebhookQueueOptions, + webhookJobOptions, +} from './webhook.queue'; + +describe('webhook queue configuration', () => { + it('retries a delivery five times with exponential backoff', () => { + expect(WEBHOOK_MAX_ATTEMPTS).toBe(5); + expect(WEBHOOK_BACKOFF_BASE_DELAY_MS).toBe(2_000); + expect(webhookJobOptions.attempts).toBe(5); + expect(webhookJobOptions.backoff).toEqual({ type: 'exponential', delay: 2_000 }); + }); + + it('keeps completed jobs bounded and failed jobs inspectable for a day', () => { + expect(webhookJobOptions.removeOnComplete).toEqual({ count: 1_000 }); + expect(webhookJobOptions.removeOnFail).toEqual({ age: 24 * 3_600 }); + }); + + it('registers the queue under the webhook name with the shared job options', () => { + const options = createWebhookQueueOptions(); + + expect(options.name).toBe(Queues.Webhooks); + expect(options.defaultJobOptions).toEqual(webhookJobOptions); + }); + + it('attaches a jittered backoff strategy so retries never fire in lockstep', () => { + const settings = createWebhookQueueOptions().settings as unknown as { + backoffStrategy: (attemptsMade: number) => number; + }; + + const firstRetry = settings.backoffStrategy(0); + expect(firstRetry).toBeGreaterThanOrEqual(2_000); + expect(firstRetry).toBeLessThan(2_400); + + // The second retry doubles the base delay while still staying inside the + // 20% jitter envelope. + const secondRetry = settings.backoffStrategy(1); + expect(secondRetry).toBeGreaterThanOrEqual(4_000); + expect(secondRetry).toBeLessThan(4_800); + }); + + it('routes permanently failed deliveries to the dead-letter queue', () => { + expect(WEBHOOK_DEAD_LETTER_QUEUE).toBe(Queues.DeadLetter); + }); + + it('uses a single well-known job name for every delivery', () => { + expect(WEBHOOK_JOB_NAME).toBe('webhook-delivery'); + }); +}); diff --git a/src/queues/webhook.queue.ts b/src/queues/webhook.queue.ts new file mode 100644 index 00000000..9fc3c8ca --- /dev/null +++ b/src/queues/webhook.queue.ts @@ -0,0 +1,54 @@ +import type { JobsOptions } from 'bullmq'; +import type { RegisterQueueOptions } from '@nestjs/bullmq'; + +import { webhookBackoffStrategy } from '../utils/backoff.util'; +import { Queues } from './queues.constants'; + +/** BullMQ job name used for every outbound webhook delivery. */ +export const WEBHOOK_JOB_NAME = 'webhook-delivery'; + +/** Total delivery attempts (1 initial + 4 retries) before a webhook is dead-lettered. */ +export const WEBHOOK_MAX_ATTEMPTS = 5; + +/** Base delay in milliseconds for the exponential backoff between attempts. */ +export const WEBHOOK_BACKOFF_BASE_DELAY_MS = 2_000; + +/** Queue that permanently failed webhook deliveries are routed to for inspection. */ +export const WEBHOOK_DEAD_LETTER_QUEUE = Queues.DeadLetter; + +/** + * Retry policy applied to every queued webhook delivery. + * + * - `attempts: 5` bounds the work spent on a dead endpoint. + * - `backoff: exponential @ 2000ms` spaces retries out (2s, 4s, 8s, 16s). + * - `removeOnFail.age: 24h` keeps exhausted jobs inspectable (and re-drivable) + * long enough for an operator to act, without growing Redis forever. + */ +export const webhookJobOptions: JobsOptions = { + attempts: WEBHOOK_MAX_ATTEMPTS, + backoff: { + type: 'exponential', + delay: WEBHOOK_BACKOFF_BASE_DELAY_MS, + }, + removeOnComplete: { count: 1_000 }, + removeOnFail: { age: 24 * 3_600 }, +}; + +/** + * Builds the BullMQ registration for the webhook queue. + * + * The custom `backoffStrategy` adds ±20% randomized jitter on top of the + * exponential delay so a fleet of failing subscribers is not retried in + * lockstep (thundering herd). BullMQ reads the strategy from + * `queue.opts.settings.backoffStrategy` at retry time; `@nestjs/bullmq` does not + * surface that field on `RegisterQueueOptions`, hence the narrow cast. + */ +export function createWebhookQueueOptions(): RegisterQueueOptions { + return { + name: Queues.Webhooks, + defaultJobOptions: webhookJobOptions, + settings: { + backoffStrategy: webhookBackoffStrategy, + } as RegisterQueueOptions['settings'], + }; +} diff --git a/src/types/http.ts b/src/types/http.ts index 115a8fe6..7c13c1b9 100644 --- a/src/types/http.ts +++ b/src/types/http.ts @@ -20,6 +20,7 @@ export interface ApiSuccessEnvelope { success: true; data: T; meta?: { + offset?: number; page?: number; limit?: number; total?: number; diff --git a/src/utils/crypto.util.spec.ts b/src/utils/crypto.util.spec.ts index 73c28180..31c5204d 100644 --- a/src/utils/crypto.util.spec.ts +++ b/src/utils/crypto.util.spec.ts @@ -1,153 +1,251 @@ -import { describe, it, expect } from 'vitest'; -import { createHmac } from 'crypto'; -import { hmacSign, safeEqual, generateToken, sha256, generateApiKey } from './crypto.util'; +import { describe, expect, it } from 'vitest'; +import { + generateToken, + sha256, + hashWithArgon2, + verifyArgon2, + generateApiKey, + hmacSign, + generateWebhookSignature, + buildWebhookHeaders, + safeEqual, +} from './crypto.util'; describe('crypto.util', () => { - describe('hmacSign', () => { - it('produces a hex-encoded HMAC-SHA256 signature', () => { - const secret = 'whsec_test-secret-key'; - const payload = '{"event":"transaction.completed","data":{}}'; - const signature = hmacSign(secret, payload); + describe('generateToken', () => { + it('should generate a random hex token of specified length', () => { + const token = generateToken(32); + expect(token).toHaveLength(64); // 32 bytes = 64 hex chars + expect(/^[a-f0-9]+$/.test(token)).toBe(true); + }); - // Verify against manual HMAC-SHA256 computation - const expected = createHmac('sha256', secret).update(payload).digest('hex'); - expect(signature).toBe(expected); + it('should generate different tokens on each call', () => { + const token1 = generateToken(16); + const token2 = generateToken(16); + expect(token1).not.toBe(token2); }); - it('returns a 64-character hex string', () => { - const signature = hmacSign('secret', 'payload'); - expect(signature).toMatch(/^[0-9a-f]{64}$/); + it('should use default byte length when not specified', () => { + const token = generateToken(); + expect(token).toHaveLength(64); // default 32 bytes }); + }); - it('produces different signatures for different secrets', () => { - const payload = '{"event":"test"}'; - const sig1 = hmacSign('secret-a', payload); - const sig2 = hmacSign('secret-b', payload); - expect(sig1).not.toBe(sig2); + describe('sha256', () => { + it('should generate consistent SHA-256 hashes', () => { + const hash1 = sha256('test'); + const hash2 = sha256('test'); + expect(hash1).toBe(hash2); }); - it('produces different signatures for different payloads', () => { - const secret = 'same-secret'; - const sig1 = hmacSign(secret, '{"a":1}'); - const sig2 = hmacSign(secret, '{"b":2}'); - expect(sig1).not.toBe(sig2); + it('should generate different hashes for different inputs', () => { + const hash1 = sha256('test1'); + const hash2 = sha256('test2'); + expect(hash1).not.toBe(hash2); }); - it('produces deterministic signatures for the same inputs', () => { - const sig1 = hmacSign('secret', 'payload'); - const sig2 = hmacSign('secret', 'payload'); - expect(sig1).toBe(sig2); + it('should produce fixed-length output', () => { + const hash = sha256('any input'); + expect(hash).toHaveLength(64); // SHA-256 produces 64 hex chars }); + }); - it('handles empty payload', () => { - const signature = hmacSign('secret', ''); - expect(signature).toMatch(/^[0-9a-f]{64}$/); + describe('hashWithArgon2', () => { + it('should generate Argon2id hash for a value', async () => { + const hash = await hashWithArgon2('password123'); + expect(hash).toBeDefined(); + expect(typeof hash).toBe('string'); + expect(hash.length).toBeGreaterThan(0); }); - it('handles Unicode payloads correctly', () => { - const signature = hmacSign('secret', '{"name":"José 🔑"}'); - expect(signature).toMatch(/^[0-9a-f]{64}$/); + it('should generate different hashes for the same input (due to salt)', async () => { + const hash1 = await hashWithArgon2('password123'); + const hash2 = await hashWithArgon2('password123'); + expect(hash1).not.toBe(hash2); }); - it('produces a signature compatible with x-astroid-signature header format', () => { - const secret = 'whsec_abc123'; - const body = JSON.stringify({ event: 'wallet.created', data: { id: 'w-1' } }); - const signature = hmacSign(secret, body); + it('should generate different hashes for different inputs', async () => { + const hash1 = await hashWithArgon2('password123'); + const hash2 = await hashWithArgon2('password456'); + expect(hash1).not.toBe(hash2); + }); - // The signature should be a valid hex string suitable for an HTTP header - expect(typeof signature).toBe('string'); - expect(signature.length).toBe(64); - expect(Buffer.from(signature, 'hex').length).toBe(32); + it('should include Argon2id identifier in hash', async () => { + const hash = await hashWithArgon2('test'); + expect(hash).toMatch(/\$argon2id\$/); }); }); - describe('safeEqual', () => { - it('returns true for identical strings', () => { - expect(safeEqual('abc', 'abc')).toBe(true); + describe('verifyArgon2', () => { + it('should verify correct password against hash', async () => { + const password = 'correct-password'; + const hash = await hashWithArgon2(password); + const isValid = await verifyArgon2(hash, password); + expect(isValid).toBe(true); + }); + + it('should reject incorrect password against hash', async () => { + const password = 'correct-password'; + const hash = await hashWithArgon2(password); + const isValid = await verifyArgon2(hash, 'wrong-password'); + expect(isValid).toBe(false); + }); + + it('should handle invalid hash gracefully', async () => { + const isValid = await verifyArgon2('invalid-hash', 'password'); + expect(isValid).toBe(false); + }); + + it('should use constant-time comparison (timing attack resistant)', async () => { + const password = 'password123'; + const hash = await hashWithArgon2(password); + + // Both should take similar time regardless of result + const start1 = Date.now(); + await verifyArgon2(hash, password); + const time1 = Date.now() - start1; + + const start2 = Date.now(); + await verifyArgon2(hash, 'wrong'); + const time2 = Date.now() - start2; + + // Times should be reasonably close (within 10x due to system variance) + expect(Math.abs(time1 - time2)).toBeLessThan(time1 * 10); }); + }); - it('returns true for identical hex signatures', () => { - const sig = hmacSign('secret', 'payload'); - expect(safeEqual(sig, sig)).toBe(true); + describe('generateApiKey', () => { + it('should generate API key with proper format', async () => { + const apiKey = await generateApiKey('live'); + expect(apiKey.raw).toMatch(/^ak_live_[a-f0-9]+$/); + expect(apiKey.prefix).toHaveLength(14); + expect(apiKey.hashedKey).toBeDefined(); + expect(apiKey.hashedKey.length).toBeGreaterThan(0); }); - it('returns false for different strings', () => { - expect(safeEqual('abc', 'def')).toBe(false); + it('should use different environment prefixes', async () => { + const liveKey = await generateApiKey('live'); + const testKey = await generateApiKey('test'); + expect(liveKey.raw).toMatch(/^ak_live_/); + expect(testKey.raw).toMatch(/^ak_test_/); }); - it('returns false for different-length strings', () => { - expect(safeEqual('abc', 'abcd')).toBe(false); + it('should use default environment when not specified', async () => { + const apiKey = await generateApiKey(); + expect(apiKey.raw).toMatch(/^ak_live_/); + }); + + it('should generate unique keys each time', async () => { + const key1 = await generateApiKey('live'); + const key2 = await generateApiKey('live'); + expect(key1.raw).not.toBe(key2.raw); + expect(key1.hashedKey).not.toBe(key2.hashedKey); }); - it('returns false for completely different signatures', () => { - const sig1 = hmacSign('secret-a', 'payload'); - const sig2 = hmacSign('secret-b', 'payload'); - expect(safeEqual(sig1, sig2)).toBe(false); + it('should use Argon2id for hashing', async () => { + const apiKey = await generateApiKey('live'); + expect(apiKey.hashedKey).toMatch(/\$argon2id\$/); }); - it('handles empty strings', () => { - expect(safeEqual('', '')).toBe(true); - expect(safeEqual('', 'a')).toBe(false); + it('should verify generated key against hash', async () => { + const apiKey = await generateApiKey('live'); + const isValid = await verifyArgon2(apiKey.hashedKey, apiKey.raw); + expect(isValid).toBe(true); }); }); - describe('generateToken', () => { - it('generates a hex-encoded token of expected length', () => { - const token = generateToken(32); - // 32 bytes = 64 hex characters - expect(token).toMatch(/^[0-9a-f]{64}$/); + describe('hmacSign', () => { + it('should generate HMAC-SHA256 signature', () => { + const signature = hmacSign('secret', 'payload'); + expect(signature).toHaveLength(64); // SHA-256 = 64 hex chars + expect(/^[a-f0-9]+$/.test(signature)).toBe(true); }); - it('generates different tokens on each call', () => { - const token1 = generateToken(16); - const token2 = generateToken(16); - expect(token1).not.toBe(token2); + it('should generate consistent signatures for same input', () => { + const sig1 = hmacSign('secret', 'payload'); + const sig2 = hmacSign('secret', 'payload'); + expect(sig1).toBe(sig2); }); - it('respects the byte length parameter', () => { - const token8 = generateToken(8); - const token16 = generateToken(16); - expect(token8.length).toBe(16); - expect(token16.length).toBe(32); + it('should generate different signatures for different secrets', () => { + const sig1 = hmacSign('secret1', 'payload'); + const sig2 = hmacSign('secret2', 'payload'); + expect(sig1).not.toBe(sig2); }); }); - describe('sha256', () => { - it('returns a 64-character hex digest', () => { - const hash = sha256('hello'); - expect(hash).toMatch(/^[0-9a-f]{64}$/); + describe('generateWebhookSignature', () => { + it('should generate webhook signature per spec', () => { + const signature = generateWebhookSignature('secret', '1234567890', '{"data":"test"}'); + expect(signature).toHaveLength(64); + expect(/^[a-f0-9]+$/.test(signature)).toBe(true); }); - it('is deterministic', () => { - expect(sha256('test')).toBe(sha256('test')); + it('should concatenate timestamp and body without delimiter', () => { + const signature = generateWebhookSignature('secret', '123', 'body'); + const manualHmac = hmacSign('secret', '123body'); + expect(signature).toBe(manualHmac); }); + }); - it('produces different hashes for different inputs', () => { - expect(sha256('a')).not.toBe(sha256('b')); + describe('buildWebhookHeaders', () => { + it('should build standard webhook headers', () => { + const headers = buildWebhookHeaders({ + signature: 'abc123', + timestamp: '1234567890', + deliveryId: 'delivery-1', + eventName: 'transaction.created', + }); + + expect(headers['x-astroid-signature']).toBe('abc123'); + expect(headers['x-astroid-timestamp']).toBe('1234567890'); + expect(headers['x-astroid-delivery']).toBe('delivery-1'); + expect(headers['x-astroid-event']).toBe('transaction.created'); + expect(headers['x-astroid-event-id']).toBe('delivery-1'); }); }); - describe('generateApiKey', () => { - it('returns raw key with expected prefix format', () => { - const key = generateApiKey('live'); - expect(key.raw).toMatch(/^ak_live_[0-9a-f]{48}$/); - expect(key.prefix).toBe(key.raw.slice(0, 14)); + describe('safeEqual', () => { + it('should return true for equal strings', () => { + expect(safeEqual('abc123', 'abc123')).toBe(true); }); - it('includes SHA-256 hash of the raw key', () => { - const key = generateApiKey(); - expect(key.hashedKey).toBe(sha256(key.raw)); + it('should return false for different strings', () => { + expect(safeEqual('abc123', 'abc456')).toBe(false); }); - it('uses the specified environment', () => { - const testKey = generateApiKey('test'); - expect(testKey.raw).toMatch(/^ak_test_[0-9a-f]{48}$/); + it('should return false for different length strings', () => { + expect(safeEqual('abc', 'abcd')).toBe(false); }); - it('generates unique keys on each call', () => { - const key1 = generateApiKey(); - const key2 = generateApiKey(); - expect(key1.raw).not.toBe(key2.raw); + it('should use constant-time comparison', () => { + const start1 = Date.now(); + safeEqual('a'.repeat(1000), 'a'.repeat(1000)); + const time1 = Date.now() - start1; + + const start2 = Date.now(); + safeEqual('a'.repeat(1000), 'b'.repeat(1000)); + const time2 = Date.now() - start2; + + // Times should be similar (within reasonable tolerance) + expect(Math.abs(time1 - time2)).toBeLessThan(10); + }); + }); + + describe('backward compatibility', () => { + it('should still support SHA-256 for non-API-key use cases', () => { + const hash = sha256('test-value'); + expect(hash).toHaveLength(64); + expect(/^[a-f0-9]+$/.test(hash)).toBe(true); + }); + + it('should distinguish between Argon2id and SHA-256 hashes', async () => { + const argonHash = await hashWithArgon2('test'); + const shaHash = sha256('test'); + + expect(argonHash).toMatch(/\$argon2id\$/); + expect(shaHash).not.toMatch(/\$argon2id\$/); + expect(argonHash).not.toBe(shaHash); }); }); }); diff --git a/src/utils/crypto.util.ts b/src/utils/crypto.util.ts index d91eedbe..c4dcd4c7 100644 --- a/src/utils/crypto.util.ts +++ b/src/utils/crypto.util.ts @@ -1,8 +1,11 @@ import { createHash, createHmac, randomBytes, timingSafeEqual } from 'crypto'; +import { argon2id, hash, verify } from 'argon2'; /** * Cryptographic helpers used for API keys, webhook signatures and refresh - * tokens. Secrets are never stored in plaintext — only SHA-256 hashes. + * tokens. Secrets are never stored in plaintext — only Argon2id hashes for + * API keys (with SHA-256 fallback for legacy keys) and SHA-256 for other + * non-reversible hashes. */ /** Generates a random URL-safe token of `bytes` entropy (hex-encoded). */ @@ -10,11 +13,44 @@ export function generateToken(bytes = 32): string { return randomBytes(bytes).toString('hex'); } -/** SHA-256 hex digest of a value — used to store non-reversible key hashes. */ +/** SHA-256 hex digest of a value — used for non-reversible key hashes (e.g., refresh tokens). */ export function sha256(value: string): string { return createHash('sha256').update(value).digest('hex'); } +/** + * Argon2id hash of a value — used for API keys for enhanced security. + * Argon2id is memory-hard and resistant to GPU/ASIC attacks. + * + * @param value - The plaintext value to hash + * @returns The Argon2id hash string + */ +export async function hashWithArgon2(value: string): Promise { + return await hash(value, { + type: argon2id, + memoryCost: 65536, // 64 MB + timeCost: 3, // 3 iterations + parallelism: 4, // 4 threads + hashLength: 32, + }); +} + +/** + * Verifies a value against an Argon2id hash. + * Uses constant-time comparison to prevent timing attacks. + * + * @param hash - The stored Argon2id hash + * @param value - The plaintext value to verify + * @returns true if the value matches the hash, false otherwise + */ +export async function verifyArgon2(hash: string, value: string): Promise { + try { + return await verify(hash, value); + } catch { + return false; + } +} + /** Computes an HMAC-SHA256 signature (hex) for webhook payload signing. */ export function hmacSign(secret: string, payload: string): string { return createHmac('sha256', secret).update(payload).digest('hex'); @@ -67,14 +103,18 @@ export interface GeneratedApiKey { raw: string; /** The short prefix stored for identification (e.g. `ak_live_abcd`). */ prefix: string; - /** The SHA-256 hash persisted in the database. */ + /** The Argon2id hash persisted in the database. */ hashedKey: string; } -/** Mints a new API key: `ak__`, returning raw + prefix + hash. */ -export function generateApiKey(environment = 'live'): GeneratedApiKey { +/** + * Mints a new API key: `ak__`, returning raw + prefix + Argon2id hash. + * Uses Argon2id for enhanced security against brute-force and GPU attacks. + */ +export async function generateApiKey(environment = 'live'): Promise { const secret = generateToken(24); const raw = `ak_${environment}_${secret}`; const prefix = raw.slice(0, 14); - return { raw, prefix, hashedKey: sha256(raw) }; + const hashedKey = await hashWithArgon2(raw); + return { raw, prefix, hashedKey }; } diff --git a/src/utils/retry.util.spec.ts b/src/utils/retry.util.spec.ts index 7a3dff02..fd9b3081 100644 --- a/src/utils/retry.util.spec.ts +++ b/src/utils/retry.util.spec.ts @@ -23,9 +23,7 @@ describe('retryWithBackoff', () => { }); it('retries on failure and succeeds on the second attempt', async () => { - const fn = vi.fn() - .mockRejectedValueOnce(new Error('transient')) - .mockResolvedValue('ok'); + const fn = vi.fn().mockRejectedValueOnce(new Error('transient')).mockResolvedValue('ok'); const promise = retryWithBackoff(fn, { maxAttempts: 3, baseDelayMs: 10 }); await vi.runAllTimersAsync(); @@ -38,6 +36,9 @@ describe('retryWithBackoff', () => { const fn = vi.fn().mockRejectedValue(boom); const promise = retryWithBackoff(fn, { maxAttempts: 3, baseDelayMs: 10 }); + // Attach a handler immediately so the rejection isn't flagged as unhandled + // while `runAllTimersAsync` drives the retry loop forward below. + promise.catch(() => {}); await vi.runAllTimersAsync(); await expect(promise).rejects.toBe(boom); expect(fn).toHaveBeenCalledTimes(3); @@ -50,6 +51,7 @@ describe('retryWithBackoff', () => { !(err instanceof Error && err.message.includes('NOT NULL')); const promise = retryWithBackoff(fn, { maxAttempts: 5, isRetryable }); + promise.catch(() => {}); await vi.runAllTimersAsync(); await expect(promise).rejects.toBe(nonRetryable); expect(fn).toHaveBeenCalledTimes(1); @@ -57,7 +59,8 @@ describe('retryWithBackoff', () => { it('calls onRetry before each retry sleep', async () => { const onRetry = vi.fn(); - const fn = vi.fn() + const fn = vi + .fn() .mockRejectedValueOnce(new Error('t1')) .mockRejectedValueOnce(new Error('t2')) .mockResolvedValue('ok'); @@ -75,9 +78,7 @@ describe('retryWithBackoff', () => { const { exponentialBackoffWithJitter } = await import('./backoff.util'); (exponentialBackoffWithJitter as ReturnType).mockReturnValue(60_000); - const fn = vi.fn() - .mockRejectedValueOnce(new Error('t')) - .mockResolvedValue('ok'); + const fn = vi.fn().mockRejectedValueOnce(new Error('t')).mockResolvedValue('ok'); const onRetry = vi.fn(); const promise = retryWithBackoff(fn, { @@ -96,6 +97,7 @@ describe('retryWithBackoff', () => { const onRetry = vi.fn(); const promise = retryWithBackoff(fn, { maxAttempts: 2, baseDelayMs: 10, onRetry }); + promise.catch(() => {}); await vi.runAllTimersAsync(); await expect(promise).rejects.toThrow(); diff --git a/src/workers/job-worker.spec.ts b/src/workers/job-worker.spec.ts index 8005666d..74db9251 100644 --- a/src/workers/job-worker.spec.ts +++ b/src/workers/job-worker.spec.ts @@ -189,9 +189,14 @@ describe('runWorkerJob', () => { expect(record).toMatchObject({ event: 'job.dead-lettered', unrecoverable: true }); }); - it('includes trace fields from the job payload', async () => { + it('includes trace fields from top-level and nested job metadata', async () => { const job = makeJob( - { organizationId: 'org-1', traceId: 'trace-abc', extra: 'noise' }, + { + organizationId: 'org-1', + metadata: { requestId: 'req-123', correlationId: 'corr-123' }, + traceId: 'trace-abc', + extra: 'noise', + }, { attemptsMade: 2, opts: { attempts: 3 } }, ); @@ -207,7 +212,12 @@ describe('runWorkerJob', () => { ).rejects.toThrow(); const record = JSON.parse(String(logger.error.mock.calls[0][0])); - expect(record.trace).toEqual({ organizationId: 'org-1', traceId: 'trace-abc' }); + expect(record.trace).toEqual({ + organizationId: 'org-1', + requestId: 'req-123', + correlationId: 'corr-123', + traceId: 'trace-abc', + }); expect(record.trace.extra).toBeUndefined(); }); diff --git a/src/workers/job-worker.ts b/src/workers/job-worker.ts index b9cc55f7..72ed023e 100644 --- a/src/workers/job-worker.ts +++ b/src/workers/job-worker.ts @@ -87,6 +87,7 @@ export async function runWorkerJob( attempt, maxAttempts, durationMs: Date.now() - startedAt, + trace: extractTrace(job.data), }); try { @@ -120,7 +121,6 @@ export async function runWorkerJob( ...(unrecoverable ? { unrecoverable: true } : {}), error: described, payload: scrubForLog(job.data), - trace: extractTrace(job.data), timestamp: new Date().toISOString(), }; @@ -166,9 +166,13 @@ function describeError(error: unknown): { name: string; message: string; stack?: function extractTrace(data: unknown): Record | undefined { if (!data || typeof data !== 'object') return undefined; const payload = data as Record; + const metadata = + payload.metadata && typeof payload.metadata === 'object' + ? (payload.metadata as Record) + : {}; const trace: Record = {}; for (const key of TRACE_KEYS) { - const value = payload[key]; + const value = payload[key] ?? metadata[key]; if (typeof value === 'string') trace[key] = value; } return Object.keys(trace).length ? trace : undefined;