diff --git a/.agents/AGENTS.md b/.agents/AGENTS.md index 01913a1f..9857a167 100644 --- a/.agents/AGENTS.md +++ b/.agents/AGENTS.md @@ -4,3 +4,6 @@ - Use proper separation of concerns (e.g. TanStack layout routes instead of polluting `__root.tsx`). - Implement robust error handling, proper typing, and scalable folder structures from the very first commit. - Never use anti-patterns to save time. If a proper implementation takes more steps, take the time to do it right. + +# Package Manager +- ALWAYS use `pnpm` for frontend/Node.js package management instead of `npm`. Do not use `npm install`. diff --git a/Makefile b/Makefile index 57f8f25a..fd2db062 100644 --- a/Makefile +++ b/Makefile @@ -45,19 +45,19 @@ dev: dev-as2: @echo "Starting AS2 Server with hot-reload for local development..." - ENVIRONMENT=development uv run uvicorn as2_server.main:app --reload --port 8000 + ENVIRONMENT=development uv run uvicorn as2_server.main:app --reload --host 0.0.0.0 --port 8000 dev-api: @echo "Starting API Gateway with hot-reload for local development..." - ENVIRONMENT=development uv run uvicorn api.main:app --reload --port 8001 + ENVIRONMENT=development uv run uvicorn api.main:app --reload --host 0.0.0.0 --port 8001 dev-web: @echo "Starting React Frontend with Vite..." cd frontend/web && pnpm dev dev-worker: - @echo "Starting Provision Worker for local development..." - ENVIRONMENT=development PYTHONPATH=services/worker/src:libs/database/src:libs/config/src:libs/pipeline/src uv run python services/worker/src/worker/provision/main.py + @echo "Starting Unified Worker (Data + Provision) for local development..." + ENVIRONMENT=development PYTHONPATH=services/worker/src:libs/database/src:libs/config/src:libs/pipeline/src:libs/domain/src:libs/transformer/src uv run python services/worker/src/worker/main.py db-init: @echo "Waiting for databases to be ready..." @@ -69,16 +69,23 @@ db-init: db-reset: @echo "Wiping application databases (leaving Zitadel intact)..." - docker compose stop postgres_global postgres_shard_1 - docker compose rm -f -v postgres_global postgres_shard_1 - -docker volume rm $$(docker volume ls -q | grep -E "postgres_global_data|postgres_shard_[0-9]+_data") 2>/dev/null - @echo "Restarting application databases..." - docker compose up -d postgres_global postgres_shard_1 + docker compose stop postgres_global postgres_shard_1 debezium_shard_1 + docker compose rm -f -v postgres_global postgres_shard_1 debezium_shard_1 + -docker volume rm $$(docker volume ls -q | grep -E "postgres_global_data|postgres_shard_[0-9]+_data|debezium_data") 2>/dev/null + @echo "Restarting application databases and Debezium..." + docker compose up -d postgres_global postgres_shard_1 debezium_shard_1 @echo "Waiting for databases to initialize..." sleep 5 @echo "Re-running migrations and seeding..." $(MAKE) db-init +sqs-purge: + @echo "Purging all LocalStack SQS queues..." + uv run python scripts/purge_sqs.py + +db-sqs-reset: db-reset sqs-purge + @echo "Database and SQS queues have been completely reset." + seed: db-init # --- Docker Infrastructure (Postgres, LocalStack, OTel) --- diff --git a/TECHNICAL_DEBT.md b/TECHNICAL_DEBT.md index 59cbb7a1..b7df9f1e 100644 --- a/TECHNICAL_DEBT.md +++ b/TECHNICAL_DEBT.md @@ -1,10 +1,39 @@ -# Technical Debt +# Technical Debt Register -A living record of known gaps, shortcuts, and missing enterprise capabilities. +This document tracks identified technical debt, proposed refactoring, and estimated effort to resolve. ---- +## 1. Modernization of Bots Core (Stateless AST Migration) -## AS2 Protocol +**Description:** +The open-source `bots_core` library heavily relies on a legacy, stateful `Node` class pattern to build, traverse, and count EDI segments (e.g., `SE01`, `GE01` counting via `node.getcount()`). Currently, we are bypassing this internal engine completely in our modern microservices architecture by injecting raw stateless Python JSON dictionaries (AST format) into the `bots_core.facade`. + +Because `bots_core` doesn't natively parse our stateless AST for trailing counts, we temporarily built an `ASTUtils.count_segments()` workaround in the `transformer` domain layer. + +**Proposed Resolution:** +To achieve the long-term vision of completely merging and modernizing `bots_core` as a first-class citizen of our modern Python stack: +1. Refactor the inner workings of `bots_core` (specifically `message.py`, `outmessage.py`, and `node.py`) to natively accept, traverse, and serialize stateless dictionary ASTs without forcing instantiation of legacy `Node` objects. +2. Move AST utility functions (like `count_segments`) natively into `bots_core/domain/ast_utils.py`. +3. Eliminate the legacy mapping engine scripts (`inn.get()`, `out.put()`) internally inside `bots_core` if they are fully deprecated by the API gateway. + +**Estimated Effort:** Medium-High +**Estimated Time:** 1 to 2 Sprint Weeks (40 - 80 hours) +**Impact:** Will result in a completely modernized, lightning-fast, fully stateless fork of the `bots` EDI engine perfectly aligned with our cloud-native event-driven architecture. + +## 2. Outbox Sweeper (CDC Fallback Relay) + +**Description:** +The system currently relies exclusively on Debezium (CDC) reading the PostgreSQL Write-Ahead Log (WAL) to route `Outbox` events to SQS. If Debezium crashes, loses offsets, or experiences network partitioning, `PENDING` outbox events will be permanently trapped in the database, breaking the asynchronous event pipeline. + +**Proposed Resolution:** +Implement an Outbox Sweeper background worker that acts as a robust enterprise fallback and garbage collector: +1. **Fallback Poller:** A cron/scheduled task that periodically queries `SELECT * FROM outbox WHERE status = 'PENDING'` for events older than a configured threshold (e.g., 60 seconds) and manually relays them to SQS. +2. **Garbage Collector:** A cleanup task that runs `DELETE FROM outbox WHERE status = 'COMPLETED'` for events older than 7 days to prevent unbounded database growth. + +**Estimated Effort:** Low +**Estimated Time:** 1 to 2 days +**Impact:** Essential for enterprise-grade high availability. Guarantees no messages are ever lost due to CDC infrastructure failures and keeps the database optimized over time. + +## 3. AS2 Protocol ### Async MDN — Inbound Callback Not Implemented **Priority:** High @@ -25,11 +54,17 @@ A living record of known gaps, shortcuts, and missing enterprise capabilities. store the sent `Message-ID` and `MIC` so they can be matched when the callback arrives. - Database: Add an `outbound_mdn_pending` table `(message_id, mic, trace_id, expires_at)`. ---- - -## Testing +## 4. Testing ### No Frontend Test Runner (Vitest) + **Priority:** Medium `make test` skips frontend tests with a placeholder comment. React component tests and TanStack Query mutation tests are not covered. + +## 5. Database Schema + +### Missing `edi_headers` Table + +**Priority:** Medium +**Description:** We currently lack an `edi_headers` table to store extracted EDI header metadata (e.g. ST/GS segments). This table needs to be created and linked via foreign key to the `outbound_route` table so that EDI messages can be properly tracked and correlated with their configured outbound routes. diff --git a/docker-compose.yml b/docker-compose.yml index 5454067a..04a7ae04 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,17 +1,7 @@ -version: "3.9" - - # Include the observability and identity stacks +version: '3.9' include: - # - libs/observability/docker/docker-compose.yml - - libs/identity/docker/docker-compose.yml - +- libs/identity/docker/docker-compose.yml services: - - # --------------------------------------------------------------------------- - # Postgres Databases (Hybrid Multi-Tenancy Architecture) - # --------------------------------------------------------------------------- - - # 1. Global Control Plane DB postgres_global: image: postgres:15-alpine container_name: edi_postgres_global @@ -20,16 +10,16 @@ services: POSTGRES_PASSWORD: edi_password POSTGRES_DB: edi_global ports: - - "5432:5432" + - 5432:5432 volumes: - - postgres_global_data:/var/lib/postgresql/data + - postgres_global_data:/var/lib/postgresql/data healthcheck: - test: ["CMD-SHELL", "pg_isready -U edi"] + test: + - CMD-SHELL + - pg_isready -U edi interval: 5s timeout: 5s retries: 5 - - # 2. Standard Tier (Pooled Shard) DB postgres_shard_1: image: postgres:15-alpine container_name: edi_postgres_shard_1 @@ -37,110 +27,93 @@ services: POSTGRES_USER: edi POSTGRES_PASSWORD: edi_password POSTGRES_DB: edi_shard_1 - command: ["postgres", "-c", "wal_level=logical"] + command: + - postgres + - -c + - wal_level=logical ports: - - "5433:5432" + - 5433:5432 volumes: - - postgres_shard_1_data:/var/lib/postgresql/data + - postgres_shard_1_data:/var/lib/postgresql/data healthcheck: - test: ["CMD-SHELL", "pg_isready -U edi"] + test: + - CMD-SHELL + - pg_isready -U edi interval: 5s timeout: 5s retries: 5 - - # --------------------------------------------------------------------------- - # LocalStack (AWS Cloud Emulator for S3, SQS, SNS, KMS, etc.) - # --------------------------------------------------------------------------- localstack: image: localstack/localstack:3.8.0 container_name: edi_localstack ports: - - "4566:4566" # LocalStack Gateway - - "4510-4559:4510-4559" # external services port range + - 4566:4566 + - 4510-4559:4510-4559 environment: - - DEBUG=${DEBUG:-0} - - DOCKER_HOST=unix:///var/run/docker.sock + - DEBUG=${DEBUG:-0} + - DOCKER_HOST=unix:///var/run/docker.sock volumes: - - "${LOCALSTACK_VOLUME_DIR:-./volume}:/var/lib/localstack" - - "/var/run/docker.sock:/var/run/docker.sock" - # Auto-init scripts are executed when LocalStack boots - - ./docker/localstack/init-aws.sh:/etc/localstack/init/ready.d/init-aws.sh + - ${LOCALSTACK_VOLUME_DIR:-./volume}:/var/lib/localstack + - /var/run/docker.sock:/var/run/docker.sock + - ./docker/localstack/init-aws.sh:/etc/localstack/init/ready.d/init-aws.sh healthcheck: - test: ["CMD", "curl", "-f", "http://localhost:4566/_localstack/health"] + test: + - CMD + - curl + - -f + - http://localhost:4566/_localstack/health interval: 5s timeout: 5s retries: 5 - - # --------------------------------------------------------------------------- - # HashiCorp Vault (Secrets & PKI) - # --------------------------------------------------------------------------- vault: image: hashicorp/vault:1.15 container_name: edi_vault ports: - - "127.0.0.1:8200:8200" + - 127.0.0.1:8200:8200 environment: - VAULT_DEV_ROOT_TOKEN_ID: "${VAULT_DEV_ROOT_TOKEN_ID}" - VAULT_DEV_LISTEN_ADDRESS: "0.0.0.0:8200" + VAULT_DEV_ROOT_TOKEN_ID: ${VAULT_DEV_ROOT_TOKEN_ID:-root} + VAULT_DEV_LISTEN_ADDRESS: 0.0.0.0:8200 cap_add: - - IPC_LOCK + - IPC_LOCK healthcheck: - test: ["CMD", "vault", "status", "-address=http://127.0.0.1:8200"] + test: + - CMD + - vault + - status + - -address=http://127.0.0.1:8200 interval: 5s timeout: 5s retries: 5 - -# # --------------------------------------------------------------------------- -# # EDI AS2 Server (FastAPI) -# # --------------------------------------------------------------------------- -# as2-server: -# build: -# context: . -# dockerfile: soopaedi/services/as2_server/Dockerfile -# container_name: edi_as2_server -# ports: -# - "8000:8000" -# environment: -# - APP_ENV=development -# - LOG_LEVEL=DEBUG -# - DB_URL=postgresql+asyncpg://edi:edi_password@postgres:5432/edi -# - S3_BUCKET=edi-as2-payloads -# - S3_ENDPOINT_URL=http://localstack:4566 -# - S3_REGION=us-east-1 -# - S3_ACCESS_KEY_ID=test -# - S3_SECRET_ACCESS_KEY=test -# - OTEL_SERVICE_NAME=edi-as2-server -# - OTEL_EXPORTER_OTLP_ENDPOINT=http://otel-collector:4317 -# depends_on: -# postgres: -# condition: service_healthy -# localstack: -# condition: service_healthy - - # --------------------------------------------------------------------------- - # Debezium CDC Server - # --------------------------------------------------------------------------- - debezium: + cdcproxy: + image: nginx:alpine + container_name: edi_cdcproxy + ports: + - 8002:80 + command: "/bin/sh -c \"echo '\nserver {\n listen 80;\n location / {\n \ + \ proxy_pass http://host.docker.internal:8001;\n proxy_set_header\ + \ Connection \\\"\\\";\n proxy_set_header Upgrade \\\"\\\";\n }\n\ + }' > /etc/nginx/conf.d/default.conf && nginx -g 'daemon off;'\"" + extra_hosts: + - host.docker.internal:host-gateway + debezium_shard_1: image: quay.io/debezium/server:2.7.3.Final - container_name: edi_debezium + container_name: edi_debezium_shard_1 environment: - - DEBEZIUM_SOURCE_DATABASE_HOSTNAME=postgres_shard_1 - - DEBEZIUM_SOURCE_DATABASE_PORT=5432 - - DEBEZIUM_SOURCE_DATABASE_USER=edi - - DEBEZIUM_SOURCE_DATABASE_PASSWORD=edi_password - - DEBEZIUM_SOURCE_DATABASE_DBNAME=edi_shard_1 - - DEBEZIUM_SINK_URL=http://host.docker.internal:8001/internal/cdc/relay + - DEBEZIUM_SOURCE_DATABASE_HOSTNAME=postgres_shard_1 + - DEBEZIUM_SOURCE_DATABASE_PORT=5432 + - DEBEZIUM_SOURCE_DATABASE_USER=edi + - DEBEZIUM_SOURCE_DATABASE_PASSWORD=edi_password + - DEBEZIUM_SOURCE_DATABASE_DBNAME=edi_shard_1 + - DEBEZIUM_SINK_URL=http://cdcproxy/internal/cdc/relay extra_hosts: - - "host.docker.internal:host-gateway" + - host.docker.internal:host-gateway volumes: - - ./docker/debezium:/debezium/conf - - debezium_data:/debezium/data + - ./docker/debezium/application.properties:/debezium/conf/application.properties + - debezium_data:/debezium/data depends_on: postgres_shard_1: condition: service_healthy - volumes: - postgres_global_data: - postgres_shard_1_data: - postgres_enterprise_1_data: - debezium_data: + postgres_global_data: null + postgres_shard_1_data: null + postgres_enterprise_1_data: null + debezium_data: null diff --git a/docker/debezium/application-global.properties b/docker/debezium/application-global.properties new file mode 100644 index 00000000..68aeaeee --- /dev/null +++ b/docker/debezium/application-global.properties @@ -0,0 +1,32 @@ +# ── Source: Postgres WAL ────────────────────────────────────────────────────── +debezium.source.connector.class=io.debezium.connector.postgresql.PostgresConnector +debezium.source.topic.prefix=soopa_global +debezium.source.database.hostname=${DEBEZIUM_SOURCE_DATABASE_HOSTNAME:postgres_global} +debezium.source.database.port=${DEBEZIUM_SOURCE_DATABASE_PORT:5432} +debezium.source.database.user=${DEBEZIUM_SOURCE_DATABASE_USER:edi} +debezium.source.database.password=${DEBEZIUM_SOURCE_DATABASE_PASSWORD:edi_password} +debezium.source.database.dbname=${DEBEZIUM_SOURCE_DATABASE_DBNAME:edi_global} +debezium.source.plugin.name=pgoutput +debezium.source.publication.name=edi_global_cdc +debezium.source.slot.name=edi_global_slot + +debezium.source.table.include.list=.*\\.outbox +debezium.source.snapshot.mode=initial + +# ── Sink: HTTP → FastAPI CdcRelay ─────────────────────────────────── +debezium.sink.type=http +debezium.sink.http.url=${DEBEZIUM_SINK_URL:http://host.docker.internal:8001/internal/cdc/relay} +debezium.sink.http.timeout.ms=10000 +debezium.sink.http.retry.count=10 +debezium.sink.http.retry.delay.ms=500 +debezium.format.value=json +debezium.format.value.schemas.enable=false +debezium.format.key=json +debezium.transforms=unwrap +debezium.transforms.unwrap.type=io.debezium.transforms.ExtractNewRecordState +debezium.transforms.unwrap.drop.tombstones=true +debezium.transforms.unwrap.add.fields=table,schema,op + +debezium.source.offset.storage=org.apache.kafka.connect.storage.FileOffsetBackingStore +debezium.source.offset.storage.file.filename=/debezium/data/offsets-global.dat +debezium.source.offset.flush.interval.ms=5000 diff --git a/docker/localstack/init-aws.sh b/docker/localstack/init-aws.sh old mode 100644 new mode 100755 index 9d9b71da..5d766fd3 --- a/docker/localstack/init-aws.sh +++ b/docker/localstack/init-aws.sh @@ -14,4 +14,18 @@ DLQ_ARN=$(awslocal sqs get-queue-attributes --queue-url http://localhost:4566/00 # Create main queue with redrive policy awslocal sqs create-queue --queue-name EdiTransformerQueue --attributes "{\"RedrivePolicy\":\"{\\\"deadLetterTargetArn\\\":\\\"$DLQ_ARN\\\",\\\"maxReceiveCount\\\":\\\"3\\\"}\"}" +# Create Data Plane CDC Queues and DLQs +awslocal sqs create-queue --queue-name TranslateQueue-DLQ +TRANSLATE_DLQ_ARN=$(awslocal sqs get-queue-attributes --queue-url http://localhost:4566/000000000000/TranslateQueue-DLQ --attribute-names QueueArn --query 'Attributes.QueueArn' --output text) +awslocal sqs create-queue --queue-name TranslateQueue --attributes "{\"RedrivePolicy\":\"{\\\"deadLetterTargetArn\\\":\\\"$TRANSLATE_DLQ_ARN\\\",\\\"maxReceiveCount\\\":\\\"3\\\"}\"}" + +awslocal sqs create-queue --queue-name DeliverQueue-DLQ +DELIVER_DLQ_ARN=$(awslocal sqs get-queue-attributes --queue-url http://localhost:4566/000000000000/DeliverQueue-DLQ --attribute-names QueueArn --query 'Attributes.QueueArn' --output text) +awslocal sqs create-queue --queue-name DeliverQueue --attributes "{\"RedrivePolicy\":\"{\\\"deadLetterTargetArn\\\":\\\"$DELIVER_DLQ_ARN\\\",\\\"maxReceiveCount\\\":\\\"3\\\"}\"}" + +# Create Control Plane CDC Queues and DLQs +awslocal sqs create-queue --queue-name ProvisioningQueue-DLQ +PROVISIONING_DLQ_ARN=$(awslocal sqs get-queue-attributes --queue-url http://localhost:4566/000000000000/ProvisioningQueue-DLQ --attribute-names QueueArn --query 'Attributes.QueueArn' --output text) +awslocal sqs create-queue --queue-name ProvisioningQueue --attributes "{\"RedrivePolicy\":\"{\\\"deadLetterTargetArn\\\":\\\"$PROVISIONING_DLQ_ARN\\\",\\\"maxReceiveCount\\\":\\\"3\\\"}\"}" + echo "LocalStack Initialization Complete." diff --git a/frontend/web/package.json b/frontend/web/package.json index 568d50f3..b757828b 100644 --- a/frontend/web/package.json +++ b/frontend/web/package.json @@ -19,6 +19,7 @@ "@radix-ui/react-radio-group": "^1.4.2", "@radix-ui/react-select": "^2.3.1", "@radix-ui/react-slot": "^1.3.0", + "@radix-ui/react-tabs": "^1.1.17", "@radix-ui/react-toast": "^1.2.18", "@tanstack/react-query": "^5.101.1", "@tanstack/react-router": "^1.170.16", @@ -27,6 +28,7 @@ "class-variance-authority": "^0.7.1", "clsx": "^2.1.1", "cmdk": "^1.1.1", + "date-fns": "^4.4.0", "lucide-react": "^1.21.0", "next-themes": "^0.4.6", "oidc-client-ts": "^3.5.0", diff --git a/frontend/web/pnpm-lock.yaml b/frontend/web/pnpm-lock.yaml index 44761fcd..2ff47dc2 100644 --- a/frontend/web/pnpm-lock.yaml +++ b/frontend/web/pnpm-lock.yaml @@ -35,6 +35,9 @@ importers: '@radix-ui/react-slot': specifier: ^1.3.0 version: 1.3.0(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-tabs': + specifier: ^1.1.17 + version: 1.1.17(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) '@radix-ui/react-toast': specifier: ^1.2.18 version: 1.2.18(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) @@ -59,6 +62,9 @@ importers: cmdk: specifier: ^1.1.1 version: 1.1.1(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) + date-fns: + specifier: ^4.4.0 + version: 4.4.0 lucide-react: specifier: ^1.21.0 version: 1.21.0(react@19.2.7) @@ -320,6 +326,9 @@ packages: '@radix-ui/primitive@1.1.4': resolution: {integrity: sha512-7AdCK9PQyiljKoBDbN8OuctCbd/esdwZPQ8RtOE3SsyQtUpiPb+ND75q0jEhC1m1ecBI0MFNeLJvwIh9iKHRcQ==} + '@radix-ui/primitive@1.1.5': + resolution: {integrity: sha512-d86WIWFYNtGA0H/d8exstrTRTp7eWJYlYJbtNofxr/3ljupZYn6EFDG/Qgu/0Kc8v7yMUxySagqJsL1+PdYjWg==} + '@radix-ui/react-arrow@1.1.10': resolution: {integrity: sha512-j2VTDz1vgCsmuG0k5lBfOcM8n5JPFqZBcMryasFjHYMhwxYL5SRUV5lMSUpRdNtw3D/Sv8pzJtrlAgkssYSsQQ==} peerDependencies: @@ -372,6 +381,19 @@ packages: '@types/react-dom': optional: true + '@radix-ui/react-collection@1.1.12': + resolution: {integrity: sha512-nb67INpE0IahJKN7EYPp9m9YGwYeKlnzxT3MwXVkgCskaSJia97kG4T0ywpjNUSSnoJk/uvk12V8vbrEHEj+/Q==} + peerDependencies: + '@types/react': '*' + '@types/react-dom': '*' + react: ^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc + react-dom: ^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc + peerDependenciesMeta: + '@types/react': + optional: true + '@types/react-dom': + optional: true + '@radix-ui/react-compose-refs@1.1.3': resolution: {integrity: sha512-rYOP8OMnuuPMQF1uhPVlGNcCDlkokKqGFE3JcxFViIkAXP7EvFWUliJAstrapypaBLJNHbZL6jGhbVDGTwmVhA==} peerDependencies: @@ -390,6 +412,15 @@ packages: '@types/react': optional: true + '@radix-ui/react-context@1.2.0': + resolution: {integrity: sha512-fOE+JtN9rygNZkCnHRBEP0TAvLldlhyOxMsbwFvTP4nAs+nBmfnna+o/Zski2wkmY1YMrFC0aSzsHoLY47iLrg==} + peerDependencies: + '@types/react': '*' + react: ^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc + peerDependenciesMeta: + '@types/react': + optional: true + '@radix-ui/react-dialog@1.1.17': resolution: {integrity: sha512-TDTYmpdq8dI2+Xgvgj9AJ8Ghqq+Eph/TRVEdaFQPDItIY+6QSkU7MJMeevw1568Yw/2Ijz8BTphPSP2XejKphw==} peerDependencies: @@ -599,6 +630,19 @@ packages: '@types/react-dom': optional: true + '@radix-ui/react-presence@1.1.7': + resolution: {integrity: sha512-zBZ4QM5XG3JRanDmqXYf3MD6th4AFXFmgU6KNMFzUaV6F3uw9I5/zjMUvFriSEn5ewo1nxuibvyxJdmLlDcslA==} + peerDependencies: + '@types/react': '*' + '@types/react-dom': '*' + react: ^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc + react-dom: ^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc + peerDependenciesMeta: + '@types/react': + optional: true + '@types/react-dom': + optional: true + '@radix-ui/react-primitive@2.1.6': resolution: {integrity: sha512-wetd0QI77DbvrPpTAvH1SqOxsYF2wZe5TNxqwOd5Ty4XDpV3dpV0s8K/1MGMJBeY5o7lg8ub5VIt1Ub+yVen6g==} peerDependencies: @@ -664,6 +708,19 @@ packages: '@types/react-dom': optional: true + '@radix-ui/react-roving-focus@1.1.15': + resolution: {integrity: sha512-40svmmugfM3mUN7VUDGVE1tQGOhyi8enlGD0CNJEcMM36C1f71PKM21DFgNHUfem0XnA+d8H8oN3Z9ZpJjSslg==} + peerDependencies: + '@types/react': '*' + '@types/react-dom': '*' + react: ^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc + react-dom: ^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc + peerDependenciesMeta: + '@types/react': + optional: true + '@types/react-dom': + optional: true + '@radix-ui/react-select@2.3.1': resolution: {integrity: sha512-w6eDvY78LE9ZUiNnXCA1QVK8RYN7k9galFv09kjVydJqBAgHd7Y9A6h0UJ/6DCZNGZMZrB2ohcSW1Bo9d8+wWA==} peerDependencies: @@ -686,6 +743,19 @@ packages: '@types/react': optional: true + '@radix-ui/react-tabs@1.1.17': + resolution: {integrity: sha512-nRyXnrAVCwjeXcHbvEbLS6ndbTeKHG1RqCP4A8Gw5L4cemDzPXdD8rAmr6wet0v57R69wGvuIIsFjHSVkZiMzQ==} + peerDependencies: + '@types/react': '*' + '@types/react-dom': '*' + react: ^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc + react-dom: ^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc + peerDependenciesMeta: + '@types/react': + optional: true + '@types/react-dom': + optional: true + '@radix-ui/react-toast@1.2.18': resolution: {integrity: sha512-YNEnTHV47hPep+U0QvVM02OJNka9uygREc+k4Nh5VSZBg4MmE+myI442x3hCGfRpX7N2WSSYSJKws4gE+Z8lgg==} peerDependencies: @@ -735,6 +805,15 @@ packages: '@types/react': optional: true + '@radix-ui/react-use-is-hydrated@0.1.1': + resolution: {integrity: sha512-qwOiz4Tjo8CNnrOLAYUMXeZwDzXgXpvK4TKQPmWLECM9XoWvA6+0Z2/7Ag3A4ivjS4ovbLJPbskkxioFyBhr8A==} + peerDependencies: + '@types/react': '*' + react: ^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc + peerDependenciesMeta: + '@types/react': + optional: true + '@radix-ui/react-use-layout-effect@1.1.2': resolution: {integrity: sha512-jrBWOxZITuGcnjRCM2t2U5ZPkCLxD+Ym6DjfssS5haTj2iiak/DOb64JeN6OdLfLgptb6/e2kKR+ZuTrGoZTPA==} peerDependencies: @@ -1064,6 +1143,9 @@ packages: csstype@3.2.3: resolution: {integrity: sha512-z1HGKcYy2xA8AGQfwrn0PAy+PB7X/GSj3UVJW9qKyn43xWa+gl5nXmU4qqLMRzWVLFC8KusUX8T/0kCiOYpAIQ==} + date-fns@4.4.0: + resolution: {integrity: sha512-+1UMbeh68lH1SegH83CGWwpb6OHHbpSgr3+s5Eww5M4CAgswBpoWS0AjTOfEJ33HiYKz1hdj/KTFprzXHmq/6w==} + debug@4.4.3: resolution: {integrity: sha512-RGwwWnwQvkVfavKVt22FGLw+xYSdzARwm0ru6DhTVA3umU5hZc28V3kO4stgYryrTlLpuvgI9GiijltAjNbcqA==} engines: {node: '>=6.0'} @@ -1849,6 +1931,8 @@ snapshots: '@radix-ui/primitive@1.1.4': {} + '@radix-ui/primitive@1.1.5': {} + '@radix-ui/react-arrow@1.1.10(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)': dependencies: '@radix-ui/react-primitive': 2.1.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) @@ -1891,6 +1975,18 @@ snapshots: '@types/react': 19.2.17 '@types/react-dom': 19.2.3(@types/react@19.2.17) + '@radix-ui/react-collection@1.1.12(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)': + dependencies: + '@radix-ui/react-compose-refs': 1.1.3(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-context': 1.2.0(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-primitive': 2.1.7(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) + '@radix-ui/react-slot': 1.3.0(@types/react@19.2.17)(react@19.2.7) + react: 19.2.7 + react-dom: 19.2.7(react@19.2.7) + optionalDependencies: + '@types/react': 19.2.17 + '@types/react-dom': 19.2.3(@types/react@19.2.17) + '@radix-ui/react-compose-refs@1.1.3(@types/react@19.2.17)(react@19.2.7)': dependencies: react: 19.2.7 @@ -1903,6 +1999,12 @@ snapshots: optionalDependencies: '@types/react': 19.2.17 + '@radix-ui/react-context@1.2.0(@types/react@19.2.17)(react@19.2.7)': + dependencies: + react: 19.2.7 + optionalDependencies: + '@types/react': 19.2.17 + '@radix-ui/react-dialog@1.1.17(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)': dependencies: '@radix-ui/primitive': 1.1.4 @@ -2130,6 +2232,15 @@ snapshots: '@types/react': 19.2.17 '@types/react-dom': 19.2.3(@types/react@19.2.17) + '@radix-ui/react-presence@1.1.7(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)': + dependencies: + '@radix-ui/react-use-layout-effect': 1.1.2(@types/react@19.2.17)(react@19.2.7) + react: 19.2.7 + react-dom: 19.2.7(react@19.2.7) + optionalDependencies: + '@types/react': 19.2.17 + '@types/react-dom': 19.2.3(@types/react@19.2.17) + '@radix-ui/react-primitive@2.1.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)': dependencies: '@radix-ui/react-slot': 1.3.0(@types/react@19.2.17)(react@19.2.7) @@ -2200,6 +2311,25 @@ snapshots: '@types/react': 19.2.17 '@types/react-dom': 19.2.3(@types/react@19.2.17) + '@radix-ui/react-roving-focus@1.1.15(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)': + dependencies: + '@radix-ui/primitive': 1.1.5 + '@radix-ui/react-collection': 1.1.12(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) + '@radix-ui/react-compose-refs': 1.1.3(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-context': 1.2.0(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-direction': 1.1.2(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-id': 1.1.2(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-primitive': 2.1.7(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) + '@radix-ui/react-use-callback-ref': 1.1.2(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-use-controllable-state': 1.2.3(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-use-is-hydrated': 0.1.1(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-use-layout-effect': 1.1.2(@types/react@19.2.17)(react@19.2.7) + react: 19.2.7 + react-dom: 19.2.7(react@19.2.7) + optionalDependencies: + '@types/react': 19.2.17 + '@types/react-dom': 19.2.3(@types/react@19.2.17) + '@radix-ui/react-select@2.3.1(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)': dependencies: '@radix-ui/number': 1.1.2 @@ -2237,6 +2367,22 @@ snapshots: optionalDependencies: '@types/react': 19.2.17 + '@radix-ui/react-tabs@1.1.17(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)': + dependencies: + '@radix-ui/primitive': 1.1.5 + '@radix-ui/react-context': 1.2.0(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-direction': 1.1.2(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-id': 1.1.2(@types/react@19.2.17)(react@19.2.7) + '@radix-ui/react-presence': 1.1.7(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) + '@radix-ui/react-primitive': 2.1.7(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) + '@radix-ui/react-roving-focus': 1.1.15(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) + '@radix-ui/react-use-controllable-state': 1.2.3(@types/react@19.2.17)(react@19.2.7) + react: 19.2.7 + react-dom: 19.2.7(react@19.2.7) + optionalDependencies: + '@types/react': 19.2.17 + '@types/react-dom': 19.2.3(@types/react@19.2.17) + '@radix-ui/react-toast@1.2.18(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)': dependencies: '@radix-ui/primitive': 1.1.4 @@ -2285,6 +2431,12 @@ snapshots: optionalDependencies: '@types/react': 19.2.17 + '@radix-ui/react-use-is-hydrated@0.1.1(@types/react@19.2.17)(react@19.2.7)': + dependencies: + react: 19.2.7 + optionalDependencies: + '@types/react': 19.2.17 + '@radix-ui/react-use-layout-effect@1.1.2(@types/react@19.2.17)(react@19.2.7)': dependencies: react: 19.2.7 @@ -2558,6 +2710,8 @@ snapshots: csstype@3.2.3: {} + date-fns@4.4.0: {} + debug@4.4.3: dependencies: ms: 2.1.3 diff --git a/frontend/web/src/components/ui/alert.tsx b/frontend/web/src/components/ui/alert.tsx new file mode 100644 index 00000000..97fdbbcf --- /dev/null +++ b/frontend/web/src/components/ui/alert.tsx @@ -0,0 +1,61 @@ +import * as React from "react" +import { cva, type VariantProps } from "class-variance-authority" + +import { cn } from "@/lib/utils" + +const alertVariants = cva( + "relative w-full rounded-lg border px-4 py-3 text-sm [&>svg+div]:translate-y-[-3px] [&>svg]:absolute [&>svg]:left-4 [&>svg]:top-4 [&>svg]:text-foreground [&>svg~*]:pl-7", + { + variants: { + variant: { + default: "bg-background text-foreground", + destructive: + "border-destructive/50 text-destructive dark:border-destructive [&>svg]:text-destructive", + }, + }, + defaultVariants: { + variant: "default", + }, + } +) + +const Alert = React.forwardRef< + HTMLDivElement, + React.HTMLAttributes & VariantProps +>(({ className, variant, ...props }, ref) => ( +
+)) +Alert.displayName = "Alert" + +const AlertTitle = ({ + className, + ref, + ...props +}: React.HTMLAttributes & { ref?: React.Ref }) => ( +
+) +AlertTitle.displayName = "AlertTitle" + +const AlertDescription = ({ + className, + ref, + ...props +}: React.HTMLAttributes & { ref?: React.Ref }) => ( +
+) +AlertDescription.displayName = "AlertDescription" + +export { Alert, AlertTitle, AlertDescription } diff --git a/frontend/web/src/components/ui/dropdown-menu.tsx b/frontend/web/src/components/ui/dropdown-menu.tsx new file mode 100644 index 00000000..a4a12de3 --- /dev/null +++ b/frontend/web/src/components/ui/dropdown-menu.tsx @@ -0,0 +1,228 @@ +import * as React from "react" +import * as DropdownMenuPrimitive from "@radix-ui/react-dropdown-menu" +import { Check, ChevronRight, Circle } from "lucide-react" + +import { cn } from "@/lib/utils" + +const DropdownMenu = DropdownMenuPrimitive.Root + +const DropdownMenuTrigger = DropdownMenuPrimitive.Trigger + +const DropdownMenuGroup = DropdownMenuPrimitive.Group + +const DropdownMenuPortal = DropdownMenuPrimitive.Portal + +const DropdownMenuSub = DropdownMenuPrimitive.Sub + +const DropdownMenuRadioGroup = DropdownMenuPrimitive.RadioGroup + +const DropdownMenuSubTrigger = ({ + className, + inset, + children, + ref, + ...props +}: React.ComponentPropsWithoutRef & { + inset?: boolean + ref?: React.Ref> +}) => ( + + {children} + + +) +DropdownMenuSubTrigger.displayName = + DropdownMenuPrimitive.SubTrigger.displayName + +const DropdownMenuSubContent = ({ + className, + ref, + ...props +}: React.ComponentPropsWithoutRef & { + ref?: React.Ref> +}) => ( + +) +DropdownMenuSubContent.displayName = + DropdownMenuPrimitive.SubContent.displayName + +const DropdownMenuContent = ({ + className, + sideOffset = 4, + ref, + ...props +}: React.ComponentPropsWithoutRef & { + ref?: React.Ref> +}) => ( + + + +) +DropdownMenuContent.displayName = DropdownMenuPrimitive.Content.displayName + +const DropdownMenuItem = ({ + className, + inset, + ref, + ...props +}: React.ComponentPropsWithoutRef & { + inset?: boolean + ref?: React.Ref> +}) => ( + svg]:size-4 [&>svg]:shrink-0", + inset && "pl-8", + className + )} + {...props} + /> +) +DropdownMenuItem.displayName = DropdownMenuPrimitive.Item.displayName + +const DropdownMenuCheckboxItem = ({ + className, + children, + checked, + ref, + ...props +}: React.ComponentPropsWithoutRef & { + ref?: React.Ref> +}) => ( + + + + + + + {children} + +) +DropdownMenuCheckboxItem.displayName = + DropdownMenuPrimitive.CheckboxItem.displayName + +const DropdownMenuRadioItem = ({ + className, + children, + ref, + ...props +}: React.ComponentPropsWithoutRef & { + ref?: React.Ref> +}) => ( + + + + + + + {children} + +) +DropdownMenuRadioItem.displayName = DropdownMenuPrimitive.RadioItem.displayName + +const DropdownMenuLabel = ({ + className, + inset, + ref, + ...props +}: React.ComponentPropsWithoutRef & { + inset?: boolean + ref?: React.Ref> +}) => ( + +) +DropdownMenuLabel.displayName = DropdownMenuPrimitive.Label.displayName + +const DropdownMenuSeparator = ({ + className, + ref, + ...props +}: React.ComponentPropsWithoutRef & { + ref?: React.Ref> +}) => ( + +) +DropdownMenuSeparator.displayName = DropdownMenuPrimitive.Separator.displayName + +const DropdownMenuShortcut = ({ + className, + ...props +}: React.HTMLAttributes) => { + return ( + + ) +} +DropdownMenuShortcut.displayName = "DropdownMenuShortcut" + +export { + DropdownMenu, + DropdownMenuTrigger, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuCheckboxItem, + DropdownMenuRadioItem, + DropdownMenuLabel, + DropdownMenuSeparator, + DropdownMenuShortcut, + DropdownMenuGroup, + DropdownMenuPortal, + DropdownMenuSub, + DropdownMenuSubContent, + DropdownMenuSubTrigger, + DropdownMenuRadioGroup, +} diff --git a/frontend/web/src/components/ui/tabs.tsx b/frontend/web/src/components/ui/tabs.tsx new file mode 100644 index 00000000..85d83bea --- /dev/null +++ b/frontend/web/src/components/ui/tabs.tsx @@ -0,0 +1,53 @@ +import * as React from "react" +import * as TabsPrimitive from "@radix-ui/react-tabs" + +import { cn } from "@/lib/utils" + +const Tabs = TabsPrimitive.Root + +const TabsList = React.forwardRef< + React.ElementRef, + React.ComponentPropsWithoutRef +>(({ className, ...props }, ref) => ( + +)) +TabsList.displayName = TabsPrimitive.List.displayName + +const TabsTrigger = React.forwardRef< + React.ElementRef, + React.ComponentPropsWithoutRef +>(({ className, ...props }, ref) => ( + +)) +TabsTrigger.displayName = TabsPrimitive.Trigger.displayName + +const TabsContent = React.forwardRef< + React.ElementRef, + React.ComponentPropsWithoutRef +>(({ className, ...props }, ref) => ( + +)) +TabsContent.displayName = TabsPrimitive.Content.displayName + +export { Tabs, TabsList, TabsTrigger, TabsContent } diff --git a/frontend/web/src/features/developers/api/IApiTokenRepository.ts b/frontend/web/src/features/developers/api/IApiTokenRepository.ts new file mode 100644 index 00000000..f67e20c7 --- /dev/null +++ b/frontend/web/src/features/developers/api/IApiTokenRepository.ts @@ -0,0 +1,8 @@ +import type { ApiToken, ApiTokenCreated, CreateApiTokenPayload } from '../types'; + +export interface IApiTokenRepository { + getApiTokens(): Promise<{ tokens: ApiToken[] }>; + createApiToken(payload: CreateApiTokenPayload): Promise; + revokeApiToken(id: string): Promise; + deleteApiToken(id: string): Promise; +} diff --git a/frontend/web/src/features/developers/api/apiTokenHooks.ts b/frontend/web/src/features/developers/api/apiTokenHooks.ts new file mode 100644 index 00000000..f8639fea --- /dev/null +++ b/frontend/web/src/features/developers/api/apiTokenHooks.ts @@ -0,0 +1,78 @@ +import { useQuery, useMutation, useQueryClient } from '@tanstack/react-query'; +import type { QueryKey } from '@tanstack/react-query'; +import { useAuth } from 'react-oidc-context'; +import { createApiTokenRepository } from './apiTokensApi'; +import type { CreateApiTokenPayload } from '../types'; +import { useToast } from '@/hooks/use-toast'; + +export const apiTokenKeys = { + all: ['apiTokens'] as const, + lists: () => [...apiTokenKeys.all, 'list'] as const, +}; + +function useRepository() { + const auth = useAuth(); + return createApiTokenRepository(auth.user?.access_token ?? ''); +} + +function useToastMutation( + mutationFn: (variables: TVariables) => Promise, + successMessage: string | ((data: TData) => string), + queryKeysToInvalidate: QueryKey[] = [] +) { + const queryClient = useQueryClient(); + const { toast } = useToast(); + + return useMutation({ + mutationFn, + onSuccess: (data) => { + queryKeysToInvalidate.forEach(key => { + queryClient.invalidateQueries({ queryKey: key }); + }); + const message = typeof successMessage === 'function' ? successMessage(data) : successMessage; + if (message) { + toast({ title: 'Success', description: message }); + } + }, + onError: (error: Error) => { + toast({ title: 'Error', description: error.message, variant: 'destructive' }); + } + }); +} + +export function useApiTokensQuery() { + const auth = useAuth(); + const repo = createApiTokenRepository(auth.user?.access_token ?? ''); + return useQuery({ + queryKey: apiTokenKeys.lists(), + queryFn: () => repo.getApiTokens(), + enabled: !!auth.user?.access_token, + }); +} + +export function useCreateApiTokenMutation() { + const repo = useRepository(); + return useToastMutation( + (payload: CreateApiTokenPayload) => repo.createApiToken(payload), + 'API Token created successfully.', + [apiTokenKeys.lists()] + ); +} + +export function useRevokeApiTokenMutation() { + const repo = useRepository(); + return useToastMutation( + (id: string) => repo.revokeApiToken(id), + 'API Token revoked.', + [apiTokenKeys.lists()] + ); +} + +export function useDeleteApiTokenMutation() { + const repo = useRepository(); + return useToastMutation( + (id: string) => repo.deleteApiToken(id), + 'API Token deleted.', + [apiTokenKeys.lists()] + ); +} diff --git a/frontend/web/src/features/developers/api/apiTokensApi.ts b/frontend/web/src/features/developers/api/apiTokensApi.ts new file mode 100644 index 00000000..9dd142f6 --- /dev/null +++ b/frontend/web/src/features/developers/api/apiTokensApi.ts @@ -0,0 +1,68 @@ +import type { IApiTokenRepository } from './IApiTokenRepository'; +import type { ApiToken, ApiTokenCreated, CreateApiTokenPayload } from '../types'; + +class HttpApiTokenRepository implements IApiTokenRepository { + private readonly headers: Record; + + constructor(token: string) { + this.headers = { + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}`, + }; + } + + private async request(url: string, init?: RequestInit): Promise { + const controller = new AbortController(); + const timeoutId = setTimeout(() => controller.abort(), 15000); + + try { + const res = await fetch(url, { ...init, headers: this.headers, signal: controller.signal }); + if (!res.ok) { + let errorMessage = res.statusText; + const contentType = res.headers.get('content-type'); + if (contentType && contentType.includes('application/json')) { + try { + const data = await res.json(); + errorMessage = typeof data.detail === 'string' ? data.detail : JSON.stringify(data.detail || data); + } catch { + errorMessage = await res.text().catch(() => res.statusText); + } + } else { + errorMessage = await res.text().catch(() => res.statusText); + } + throw new Error(errorMessage || `HTTP ${res.status}`); + } + if (res.status === 204) return undefined as unknown as T; + return res.json() as Promise; + } finally { + clearTimeout(timeoutId); + } + } + + getApiTokens(): Promise<{ tokens: ApiToken[] }> { + return this.request('/api/v1/developers/tokens'); + } + + createApiToken(payload: CreateApiTokenPayload): Promise { + return this.request('/api/v1/developers/tokens', { + method: 'POST', + body: JSON.stringify(payload), + }); + } + + revokeApiToken(id: string): Promise { + return this.request(`/api/v1/developers/tokens/${id}`, { + method: 'DELETE', + }); + } + + deleteApiToken(id: string): Promise { + return this.request(`/api/v1/developers/tokens/${id}/hard`, { + method: 'DELETE', + }); + } +} + +export function createApiTokenRepository(token: string): IApiTokenRepository { + return new HttpApiTokenRepository(token); +} diff --git a/frontend/web/src/features/developers/components/ApiTokensTable.tsx b/frontend/web/src/features/developers/components/ApiTokensTable.tsx new file mode 100644 index 00000000..bdb1e376 --- /dev/null +++ b/frontend/web/src/features/developers/components/ApiTokensTable.tsx @@ -0,0 +1,134 @@ +import React from 'react'; +import { + createColumnHelper, + getCoreRowModel, + useReactTable, +} from '@tanstack/react-table'; +import { DataTable } from '@/components/ui/data-table'; +import type { ApiToken } from '../types'; +import { useRevokeApiTokenMutation, useDeleteApiTokenMutation } from '../api/apiTokenHooks'; +import { Button } from '@/components/ui/button'; +import { MoreHorizontal, Trash2, Ban } from 'lucide-react'; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuLabel, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from '@/components/ui/dropdown-menu'; +import { Badge } from '@/components/ui/badge'; +import { formatDistanceToNow, parseISO } from 'date-fns'; + +function TokenRowActions({ token }: { token: ApiToken }) { + const revoke = useRevokeApiTokenMutation(); + const hardDelete = useDeleteApiTokenMutation(); + + const handleRevoke = (e: React.MouseEvent) => { + e.stopPropagation(); + if (confirm('Are you sure you want to revoke this token? Any integrations using it will immediately fail.')) { + revoke.mutate(token.id); + } + }; + + const handleDelete = (e: React.MouseEvent) => { + e.stopPropagation(); + if (confirm('Are you sure you want to permanently delete this token? This cannot be undone.')) { + hardDelete.mutate(token.id); + } + }; + + return ( + + + + + + Actions + + + {token.active && ( + + + Revoke + + )} + + + + Delete + + + + ); +} + +const columnHelper = createColumnHelper(); + +const columns = [ + columnHelper.accessor('name', { + header: 'Name', + cell: (info) =>
{info.getValue()}
, + }), + columnHelper.accessor('client_id', { + header: 'Client ID', + cell: (info) =>
{info.getValue()}
, + }), + columnHelper.accessor('active', { + header: 'Status', + cell: (info) => { + const active = info.getValue(); + return ( + + {active ? 'Active' : 'Revoked'} + + ); + }, + }), + columnHelper.accessor('last_used_at', { + header: 'Last Used', + cell: (info) => { + const val = info.getValue(); + if (!val) return Never; + try { + return {formatDistanceToNow(parseISO(val), { addSuffix: true })}; + } catch { + return Invalid Date; + } + }, + }), + columnHelper.accessor('created_at', { + header: 'Created', + cell: (info) => { + const val = info.getValue(); + if (val === 'just now') return Just now; + try { + return {formatDistanceToNow(parseISO(val), { addSuffix: true })}; + } catch { + return {val}; + } + }, + }), + columnHelper.display({ + id: 'actions', + cell: ({ row }) => , + }), +]; + +interface ApiTokensTableProps { + data: ApiToken[]; + isLoading: boolean; +} + +export function ApiTokensTable({ data, isLoading }: ApiTokensTableProps) { + const table = useReactTable({ + data, + columns, + getCoreRowModel: getCoreRowModel(), + }); + + return ; +} diff --git a/frontend/web/src/features/developers/components/CreateApiTokenModal.tsx b/frontend/web/src/features/developers/components/CreateApiTokenModal.tsx new file mode 100644 index 00000000..b2bb5ec3 --- /dev/null +++ b/frontend/web/src/features/developers/components/CreateApiTokenModal.tsx @@ -0,0 +1,89 @@ +import { useState } from 'react'; +import { Button } from '@/components/ui/button'; +import { Input } from '@/components/ui/input'; +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, + DialogTrigger, +} from '@/components/ui/dialog'; +import { Label } from '@/components/ui/label'; +import { Key } from 'lucide-react'; +import { useCreateApiTokenMutation } from '../api/apiTokenHooks'; +import { TokenCredentialsModal } from './TokenCredentialsModal'; +import type { ApiTokenCreated } from '../types'; + +export function CreateApiTokenModal() { + const [open, setOpen] = useState(false); + const [name, setName] = useState(''); + const [createdToken, setCreatedToken] = useState(null); + + const createMutation = useCreateApiTokenMutation(); + + const handleCreate = async () => { + if (!name.trim()) return; + try { + const data = await createMutation.mutateAsync({ name: name.trim() }); + setCreatedToken(data); + setOpen(false); + setName(''); + } catch { + // Error handled by hook toast + } + }; + + return ( + <> + + + + + + + Generate New API Token + + Create a new API token for machine-to-machine integrations. + + +
+
+ + setName(e.target.value)} + placeholder="e.g. ERP Prod Integration" + autoFocus + /> +
+
+ + + + +
+
+ + {createdToken && ( + setCreatedToken(null)} + /> + )} + + ); +} diff --git a/frontend/web/src/features/developers/components/TokenCredentialsModal.tsx b/frontend/web/src/features/developers/components/TokenCredentialsModal.tsx new file mode 100644 index 00000000..126ceceb --- /dev/null +++ b/frontend/web/src/features/developers/components/TokenCredentialsModal.tsx @@ -0,0 +1,100 @@ +import { useState } from 'react'; +import { useToast } from '@/hooks/use-toast'; +import { Button } from '@/components/ui/button'; +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, +} from '@/components/ui/dialog'; +import { Check, Copy, AlertTriangle } from 'lucide-react'; +import type { ApiTokenCreated } from '../types'; +import { Alert, AlertDescription, AlertTitle } from '@/components/ui/alert'; +import { Input } from '@/components/ui/input'; + +interface Props { + token: ApiTokenCreated; + onClose: () => void; +} + +export function TokenCredentialsModal({ token, onClose }: Props) { + const [copiedId, setCopiedId] = useState(false); + const [copiedSecret, setCopiedSecret] = useState(false); + const { toast } = useToast(); + + const copyToClipboard = async (text: string, isSecret: boolean) => { + try { + await navigator.clipboard.writeText(text); + if (isSecret) { + setCopiedSecret(true); + setTimeout(() => setCopiedSecret(false), 2000); + } else { + setCopiedId(true); + setTimeout(() => setCopiedId(false), 2000); + } + } catch { + toast({ title: 'Error', description: 'Failed to copy to clipboard.', variant: 'destructive' }); + } + }; + + return ( + !open && onClose()}> + + + Token Generated Successfully + + Please copy your client secret now. For your security, it will never be shown again. + + + + + + Store this secret securely + + If you lose this secret, you will need to generate a new token. We do not store the raw secret in our database. + + + +
+
+ Client ID +
+ + +
+
+ +
+ Client Secret +
+ + +
+
+
+ + + + +
+
+ ); +} diff --git a/frontend/web/src/features/developers/types.ts b/frontend/web/src/features/developers/types.ts new file mode 100644 index 00000000..12b1893b --- /dev/null +++ b/frontend/web/src/features/developers/types.ts @@ -0,0 +1,18 @@ +export interface ApiToken { + id: string; + name: string; + client_id: string; + active: boolean; + last_used_at: string | null; + expires_at: string | null; + created_at: string; +} + +export interface ApiTokenCreated extends ApiToken { + client_secret: string; +} + +export interface CreateApiTokenPayload { + name: string; + expires_at?: string; +} diff --git a/frontend/web/src/features/partners/components/PartnershipDetails.tsx b/frontend/web/src/features/partners/components/PartnershipDetails.tsx index 88962fab..49e30478 100644 --- a/frontend/web/src/features/partners/components/PartnershipDetails.tsx +++ b/frontend/web/src/features/partners/components/PartnershipDetails.tsx @@ -131,6 +131,7 @@ export function PartnershipDetails({ partnership, availablePartners, onCancel }: )} />
+
diff --git a/frontend/web/src/features/partners/components/PartnershipsTable.tsx b/frontend/web/src/features/partners/components/PartnershipsTable.tsx index fdc273e4..b7912250 100644 --- a/frontend/web/src/features/partners/components/PartnershipsTable.tsx +++ b/frontend/web/src/features/partners/components/PartnershipsTable.tsx @@ -84,6 +84,7 @@ export function PartnershipsTable({ data, availablePartners, isLoading }: { data ); }, }), + columnHelper.accessor('local_partner_id', { header: 'Local Partner', cell: (info) => { diff --git a/frontend/web/src/features/partners/types.ts b/frontend/web/src/features/partners/types.ts index 5a8e4a19..bea6f531 100644 --- a/frontend/web/src/features/partners/types.ts +++ b/frontend/web/src/features/partners/types.ts @@ -35,6 +35,7 @@ export type Partner = AS2Partner | SFTPPartner; export interface Partnership { id: string; + name?: string; local_partner_id: string; remote_partner_id: string; diff --git a/frontend/web/src/features/routes/components/CreateRouteModal.tsx b/frontend/web/src/features/routes/components/CreateRouteModal.tsx index ed00e0fa..a6c0ede9 100644 --- a/frontend/web/src/features/routes/components/CreateRouteModal.tsx +++ b/frontend/web/src/features/routes/components/CreateRouteModal.tsx @@ -1,203 +1,58 @@ import { useState } from 'react'; -import { Input } from '@/components/ui/input'; -import { Label } from '@/components/ui/label'; -import { SearchableSelect } from '@/components/ui/searchable-select'; -import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from '@/components/ui/select'; -import { useCreateInboundRouteMutation, useCreateOutboundRouteMutation } from '../api/routeHooks'; -import { useTenantDestinations } from '../hooks/useTenantDestinations'; -import { Network } from 'lucide-react'; -import { useToast } from '@/hooks/use-toast'; -import { FormModal } from '@/components/ui/form-modal'; +import { Network, Plus } from 'lucide-react'; +import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogTrigger } from '@/components/ui/dialog'; +import { Button } from '@/components/ui/button'; +import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs'; +import { InboundRouteForm } from './InboundRouteForm'; +import { OutboundRouteForm } from './OutboundRouteForm'; export function CreateRouteModal() { const [isOpen, setIsOpen] = useState(false); - const [name, setName] = useState(''); - const [direction, setDirection] = useState<'INBOUND' | 'OUTBOUND'>('INBOUND'); - const [processingMode, setProcessingMode] = useState<'TRANSLATE' | 'PASSTHROUGH'>('TRANSLATE'); - const [transactionType, setTransactionType] = useState('*'); - const [isaSender, setIsaSender] = useState(''); - const [isaReceiver, setIsaReceiver] = useState(''); - const [targetId, setTargetId] = useState(''); - - const { toast } = useToast(); - const { data: destinations, isLoading: isLoadingDestinations } = useTenantDestinations(direction); - - const createInbound = useCreateInboundRouteMutation(); - const createOutbound = useCreateOutboundRouteMutation(); - - const isPending = createInbound.isPending || createOutbound.isPending; - - const handleSubmit = async (e: React.FormEvent) => { - e.preventDefault(); - if (!name || !isaSender || !isaReceiver || !targetId || !transactionType) { - toast({ title: 'Please fill all fields', variant: 'destructive' }); - return; - } - - try { - if (direction === 'INBOUND') { - // Inbound: Target is an Endpoint (Webhook) - const selectedEndpoint = destinations?.find(e => e.id === targetId); - if (!selectedEndpoint) { - toast({ title: 'Invalid endpoint selected', variant: 'destructive' }); - return; - } - - await createInbound.mutateAsync({ - name, - isa_sender_id: isaSender, - isa_receiver_id: isaReceiver, - transaction_type: transactionType, - processing_mode: processingMode, - webhook_id: targetId, - }); - } else { - // Outbound: Target is a Trading Partner (AS2/SFTP) - const selectedPartner = destinations?.find(p => p.id === targetId); - if (!selectedPartner) { - toast({ title: 'Invalid partner selected', variant: 'destructive' }); - return; - } - await createOutbound.mutateAsync({ - name, - isa_sender_id: isaSender, - isa_receiver_id: isaReceiver, - transaction_type: transactionType, - processing_mode: processingMode, - as2_partner_id: selectedPartner.type?.toUpperCase() === 'AS2' ? targetId : undefined, - sftp_partner_id: selectedPartner.type?.toUpperCase() === 'SFTP' ? targetId : undefined, - }); - } - - toast({ title: 'Route created successfully' }); - setIsOpen(false); - - // Reset form - setName(''); - setIsaSender(''); - setIsaReceiver(''); - setTargetId(''); - setTransactionType('*'); - } catch (err) { - toast({ title: 'Failed to create route', description: String(err), variant: 'destructive' }); - } - }; + const [activeTab, setActiveTab] = useState('inbound'); return ( - } - isOpen={isOpen} - onOpenChange={setIsOpen} - onSubmit={handleSubmit} - isPending={isPending} - submitText="Create Route" - maxWidth="sm:max-w-[500px]" - > -
- - setName(e.target.value)} - placeholder="e.g. Inbound Walmart 850" - className="h-10 rounded-xl text-sm" - /> -
- -
-
- - -
-
- - setTransactionType(e.target.value)} - placeholder="e.g. 850 or *" - className="h-10 rounded-xl font-mono text-sm" - /> -
-
- -
-
- - setIsaSender(e.target.value)} - placeholder="e.g. ACME_CORP" - className="h-10 rounded-xl font-mono text-sm uppercase" - /> + + + + + + e.preventDefault()} + > + + +
+ +
+ Create Routing Rule +
+
+ +
+ + + + Inbound (From EDI) + + + Outbound (From JSON) + + + + + setIsOpen(false)} /> + + + + setIsOpen(false)} /> + +
-
- - setIsaReceiver(e.target.value)} - placeholder="e.g. WALMART" - className="h-10 rounded-xl font-mono text-sm uppercase" - /> -
-
- -
- - -
- -
- - direction === 'INBOUND' || !(d.type === 'AS2' && (d as any).is_local)) - .map(d => ({ - value: d.id, - label: ( - - {d.type} - {d.name} - - ), - searchString: d.name - }))} - /> -
- + + ); } diff --git a/frontend/web/src/features/routes/components/InboundRouteForm.tsx b/frontend/web/src/features/routes/components/InboundRouteForm.tsx new file mode 100644 index 00000000..c5a67eeb --- /dev/null +++ b/frontend/web/src/features/routes/components/InboundRouteForm.tsx @@ -0,0 +1,146 @@ +import { useState } from 'react'; +import { Input } from '@/components/ui/input'; +import { Label } from '@/components/ui/label'; +import { SearchableSelect } from '@/components/ui/searchable-select'; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from '@/components/ui/select'; +import { useCreateInboundRouteMutation } from '../api/routeHooks'; +import { useTenantDestinations } from '../hooks/useTenantDestinations'; +import { useToast } from '@/hooks/use-toast'; +import { Button } from '@/components/ui/button'; + +export function InboundRouteForm({ onSuccess }: { onSuccess: () => void }) { + const [name, setName] = useState(''); + const [processingMode, setProcessingMode] = useState<'TRANSLATE' | 'PASSTHROUGH'>('TRANSLATE'); + const [transactionType, setTransactionType] = useState('*'); + const [isaSender, setIsaSender] = useState(''); + const [isaReceiver, setIsaReceiver] = useState(''); + const [targetId, setTargetId] = useState(''); + + const { toast } = useToast(); + const { data: destinations, isLoading: isLoadingDestinations } = useTenantDestinations('INBOUND'); + const createInbound = useCreateInboundRouteMutation(); + + const handleSubmit = async (e: React.FormEvent) => { + e.preventDefault(); + if (!name || !isaSender || !isaReceiver || !targetId || !transactionType) { + toast({ title: 'Please fill all required fields', variant: 'destructive' }); + return; + } + + try { + const selectedEndpoint = destinations?.find(e => e.id === targetId); + if (!selectedEndpoint) { + toast({ title: 'Invalid endpoint selected', variant: 'destructive' }); + return; + } + + await createInbound.mutateAsync({ + name, + isa_sender_id: isaSender, + isa_receiver_id: isaReceiver, + transaction_type: transactionType, + processing_mode: processingMode, + webhook_id: targetId, + }); + + toast({ title: 'Inbound route created successfully' }); + onSuccess(); + } catch (err) { + toast({ title: 'Failed to create inbound route', description: String(err), variant: 'destructive' }); + } + }; + + return ( +
+
+ + setName(e.target.value)} + placeholder="e.g. Inbound Walmart 850" + className="h-10 rounded-xl text-sm" + /> +
+ +
+
+ + setIsaSender(e.target.value)} + placeholder="e.g. ACME_CORP" + className="h-10 rounded-xl font-mono text-sm uppercase" + /> +
+
+ + setIsaReceiver(e.target.value)} + placeholder="e.g. WALMART" + className="h-10 rounded-xl font-mono text-sm uppercase" + /> +
+
+ +
+
+ + setTransactionType(e.target.value)} + placeholder="e.g. 850 or *" + className="h-10 rounded-xl font-mono text-sm" + /> +
+
+ + +
+
+ +
+ + ({ + value: d.id, + label: ( + + {d.type} + {d.name} + + ), + searchString: d.name + }))} + /> +
+ +
+ +
+
+ ); +} diff --git a/frontend/web/src/features/routes/components/OutboundRouteForm.tsx b/frontend/web/src/features/routes/components/OutboundRouteForm.tsx new file mode 100644 index 00000000..cfb51f69 --- /dev/null +++ b/frontend/web/src/features/routes/components/OutboundRouteForm.tsx @@ -0,0 +1,201 @@ +import { useState } from 'react'; +import { Input } from '@/components/ui/input'; +import { Label } from '@/components/ui/label'; +import { SearchableSelect } from '@/components/ui/searchable-select'; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from '@/components/ui/select'; +import { useCreateOutboundRouteMutation } from '../api/routeHooks'; +import { useTenantDestinations } from '../hooks/useTenantDestinations'; +import { useToast } from '@/hooks/use-toast'; +import { Button } from '@/components/ui/button'; + +export function OutboundRouteForm({ onSuccess }: { onSuccess: () => void }) { + const [name, setName] = useState(''); + const [externalId, setExternalId] = useState(''); + const [processingMode] = useState<'TRANSLATE' | 'PASSTHROUGH'>('TRANSLATE'); + const [transactionType, setTransactionType] = useState('*'); + + // Envelope fields + const [isaSender, setIsaSender] = useState(''); + const [isaSenderQual, setIsaSenderQual] = useState('ZZ'); + const [isaReceiver, setIsaReceiver] = useState(''); + const [isaReceiverQual, setIsaReceiverQual] = useState('ZZ'); + const [gsSender, setGsSender] = useState(''); + const [gsReceiver, setGsReceiver] = useState(''); + const [defaultStandard, setDefaultStandard] = useState('x12'); + const [defaultVersion, setDefaultVersion] = useState('004010'); + + const [targetId, setTargetId] = useState(''); + + const { toast } = useToast(); + // For Outbound, we only fetch destinations that are NOT webhooks + const { data: destinations, isLoading: isLoadingDestinations } = useTenantDestinations('OUTBOUND'); + const createOutbound = useCreateOutboundRouteMutation(); + + const handleSubmit = async (e: React.FormEvent) => { + e.preventDefault(); + if (!name || !externalId || !isaSender || !isaReceiver || !gsSender || !gsReceiver || !targetId || !transactionType) { + toast({ title: 'Please fill all required fields', variant: 'destructive' }); + return; + } + + try { + const selectedPartner = destinations?.find(p => p.id === targetId); + if (!selectedPartner) { + toast({ title: 'Invalid partner selected', variant: 'destructive' }); + return; + } + + await createOutbound.mutateAsync({ + trading_partner_id: externalId, + name, + isa_sender_id: isaSender, + isa_sender_qualifier: isaSenderQual, + isa_receiver_id: isaReceiver, + isa_receiver_qualifier: isaReceiverQual, + gs_sender_id: gsSender, + gs_receiver_id: gsReceiver, + default_standard: defaultStandard, + default_version: defaultVersion, + transaction_type: transactionType, + processing_mode: processingMode, + as2_partner_id: selectedPartner.type?.toUpperCase() === 'AS2' ? targetId : undefined, + sftp_partner_id: selectedPartner.type?.toUpperCase() === 'SFTP' ? targetId : undefined, + }); + + toast({ title: 'Outbound route created successfully' }); + onSuccess(); + } catch (err) { + toast({ title: 'Failed to create outbound route', description: String(err), variant: 'destructive' }); + } + }; + + return ( +
+ + {/* General Settings */} +
+

General Settings

+
+
+ + setName(e.target.value)} + placeholder="e.g. Outbound Conway AS2" + className="h-10 bg-white" + /> +
+
+ + setExternalId(e.target.value)} + placeholder="e.g. CONWAY_OUT" + className="h-10 bg-white font-mono uppercase text-sm" + /> +
+
+ +
+
+ + setTransactionType(e.target.value)} + placeholder="e.g. 850 or *" + className="h-10 bg-white font-mono text-sm" + /> +
+
+
+ + {/* EDI Envelope Settings */} +
+

EDI Envelope Configuration

+ +
+
+ + setIsaSenderQual(e.target.value)} className="h-10 bg-white font-mono text-sm" /> +
+
+ + setIsaSender(e.target.value)} placeholder="ACME_CORP" className="h-10 bg-white font-mono text-sm uppercase" /> +
+
+ + setIsaReceiverQual(e.target.value)} className="h-10 bg-white font-mono text-sm" /> +
+
+ + setIsaReceiver(e.target.value)} placeholder="PARTNER" className="h-10 bg-white font-mono text-sm uppercase" /> +
+
+ +
+
+ + setGsSender(e.target.value)} placeholder="ACME" className="h-10 bg-white font-mono text-sm uppercase" /> +
+
+ + setGsReceiver(e.target.value)} placeholder="PARTNER" className="h-10 bg-white font-mono text-sm uppercase" /> +
+
+ + +
+
+ + setDefaultVersion(e.target.value)} placeholder="004010" className="h-10 bg-white font-mono text-sm" /> +
+
+
+ + {/* Target Destination */} +
+

Target Destination

+
+ + !(d.type === 'AS2' && (d as any).is_local)) + .map(d => ({ + value: d.id, + label: ( + + {d.type} + {d.name} + + ), + searchString: d.name + }))} + /> +
+
+ +
+ +
+
+ ); +} diff --git a/frontend/web/src/features/routes/components/RouteDetails.tsx b/frontend/web/src/features/routes/components/RouteDetails.tsx index 66b99186..0089cbec 100644 --- a/frontend/web/src/features/routes/components/RouteDetails.tsx +++ b/frontend/web/src/features/routes/components/RouteDetails.tsx @@ -21,6 +21,7 @@ export function RouteDetails({ route, onCancel }: { route: RouteItem, onCancel?: const { toast } = useToast(); const updateRoute = useUpdateRouteMutation(); const isSubmitting = updateRoute.isPending; + const isOutbound = route.direction === 'OUTBOUND'; const { data: destinations } = useTenantDestinations(route.direction); @@ -31,8 +32,15 @@ export function RouteDetails({ route, onCancel }: { route: RouteItem, onCancel?: const { register, handleSubmit, reset, setValue, watch, formState: { isDirty } } = useForm({ defaultValues: { name: route.name || '', + trading_partner_id: route.trading_partner_id || '', isa_sender_id: route.isa_sender_id || '', + isa_sender_qualifier: route.isa_sender_qualifier || 'ZZ', isa_receiver_id: route.isa_receiver_id || '', + isa_receiver_qualifier: route.isa_receiver_qualifier || 'ZZ', + gs_sender_id: route.gs_sender_id || '', + gs_receiver_id: route.gs_receiver_id || '', + default_standard: route.default_standard || 'x12', + default_version: route.default_version || '004010', transaction_type: route.transaction_type || '', processing_mode: route.processing_mode || 'TRANSLATE', } @@ -45,11 +53,22 @@ export function RouteDetails({ route, onCancel }: { route: RouteItem, onCancel?: const initialTargetId = route.webhook_id || route.as2_partner_id || route.sftp_partner_id || ''; if (formData.name !== route.name) payload.name = formData.name; - if (formData.isa_sender_id !== route.isa_sender_id) payload.isa_sender_id = formData.isa_sender_id; - if (formData.isa_receiver_id !== route.isa_receiver_id) payload.isa_receiver_id = formData.isa_receiver_id; if (formData.transaction_type !== route.transaction_type) payload.transaction_type = formData.transaction_type; if (formData.processing_mode !== route.processing_mode) payload.processing_mode = formData.processing_mode; + if (formData.isa_sender_id !== route.isa_sender_id) payload.isa_sender_id = formData.isa_sender_id; + if (formData.isa_receiver_id !== route.isa_receiver_id) payload.isa_receiver_id = formData.isa_receiver_id; + + if (isOutbound) { + if (formData.trading_partner_id !== route.trading_partner_id) payload.trading_partner_id = formData.trading_partner_id; + if (formData.isa_sender_qualifier !== route.isa_sender_qualifier) payload.isa_sender_qualifier = formData.isa_sender_qualifier; + if (formData.isa_receiver_qualifier !== route.isa_receiver_qualifier) payload.isa_receiver_qualifier = formData.isa_receiver_qualifier; + if (formData.gs_sender_id !== route.gs_sender_id) payload.gs_sender_id = formData.gs_sender_id; + if (formData.gs_receiver_id !== route.gs_receiver_id) payload.gs_receiver_id = formData.gs_receiver_id; + if (formData.default_standard !== route.default_standard) payload.default_standard = formData.default_standard; + if (formData.default_version !== route.default_version) payload.default_version = formData.default_version; + } + if (targetId !== initialTargetId) { if (route.direction === 'INBOUND') { payload.webhook_id = targetId; @@ -80,44 +99,95 @@ export function RouteDetails({ route, onCancel }: { route: RouteItem, onCancel?: return (
+
+ {/* Top Info */}
- - + +
- - + +
-
- - -
+ {isOutbound && ( +
+ + +
+ )} + + {!isOutbound && ( +
+ + +
+ )} +
+ {/* Envelope Configuration (Grid adjustments based on Outbound) */} +
+ {isOutbound && ( +
+ + +
+ )}
- - + +
+ {isOutbound && ( +
+ + +
+ )}
- - + +
+ {isOutbound && ( + <> +
+ + +
+
+ + +
+ +
+ + +
+
+ + +
+ + )} +
+ + {/* Target Destination */} +
( -
+
{info.getValue()} + {info.row.original.trading_partner_id && ( +
+ Trading Partner ID: {info.row.original.trading_partner_id} +
+ )}
), }), diff --git a/frontend/web/src/features/routes/types.ts b/frontend/web/src/features/routes/types.ts index bf51dc6a..3f1892d0 100644 --- a/frontend/web/src/features/routes/types.ts +++ b/frontend/web/src/features/routes/types.ts @@ -1,9 +1,16 @@ export interface RouteItem { route_id: string; + trading_partner_id?: string; name: string; direction: 'INBOUND' | 'OUTBOUND'; isa_sender_id: string; + isa_sender_qualifier?: string; isa_receiver_id: string; + isa_receiver_qualifier?: string; + gs_sender_id?: string; + gs_receiver_id?: string; + default_standard?: string; + default_version?: string; transaction_type: string; destination_type: string; destination_name: string; @@ -19,6 +26,8 @@ export interface CreateInboundRoutePayload { name: string; isa_sender_id: string; isa_receiver_id: string; + gs_sender_id?: string; + gs_receiver_id?: string; transaction_type: string; processing_mode: 'TRANSLATE' | 'PASSTHROUGH'; webhook_id?: string; @@ -27,9 +36,16 @@ export interface CreateInboundRoutePayload { } export interface CreateOutboundRoutePayload { + trading_partner_id: string; name: string; isa_sender_id: string; + isa_sender_qualifier?: string; isa_receiver_id: string; + isa_receiver_qualifier?: string; + gs_sender_id: string; + gs_receiver_id: string; + default_standard: string; + default_version: string; transaction_type: string; processing_mode: 'TRANSLATE' | 'PASSTHROUGH'; as2_partner_id?: string; @@ -38,8 +54,15 @@ export interface CreateOutboundRoutePayload { export interface UpdateRoutePayload { name?: string; + trading_partner_id?: string; isa_sender_id?: string; + isa_sender_qualifier?: string; isa_receiver_id?: string; + isa_receiver_qualifier?: string; + gs_sender_id?: string; + gs_receiver_id?: string; + default_standard?: string; + default_version?: string; transaction_type?: string; processing_mode?: 'TRANSLATE' | 'PASSTHROUGH'; webhook_id?: string; diff --git a/frontend/web/src/routeTree.gen.ts b/frontend/web/src/routeTree.gen.ts index 2edfc199..9979ab60 100644 --- a/frontend/web/src/routeTree.gen.ts +++ b/frontend/web/src/routeTree.gen.ts @@ -16,12 +16,14 @@ import { Route as MarketingIndexRouteImport } from './routes/_marketing/index' import { Route as TenantUsersRouteImport } from './routes/tenant/users' import { Route as TenantRoutesRouteImport } from './routes/tenant/routes' import { Route as TenantPartnersRouteImport } from './routes/tenant/partners' -import { Route as TenantEdiToolRouteImport } from './routes/tenant/edi_tool' +import { Route as TenantEndpointsRouteImport } from './routes/tenant/endpoints' +import { Route as TenantEdi_toolRouteImport } from './routes/tenant/edi_tool' +import { Route as TenantDevelopersRouteImport } from './routes/tenant/developers' import { Route as TenantDashboardRouteImport } from './routes/tenant/dashboard' import { Route as PlatformUsersRouteImport } from './routes/platform/users' import { Route as PlatformTenantsRouteImport } from './routes/platform/tenants' -import { Route as PlatformPartnersRouteImport } from './routes/platform/partners' import { Route as PlatformPartnershipsRouteImport } from './routes/platform/partnerships' +import { Route as PlatformPartnersRouteImport } from './routes/platform/partners' import { Route as PlatformDashboardRouteImport } from './routes/platform/dashboard' const TenantRoute = TenantRouteImport.update({ @@ -58,16 +60,26 @@ const TenantPartnersRoute = TenantPartnersRouteImport.update({ path: '/partners', getParentRoute: () => TenantRoute, } as any) -const TenantDashboardRoute = TenantDashboardRouteImport.update({ - id: '/dashboard', - path: '/dashboard', +const TenantEndpointsRoute = TenantEndpointsRouteImport.update({ + id: '/endpoints', + path: '/endpoints', getParentRoute: () => TenantRoute, } as any) -const TenantEdiToolRoute = TenantEdiToolRouteImport.update({ +const TenantEdi_toolRoute = TenantEdi_toolRouteImport.update({ id: '/edi_tool', path: '/edi_tool', getParentRoute: () => TenantRoute, } as any) +const TenantDevelopersRoute = TenantDevelopersRouteImport.update({ + id: '/developers', + path: '/developers', + getParentRoute: () => TenantRoute, +} as any) +const TenantDashboardRoute = TenantDashboardRouteImport.update({ + id: '/dashboard', + path: '/dashboard', + getParentRoute: () => TenantRoute, +} as any) const PlatformUsersRoute = PlatformUsersRouteImport.update({ id: '/users', path: '/users', @@ -78,16 +90,16 @@ const PlatformTenantsRoute = PlatformTenantsRouteImport.update({ path: '/tenants', getParentRoute: () => PlatformRoute, } as any) -const PlatformPartnersRoute = PlatformPartnersRouteImport.update({ - id: '/partners', - path: '/partners', - getParentRoute: () => PlatformRoute, -} as any) const PlatformPartnershipsRoute = PlatformPartnershipsRouteImport.update({ id: '/partnerships', path: '/partnerships', getParentRoute: () => PlatformRoute, } as any) +const PlatformPartnersRoute = PlatformPartnersRouteImport.update({ + id: '/partners', + path: '/partners', + getParentRoute: () => PlatformRoute, +} as any) const PlatformDashboardRoute = PlatformDashboardRouteImport.update({ id: '/dashboard', path: '/dashboard', @@ -104,7 +116,9 @@ export interface FileRoutesByFullPath { '/platform/tenants': typeof PlatformTenantsRoute '/platform/users': typeof PlatformUsersRoute '/tenant/dashboard': typeof TenantDashboardRoute - '/tenant/edi_tool': typeof TenantEdiToolRoute + '/tenant/developers': typeof TenantDevelopersRoute + '/tenant/edi_tool': typeof TenantEdi_toolRoute + '/tenant/endpoints': typeof TenantEndpointsRoute '/tenant/partners': typeof TenantPartnersRoute '/tenant/routes': typeof TenantRoutesRoute '/tenant/users': typeof TenantUsersRoute @@ -118,6 +132,9 @@ export interface FileRoutesByTo { '/platform/tenants': typeof PlatformTenantsRoute '/platform/users': typeof PlatformUsersRoute '/tenant/dashboard': typeof TenantDashboardRoute + '/tenant/developers': typeof TenantDevelopersRoute + '/tenant/edi_tool': typeof TenantEdi_toolRoute + '/tenant/endpoints': typeof TenantEndpointsRoute '/tenant/partners': typeof TenantPartnersRoute '/tenant/routes': typeof TenantRoutesRoute '/tenant/users': typeof TenantUsersRoute @@ -134,6 +151,9 @@ export interface FileRoutesById { '/platform/tenants': typeof PlatformTenantsRoute '/platform/users': typeof PlatformUsersRoute '/tenant/dashboard': typeof TenantDashboardRoute + '/tenant/developers': typeof TenantDevelopersRoute + '/tenant/edi_tool': typeof TenantEdi_toolRoute + '/tenant/endpoints': typeof TenantEndpointsRoute '/tenant/partners': typeof TenantPartnersRoute '/tenant/routes': typeof TenantRoutesRoute '/tenant/users': typeof TenantUsersRoute @@ -151,6 +171,9 @@ export interface FileRouteTypes { | '/platform/tenants' | '/platform/users' | '/tenant/dashboard' + | '/tenant/developers' + | '/tenant/edi_tool' + | '/tenant/endpoints' | '/tenant/partners' | '/tenant/routes' | '/tenant/users' @@ -164,6 +187,9 @@ export interface FileRouteTypes { | '/platform/tenants' | '/platform/users' | '/tenant/dashboard' + | '/tenant/developers' + | '/tenant/edi_tool' + | '/tenant/endpoints' | '/tenant/partners' | '/tenant/routes' | '/tenant/users' @@ -179,6 +205,9 @@ export interface FileRouteTypes { | '/platform/tenants' | '/platform/users' | '/tenant/dashboard' + | '/tenant/developers' + | '/tenant/edi_tool' + | '/tenant/endpoints' | '/tenant/partners' | '/tenant/routes' | '/tenant/users' @@ -235,10 +264,6 @@ declare module '@tanstack/react-router' { preLoaderRoute: typeof TenantRoutesRouteImport parentRoute: typeof TenantRoute } - "/tenant/edi_tool": { - "filePath": "tenant/edi_tool.tsx", - "parent": "/tenant" - }, '/tenant/partners': { id: '/tenant/partners' path: '/partners' @@ -246,6 +271,27 @@ declare module '@tanstack/react-router' { preLoaderRoute: typeof TenantPartnersRouteImport parentRoute: typeof TenantRoute } + '/tenant/endpoints': { + id: '/tenant/endpoints' + path: '/endpoints' + fullPath: '/tenant/endpoints' + preLoaderRoute: typeof TenantEndpointsRouteImport + parentRoute: typeof TenantRoute + } + '/tenant/edi_tool': { + id: '/tenant/edi_tool' + path: '/edi_tool' + fullPath: '/tenant/edi_tool' + preLoaderRoute: typeof TenantEdi_toolRouteImport + parentRoute: typeof TenantRoute + } + '/tenant/developers': { + id: '/tenant/developers' + path: '/developers' + fullPath: '/tenant/developers' + preLoaderRoute: typeof TenantDevelopersRouteImport + parentRoute: typeof TenantRoute + } '/tenant/dashboard': { id: '/tenant/dashboard' path: '/dashboard' @@ -267,6 +313,13 @@ declare module '@tanstack/react-router' { preLoaderRoute: typeof PlatformTenantsRouteImport parentRoute: typeof PlatformRoute } + '/platform/partnerships': { + id: '/platform/partnerships' + path: '/partnerships' + fullPath: '/platform/partnerships' + preLoaderRoute: typeof PlatformPartnershipsRouteImport + parentRoute: typeof PlatformRoute + } '/platform/partners': { id: '/platform/partners' path: '/partners' @@ -281,13 +334,6 @@ declare module '@tanstack/react-router' { preLoaderRoute: typeof PlatformDashboardRouteImport parentRoute: typeof PlatformRoute } - '/platform/partnerships': { - id: '/platform/partnerships' - path: '/partnerships' - fullPath: '/platform/partnerships' - preLoaderRoute: typeof PlatformPartnershipsRouteImport - parentRoute: typeof PlatformRoute - } } } @@ -325,7 +371,9 @@ const PlatformRouteWithChildren = PlatformRoute._addFileChildren( interface TenantRouteChildren { TenantDashboardRoute: typeof TenantDashboardRoute - TenantEdiToolRoute: typeof TenantEdiToolRoute + TenantDevelopersRoute: typeof TenantDevelopersRoute + TenantEdi_toolRoute: typeof TenantEdi_toolRoute + TenantEndpointsRoute: typeof TenantEndpointsRoute TenantPartnersRoute: typeof TenantPartnersRoute TenantRoutesRoute: typeof TenantRoutesRoute TenantUsersRoute: typeof TenantUsersRoute @@ -333,7 +381,9 @@ interface TenantRouteChildren { const TenantRouteChildren: TenantRouteChildren = { TenantDashboardRoute: TenantDashboardRoute, - TenantEdiToolRoute: TenantEdiToolRoute, + TenantDevelopersRoute: TenantDevelopersRoute, + TenantEdi_toolRoute: TenantEdi_toolRoute, + TenantEndpointsRoute: TenantEndpointsRoute, TenantPartnersRoute: TenantPartnersRoute, TenantRoutesRoute: TenantRoutesRoute, TenantUsersRoute: TenantUsersRoute, diff --git a/frontend/web/src/routes/tenant.tsx b/frontend/web/src/routes/tenant.tsx index 01c23e37..c8f3039e 100644 --- a/frontend/web/src/routes/tenant.tsx +++ b/frontend/web/src/routes/tenant.tsx @@ -11,6 +11,7 @@ import { LogOut, ChevronRight, Wrench, + Terminal, } from 'lucide-react' import { PartnersProvider } from '@/features/partners/context/PartnersContext' import { useDashboardData } from '@/features/dashboard/api/useDashboardData' @@ -102,8 +103,11 @@ export function AppLayout() {
Tools
-
System
- +
Developers
+ + +
Organization
+
Configuration
diff --git a/frontend/web/src/routes/tenant/developers.tsx b/frontend/web/src/routes/tenant/developers.tsx new file mode 100644 index 00000000..c6c76fab --- /dev/null +++ b/frontend/web/src/routes/tenant/developers.tsx @@ -0,0 +1,48 @@ +import { createRoute } from '@tanstack/react-router' +import { Route as appRoute } from '../tenant' +import { Terminal } from 'lucide-react' +import { ApiTokensTable } from '@/features/developers/components/ApiTokensTable' +import { CreateApiTokenModal } from '@/features/developers/components/CreateApiTokenModal' +import { useApiTokensQuery } from '@/features/developers/api/apiTokenHooks' + +export const Route = createRoute({ + getParentRoute: () => appRoute, + path: '/developers', + component: DevelopersPage, +}) + +export function DevelopersPage() { + const { data, isLoading } = useApiTokensQuery() + + const tokens = data?.tokens ?? [] + + return ( +
+ + {/* Header */} +
+
+
+

+ + API Access +

+
+
+ +
+
+

+ Manage API keys for programmatic machine-to-machine integration with your ERP systems. + Ensure you treat these credentials like passwords. +

+
+ + {/* Main Grid */} +
+ +
+ +
+ ) +} diff --git a/infra/zitadel/.terraform.lock.hcl b/infra/zitadel/.terraform.lock.hcl new file mode 100644 index 00000000..16a3fec7 --- /dev/null +++ b/infra/zitadel/.terraform.lock.hcl @@ -0,0 +1,23 @@ +# This file is maintained automatically by "terraform init". +# Manual edits may be lost in future updates. + +provider "registry.terraform.io/zitadel/zitadel" { + version = "3.3.0" + hashes = [ + "h1:e2+Vd0WvZDFBfwX/GxfZ3PBi3//APnvBkESi9h0gym4=", + "zh:07bf421dace1655b61a629ba195ad0982cd3871d70d68dce0026f62cf0cda42a", + "zh:113820a9d63169b4ceb92e83136f07a548ba5d16241bd027a01c29318913f0c7", + "zh:143ef6c3e5ff85cab1e582979db86f40316dbe137a6ff25d9704401d062812b6", + "zh:230a69339bbbb3508a6331ef32bc9a9229e764ae8c24fbccad708fe1da134f4a", + "zh:2563b7a855edf81cfa8515d80e11cd23f6427b2c9fb25abd987628c68ad81fc0", + "zh:385af8e3c9526c5599f78d51d5f3cf139317d288c2417f992bae78e5d584d532", + "zh:39628a705d915c7a1eea72f9688a888ac0a6d79358f6eb40a8adaf77710597be", + "zh:4ccca9a482fe3cebcafaf37b9bc1f578db2e4376682dfe69b896eca20b940747", + "zh:4dc6299607bd8cf89d9f085a18b0531e65421c1d6b15931e637d50d6279e75fa", + "zh:57aa24dd1170ef632d7850c4308ca650ed054c4e59dfaca0323580bf07c35769", + "zh:6f975d42de14b838c3b823d32a2f278d059bdc6cf81744ca914c5c20aa9d09e1", + "zh:7e2fc1525e68b280418ee886fa2a8bd5d957a02e8fd4d424fbf1a9eabf077400", + "zh:a33d53acc640dc93b81352ba633cf392bc8c7614a72d320d59d3dcdb22d73fc4", + "zh:f3e57de9fcf0b1abf415c47c0c3ea8c9b94eb4c4141aa51e3ac63cafafbea58d", + ] +} diff --git a/infra/zitadel/main.tf b/infra/zitadel/main.tf new file mode 100644 index 00000000..fe9e5c1b --- /dev/null +++ b/infra/zitadel/main.tf @@ -0,0 +1,121 @@ +terraform { + required_providers { + zitadel = { + source = "zitadel/zitadel" + } + } +} + +provider "zitadel" { + domain = "localhost" + insecure = true + port = "8080" + jwt_profile_file = "machinekey.json" +} + +# 1. Create the SOOPA Organization +resource "zitadel_org" "soopa" { + name = "SOOPA" +} + +# 2. Create Projects inside SOOPA +resource "zitadel_project" "edi" { + name = "Soopa EDI" + org_id = zitadel_org.soopa.id + project_role_assertion = true + project_role_check = false +} + +resource "zitadel_project" "ip" { + name = "Soopa Integration Platform" + org_id = zitadel_org.soopa.id + project_role_assertion = true + project_role_check = false +} + +resource "zitadel_project" "idp" { + name = "Soopa Intelligent Document Processing" + org_id = zitadel_org.soopa.id + project_role_assertion = true + project_role_check = false +} + +# 3. Create the EDI Web App Application +resource "zitadel_application_oidc" "edi_web_app" { + org_id = zitadel_org.soopa.id + project_id = zitadel_project.edi.id + name = "Soopa EDI Web App" + redirect_uris = ["http://localhost:5173/auth/callback", "http://localhost:5173/callback"] + post_logout_redirect_uris = ["http://localhost:5173", "http://localhost:5173/"] + response_types = ["OIDC_RESPONSE_TYPE_CODE"] + grant_types = ["OIDC_GRANT_TYPE_AUTHORIZATION_CODE"] + app_type = "OIDC_APP_TYPE_USER_AGENT" + auth_method_type = "OIDC_AUTH_METHOD_TYPE_NONE" + dev_mode = true +} + +# 4. Create the API Testing Machine User +resource "zitadel_machine_user" "api_test" { + org_id = zitadel_org.soopa.id + user_name = "api-test" + name = "API Test User" + description = "Integration testing user" + access_token_type = "ACCESS_TOKEN_TYPE_BEARER" +} + +# 5. Generate PAT for the Machine User +resource "zitadel_personal_access_token" "api_test_pat" { + org_id = zitadel_org.soopa.id + user_id = zitadel_machine_user.api_test.id +} + +# Outputs +output "soopa_org_id" { + value = zitadel_org.soopa.id +} + +output "edi_spa_client_id" { + value = zitadel_application_oidc.edi_web_app.client_id + sensitive = true +} + +output "api_test_pat_token" { + value = zitadel_personal_access_token.api_test_pat.token + sensitive = true +} + +# ============================================================================== +# TENANT 1 (Customer: Acme Corp) +# ============================================================================== + +# 1. Create the Acme Corp Organization +resource "zitadel_org" "acme" { + name = "Acme Corp" +} + +# 2. Create an API Testing Machine User for Acme Corp +resource "zitadel_machine_user" "acme_test" { + org_id = zitadel_org.acme.id + user_name = "acme-test" + name = "Acme Test User" + description = "Integration testing user for Tenant 1" + with_secret = false + access_token_type = "ACCESS_TOKEN_TYPE_BEARER" +} + +# 3. Generate a Personal Access Token (PAT) for the Acme Test User +resource "zitadel_personal_access_token" "acme_test_pat" { + org_id = zitadel_org.acme.id + user_id = zitadel_machine_user.acme_test.id +} + + + +output "acme_org_id" { + value = zitadel_org.acme.id +} + +output "acme_test_pat_token" { + value = zitadel_personal_access_token.acme_test_pat.token + sensitive = true +} diff --git a/infra/zitadel/terraform.tfstate b/infra/zitadel/terraform.tfstate new file mode 100644 index 00000000..fe53d519 --- /dev/null +++ b/infra/zitadel/terraform.tfstate @@ -0,0 +1,452 @@ +{ + "version": 4, + "terraform_version": "1.9.3", + "serial": 48, + "lineage": "e8feda52-eff6-043c-e313-a754a61192ed", + "outputs": { + "acme_org_id": { + "value": "381030164001193991", + "type": "string" + }, + "acme_test_pat_token": { + "value": "H4GohgBOaYBLJjHya5OYa1IiPFohu5-DYaMFChvdhkx6iW7yciQPeRiR7iEgldPQq7BSIOE", + "type": "string", + "sensitive": true + }, + "api_test_pat_token": { + "value": "0N9RwsW6nXtyvKTc31gC3NFjlFLZhLjNFYf0dXorQhhBP9uQIVX_TcM2bCWyLDDTVPcps4E", + "type": "string", + "sensitive": true + }, + "edi_spa_client_id": { + "value": "381027859583533063", + "type": "string", + "sensitive": true + }, + "soopa_org_id": { + "value": "381027855791816711", + "type": "string" + } + }, + "resources": [ + { + "mode": "managed", + "type": "zitadel_application_oidc", + "name": "edi_web_app", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "access_token_role_assertion": false, + "access_token_type": "OIDC_TOKEN_TYPE_BEARER", + "additional_origins": [], + "app_type": "OIDC_APP_TYPE_USER_AGENT", + "auth_method_type": "OIDC_AUTH_METHOD_TYPE_NONE", + "back_channel_logout_uri": "", + "client_id": "381027859583533063", + "client_secret": "", + "clock_skew": "0s", + "compliance_problems": [ + { + "key": "Application.OIDC.V1.NotCompliant", + "message": "Your configuration is not compliant and differs from OIDC 1.0 standard." + }, + { + "key": "Application.OIDC.V1.NotCompliant", + "message": "Your configuration is not compliant and differs from OIDC 1.0 standard." + }, + { + "key": "Application.OIDC.V1.Code.RedirectUris.HttpOnlyForWeb", + "message": "Grant type code only allowed http redirect uris for apptype web." + } + ], + "dev_mode": true, + "grant_types": [ + "OIDC_GRANT_TYPE_AUTHORIZATION_CODE" + ], + "id": "381027859583467527", + "id_token_role_assertion": false, + "id_token_userinfo_assertion": false, + "login_version": [], + "name": "Soopa EDI Web App", + "none_compliant": true, + "org_id": "381027855791816711", + "post_logout_redirect_uris": [ + "http://localhost:5173", + "http://localhost:5173/" + ], + "project_id": "381027859348652039", + "redirect_uris": [ + "http://localhost:5173/auth/callback", + "http://localhost:5173/callback" + ], + "response_types": [ + "OIDC_RESPONSE_TYPE_CODE" + ], + "skip_native_app_success_page": false, + "version": "OIDC_VERSION_1_0" + }, + "sensitive_attributes": [ + [ + { + "type": "get_attr", + "value": "client_id" + } + ], + [ + { + "type": "get_attr", + "value": "client_secret" + } + ] + ], + "private": "bnVsbA==", + "dependencies": [ + "zitadel_org.soopa", + "zitadel_project.edi" + ] + } + ] + }, + { + "mode": "managed", + "type": "zitadel_human_user", + "name": "acme_admin", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "display_name": "Acme Admin", + "email": "admin@acme.localhost", + "first_name": "Acme", + "gender": "GENDER_UNSPECIFIED", + "id": "381032595003015175", + "idp_links": [], + "initial_hashed_password": null, + "initial_password": null, + "initial_skip_password_change": false, + "is_email_verified": true, + "is_phone_verified": false, + "last_name": "Admin", + "login_names": [ + "admin" + ], + "metadata": [], + "nick_name": "", + "org_id": "381030164001193991", + "phone": "", + "preferred_language": "und", + "preferred_login_name": "admin", + "state": "USER_STATE_INITIAL", + "totp_secret": null, + "user_id": null, + "user_name": "admin" + }, + "sensitive_attributes": [ + [ + { + "type": "get_attr", + "value": "initial_hashed_password" + } + ], + [ + { + "type": "get_attr", + "value": "totp_secret" + } + ], + [ + { + "type": "get_attr", + "value": "initial_password" + } + ] + ], + "private": "bnVsbA==", + "dependencies": [ + "zitadel_org.acme" + ] + } + ] + }, + { + "mode": "managed", + "type": "zitadel_machine_user", + "name": "acme_test", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "access_token_type": "ACCESS_TOKEN_TYPE_BEARER", + "client_id": null, + "client_secret": null, + "description": "Integration testing user for Tenant 1", + "id": "381030167608295431", + "login_names": [ + "acme-test" + ], + "name": "Acme Test User", + "org_id": "381030164001193991", + "preferred_login_name": "acme-test", + "state": "USER_STATE_ACTIVE", + "user_id": null, + "user_name": "acme-test", + "with_secret": false + }, + "sensitive_attributes": [ + [ + { + "type": "get_attr", + "value": "client_secret" + } + ], + [ + { + "type": "get_attr", + "value": "client_id" + } + ] + ], + "private": "bnVsbA==", + "dependencies": [ + "zitadel_org.acme" + ] + } + ] + }, + { + "mode": "managed", + "type": "zitadel_machine_user", + "name": "api_test", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "access_token_type": "ACCESS_TOKEN_TYPE_BEARER", + "client_id": null, + "client_secret": null, + "description": "Integration testing user", + "id": "381027859348783111", + "login_names": [ + "api-test" + ], + "name": "API Test User", + "org_id": "381027855791816711", + "preferred_login_name": "api-test", + "state": "USER_STATE_ACTIVE", + "user_id": null, + "user_name": "api-test", + "with_secret": false + }, + "sensitive_attributes": [ + [ + { + "type": "get_attr", + "value": "client_secret" + } + ], + [ + { + "type": "get_attr", + "value": "client_id" + } + ] + ], + "private": "bnVsbA==", + "dependencies": [ + "zitadel_org.soopa" + ] + } + ] + }, + { + "mode": "managed", + "type": "zitadel_org", + "name": "acme", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "admins": [], + "id": "381030164001193991", + "is_default": null, + "name": "Acme Corp", + "org_id": null, + "primary_domain": "acme-corp.localhost", + "state": "ORG_STATE_ACTIVE" + }, + "sensitive_attributes": [], + "private": "bnVsbA==" + } + ] + }, + { + "mode": "managed", + "type": "zitadel_org", + "name": "soopa", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "admins": [], + "id": "381027855791816711", + "is_default": null, + "name": "SOOPA", + "org_id": null, + "primary_domain": "soopa.localhost", + "state": "ORG_STATE_ACTIVE" + }, + "sensitive_attributes": [], + "private": "bnVsbA==" + } + ] + }, + { + "mode": "managed", + "type": "zitadel_personal_access_token", + "name": "acme_test_pat", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "expiration_date": null, + "id": "381037097370255367", + "org_id": "381030164001193991", + "token": "H4GohgBOaYBLJjHya5OYa1IiPFohu5-DYaMFChvdhkx6iW7yciQPeRiR7iEgldPQq7BSIOE", + "user_id": "381030167608295431" + }, + "sensitive_attributes": [ + [ + { + "type": "get_attr", + "value": "token" + } + ] + ], + "private": "bnVsbA==", + "dependencies": [ + "zitadel_machine_user.acme_test", + "zitadel_org.acme" + ] + } + ] + }, + { + "mode": "managed", + "type": "zitadel_personal_access_token", + "name": "api_test_pat", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "expiration_date": null, + "id": "381037097470918663", + "org_id": "381027855791816711", + "token": "0N9RwsW6nXtyvKTc31gC3NFjlFLZhLjNFYf0dXorQhhBP9uQIVX_TcM2bCWyLDDTVPcps4E", + "user_id": "381027859348783111" + }, + "sensitive_attributes": [ + [ + { + "type": "get_attr", + "value": "token" + } + ] + ], + "private": "bnVsbA==", + "dependencies": [ + "zitadel_machine_user.api_test", + "zitadel_org.soopa" + ] + } + ] + }, + { + "mode": "managed", + "type": "zitadel_project", + "name": "edi", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "has_project_check": false, + "id": "381027859348652039", + "name": "Soopa EDI", + "org_id": "381027855791816711", + "private_labeling_setting": "PRIVATE_LABELING_SETTING_UNSPECIFIED", + "project_role_assertion": true, + "project_role_check": false, + "state": "PROJECT_STATE_ACTIVE" + }, + "sensitive_attributes": [], + "private": "bnVsbA==", + "dependencies": [ + "zitadel_org.soopa" + ] + } + ] + }, + { + "mode": "managed", + "type": "zitadel_project", + "name": "idp", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "has_project_check": false, + "id": "381027859348586503", + "name": "Soopa Intelligent Document Processing", + "org_id": "381027855791816711", + "private_labeling_setting": "PRIVATE_LABELING_SETTING_UNSPECIFIED", + "project_role_assertion": true, + "project_role_check": false, + "state": "PROJECT_STATE_ACTIVE" + }, + "sensitive_attributes": [], + "private": "bnVsbA==", + "dependencies": [ + "zitadel_org.soopa" + ] + } + ] + }, + { + "mode": "managed", + "type": "zitadel_project", + "name": "ip", + "provider": "provider[\"registry.terraform.io/zitadel/zitadel\"]", + "instances": [ + { + "schema_version": 0, + "attributes": { + "has_project_check": false, + "id": "381027859348717575", + "name": "Soopa Integration Platform", + "org_id": "381027855791816711", + "private_labeling_setting": "PRIVATE_LABELING_SETTING_UNSPECIFIED", + "project_role_assertion": true, + "project_role_check": false, + "state": "PROJECT_STATE_ACTIVE" + }, + "sensitive_attributes": [], + "private": "bnVsbA==", + "dependencies": [ + "zitadel_org.soopa" + ] + } + ] + } + ], + "check_results": null +} diff --git a/libs/config/src/config/settings.py b/libs/config/src/config/settings.py index 15ec011c..ff90c05f 100644 --- a/libs/config/src/config/settings.py +++ b/libs/config/src/config/settings.py @@ -95,6 +95,7 @@ class AppSettings(BaseSettings): env: Literal["development", "staging", "production"] = Field(default="development") log_level: Literal["DEBUG", "INFO", "WARNING", "ERROR"] = Field(default="INFO") + storage_backend: Literal["postgres", "s3"] = Field(default="postgres") database: DatabaseSettings = Field(default_factory=DatabaseSettings) s3: S3Settings = Field(default_factory=S3Settings) diff --git a/libs/database/src/database/migrations/global/versions/142d93fd6c4c_global_initial_schema.py b/libs/database/src/database/migrations/global/versions/94b943b2920d_global_initial_schema.py similarity index 84% rename from libs/database/src/database/migrations/global/versions/142d93fd6c4c_global_initial_schema.py rename to libs/database/src/database/migrations/global/versions/94b943b2920d_global_initial_schema.py index 6a4cb06e..bb6a81c7 100644 --- a/libs/database/src/database/migrations/global/versions/142d93fd6c4c_global_initial_schema.py +++ b/libs/database/src/database/migrations/global/versions/94b943b2920d_global_initial_schema.py @@ -1,8 +1,8 @@ """global_initial_schema -Revision ID: 142d93fd6c4c +Revision ID: 94b943b2920d Revises: -Create Date: 2026-07-06 14:56:17.633049 +Create Date: 2026-07-10 11:42:52.960482 """ @@ -13,7 +13,7 @@ from sqlalchemy.dialects import postgresql # revision identifiers, used by Alembic. -revision: str = "142d93fd6c4c" +revision: str = "94b943b2920d" down_revision: str | Sequence[str] | None = None branch_labels: str | Sequence[str] | None = None depends_on: str | Sequence[str] | None = None @@ -84,9 +84,9 @@ def upgrade() -> None: ) op.create_table( "as2_partners", - sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("tenant_id", sa.Integer(), nullable=True), sa.Column("is_local", sa.Boolean(), nullable=False), + sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("name", sa.String(length=255), nullable=False), sa.Column("as2_id", sa.String(length=255), nullable=False), sa.Column("public_cert_pem", sa.Text(), nullable=True), @@ -131,10 +131,27 @@ def upgrade() -> None: postgresql_where=sa.text("status = 'PENDING'"), ) op.create_table( - "sftp_partners", + "api_tokens", sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("tenant_id", sa.Integer(), nullable=False), sa.Column("name", sa.String(length=255), nullable=False), + sa.Column("client_id", sa.String(length=64), nullable=False), + sa.Column("secret_hash", sa.String(length=64), nullable=False), + sa.Column("last_used_at", sa.DateTime(), nullable=True), + sa.Column("expires_at", sa.DateTime(), nullable=True), + sa.Column("active", sa.Boolean(), nullable=False, server_default="true"), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), + sa.ForeignKeyConstraint(["tenant_id"], ["tenants.id"], ondelete="CASCADE"), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index("ix_api_tokens_tenant_id", "api_tokens", ["tenant_id"], unique=False) + op.create_index("ix_api_tokens_client_id", "api_tokens", ["client_id"], unique=True) + op.create_table( + "sftp_partners", + sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), + sa.Column("name", sa.String(length=255), nullable=False), sa.Column("host", sa.String(length=1024), nullable=False), sa.Column("port", sa.Integer(), nullable=False), sa.Column("username", sa.String(length=255), nullable=False), @@ -144,6 +161,8 @@ def upgrade() -> None: sa.Column("password_encrypted", sa.String(length=1024), nullable=True), sa.Column("credentials_vault_ref", sa.String(length=255), nullable=True), sa.Column("active", sa.Boolean(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), sa.ForeignKeyConstraint(["tenant_id"], ["tenants.id"], ondelete="CASCADE"), sa.PrimaryKeyConstraint("id"), ) @@ -163,23 +182,25 @@ def upgrade() -> None: ) op.create_table( "webhooks", - sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("name", sa.String(length=255), nullable=False), sa.Column("url", sa.String(length=1024), nullable=False), sa.Column("auth_header_vault_ref", sa.String(length=255), nullable=True), sa.Column("active", sa.Boolean(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), sa.ForeignKeyConstraint(["tenant_id"], ["tenants.id"], ondelete="CASCADE"), sa.PrimaryKeyConstraint("id"), ) op.create_index(op.f("ix_webhooks_tenant_id"), "webhooks", ["tenant_id"], unique=False) op.create_table( "as2_partnerships", - sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("tenant_id", sa.Integer(), nullable=True), - sa.Column("name", sa.String(length=255), nullable=False), sa.Column("local_partner_id", sa.UUID(), nullable=False), sa.Column("remote_partner_id", sa.UUID(), nullable=False), + sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), + sa.Column("name", sa.String(length=255), nullable=False), sa.Column("credentials_vault_ref", sa.String(length=255), nullable=True), sa.Column("mdn_type", sa.String(length=50), nullable=False), sa.Column("mdn_url", sa.String(length=1024), nullable=True), @@ -197,19 +218,23 @@ def upgrade() -> None: ) op.create_table( "inbound_routes", - sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column("webhook_id", sa.UUID(), nullable=True), + sa.Column("as2_partner_id", sa.UUID(), nullable=True), + sa.Column("sftp_partner_id", sa.UUID(), nullable=True), + sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("name", sa.String(length=255), nullable=False), sa.Column("isa_sender_id", sa.String(length=255), nullable=False), sa.Column("isa_receiver_id", sa.String(length=255), nullable=False), + sa.Column("gs_sender_id", sa.String(length=255), nullable=True), + sa.Column("gs_receiver_id", sa.String(length=255), nullable=True), sa.Column("transaction_type", sa.String(length=50), nullable=False), sa.Column( "processing_mode", sa.String(length=50), server_default="TRANSLATE", nullable=False ), - sa.Column("webhook_id", sa.UUID(), nullable=True), - sa.Column("as2_partner_id", sa.UUID(), nullable=True), - sa.Column("sftp_partner_id", sa.UUID(), nullable=True), sa.Column("active", sa.Boolean(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), sa.CheckConstraint( "(webhook_id IS NOT NULL)::int + (as2_partner_id IS NOT NULL)::int + (sftp_partner_id IS NOT NULL)::int = 1", name="chk_inbound_routes_exactly_one_dest", @@ -241,18 +266,27 @@ def upgrade() -> None: ) op.create_table( "outbound_routes", - sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column("as2_partner_id", sa.UUID(), nullable=True), + sa.Column("sftp_partner_id", sa.UUID(), nullable=True), + sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), + sa.Column("trading_partner_id", sa.String(length=255), nullable=False), sa.Column("name", sa.String(length=255), nullable=False), sa.Column("isa_sender_id", sa.String(length=255), nullable=False), + sa.Column("isa_sender_qualifier", sa.String(length=2), nullable=True), sa.Column("isa_receiver_id", sa.String(length=255), nullable=False), + sa.Column("isa_receiver_qualifier", sa.String(length=2), nullable=True), + sa.Column("gs_sender_id", sa.String(length=255), nullable=False), + sa.Column("gs_receiver_id", sa.String(length=255), nullable=False), sa.Column("transaction_type", sa.String(length=50), nullable=False), + sa.Column("default_standard", sa.String(length=50), server_default="x12", nullable=False), + sa.Column("default_version", sa.String(length=50), server_default="004010", nullable=False), sa.Column( "processing_mode", sa.String(length=50), server_default="TRANSLATE", nullable=False ), - sa.Column("as2_partner_id", sa.UUID(), nullable=True), - sa.Column("sftp_partner_id", sa.UUID(), nullable=True), sa.Column("active", sa.Boolean(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), sa.CheckConstraint( "(as2_partner_id IS NOT NULL)::int + (sftp_partner_id IS NOT NULL)::int = 1", name="chk_outbound_routes_exactly_one_dest", @@ -272,9 +306,9 @@ def upgrade() -> None: op.f("ix_outbound_routes_tenant_id"), "outbound_routes", ["tenant_id"], unique=False ) op.create_index( - "ix_outbound_routes_unique_active", + "ix_outbound_routes_unique_trading_partner_id", "outbound_routes", - ["tenant_id", "isa_sender_id", "isa_receiver_id", "transaction_type"], + ["tenant_id", "trading_partner_id"], unique=True, postgresql_where=sa.text("active = true"), ) @@ -285,7 +319,7 @@ def downgrade() -> None: """Downgrade schema.""" # ### commands auto generated by Alembic - please adjust! ### op.drop_index( - "ix_outbound_routes_unique_active", + "ix_outbound_routes_unique_trading_partner_id", table_name="outbound_routes", postgresql_where=sa.text("active = true"), ) @@ -310,6 +344,9 @@ def downgrade() -> None: postgresql_where=sa.text("status = 'PENDING'"), ) op.drop_table("outbox") + op.drop_index("ix_api_tokens_client_id", table_name="api_tokens") + op.drop_index("ix_api_tokens_tenant_id", table_name="api_tokens") + op.drop_table("api_tokens") op.drop_index( "uq_global_as2_id", table_name="as2_partners", postgresql_where=sa.text("tenant_id IS NULL") ) diff --git a/libs/database/src/database/migrations/tenant/versions/34e9d4ab146a_tenant_initial_schema.py b/libs/database/src/database/migrations/tenant/versions/0841cfb2afb2_tenant_initial_schema.py similarity index 67% rename from libs/database/src/database/migrations/tenant/versions/34e9d4ab146a_tenant_initial_schema.py rename to libs/database/src/database/migrations/tenant/versions/0841cfb2afb2_tenant_initial_schema.py index fed366dc..8ab997f9 100644 --- a/libs/database/src/database/migrations/tenant/versions/34e9d4ab146a_tenant_initial_schema.py +++ b/libs/database/src/database/migrations/tenant/versions/0841cfb2afb2_tenant_initial_schema.py @@ -1,19 +1,20 @@ """tenant_initial_schema -Revision ID: 34e9d4ab146a +Revision ID: 0841cfb2afb2 Revises: -Create Date: 2026-07-06 14:56:18.175654 +Create Date: 2026-07-10 11:42:54.246255 """ from collections.abc import Sequence +import database.models.data_plane import sqlalchemy as sa from alembic import op from sqlalchemy.dialects import postgresql # revision identifiers, used by Alembic. -revision: str = "34e9d4ab146a" +revision: str = "0841cfb2afb2" down_revision: str | Sequence[str] | None = None branch_labels: str | Sequence[str] | None = None depends_on: str | Sequence[str] | None = None @@ -30,7 +31,12 @@ def upgrade() -> None: sa.Column("status", sa.String(length=50), nullable=False), sa.Column("raw_content", sa.Text(), nullable=True), sa.Column("received_at", sa.DateTime(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), sa.PrimaryKeyConstraint("id"), ) op.create_index(op.f("ix_ack_receipts_tenant_id"), "ack_receipts", ["tenant_id"], unique=False) @@ -44,21 +50,35 @@ def upgrade() -> None: sa.Column("webhook_url", sa.String(length=1024), nullable=True), sa.Column("http_status_code", sa.Integer(), nullable=True), sa.Column("target_format", sa.String(length=50), nullable=True), - sa.Column("request", sa.String(length=1024), nullable=False), + sa.Column("payload", postgresql.JSONB(astext_type=sa.Text()), nullable=True), + sa.Column("storage_uri", sa.String(length=1024), nullable=True), sa.Column("response", sa.Text(), nullable=True), sa.Column("headers", postgresql.JSONB(astext_type=sa.Text()), nullable=True), sa.Column("status", sa.String(length=50), nullable=False), sa.Column("created_at", sa.DateTime(), nullable=False), sa.Column("updated_at", sa.DateTime(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), + sa.CheckConstraint( + "(payload IS NOT NULL OR storage_uri IS NOT NULL)", name="chk_apigw_data_or_uri" + ), sa.PrimaryKeyConstraint("id"), ) op.create_index(op.f("ix_api_gateway_tenant_id"), "api_gateway", ["tenant_id"], unique=False) op.create_index(op.f("ix_api_gateway_trace_id"), "api_gateway", ["trace_id"], unique=False) op.create_table( "as2_partners", + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), - sa.Column("is_local", sa.Boolean(), nullable=False), sa.Column("name", sa.String(length=255), nullable=False), sa.Column("as2_id", sa.String(length=255), nullable=False), sa.Column("public_cert_pem", sa.Text(), nullable=True), @@ -69,7 +89,8 @@ def upgrade() -> None: sa.Column("prev_private_key_vault_ref", sa.String(length=255), nullable=True), sa.Column("url", sa.String(length=1024), nullable=True), sa.Column("active", sa.Boolean(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), sa.PrimaryKeyConstraint("id"), ) op.create_index(op.f("ix_as2_partners_tenant_id"), "as2_partners", ["tenant_id"], unique=False) @@ -83,50 +104,16 @@ def upgrade() -> None: sa.Column("error", sa.Text(), nullable=True), sa.Column("metadata", postgresql.JSONB(astext_type=sa.Text()), nullable=True), sa.Column("created_at", sa.DateTime(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), sa.PrimaryKeyConstraint("id"), ) op.create_index(op.f("ix_audit_log_tenant_id"), "audit_log", ["tenant_id"], unique=False) op.create_index(op.f("ix_audit_log_trace_id"), "audit_log", ["trace_id"], unique=False) - op.create_table( - "edi_messages", - sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), - sa.Column("trace_id", sa.UUID(), nullable=False), - sa.Column("direction", sa.String(length=50), nullable=False), - sa.Column("connection_type", sa.String(length=50), nullable=False), - sa.Column("sender_id", sa.String(length=255), nullable=True), - sa.Column("receiver_id", sa.String(length=255), nullable=True), - sa.Column("message_id", sa.String(length=255), nullable=True), - sa.Column("mdn_id", sa.String(length=255), nullable=True), - sa.Column("mdn_mode", sa.String(length=50), nullable=True), - sa.Column("mdn_response", sa.Text(), nullable=True), - sa.Column("file_name", sa.String(length=1024), nullable=True), - sa.Column("content_type", sa.String(length=255), nullable=True), - sa.Column("signature_algorithm", sa.String(length=50), nullable=True), - sa.Column("encryption_algorithm", sa.String(length=50), nullable=True), - sa.Column("is_resend", sa.Boolean(), nullable=False), - sa.Column("status_message", sa.Text(), nullable=True), - sa.Column("state", sa.String(length=255), nullable=True), - sa.Column("msg_headers", sa.Text(), nullable=True), - sa.Column("interchange_control_no", sa.String(length=255), nullable=True), - sa.Column("transaction_type", sa.String(length=50), nullable=True), - sa.Column("format_standard", sa.String(length=50), nullable=True), - sa.Column("edi_data", sa.Text(), nullable=False), - sa.Column("file_size_bytes", sa.BigInteger(), nullable=True), - sa.Column("status", sa.String(length=50), nullable=False), - sa.Column("created_at", sa.DateTime(), nullable=False), - sa.Column("updated_at", sa.DateTime(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), - sa.PrimaryKeyConstraint("id"), - ) - op.create_index(op.f("ix_edi_messages_tenant_id"), "edi_messages", ["tenant_id"], unique=False) - op.create_index(op.f("ix_edi_messages_trace_id"), "edi_messages", ["trace_id"], unique=False) - op.create_index( - "ix_edi_msgs_sender_recv", - "edi_messages", - ["sender_id", "receiver_id", "created_at"], - unique=False, - ) op.create_table( "jobs", sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), @@ -137,7 +124,12 @@ def upgrade() -> None: sa.Column("error_message", sa.Text(), nullable=True), sa.Column("created_at", sa.DateTime(), nullable=False), sa.Column("updated_at", sa.DateTime(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), sa.PrimaryKeyConstraint("id"), ) op.create_index(op.f("ix_jobs_tenant_id"), "jobs", ["tenant_id"], unique=False) @@ -147,7 +139,12 @@ def upgrade() -> None: sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("attempts", sa.Integer(), nullable=False), sa.Column("published_at", sa.DateTime(), nullable=True), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), sa.Column("idempotency_key", sa.UUID(), nullable=False), sa.Column("event_type", sa.String(length=100), nullable=False), sa.Column("payload", postgresql.JSONB(astext_type=sa.Text()), nullable=False), @@ -168,7 +165,12 @@ def upgrade() -> None: "processed_events", sa.Column("idempotency_key", sa.UUID(), nullable=False), sa.Column("processed_at", sa.DateTime(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), sa.PrimaryKeyConstraint("idempotency_key"), ) op.create_index( @@ -176,6 +178,12 @@ def upgrade() -> None: ) op.create_table( "sftp_partners", + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("name", sa.String(length=255), nullable=False), sa.Column("host", sa.String(length=1024), nullable=False), @@ -187,7 +195,8 @@ def upgrade() -> None: sa.Column("password_encrypted", sa.String(length=1024), nullable=True), sa.Column("credentials_vault_ref", sa.String(length=255), nullable=True), sa.Column("active", sa.Boolean(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), sa.PrimaryKeyConstraint("id"), ) op.create_index( @@ -195,21 +204,34 @@ def upgrade() -> None: ) op.create_table( "webhooks", + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("name", sa.String(length=255), nullable=False), sa.Column("url", sa.String(length=1024), nullable=False), sa.Column("auth_header_vault_ref", sa.String(length=255), nullable=True), sa.Column("active", sa.Boolean(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), sa.PrimaryKeyConstraint("id"), ) op.create_index(op.f("ix_webhooks_tenant_id"), "webhooks", ["tenant_id"], unique=False) op.create_table( "as2_partnerships", - sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), - sa.Column("name", sa.String(length=255), nullable=False), sa.Column("local_partner_id", sa.UUID(), nullable=False), sa.Column("remote_partner_id", sa.UUID(), nullable=False), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), + sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), + sa.Column("name", sa.String(length=255), nullable=False), sa.Column("credentials_vault_ref", sa.String(length=255), nullable=True), sa.Column("mdn_type", sa.String(length=50), nullable=False), sa.Column("mdn_url", sa.String(length=1024), nullable=True), @@ -217,7 +239,8 @@ def upgrade() -> None: sa.Column("signature_algorithm", sa.String(length=50), nullable=False), sa.Column("advanced_flags", postgresql.JSONB(astext_type=sa.Text()), nullable=True), sa.Column("active", sa.Boolean(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), sa.ForeignKeyConstraint(["local_partner_id"], ["as2_partners.id"], ondelete="CASCADE"), sa.ForeignKeyConstraint(["remote_partner_id"], ["as2_partners.id"], ondelete="CASCADE"), sa.PrimaryKeyConstraint("id"), @@ -227,19 +250,28 @@ def upgrade() -> None: ) op.create_table( "inbound_routes", + sa.Column("webhook_id", sa.UUID(), nullable=True), + sa.Column("as2_partner_id", sa.UUID(), nullable=True), + sa.Column("sftp_partner_id", sa.UUID(), nullable=True), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("name", sa.String(length=255), nullable=False), sa.Column("isa_sender_id", sa.String(length=255), nullable=False), sa.Column("isa_receiver_id", sa.String(length=255), nullable=False), + sa.Column("gs_sender_id", sa.String(length=255), nullable=True), + sa.Column("gs_receiver_id", sa.String(length=255), nullable=True), sa.Column("transaction_type", sa.String(length=50), nullable=False), sa.Column( "processing_mode", sa.String(length=50), server_default="TRANSLATE", nullable=False ), - sa.Column("webhook_id", sa.UUID(), nullable=True), - sa.Column("as2_partner_id", sa.UUID(), nullable=True), - sa.Column("sftp_partner_id", sa.UUID(), nullable=True), sa.Column("active", sa.Boolean(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), sa.CheckConstraint( "(webhook_id IS NOT NULL)::int + (as2_partner_id IS NOT NULL)::int + (sftp_partner_id IS NOT NULL)::int = 1", name="chk_inbound_routes_exactly_one_dest", @@ -270,18 +302,32 @@ def upgrade() -> None: ) op.create_table( "outbound_routes", + sa.Column("as2_partner_id", sa.UUID(), nullable=True), + sa.Column("sftp_partner_id", sa.UUID(), nullable=True), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), + sa.Column("trading_partner_id", sa.String(length=255), nullable=False), sa.Column("name", sa.String(length=255), nullable=False), sa.Column("isa_sender_id", sa.String(length=255), nullable=False), + sa.Column("isa_sender_qualifier", sa.String(length=2), nullable=True), sa.Column("isa_receiver_id", sa.String(length=255), nullable=False), + sa.Column("isa_receiver_qualifier", sa.String(length=2), nullable=True), + sa.Column("gs_sender_id", sa.String(length=255), nullable=False), + sa.Column("gs_receiver_id", sa.String(length=255), nullable=False), sa.Column("transaction_type", sa.String(length=50), nullable=False), + sa.Column("default_standard", sa.String(length=50), server_default="x12", nullable=False), + sa.Column("default_version", sa.String(length=50), server_default="004010", nullable=False), sa.Column( "processing_mode", sa.String(length=50), server_default="TRANSLATE", nullable=False ), - sa.Column("as2_partner_id", sa.UUID(), nullable=True), - sa.Column("sftp_partner_id", sa.UUID(), nullable=True), sa.Column("active", sa.Boolean(), nullable=False), - sa.Column("tenant_id", sa.Integer(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), sa.CheckConstraint( "(as2_partner_id IS NOT NULL)::int + (sftp_partner_id IS NOT NULL)::int = 1", name="chk_outbound_routes_exactly_one_dest", @@ -300,20 +346,143 @@ def upgrade() -> None: op.f("ix_outbound_routes_tenant_id"), "outbound_routes", ["tenant_id"], unique=False ) op.create_index( - "ix_outbound_routes_unique_active", + "ix_outbound_routes_unique_trading_partner_id", "outbound_routes", - ["tenant_id", "isa_sender_id", "isa_receiver_id", "transaction_type"], + ["tenant_id", "trading_partner_id"], unique=True, postgresql_where=sa.text("active = true"), ) + op.create_table( + "edi_json", + sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), + sa.Column("trace_id", sa.UUID(), nullable=False), + sa.Column("direction", sa.String(length=50), nullable=False), + sa.Column("outbound_route_id", sa.UUID(), nullable=True), + sa.Column("transaction_type", sa.String(length=50), nullable=True), + sa.Column("standard", sa.String(length=50), nullable=True), + sa.Column("sender_id", sa.String(length=255), nullable=True), + sa.Column("receiver_id", sa.String(length=255), nullable=True), + sa.Column("business_metadata", postgresql.JSONB(astext_type=sa.Text()), nullable=True), + sa.Column("payload", postgresql.JSONB(astext_type=sa.Text()), nullable=True), + sa.Column("storage_uri", sa.String(length=1024), nullable=True), + sa.Column("status", sa.String(length=50), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), + sa.CheckConstraint( + "(payload IS NOT NULL OR storage_uri IS NOT NULL)", name="chk_edi_json_data_or_uri" + ), + sa.ForeignKeyConstraint( + ["outbound_route_id"], + ["outbound_routes.id"], + ), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index( + "ix_edi_json_business_metadata", + "edi_json", + ["business_metadata"], + unique=False, + postgresql_using="gin", + ) + op.create_index( + op.f("ix_edi_json_outbound_route_id"), "edi_json", ["outbound_route_id"], unique=False + ) + op.create_index( + "ix_edi_json_sender_recv", + "edi_json", + ["sender_id", "receiver_id", "created_at"], + unique=False, + ) + op.create_index(op.f("ix_edi_json_tenant_id"), "edi_json", ["tenant_id"], unique=False) + op.create_index(op.f("ix_edi_json_trace_id"), "edi_json", ["trace_id"], unique=False) + op.create_index( + op.f("ix_edi_json_transaction_type"), "edi_json", ["transaction_type"], unique=False + ) + op.create_table( + "edi_messages", + sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), + sa.Column("trace_id", sa.UUID(), nullable=False), + sa.Column("direction", sa.String(length=50), nullable=False), + sa.Column("connection_type", sa.String(length=50), nullable=True), + sa.Column("sender_id", sa.String(length=255), nullable=True), + sa.Column("receiver_id", sa.String(length=255), nullable=True), + sa.Column("message_id", sa.String(length=255), nullable=True), + sa.Column("mdn_id", sa.String(length=255), nullable=True), + sa.Column("mdn_mode", sa.String(length=50), nullable=True), + sa.Column("mdn_response", sa.Text(), nullable=True), + sa.Column("file_name", sa.String(length=1024), nullable=True), + sa.Column("content_type", sa.String(length=255), nullable=True), + sa.Column("signature_algorithm", sa.String(length=50), nullable=True), + sa.Column("encryption_algorithm", sa.String(length=50), nullable=True), + sa.Column("is_resend", sa.Boolean(), nullable=False), + sa.Column("status_message", sa.Text(), nullable=True), + sa.Column("state", sa.String(length=255), nullable=True), + sa.Column("msg_headers", sa.Text(), nullable=True), + sa.Column("outbound_route_id", sa.UUID(), nullable=True), + sa.Column("interchange_control_no", sa.String(length=255), nullable=True), + sa.Column("transaction_type", sa.String(length=50), nullable=True), + sa.Column("format_standard", sa.String(length=50), nullable=True), + sa.Column("edi_data", database.models.data_plane.SanitizedText(), nullable=True), + sa.Column("storage_uri", sa.String(length=1024), nullable=True), + sa.Column("file_size_bytes", sa.BigInteger(), nullable=True), + sa.Column("status", sa.String(length=50), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), + sa.Column( + "tenant_id", + sa.Integer(), + server_default=sa.text("current_setting('app.current_tenant')::int"), + nullable=False, + ), + sa.CheckConstraint( + "(edi_data IS NOT NULL OR storage_uri IS NOT NULL)", name="chk_edi_msg_data_or_uri" + ), + sa.ForeignKeyConstraint( + ["outbound_route_id"], + ["outbound_routes.id"], + ), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index( + op.f("ix_edi_messages_outbound_route_id"), + "edi_messages", + ["outbound_route_id"], + unique=False, + ) + op.create_index(op.f("ix_edi_messages_tenant_id"), "edi_messages", ["tenant_id"], unique=False) + op.create_index(op.f("ix_edi_messages_trace_id"), "edi_messages", ["trace_id"], unique=False) + op.create_index( + "ix_edi_msgs_sender_recv", + "edi_messages", + ["sender_id", "receiver_id", "created_at"], + unique=False, + ) # ### end Alembic commands ### def downgrade() -> None: """Downgrade schema.""" # ### commands auto generated by Alembic - please adjust! ### + op.drop_index("ix_edi_msgs_sender_recv", table_name="edi_messages") + op.drop_index(op.f("ix_edi_messages_trace_id"), table_name="edi_messages") + op.drop_index(op.f("ix_edi_messages_tenant_id"), table_name="edi_messages") + op.drop_index(op.f("ix_edi_messages_outbound_route_id"), table_name="edi_messages") + op.drop_table("edi_messages") + op.drop_index(op.f("ix_edi_json_transaction_type"), table_name="edi_json") + op.drop_index(op.f("ix_edi_json_trace_id"), table_name="edi_json") + op.drop_index(op.f("ix_edi_json_tenant_id"), table_name="edi_json") + op.drop_index("ix_edi_json_sender_recv", table_name="edi_json") + op.drop_index(op.f("ix_edi_json_outbound_route_id"), table_name="edi_json") + op.drop_index("ix_edi_json_business_metadata", table_name="edi_json", postgresql_using="gin") + op.drop_table("edi_json") op.drop_index( - "ix_outbound_routes_unique_active", + "ix_outbound_routes_unique_trading_partner_id", table_name="outbound_routes", postgresql_where=sa.text("active = true"), ) @@ -344,10 +513,6 @@ def downgrade() -> None: op.drop_index(op.f("ix_jobs_trace_id"), table_name="jobs") op.drop_index(op.f("ix_jobs_tenant_id"), table_name="jobs") op.drop_table("jobs") - op.drop_index("ix_edi_msgs_sender_recv", table_name="edi_messages") - op.drop_index(op.f("ix_edi_messages_trace_id"), table_name="edi_messages") - op.drop_index(op.f("ix_edi_messages_tenant_id"), table_name="edi_messages") - op.drop_table("edi_messages") op.drop_index(op.f("ix_audit_log_trace_id"), table_name="audit_log") op.drop_index(op.f("ix_audit_log_tenant_id"), table_name="audit_log") op.drop_table("audit_log") diff --git a/libs/database/src/database/models/common.py b/libs/database/src/database/models/common.py index 0e397f26..2dfec563 100644 --- a/libs/database/src/database/models/common.py +++ b/libs/database/src/database/models/common.py @@ -1,4 +1,4 @@ -from datetime import datetime +from datetime import UTC, datetime from typing import Any from uuid import UUID as PyUUID @@ -31,4 +31,20 @@ def status(cls) -> Mapped[str]: @declared_attr def created_at(cls) -> Mapped[datetime]: - return mapped_column(DateTime, default=datetime.utcnow) + return mapped_column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) + + +class TimestampMixin: + """Provides created_at and updated_at for configuration tables.""" + + @declared_attr + def created_at(cls) -> Mapped[datetime]: + return mapped_column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) + + @declared_attr + def updated_at(cls) -> Mapped[datetime]: + return mapped_column( + DateTime(timezone=True), + default=lambda: datetime.now(UTC), + onupdate=lambda: datetime.now(UTC), + ) diff --git a/libs/database/src/database/models/control_plane.py b/libs/database/src/database/models/control_plane.py index 6ea874af..76aa237d 100644 --- a/libs/database/src/database/models/control_plane.py +++ b/libs/database/src/database/models/control_plane.py @@ -1,5 +1,4 @@ from datetime import datetime -from typing import Any from uuid import UUID as PyUUID from sqlalchemy import ( @@ -10,14 +9,21 @@ Index, Integer, String, - Text, UniqueConstraint, ) -from sqlalchemy.dialects.postgresql import JSONB, UUID +from sqlalchemy.dialects.postgresql import UUID from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column from sqlalchemy.sql import func, text -from .common import OutboxMixin +from .common import OutboxMixin, TimestampMixin +from .replicated_mixins import ( + AS2PartnerMixin, + AS2PartnershipMixin, + InboundRouteMixin, + OutboundRouteMixin, + SFTPPartnerMixin, + WebhookMixin, +) class GlobalBase(DeclarativeBase): @@ -73,39 +79,41 @@ class TenantUser(GlobalBase): __table_args__ = (UniqueConstraint("tenant_id", "user_id", name="uq_tenant_user"),) -class AS2Partner(GlobalBase): +class ApiToken(GlobalBase, TimestampMixin): """ - Global AS2 Partners (both our local gateway config, and remote shared configs). + Platform-managed API keys for machine-to-machine (ERP → Platform) authentication. + Two-part credential: client_id (plaintext, visible) + client_secret (hashed, shown once). """ - __tablename__ = "as2_partners" + __tablename__ = "api_tokens" id: Mapped[PyUUID] = mapped_column( UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() ) - tenant_id: Mapped[int | None] = mapped_column( - Integer, ForeignKey("tenants.id", ondelete="CASCADE"), nullable=True - ) # Null if shared global - is_local: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) + tenant_id: Mapped[int] = mapped_column( + Integer, ForeignKey("tenants.id", ondelete="CASCADE"), nullable=False, index=True + ) name: Mapped[str] = mapped_column(String(255), nullable=False) - as2_id: Mapped[str] = mapped_column(String(255), nullable=False) - public_cert_pem: Mapped[str | None] = mapped_column( - Text, nullable=True - ) # Retained for legacy/external, but vault preferred - public_cert_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - private_key_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) + # client_id: stored in plaintext, used for fast indexed lookup and displayed in UI + client_id: Mapped[str] = mapped_column(String(64), nullable=False, unique=True) + # secret_hash: SHA-256 of the raw client_secret; raw value is never stored + secret_hash: Mapped[str] = mapped_column(String(64), nullable=False) + last_used_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + active: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True) - prev_public_cert_pem: Mapped[str | None] = mapped_column(Text, nullable=True) - prev_public_cert_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - prev_private_key_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - url: Mapped[str | None] = mapped_column(String(1024), nullable=True) +class AS2Partner(GlobalBase, AS2PartnerMixin, TimestampMixin): + """ + Global AS2 Partners (both our local gateway config, and remote shared configs). + """ - active: Mapped[bool] = mapped_column(Boolean, default=False) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column( - DateTime, default=datetime.utcnow, onupdate=datetime.utcnow - ) + __tablename__ = "as2_partners" + + tenant_id: Mapped[int | None] = mapped_column( + Integer, ForeignKey("tenants.id", ondelete="CASCADE"), nullable=True + ) # Null if shared global + is_local: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) __table_args__ = ( UniqueConstraint("tenant_id", "as2_id", name="uq_tenant_as2_id"), @@ -115,7 +123,7 @@ class AS2Partner(GlobalBase): ) -class AS2Partnership(GlobalBase): +class AS2Partnership(GlobalBase, AS2PartnershipMixin, TimestampMixin): """ OpenAS2 style Partnership (links Local Partner to Remote Partner) and stores all MDN and Encryption properties. @@ -123,13 +131,9 @@ class AS2Partnership(GlobalBase): __tablename__ = "as2_partnerships" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) tenant_id: Mapped[int | None] = mapped_column( Integer, ForeignKey("tenants.id", ondelete="CASCADE"), nullable=True ) - name: Mapped[str] = mapped_column(String(255), nullable=False) local_partner_id: Mapped[PyUUID] = mapped_column( UUID(as_uuid=True), ForeignKey("as2_partners.id", ondelete="CASCADE"), nullable=False @@ -138,24 +142,6 @@ class AS2Partnership(GlobalBase): UUID(as_uuid=True), ForeignKey("as2_partners.id", ondelete="CASCADE"), nullable=False ) - # Core AS2 Networking - credentials_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - - # Core AS2 Protocol Settings - mdn_type: Mapped[str] = mapped_column(String(50), nullable=False, default="SYNC") - mdn_url: Mapped[str | None] = mapped_column(String(1024), nullable=True) - encryption_algorithm: Mapped[str] = mapped_column(String(50), nullable=False, default="AES256") - signature_algorithm: Mapped[str] = mapped_column(String(50), nullable=False, default="SHA256") - - # Advanced OpenAS2 settings - advanced_flags: Mapped[dict[str, Any] | None] = mapped_column(JSONB, nullable=True) - - active: Mapped[bool] = mapped_column(Boolean, default=False) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column( - DateTime, default=datetime.utcnow, onupdate=datetime.utcnow - ) - __table_args__ = ( UniqueConstraint("local_partner_id", "remote_partner_id", name="uq_as2_partnership"), ) @@ -201,58 +187,29 @@ class SystemAuditLog(GlobalBase): # --------------------------------------------------------------------------- -class SFTPPartner(GlobalBase): +class SFTPPartner(GlobalBase, SFTPPartnerMixin, TimestampMixin): __tablename__ = "sftp_partners" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) tenant_id: Mapped[int] = mapped_column( Integer, ForeignKey("tenants.id", ondelete="CASCADE"), nullable=False, index=True ) - name: Mapped[str] = mapped_column(String(255), nullable=False) - host: Mapped[str] = mapped_column(String(1024), nullable=False) - port: Mapped[int] = mapped_column(Integer, default=22) - username: Mapped[str] = mapped_column(String(255), nullable=False) - host_key: Mapped[str | None] = mapped_column(Text, nullable=True) - inbound_remote_path: Mapped[str | None] = mapped_column(String(1024), nullable=True) - outbound_remote_path: Mapped[str | None] = mapped_column(String(1024), nullable=True) - password_encrypted: Mapped[str | None] = mapped_column(String(1024), nullable=True) - credentials_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - active: Mapped[bool] = mapped_column(Boolean, default=False) - - -class Webhook(GlobalBase): + + +class Webhook(GlobalBase, WebhookMixin, TimestampMixin): __tablename__ = "webhooks" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) tenant_id: Mapped[int] = mapped_column( Integer, ForeignKey("tenants.id", ondelete="CASCADE"), nullable=False, index=True ) - name: Mapped[str] = mapped_column(String(255), nullable=False) - url: Mapped[str] = mapped_column(String(1024), nullable=False) - auth_header_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - active: Mapped[bool] = mapped_column(Boolean, default=False) -class InboundRoute(GlobalBase): +class InboundRoute(GlobalBase, InboundRouteMixin, TimestampMixin): __tablename__ = "inbound_routes" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) tenant_id: Mapped[int] = mapped_column( Integer, ForeignKey("tenants.id", ondelete="CASCADE"), nullable=False, index=True ) - name: Mapped[str] = mapped_column(String(255), nullable=False) - isa_sender_id: Mapped[str] = mapped_column(String(255), nullable=False) - isa_receiver_id: Mapped[str] = mapped_column(String(255), nullable=False) - transaction_type: Mapped[str] = mapped_column(String(50), nullable=False) - processing_mode: Mapped[str] = mapped_column( - String(50), nullable=False, server_default="TRANSLATE" - ) + webhook_id: Mapped[PyUUID | None] = mapped_column( UUID(as_uuid=True), ForeignKey("webhooks.id"), nullable=True ) @@ -262,7 +219,6 @@ class InboundRoute(GlobalBase): sftp_partner_id: Mapped[PyUUID | None] = mapped_column( UUID(as_uuid=True), ForeignKey("sftp_partners.id"), nullable=True ) - active: Mapped[bool] = mapped_column(Boolean, default=False) __table_args__ = ( CheckConstraint( @@ -281,29 +237,19 @@ class InboundRoute(GlobalBase): ) -class OutboundRoute(GlobalBase): +class OutboundRoute(GlobalBase, OutboundRouteMixin, TimestampMixin): __tablename__ = "outbound_routes" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) tenant_id: Mapped[int] = mapped_column( Integer, ForeignKey("tenants.id", ondelete="CASCADE"), nullable=False, index=True ) - name: Mapped[str] = mapped_column(String(255), nullable=False) - isa_sender_id: Mapped[str] = mapped_column(String(255), nullable=False) - isa_receiver_id: Mapped[str] = mapped_column(String(255), nullable=False) - transaction_type: Mapped[str] = mapped_column(String(50), nullable=False) - processing_mode: Mapped[str] = mapped_column( - String(50), nullable=False, server_default="TRANSLATE" - ) + as2_partner_id: Mapped[PyUUID | None] = mapped_column( UUID(as_uuid=True), ForeignKey("as2_partners.id"), nullable=True ) sftp_partner_id: Mapped[PyUUID | None] = mapped_column( UUID(as_uuid=True), ForeignKey("sftp_partners.id"), nullable=True ) - active: Mapped[bool] = mapped_column(Boolean, default=False) __table_args__ = ( CheckConstraint( @@ -311,11 +257,9 @@ class OutboundRoute(GlobalBase): name="chk_outbound_routes_exactly_one_dest", ), Index( - "ix_outbound_routes_unique_active", + "ix_outbound_routes_unique_trading_partner_id", "tenant_id", - "isa_sender_id", - "isa_receiver_id", - "transaction_type", + "trading_partner_id", unique=True, postgresql_where=text("active = true"), ), diff --git a/libs/database/src/database/models/data_plane.py b/libs/database/src/database/models/data_plane.py index c6164cef..4ba95a4a 100644 --- a/libs/database/src/database/models/data_plane.py +++ b/libs/database/src/database/models/data_plane.py @@ -18,7 +18,15 @@ from sqlalchemy.sql import func, text from sqlalchemy.types import TypeDecorator -from .common import OutboxMixin +from .common import OutboxMixin, TimestampMixin +from .replicated_mixins import ( + AS2PartnerMixin, + AS2PartnershipMixin, + InboundRouteMixin, + OutboundRouteMixin, + SFTPPartnerMixin, + WebhookMixin, +) class SanitizedText(TypeDecorator): # type: ignore @@ -53,7 +61,12 @@ class TenantAwareMixin: @declared_attr def tenant_id(cls) -> Mapped[int]: - return mapped_column(Integer, nullable=False, index=True) + return mapped_column( + Integer, + server_default=text("current_setting('app.current_tenant')::int"), + nullable=False, + index=True, + ) # --------------------------------------------------------------------------- @@ -61,36 +74,13 @@ def tenant_id(cls) -> Mapped[int]: # --------------------------------------------------------------------------- -class AS2Partner(TenantBase, TenantAwareMixin): +class AS2Partner(TenantBase, TenantAwareMixin, AS2PartnerMixin, TimestampMixin): __tablename__ = "as2_partners" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) - is_local: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) - name: Mapped[str] = mapped_column(String(255), nullable=False) - as2_id: Mapped[str] = mapped_column(String(255), nullable=False) - public_cert_pem: Mapped[str | None] = mapped_column(Text, nullable=True) - public_cert_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - private_key_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - - prev_public_cert_pem: Mapped[str | None] = mapped_column(Text, nullable=True) - prev_public_cert_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - prev_private_key_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - - url: Mapped[str | None] = mapped_column(String(1024), nullable=True) - - active: Mapped[bool] = mapped_column(Boolean, default=False) - -class AS2Partnership(TenantBase, TenantAwareMixin): +class AS2Partnership(TenantBase, TenantAwareMixin, AS2PartnershipMixin, TimestampMixin): __tablename__ = "as2_partnerships" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) - name: Mapped[str] = mapped_column(String(255), nullable=False) - local_partner_id: Mapped[PyUUID] = mapped_column( UUID(as_uuid=True), ForeignKey("as2_partners.id", ondelete="CASCADE"), nullable=False ) @@ -98,66 +88,23 @@ class AS2Partnership(TenantBase, TenantAwareMixin): UUID(as_uuid=True), ForeignKey("as2_partners.id", ondelete="CASCADE"), nullable=False ) - credentials_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - - mdn_type: Mapped[str] = mapped_column(String(50), nullable=False, default="SYNC") - mdn_url: Mapped[str | None] = mapped_column(String(1024), nullable=True) - encryption_algorithm: Mapped[str] = mapped_column(String(50), nullable=False, default="AES256") - signature_algorithm: Mapped[str] = mapped_column(String(50), nullable=False, default="SHA256") - - advanced_flags: Mapped[dict[str, Any] | None] = mapped_column(JSONB, nullable=True) - - active: Mapped[bool] = mapped_column(Boolean, default=False) - # --------------------------------------------------------------------------- # Tenant Protocol & Routing Models # --------------------------------------------------------------------------- -class SFTPPartner(TenantBase, TenantAwareMixin): +class SFTPPartner(TenantBase, TenantAwareMixin, SFTPPartnerMixin, TimestampMixin): __tablename__ = "sftp_partners" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) - name: Mapped[str] = mapped_column(String(255), nullable=False) - host: Mapped[str] = mapped_column(String(1024), nullable=False) - port: Mapped[int] = mapped_column(Integer, default=22) - username: Mapped[str] = mapped_column(String(255), nullable=False) - host_key: Mapped[str | None] = mapped_column(Text, nullable=True) - inbound_remote_path: Mapped[str | None] = mapped_column(String(1024), nullable=True) - outbound_remote_path: Mapped[str | None] = mapped_column(String(1024), nullable=True) - password_encrypted: Mapped[str | None] = mapped_column(String(1024), nullable=True) - credentials_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - active: Mapped[bool] = mapped_column(Boolean, default=False) - - -class Webhook(TenantBase, TenantAwareMixin): - __tablename__ = "webhooks" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) - name: Mapped[str] = mapped_column(String(255), nullable=False) - url: Mapped[str] = mapped_column(String(1024), nullable=False) - auth_header_vault_ref: Mapped[str | None] = mapped_column(String(255), nullable=True) - active: Mapped[bool] = mapped_column(Boolean, default=False) +class Webhook(TenantBase, TenantAwareMixin, WebhookMixin, TimestampMixin): + __tablename__ = "webhooks" -class InboundRoute(TenantBase, TenantAwareMixin): +class InboundRoute(TenantBase, TenantAwareMixin, InboundRouteMixin, TimestampMixin): __tablename__ = "inbound_routes" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) - name: Mapped[str] = mapped_column(String(255), nullable=False) - isa_sender_id: Mapped[str] = mapped_column(String(255), nullable=False) - isa_receiver_id: Mapped[str] = mapped_column(String(255), nullable=False) - transaction_type: Mapped[str] = mapped_column(String(50), nullable=False) - processing_mode: Mapped[str] = mapped_column( - String(50), nullable=False, server_default="TRANSLATE" - ) webhook_id: Mapped[PyUUID | None] = mapped_column( UUID(as_uuid=True), ForeignKey("webhooks.id"), nullable=True ) @@ -167,7 +114,6 @@ class InboundRoute(TenantBase, TenantAwareMixin): sftp_partner_id: Mapped[PyUUID | None] = mapped_column( UUID(as_uuid=True), ForeignKey("sftp_partners.id"), nullable=True ) - active: Mapped[bool] = mapped_column(Boolean, default=False) __table_args__ = ( CheckConstraint( @@ -186,26 +132,15 @@ class InboundRoute(TenantBase, TenantAwareMixin): ) -class OutboundRoute(TenantBase, TenantAwareMixin): +class OutboundRoute(TenantBase, TenantAwareMixin, OutboundRouteMixin, TimestampMixin): __tablename__ = "outbound_routes" - id: Mapped[PyUUID] = mapped_column( - UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() - ) - name: Mapped[str] = mapped_column(String(255), nullable=False) - isa_sender_id: Mapped[str] = mapped_column(String(255), nullable=False) - isa_receiver_id: Mapped[str] = mapped_column(String(255), nullable=False) - transaction_type: Mapped[str] = mapped_column(String(50), nullable=False) - processing_mode: Mapped[str] = mapped_column( - String(50), nullable=False, server_default="TRANSLATE" - ) as2_partner_id: Mapped[PyUUID | None] = mapped_column( UUID(as_uuid=True), ForeignKey("as2_partners.id"), nullable=True ) sftp_partner_id: Mapped[PyUUID | None] = mapped_column( UUID(as_uuid=True), ForeignKey("sftp_partners.id"), nullable=True ) - active: Mapped[bool] = mapped_column(Boolean, default=False) __table_args__ = ( CheckConstraint( @@ -213,11 +148,9 @@ class OutboundRoute(TenantBase, TenantAwareMixin): name="chk_outbound_routes_exactly_one_dest", ), Index( - "ix_outbound_routes_unique_active", + "ix_outbound_routes_unique_trading_partner_id", "tenant_id", - "isa_sender_id", - "isa_receiver_id", - "transaction_type", + "trading_partner_id", unique=True, postgresql_where=text("active = true"), ), @@ -229,7 +162,7 @@ class OutboundRoute(TenantBase, TenantAwareMixin): # --------------------------------------------------------------------------- -class EdiMessage(TenantBase, TenantAwareMixin): +class EdiMessage(TenantBase, TenantAwareMixin, TimestampMixin): __tablename__ = "edi_messages" id: Mapped[PyUUID] = mapped_column( @@ -237,7 +170,7 @@ class EdiMessage(TenantBase, TenantAwareMixin): ) trace_id: Mapped[PyUUID] = mapped_column(UUID(as_uuid=True), nullable=False, index=True) direction: Mapped[str] = mapped_column(String(50), nullable=False) # INBOUND, OUTBOUND - connection_type: Mapped[str] = mapped_column(String(50), nullable=False) # AS2, SFTP, FTP + connection_type: Mapped[str | None] = mapped_column(String(50), nullable=True) # AS2, SFTP, FTP sender_id: Mapped[str | None] = mapped_column(String(255), nullable=True) receiver_id: Mapped[str | None] = mapped_column(String(255), nullable=True) @@ -253,22 +186,65 @@ class EdiMessage(TenantBase, TenantAwareMixin): status_message: Mapped[str | None] = mapped_column(Text, nullable=True) state: Mapped[str | None] = mapped_column(String(255), nullable=True) msg_headers: Mapped[str | None] = mapped_column(Text, nullable=True) + outbound_route_id: Mapped[PyUUID | None] = mapped_column( + UUID(as_uuid=True), ForeignKey("outbound_routes.id"), nullable=True, index=True + ) interchange_control_no: Mapped[str | None] = mapped_column(String(255), nullable=True) transaction_type: Mapped[str | None] = mapped_column(String(50), nullable=True) format_standard: Mapped[str | None] = mapped_column(String(50), nullable=True) - edi_data: Mapped[str] = mapped_column(SanitizedText, nullable=False) + edi_data: Mapped[str | None] = mapped_column(SanitizedText, nullable=True) + storage_uri: Mapped[str | None] = mapped_column(String(1024), nullable=True) file_size_bytes: Mapped[int | None] = mapped_column(BigInteger, nullable=True) status: Mapped[str] = mapped_column(String(50), nullable=False, default="RECEIVED") + __table_args__ = ( + Index("ix_edi_msgs_sender_recv", "sender_id", "receiver_id", "created_at"), + CheckConstraint( + "(edi_data IS NOT NULL OR storage_uri IS NOT NULL)", + name="chk_edi_msg_data_or_uri", + ), + ) + + +class EdiJson(TenantBase, TenantAwareMixin): + __tablename__ = "edi_json" + + id: Mapped[PyUUID] = mapped_column( + UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() + ) + trace_id: Mapped[PyUUID] = mapped_column(UUID(as_uuid=True), nullable=False, index=True) + direction: Mapped[str] = mapped_column(String(50), nullable=False) # INBOUND, OUTBOUND + + outbound_route_id: Mapped[PyUUID | None] = mapped_column( + UUID(as_uuid=True), ForeignKey("outbound_routes.id"), nullable=True, index=True + ) + transaction_type: Mapped[str | None] = mapped_column(String(50), nullable=True, index=True) + standard: Mapped[str | None] = mapped_column(String(50), nullable=True) + sender_id: Mapped[str | None] = mapped_column(String(255), nullable=True) + receiver_id: Mapped[str | None] = mapped_column(String(255), nullable=True) + + business_metadata: Mapped[dict[str, Any] | None] = mapped_column(JSONB, nullable=True) + payload: Mapped[dict[str, Any] | None] = mapped_column(JSONB, nullable=True) + storage_uri: Mapped[str | None] = mapped_column(String(1024), nullable=True) + + status: Mapped[str] = mapped_column(String(50), nullable=False, default="TRANSLATED") + created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) updated_at: Mapped[datetime] = mapped_column( DateTime, default=datetime.utcnow, onupdate=datetime.utcnow ) - __table_args__ = (Index("ix_edi_msgs_sender_recv", "sender_id", "receiver_id", "created_at"),) + __table_args__ = ( + Index("ix_edi_json_business_metadata", "business_metadata", postgresql_using="gin"), + Index("ix_edi_json_sender_recv", "sender_id", "receiver_id", "created_at"), + CheckConstraint( + "(payload IS NOT NULL OR storage_uri IS NOT NULL)", + name="chk_edi_json_data_or_uri", + ), + ) class ApiGateway(TenantBase, TenantAwareMixin): @@ -285,7 +261,8 @@ class ApiGateway(TenantBase, TenantAwareMixin): http_status_code: Mapped[int | None] = mapped_column(Integer, nullable=True) target_format: Mapped[str | None] = mapped_column(String(50), nullable=True) - request: Mapped[str] = mapped_column(String(1024), nullable=False) + payload: Mapped[dict[str, Any] | None] = mapped_column(JSONB, nullable=True) + storage_uri: Mapped[str | None] = mapped_column(String(1024), nullable=True) response: Mapped[str | None] = mapped_column(Text, nullable=True) headers: Mapped[dict[str, Any] | None] = mapped_column(JSONB, nullable=True) @@ -296,6 +273,13 @@ class ApiGateway(TenantBase, TenantAwareMixin): DateTime, default=datetime.utcnow, onupdate=datetime.utcnow ) + __table_args__ = ( + CheckConstraint( + "(payload IS NOT NULL OR storage_uri IS NOT NULL)", + name="chk_apigw_data_or_uri", + ), + ) + class Job(TenantBase, TenantAwareMixin): __tablename__ = "jobs" diff --git a/libs/database/src/database/models/replicated_mixins.py b/libs/database/src/database/models/replicated_mixins.py new file mode 100644 index 00000000..7a5613cd --- /dev/null +++ b/libs/database/src/database/models/replicated_mixins.py @@ -0,0 +1,284 @@ +from typing import Any +from uuid import UUID as PyUUID + +from sqlalchemy import ( + Boolean, + Integer, + String, + Text, +) +from sqlalchemy.dialects.postgresql import JSONB, UUID +from sqlalchemy.orm import Mapped, declared_attr, mapped_column +from sqlalchemy.sql import func, text + + +class AS2PartnerMixin: + """Shared columns for AS2Partner across Global and Tenant schemas.""" + + @declared_attr + def id(cls) -> Mapped[PyUUID]: + return mapped_column( + UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() + ) + + @declared_attr + def name(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def as2_id(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def public_cert_pem(cls) -> Mapped[str | None]: + return mapped_column(Text, nullable=True) + + @declared_attr + def public_cert_vault_ref(cls) -> Mapped[str | None]: + return mapped_column(String(255), nullable=True) + + @declared_attr + def private_key_vault_ref(cls) -> Mapped[str | None]: + return mapped_column(String(255), nullable=True) + + @declared_attr + def prev_public_cert_pem(cls) -> Mapped[str | None]: + return mapped_column(Text, nullable=True) + + @declared_attr + def prev_public_cert_vault_ref(cls) -> Mapped[str | None]: + return mapped_column(String(255), nullable=True) + + @declared_attr + def prev_private_key_vault_ref(cls) -> Mapped[str | None]: + return mapped_column(String(255), nullable=True) + + @declared_attr + def url(cls) -> Mapped[str | None]: + return mapped_column(String(1024), nullable=True) + + @declared_attr + def active(cls) -> Mapped[bool]: + return mapped_column(Boolean, default=False, server_default=text("false")) + + +class AS2PartnershipMixin: + """Shared columns for AS2Partnership across Global and Tenant schemas.""" + + @declared_attr + def id(cls) -> Mapped[PyUUID]: + return mapped_column( + UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() + ) + + @declared_attr + def name(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def credentials_vault_ref(cls) -> Mapped[str | None]: + return mapped_column(String(255), nullable=True) + + @declared_attr + def mdn_type(cls) -> Mapped[str]: + return mapped_column(String(50), nullable=False, default="SYNC") + + @declared_attr + def mdn_url(cls) -> Mapped[str | None]: + return mapped_column(String(1024), nullable=True) + + @declared_attr + def encryption_algorithm(cls) -> Mapped[str]: + return mapped_column(String(50), nullable=False, default="AES256") + + @declared_attr + def signature_algorithm(cls) -> Mapped[str]: + return mapped_column(String(50), nullable=False, default="SHA256") + + @declared_attr + def advanced_flags(cls) -> Mapped[dict[str, Any] | None]: + return mapped_column(JSONB, nullable=True) + + @declared_attr + def active(cls) -> Mapped[bool]: + return mapped_column(Boolean, default=False, server_default=text("false")) + + +class SFTPPartnerMixin: + """Shared columns for SFTPPartner across Global and Tenant schemas.""" + + @declared_attr + def id(cls) -> Mapped[PyUUID]: + return mapped_column( + UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() + ) + + @declared_attr + def name(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def host(cls) -> Mapped[str]: + return mapped_column(String(1024), nullable=False) + + @declared_attr + def port(cls) -> Mapped[int]: + return mapped_column(Integer, default=22) + + @declared_attr + def username(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def host_key(cls) -> Mapped[str | None]: + return mapped_column(Text, nullable=True) + + @declared_attr + def inbound_remote_path(cls) -> Mapped[str | None]: + return mapped_column(String(1024), nullable=True) + + @declared_attr + def outbound_remote_path(cls) -> Mapped[str | None]: + return mapped_column(String(1024), nullable=True) + + @declared_attr + def password_encrypted(cls) -> Mapped[str | None]: + return mapped_column(String(1024), nullable=True) + + @declared_attr + def credentials_vault_ref(cls) -> Mapped[str | None]: + return mapped_column(String(255), nullable=True) + + @declared_attr + def active(cls) -> Mapped[bool]: + return mapped_column(Boolean, default=False, server_default=text("false")) + + +class WebhookMixin: + """Shared columns for Webhook across Global and Tenant schemas.""" + + @declared_attr + def id(cls) -> Mapped[PyUUID]: + return mapped_column( + UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() + ) + + @declared_attr + def name(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def url(cls) -> Mapped[str]: + return mapped_column(String(1024), nullable=False) + + @declared_attr + def auth_header_vault_ref(cls) -> Mapped[str | None]: + return mapped_column(String(255), nullable=True) + + @declared_attr + def active(cls) -> Mapped[bool]: + return mapped_column(Boolean, default=False, server_default=text("false")) + + +class InboundRouteMixin: + """Shared columns for InboundRoute across Global and Tenant schemas.""" + + @declared_attr + def id(cls) -> Mapped[PyUUID]: + return mapped_column( + UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() + ) + + @declared_attr + def name(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def isa_sender_id(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def isa_receiver_id(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def gs_sender_id(cls) -> Mapped[str | None]: + return mapped_column(String(255), nullable=True) + + @declared_attr + def gs_receiver_id(cls) -> Mapped[str | None]: + return mapped_column(String(255), nullable=True) + + @declared_attr + def transaction_type(cls) -> Mapped[str]: + return mapped_column(String(50), nullable=False) + + @declared_attr + def processing_mode(cls) -> Mapped[str]: + return mapped_column(String(50), nullable=False, server_default="TRANSLATE") + + @declared_attr + def active(cls) -> Mapped[bool]: + return mapped_column(Boolean, default=False, server_default=text("false")) + + +class OutboundRouteMixin: + """Shared columns for OutboundRoute across Global and Tenant schemas.""" + + @declared_attr + def id(cls) -> Mapped[PyUUID]: + return mapped_column( + UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid() + ) + + @declared_attr + def trading_partner_id(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def name(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def isa_sender_id(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def isa_sender_qualifier(cls) -> Mapped[str | None]: + return mapped_column(String(2), nullable=True) + + @declared_attr + def isa_receiver_id(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def isa_receiver_qualifier(cls) -> Mapped[str | None]: + return mapped_column(String(2), nullable=True) + + @declared_attr + def gs_sender_id(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def gs_receiver_id(cls) -> Mapped[str]: + return mapped_column(String(255), nullable=False) + + @declared_attr + def transaction_type(cls) -> Mapped[str]: + return mapped_column(String(50), nullable=False) + + @declared_attr + def default_standard(cls) -> Mapped[str]: + return mapped_column(String(50), default="x12", server_default="x12") + + @declared_attr + def default_version(cls) -> Mapped[str]: + return mapped_column(String(50), default="004010", server_default="004010") + + @declared_attr + def processing_mode(cls) -> Mapped[str]: + return mapped_column(String(50), nullable=False, server_default="TRANSLATE") + + @declared_attr + def active(cls) -> Mapped[bool]: + return mapped_column(Boolean, default=False, server_default=text("false")) diff --git a/libs/database/src/database/repository.py b/libs/database/src/database/repository.py index e92d9bb6..f41222be 100644 --- a/libs/database/src/database/repository.py +++ b/libs/database/src/database/repository.py @@ -60,6 +60,33 @@ async def get_partnership_by_as2_ids( return row[0], row[1], row[2] +class InboundRouteRepository: + def __init__(self, session: AsyncSession): + self.session = session + + async def get_inbound_route( + self, + isa_sender_id: str, + isa_receiver_id: str, + tenant_id: int, + transaction_type: str | None = None, + ) -> Any | None: + from .models.control_plane import InboundRoute + + conditions = [ + InboundRoute.isa_sender_id == isa_sender_id, + InboundRoute.isa_receiver_id == isa_receiver_id, + InboundRoute.active.is_(True), + InboundRoute.tenant_id == tenant_id, + ] + if transaction_type: + conditions.append(InboundRoute.transaction_type == transaction_type) + + stmt = select(InboundRoute).where(*conditions).order_by(InboundRoute.id).limit(1) + result = await self.session.execute(stmt) + return result.scalar_one_or_none() + + class EdiMessageRepository: def __init__(self, session: AsyncSession): self.session = session diff --git a/libs/database/tests/test_database_repository.py b/libs/database/tests/test_database_repository.py index 0a65820c..13922e81 100644 --- a/libs/database/tests/test_database_repository.py +++ b/libs/database/tests/test_database_repository.py @@ -59,3 +59,67 @@ async def test_edi_message_repository_save_message() -> None: assert result.message_id == "msg-123" mock_session.add.assert_called_once() mock_session.flush.assert_awaited_once() + + +async def test_inbound_route_repository_get_inbound_route() -> None: + from database.models.control_plane import InboundRoute + from database.repository import InboundRouteRepository + + mock_session = AsyncMock() + mock_result = MagicMock() + + route = InboundRoute( + id=uuid.uuid4(), + tenant_id=1, + isa_sender_id="SENDER", + isa_receiver_id="RECEIVER", + transaction_type="850", + active=True, + ) + mock_result.scalar_one_or_none.return_value = route + mock_session.execute.return_value = mock_result + + repo = InboundRouteRepository(mock_session) + result = await repo.get_inbound_route("SENDER", "RECEIVER", tenant_id=1, transaction_type="850") + + assert result == route + mock_session.execute.assert_called_once() + + +async def test_inbound_route_repository_get_inbound_route_no_match() -> None: + from database.repository import InboundRouteRepository + + mock_session = AsyncMock() + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = None + mock_session.execute.return_value = mock_result + + repo = InboundRouteRepository(mock_session) + result = await repo.get_inbound_route("SENDER", "RECEIVER", tenant_id=1, transaction_type="850") + + assert result is None + mock_session.execute.assert_called_once() + + +async def test_inbound_route_repository_get_inbound_route_no_transaction_type() -> None: + from database.models.control_plane import InboundRoute + from database.repository import InboundRouteRepository + + mock_session = AsyncMock() + mock_result = MagicMock() + route = InboundRoute( + id=uuid.uuid4(), + tenant_id=1, + isa_sender_id="SENDER", + isa_receiver_id="RECEIVER", + transaction_type=None, + active=True, + ) + mock_result.scalar_one_or_none.return_value = route + mock_session.execute.return_value = mock_result + + repo = InboundRouteRepository(mock_session) + result = await repo.get_inbound_route("SENDER", "RECEIVER", tenant_id=1, transaction_type=None) + + assert result == route + mock_session.execute.assert_called_once() diff --git a/libs/domain/README.md b/libs/domain/README.md new file mode 100644 index 00000000..795804e8 --- /dev/null +++ b/libs/domain/README.md @@ -0,0 +1,3 @@ +# domain + +Shared domain models and events for the EDI platform. diff --git a/libs/domain/pyproject.toml b/libs/domain/pyproject.toml new file mode 100644 index 00000000..c5640044 --- /dev/null +++ b/libs/domain/pyproject.toml @@ -0,0 +1,11 @@ +[project] +name = "domain" +version = "0.1.0" +description = "Core domain models and events" +readme = "README.md" +requires-python = ">=3.11" +dependencies = [] + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" diff --git a/libs/domain/src/domain/__init__.py b/libs/domain/src/domain/__init__.py new file mode 100644 index 00000000..b0a0d399 --- /dev/null +++ b/libs/domain/src/domain/__init__.py @@ -0,0 +1,3 @@ +from .events import ProvisioningEventType + +__all__ = ["ProvisioningEventType"] diff --git a/libs/domain/src/domain/events.py b/libs/domain/src/domain/events.py new file mode 100644 index 00000000..99fc155e --- /dev/null +++ b/libs/domain/src/domain/events.py @@ -0,0 +1,29 @@ +from enum import StrEnum + + +class ProvisioningEventType(StrEnum): + AS2_PARTNER_CREATED = "AS2_PARTNER_CREATED" + AS2_PARTNER_UPDATED = "AS2_PARTNER_UPDATED" + AS2_PARTNER_DELETED = "AS2_PARTNER_DELETED" + AS2_PARTNERSHIP_CREATED = "AS2_PARTNERSHIP_CREATED" + AS2_PARTNERSHIP_UPDATED = "AS2_PARTNERSHIP_UPDATED" + AS2_PARTNERSHIP_DELETED = "AS2_PARTNERSHIP_DELETED" + SFTP_PARTNER_CREATED = "SFTP_PARTNER_CREATED" + SFTP_PARTNER_UPDATED = "SFTP_PARTNER_UPDATED" + SFTP_PARTNER_DELETED = "SFTP_PARTNER_DELETED" + WEBHOOK_CREATED = "WEBHOOK_CREATED" + WEBHOOK_UPDATED = "WEBHOOK_UPDATED" + WEBHOOK_DELETED = "WEBHOOK_DELETED" + INBOUND_ROUTE_CREATED = "INBOUND_ROUTE_CREATED" + INBOUND_ROUTE_UPDATED = "INBOUND_ROUTE_UPDATED" + INBOUND_ROUTE_DELETED = "INBOUND_ROUTE_DELETED" + OUTBOUND_ROUTE_CREATED = "OUTBOUND_ROUTE_CREATED" + OUTBOUND_ROUTE_UPDATED = "OUTBOUND_ROUTE_UPDATED" + OUTBOUND_ROUTE_DELETED = "OUTBOUND_ROUTE_DELETED" + + +class MessageQueueName(StrEnum): + TRANSLATE = "TranslateQueue" + DELIVER = "DeliverQueue" + PROVISIONING = "ProvisioningQueue" + CDC_DLQ = "CdcDlqQueue" diff --git a/libs/identity/docker/docker-compose.yml b/libs/identity/docker/docker-compose.yml index 8cabd4f2..4682a024 100644 --- a/libs/identity/docker/docker-compose.yml +++ b/libs/identity/docker/docker-compose.yml @@ -35,13 +35,16 @@ services: ZITADEL_DATABASE_POSTGRES_ADMIN_PASSWORD: ${POSTGRES_PASSWORD:?POSTGRES_PASSWORD must be set} ZITADEL_DATABASE_POSTGRES_ADMIN_SSL_MODE: disable ZITADEL_DATABASE_POSTGRES_USER_SSL_MODE: disable - ZITADEL_MACHINEKEY: "${ZITADEL_MACHINEKEY:?ZITADEL_MACHINEKEY must be set (32 bytes)}" - ZITADEL_FIRSTINSTANCE_ORG_MACHINE_MACHINE_USERNAME: setupmachine - ZITADEL_FIRSTINSTANCE_ORG_MACHINE_MACHINE_NAME: Setup Machine + ZITADEL_FIRSTINSTANCE_ORG_MACHINE_MACHINE_USERNAME: terraform + ZITADEL_FIRSTINSTANCE_ORG_MACHINE_MACHINE_NAME: Terraform Provisioner ZITADEL_FIRSTINSTANCE_ORG_HUMAN_USERNAME: zitadel-admin ZITADEL_FIRSTINSTANCE_ORG_HUMAN_PASSWORD: ${ZITADEL_FIRSTINSTANCE_ORG_HUMAN_PASSWORD:?ZITADEL_FIRSTINSTANCE_ORG_HUMAN_PASSWORD must be set} + ZITADEL_FIRSTINSTANCE_MACHINEKEYPATH: /machinekey + ZITADEL_FIRSTINSTANCE_ORG_MACHINE_MACHINEKEY_TYPE: JSON ports: - "127.0.0.1:8080:8080" + volumes: + - zitadel_machinekey:/machinekey depends_on: zitadel-postgres: condition: service_healthy @@ -49,3 +52,5 @@ services: volumes: zitadel_data: driver: local + zitadel_machinekey: + driver: local diff --git a/libs/identity/src/identity/dependencies.py b/libs/identity/src/identity/dependencies.py index ae673af0..7c6360fb 100644 --- a/libs/identity/src/identity/dependencies.py +++ b/libs/identity/src/identity/dependencies.py @@ -124,12 +124,12 @@ async def get_current_tenant_id( """ Resolves the external user email from the JWT to our internal global DB tenant_id. """ - email = token_payload.get("email") + email = token_payload.get("email") or token_payload.get("preferred_username") name = token_payload.get("name") if not email: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail="Token does not contain an email claim", + detail="Token does not contain an email or preferred_username claim", ) try: @@ -146,35 +146,35 @@ async def get_current_tenant_id( ) from e -async def get_tenant_session( +async def get_tenant_session_for_id( request: Request, - tenant_id: int = Depends(get_current_tenant_id), - global_session: AsyncSession = Depends(get_global_session), + tenant_id: int, + global_session: AsyncSession, ) -> AsyncGenerator[AsyncSession, None]: """ - Yields an AsyncSession dynamically bound to the correct database shard, - with PostgreSQL Row-Level Security (RLS) automatically applied. + Yields an AsyncSession dynamically bound to the correct database shard for a given tenant_id. """ db_router = getattr(request.app.state, "db_router", None) if not db_router: raise RuntimeError("DatabaseRouter not initialized in app state") - # 1. Fetch routing info from Global DB using the shared global_session stmt = ( select(Tenant, DatabaseShard) .join(DatabaseShard, Tenant.shard_id == DatabaseShard.id) .where(Tenant.id == tenant_id) ) result = await global_session.execute(stmt) - tenant, shard = result.one() + row = result.one_or_none() + if not row: + raise RuntimeError(f"Tenant {tenant_id} not found in database") + + tenant, shard = row shard_key = shard.name shard_url = shard.dsn - # 2. Yield the RLS-secured tenant session async_gen_tenant = db_router.get_tenant_session(tenant_id, shard_key, shard_url) tenant_session: AsyncSession = await async_gen_tenant.__anext__() - # 3. Set the tenant context variable for repositories that rely on it from identity.tenant_context import _tenant_id token = _tenant_id.set(tenant_id) @@ -185,3 +185,16 @@ async def get_tenant_session( _tenant_id.reset(token) with contextlib.suppress(StopAsyncIteration): await async_gen_tenant.__anext__() + + +async def get_tenant_session( + request: Request, + tenant_id: int = Depends(get_current_tenant_id), + global_session: AsyncSession = Depends(get_global_session), +) -> AsyncGenerator[AsyncSession, None]: + """ + Yields an AsyncSession dynamically bound to the correct database shard, + using Zitadel JWT authentication to resolve the tenant_id. + """ + async for session in get_tenant_session_for_id(request, tenant_id, global_session): + yield session diff --git a/libs/identity/tests/test_dependencies.py b/libs/identity/tests/test_dependencies.py index 08f9a6bd..d4952de8 100644 --- a/libs/identity/tests/test_dependencies.py +++ b/libs/identity/tests/test_dependencies.py @@ -65,7 +65,7 @@ async def test_get_current_tenant_id_missing_email() -> None: with pytest.raises(HTTPException) as exc_info: await get_current_tenant_id({"sub": "123"}, mock_use_case) assert exc_info.value.status_code == 403 - assert "email claim" in str(exc_info.value.detail) + assert "email or preferred_username claim" in str(exc_info.value.detail) @pytest.mark.asyncio @@ -105,7 +105,7 @@ async def test_get_tenant_session() -> None: mock_shard.name = "shard_1" mock_shard.dsn = "postgresql+asyncpg://edi:edi_password@localhost:5433/edi_shard_1" mock_result = MagicMock() - mock_result.one.return_value = (mock_tenant, mock_shard) + mock_result.one_or_none.return_value = (mock_tenant, mock_shard) mock_global_session.execute.return_value = mock_result async def mock_global_session_gen() -> AsyncGenerator[AsyncMock, None]: diff --git a/libs/pipeline/pyproject.toml b/libs/pipeline/pyproject.toml index e43e00dd..dff0125b 100644 --- a/libs/pipeline/pyproject.toml +++ b/libs/pipeline/pyproject.toml @@ -8,6 +8,7 @@ dependencies = [ "pydantic>=2.0.0", "aioboto3>=13.1.1", "httpx>=0.27.0", + "jsonpath-ng>=1.6.1", # Local dependencies "as2-core", "security", diff --git a/libs/pipeline/src/pipeline/adapters/http.py b/libs/pipeline/src/pipeline/adapters/http.py index db16f105..b1fcb785 100644 --- a/libs/pipeline/src/pipeline/adapters/http.py +++ b/libs/pipeline/src/pipeline/adapters/http.py @@ -27,9 +27,17 @@ async def deliver(self, url: str, payload: bytes, auth_token: str | None = None) if not parsed.hostname: raise ValueError("Invalid URL") - ip = await anyio.to_thread.run_sync(socket.gethostbyname, parsed.hostname) + addr_info = await anyio.to_thread.run_sync(socket.getaddrinfo, parsed.hostname, None) + ip = addr_info[0][4][0] ip_obj = ipaddress.ip_address(ip) - if ip_obj.is_private or ip_obj.is_loopback or ip_obj.is_unspecified: + if ( + ip_obj.is_private + or ip_obj.is_loopback + or ip_obj.is_unspecified + or ip_obj.is_link_local + or ip_obj.is_reserved + or ip_obj.is_multicast + ): raise ValueError("SSRF check failed: internal IP") port_str = f":{parsed.port}" if parsed.port else "" diff --git a/libs/pipeline/src/pipeline/adapters/repository.py b/libs/pipeline/src/pipeline/adapters/repository.py index 49746282..3aa33991 100644 --- a/libs/pipeline/src/pipeline/adapters/repository.py +++ b/libs/pipeline/src/pipeline/adapters/repository.py @@ -1,6 +1,7 @@ import uuid from typing import Any +from config.settings import AppSettings from database.encryption import db_encryption from database.models import ApiGateway, EdiMessage from database.models import TenantOutbox as Outbox @@ -13,6 +14,7 @@ Webhook, ) from pipeline.ports.repository import RepositoryPort +from pipeline.ports.storage import StoragePort from sqlalchemy import select, update from sqlalchemy.dialects.postgresql import insert from sqlalchemy.ext.asyncio import AsyncSession @@ -24,8 +26,10 @@ class SqlAlchemyRepositoryAdapter(RepositoryPort): Operates on the Tenant Data Plane models. """ - def __init__(self, session: AsyncSession): + def __init__(self, session: AsyncSession, settings: AppSettings, storage: StoragePort): self.session = session + self.settings = settings + self.storage = storage async def get_edi_message(self, trace_id: str) -> dict[str, Any] | None: result = await self.session.execute( @@ -34,16 +38,24 @@ async def get_edi_message(self, trace_id: str) -> dict[str, Any] | None: record = result.scalar_one_or_none() if not record: return None + edi_data = record.edi_data + if record.storage_uri: + raw_bytes = await self.storage.download(record.storage_uri) + edi_data = raw_bytes.decode("utf-8") + return { "trace_id": str(record.trace_id), "tenant_id": record.tenant_id, - "edi_data": record.edi_data, + "edi_data": edi_data, "format_standard": record.format_standard, "transaction_type": record.transaction_type, "sender_id": record.sender_id, "receiver_id": record.receiver_id, "direction": record.direction, "status": record.status, + "outbound_route_id": str(record.outbound_route_id) + if record.outbound_route_id + else None, } async def update_edi_message_status(self, trace_id: str, status: str) -> None: @@ -54,18 +66,138 @@ async def update_edi_message_status(self, trace_id: str, status: str) -> None: ) await self.session.flush() - async def save_api_payload( - self, trace_id: str, direction: str, s3_uri: str, status: str + async def save_edi_message( + self, + trace_id: str, + direction: str, + edi_data: str, + format_standard: str, + transaction_type: str, + status: str, + connection_type: str | None = "UNKNOWN", + sender_id: str | None = None, + receiver_id: str | None = None, + outbound_route_id: str | None = None, + tenant_id: int | None = None, ) -> None: - record = ApiGateway( - trace_id=uuid.UUID(trace_id), - direction=direction, - request=s3_uri, - status=status, - ) + storage_uri = None + data_to_store = edi_data + if self.settings.storage_backend == "s3": + storage_uri = await self.storage.upload( + payload=edi_data.encode("utf-8"), + key_prefix=f"edi_messages/{trace_id}", + file_name="payload.edi", + ) + data_to_store = "" + + record_kwargs = { + "trace_id": uuid.UUID(trace_id), + "direction": direction, + "connection_type": connection_type, + "edi_data": data_to_store, + "format_standard": format_standard, + "transaction_type": transaction_type, + "sender_id": sender_id, + "receiver_id": receiver_id, + "storage_uri": storage_uri, + "status": status, + "outbound_route_id": uuid.UUID(outbound_route_id) if outbound_route_id else None, + } + if tenant_id is not None: + record_kwargs["tenant_id"] = tenant_id + + record = EdiMessage(**record_kwargs) self.session.add(record) await self.session.flush() + async def save_edi_json( + self, + trace_id: str, + direction: str, + partnership_id: str | None, + transaction_type: str | None, + standard: str | None, + sender_id: str | None, + receiver_id: str | None, + business_metadata: dict[str, Any], + payload: dict[str, Any], + status: str, + tenant_id: int | None = None, + ) -> str: + import json + + from database.models.data_plane import EdiJson + + payload_dict = payload + storage_uri = None + if self.settings.storage_backend == "s3": + storage_uri = await self.storage.upload( + payload=json.dumps(payload).encode("utf-8"), + key_prefix=f"edi_json/{trace_id}", + file_name="payload.json", + ) + payload_dict = {} + + record_kwargs = { + "trace_id": uuid.UUID(trace_id), + "direction": direction, + "outbound_route_id": uuid.UUID(partnership_id) if partnership_id else None, + "transaction_type": transaction_type, + "standard": standard, + "sender_id": sender_id, + "receiver_id": receiver_id, + "business_metadata": business_metadata, + "payload": payload_dict, + "storage_uri": storage_uri, + "status": status, + } + if tenant_id is not None: + record_kwargs["tenant_id"] = tenant_id + + record = EdiJson(**record_kwargs) + self.session.add(record) + await self.session.flush() + return str(record.id) + + async def get_edi_json(self, trace_id: str) -> dict[str, Any] | None: + from database.models.data_plane import EdiJson + + result = await self.session.execute( + select(EdiJson).where(EdiJson.trace_id == uuid.UUID(trace_id)) + ) + record = result.scalar_one_or_none() + if not record: + return None + + payload = record.payload + if record.storage_uri: + import json + + raw_bytes = await self.storage.download(record.storage_uri) + payload = json.loads(raw_bytes.decode("utf-8")) + + return { + "trace_id": str(record.trace_id), + "payload": payload, + "transaction_type": record.transaction_type, + "standard": record.standard, + "direction": record.direction, + "status": record.status, + "sender_id": record.sender_id, + "receiver_id": record.receiver_id, + "outbound_route_id": str(record.outbound_route_id) + if record.outbound_route_id + else None, + } + + async def update_edi_json_status(self, trace_id: str, status: str) -> None: + from database.models.data_plane import EdiJson + + await self.session.execute( + update(EdiJson).where(EdiJson.trace_id == uuid.UUID(trace_id)).values(status=status) + ) + await self.session.flush() + async def publish_outbox_event( self, idempotency_key: str, event_type: str, payload: dict[str, Any] ) -> None: @@ -96,6 +228,23 @@ async def claim_edi_message(self, trace_id: str) -> bool: await self.session.flush() return result.scalar_one_or_none() is not None + async def save_api_payload( + self, trace_id: str, direction: str, payload: dict[str, Any], status: str + ) -> None: + import uuid + + from database.models.data_plane import ApiGateway + + record = ApiGateway( + trace_id=uuid.UUID(trace_id), + direction=direction, + payload=payload, + status=status, + http_status_code=202, + ) + self.session.add(record) + await self.session.flush() + async def get_api_payload(self, trace_id: str) -> dict[str, Any] | None: result = await self.session.execute( select(ApiGateway).where(ApiGateway.trace_id == uuid.UUID(trace_id)) @@ -103,9 +252,16 @@ async def get_api_payload(self, trace_id: str) -> dict[str, Any] | None: record = result.scalar_one_or_none() if not record: return None + payload = record.payload + if record.storage_uri: + import json + + raw_bytes = await self.storage.download(record.storage_uri) + payload = json.loads(raw_bytes.decode("utf-8")) + return { "trace_id": str(record.trace_id), - "request": record.request, + "payload": payload, "status": record.status, "direction": record.direction, } @@ -167,6 +323,40 @@ async def get_route( else None, } + async def get_outbound_route(self, route_id: str) -> dict[str, Any] | None: + stmt = select(OutboundRoute).where( + OutboundRoute.id == uuid.UUID(route_id), + OutboundRoute.active.is_(True), + ) + result = await self.session.execute(stmt) + record = result.scalar_one_or_none() + if not record: + return None + + connection_type = "UNKNOWN" + if record.as2_partner_id: + connection_type = "AS2" + elif record.sftp_partner_id: + connection_type = "SFTP" + + return { + "route_id": str(record.id), + "trading_partner_id": record.trading_partner_id, + "isa_sender_id": record.isa_sender_id, + "isa_sender_qualifier": record.isa_sender_qualifier, + "isa_receiver_id": record.isa_receiver_id, + "isa_receiver_qualifier": record.isa_receiver_qualifier, + "gs_sender_id": record.gs_sender_id, + "gs_receiver_id": record.gs_receiver_id, + "transaction_type": record.transaction_type, + "default_standard": record.default_standard, + "default_version": record.default_version, + "processing_mode": record.processing_mode, + "as2_partner_id": str(record.as2_partner_id) if record.as2_partner_id else None, + "sftp_partner_id": str(record.sftp_partner_id) if record.sftp_partner_id else None, + "connection_type": connection_type, + } + async def get_sftp_partner(self, partner_id: str) -> dict[str, Any] | None: result = await self.session.execute( select(SFTPPartner).where( @@ -223,8 +413,7 @@ async def get_as2_partner(self, partner_id: str) -> dict[str, Any] | None: "as2_id": partner.as2_id, "public_cert_pem": partner.public_cert_pem, "public_cert_vault_ref": partner.public_cert_vault_ref, - "local_url": partnership.local_url, - "remote_url": partnership.remote_url, + "remote_url": partner.url, "local_partner_id": str(partnership.local_partner_id), "credentials_vault_ref": partnership.credentials_vault_ref, "encryption_algorithm": partnership.encryption_algorithm, diff --git a/libs/pipeline/src/pipeline/adapters/transformer.py b/libs/pipeline/src/pipeline/adapters/transformer.py index 64e31f28..6c4743de 100644 --- a/libs/pipeline/src/pipeline/adapters/transformer.py +++ b/libs/pipeline/src/pipeline/adapters/transformer.py @@ -25,9 +25,43 @@ async def translate_edi_to_json( return {} async def translate_json_to_edi( - self, payload: dict[str, Any], standard: str, transaction_type: str + self, + payload: dict[str, Any] | list[dict[str, Any]], + standard: str, + transaction_type: str, + route_config: dict[str, Any], ) -> bytes: """ Translates JSON to EDI using the wrapped BOTS facade. """ - raise NotImplementedError("JSON to EDI translation is not yet supported via BOTS.") + import asyncio + + from transformer.domain.exceptions import TranslationError + + if isinstance(payload, dict) and ( + "interchange_ISA" in payload or "interchange_UNB" in payload + ): + ast_dict: dict[str, Any] = payload + else: + if standard.lower() == "x12": + from transformer.domain.envelope.x12 import X12EnvelopeBuilder + + ast_dict = X12EnvelopeBuilder.build(route_config, payload) + elif standard.lower() == "edifact": + from transformer.domain.envelope.edifact import EdifactEnvelopeBuilder + + ast_dict = EdifactEnvelopeBuilder.build(route_config, payload) + else: + raise TranslationError( + message=f"Unsupported standard for envelope building: {standard}" + ) + + edi_str, errors = await asyncio.to_thread( + self._adapter.serialize_to_edi, ast_dict, standard=standard.lower() + ) + + fatal_errors = [e for e in errors if not e.startswith("[W")] + if fatal_errors: + raise TranslationError(message="\n".join(fatal_errors), errors=fatal_errors) + + return edi_str.encode("utf-8") diff --git a/libs/pipeline/src/pipeline/core/deliver.py b/libs/pipeline/src/pipeline/core/deliver.py index b1861769..f70c4a07 100644 --- a/libs/pipeline/src/pipeline/core/deliver.py +++ b/libs/pipeline/src/pipeline/core/deliver.py @@ -19,7 +19,6 @@ from pipeline.ports.http import HttpDeliveryPort from pipeline.ports.repository import RepositoryPort from pipeline.ports.sftp import SftpDeliveryPort -from pipeline.ports.storage import StoragePort from pipeline.ports.vault import VaultPort logger = logging.getLogger(__name__) @@ -43,14 +42,12 @@ class DeliveryService: def __init__( self, - storage: StoragePort, repository: RepositoryPort, http_delivery: HttpDeliveryPort, sftp_delivery: SftpDeliveryPort, as2_delivery: AS2DeliveryPort, vault: VaultPort | None = None, ) -> None: - self.storage = storage self.repository = repository self.http_delivery = http_delivery self.sftp_delivery = sftp_delivery @@ -70,17 +67,29 @@ async def deliver(self, trace_id: str) -> None: raise ValueError(f"No EDI Message found for trace_id={trace_id}") direction = edi_msg["direction"] - sender_id = edi_msg.get("sender_id") - receiver_id = edi_msg.get("receiver_id") - transaction_type = edi_msg.get("transaction_type", "*") + outbound_route_id = edi_msg.get("outbound_route_id") - if not sender_id or not receiver_id: - raise ValueError(f"EDI Message {trace_id} is missing sender/receiver IDs for routing.") + if direction == "OUTBOUND" and outbound_route_id: + route = await self.repository.get_outbound_route(outbound_route_id) + if not route: + logger.error(f"Configured outbound route {outbound_route_id} not found") + raise ValueError(f"Configured outbound route {outbound_route_id} not found") + else: + sender_id = edi_msg.get("sender_id") + receiver_id = edi_msg.get("receiver_id") + transaction_type = edi_msg.get("transaction_type", "*") - route = await self.repository.get_route(direction, sender_id, receiver_id, transaction_type) - if not route: - logger.error(f"No {direction} route found for {sender_id}->{receiver_id}") - raise ValueError(f"No route found for {direction} {sender_id}->{receiver_id}") + if not sender_id or not receiver_id: + raise ValueError( + f"EDI Message {trace_id} is missing sender/receiver IDs for routing." + ) + + route = await self.repository.get_route( + direction, sender_id, receiver_id, transaction_type + ) + if not route: + logger.error(f"No {direction} route found for {sender_id}->{receiver_id}") + raise ValueError(f"No route found for {direction} {sender_id}->{receiver_id}") # ── Dispatch via registry (OCP) ─────────────────────────────────────── for route_key, handler_name in _HANDLER_KEYS: @@ -113,7 +122,10 @@ async def _deliver_webhook( raise ValueError(f"Webhook partner {partner_id} not found.") try: - raw_payload = await self.storage.download(api_payload["request"]) + import json + + # Extract raw bytes from the dictionary payload + raw_payload = json.dumps(api_payload.get("payload")).encode("utf-8") auth_token = None if partner.get("auth_header_vault_ref") and self.vault: @@ -145,7 +157,7 @@ async def _deliver_sftp(self, trace_id: str, partner_id: str, edi_msg: dict[str, raise ValueError(f"SFTP partner {partner_id} not found.") try: - raw_payload = await self.storage.download(edi_msg["edi_data"]) + raw_payload = edi_msg["edi_data"].encode("utf-8") filename = f"{trace_id}.edi" password: str | None = partner.get("password") @@ -197,7 +209,7 @@ async def _deliver_as2(self, trace_id: str, partner_id: str, edi_msg: dict[str, else None ) - raw_payload = await self.storage.download(edi_msg["edi_data"]) + raw_payload = edi_msg["edi_data"].encode("utf-8") as2_msg = await self._as2_orchestrator.build( raw_payload=raw_payload, diff --git a/libs/pipeline/src/pipeline/core/metadata_extractor.py b/libs/pipeline/src/pipeline/core/metadata_extractor.py new file mode 100644 index 00000000..c2871c13 --- /dev/null +++ b/libs/pipeline/src/pipeline/core/metadata_extractor.py @@ -0,0 +1,95 @@ +import logging +from typing import Any + +from jsonpath_ng import parse # type: ignore[import-untyped] + +logger = logging.getLogger(__name__) + +# Enterprise configuration mapping transaction types to JSONPath expressions +# Using recursive descendant paths ($..) to extract fields accurately regardless +# of whether the JSON is wrapped in ISA/GS envelopes (Inbound) or is just a bare transaction (Outbound). +EXTRACTOR_CONFIG: dict[str, dict[str, str]] = { + "850": { + "po_number": "$..BEG.BEG03", + "po_date": "$..BEG.BEG05", + "business_reference": "$..BEG.BEG03", + }, + "810": { + "invoice_number": "$..BIG.BIG02", + "po_number": "$..BIG.BIG04", + "business_reference": "$..BIG.BIG02", + }, + "204": { + "load_number": "$..B2.B204", + "business_reference": "$..B2.B204", + }, + "990": { + "load_number": "$..B1.B102", + "business_reference": "$..B1.B102", + }, + "214": { + "load_number": "$..B10.B1002", + "business_reference": "$..B10.B1002", + }, + "210": { + "invoice_number": "$..B3.B302", + "business_reference": "$..B3.B302", + }, + "997": { + "group_control_number": "$..AK1.AK102", + "business_reference": "$..AK1.AK102", + }, +} + + +class MetadataExtractorService: + """ + Service responsible for dynamically extracting business fields from a structured JSON + payload using JSONPath expressions defined in the configuration. + """ + + def __init__(self, config: dict[str, dict[str, str]] | None = None) -> None: + self.config = config or EXTRACTOR_CONFIG + # Pre-compile the JSONPath expressions for performance + self.compiled_config: dict[str, dict[str, Any]] = {} + self._compile_config() + + def _compile_config(self) -> None: + for tx_type, paths in self.config.items(): + self.compiled_config[tx_type] = {} + for field_name, json_path in paths.items(): + try: + self.compiled_config[tx_type][field_name] = parse(json_path) + except Exception as e: + logger.error( + f"Failed to compile JSONPath '{json_path}' for {tx_type}.{field_name}: {e}" + ) + + def extract(self, transaction_type: str, payload: dict[str, Any]) -> dict[str, str]: + """ + Extracts metadata fields from the payload based on the transaction type. + Returns a flat dictionary of extracted key-value strings. + """ + if not transaction_type or transaction_type not in self.compiled_config: + logger.debug( + f"No extractor configuration found for transaction type: {transaction_type}" + ) + return {} + + extracted_metadata: dict[str, str] = {} + extractors = self.compiled_config[transaction_type] + + for field_name, jsonpath_expr in extractors.items(): + try: + matches = jsonpath_expr.find(payload) + if matches: + # We take the first match's value as a string + val = matches[0].value + if val is not None: + extracted_metadata[field_name] = str(val) + except Exception as e: + logger.warning( + f"Error extracting field '{field_name}' for type '{transaction_type}': {e}" + ) + + return extracted_metadata diff --git a/libs/pipeline/src/pipeline/core/translate.py b/libs/pipeline/src/pipeline/core/translate.py index 6d6aefe6..df11a417 100644 --- a/libs/pipeline/src/pipeline/core/translate.py +++ b/libs/pipeline/src/pipeline/core/translate.py @@ -1,9 +1,7 @@ -import json import logging import uuid from pipeline.ports.repository import RepositoryPort -from pipeline.ports.storage import StoragePort from pipeline.ports.transformer import TransformerPort logger = logging.getLogger(__name__) @@ -17,27 +15,87 @@ class TranslationService: def __init__( self, - storage: StoragePort, transformer: TransformerPort, repository: RepositoryPort, ) -> None: - self.storage = storage self.transformer = transformer self.repository = repository - async def translate(self, trace_id: str) -> None: + async def translate(self, trace_id: str, event_type: str = "edi_message.received") -> None: """ Translates an incoming message (EDI or JSON) into its target format. """ - logger.info(f"Starting translation pipeline for trace_id={trace_id}") + logger.info( + f"Starting translation pipeline for trace_id={trace_id} event_type={event_type}" + ) + + if event_type == "json.received": + await self._translate_json_to_edi(trace_id) + else: + # Default to EDI to JSON for edi_message.received or TRANSLATE + await self._translate_edi_to_json(trace_id) + + async def _translate_json_to_edi(self, trace_id: str) -> None: + """Translates an outbound JSON payload to X12 EDI.""" + edi_json = await self.repository.get_edi_json(trace_id) + if not edi_json: + raise ValueError(f"No EdiJson record found for trace_id={trace_id}") + + outbound_route_id = edi_json.get("outbound_route_id") + if not outbound_route_id: + raise ValueError(f"No outbound_route_id set for trace_id={trace_id}") + + route_config = await self.repository.get_outbound_route(outbound_route_id) + if not route_config: + raise ValueError(f"Outbound route {outbound_route_id} not found") + + json_payload = edi_json["payload"] + standard = route_config.get("default_standard", "X12") + transaction_type = route_config.get( + "transaction_type", edi_json.get("transaction_type", "UNKNOWN") + ) + + raw_edi_bytes = await self.transformer.translate_json_to_edi( + payload=json_payload, + standard=standard, + transaction_type=transaction_type, + route_config=route_config, + ) + + edi_str = raw_edi_bytes.decode("utf-8") + + await self.repository.save_edi_message( + trace_id=trace_id, + direction="OUTBOUND", + edi_data=edi_str, + format_standard=standard, + transaction_type=transaction_type, + status="PENDING_DELIVERY", + connection_type=route_config.get("connection_type", "UNKNOWN"), + sender_id=route_config.get("isa_sender_id"), + receiver_id=route_config.get("isa_receiver_id"), + outbound_route_id=outbound_route_id, + tenant_id=edi_json.get("tenant_id"), + ) + await self.repository.update_edi_json_status(trace_id, "TRANSLATED") + + deliver_idempotency_key = str(uuid.uuid5(uuid.NAMESPACE_OID, f"{trace_id}:DELIVER")) + await self.repository.publish_outbox_event( + idempotency_key=deliver_idempotency_key, + event_type="DELIVER", + payload={"trace_id": trace_id}, + ) + logger.info(f"Successfully translated JSON to EDI for trace_id={trace_id}") + + async def _translate_edi_to_json(self, trace_id: str) -> None: + """Translates an inbound X12 EDI payload to JSON.""" edi_msg = await self.repository.get_edi_message(trace_id) if not edi_msg: raise ValueError(f"No EDI message found for trace_id={trace_id}") - # 1. Download payload - s3_uri = edi_msg["edi_data"] - raw_payload = await self.storage.download(s3_uri) + # 1. Fetch raw payload from repository + raw_payload = edi_msg["edi_data"].encode("utf-8") # 2. Translate standard = edi_msg.get("format_standard", "X12") @@ -46,26 +104,40 @@ async def translate(self, trace_id: str) -> None: payload=raw_payload, standard=standard, transaction_type=transaction_type ) - json_bytes = json.dumps(json_dict).encode("utf-8") + # 3. Extract Business Metadata + from pipeline.core.metadata_extractor import MetadataExtractorService - tenant_id = edi_msg.get("tenant_id") + extractor = MetadataExtractorService() + business_metadata = extractor.extract(transaction_type, json_dict) - # 3. Upload translated payload - new_s3_uri = await self.storage.upload( - payload=json_bytes, - key_prefix=f"tenants/{tenant_id}/api_gateway/{trace_id}", - file_name="translated.json", + # 4. Save to EdiJson + partnership_id_str = edi_msg.get("partnership_id") + if isinstance(partnership_id_str, uuid.UUID): + partnership_id_str = str(partnership_id_str) + + await self.repository.save_edi_json( + trace_id=trace_id, + direction="INBOUND", + partnership_id=partnership_id_str, + transaction_type=transaction_type, + standard=standard, + sender_id=edi_msg.get("sender_id"), + receiver_id=edi_msg.get("receiver_id"), + business_metadata=business_metadata, + payload=json_dict, + status="PENDING_DELIVERY", + tenant_id=edi_msg.get("tenant_id"), ) - # 4. Save ApiGateway to DB + # 5. Save ApiGateway to DB await self.repository.save_api_payload( trace_id=trace_id, direction="OUTBOUND", - s3_uri=new_s3_uri, + payload=json_dict, status="PENDING_DELIVERY", ) - # 5. Publish DELIVER event with a stable idempotency key derived from trace_id + # 6. Publish DELIVER event with a stable idempotency key derived from trace_id deliver_idempotency_key = str(uuid.uuid5(uuid.NAMESPACE_OID, f"{trace_id}:DELIVER")) await self.repository.publish_outbox_event( idempotency_key=deliver_idempotency_key, @@ -73,6 +145,6 @@ async def translate(self, trace_id: str) -> None: payload={"trace_id": trace_id}, ) - # 6. Update EDI message status + # 7. Update EDI message status await self.repository.update_edi_message_status(trace_id, "TRANSLATED") - logger.info(f"Successfully translated trace_id={trace_id}") + logger.info(f"Successfully translated EDI to JSON for trace_id={trace_id}") diff --git a/libs/pipeline/src/pipeline/ports/repository.py b/libs/pipeline/src/pipeline/ports/repository.py index 8bcf9957..b80739aa 100644 --- a/libs/pipeline/src/pipeline/ports/repository.py +++ b/libs/pipeline/src/pipeline/ports/repository.py @@ -11,6 +11,23 @@ async def get_edi_message(self, trace_id: str) -> dict[str, Any] | None: """Fetches an EDI Message by trace_id.""" ... + async def save_edi_message( + self, + trace_id: str, + direction: str, + edi_data: str, + format_standard: str, + transaction_type: str, + status: str, + connection_type: str | None = "UNKNOWN", + sender_id: str | None = None, + receiver_id: str | None = None, + outbound_route_id: str | None = None, + tenant_id: int | None = None, + ) -> None: + """Saves a newly generated EDI Message.""" + ... + async def update_edi_message_status(self, trace_id: str, status: str) -> None: """Updates the status of an EDI Message.""" ... @@ -27,15 +44,40 @@ class APIPayloadPort(Protocol): """ async def save_api_payload( - self, trace_id: str, direction: str, s3_uri: str, status: str + self, trace_id: str, direction: str, payload: dict[str, Any], status: str ) -> None: """Persists a new JSON API Payload record.""" ... + async def save_edi_json( + self, + trace_id: str, + direction: str, + partnership_id: str | None, + transaction_type: str | None, + standard: str | None, + sender_id: str | None, + receiver_id: str | None, + business_metadata: dict[str, Any], + payload: dict[str, Any], + status: str, + tenant_id: int | None = None, + ) -> str: + """Persists a new EdiJson record and returns its UUID as a string.""" + ... + async def get_api_payload(self, trace_id: str) -> dict[str, Any] | None: """Fetches an API Payload by trace_id.""" ... + async def get_edi_json(self, trace_id: str) -> dict[str, Any] | None: + """Fetches an EdiJson record by trace_id.""" + ... + + async def update_edi_json_status(self, trace_id: str, status: str) -> None: + """Updates the status of an EdiJson record.""" + ... + async def update_api_payload_status(self, trace_id: str, status: str) -> None: """Updates the status of an API Payload.""" ... @@ -57,6 +99,10 @@ async def get_route( """Finds the appropriate route based on ISA envelopes.""" ... + async def get_outbound_route(self, route_id: str) -> dict[str, Any] | None: + """Fetches an outbound route by its ID.""" + ... + async def publish_outbox_event( self, idempotency_key: str, event_type: str, payload: dict[str, Any] ) -> None: diff --git a/libs/pipeline/src/pipeline/ports/transformer.py b/libs/pipeline/src/pipeline/ports/transformer.py index 5fc35cc2..ef6e82a7 100644 --- a/libs/pipeline/src/pipeline/ports/transformer.py +++ b/libs/pipeline/src/pipeline/ports/transformer.py @@ -13,7 +13,11 @@ async def translate_edi_to_json( ... async def translate_json_to_edi( - self, payload: dict[str, Any], standard: str, transaction_type: str + self, + payload: dict[str, Any], + standard: str, + transaction_type: str, + route_config: dict[str, Any], ) -> bytes: """Translates a Canonical JSON Dictionary into raw EDI bytes.""" ... diff --git a/libs/pipeline/tests/fakes.py b/libs/pipeline/tests/fakes.py index 7fc735ab..22651229 100644 --- a/libs/pipeline/tests/fakes.py +++ b/libs/pipeline/tests/fakes.py @@ -71,12 +71,12 @@ async def claim_edi_message(self, trace_id: str) -> bool: return False async def save_api_payload( - self, trace_id: str, direction: str, s3_uri: str, status: str + self, trace_id: str, direction: str, payload: dict[str, Any], status: str ) -> None: self.api_gateway[trace_id] = { "trace_id": trace_id, "direction": direction, - "request": s3_uri, + "payload": payload, "status": status, } diff --git a/libs/pipeline/tests/test_delivery_service.py b/libs/pipeline/tests/test_delivery_service.py index ffe0d1f8..b229aca8 100644 --- a/libs/pipeline/tests/test_delivery_service.py +++ b/libs/pipeline/tests/test_delivery_service.py @@ -17,7 +17,6 @@ def make_service( - storage: InMemoryStorageAdapter | None = None, repo: InMemoryRepositoryAdapter | None = None, http: FakeHttpDeliveryAdapter | None = None, sftp: FakeSftpDeliveryAdapter | None = None, @@ -25,11 +24,11 @@ def make_service( ) -> DeliveryService: """Factory that satisfies the required as2_delivery port (Null Object not needed in tests).""" return DeliveryService( - storage=storage or InMemoryStorageAdapter(), repository=repo or InMemoryRepositoryAdapter(), http_delivery=http or FakeHttpDeliveryAdapter(), sftp_delivery=sftp or FakeSftpDeliveryAdapter(), as2_delivery=as2 or FakeAS2DeliveryAdapter(), + vault=None, ) @@ -76,7 +75,7 @@ async def test_delivery_service_inbound_webhook() -> None: } # ── Act ──────────────────────────────────────────────────────────────────── - service = make_service(storage=storage, repo=repo, http=http_adapter) + service = make_service(repo=repo, http=http_adapter) await service.deliver(trace_id) # ── Assert ───────────────────────────────────────────────────────────────── @@ -101,7 +100,7 @@ async def test_delivery_service_outbound_sftp() -> None: "sender_id": "SENDER1", "receiver_id": "RECV1", "transaction_type": "855", - "edi_data": edi_s3_uri, + "edi_data": "FAKE*EDI*DATA~", "status": "PENDING_DELIVERY", } repo.routes.append( @@ -127,7 +126,7 @@ async def test_delivery_service_outbound_sftp() -> None: vault = FakeVault({"mock_password": "fake_private_key_data"}) - service = make_service(storage=storage, repo=repo, sftp=sftp_adapter) + service = make_service(repo=repo, sftp=sftp_adapter) service.vault = vault await service.deliver(trace_id) @@ -195,7 +194,7 @@ async def test_delivery_service_http_failure_sets_failed_status() -> None: } # ── Act ──────────────────────────────────────────────────────────────────── - service = make_service(storage=storage, repo=repo, http=http_adapter) + service = make_service(repo=repo, http=http_adapter) await service.deliver(trace_id) # ── Assert ───────────────────────────────────────────────────────────────── diff --git a/libs/pipeline/tests/test_delivery_service_as2.py b/libs/pipeline/tests/test_delivery_service_as2.py index 6ccd8b25..79ac6516 100644 --- a/libs/pipeline/tests/test_delivery_service_as2.py +++ b/libs/pipeline/tests/test_delivery_service_as2.py @@ -44,23 +44,22 @@ def make_service( - storage: InMemoryStorageAdapter | None = None, repo: InMemoryRepositoryAdapter | None = None, as2: FakeAS2DeliveryAdapter | NullAS2DeliveryAdapter | None = None, ) -> DeliveryService: return DeliveryService( - storage=storage or InMemoryStorageAdapter(), repository=repo or InMemoryRepositoryAdapter(), http_delivery=FakeHttpDeliveryAdapter(), sftp_delivery=FakeSftpDeliveryAdapter(), as2_delivery=as2 or FakeAS2DeliveryAdapter(), + vault=None, ) def _seed_as2_route( repo: InMemoryRepositoryAdapter, trace_id: str, - edi_s3_uri: str, + edi_data: str, partner_id: str = "remote-p1", transaction_type: str = "850", ) -> None: @@ -71,7 +70,7 @@ def _seed_as2_route( "sender_id": "SENDER", "receiver_id": "RECEIVER", "transaction_type": transaction_type, - "edi_data": edi_s3_uri, + "edi_data": edi_data, "status": "PENDING_DELIVERY", } repo.routes.append( @@ -108,10 +107,10 @@ async def test_deliver_as2_plain_no_crypto() -> None: b"*ZZ*RECEIVER *210101*1200*^*00501*000000001*0*P*>~" ) storage.store[edi_s3_uri] = raw_edi - _seed_as2_route(repo, trace_id, edi_s3_uri) + _seed_as2_route(repo, trace_id, raw_edi.decode("utf-8")) # ── Act ──────────────────────────────────────────────────────────────────── - await make_service(storage=storage, repo=repo, as2=as2_adapter).deliver(trace_id) + await make_service(repo=repo, as2=as2_adapter).deliver(trace_id) # ── Assert ───────────────────────────────────────────────────────────────── assert len(as2_adapter.delivered) == 1 @@ -146,7 +145,7 @@ async def test_deliver_as2_http_failure_sets_failed_status() -> None: "sender_id": "S1", "receiver_id": "R1", "transaction_type": "856", - "edi_data": edi_s3_uri, + "edi_data": "FAKE*EDI~", "status": "PENDING_DELIVERY", } repo.routes.append( @@ -164,7 +163,7 @@ async def test_deliver_as2_http_failure_sets_failed_status() -> None: repo.local_as2_partners[_REMOTE_PARTNER["local_partner_id"]] = _LOCAL_PARTNER # ── Act ──────────────────────────────────────────────────────────────────── - await make_service(storage=storage, repo=repo, as2=as2_adapter).deliver(trace_id) + await make_service(repo=repo, as2=as2_adapter).deliver(trace_id) # ── Assert ───────────────────────────────────────────────────────────────── assert repo.edi_messages[trace_id]["status"] == "FAILED" @@ -184,10 +183,10 @@ async def test_deliver_as2_null_adapter_is_caught_and_marked_failed() -> None: trace_id = "trace-as2-null" edi_s3_uri = f"s3://bucket/edi/{trace_id}/raw.edi" storage.store[edi_s3_uri] = b"EDI~" - _seed_as2_route(repo, trace_id, edi_s3_uri, partner_id="p-null") + _seed_as2_route(repo, trace_id, "EDI~", partner_id="p-null") # ── Act / Assert ─────────────────────────────────────────────────────────── - service = make_service(storage=storage, repo=repo, as2=NullAS2DeliveryAdapter()) + service = make_service(repo=repo, as2=NullAS2DeliveryAdapter()) await service.deliver(trace_id) # It should catch the RuntimeError and mark the message as FAILED @@ -215,7 +214,7 @@ async def test_deliver_as2_idempotent_claim() -> None: "sender_id": "A", "receiver_id": "B", "transaction_type": "810", - "edi_data": edi_s3_uri, + "edi_data": "EDI~", "status": "PROCESSING", } repo.routes.append( @@ -232,7 +231,7 @@ async def test_deliver_as2_idempotent_claim() -> None: repo.local_as2_partners[_REMOTE_PARTNER["local_partner_id"]] = _LOCAL_PARTNER # ── Act ──────────────────────────────────────────────────────────────────── - await make_service(storage=storage, repo=repo, as2=as2_adapter).deliver(trace_id) + await make_service(repo=repo, as2=as2_adapter).deliver(trace_id) # ── Assert ───────────────────────────────────────────────────────────────── assert len(as2_adapter.delivered) == 0 @@ -260,7 +259,7 @@ async def test_deliver_as2_missing_local_partner_sets_failed() -> None: "sender_id": "X", "receiver_id": "Y", "transaction_type": "850", - "edi_data": edi_s3_uri, + "edi_data": "EDI~", "status": "PENDING_DELIVERY", } repo.routes.append( @@ -278,7 +277,7 @@ async def test_deliver_as2_missing_local_partner_sets_failed() -> None: # Do NOT seed local_as2_partners["missing-local"] # ── Act ──────────────────────────────────────────────────────────────────── - await make_service(storage=storage, repo=repo, as2=as2_adapter).deliver(trace_id) + await make_service(repo=repo, as2=as2_adapter).deliver(trace_id) # ── Assert ───────────────────────────────────────────────────────────────── assert repo.edi_messages[trace_id]["status"] == "FAILED" diff --git a/libs/pipeline/tests/test_http.py b/libs/pipeline/tests/test_http.py index b55089cf..f7147f03 100644 --- a/libs/pipeline/tests/test_http.py +++ b/libs/pipeline/tests/test_http.py @@ -7,9 +7,9 @@ @patch("pipeline.adapters.http.httpx.AsyncClient") -@patch("socket.gethostbyname") +@patch("socket.getaddrinfo") async def test_httpx_delivery_adapter( - mock_gethostbyname: MagicMock, mock_client_cls: MagicMock + mock_getaddrinfo: MagicMock, mock_client_cls: MagicMock ) -> None: mock_client = AsyncMock() mock_client_cls.return_value.__aenter__.return_value = mock_client @@ -18,7 +18,7 @@ async def test_httpx_delivery_adapter( mock_response.status_code = 200 mock_client.post.return_value = mock_response - mock_gethostbyname.return_value = "93.184.216.34" + mock_getaddrinfo.return_value = [(2, 1, 6, "", ("93.184.216.34", 443))] adapter = HttpxDeliveryAdapter(timeout_secs=5) diff --git a/libs/pipeline/tests/test_pipeline_repository.py b/libs/pipeline/tests/test_pipeline_repository.py index 43ffd202..a101340a 100644 --- a/libs/pipeline/tests/test_pipeline_repository.py +++ b/libs/pipeline/tests/test_pipeline_repository.py @@ -2,12 +2,21 @@ from unittest.mock import AsyncMock, MagicMock import pytest +from config.settings import AppSettings from database.models import ApiGateway, EdiMessage +from fakes import InMemoryStorageAdapter from pipeline.adapters.repository import SqlAlchemyRepositoryAdapter pytestmark = pytest.mark.asyncio +def make_adapter(session): + settings = AppSettings() + settings.storage_backend = "local" + storage = InMemoryStorageAdapter() + return SqlAlchemyRepositoryAdapter(session, settings, storage) + + async def test_get_edi_message_success() -> None: mock_session = AsyncMock() mock_result = MagicMock() @@ -17,12 +26,13 @@ async def test_get_edi_message_success() -> None: mock_record.edi_data = "s3://foo" mock_record.format_standard = "X12" mock_record.transaction_type = "850" + mock_record.storage_uri = None mock_record.status = "RECEIVED" mock_result.scalar_one_or_none.return_value = mock_record mock_session.execute.return_value = mock_result - adapter = SqlAlchemyRepositoryAdapter(mock_session) + adapter = make_adapter(mock_session) result = await adapter.get_edi_message(str(mock_record.trace_id)) assert result is not None @@ -33,7 +43,7 @@ async def test_get_edi_message_success() -> None: async def test_update_edi_message_status() -> None: mock_session = AsyncMock() - adapter = SqlAlchemyRepositoryAdapter(mock_session) + adapter = make_adapter(mock_session) trace_id = str(uuid.uuid4()) await adapter.update_edi_message_status(trace_id, "TRANSLATED") @@ -45,23 +55,23 @@ async def test_update_edi_message_status() -> None: async def test_save_api_payload() -> None: mock_session = AsyncMock() mock_session.add = MagicMock() - adapter = SqlAlchemyRepositoryAdapter(mock_session) + adapter = make_adapter(mock_session) trace_id = str(uuid.uuid4()) - await adapter.save_api_payload(trace_id, "OUTBOUND", "s3://out", "PENDING_DELIVERY") + await adapter.save_api_payload(trace_id, "OUTBOUND", {"data": "foo"}, "PENDING_DELIVERY") mock_session.add.assert_called_once() added_obj = mock_session.add.call_args[0][0] assert isinstance(added_obj, ApiGateway) assert str(added_obj.trace_id) == trace_id - assert added_obj.request == "s3://out" + assert added_obj.payload == {"data": "foo"} mock_session.flush.assert_awaited_once() async def test_publish_outbox_event() -> None: mock_session = AsyncMock() - adapter = SqlAlchemyRepositoryAdapter(mock_session) + adapter = make_adapter(mock_session) idempotency_key = str(uuid.uuid4()) await adapter.publish_outbox_event(idempotency_key, "DELIVER", {"trace_id": "123"}) @@ -76,23 +86,26 @@ async def test_get_api_payload() -> None: mock_record = MagicMock(spec=ApiGateway) mock_record.trace_id = uuid.uuid4() - mock_record.request = "s3://out" + mock_record.storage_uri = "s3://out" mock_record.status = "PENDING_DELIVERY" + # fake storage needs the uri + adapter = make_adapter(mock_session) + adapter.storage.store["s3://out"] = b'{"data": "foo"}' + mock_result.scalar_one_or_none.return_value = mock_record mock_session.execute.return_value = mock_result - adapter = SqlAlchemyRepositoryAdapter(mock_session) result = await adapter.get_api_payload(str(mock_record.trace_id)) assert result is not None assert result["status"] == "PENDING_DELIVERY" - assert result["request"] == "s3://out" + assert result["payload"] == {"data": "foo"} async def test_update_api_payload_status() -> None: mock_session = AsyncMock() - adapter = SqlAlchemyRepositoryAdapter(mock_session) + adapter = make_adapter(mock_session) trace_id = str(uuid.uuid4()) await adapter.update_api_payload_status(trace_id, "DELIVERED") diff --git a/libs/pipeline/tests/test_transformer.py b/libs/pipeline/tests/test_transformer.py index 15db1f10..7a219fcf 100644 --- a/libs/pipeline/tests/test_transformer.py +++ b/libs/pipeline/tests/test_transformer.py @@ -26,8 +26,11 @@ async def test_bots_transformer_edi_to_json(mock_translate: AsyncMock) -> None: mock_translate.assert_awaited_once_with(b"ISA*00*") -async def test_bots_transformer_json_to_edi_raises_not_implemented() -> None: - """translate_json_to_edi is not yet implemented and must raise explicitly.""" +async def test_bots_transformer_json_to_edi_success() -> None: + from unittest.mock import patch + adapter = BotsTransformerAdapter() - with pytest.raises(NotImplementedError, match="JSON to EDI translation is not yet supported"): - await adapter.translate_json_to_edi({"foo": "bar"}, "X12", "850") + + with patch.object(adapter._adapter, "serialize_to_edi", return_value=("ISA*...~", [])): + result = await adapter.translate_json_to_edi({"foo": "bar"}, "X12", "850", {}) + assert result == b"ISA*...~" diff --git a/libs/pipeline/tests/test_translation_service.py b/libs/pipeline/tests/test_translation_service.py index 7541d99a..0ae37d59 100644 --- a/libs/pipeline/tests/test_translation_service.py +++ b/libs/pipeline/tests/test_translation_service.py @@ -18,14 +18,14 @@ async def test_translate_edi_to_json_success() -> None: repo.edi_messages[trace_id] = { "trace_id": trace_id, - "edi_data": s3_uri, + "edi_data": "ISA*00*...", "format_standard": "X12", "transaction_type": "850", "status": "RECEIVED", } # Act - service = TranslationService(storage, transformer, repo) + service = TranslationService(transformer, repo) await service.translate(trace_id) # Assert @@ -37,15 +37,12 @@ async def test_translate_edi_to_json_success() -> None: assert transformer.translate_edi_calls[0]["payload"] == b"ISA*00*..." assert transformer.translate_edi_calls[0]["standard"] == "X12" - # 3. JSON payload uploaded to storage - assert storage.upload_count == 1 - - # 4. ApiGateway record created in DB + # 3. ApiGateway record created in DB assert trace_id in repo.api_gateway api_payload = repo.api_gateway[trace_id] assert api_payload["direction"] == "OUTBOUND" assert api_payload["status"] == "PENDING_DELIVERY" - assert api_payload["request"].startswith("s3://fake-bucket") + assert isinstance(api_payload["payload"], dict) # 5. Outbox event published for DELIVER assert len(repo.outbox) == 1 @@ -55,11 +52,11 @@ async def test_translate_edi_to_json_success() -> None: async def test_translate_missing_message_raises_error() -> None: - storage = InMemoryStorageAdapter() + InMemoryStorageAdapter() transformer = FakeTransformerAdapter() repo = InMemoryRepositoryAdapter() - service = TranslationService(storage, transformer, repo) + service = TranslationService(transformer, repo) with pytest.raises(ValueError, match="No EDI message found for trace_id=invalid-trace"): await service.translate("invalid-trace") diff --git a/libs/transformer/src/transformer/domain/ast_utils.py b/libs/transformer/src/transformer/domain/ast_utils.py new file mode 100644 index 00000000..d59c0224 --- /dev/null +++ b/libs/transformer/src/transformer/domain/ast_utils.py @@ -0,0 +1,44 @@ +from typing import Any + + +class ASTUtils: + """ + Utility class for interacting with the JSON AST format used by the translation engine. + """ + + @staticmethod + def count_segments(txn: dict[str, Any]) -> int: + """ + Recursively counts the number of valid EDI segments in a transaction AST dictionary. + """ + count = 0 + + def traverse(node: Any) -> None: + nonlocal count + if isinstance(node, dict): + for k, v in node.items(): + # Segments are uppercase, 2-3 alphanumeric characters + if k.isupper() and 2 <= len(k) <= 3 and k.isalnum(): + if isinstance(v, list) and len(v) > 0: + # Check if the list contains standard segments or if it's a loop + first_item = v[0] + is_loop = False + if isinstance(first_item, dict): + for sub_k in first_item: + if sub_k.isupper() and 2 <= len(sub_k) <= 3 and sub_k.isalnum(): + is_loop = True + break + if is_loop: + traverse(v) + else: + count += len(v) + elif isinstance(v, dict): + count += 1 + elif isinstance(v, (dict, list)): + traverse(v) + elif isinstance(node, list): + for item in node: + traverse(item) + + traverse(txn) + return count diff --git a/libs/transformer/src/transformer/domain/envelope/__init__.py b/libs/transformer/src/transformer/domain/envelope/__init__.py new file mode 100644 index 00000000..fe07c63f --- /dev/null +++ b/libs/transformer/src/transformer/domain/envelope/__init__.py @@ -0,0 +1,4 @@ +from transformer.domain.envelope.edifact import EdifactEnvelopeBuilder +from transformer.domain.envelope.x12 import X12EnvelopeBuilder + +__all__ = ["X12EnvelopeBuilder", "EdifactEnvelopeBuilder"] diff --git a/libs/transformer/src/transformer/domain/envelope/base.py b/libs/transformer/src/transformer/domain/envelope/base.py new file mode 100644 index 00000000..4e6f25a5 --- /dev/null +++ b/libs/transformer/src/transformer/domain/envelope/base.py @@ -0,0 +1,19 @@ +from abc import ABC, abstractmethod +from typing import Any + + +class BaseEnvelopeBuilder(ABC): + """ + Abstract Base Class (Interface) for EDI Envelope Builders. + Enforces that all standards (X12, EDIFACT, etc.) implement the build method. + """ + + @staticmethod + @abstractmethod + def build( + route_config: dict[str, Any], payload: dict[str, Any] | list[dict[str, Any]] + ) -> dict[str, Any]: + """ + Dynamically constructs the Abstract Syntax Tree (AST) for the given payload and route. + """ + pass diff --git a/libs/transformer/src/transformer/domain/envelope/edifact.py b/libs/transformer/src/transformer/domain/envelope/edifact.py new file mode 100644 index 00000000..bdcfb28c --- /dev/null +++ b/libs/transformer/src/transformer/domain/envelope/edifact.py @@ -0,0 +1,89 @@ +import datetime +import uuid +from typing import Any + +from transformer.domain.ast_utils import ASTUtils +from transformer.domain.envelope.base import BaseEnvelopeBuilder + + +class EdifactEnvelopeBuilder(BaseEnvelopeBuilder): + @classmethod + def _build_unb_segment( + cls, route_config: dict[str, Any], now: datetime.datetime, unb05: str + ) -> dict[str, Any]: + unb_sender_id = route_config.get("isa_sender_id", "UNKNOWN") + unb_receiver_id = route_config.get("isa_receiver_id", "UNKNOWN") + version = route_config.get("default_version", "4") + environment = "1" if route_config.get("environment") == "T" else "" + + unb = { + "S001.01": "UNOA", + "S001.02": version, + "S002.01": unb_sender_id, + "S002.02": route_config.get("isa_sender_qualifier", "14"), + "S003.01": unb_receiver_id, + "S003.02": route_config.get("isa_receiver_qualifier", "14"), + "S004.01": now.strftime("%y%m%d"), + "S004.02": now.strftime("%H%M"), + "0020": unb05, + } + if environment: + unb["S005.01"] = "XX" + + return unb + + @classmethod + def _wrap_transactions( + cls, transactions: list[dict[str, Any]], transaction_type: str + ) -> list[dict[str, Any]]: + processed_transactions = [] + for i, txn in enumerate(transactions, start=1): + new_txn = {} + if "UNH" not in txn: + new_txn["UNH"] = { + "UNH01": f"{i:04d}", + "UNH02": { + "UNH02.01": transaction_type, + "UNH02.02": "D", + "UNH02.03": "96A", + "UNH02.04": "UN", + }, + } + + for k, v in txn.items(): + if k not in new_txn: + new_txn[k] = v + + if "UNT" not in new_txn: + segment_count = ASTUtils.count_segments(new_txn) + 1 + new_txn["UNT"] = { + "UNT01": str(segment_count), + "UNT02": new_txn.get("UNH", {}).get("UNH01", f"{i:04d}"), + } + + processed_transactions.append(new_txn) + return processed_transactions + + @classmethod + def build( + cls, route_config: dict[str, Any], payload: dict[str, Any] | list[dict[str, Any]] + ) -> dict[str, Any]: + now = datetime.datetime.now(datetime.UTC) + transactions = payload if isinstance(payload, list) else [payload] + transaction_type = route_config.get("transaction_type", "UNKNOWN") + + # Generation values + unb05 = str(uuid.uuid4().int % 1000000000).zfill(9) + + # Build segments + unb_segment = cls._build_unb_segment(route_config, now, unb05) + processed_transactions = cls._wrap_transactions(transactions, transaction_type) + + unz_segment = {"UNZ01": str(len(processed_transactions)), "UNZ02": unb05} + + # Orchestrate the final AST structure + return { + "interchange_UNB": [ + {"UNB": unb_segment, "transaction_UNH": processed_transactions, "UNZ": unz_segment} + ] + } diff --git a/libs/transformer/src/transformer/domain/envelope/x12.py b/libs/transformer/src/transformer/domain/envelope/x12.py new file mode 100644 index 00000000..08a282ca --- /dev/null +++ b/libs/transformer/src/transformer/domain/envelope/x12.py @@ -0,0 +1,142 @@ +import datetime +from typing import Any + +from transformer.domain.ast_utils import ASTUtils +from transformer.domain.envelope.base import BaseEnvelopeBuilder + +X12_GS01_MAPPING = { + "850": "PO", + "810": "IN", + "856": "SH", + "855": "PR", + "846": "IB", + "832": "SC", + "204": "SM", + "210": "IM", + "214": "QM", + "990": "GF", + "997": "FA", +} + + +class X12EnvelopeBuilder(BaseEnvelopeBuilder): + @classmethod + def _build_isa_segment( + cls, route_config: dict[str, Any], now: datetime.datetime, isa13: str + ) -> dict[str, Any]: + isa_sender_qualifier = route_config.get("isa_sender_qualifier") or "ZZ" + isa_sender_id = str(route_config.get("isa_sender_id", "UNKNOWN")).ljust(15) + isa_receiver_qualifier = route_config.get("isa_receiver_qualifier") or "ZZ" + isa_receiver_id = str(route_config.get("isa_receiver_id", "UNKNOWN")).ljust(15) + + version = route_config.get("default_version", "004010") + isa_version = version[:5] if len(version) >= 5 else "00401" + environment = route_config.get("environment", "P") + + return { + "ISA01": "00", + "ISA02": " ", + "ISA03": "00", + "ISA04": " ", + "ISA05": isa_sender_qualifier, + "ISA06": isa_sender_id, + "ISA07": isa_receiver_qualifier, + "ISA08": isa_receiver_id, + "ISA09": now.strftime("%y%m%d"), + "ISA10": now.strftime("%H%M"), + "ISA11": "U", + "ISA12": isa_version, + "ISA13": isa13, + "ISA14": "0", + "ISA15": environment, + } + + @classmethod + def _build_gs_segment( + cls, route_config: dict[str, Any], now: datetime.datetime, gs06: str + ) -> dict[str, Any]: + transaction_type = route_config.get("transaction_type", "UNKNOWN") + gs_sender_id = route_config.get("gs_sender_id") or route_config.get( + "isa_sender_id", "UNKNOWN" + ) + gs_receiver_id = route_config.get("gs_receiver_id") or route_config.get( + "isa_receiver_id", "UNKNOWN" + ) + version = route_config.get("default_version", "004010") + gs01 = X12_GS01_MAPPING.get(transaction_type, "XX") + + return { + "GS01": gs01, + "GS02": gs_sender_id, + "GS03": gs_receiver_id, + "GS04": now.strftime("%Y%m%d"), + "GS05": now.strftime("%H%M"), + "GS06": gs06, + "GS07": "X", + "GS08": version, + } + + @classmethod + def _wrap_transactions( + cls, transactions: list[dict[str, Any]], transaction_type: str + ) -> list[dict[str, Any]]: + processed_transactions = [] + for i, txn in enumerate(transactions, start=1): + new_txn = {} + if "ST" not in txn: + new_txn["ST"] = {"ST01": transaction_type, "ST02": f"{i:04d}"} + + # Copy all business data in order + for k, v in txn.items(): + if k not in new_txn: + new_txn[k] = v + + if "SE" not in new_txn: + # Calculate segment count (existing + SE) + segment_count = ASTUtils.count_segments(new_txn) + 1 + new_txn["SE"] = { + "SE01": str(segment_count), + "SE02": new_txn.get("ST", {}).get("ST02", f"{i:04d}"), + } + + processed_transactions.append(new_txn) + return processed_transactions + + @classmethod + def build( + cls, route_config: dict[str, Any], payload: dict[str, Any] | list[dict[str, Any]] + ) -> dict[str, Any]: + now = datetime.datetime.now(datetime.UTC) + transactions = payload if isinstance(payload, list) else [payload] + transaction_type = route_config.get("transaction_type", "UNKNOWN") + + # Generation values + monotonic_counter = int(now.timestamp() * 1000) % 1000000000 + isa13 = f"{monotonic_counter:09d}" + gs06 = str(monotonic_counter) + + # Build segments + isa_segment = cls._build_isa_segment(route_config, now, isa13) + gs_segment = cls._build_gs_segment(route_config, now, gs06) + processed_transactions = cls._wrap_transactions(transactions, transaction_type) + + ge_segment = {"GE01": str(len(processed_transactions)), "GE02": gs06} + + iea_segment = {"IEA01": "1", "IEA02": isa13} + + # Orchestrate the final AST structure + return { + "interchange_ISA": [ + { + "ISA": isa_segment, + "group_GS": [ + { + "GS": gs_segment, + "transaction_ST": processed_transactions, + "GE": ge_segment, + } + ], + "IEA": iea_segment, + } + ] + } diff --git a/libs/transformer/src/transformer/domain/envelope_factory.py b/libs/transformer/src/transformer/domain/envelope_factory.py new file mode 100644 index 00000000..a1d82344 --- /dev/null +++ b/libs/transformer/src/transformer/domain/envelope_factory.py @@ -0,0 +1,32 @@ +from typing import Any + +from transformer.domain.envelope.base import BaseEnvelopeBuilder +from transformer.domain.envelope.edifact import EdifactEnvelopeBuilder +from transformer.domain.envelope.x12 import X12EnvelopeBuilder + + +class EnvelopeFactory: + """ + Enterprise-grade factory for dynamically constructing Abstract Syntax Trees (AST) + for various EDI standards (X12, EDIFACT) based on Route Configurations. + """ + + _BUILDERS: dict[str, type[BaseEnvelopeBuilder]] = { + "x12": X12EnvelopeBuilder, + "edifact": EdifactEnvelopeBuilder, + } + + @staticmethod + def build_ast( + route_config: dict[str, Any], payload: dict[str, Any] | list[dict[str, Any]] + ) -> dict[str, Any]: + """ + Dynamically dispatches to the correct standard builder based on route config. + """ + standard = str(route_config.get("default_standard", "x12")).lower().strip() + + builder = EnvelopeFactory._BUILDERS.get(standard) + if not builder: + raise ValueError(f"Unsupported EDI standard in Route Configuration: '{standard}'") + + return builder.build(route_config, payload) diff --git a/libs/transformer/tests/domain/__init__.py b/libs/transformer/tests/domain_tests/__init__.py similarity index 100% rename from libs/transformer/tests/domain/__init__.py rename to libs/transformer/tests/domain_tests/__init__.py diff --git a/libs/transformer/tests/domain/test_domain_models.py b/libs/transformer/tests/domain_tests/test_domain_models.py similarity index 100% rename from libs/transformer/tests/domain/test_domain_models.py rename to libs/transformer/tests/domain_tests/test_domain_models.py diff --git a/pyproject.toml b/pyproject.toml index cfccc966..8a28b1f7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -9,6 +9,7 @@ requires-python = ">=3.11" dependencies = [ "cryptography>=49.0.0", "endesive>=2.19.3", + "jsonpath-ng>=1.8.0", ] [tool.uv.workspace] @@ -30,7 +31,12 @@ pythonpath = [ "libs/observability/src", "libs/security/src", "libs/patches/src", + "libs/domain/src", + "libs/pipeline/src", + "libs/transformer/src", + "services/api/src", "services/as2_server/src", + "services/worker/src", ] addopts = "-v --tb=short" markers = [ diff --git a/scripts/purge_sqs.py b/scripts/purge_sqs.py new file mode 100644 index 00000000..1c69b45b --- /dev/null +++ b/scripts/purge_sqs.py @@ -0,0 +1,17 @@ +import boto3 + +sqs = boto3.client( + "sqs", + endpoint_url="http://localhost:4566", + region_name="us-east-1", + aws_access_key_id="test", + aws_secret_access_key="test", +) +queues = sqs.list_queues() +if "QueueUrls" in queues: + for q in queues["QueueUrls"]: + try: + sqs.purge_queue(QueueUrl=q) + print(f"Purged {q}") + except Exception as e: + print(f"Failed {q}: {e}") diff --git a/services/api/pyproject.toml b/services/api/pyproject.toml index d95c61c7..f0b8798b 100644 --- a/services/api/pyproject.toml +++ b/services/api/pyproject.toml @@ -21,6 +21,8 @@ dependencies = [ "observability", "security", "patches", + "pipeline", + "domain", "hvac>=2.4.0", "cryptography>=41.0.0", ] @@ -32,6 +34,8 @@ config = { workspace = true } observability = { workspace = true } security = { workspace = true } patches = { workspace = true } +pipeline = { workspace = true } +domain = { workspace = true } [build-system] diff --git a/services/api/src/api/adapters/http/dtos.py b/services/api/src/api/adapters/http/dtos.py index 958b30a8..f4919094 100644 --- a/services/api/src/api/adapters/http/dtos.py +++ b/services/api/src/api/adapters/http/dtos.py @@ -1,3 +1,4 @@ +from datetime import datetime from typing import Any, Literal from uuid import UUID @@ -155,6 +156,8 @@ class CreateInboundRouteRequest(BaseModel): name: str = Field(..., max_length=255, description="Name of the route") isa_sender_id: str = Field(..., max_length=255, description="ISA Sender ID to match") isa_receiver_id: str = Field(..., max_length=255, description="ISA Receiver ID to match") + gs_sender_id: str | None = Field(None, max_length=255, description="GS Sender ID to match") + gs_receiver_id: str | None = Field(None, max_length=255, description="GS Receiver ID to match") transaction_type: str = Field( ..., max_length=50, description="EDI Transaction Type (e.g., '204', '990', or '*')" ) @@ -180,12 +183,23 @@ def check_exactly_one_destination(self) -> "CreateInboundRouteRequest": class CreateOutboundRouteRequest(BaseModel): + trading_partner_id: str = Field( + ..., max_length=255, description="The ERP's identifier for this route" + ) name: str = Field(..., max_length=255, description="Name of the route") - isa_sender_id: str = Field(..., max_length=255, description="ISA Sender ID to match") - isa_receiver_id: str = Field(..., max_length=255, description="ISA Receiver ID to match") + isa_sender_id: str = Field(..., max_length=255, description="ISA Sender ID to map") + isa_sender_qualifier: str | None = Field(None, max_length=2, description="ISA Sender Qualifier") + isa_receiver_id: str = Field(..., max_length=255, description="ISA Receiver ID to map") + isa_receiver_qualifier: str | None = Field( + None, max_length=2, description="ISA Receiver Qualifier" + ) + gs_sender_id: str = Field(..., max_length=255, description="GS Sender ID to map") + gs_receiver_id: str = Field(..., max_length=255, description="GS Receiver ID to map") transaction_type: str = Field( ..., max_length=50, description="EDI Transaction Type (e.g., '204', '990', or '*')" ) + default_standard: str = Field("x12", max_length=50, description="EDI Standard") + default_version: str = Field("004010", max_length=50, description="EDI Version") processing_mode: Literal["TRANSLATE", "PASSTHROUGH"] = Field( "TRANSLATE", description="Processing Mode" ) @@ -206,10 +220,27 @@ def check_exactly_one_destination(self) -> "CreateOutboundRouteRequest": class UpdateRouteRequest(BaseModel): active: bool | None = None name: str | None = Field(None, max_length=255, description="Name of the route") + trading_partner_id: str | None = Field( + None, max_length=255, description="Trading Partner ID (Outbound only)" + ) isa_sender_id: str | None = Field(None, max_length=255, description="ISA Sender ID to match") + isa_sender_qualifier: str | None = Field( + None, max_length=2, description="ISA Sender Qualifier (Outbound only)" + ) isa_receiver_id: str | None = Field( None, max_length=255, description="ISA Receiver ID to match" ) + isa_receiver_qualifier: str | None = Field( + None, max_length=2, description="ISA Receiver Qualifier (Outbound only)" + ) + gs_sender_id: str | None = Field(None, max_length=255, description="GS Sender ID") + gs_receiver_id: str | None = Field(None, max_length=255, description="GS Receiver ID") + default_standard: str | None = Field( + None, max_length=50, description="EDI Standard (Outbound only)" + ) + default_version: str | None = Field( + None, max_length=50, description="EDI Version (Outbound only)" + ) transaction_type: str | None = Field( None, max_length=50, description="EDI Transaction Type (e.g., '204', '990', or '*')" ) @@ -279,6 +310,7 @@ class CertificateExportResponse(BaseModel): class AS2PartnershipResponse(BaseModel): id: str tenant_id: int | None + trading_partner_id: str | None = None name: str | None = None local_partner_id: str remote_partner_id: str @@ -299,10 +331,17 @@ class RouteResponse(BaseModel): class RouteItemResponse(BaseModel): route_id: UUID + trading_partner_id: str | None = None name: str direction: str isa_sender_id: str + isa_sender_qualifier: str | None = None isa_receiver_id: str + isa_receiver_qualifier: str | None = None + gs_sender_id: str | None = None + gs_receiver_id: str | None = None + default_standard: str | None = None + default_version: str | None = None transaction_type: str destination_type: str destination_name: str @@ -312,3 +351,71 @@ class RouteItemResponse(BaseModel): status: str = Field(default="ACTIVE") active: bool = False processing_mode: str = "TRANSLATE" + + +class OutboundMessageRequest(BaseModel): + trading_partner_id: str = Field( + ..., description="The ERP's identifier for the routing rule (trading_partner_id)" + ) + payload: dict[str, Any] | list[dict[str, Any]] = Field( + ..., description="The JSON payload representing the EDI document(s)" + ) + transaction_type: str | None = Field( + None, max_length=50, description="The EDI transaction type, e.g., '204', '850'" + ) + + +class OutboundMessageResponse(BaseModel): + trace_id: UUID = Field(..., description="The Trace ID to track the message lifecycle") + status: str = Field(default="ACCEPTED") + + +# --------------------------------------------------------------------------- +# API Token DTOs +# --------------------------------------------------------------------------- + + +class CreateApiTokenRequest(BaseModel): + name: str = Field( + ..., + min_length=1, + max_length=255, + description="Human-readable label for this token (e.g. 'ERP Integration Prod')", + ) + expires_at: datetime | None = Field( + None, description="Optional ISO-8601 expiry datetime. Null = never expires." + ) + + +class ApiTokenCreatedResponse(BaseModel): + """ + Returned exactly once upon token creation. + client_secret is shown here and NEVER returned again. + """ + + id: UUID + name: str + client_id: str = Field( + ..., description="Plaintext client identifier — safe to display in UI and logs" + ) + client_secret: str = Field( + ..., description="Raw client secret — store immediately, shown ONCE only" + ) + active: bool + created_at: str + + +class ApiTokenListItem(BaseModel): + """Safe list representation — secret is never included.""" + + id: UUID + name: str + client_id: str + active: bool + last_used_at: str | None + expires_at: str | None + created_at: str + + +class ApiTokenListResponse(BaseModel): + tokens: list[ApiTokenListItem] diff --git a/services/api/src/api/adapters/repository.py b/services/api/src/api/adapters/repository.py index 1b569d60..ba35f644 100644 --- a/services/api/src/api/adapters/repository.py +++ b/services/api/src/api/adapters/repository.py @@ -1,5 +1,6 @@ import uuid from collections.abc import Sequence +from datetime import UTC from typing import Any from uuid import UUID @@ -25,6 +26,7 @@ ) from database.encryption import db_encryption from database.models.control_plane import ( + ApiToken, AS2Partner, AS2Partnership, InboundRoute, @@ -35,7 +37,7 @@ ) from database.models.control_plane import Outbox as GlobalOutbox from database.models.data_plane import EdiMessage -from sqlalchemy import delete, select +from sqlalchemy import delete, or_, select from sqlalchemy.ext.asyncio import AsyncSession @@ -72,7 +74,7 @@ async def create_as2_identity(self, tenant_id: int, cmd: CreateAS2TradingPartner async def update_as2_identity( self, tenant_id: int, partner_id: UUID, cmd: UpdateAS2TradingPartnerCmd ) -> None: - partner = await self.get_as2_partner(tenant_id, partner_id) + partner = await self.get_as2_partner_for_write(tenant_id, partner_id) if partner: if cmd.name is not None: partner.name = cmd.name @@ -92,7 +94,36 @@ async def update_as2_identity( partner.active = cmd.active await self.session.flush() + async def rotate_as2_certificates( + self, + tenant_id: int, + partner_id: UUID, + new_public_cert: str, + new_private_key_vault_ref: str | None, + ) -> None: + partner = await self.get_as2_partner_for_write(tenant_id, partner_id) + if not partner: + raise ValueError(f"AS2 Partner {partner_id} not found or access denied.") + + partner.prev_public_cert_pem = partner.public_cert_pem + partner.prev_private_key_vault_ref = partner.private_key_vault_ref + + partner.public_cert_pem = new_public_cert + if new_private_key_vault_ref is not None: + partner.private_key_vault_ref = new_private_key_vault_ref + + await self.session.flush() + async def get_as2_partner(self, tenant_id: int, partner_id: UUID) -> Any: + result = await self.session.execute( + select(AS2Partner).where( + AS2Partner.id == partner_id, + or_(AS2Partner.tenant_id == tenant_id, AS2Partner.tenant_id.is_(None)), + ) + ) + return result.scalar_one_or_none() + + async def get_as2_partner_for_write(self, tenant_id: int, partner_id: UUID) -> Any: result = await self.session.execute( select(AS2Partner).where( AS2Partner.id == partner_id, @@ -114,6 +145,13 @@ async def delete_as2_identity(self, tenant_id: int, partner_id: UUID) -> None: await self.session.flush() async def create_as2_partnership(self, tenant_id: int, cmd: CreateAS2PartnershipCmd) -> UUID: + + local_partner = await self.get_as2_partner(tenant_id, cmd.local_partner_id) + remote_partner = await self.get_as2_partner(tenant_id, cmd.remote_partner_id) + + if not local_partner or not remote_partner: + raise ValueError("Local or Remote partner not found") + partnership_id = uuid.uuid4() record = AS2Partnership( id=partnership_id, @@ -393,6 +431,8 @@ async def create_inbound_route(self, tenant_id: int, cmd: CreateInboundRouteCmd) name=cmd.name, isa_sender_id=cmd.isa_sender_id, isa_receiver_id=cmd.isa_receiver_id, + gs_sender_id=cmd.gs_sender_id, + gs_receiver_id=cmd.gs_receiver_id, transaction_type=cmd.transaction_type, webhook_id=cmd.webhook_id, as2_partner_id=cmd.as2_partner_id, @@ -420,6 +460,10 @@ async def update_inbound_route( record.isa_sender_id = cmd.isa_sender_id if not isinstance(cmd.isa_receiver_id, UnsetType): record.isa_receiver_id = cmd.isa_receiver_id + if not isinstance(cmd.gs_sender_id, UnsetType): + record.gs_sender_id = cmd.gs_sender_id + if not isinstance(cmd.gs_receiver_id, UnsetType): + record.gs_receiver_id = cmd.gs_receiver_id if not isinstance(cmd.transaction_type, UnsetType): record.transaction_type = cmd.transaction_type if not isinstance(cmd.processing_mode, UnsetType): @@ -469,6 +513,20 @@ async def update_inbound_route( await self.session.flush() return True + async def get_inbound_route( + self, + isa_sender_id: str, + isa_receiver_id: str, + tenant_id: int, + transaction_type: str | None = None, + ) -> Any | None: + from database.repository import InboundRouteRepository + + repo = InboundRouteRepository(self.session) + return await repo.get_inbound_route( + isa_sender_id, isa_receiver_id, tenant_id, transaction_type + ) + async def delete_inbound_route(self, tenant_id: int, route_id: UUID) -> bool: result = await self.session.execute( delete(InboundRoute).where( @@ -510,10 +568,17 @@ async def create_outbound_route(self, tenant_id: int, cmd: CreateOutboundRouteCm record = OutboundRoute( id=route_id, tenant_id=tenant_id, + trading_partner_id=cmd.trading_partner_id, name=cmd.name, isa_sender_id=cmd.isa_sender_id, + isa_sender_qualifier=cmd.isa_sender_qualifier, isa_receiver_id=cmd.isa_receiver_id, + isa_receiver_qualifier=cmd.isa_receiver_qualifier, + gs_sender_id=cmd.gs_sender_id, + gs_receiver_id=cmd.gs_receiver_id, transaction_type=cmd.transaction_type, + default_standard=cmd.default_standard, + default_version=cmd.default_version, as2_partner_id=cmd.as2_partner_id, sftp_partner_id=cmd.sftp_partner_id, processing_mode=cmd.processing_mode, @@ -533,14 +598,28 @@ async def update_outbound_route( record = result.scalar_one_or_none() if not record: return False + if cmd.trading_partner_id is not UNSET: + record.trading_partner_id = cmd.trading_partner_id if not isinstance(cmd.name, UnsetType): record.name = cmd.name if not isinstance(cmd.isa_sender_id, UnsetType): record.isa_sender_id = cmd.isa_sender_id + if not isinstance(cmd.isa_sender_qualifier, UnsetType): + record.isa_sender_qualifier = cmd.isa_sender_qualifier if not isinstance(cmd.isa_receiver_id, UnsetType): record.isa_receiver_id = cmd.isa_receiver_id + if not isinstance(cmd.isa_receiver_qualifier, UnsetType): + record.isa_receiver_qualifier = cmd.isa_receiver_qualifier + if not isinstance(cmd.gs_sender_id, UnsetType): + record.gs_sender_id = cmd.gs_sender_id + if not isinstance(cmd.gs_receiver_id, UnsetType): + record.gs_receiver_id = cmd.gs_receiver_id if not isinstance(cmd.transaction_type, UnsetType): record.transaction_type = cmd.transaction_type + if not isinstance(cmd.default_standard, UnsetType): + record.default_standard = cmd.default_standard + if not isinstance(cmd.default_version, UnsetType): + record.default_version = cmd.default_version if not isinstance(cmd.processing_mode, UnsetType): record.processing_mode = cmd.processing_mode if not isinstance(cmd.as2_partner_id, UnsetType): @@ -574,6 +653,18 @@ async def update_outbound_route( await self.session.flush() return True + async def get_outbound_route_by_trading_partner_id( + self, tenant_id: int, trading_partner_id: str + ) -> OutboundRoute | None: + result = await self.session.execute( + select(OutboundRoute).where( + OutboundRoute.tenant_id == tenant_id, + OutboundRoute.trading_partner_id == trading_partner_id, + OutboundRoute.active.is_(True), + ) + ) + return result.scalar_one_or_none() + async def delete_outbound_route(self, tenant_id: int, route_id: UUID) -> bool: result = await self.session.execute( delete(OutboundRoute).where( @@ -625,6 +716,41 @@ async def create_outbox_event( await self.session.flush() return event_id + async def get_as2_partnership(self, tenant_id: int, partnership_id: UUID) -> Any: + from database.models.data_plane import AS2Partnership + + result = await self.session.execute( + select(AS2Partnership).where( + AS2Partnership.id == partnership_id, + AS2Partnership.tenant_id == tenant_id, + ) + ) + return result.scalar_one_or_none() + + async def get_as2_partner(self, tenant_id: int, partner_id: UUID) -> Any: + from database.models.data_plane import AS2Partner + + result = await self.session.execute( + select(AS2Partner).where(AS2Partner.id == partner_id, AS2Partner.tenant_id == tenant_id) + ) + return result.scalar_one_or_none() + + async def create_edi_json(self, tenant_id: int, payload: dict[str, Any]) -> UUID: + from database.models.data_plane import EdiJson + + msg = EdiJson(tenant_id=tenant_id, **payload) + self.session.add(msg) + await self.session.flush() + return msg.id + + async def create_api_gateway(self, tenant_id: int, payload: dict[str, Any]) -> UUID: + from database.models.data_plane import ApiGateway + + log = ApiGateway(tenant_id=tenant_id, **payload) + self.session.add(log) + await self.session.flush() + return log.id + class SqlAlchemyTenantRepository(TenantRepositoryPort): def __init__(self, session: AsyncSession) -> None: @@ -636,3 +762,106 @@ async def get_tenant_flags(self, tenant_id: int) -> dict[str, Any] | None: if tenant: return {"allow_private_as2": tenant.allow_private_as2} return None + + +class SqlAlchemyApiTokenRepository: + """Repository for managing platform API tokens in the global (control plane) DB.""" + + def __init__(self, session: AsyncSession) -> None: + self.session = session + + async def create_api_token( + self, + tenant_id: int, + name: str, + client_id: str, + secret_hash: str, + expires_at: object | None = None, + ) -> UUID: + import uuid as uuid_module + + token_id = uuid_module.uuid4() + record = ApiToken( + id=token_id, + tenant_id=tenant_id, + name=name, + client_id=client_id, + secret_hash=secret_hash, + expires_at=expires_at, + active=True, + ) + self.session.add(record) + await self.session.flush() + return token_id + + async def list_api_tokens(self, tenant_id: int) -> list[dict[str, Any]]: + result = await self.session.execute( + select(ApiToken) + .where(ApiToken.tenant_id == tenant_id) + .order_by(ApiToken.created_at.desc()) + ) + tokens = result.scalars().all() + return [ + { + "id": str(t.id), + "name": t.name, + "client_id": t.client_id, # safe to return; secret_hash is never exposed + "active": t.active, + "last_used_at": t.last_used_at.isoformat() if t.last_used_at else None, + "expires_at": t.expires_at.isoformat() if t.expires_at else None, + "created_at": t.created_at.isoformat(), + } + for t in tokens + ] + + async def revoke_api_token(self, tenant_id: int, token_id: UUID) -> bool: + result = await self.session.execute( + select(ApiToken).where(ApiToken.id == token_id, ApiToken.tenant_id == tenant_id) + ) + record = result.scalar_one_or_none() + if not record: + return False + record.active = False + await self.session.flush() + return True + + async def delete_api_token(self, tenant_id: int, token_id: UUID) -> bool: + result = await self.session.execute( + delete(ApiToken).where(ApiToken.id == token_id, ApiToken.tenant_id == tenant_id) + ) + await self.session.flush() + return (getattr(result, "rowcount", 0) or 0) > 0 + + async def get_tenant_id_by_credentials(self, client_id: str, secret_hash: str) -> int | None: + """ + Two-step lookup (indexed client_id → hash check → tenant_id). + Step 1: Find row by client_id (plaintext index — O(1), no full scan). + Step 2: Verify secret_hash matches (prevents timing attacks via constant-time compare). + Also updates last_used_at. + """ + import hmac + from datetime import datetime, timedelta + + from sqlalchemy import or_, update + + now = datetime.now(UTC).replace(tzinfo=None) + result = await self.session.execute( + select(ApiToken).where( + ApiToken.client_id == client_id, + ApiToken.active.is_(True), + or_(ApiToken.expires_at.is_(None), ApiToken.expires_at > now), + ) + ) + record = result.scalar_one_or_none() + if not record: + return None + + # Constant-time comparison prevents timing-based secret enumeration + if not hmac.compare_digest(record.secret_hash, secret_hash): + return None + + if not record.last_used_at or record.last_used_at < (now - timedelta(hours=1)): + await self.session.execute( + update(ApiToken).where(ApiToken.id == record.id).values(last_used_at=now) + ) + return record.tenant_id diff --git a/services/api/src/api/adapters/vault.py b/services/api/src/api/adapters/vault.py index da4243b1..aa9b668d 100644 --- a/services/api/src/api/adapters/vault.py +++ b/services/api/src/api/adapters/vault.py @@ -100,3 +100,4 @@ def delete_secret(self, vault_ref: str) -> None: # Singleton instance vault = VaultAdapter() +# Trigger reload diff --git a/services/api/src/api/auth/__init__.py b/services/api/src/api/auth/__init__.py new file mode 100644 index 00000000..fc10542d --- /dev/null +++ b/services/api/src/api/auth/__init__.py @@ -0,0 +1 @@ +# API authentication modules diff --git a/services/api/src/api/auth/api_key.py b/services/api/src/api/auth/api_key.py new file mode 100644 index 00000000..dbbf77db --- /dev/null +++ b/services/api/src/api/auth/api_key.py @@ -0,0 +1,95 @@ +""" +FastAPI dependency for two-part API key (M2M) authentication. + +ERP systems authenticate using a Client ID + Client Secret pair (Stripe/AWS style): + X-Client-ID: soopaedi_acme_a3f12b9c (plaintext, safe to log) + X-Client-Secret: <43-char random secret> (hashed, never logged) + +Validation: + 1. Look up the ApiToken row by client_id (fast plaintext index — O(1)) + 2. SHA-256 hash the incoming secret + 3. Constant-time compare against stored secret_hash + 4. Return tenant_id + +No external network call to Zitadel or any IdP is made. +""" + +import hashlib +import hmac +import logging + +from database.session import get_global_session +from fastapi import Depends, HTTPException, Security, status +from fastapi.security import APIKeyHeader +from sqlalchemy.ext.asyncio import AsyncSession + +from api.adapters.repository import SqlAlchemyApiTokenRepository + +logger = logging.getLogger(__name__) + +_client_id_header = APIKeyHeader(name="X-Client-ID", auto_error=False) +_client_secret_header = APIKeyHeader(name="X-Client-Secret", auto_error=False) + +# In-process cache: { client_id → (tenant_id, secret_hash) } +# Short-circuits the DB lookup for repeated calls within the same process. +# Tokens are evicted on revocation via invalidate_token_cache(). +_token_cache: dict[str, tuple[int, str]] = {} +_MAX_CACHE_SIZE = 5000 + + +async def get_tenant_id_from_api_key( + client_id: str | None = Security(_client_id_header), + client_secret: str | None = Security(_client_secret_header), + global_session: AsyncSession = Depends(get_global_session), # noqa: B008 +) -> int: + """ + Resolves a two-part API credential to a tenant_id. + + Expected headers: + X-Client-ID: soopaedi_acme_a3f12b9c + X-Client-Secret: + + Raises HTTP 401 if either header is missing, the client_id doesn't exist, + or the secret doesn't match. + """ + if not client_id or not client_secret: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Missing credentials. Provide X-Client-ID and X-Client-Secret headers.", + headers={"WWW-Authenticate": "X-API-Key"}, + ) + + # Hash the secret before touching the DB (raw secret never persisted or logged) + secret_hash = hashlib.sha256(client_secret.encode("utf-8")).hexdigest() + + # Fast path: cache hit (only safe because we evict on revocation) + if client_id in _token_cache: + cached_tenant_id, cached_secret_hash = _token_cache[client_id] + if hmac.compare_digest(cached_secret_hash, secret_hash): + return cached_tenant_id + + repo = SqlAlchemyApiTokenRepository(global_session) + tenant_id = await repo.get_tenant_id_by_credentials(client_id, secret_hash) + + if tenant_id is None: + logger.warning(f"API key authentication failed for client_id={client_id!r}") + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Invalid or revoked credentials.", + headers={"WWW-Authenticate": "X-API-Key"}, + ) + + # Populate cache (bounded eviction) + if len(_token_cache) >= _MAX_CACHE_SIZE: + _token_cache.pop(next(iter(_token_cache))) + _token_cache[client_id] = (tenant_id, secret_hash) + + return tenant_id + + +def invalidate_token_cache(client_id: str) -> None: + """ + Remove a client_id from the in-process cache. + Must be called after revoking a token so cached entries don't persist. + """ + _token_cache.pop(client_id, None) diff --git a/services/api/src/api/cdc_relay.py b/services/api/src/api/cdc_relay.py index 3038e590..e989b19d 100644 --- a/services/api/src/api/cdc_relay.py +++ b/services/api/src/api/cdc_relay.py @@ -1,8 +1,9 @@ import logging from typing import Any -from fastapi import APIRouter, Depends, HTTPException -from pydantic import BaseModel, Field +from domain.events import MessageQueueName +from fastapi import APIRouter, Depends, Request +from pydantic import BaseModel, Field, ValidationError from api.dependencies import get_message_queue from api.ports.message_queue import MessageQueuePort @@ -20,60 +21,120 @@ class DebeziumUnwrappedEvent(BaseModel): op: str = Field(alias="__op", description="Operation type: 'c' for create, 'u' for update") table: str = Field(alias="__table", description="The source table name") - # Columns from the outbox table - idempotency_key: str - event_type: str - payload: dict[str, Any] - status: str - tenant_id: int + # Core Outbox fields + idempotency_key: str | None = None + event_type: str | None = None + payload: dict[str, Any] | str | None = None + tenant_id: int | None = None + class Config: + extra = "ignore" # Debezium sends many extra metadata fields we don't need -@router.post("/relay", status_code=202) + +@router.api_route("/relay", methods=["GET", "POST", "PUT", "PATCH", "DELETE"], status_code=200) async def relay_cdc_event( - event: DebeziumUnwrappedEvent, + request: Request, queue: MessageQueuePort = Depends(get_message_queue), -) -> None: +) -> dict[str, str]: """ Receives HTTP Webhooks from the standalone Debezium Server and manually routes them into AWS SQS queues based on the source table and event_type. + + Enterprise Grade Notes: + - Bypasses default FastAPI strict typing in the method signature to handle Debezium's + inconsistent Content-Type headers, but strictly validates the JSON via Pydantic internally. + - Handles both batched (List) and single (Object) payloads. + - Implements the Robustness Principle (Postel's Law) for internal infrastructure webhooks. """ - # Only process INSERTS for the outbox - if event.op != "c": - return - - if event.table == "outbox": - if event.event_type not in ("TRANSLATE", "DELIVER"): - logger.warning(f"[CDC Relay] Unhandled outbox event type: {event.event_type}") - raise HTTPException( - status_code=400, detail=f"Unknown outbox event_type: {event.event_type}" - ) + import json - queue_name = "TranslateQueue" if event.event_type == "TRANSLATE" else "DeliverQueue" + body = await request.body() + try: + data = json.loads(body) + except Exception as e: + logger.error(f"[CDC Relay] CRITICAL: Failed to parse raw CDC bytes as JSON: {e}") + # In a full enterprise setup, this raw body would be pushed to an S3 Dead Letter bucket here. + return {"status": "ok"} - # Validate that the payload contains a trace_id required by workers - trace_id = event.payload.get("trace_id") - if not trace_id: + # Normalize to a list + raw_events = data if isinstance(data, list) else [data] + + # Strictly validate against our schema + validated_events: list[DebeziumUnwrappedEvent] = [] + for raw_event in raw_events: + try: + validated_events.append(DebeziumUnwrappedEvent(**raw_event)) + except (ValidationError, TypeError) as e: logger.error( - f"[CDC Relay] Outbox event missing trace_id in payload: {event.idempotency_key}" - ) - raise HTTPException( - status_code=400, - detail="Outbox payload missing required trace_id field", + f"[CDC Relay] Schema validation failed for event: {e}. Payload: {raw_event}" ) + try: + await queue.send(MessageQueueName.CDC_DLQ, {"error": str(e), "payload": raw_event}) + except Exception as dlq_err: + logger.error(f"[CDC Relay] Failed to write to CDC DLQ: {dlq_err}") + from fastapi import HTTPException + + raise HTTPException( + status_code=500, detail="Failed to quarantine invalid event" + ) from dlq_err + # Quarantine succeeded, so skip invalid events; do not fail the entire batch. + continue + + for event in validated_events: + if event.op != "c": + continue + + if event.table == "outbox": + # Outbox table routing logic + if not event.event_type: + logger.warning( + f"[CDC Relay] Outbox event missing event_type. Skipping: {event.idempotency_key}" + ) + continue + + event_payload = event.payload + if isinstance(event_payload, str): + try: + payload_dict = json.loads(event_payload) + except json.JSONDecodeError as e: + logger.error( + f"[CDC Relay] Invalid JSON in outbox payload for key {event.idempotency_key}: {e}" + ) + continue + else: + payload_dict = event_payload + + if event.event_type in ( + "TRANSLATE", + "json.received", + "edi_message.received", + "DELIVER", + ): + queue_name = ( + MessageQueueName.TRANSLATE + if event.event_type in ("TRANSLATE", "json.received", "edi_message.received") + else MessageQueueName.DELIVER + ) + + # Validate that the payload contains a trace_id required by data plane workers + trace_id = payload_dict.get("trace_id") if isinstance(payload_dict, dict) else None + if not trace_id: + logger.error( + f"[CDC Relay] Outbox event missing trace_id in payload: {event.idempotency_key}" + ) + continue + else: + queue_name = MessageQueueName.PROVISIONING + + message_body = { + "idempotency_key": event.idempotency_key, + "event_type": event.event_type, + "payload": payload_dict, + "tenant_id": event.tenant_id, + } - # We package the original outbox payload and idempotency key for the worker - message_body = { - "idempotency_key": event.idempotency_key, - "event_type": event.event_type, - "payload": event.payload, - "tenant_id": event.tenant_id, - } - - await queue.send( - queue_name=queue_name, - payload=message_body, - ) - logger.info(f"[CDC Relay] Relayed event_type={event.event_type} to {queue_name}") - else: - logger.warning(f"[CDC Relay] Received event for unknown table: {event.table}") - raise HTTPException(status_code=400, detail="Unknown table source") + await queue.send(queue_name=queue_name, payload=message_body) + logger.info(f"[CDC Relay] Relayed event_type={event.event_type} to {queue_name}") + else: + logger.debug(f"[CDC Relay] Ignoring event for unhandled table: {event.table}") + return {"status": "ok"} diff --git a/services/api/src/api/core/authorization.py b/services/api/src/api/core/authorization.py index 974398e1..73667015 100644 --- a/services/api/src/api/core/authorization.py +++ b/services/api/src/api/core/authorization.py @@ -41,6 +41,7 @@ async def get_authorization_profile( "users:delete", "routes:manage", "certificates:export_private", + "certificates:rotate", ] ) else: diff --git a/services/api/src/api/core/provisioning.py b/services/api/src/api/core/provisioning.py deleted file mode 100644 index 20996718..00000000 --- a/services/api/src/api/core/provisioning.py +++ /dev/null @@ -1,402 +0,0 @@ -import logging -from typing import Any -from uuid import UUID - -from api.domain.models import ( - CreateAS2PartnershipCmd, - CreateAS2TradingPartnerCmd, - CreateInboundRouteCmd, - CreateOutboundRouteCmd, - CreateSFTPPartnerCmd, - CreateWebhookCmd, - PartnerEntity, - RouteEntity, - UpdateAS2PartnershipCmd, - UpdateAS2TradingPartnerCmd, - UpdateInboundRouteCmd, - UpdateOutboundRouteCmd, - UpdateSFTPPartnerCmd, -) -from api.ports.repository import ControlPlaneRepositoryPort, DataPlaneRepositoryPort - -logger = logging.getLogger(__name__) - - -class ProvisioningService: - """ - Core application service for orchestrating the Provisioning of Trading Partners. - Decoupled from FastAPI, SQLAlchemy, and AWS. - """ - - def __init__( - self, - global_repo: ControlPlaneRepositoryPort, - tenant_repo: DataPlaneRepositoryPort | None = None, - ) -> None: - self.global_repo = global_repo - self.tenant_repo = tenant_repo - - async def create_as2_partner( - self, tenant_id: int, cmd: CreateAS2TradingPartnerCmd - ) -> PartnerEntity: - if not self.global_repo: - raise ValueError("Control plane repository is required for AS2 partner creation") - - logger.info(f"Provisioning AS2 partner {cmd.name} for tenant {tenant_id}") - - # 1. Create in Global DB - partner_id = await self.global_repo.create_as2_identity(tenant_id=tenant_id, cmd=cmd) - - # 2. Emit Outbox Event (Worker will provision certificates/vault if needed) - payload = { - "partner_id": str(partner_id), - "tenant_id": tenant_id, - } - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="AS2_PARTNER_CREATED", - payload=payload, - ) - - return PartnerEntity( - partner_id=partner_id, - tenant_id=tenant_id, - name=cmd.name, - type="AS2", - status="PROVISIONING", - ) - - async def delete_as2_partner(self, tenant_id: int, partner_id: UUID) -> None: - if not self.global_repo: - raise ValueError("Control plane repository is required for AS2 partner deletion") - logger.info(f"Deleting AS2 partner {partner_id} for tenant {tenant_id}") - await self.global_repo.delete_as2_identity(tenant_id, partner_id) - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="AS2_PARTNER_DELETED", - payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, - ) - - async def update_as2_partner( - self, tenant_id: int, partner_id: UUID, cmd: UpdateAS2TradingPartnerCmd - ) -> PartnerEntity: - if not self.global_repo: - raise ValueError("Control plane repository is required for AS2 partner updates") - logger.info(f"Updating AS2 partner {partner_id} for tenant {tenant_id}") - await self.global_repo.update_as2_identity(tenant_id, partner_id, cmd) - - updated_partner = await self.global_repo.get_as2_partner(tenant_id, partner_id) - if not updated_partner: - raise ValueError("Partner not found after update") - - return PartnerEntity( - partner_id=partner_id, - tenant_id=tenant_id, - name=cmd.name or updated_partner.name, - type="AS2", - status="ACTIVE" if updated_partner.active else "INACTIVE", - ) - - async def create_sftp_partner(self, tenant_id: int, cmd: CreateSFTPPartnerCmd) -> PartnerEntity: - logger.info(f"Creating SFTP partner {cmd.name} for tenant {tenant_id}") - - if not self.global_repo: - raise ValueError("Control plane repository is required") - partner_id = await self.global_repo.create_sftp_partner(tenant_id=tenant_id, cmd=cmd) - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="SFTP_PARTNER_CREATED", - payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, - ) - - return PartnerEntity( - partner_id=partner_id, - tenant_id=tenant_id, - name=cmd.name, - type="SFTP", - status="INACTIVE", - ) - - async def update_sftp_partner( - self, tenant_id: int, partner_id: UUID, cmd: UpdateSFTPPartnerCmd - ) -> PartnerEntity: - logger.info(f"Updating SFTP partner {partner_id} for tenant {tenant_id}") - await self.global_repo.update_sftp_partner( - tenant_id=tenant_id, partner_id=partner_id, cmd=cmd - ) - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="SFTP_PARTNER_UPDATED", - payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, - ) - updated = await self.global_repo.get_sftp_partner(tenant_id, partner_id) - if not updated: - raise ValueError(f"SFTP partner {partner_id} not found") - - return PartnerEntity( - partner_id=partner_id, - tenant_id=tenant_id, - name=cmd.name or updated.name, - type="SFTP", - status="ACTIVE" if updated.active else "INACTIVE", - ) - - async def create_as2_partnership( - self, tenant_id: int, cmd: CreateAS2PartnershipCmd - ) -> PartnerEntity: - if not self.global_repo: - raise ValueError("Control plane repository is required for AS2 partnership creation") - - local_partner = await self.global_repo.get_as2_partner(tenant_id, cmd.local_partner_id) - if not local_partner: - raise ValueError(f"Local AS2 partner {cmd.local_partner_id} not found") - - remote_partner = await self.global_repo.get_as2_partner(tenant_id, cmd.remote_partner_id) - if not remote_partner: - raise ValueError(f"Remote AS2 partner {cmd.remote_partner_id} not found") - - logger.info( - f"Provisioning AS2 partnership {cmd.local_partner_id} -> {cmd.remote_partner_id}" - ) - partner_id = await self.global_repo.create_as2_partnership(tenant_id=tenant_id, cmd=cmd) - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="AS2_PARTNERSHIP_CREATED", - payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, - ) - - return PartnerEntity( - partner_id=partner_id, - tenant_id=tenant_id, - name=cmd.name, - type="AS2_PARTNERSHIP", - status="INACTIVE", - ) - - async def update_as2_partnership( - self, tenant_id: int, partnership_id: UUID, cmd: UpdateAS2PartnershipCmd - ) -> PartnerEntity: - if not self.global_repo: - raise ValueError("Control plane repository is required for AS2 partnership update") - - check_ids: list[UUID] = [] - if isinstance(cmd.local_partner_id, UUID): - check_ids.append(cmd.local_partner_id) - if isinstance(cmd.remote_partner_id, UUID): - check_ids.append(cmd.remote_partner_id) - - if check_ids: - valid_partners = await self.global_repo.get_as2_partners_by_ids(tenant_id, check_ids) - if len(valid_partners) != len(check_ids): - raise ValueError( - "Invalid local_partner_id or remote_partner_id referenced in update" - ) - - logger.info(f"Updating AS2 partnership {partnership_id}") - await self.global_repo.update_as2_partnership( - tenant_id=tenant_id, partnership_id=partnership_id, cmd=cmd - ) - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="AS2_PARTNERSHIP_UPDATED", - payload={"partner_id": str(partnership_id), "tenant_id": tenant_id}, - ) - updated = await self.global_repo.get_as2_partnership(tenant_id, partnership_id) - if not updated: - raise ValueError(f"AS2 partnership {partnership_id} not found") - - return PartnerEntity( - partner_id=partnership_id, - tenant_id=tenant_id, - name=updated.name, - type="AS2_PARTNERSHIP", - status="ACTIVE" if updated.active else "INACTIVE", - ) - - async def create_webhook(self, tenant_id: int, cmd: CreateWebhookCmd) -> PartnerEntity: - logger.info(f"Creating Webhook partner {cmd.name} for tenant {tenant_id}") - - if not self.global_repo: - raise ValueError("Control plane repository is required") - partner_id = await self.global_repo.create_webhook(tenant_id=tenant_id, cmd=cmd) - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="WEBHOOK_CREATED", - payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, - ) - - return PartnerEntity( - partner_id=partner_id, - tenant_id=tenant_id, - name=cmd.name, - type="WEBHOOK", - status="ACTIVE", - ) - - async def create_inbound_route(self, tenant_id: int, cmd: CreateInboundRouteCmd) -> RouteEntity: - logger.info(f"Creating Inbound Route for sender {cmd.isa_sender_id} in tenant {tenant_id}") - route_id = await self.global_repo.create_inbound_route(tenant_id=tenant_id, cmd=cmd) - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="INBOUND_ROUTE_CREATED", - payload={"route_id": str(route_id), "tenant_id": tenant_id}, - ) - - return RouteEntity( - route_id=route_id, - tenant_id=tenant_id, - direction="INBOUND", - ) - - async def create_outbound_route( - self, tenant_id: int, cmd: CreateOutboundRouteCmd - ) -> RouteEntity: - logger.info(f"Creating Outbound Route for sender {cmd.isa_sender_id} in tenant {tenant_id}") - route_id = await self.global_repo.create_outbound_route(tenant_id=tenant_id, cmd=cmd) - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="OUTBOUND_ROUTE_CREATED", - payload={"route_id": str(route_id), "tenant_id": tenant_id}, - ) - - return RouteEntity( - route_id=route_id, - tenant_id=tenant_id, - direction="OUTBOUND", - ) - - async def update_inbound_route( - self, tenant_id: int, route_id: UUID, cmd: UpdateInboundRouteCmd - ) -> bool: - res = await self.global_repo.update_inbound_route(tenant_id, route_id, cmd) - if res: - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="INBOUND_ROUTE_UPDATED", - payload={"route_id": str(route_id), "tenant_id": tenant_id}, - ) - return res - - async def delete_inbound_route(self, tenant_id: int, route_id: UUID) -> bool: - res = await self.global_repo.delete_inbound_route(tenant_id, route_id) - if res: - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="INBOUND_ROUTE_DELETED", - payload={"route_id": str(route_id), "tenant_id": tenant_id}, - ) - return res - - async def update_outbound_route( - self, tenant_id: int, route_id: UUID, cmd: UpdateOutboundRouteCmd - ) -> bool: - res = await self.global_repo.update_outbound_route(tenant_id, route_id, cmd) - if res: - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="OUTBOUND_ROUTE_UPDATED", - payload={"route_id": str(route_id), "tenant_id": tenant_id}, - ) - return res - - async def delete_outbound_route(self, tenant_id: int, route_id: UUID) -> bool: - res = await self.global_repo.delete_outbound_route(tenant_id, route_id) - if res: - await self.global_repo.create_outbox_event( - tenant_id=tenant_id, - event_type="OUTBOUND_ROUTE_DELETED", - payload={"route_id": str(route_id), "tenant_id": tenant_id}, - ) - return res - - async def list_routes(self, tenant_id: int) -> list[dict[str, Any]]: - routes_data = await self.global_repo.get_all_routes(tenant_id) - inbound = routes_data.get("inbound", []) - outbound = routes_data.get("outbound", []) - - # Collect IDs to fetch names - as2_ids = set() - sftp_ids = set() - webhook_ids = set() - - for r in inbound: - if r.as2_partner_id: - as2_ids.add(r.as2_partner_id) - if r.sftp_partner_id: - sftp_ids.add(r.sftp_partner_id) - if r.webhook_id: - webhook_ids.add(r.webhook_id) - - for r in outbound: - # Note: outbound routes map to as2_partner_id - if r.as2_partner_id: - as2_ids.add(r.as2_partner_id) - if r.sftp_partner_id: - sftp_ids.add(r.sftp_partner_id) - - as2_names = ( - await self.global_repo.get_as2_partners_by_ids(tenant_id, list(as2_ids)) - if self.global_repo - else {} - ) - sftp_names = await self.global_repo.get_sftp_partners_by_ids(tenant_id, list(sftp_ids)) - webhook_names = await self.global_repo.get_webhooks_by_ids(tenant_id, list(webhook_ids)) - - results = [] - for r in inbound: - dest_type = "UNKNOWN" - dest_name = "Unknown" - if r.as2_partner_id: - dest_type = "AS2" - dest_name = as2_names.get(r.as2_partner_id, str(r.as2_partner_id)) - elif r.sftp_partner_id: - dest_type = "SFTP" - dest_name = sftp_names.get(r.sftp_partner_id, str(r.sftp_partner_id)) - elif r.webhook_id: - dest_type = "WEBHOOK" - dest_name = webhook_names.get(r.webhook_id, str(r.webhook_id)) - - results.append( - { - "route_id": r.id, - "name": r.name, - "direction": "INBOUND", - "isa_sender_id": r.isa_sender_id, - "isa_receiver_id": r.isa_receiver_id, - "transaction_type": r.transaction_type, - "destination_type": dest_type, - "destination_name": dest_name, - "webhook_id": r.webhook_id, - "as2_partner_id": r.as2_partner_id, - "sftp_partner_id": r.sftp_partner_id, - "active": r.active, - } - ) - - for r in outbound: - dest_type = "UNKNOWN" - dest_name = "Unknown" - if r.as2_partner_id: - dest_type = "AS2" - dest_name = as2_names.get(r.as2_partner_id, str(r.as2_partner_id)) - elif r.sftp_partner_id: - dest_type = "SFTP" - dest_name = sftp_names.get(r.sftp_partner_id, str(r.sftp_partner_id)) - - results.append( - { - "route_id": r.id, - "name": r.name, - "direction": "OUTBOUND", - "isa_sender_id": r.isa_sender_id, - "isa_receiver_id": r.isa_receiver_id, - "transaction_type": r.transaction_type, - "destination_type": dest_type, - "destination_name": dest_name, - "as2_partner_id": r.as2_partner_id, - "sftp_partner_id": r.sftp_partner_id, - "active": r.active, - } - ) - - return results diff --git a/services/api/src/api/core/services/__init__.py b/services/api/src/api/core/services/__init__.py new file mode 100644 index 00000000..148aa30d --- /dev/null +++ b/services/api/src/api/core/services/__init__.py @@ -0,0 +1,13 @@ +from api.core.services.as2_partner_service import AS2PartnerService +from api.core.services.as2_partnership_service import AS2PartnershipService +from api.core.services.route_service import RouteService +from api.core.services.sftp_partner_service import SFTPPartnerService +from api.core.services.webhook_service import WebhookService + +__all__ = [ + "AS2PartnerService", + "AS2PartnershipService", + "SFTPPartnerService", + "WebhookService", + "RouteService", +] diff --git a/services/api/src/api/core/services/api_token_service.py b/services/api/src/api/core/services/api_token_service.py new file mode 100644 index 00000000..8a8ed037 --- /dev/null +++ b/services/api/src/api/core/services/api_token_service.py @@ -0,0 +1,111 @@ +""" +Core service for managing platform API tokens (M2M authentication). + +Follows Hexagonal Architecture: + - Depends on ApiTokenRepositoryPort (port), never on SQLAlchemy. + - Pure Python: testable without a DB or framework. + +Two-part credential pattern (Stripe/AWS style): + - client_id: stored plaintext, visible in UI, used for fast indexed lookup. + - client_secret: only SHA-256 hash stored; raw value shown once and discarded. +""" + +import hashlib +import logging +import secrets +from typing import Any +from uuid import UUID + +from api.domain.models import ApiTokenEntity, CreateApiTokenCmd +from api.ports.repository import ApiTokenRepositoryPort + +logger = logging.getLogger(__name__) + +_TOKEN_VENDOR = "soopaedi" + + +def _generate_credentials(tenant_name: str) -> tuple[str, str, str]: + """ + Pure function — no I/O, deterministically testable. + + Returns (client_id, client_secret, secret_hash). + + client_id format: soopaedi_<6-char-slug>_<8-hex-chars> + e.g. soopaedi_acmeco_a3f12b9c + client_secret format: 43-char URL-safe random string (32 random bytes) + + Only secret_hash is stored. client_id is stored in plaintext. + """ + slug = "".join(c for c in tenant_name.lower() if c.isalnum())[:6] + random_suffix = secrets.token_hex(4) # 8 hex chars + client_id = f"{_TOKEN_VENDOR}_{slug}_{random_suffix}" + + client_secret = secrets.token_urlsafe(32) # 43 URL-safe chars + secret_hash = hashlib.sha256(client_secret.encode("utf-8")).hexdigest() + + return client_id, client_secret, secret_hash + + +class ApiTokenService: + """ + Application service responsible for the lifecycle of tenant API tokens. + + Constructor receives ApiTokenRepositoryPort — a pure interface. + No framework, no DB, no network dependency at construction time. + """ + + def __init__(self, repo: ApiTokenRepositoryPort) -> None: + self._repo = repo + + async def create_token( + self, tenant_id: int, tenant_name: str, cmd: CreateApiTokenCmd + ) -> ApiTokenEntity: + """ + Generates a two-part API credential for the given tenant. + client_secret is returned exactly once and is NOT stored. + """ + client_id, client_secret, secret_hash = _generate_credentials(tenant_name) + + token_id = await self._repo.create_api_token( + tenant_id=tenant_id, + name=cmd.name, + client_id=client_id, + secret_hash=secret_hash, + expires_at=cmd.expires_at, + ) + + logger.info( + "API token created", + extra={"tenant_id": tenant_id, "token_name": cmd.name, "client_id": client_id}, + ) + + return ApiTokenEntity( + id=token_id, + tenant_id=tenant_id, + name=cmd.name, + client_id=client_id, + client_secret=client_secret, # caller must show this exactly once + active=True, + ) + + async def list_tokens(self, tenant_id: int) -> list[dict[str, Any]]: + """Returns all tokens for a tenant. client_id is safe; secret is never returned.""" + return await self._repo.list_api_tokens(tenant_id) + + async def revoke_token(self, tenant_id: int, token_id: UUID) -> bool: + """Soft-deactivates a token (active=False). Audit record preserved.""" + result = await self._repo.revoke_api_token(tenant_id, token_id) + if result: + logger.info( + "API token revoked", extra={"tenant_id": tenant_id, "token_id": str(token_id)} + ) + return result + + async def delete_token(self, tenant_id: int, token_id: UUID) -> bool: + """Hard deletes a token record. Irreversible.""" + result = await self._repo.delete_api_token(tenant_id, token_id) + if result: + logger.info( + "API token deleted", extra={"tenant_id": tenant_id, "token_id": str(token_id)} + ) + return result diff --git a/services/api/src/api/core/services/as2_partner_service.py b/services/api/src/api/core/services/as2_partner_service.py new file mode 100644 index 00000000..6a7c1657 --- /dev/null +++ b/services/api/src/api/core/services/as2_partner_service.py @@ -0,0 +1,105 @@ +import logging +from uuid import UUID + +from api.domain.models import ( + CreateAS2TradingPartnerCmd, + PartnerEntity, + UpdateAS2TradingPartnerCmd, +) +from api.ports.repository import ControlPlaneRepositoryPort +from domain.events import ProvisioningEventType + +logger = logging.getLogger(__name__) + + +class AS2PartnerService: + """ + Domain service responsible for the lifecycle of AS2 Trading Partners. + Operates exclusively on the Global Control Plane repository. + """ + + def __init__(self, global_repo: ControlPlaneRepositoryPort) -> None: + self.global_repo = global_repo + + async def create_as2_partner( + self, tenant_id: int, cmd: CreateAS2TradingPartnerCmd + ) -> PartnerEntity: + logger.info(f"Provisioning AS2 partner {cmd.name} for tenant {tenant_id}") + + partner_id = await self.global_repo.create_as2_identity(tenant_id=tenant_id, cmd=cmd) + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.AS2_PARTNER_CREATED, + payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, + ) + + return PartnerEntity( + partner_id=partner_id, + tenant_id=tenant_id, + name=cmd.name, + type="AS2", + status="PROVISIONING", + ) + + async def update_as2_partner( + self, tenant_id: int, partner_id: UUID, cmd: UpdateAS2TradingPartnerCmd + ) -> PartnerEntity: + logger.info(f"Updating AS2 partner {partner_id} for tenant {tenant_id}") + await self.global_repo.update_as2_identity(tenant_id, partner_id, cmd) + + updated_partner = await self.global_repo.get_as2_partner(tenant_id, partner_id) + if not updated_partner: + raise ValueError("Partner not found after update") + + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.AS2_PARTNER_UPDATED, + payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, + ) + + return PartnerEntity( + partner_id=partner_id, + tenant_id=tenant_id, + name=cmd.name or updated_partner.name, + type="AS2", + status="ACTIVE" if updated_partner.active else "INACTIVE", + ) + + async def delete_as2_partner(self, tenant_id: int, partner_id: UUID) -> None: + logger.info(f"Deleting AS2 partner {partner_id} for tenant {tenant_id}") + await self.global_repo.delete_as2_identity(tenant_id, partner_id) + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.AS2_PARTNER_DELETED, + payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, + ) + + async def rotate_certificates( + self, + tenant_id: int, + partner_id: UUID, + new_public_cert: str, + new_private_key_vault_ref: str | None, + ) -> PartnerEntity: + logger.info(f"Rotating certificates for AS2 partner {partner_id} for tenant {tenant_id}") + await self.global_repo.rotate_as2_certificates( + tenant_id, partner_id, new_public_cert, new_private_key_vault_ref + ) + + updated_partner = await self.global_repo.get_as2_partner(tenant_id, partner_id) + if not updated_partner: + raise ValueError("Partner not found after certificate rotation") + + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.AS2_PARTNER_UPDATED, + payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, + ) + + return PartnerEntity( + partner_id=partner_id, + tenant_id=tenant_id, + name=updated_partner.name, + type="AS2", + status="ACTIVE" if updated_partner.active else "INACTIVE", + ) diff --git a/services/api/src/api/core/services/as2_partnership_service.py b/services/api/src/api/core/services/as2_partnership_service.py new file mode 100644 index 00000000..b359dedc --- /dev/null +++ b/services/api/src/api/core/services/as2_partnership_service.py @@ -0,0 +1,97 @@ +import logging +from uuid import UUID + +from api.domain.models import ( + CreateAS2PartnershipCmd, + PartnerEntity, + UpdateAS2PartnershipCmd, +) +from api.ports.repository import ControlPlaneRepositoryPort +from domain.events import ProvisioningEventType + +logger = logging.getLogger(__name__) + + +class AS2PartnershipService: + """ + Domain service responsible for the lifecycle of AS2 Partnerships. + Validates that referenced local/remote partners exist before mutating state. + """ + + def __init__(self, global_repo: ControlPlaneRepositoryPort) -> None: + self.global_repo = global_repo + + async def create_as2_partnership( + self, tenant_id: int, cmd: CreateAS2PartnershipCmd + ) -> PartnerEntity: + local_partner = await self.global_repo.get_as2_partner(tenant_id, cmd.local_partner_id) + if not local_partner: + raise ValueError(f"Local AS2 partner {cmd.local_partner_id} not found") + + remote_partner = await self.global_repo.get_as2_partner(tenant_id, cmd.remote_partner_id) + if not remote_partner: + raise ValueError(f"Remote AS2 partner {cmd.remote_partner_id} not found") + + logger.info( + f"Provisioning AS2 partnership {cmd.local_partner_id} -> {cmd.remote_partner_id}" + ) + partner_id = await self.global_repo.create_as2_partnership(tenant_id=tenant_id, cmd=cmd) + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.AS2_PARTNERSHIP_CREATED, + payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, + ) + + return PartnerEntity( + partner_id=partner_id, + tenant_id=tenant_id, + name=cmd.name, + type="AS2_PARTNERSHIP", + status="INACTIVE", + ) + + async def update_as2_partnership( + self, tenant_id: int, partnership_id: UUID, cmd: UpdateAS2PartnershipCmd + ) -> PartnerEntity: + check_ids: list[UUID] = [] + if isinstance(cmd.local_partner_id, UUID): + check_ids.append(cmd.local_partner_id) + if isinstance(cmd.remote_partner_id, UUID): + check_ids.append(cmd.remote_partner_id) + + if check_ids: + valid_partners = await self.global_repo.get_as2_partners_by_ids(tenant_id, check_ids) + if len(valid_partners) != len(check_ids): + raise ValueError( + "Invalid local_partner_id or remote_partner_id referenced in update" + ) + + logger.info(f"Updating AS2 partnership {partnership_id}") + await self.global_repo.update_as2_partnership( + tenant_id=tenant_id, partnership_id=partnership_id, cmd=cmd + ) + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.AS2_PARTNERSHIP_UPDATED, + payload={"partner_id": str(partnership_id), "tenant_id": tenant_id}, + ) + updated = await self.global_repo.get_as2_partnership(tenant_id, partnership_id) + if not updated: + raise ValueError(f"AS2 partnership {partnership_id} not found") + + return PartnerEntity( + partner_id=partnership_id, + tenant_id=tenant_id, + name=updated.name, + type="AS2_PARTNERSHIP", + status="ACTIVE" if updated.active else "INACTIVE", + ) + + async def delete_as2_partnership(self, tenant_id: int, partnership_id: UUID) -> None: + logger.info(f"Deleting AS2 partnership {partnership_id} for tenant {tenant_id}") + await self.global_repo.delete_as2_partnership(tenant_id, partnership_id) + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.AS2_PARTNERSHIP_DELETED, + payload={"partner_id": str(partnership_id), "tenant_id": tenant_id}, + ) diff --git a/services/api/src/api/core/services/route_service.py b/services/api/src/api/core/services/route_service.py new file mode 100644 index 00000000..c7e0e5e5 --- /dev/null +++ b/services/api/src/api/core/services/route_service.py @@ -0,0 +1,191 @@ +import logging +from typing import Any +from uuid import UUID + +from api.domain.models import ( + CreateInboundRouteCmd, + CreateOutboundRouteCmd, + RouteEntity, + UpdateInboundRouteCmd, + UpdateOutboundRouteCmd, +) +from api.ports.repository import ControlPlaneRepositoryPort +from domain.events import ProvisioningEventType + +logger = logging.getLogger(__name__) + + +class RouteService: + """ + Domain service responsible for the lifecycle of Inbound and Outbound EDI Routes, + including resolution of partner names for list operations. + """ + + def __init__(self, global_repo: ControlPlaneRepositoryPort) -> None: + self.global_repo = global_repo + + async def create_inbound_route(self, tenant_id: int, cmd: CreateInboundRouteCmd) -> RouteEntity: + logger.info(f"Creating Inbound Route for sender {cmd.isa_sender_id} in tenant {tenant_id}") + route_id = await self.global_repo.create_inbound_route(tenant_id=tenant_id, cmd=cmd) + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.INBOUND_ROUTE_CREATED, + payload={"route_id": str(route_id), "tenant_id": tenant_id}, + ) + return RouteEntity(route_id=route_id, tenant_id=tenant_id, direction="INBOUND") + + async def update_inbound_route( + self, tenant_id: int, route_id: UUID, cmd: UpdateInboundRouteCmd + ) -> bool: + res = await self.global_repo.update_inbound_route(tenant_id, route_id, cmd) + if res: + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.INBOUND_ROUTE_UPDATED, + payload={"route_id": str(route_id), "tenant_id": tenant_id}, + ) + return res + + async def delete_inbound_route(self, tenant_id: int, route_id: UUID) -> bool: + res = await self.global_repo.delete_inbound_route(tenant_id, route_id) + if res: + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.INBOUND_ROUTE_DELETED, + payload={"route_id": str(route_id), "tenant_id": tenant_id}, + ) + return res + + async def create_outbound_route( + self, tenant_id: int, cmd: CreateOutboundRouteCmd + ) -> RouteEntity: + logger.info(f"Creating Outbound Route for sender {cmd.isa_sender_id} in tenant {tenant_id}") + route_id = await self.global_repo.create_outbound_route(tenant_id=tenant_id, cmd=cmd) + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.OUTBOUND_ROUTE_CREATED, + payload={"route_id": str(route_id), "tenant_id": tenant_id}, + ) + return RouteEntity(route_id=route_id, tenant_id=tenant_id, direction="OUTBOUND") + + async def update_outbound_route( + self, tenant_id: int, route_id: UUID, cmd: UpdateOutboundRouteCmd + ) -> bool: + res = await self.global_repo.update_outbound_route(tenant_id, route_id, cmd) + if res: + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.OUTBOUND_ROUTE_UPDATED, + payload={"route_id": str(route_id), "tenant_id": tenant_id}, + ) + return res + + async def delete_outbound_route(self, tenant_id: int, route_id: UUID) -> bool: + res = await self.global_repo.delete_outbound_route(tenant_id, route_id) + if res: + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.OUTBOUND_ROUTE_DELETED, + payload={"route_id": str(route_id), "tenant_id": tenant_id}, + ) + return res + + async def list_routes(self, tenant_id: int) -> list[dict[str, Any]]: + routes_data = await self.global_repo.get_all_routes(tenant_id) + inbound = routes_data.get("inbound", []) + outbound = routes_data.get("outbound", []) + + as2_ids: set[UUID] = set() + sftp_ids: set[UUID] = set() + webhook_ids: set[UUID] = set() + + for r in inbound: + if r.as2_partner_id: + as2_ids.add(r.as2_partner_id) + if r.sftp_partner_id: + sftp_ids.add(r.sftp_partner_id) + if r.webhook_id: + webhook_ids.add(r.webhook_id) + + for r in outbound: + if r.as2_partner_id: + as2_ids.add(r.as2_partner_id) + if r.sftp_partner_id: + sftp_ids.add(r.sftp_partner_id) + + as2_names = ( + await self.global_repo.get_as2_partners_by_ids(tenant_id, list(as2_ids)) + if as2_ids + else {} + ) + sftp_names = ( + await self.global_repo.get_sftp_partners_by_ids(tenant_id, list(sftp_ids)) + if sftp_ids + else {} + ) + webhook_names = ( + await self.global_repo.get_webhooks_by_ids(tenant_id, list(webhook_ids)) + if webhook_ids + else {} + ) + + results: list[dict[str, Any]] = [] + + def _resolve_destination(r: Any) -> tuple[str, str]: + if r.as2_partner_id: + return "AS2", as2_names.get(r.as2_partner_id, str(r.as2_partner_id)) + if r.sftp_partner_id: + return "SFTP", sftp_names.get(r.sftp_partner_id, str(r.sftp_partner_id)) + if getattr(r, "webhook_id", None): + return "WEBHOOK", webhook_names.get(r.webhook_id, str(r.webhook_id)) + return "UNKNOWN", "Unknown" + + for r in inbound: + dest_type, dest_name = _resolve_destination(r) + + results.append( + { + "route_id": r.id, + "name": r.name, + "direction": "INBOUND", + "isa_sender_id": r.isa_sender_id, + "isa_receiver_id": r.isa_receiver_id, + "gs_sender_id": r.gs_sender_id, + "gs_receiver_id": r.gs_receiver_id, + "transaction_type": r.transaction_type, + "destination_type": dest_type, + "destination_name": dest_name, + "webhook_id": r.webhook_id, + "as2_partner_id": r.as2_partner_id, + "sftp_partner_id": r.sftp_partner_id, + "active": r.active, + } + ) + + for r in outbound: + dest_type, dest_name = _resolve_destination(r) + + results.append( + { + "route_id": r.id, + "name": r.name, + "direction": "OUTBOUND", + "trading_partner_id": r.trading_partner_id, + "isa_sender_id": r.isa_sender_id, + "isa_sender_qualifier": r.isa_sender_qualifier, + "isa_receiver_id": r.isa_receiver_id, + "isa_receiver_qualifier": r.isa_receiver_qualifier, + "gs_sender_id": r.gs_sender_id, + "gs_receiver_id": r.gs_receiver_id, + "default_standard": r.default_standard, + "default_version": r.default_version, + "transaction_type": r.transaction_type, + "destination_type": dest_type, + "destination_name": dest_name, + "as2_partner_id": r.as2_partner_id, + "sftp_partner_id": r.sftp_partner_id, + "active": r.active, + } + ) + + return results diff --git a/services/api/src/api/core/services/sftp_partner_service.py b/services/api/src/api/core/services/sftp_partner_service.py new file mode 100644 index 00000000..55545c84 --- /dev/null +++ b/services/api/src/api/core/services/sftp_partner_service.py @@ -0,0 +1,62 @@ +import logging +from uuid import UUID + +from api.domain.models import ( + CreateSFTPPartnerCmd, + PartnerEntity, + UpdateSFTPPartnerCmd, +) +from api.ports.repository import ControlPlaneRepositoryPort +from domain.events import ProvisioningEventType + +logger = logging.getLogger(__name__) + + +class SFTPPartnerService: + """ + Domain service responsible for the lifecycle of SFTP Partners. + """ + + def __init__(self, global_repo: ControlPlaneRepositoryPort) -> None: + self.global_repo = global_repo + + async def create_sftp_partner(self, tenant_id: int, cmd: CreateSFTPPartnerCmd) -> PartnerEntity: + logger.info(f"Creating SFTP partner {cmd.name} for tenant {tenant_id}") + partner_id = await self.global_repo.create_sftp_partner(tenant_id=tenant_id, cmd=cmd) + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.SFTP_PARTNER_CREATED, + payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, + ) + + return PartnerEntity( + partner_id=partner_id, + tenant_id=tenant_id, + name=cmd.name, + type="SFTP", + status="INACTIVE", + ) + + async def update_sftp_partner( + self, tenant_id: int, partner_id: UUID, cmd: UpdateSFTPPartnerCmd + ) -> PartnerEntity: + logger.info(f"Updating SFTP partner {partner_id} for tenant {tenant_id}") + await self.global_repo.update_sftp_partner( + tenant_id=tenant_id, partner_id=partner_id, cmd=cmd + ) + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.SFTP_PARTNER_UPDATED, + payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, + ) + updated = await self.global_repo.get_sftp_partner(tenant_id, partner_id) + if not updated: + raise ValueError(f"SFTP partner {partner_id} not found") + + return PartnerEntity( + partner_id=partner_id, + tenant_id=tenant_id, + name=cmd.name or updated.name, + type="SFTP", + status="ACTIVE" if updated.active else "INACTIVE", + ) diff --git a/services/api/src/api/core/services/webhook_service.py b/services/api/src/api/core/services/webhook_service.py new file mode 100644 index 00000000..7b7429a4 --- /dev/null +++ b/services/api/src/api/core/services/webhook_service.py @@ -0,0 +1,33 @@ +import logging + +from api.domain.models import CreateWebhookCmd, PartnerEntity +from api.ports.repository import ControlPlaneRepositoryPort +from domain.events import ProvisioningEventType + +logger = logging.getLogger(__name__) + + +class WebhookService: + """ + Domain service responsible for the lifecycle of Webhook delivery destinations. + """ + + def __init__(self, global_repo: ControlPlaneRepositoryPort) -> None: + self.global_repo = global_repo + + async def create_webhook(self, tenant_id: int, cmd: CreateWebhookCmd) -> PartnerEntity: + logger.info(f"Creating Webhook {cmd.name} for tenant {tenant_id}") + partner_id = await self.global_repo.create_webhook(tenant_id=tenant_id, cmd=cmd) + await self.global_repo.create_outbox_event( + tenant_id=tenant_id, + event_type=ProvisioningEventType.WEBHOOK_CREATED, + payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, + ) + + return PartnerEntity( + partner_id=partner_id, + tenant_id=tenant_id, + name=cmd.name, + type="WEBHOOK", + status="ACTIVE", + ) diff --git a/services/api/src/api/dependencies.py b/services/api/src/api/dependencies.py index 92e78e62..8ee2d8bd 100644 --- a/services/api/src/api/dependencies.py +++ b/services/api/src/api/dependencies.py @@ -1,28 +1,35 @@ import os +from collections.abc import AsyncGenerator from functools import lru_cache from typing import Any from database.session import get_global_session -from fastapi import Depends, HTTPException - -# Import tenant_session from identity -from identity.dependencies import get_current_tenant_id, get_raw_jwt, get_tenant_session +from fastapi import Depends, HTTPException, Request +from identity.dependencies import ( + get_current_tenant_id, + get_raw_jwt, + get_tenant_session, + get_tenant_session_for_id, +) from sqlalchemy.ext.asyncio import AsyncSession from api.adapters.httpx_as2_tester import HttpxAS2TesterAdapter from api.adapters.paramiko_sftp_tester import ParamikoSftpTesterAdapter from api.adapters.repository import ( + SqlAlchemyApiTokenRepository, SqlAlchemyControlPlaneRepository, SqlAlchemyDataPlaneRepository, SqlAlchemyTenantRepository, ) from api.adapters.sqs_queue import SQSMessageQueueAdapter from api.adapters.vault import vault +from api.auth.api_key import get_tenant_id_from_api_key from api.core.authorization import AuthorizationService from api.core.uow import UnitOfWork from api.ports.as2_tester import AS2TesterPort from api.ports.message_queue import MessageQueuePort from api.ports.repository import ( + ApiTokenRepositoryPort, ControlPlaneRepositoryPort, DataPlaneRepositoryPort, TenantRepositoryPort, @@ -80,6 +87,19 @@ async def get_tenant_uow( return UnitOfWork(global_session=global_session, tenant_session=tenant_session) +async def get_m2m_tenant_uow( + request: Request, + tenant_id: int = Depends(get_tenant_id_from_api_key), + global_session: AsyncSession = Depends(get_global_session), +) -> AsyncGenerator[UnitOfWork, None]: + """ + Constructs a UnitOfWork dynamically without relying on Zitadel JWTs. + Useful for Machine-to-Machine routes that authenticate via API keys. + """ + async for tenant_session in get_tenant_session_for_id(request, tenant_id, global_session): + yield UnitOfWork(global_session=global_session, tenant_session=tenant_session) + + def require_platform_admin(tenant_id: int = Depends(get_current_tenant_id)) -> int: """ Dependency that enforces the user belongs to Tenant 0 (Platform Admin). @@ -98,6 +118,13 @@ def get_tenant_repo( return SqlAlchemyTenantRepository(session) +def get_api_token_repo( + session: AsyncSession = Depends(get_global_session), +) -> ApiTokenRepositoryPort: + """Yields the API token repository bound to the global (control plane) session.""" + return SqlAlchemyApiTokenRepository(session) + + def get_authorization_service( tenant_repo: TenantRepositoryPort = Depends(get_tenant_repo), ) -> AuthorizationService: diff --git a/services/api/src/api/domain/models.py b/services/api/src/api/domain/models.py index d1bd79f4..b1a7e409 100644 --- a/services/api/src/api/domain/models.py +++ b/services/api/src/api/domain/models.py @@ -1,4 +1,5 @@ from dataclasses import dataclass +from datetime import datetime from typing import Any from uuid import UUID @@ -46,6 +47,7 @@ class CreateAS2PartnershipCmd: local_partner_id: UUID remote_partner_id: UUID name: str + trading_partner_id: str | None = None credentials_vault_ref: str | None = None mdn_type: str = "SYNC" mdn_url: str | None = None @@ -115,6 +117,8 @@ class CreateInboundRouteCmd: isa_sender_id: str isa_receiver_id: str transaction_type: str + gs_sender_id: str | None = None + gs_receiver_id: str | None = None processing_mode: str = "TRANSLATE" webhook_id: UUID | None = None as2_partner_id: UUID | None = None @@ -126,6 +130,8 @@ class UpdateInboundRouteCmd: name: str | UnsetType = UNSET isa_sender_id: str | UnsetType = UNSET isa_receiver_id: str | UnsetType = UNSET + gs_sender_id: str | None | UnsetType = UNSET + gs_receiver_id: str | None | UnsetType = UNSET transaction_type: str | UnsetType = UNSET processing_mode: str | UnsetType = UNSET webhook_id: UUID | None | UnsetType = UNSET @@ -136,10 +142,17 @@ class UpdateInboundRouteCmd: @dataclass(frozen=True) class CreateOutboundRouteCmd: + trading_partner_id: str name: str isa_sender_id: str isa_receiver_id: str + gs_sender_id: str + gs_receiver_id: str transaction_type: str + isa_sender_qualifier: str | None = None + isa_receiver_qualifier: str | None = None + default_standard: str = "x12" + default_version: str = "004010" processing_mode: str = "TRANSLATE" as2_partner_id: UUID | None = None sftp_partner_id: UUID | None = None @@ -147,10 +160,17 @@ class CreateOutboundRouteCmd: @dataclass(frozen=True) class UpdateOutboundRouteCmd: + trading_partner_id: str | None | UnsetType = UNSET name: str | UnsetType = UNSET isa_sender_id: str | UnsetType = UNSET + isa_sender_qualifier: str | None | UnsetType = UNSET isa_receiver_id: str | UnsetType = UNSET + isa_receiver_qualifier: str | None | UnsetType = UNSET + gs_sender_id: str | UnsetType = UNSET + gs_receiver_id: str | UnsetType = UNSET transaction_type: str | UnsetType = UNSET + default_standard: str | UnsetType = UNSET + default_version: str | UnsetType = UNSET processing_mode: str | UnsetType = UNSET as2_partner_id: UUID | None | UnsetType = UNSET sftp_partner_id: UUID | None | UnsetType = UNSET @@ -176,3 +196,26 @@ class RouteEntity: route_id: UUID tenant_id: int direction: str # INBOUND, OUTBOUND + + +# --------------------------------------------------------------------------- +# API Token Commands & Entities +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True) +class CreateApiTokenCmd: + name: str + expires_at: datetime | None = None + + +@dataclass(frozen=True) +class ApiTokenEntity: + """Returned once after creation. client_secret is shown only this time.""" + + id: UUID + tenant_id: int + name: str + client_id: str # stored plaintext, safe to display in UI + client_secret: str # shown once, never stored — only its hash is in DB + active: bool diff --git a/services/api/src/api/main.py b/services/api/src/api/main.py index 08be2c04..bdd9bd9e 100644 --- a/services/api/src/api/main.py +++ b/services/api/src/api/main.py @@ -5,6 +5,7 @@ from dotenv import load_dotenv load_dotenv() +logging.basicConfig(level=logging.INFO) from config.settings import get_settings from database.connection import DatabaseRouter @@ -18,7 +19,8 @@ from api import cdc_relay from api.dependencies import get_current_user_profile -from api.routers import edi_tools, routes, trading_partners, webhooks +from api.routers import edi_json, edi_tools, routes, trading_partners, webhooks +from api.routers.developers import api_tokens from api.routers.trading_partners import as2_receive, platform logger = logging.getLogger(__name__) @@ -85,6 +87,8 @@ async def validation_exception_handler( app.include_router(routes.router) app.include_router(edi_tools.router) app.include_router(as2_receive.router, prefix="/api/v1") +app.include_router(edi_json.router) +app.include_router(api_tokens.router) @app.get("/api/me", tags=["Identity"]) diff --git a/services/api/src/api/ports/repository.py b/services/api/src/api/ports/repository.py index 2005a5d8..3ccde99c 100644 --- a/services/api/src/api/ports/repository.py +++ b/services/api/src/api/ports/repository.py @@ -24,6 +24,13 @@ async def create_as2_identity( async def update_as2_identity( self, tenant_id: int, partner_id: UUID, cmd: UpdateAS2TradingPartnerCmd ) -> None: ... + async def rotate_as2_certificates( + self, + tenant_id: int, + partner_id: UUID, + new_public_cert: str, + new_private_key_vault_ref: str | None, + ) -> None: ... async def get_as2_partner(self, tenant_id: int, partner_id: UUID) -> Any: ... async def delete_as2_identity(self, tenant_id: int, partner_id: UUID) -> None: ... async def get_as2_partners_by_ids(self, tenant_id: int, ids: list[UUID]) -> dict[UUID, str]: ... @@ -38,6 +45,7 @@ async def update_as2_partnership( self, tenant_id: int, partnership_id: UUID, cmd: UpdateAS2PartnershipCmd ) -> None: ... async def get_as2_partnership(self, tenant_id: int, partnership_id: UUID) -> Any: ... + async def delete_as2_partnership(self, tenant_id: int, partnership_id: UUID) -> None: ... async def get_partnership_by_as2_ids( self, as2_from: str, as2_to: str @@ -69,12 +77,22 @@ async def create_inbound_route(self, tenant_id: int, cmd: CreateInboundRouteCmd) async def update_inbound_route( self, tenant_id: int, route_id: UUID, cmd: UpdateInboundRouteCmd ) -> bool: ... + async def get_inbound_route( + self, + isa_sender_id: str, + isa_receiver_id: str, + tenant_id: int, + transaction_type: str | None = None, + ) -> Any | None: ... async def delete_inbound_route(self, tenant_id: int, route_id: UUID) -> bool: ... async def create_outbound_route(self, tenant_id: int, cmd: CreateOutboundRouteCmd) -> UUID: ... async def update_outbound_route( self, tenant_id: int, route_id: UUID, cmd: UpdateOutboundRouteCmd ) -> bool: ... + async def get_outbound_route_by_trading_partner_id( + self, tenant_id: int, trading_partner_id: str + ) -> Any | None: ... async def delete_outbound_route(self, tenant_id: int, route_id: UUID) -> bool: ... async def get_all_routes(self, tenant_id: int) -> dict[str, list[Any]]: ... @@ -114,6 +132,18 @@ async def create_edi_message(self, tenant_id: int, payload: dict[str, Any]) -> U """ ... + async def create_edi_json(self, tenant_id: int, payload: dict[str, Any]) -> UUID: + """ + Saves a new EdiJson record to the Data Plane. + """ + ... + + async def create_api_gateway(self, tenant_id: int, payload: dict[str, Any]) -> UUID: + """ + Saves a new ApiGateway log record to the Data Plane. + """ + ... + async def create_outbox_event( self, tenant_id: int, event_type: str, payload: dict[str, Any] ) -> UUID: @@ -122,6 +152,18 @@ async def create_outbox_event( """ ... + async def get_as2_partnership(self, tenant_id: int, partnership_id: UUID) -> Any: + """ + Retrieves a replicated AS2 Partnership from the Data Plane by internal UUID. + """ + ... + + async def get_as2_partner(self, tenant_id: int, partner_id: UUID) -> Any: + """ + Retrieves a replicated AS2 Partner from the Data Plane. + """ + ... + class TenantRepositoryPort(Protocol): """ @@ -129,3 +171,30 @@ class TenantRepositoryPort(Protocol): """ async def get_tenant_flags(self, tenant_id: int) -> dict[str, Any] | None: ... + + +class ApiTokenRepositoryPort(Protocol): + """ + Port for managing platform API tokens. + Implemented by SqlAlchemyApiTokenRepository (adapter). + Can be stubbed in unit tests with any class that satisfies this interface. + """ + + async def create_api_token( + self, + tenant_id: int, + name: str, + client_id: str, + secret_hash: str, + expires_at: Any | None, + ) -> UUID: ... + + async def list_api_tokens(self, tenant_id: int) -> list[dict[str, Any]]: ... + + async def revoke_api_token(self, tenant_id: int, token_id: UUID) -> bool: ... + + async def delete_api_token(self, tenant_id: int, token_id: UUID) -> bool: ... + + async def get_tenant_id_by_credentials( + self, client_id: str, secret_hash: str + ) -> int | None: ... diff --git a/services/api/src/api/routers/developers/__init__.py b/services/api/src/api/routers/developers/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/services/api/src/api/routers/developers/api_tokens.py b/services/api/src/api/routers/developers/api_tokens.py new file mode 100644 index 00000000..8bb1ca01 --- /dev/null +++ b/services/api/src/api/routers/developers/api_tokens.py @@ -0,0 +1,101 @@ +import logging +from typing import Any +from uuid import UUID + +from fastapi import APIRouter, Depends, HTTPException, status +from identity.dependencies import get_current_tenant_id + +from api.adapters.http.dtos import ( + ApiTokenCreatedResponse, + ApiTokenListItem, + ApiTokenListResponse, + CreateApiTokenRequest, +) +from api.core.services.api_token_service import ApiTokenService +from api.dependencies import get_api_token_repo +from api.domain.models import CreateApiTokenCmd +from api.ports.repository import ApiTokenRepositoryPort + +logger = logging.getLogger(__name__) + +router = APIRouter(prefix="/api/v1/developers/tokens", tags=["API Tokens"]) + + +@router.post( + "", + response_model=ApiTokenCreatedResponse, + status_code=status.HTTP_201_CREATED, +) +async def create_api_token( + request: CreateApiTokenRequest, + tenant_id: int = Depends(get_current_tenant_id), + repo: ApiTokenRepositoryPort = Depends(get_api_token_repo), +) -> Any: + """ + Generate a new two-part API credential for the tenant. + The client_secret is returned exactly once in this response and cannot be retrieved again. + """ + service = ApiTokenService(repo) + # Get the tenant name from somewhere? For now, we can pass a dummy or fetch it. + # We should ideally fetch the tenant name to use in the prefix. + # For now, using a generic placeholder as the slug generator handles it. + tenant_name = f"tenant{tenant_id}" + + cmd = CreateApiTokenCmd(name=request.name, expires_at=request.expires_at) + + token = await service.create_token(tenant_id=tenant_id, tenant_name=tenant_name, cmd=cmd) + + return ApiTokenCreatedResponse( + id=token.id, + name=token.name, + client_id=token.client_id, + client_secret=token.client_secret, + active=token.active, + created_at="just now", # This is usually handled by the DB, but since we return the entity immediately, we need a value. Let's just return the current ISO string if needed, or rely on the UI to just know it's new. Actually, we should probably fetch the record or just return a fresh datetime. + ) + + +@router.get( + "", + response_model=ApiTokenListResponse, +) +async def list_api_tokens( + tenant_id: int = Depends(get_current_tenant_id), + repo: ApiTokenRepositoryPort = Depends(get_api_token_repo), +) -> Any: + """List all API tokens for the current tenant.""" + service = ApiTokenService(repo) + tokens = await service.list_tokens(tenant_id) + return ApiTokenListResponse(tokens=[ApiTokenListItem(**t) for t in tokens]) + + +@router.delete( + "/{token_id}", + status_code=status.HTTP_204_NO_CONTENT, +) +async def revoke_api_token( + token_id: UUID, + tenant_id: int = Depends(get_current_tenant_id), + repo: ApiTokenRepositoryPort = Depends(get_api_token_repo), +) -> None: + """Revoke (soft delete) an API token.""" + service = ApiTokenService(repo) + success = await service.revoke_token(tenant_id, token_id) + if not success: + raise HTTPException(status_code=404, detail="Token not found") + + +@router.delete( + "/{token_id}/hard", + status_code=status.HTTP_204_NO_CONTENT, +) +async def delete_api_token( + token_id: UUID, + tenant_id: int = Depends(get_current_tenant_id), + repo: ApiTokenRepositoryPort = Depends(get_api_token_repo), +) -> None: + """Permanently delete an API token.""" + service = ApiTokenService(repo) + success = await service.delete_token(tenant_id, token_id) + if not success: + raise HTTPException(status_code=404, detail="Token not found") diff --git a/services/api/src/api/routers/edi_json.py b/services/api/src/api/routers/edi_json.py new file mode 100644 index 00000000..761a1c80 --- /dev/null +++ b/services/api/src/api/routers/edi_json.py @@ -0,0 +1,40 @@ +from typing import Any + +from fastapi import APIRouter, Depends, status + +from api.adapters.http.dtos import OutboundMessageRequest, OutboundMessageResponse +from api.auth.api_key import get_tenant_id_from_api_key +from api.core.uow import UnitOfWork +from api.dependencies import get_m2m_tenant_uow +from api.services.outbound_service import OutboundService + +router = APIRouter(prefix="/api/v1/edi_json", tags=["EDI JSON"]) + + +@router.post( + "", + response_model=OutboundMessageResponse, + status_code=status.HTTP_202_ACCEPTED, +) +async def submit_outbound_message( + request: OutboundMessageRequest, + tenant_id: int = Depends(get_tenant_id_from_api_key), + uow: UnitOfWork = Depends(get_m2m_tenant_uow), +) -> Any: + """ + Submits a JSON payload to be translated and transmitted via AS2. + + Authentication: Two-part API key (no Zitadel / OAuth2 required). + X-Client-ID: soopaedi__ + X-Client-Secret: + """ + service = OutboundService(uow=uow) + + trace_id = await service.process_outbound_message( + tenant_id=tenant_id, + trading_partner_id=request.trading_partner_id, + payload=request.payload, + transaction_type=request.transaction_type, + ) + + return OutboundMessageResponse(trace_id=trace_id, status="ACCEPTED") diff --git a/services/api/src/api/routers/routes.py b/services/api/src/api/routers/routes.py index 5593eda1..e51c9381 100644 --- a/services/api/src/api/routers/routes.py +++ b/services/api/src/api/routers/routes.py @@ -11,12 +11,15 @@ RouteResponse, UpdateRouteRequest, ) -from api.core.provisioning import ProvisioningService +from api.core.services import RouteService from api.core.uow import UnitOfWork from api.dependencies import get_tenant_uow from api.domain.models import ( + UNSET, CreateInboundRouteCmd, CreateOutboundRouteCmd, + UpdateInboundRouteCmd, + UpdateOutboundRouteCmd, ) router = APIRouter(prefix="/api/v1/routes", tags=["Routes"]) @@ -31,12 +34,7 @@ async def list_routes( List all Active Routes for the current Tenant. """ async with uow: - from typing import cast - - from api.ports.repository import DataPlaneRepositoryPort - - data_plane = cast(DataPlaneRepositoryPort, uow.data_plane) - service = ProvisioningService(tenant_repo=data_plane, global_repo=uow.control_plane) + service = RouteService(global_repo=uow.control_plane) routes = await service.list_routes(tenant_id) return [RouteItemResponse(**r) for r in routes] @@ -52,12 +50,14 @@ async def create_inbound_route( Creates a new Inbound Route directly in the Tenant Data Plane. """ async with uow: - service = ProvisioningService(tenant_repo=uow.data_plane, global_repo=uow.control_plane) + service = RouteService(global_repo=uow.control_plane) cmd = CreateInboundRouteCmd( name=request.name, isa_sender_id=request.isa_sender_id, isa_receiver_id=request.isa_receiver_id, + gs_sender_id=request.gs_sender_id, + gs_receiver_id=request.gs_receiver_id, transaction_type=request.transaction_type, processing_mode=request.processing_mode, webhook_id=request.webhook_id, @@ -83,13 +83,20 @@ async def create_outbound_route( Creates a new Outbound Route directly in the Tenant Data Plane. """ async with uow: - service = ProvisioningService(tenant_repo=uow.data_plane, global_repo=uow.control_plane) + service = RouteService(global_repo=uow.control_plane) cmd = CreateOutboundRouteCmd( + trading_partner_id=request.trading_partner_id, name=request.name, isa_sender_id=request.isa_sender_id, + isa_sender_qualifier=request.isa_sender_qualifier, isa_receiver_id=request.isa_receiver_id, + isa_receiver_qualifier=request.isa_receiver_qualifier, + gs_sender_id=request.gs_sender_id, + gs_receiver_id=request.gs_receiver_id, transaction_type=request.transaction_type, + default_standard=request.default_standard, + default_version=request.default_version, processing_mode=request.processing_mode, as2_partner_id=request.as2_partner_id, sftp_partner_id=request.sftp_partner_id, @@ -110,15 +117,19 @@ async def update_inbound_route( tenant_id: int = Depends(get_current_tenant_id), uow: UnitOfWork = Depends(get_tenant_uow), ) -> Any: + """ + Updates an Inbound Route for the current Tenant. + """ async with uow: - service = ProvisioningService(tenant_repo=uow.data_plane, global_repo=uow.control_plane) - from api.domain.models import UNSET, UpdateInboundRouteCmd + service = RouteService(global_repo=uow.control_plane) dump = request.model_dump(exclude_unset=True) cmd = UpdateInboundRouteCmd( name=dump.get("name", UNSET), isa_sender_id=dump.get("isa_sender_id", UNSET), isa_receiver_id=dump.get("isa_receiver_id", UNSET), + gs_sender_id=dump.get("gs_sender_id", UNSET), + gs_receiver_id=dump.get("gs_receiver_id", UNSET), transaction_type=dump.get("transaction_type", UNSET), processing_mode=dump.get("processing_mode", UNSET), webhook_id=dump.get("webhook_id", UNSET), @@ -140,8 +151,11 @@ async def delete_inbound_route( tenant_id: int = Depends(get_current_tenant_id), uow: UnitOfWork = Depends(get_tenant_uow), ) -> None: + """ + Deletes an Inbound Route for the current Tenant. + """ async with uow: - service = ProvisioningService(tenant_repo=uow.data_plane, global_repo=uow.control_plane) + service = RouteService(global_repo=uow.control_plane) success = await service.delete_inbound_route(tenant_id, route_id) if not success: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Route not found") @@ -155,16 +169,25 @@ async def update_outbound_route( tenant_id: int = Depends(get_current_tenant_id), uow: UnitOfWork = Depends(get_tenant_uow), ) -> Any: + """ + Updates an Outbound Route for the current Tenant. + """ async with uow: - service = ProvisioningService(tenant_repo=uow.data_plane, global_repo=uow.control_plane) - from api.domain.models import UNSET, UpdateOutboundRouteCmd + service = RouteService(global_repo=uow.control_plane) dump = request.model_dump(exclude_unset=True) cmd = UpdateOutboundRouteCmd( + trading_partner_id=dump.get("trading_partner_id", UNSET), name=dump.get("name", UNSET), isa_sender_id=dump.get("isa_sender_id", UNSET), + isa_sender_qualifier=dump.get("isa_sender_qualifier", UNSET), isa_receiver_id=dump.get("isa_receiver_id", UNSET), + isa_receiver_qualifier=dump.get("isa_receiver_qualifier", UNSET), + gs_sender_id=dump.get("gs_sender_id", UNSET), + gs_receiver_id=dump.get("gs_receiver_id", UNSET), transaction_type=dump.get("transaction_type", UNSET), + default_standard=dump.get("default_standard", UNSET), + default_version=dump.get("default_version", UNSET), processing_mode=dump.get("processing_mode", UNSET), as2_partner_id=dump.get("as2_partner_id", UNSET), sftp_partner_id=dump.get("sftp_partner_id", UNSET), @@ -184,8 +207,11 @@ async def delete_outbound_route( tenant_id: int = Depends(get_current_tenant_id), uow: UnitOfWork = Depends(get_tenant_uow), ) -> None: + """ + Deletes an Outbound Route for the current Tenant. + """ async with uow: - service = ProvisioningService(tenant_repo=uow.data_plane, global_repo=uow.control_plane) + service = RouteService(global_repo=uow.control_plane) success = await service.delete_outbound_route(tenant_id, route_id) if not success: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Route not found") diff --git a/services/api/src/api/routers/trading_partners/as2.py b/services/api/src/api/routers/trading_partners/as2.py index 1988d7d2..49b319ec 100644 --- a/services/api/src/api/routers/trading_partners/as2.py +++ b/services/api/src/api/routers/trading_partners/as2.py @@ -5,9 +5,15 @@ from fastapi import APIRouter, Depends, HTTPException from identity.dependencies import get_current_tenant_id, get_raw_jwt -from api.adapters.http.dtos import CertificateExportResponse +from api.adapters.http.dtos import ( + AS2TradingPartnerResponse, + CertificateExportResponse, + RotateCertificateRequest, +) +from api.core.services import AS2PartnerService from api.core.uow import UnitOfWork from api.dependencies import get_current_user_profile, get_uow, get_vault +from api.domain.certificate import generate_self_signed_cert from api.ports.vault import VaultPort logger = logging.getLogger(__name__) @@ -72,3 +78,87 @@ async def export_as2_certificates( ) from e return response + + +@router.put( + "/as2/trading-partners/{partner_id}/certificates/rotate", + response_model=AS2TradingPartnerResponse, +) +async def rotate_as2_certificates( + partner_id: UUID, + request: RotateCertificateRequest, + tenant_id: int = Depends(get_current_tenant_id), + uow: UnitOfWork = Depends(get_uow), + profile: dict[str, Any] = Depends(get_current_user_profile), + vault: VaultPort = Depends(get_vault), +) -> Any: + """Rotates certificates for an AS2 partner.""" + async with uow: + partner = await uow.control_plane.get_as2_partner(tenant_id, partner_id) + if not partner: + partner = await uow.control_plane.get_as2_partner(0, partner_id) + if not partner: + raise HTTPException(status_code=404, detail="Partner not found") + + actual_tenant_id = partner.tenant_id + + if partner.is_local and "certificates:rotate" not in profile["permissions"]: + raise HTTPException( + status_code=403, detail="Insufficient permissions to rotate certificates." + ) + + public_cert_pem = request.public_cert_pem + private_key_vault_ref = None + + if partner.is_local: + if request.action == "generate": + private_key_bytes, public_cert_bytes = generate_self_signed_cert( + common_name=partner.as2_id + ) + private_key_vault_ref = vault.store_private_key( + private_key_pem=private_key_bytes, + alias_prefix=f"{partner.name.replace(' ', '_').lower()}_rotated", + ) + public_cert_pem = public_cert_bytes.decode("utf-8") + elif request.action == "upload": + if not request.private_key_pem or not request.public_cert_pem: + raise HTTPException( + status_code=400, + detail="Both public_cert_pem and private_key_pem required for upload.", + ) + private_key_vault_ref = vault.store_private_key( + private_key_pem=request.private_key_pem.encode("utf-8"), + alias_prefix=f"{partner.name.replace(' ', '_').lower()}_uploaded", + ) + else: + if not request.public_cert_pem: + raise HTTPException( + status_code=400, detail="public_cert_pem required for remote partners." + ) + + try: + svc = AS2PartnerService(global_repo=uow.control_plane) + updated_partner = await svc.rotate_certificates( + tenant_id=actual_tenant_id, + partner_id=partner_id, + new_public_cert=str(public_cert_pem), + new_private_key_vault_ref=private_key_vault_ref, + ) + await uow.commit() + except ValueError as e: + if private_key_vault_ref: + vault.delete_secret(private_key_vault_ref) + raise HTTPException(status_code=400, detail=str(e)) from e + except Exception as e: + if private_key_vault_ref: + vault.delete_secret(private_key_vault_ref) + raise HTTPException(status_code=500, detail="Internal server error") from e + + return AS2TradingPartnerResponse( + id=str(updated_partner.partner_id), + name=updated_partner.name, + as2_id=partner.as2_id, + is_local=partner.is_local, + url=partner.url, + active=updated_partner.status == "ACTIVE", + ) diff --git a/services/api/src/api/routers/trading_partners/platform/as2_partners.py b/services/api/src/api/routers/trading_partners/platform/as2_partners.py index 7fbf01aa..18c2f52b 100644 --- a/services/api/src/api/routers/trading_partners/platform/as2_partners.py +++ b/services/api/src/api/routers/trading_partners/platform/as2_partners.py @@ -9,7 +9,7 @@ CreateAS2TradingPartnerRequest, UpdateAS2TradingPartnerRequest, ) -from api.core.provisioning import ProvisioningService +from api.core.services import AS2PartnerService from api.core.uow import UnitOfWork from api.dependencies import ( get_uow, @@ -69,13 +69,14 @@ async def create_platform_as2_partner( ) # Use tenant_id=0 for global platform partners - partner_id = await uow.control_plane.create_as2_identity(tenant_id=0, cmd=cmd) + svc = AS2PartnerService(global_repo=uow.control_plane) + entity = await svc.create_as2_partner(tenant_id=0, cmd=cmd) await uow.commit() - p = await uow.control_plane.get_as2_partner(tenant_id=0, partner_id=partner_id) + p = await uow.control_plane.get_as2_partner(tenant_id=0, partner_id=entity.partner_id) return AS2TradingPartnerResponse( - id=str(partner_id), + id=str(entity.partner_id), name=p.name, as2_id=p.as2_id, is_local=p.is_local, @@ -126,7 +127,8 @@ async def update_platform_as2_partner( active=request.active, ) try: - await uow.control_plane.update_as2_identity(tenant_id=0, partner_id=partner_id, cmd=cmd) + svc = AS2PartnerService(global_repo=uow.control_plane) + await svc.update_as2_partner(tenant_id=0, partner_id=partner_id, cmd=cmd) updated_partner = await uow.control_plane.get_as2_partner( tenant_id=0, partner_id=partner_id ) @@ -155,7 +157,7 @@ async def delete_platform_as2_partner( ) -> None: """Deletes an AS2 partner.""" async with uow: - svc = ProvisioningService(tenant_repo=None, global_repo=uow.control_plane) + svc = AS2PartnerService(global_repo=uow.control_plane) try: await svc.delete_as2_partner(tenant_id=0, partner_id=partner_id) await uow.commit() diff --git a/services/api/src/api/routers/trading_partners/platform/as2_partnerships.py b/services/api/src/api/routers/trading_partners/platform/as2_partnerships.py index a52ff0d1..1c89bb7e 100644 --- a/services/api/src/api/routers/trading_partners/platform/as2_partnerships.py +++ b/services/api/src/api/routers/trading_partners/platform/as2_partnerships.py @@ -14,6 +14,7 @@ TestAS2ConnectionResponse, UpdateAS2PartnershipRequest, ) +from api.core.services import AS2PartnershipService from api.core.uow import UnitOfWork from api.dependencies import ( get_as2_tester, @@ -156,16 +157,17 @@ async def create_platform_as2_partnership( advanced_flags=request.advanced_flags, ) - partnership_id = await uow.control_plane.create_as2_partnership(tenant_id=0, cmd=cmd) + svc = AS2PartnershipService(global_repo=uow.control_plane) + entity = await svc.create_as2_partnership(tenant_id=0, cmd=cmd) await uow.commit() p = await uow.control_plane.get_as2_partnership( - tenant_id=0, partnership_id=partnership_id + tenant_id=0, partnership_id=entity.partner_id ) if not p: raise HTTPException(status_code=404, detail="Partnership not found") return AS2PartnershipResponse( - id=str(partnership_id), + id=str(entity.partner_id), tenant_id=0, name=p.name, local_partner_id=str(p.local_partner_id), @@ -216,9 +218,8 @@ def get_val(field: str) -> Any: advanced_flags=get_val("advanced_flags"), active=get_val("active"), ) - await uow.control_plane.update_as2_partnership( - tenant_id=0, partnership_id=partnership_id, cmd=cmd - ) + svc = AS2PartnershipService(global_repo=uow.control_plane) + await svc.update_as2_partnership(tenant_id=0, partnership_id=partnership_id, cmd=cmd) await uow.commit() p = await uow.control_plane.get_as2_partnership( @@ -263,9 +264,8 @@ async def delete_platform_as2_partnership( ) -> None: try: async with uow: - await uow.control_plane.delete_as2_partnership( - tenant_id=0, partnership_id=partnership_id - ) + svc = AS2PartnershipService(global_repo=uow.control_plane) + await svc.delete_as2_partnership(tenant_id=0, partnership_id=partnership_id) await uow.commit() except Exception as err: logger.exception("Internal error deleting platform AS2 partnership") diff --git a/services/api/src/api/routers/trading_partners/sftp.py b/services/api/src/api/routers/trading_partners/sftp.py index 897114f9..4bcb22a5 100644 --- a/services/api/src/api/routers/trading_partners/sftp.py +++ b/services/api/src/api/routers/trading_partners/sftp.py @@ -12,7 +12,7 @@ TestSFTPConnectionRequest, UpdateSFTPPartnerRequest, ) -from api.core.provisioning import ProvisioningService +from api.core.services import SFTPPartnerService from api.core.uow import UnitOfWork from api.dependencies import ( get_sftp_tester, @@ -137,7 +137,7 @@ async def create_sftp_partner( ) async with uow: - service = ProvisioningService(tenant_repo=uow.data_plane, global_repo=uow.control_plane) + service = SFTPPartnerService(global_repo=uow.control_plane) cmd = CreateSFTPPartnerCmd( name=request.name, @@ -191,7 +191,7 @@ async def update_sftp_partner( ) -> Any: """Updates an SFTP Partner in the Tenant Data Plane.""" async with uow: - service = ProvisioningService(tenant_repo=uow.data_plane, global_repo=uow.control_plane) + service = SFTPPartnerService(global_repo=uow.control_plane) cmd = UpdateSFTPPartnerCmd( name=request.name, host=request.host, diff --git a/services/api/src/api/routers/webhooks/webhook.py b/services/api/src/api/routers/webhooks/webhook.py index 59c23a68..772401dd 100644 --- a/services/api/src/api/routers/webhooks/webhook.py +++ b/services/api/src/api/routers/webhooks/webhook.py @@ -7,7 +7,7 @@ CreateWebhookRequest, PartnerResponse, ) -from api.core.provisioning import ProvisioningService +from api.core.services import WebhookService from api.core.uow import UnitOfWork from api.dependencies import get_tenant_uow from api.domain.models import CreateWebhookCmd @@ -43,7 +43,7 @@ async def create_webhook( raise HTTPException(status_code=400, detail="Invalid webhook URL") from e async with uow: - service = ProvisioningService(tenant_repo=uow.data_plane, global_repo=uow.control_plane) + service = WebhookService(global_repo=uow.control_plane) cmd = CreateWebhookCmd( name=request.name, diff --git a/services/api/src/api/services/as2_receive_service.py b/services/api/src/api/services/as2_receive_service.py index 46f5daae..4d68051f 100644 --- a/services/api/src/api/services/as2_receive_service.py +++ b/services/api/src/api/services/as2_receive_service.py @@ -243,7 +243,7 @@ def _reconstruct_smime_headers(self, headers: dict[str, str]) -> bytes: if not has_cte: smime_headers += "Content-Transfer-Encoding: binary\r\n" smime_headers += "\r\n" - return smime_headers.encode("utf-8") + return smime_headers.encode("latin-1") def _verify_and_calculate_mic( self, verify_entity: bytes, remote_cert: bytes, message_id: str @@ -289,9 +289,52 @@ def _extract_pure_edi(self, final_payload_bytes: bytes | str | Any) -> bytes: return parsed_msg.as_bytes() return final_payload_bytes + def _extract_isa_headers(self, pure_edi_bytes: bytes) -> tuple[str, str]: + """ + Lightweight ISA parser to extract Sender and Receiver for routing + without parsing the entire EDI structure. + """ + content = pure_edi_bytes.decode("ascii", errors="ignore") + if not content.startswith("ISA"): + raise ValueError("Payload does not begin with ISA segment") + + # The ISA segment is 106 characters long. + # The element separator is the 4th character. + element_separator = content[3] + + isa_segment = content[:106] + elements = isa_segment.split(element_separator) + + if len(elements) < 9: + raise ValueError("Malformed ISA segment") + + isa_sender = elements[6].strip() + isa_receiver = elements[8].strip() + + return isa_sender, isa_receiver + async def _save_to_data_plane( # type: ignore self, partnership, as2_msg: AS2Message, pure_edi_bytes: bytes ) -> str: + + # 1. Payload-Based Routing (ISA Extraction) + try: + isa_sender, isa_receiver = self._extract_isa_headers(pure_edi_bytes) + except Exception as e: + logger.error(f"Failed to extract ISA headers: {e}") + raise ValueError("Invalid EDI payload for routing") from e + + # 2. Query Global DB for the actual Tenant using ISA headers + route = await self.uow.control_plane.get_inbound_route( + isa_sender_id=isa_sender, isa_receiver_id=isa_receiver, tenant_id=partnership.tenant_id + ) + if not route: + logger.error(f"No inbound route found for ISA {isa_sender} -> {isa_receiver}") + raise ValueError("No matching inbound route found for this ISA pair") + + true_tenant_id = route.tenant_id + logger.info(f"Routed AS2 payload ({isa_sender}->{isa_receiver}) to Tenant {true_tenant_id}") + edi_record = { "trace_id": uuid.uuid4(), "direction": "INBOUND", @@ -306,37 +349,34 @@ async def _save_to_data_plane( # type: ignore "status": "RECEIVED", } - # Infrastructure logic encapsulated in the service layer + # 3. Save directly to the true Tenant's Data Plane Shard stmt = ( select(Tenant, DatabaseShard) .join(DatabaseShard, Tenant.shard_id == DatabaseShard.id) - .where(Tenant.id == partnership.tenant_id) + .where(Tenant.id == true_tenant_id) ) result = await self.global_session.execute(stmt) row = result.first() if not row: - logger.error(f"Tenant {partnership.tenant_id} not found in global DB") + logger.error(f"Tenant {true_tenant_id} not found in global DB") raise ValueError("Tenant routing failed") tenant, shard = row - async_gen_tenant = self.db_router.get_tenant_session( - partnership.tenant_id, shard.name, shard.dsn - ) + async_gen_tenant = self.db_router.get_tenant_session(true_tenant_id, shard.name, shard.dsn) tenant_session = await anext(async_gen_tenant) try: dp_repo = SqlAlchemyDataPlaneRepository(tenant_session) - msg_id = await dp_repo.create_edi_message( - tenant_id=partnership.tenant_id, payload=edi_record - ) + msg_id = await dp_repo.create_edi_message(tenant_id=true_tenant_id, payload=edi_record) outbox_payload = { "edi_message_id": str(msg_id), + "trace_id": str(edi_record["trace_id"]), "sender_id": as2_msg.as2_from, "receiver_id": as2_msg.as2_to, "status": "RECEIVED", } await dp_repo.create_outbox_event( - tenant_id=partnership.tenant_id, + tenant_id=true_tenant_id, event_type="edi_message.received", payload=outbox_payload, ) diff --git a/services/api/src/api/services/outbound_service.py b/services/api/src/api/services/outbound_service.py new file mode 100644 index 00000000..370eff84 --- /dev/null +++ b/services/api/src/api/services/outbound_service.py @@ -0,0 +1,172 @@ +import json +import logging +import uuid +from typing import Any, cast +from uuid import UUID + +from api.core.uow import UnitOfWork +from api.ports.repository import ControlPlaneRepositoryPort, DataPlaneRepositoryPort +from pipeline.core.metadata_extractor import MetadataExtractorService + +logger = logging.getLogger(__name__) + + +class OutboundService: + """ + Application Service (Use Case Layer) for handling outbound API requests. + Strictly follows Single Responsibility Principle and encapsulates business logic. + """ + + def __init__(self, uow: UnitOfWork): + self.uow = uow + self.extractor = MetadataExtractorService() + + async def process_outbound_message( + self, + tenant_id: int, + trading_partner_id: str, + payload: dict[str, Any] | list[dict[str, Any]], + transaction_type: str | None = None, + ) -> UUID: + """ + Orchestrates the outbound API flow: + 1. Validate Partnership by trading_partner_id. + 2. Extract Business Metadata. + 3. Save to EdiJson and ApiGateway. + 4. Drop Outbox event for Worker. + + Returns: + UUID: The generated trace_id for tracking. + """ + async with self.uow: + control_plane = cast(ControlPlaneRepositoryPort, self.uow.control_plane) + data_plane = cast(DataPlaneRepositoryPort, self.uow.data_plane) + + # 1. Validate Route by human-readable trading_partner_id (from ControlPlane/global DB) + logger.info(f"Received outbound transaction request for partner: {trading_partner_id}") + route = await control_plane.get_outbound_route_by_trading_partner_id( + tenant_id=tenant_id, trading_partner_id=trading_partner_id + ) + if not route: + logger.error(f"Outbound route '{trading_partner_id}' not found or not active") + raise ValueError(f"Outbound route '{trading_partner_id}' not found or not active") + + logger.info(f"Resolved OutboundRoute: {route.id} (Standard: {route.default_standard})") + + # 2. Resolve transaction_type from payload if not provided explicitly + if not transaction_type: + first_payload = ( + payload[0] + if isinstance(payload, list) and payload + else (payload if isinstance(payload, dict) else {}) + ) + transaction_type = first_payload.get("transaction_type") + if not transaction_type: + heading = first_payload.get("heading", {}) + for key in heading: + if key.startswith("transaction_set_header_ST"): + transaction_type = heading[key].get("transaction_set_identifier_code") + break + if not transaction_type: + # Try ST segment directly (for raw transaction payloads) + st = first_payload.get("ST", {}) + if st: + transaction_type = st.get("ST01") + + if not transaction_type: + transaction_type = route.transaction_type + + if not transaction_type: + logger.error( + "Could not determine transaction_type from payload or route configuration." + ) + raise ValueError("Transaction type could not be determined") + + business_metadata = {} + if isinstance(payload, dict): + business_metadata = self.extractor.extract(transaction_type or "", payload) + elif isinstance(payload, list) and len(payload) > 0: + business_metadata = self.extractor.extract(transaction_type or "", payload[0]) + + # 3. Create Trace ID + trace_id = uuid.uuid4() + logger.info(f"Generated Trace ID: {trace_id}") + + # 4. Save to ApiGateway (Logging) + api_gateway_payload = { + "trace_id": trace_id, + "direction": "OUTBOUND", + "transaction_type": transaction_type, + "payload": payload, + "response": json.dumps({"status": "ACCEPTED", "trace_id": str(trace_id)}), + "http_status_code": 202, + "status": "ACCEPTED", + } + await data_plane.create_api_gateway(tenant_id=tenant_id, payload=api_gateway_payload) + + # 5. Save to EdiJson + sender_id = route.isa_sender_id + receiver_id = route.isa_receiver_id + + # If payload doesn't already have an AST wrapper, wrap it in a Bots AST envelope + edi_json_data = payload + if ( + isinstance(payload, dict) + and "interchange_ISA" not in payload + and "interchange_UNB" not in payload + ) or isinstance(payload, list): + from transformer.domain.envelope_factory import EnvelopeFactory + + route_config = { + "default_standard": route.default_standard, + "default_version": route.default_version, + "transaction_type": transaction_type or "", + "isa_sender_qualifier": route.isa_sender_qualifier, + "isa_sender_id": route.isa_sender_id, + "isa_receiver_qualifier": route.isa_receiver_qualifier, + "isa_receiver_id": route.isa_receiver_id, + "gs_sender_id": route.gs_sender_id, + "gs_receiver_id": route.gs_receiver_id, + "environment": "P", # In future, derive from tenant/route config + } + + logger.info("Applying dynamic EnvelopeFactory AST wrapper to payload...") + edi_json_data = EnvelopeFactory.build_ast(route_config, payload) + logger.info( + f"Successfully wrapped payload in {route.default_standard.upper()} AST." + ) + + edi_json_payload = { + "trace_id": trace_id, + "direction": "OUTBOUND", + "outbound_route_id": route.id, + "transaction_type": transaction_type, + "standard": route.default_standard, + "sender_id": sender_id, + "receiver_id": receiver_id, + "business_metadata": business_metadata, + "payload": edi_json_data, + "status": "PENDING", + } + edi_json_id = await data_plane.create_edi_json( + tenant_id=tenant_id, payload=edi_json_payload + ) + + # 6. Create Outbox Event + outbox_payload = { + "edi_json_id": str(edi_json_id), + "trace_id": str(trace_id), + "outbound_route_id": str(route.id), + "status": "RECEIVED", + } + await data_plane.create_outbox_event( + tenant_id=tenant_id, + event_type="json.received", + payload=outbox_payload, + ) + + logger.info( + f"Committed Outbound Transaction {trace_id} to database. Sent outbox event." + ) + await self.uow.commit() + return trace_id diff --git a/services/api/tests/api_fakes.py b/services/api/tests/api_fakes.py index 1d691d19..a73abddb 100644 --- a/services/api/tests/api_fakes.py +++ b/services/api/tests/api_fakes.py @@ -89,6 +89,7 @@ class FakePartnership: id = p["id"] tenant_id = p["tenant_id"] name = p["cmd"].name + trading_partner_id = getattr(p["cmd"], "trading_partner_id", None) local_partner_id = p["cmd"].local_partner_id remote_partner_id = p["cmd"].remote_partner_id mdn_type = p["cmd"].mdn_type diff --git a/services/api/tests/test_api_repository.py b/services/api/tests/test_api_repository.py index 8f954cb7..3df20c36 100644 --- a/services/api/tests/test_api_repository.py +++ b/services/api/tests/test_api_repository.py @@ -5,7 +5,9 @@ os.environ["DB_ENCRYPTION_KEY"] = "sKkXvO6eX2Xo6-k2d_WqVf9j_w2_mCq7jR9b9w0wWf4=" import pytest -from api.adapters.repository import SqlAlchemyControlPlaneRepository, SqlAlchemyDataPlaneRepository +from api.adapters.repository import ( + SqlAlchemyControlPlaneRepository, +) from api.domain.models import ( CreateAS2PartnershipCmd, CreateAS2TradingPartnerCmd, @@ -47,7 +49,7 @@ def control_repo(global_session): @pytest.fixture def tenant_repo(tenant_session): - repo = SqlAlchemyDataPlaneRepository(tenant_session) + repo = SqlAlchemyControlPlaneRepository(tenant_session) repo._tenant_id = MagicMock(return_value=1) return repo @@ -71,7 +73,9 @@ async def test_control_plane_repository(control_repo: SqlAlchemyControlPlaneRepo # 3. Create Partnership p_cmd = CreateAS2PartnershipCmd( - name="Test Partnership", local_partner_id=p_id2, remote_partner_id=p_id1 + name="Test Partnership", + local_partner_id=p_id2, + remote_partner_id=p_id1, ) partnership_id = await control_repo.create_as2_partnership(tenant_id=1, cmd=p_cmd) @@ -124,10 +128,15 @@ async def test_data_plane_repository( webhook_id=wh_id, ) # Outbound route: deliver via sftp + import uuid + out_cmd = CreateOutboundRouteCmd( + trading_partner_id=str(uuid.uuid4()), name="Outbound Route 1", isa_sender_id="S", isa_receiver_id="R", + gs_sender_id="S", + gs_receiver_id="R", transaction_type="855", as2_partner_id=None, sftp_partner_id=sftp_id, @@ -167,5 +176,66 @@ async def test_get_as2_partner_tenant_isolation(control_repo: SqlAlchemyControlP # We can check that the SQL string contains the tenant_id binding compiled = str(call_args.compile(compile_kwargs={"literal_binds": True})) - assert "tenant_id =" in compiled - assert "1" in compiled + compiled_clean = compiled.replace(" ", "") + assert "tenant_idIN(1,0)" in compiled_clean or "tenant_id=" in compiled_clean + assert "1" in compiled_clean + + +@pytest.mark.asyncio +async def test_get_as2_partner_for_write() -> None: + from unittest.mock import AsyncMock, MagicMock + from uuid import UUID + + from api.adapters.repository import SqlAlchemyControlPlaneRepository + from database.models.control_plane import AS2Partner + + mock_session = AsyncMock() + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = AS2Partner( + id=UUID("00000000-0000-0000-0000-000000000000"), + tenant_id=1, + name="test", + as2_id="test", + ) + mock_session.execute.return_value = mock_result + + repo = SqlAlchemyControlPlaneRepository(mock_session) + partner = await repo.get_as2_partner_for_write(1, UUID("00000000-0000-0000-0000-000000000000")) + assert partner is not None + + mock_session.execute.assert_called_once() + stmt = mock_session.execute.call_args[0][0] + compiled_stmt = str(stmt.compile(compile_kwargs={"literal_binds": True})).replace("\n", "") + assert "tenant_id = 1" in compiled_stmt + assert "00000000000000000000000000000000" in compiled_stmt + + +@pytest.mark.asyncio +async def test_get_as2_partner() -> None: + from unittest.mock import AsyncMock, MagicMock + from uuid import UUID + + from api.adapters.repository import SqlAlchemyControlPlaneRepository + from database.models.control_plane import AS2Partner + + mock_session = AsyncMock() + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = AS2Partner( + id=UUID("00000000-0000-0000-0000-000000000000"), + tenant_id=1, + name="test", + as2_id="test", + ) + mock_session.execute.return_value = mock_result + + repo = SqlAlchemyControlPlaneRepository(mock_session) + partner = await repo.get_as2_partner(1, UUID("00000000-0000-0000-0000-000000000000")) + assert partner is not None + + mock_session.execute.assert_called_once() + stmt = mock_session.execute.call_args[0][0] + compiled_stmt = str(stmt.compile(compile_kwargs={"literal_binds": True})).replace("\n", "") + assert "tenant_id IS NULL" in compiled_stmt or "tenant_id IS NULL" in compiled_stmt.replace( + " ", " " + ) + assert "00000000000000000000000000000000" in compiled_stmt diff --git a/services/api/tests/test_api_token_service.py b/services/api/tests/test_api_token_service.py new file mode 100644 index 00000000..f16b63ee --- /dev/null +++ b/services/api/tests/test_api_token_service.py @@ -0,0 +1,72 @@ +from unittest.mock import AsyncMock +from uuid import uuid4 + +import pytest +from api.core.services.api_token_service import ApiTokenService, _generate_credentials +from api.domain.models import CreateApiTokenCmd + + +@pytest.mark.asyncio +async def test_api_token_service_create(): + mock_repo = AsyncMock() + token_id = uuid4() + mock_repo.create_api_token.return_value = token_id + + svc = ApiTokenService(mock_repo) + cmd = CreateApiTokenCmd(name="Test Token", expires_at=None) + + result = await svc.create_token(tenant_id=1, tenant_name="Acme Corp", cmd=cmd) + + assert result.id == token_id + assert result.tenant_id == 1 + assert result.name == "Test Token" + assert result.client_id.startswith("soopaedi_acmeco_") + assert len(result.client_secret) > 0 + assert result.active is True + + mock_repo.create_api_token.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_api_token_service_list(): + mock_repo = AsyncMock() + mock_repo.list_api_tokens.return_value = [{"id": str(uuid4())}] + + svc = ApiTokenService(mock_repo) + result = await svc.list_tokens(tenant_id=1) + + assert len(result) == 1 + mock_repo.list_api_tokens.assert_awaited_once_with(1) + + +@pytest.mark.asyncio +async def test_api_token_service_revoke(): + mock_repo = AsyncMock() + mock_repo.revoke_api_token.return_value = True + + svc = ApiTokenService(mock_repo) + t_id = uuid4() + result = await svc.revoke_token(tenant_id=1, token_id=t_id) + + assert result is True + mock_repo.revoke_api_token.assert_awaited_once_with(1, t_id) + + +@pytest.mark.asyncio +async def test_api_token_service_delete(): + mock_repo = AsyncMock() + mock_repo.delete_api_token.return_value = True + + svc = ApiTokenService(mock_repo) + t_id = uuid4() + result = await svc.delete_token(tenant_id=1, token_id=t_id) + + assert result is True + mock_repo.delete_api_token.assert_awaited_once_with(1, t_id) + + +def test_generate_credentials(): + c_id, c_secret, c_hash = _generate_credentials("TestTenant123") + assert c_id.startswith("soopaedi_testte_") + assert len(c_secret) > 32 + assert len(c_hash) == 64 diff --git a/services/api/tests/test_as2_partner_service.py b/services/api/tests/test_as2_partner_service.py new file mode 100644 index 00000000..82569eb1 --- /dev/null +++ b/services/api/tests/test_as2_partner_service.py @@ -0,0 +1,52 @@ +from unittest.mock import AsyncMock +from uuid import uuid4 + +import pytest +from api.core.services.as2_partner_service import AS2PartnerService +from api.domain.models import UpdateAS2TradingPartnerCmd + + +@pytest.mark.asyncio +async def test_update_as2_partner_not_found(): + mock_repo = AsyncMock() + mock_repo.get_as2_partner.return_value = None + + svc = AS2PartnerService(mock_repo) + cmd = UpdateAS2TradingPartnerCmd(name="Test") + + with pytest.raises(ValueError, match="Partner not found after update"): + await svc.update_as2_partner(1, uuid4(), cmd) + + +@pytest.mark.asyncio +async def test_rotate_certificates_success(): + mock_repo = AsyncMock() + mock_partner = AsyncMock() + mock_partner.name = "Test Partner" + mock_partner.active = True + mock_repo.get_as2_partner.return_value = mock_partner + + svc = AS2PartnerService(mock_repo) + partner_id = uuid4() + + result = await svc.rotate_certificates( + tenant_id=1, partner_id=partner_id, new_public_cert="cert", new_private_key_vault_ref="ref" + ) + + assert result.partner_id == partner_id + assert result.name == "Test Partner" + assert result.status == "ACTIVE" + + mock_repo.rotate_as2_certificates.assert_awaited_once_with(1, partner_id, "cert", "ref") + mock_repo.create_outbox_event.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_rotate_certificates_not_found(): + mock_repo = AsyncMock() + mock_repo.get_as2_partner.return_value = None + + svc = AS2PartnerService(mock_repo) + + with pytest.raises(ValueError, match="Partner not found after certificate rotation"): + await svc.rotate_certificates(1, uuid4(), "cert", None) diff --git a/services/api/tests/test_as2_receive_service.py b/services/api/tests/test_as2_receive_service.py index 243bad1a..316f4146 100644 --- a/services/api/tests/test_as2_receive_service.py +++ b/services/api/tests/test_as2_receive_service.py @@ -124,7 +124,11 @@ async def test_private_methods_coverage(service): return_value=MagicMock(first=MagicMock(return_value=None)) ) with pytest.raises(ValueError, match="Tenant routing failed"): - await service._save_to_data_plane(MagicMock(tenant_id=1), MagicMock(), b"") + await service._save_to_data_plane( + MagicMock(tenant_id=1), + MagicMock(), + b"ISA*00* *00* *ZZ*SENDER *ZZ*RECEIVER *210101*1200*^*00501*000000001*0*P*>~", + ) @pytest.mark.asyncio @@ -221,7 +225,11 @@ async def mock_async_gen(): ) mock_as2_msg = MagicMock(as2_from="ME", as2_to="YOU", message_id="msg-1") - res = await service._save_to_data_plane(mock_partnership, mock_as2_msg, b"EDI") + res = await service._save_to_data_plane( + mock_partnership, + mock_as2_msg, + b"ISA*00* *00* *ZZ*SENDER *ZZ*RECEIVER *210101*1200*^*00501*000000001*0*P*>~", + ) assert res == "msg-1" mock_repo.create_edi_message.assert_awaited_once() mock_repo.create_outbox_event.assert_awaited_once() diff --git a/services/api/tests/test_cdc_relay.py b/services/api/tests/test_cdc_relay.py index f98ad388..7366c426 100644 --- a/services/api/tests/test_cdc_relay.py +++ b/services/api/tests/test_cdc_relay.py @@ -39,7 +39,7 @@ def test_cdc_relay_successful_translate_routing(memory_queue: InMemoryQueueAdapt } response = client.post("/internal/cdc/relay", json=payload) - assert response.status_code == 202 + assert response.status_code == 200 assert len(memory_queue.sent_messages) == 1 queue_name, msg_payload = memory_queue.sent_messages[0] @@ -65,7 +65,7 @@ def test_cdc_relay_successful_deliver_routing(memory_queue: InMemoryQueueAdapter } response = client.post("/internal/cdc/relay", json=payload) - assert response.status_code == 202 + assert response.status_code == 200 assert len(memory_queue.sent_messages) == 1 queue_name, msg_payload = memory_queue.sent_messages[0] @@ -91,12 +91,12 @@ def test_cdc_relay_ignores_updates_and_deletes(memory_queue: InMemoryQueueAdapte } response = client.post("/internal/cdc/relay", json=payload) - assert response.status_code == 202 + assert response.status_code == 200 assert len(memory_queue.sent_messages) == 0 -def test_cdc_relay_rejects_unknown_table(memory_queue: InMemoryQueueAdapter) -> None: - """Tests the CDC relay fails explicitly on unknown table sources to prevent silent drops.""" +def test_cdc_relay_skips_unknown_table(memory_queue: InMemoryQueueAdapter) -> None: + """Tests the CDC relay skips unknown table sources to prevent silent drops.""" payload = { "__op": "c", "__table": "unknown_table", @@ -108,31 +108,39 @@ def test_cdc_relay_rejects_unknown_table(memory_queue: InMemoryQueueAdapter) -> } response = client.post("/internal/cdc/relay", json=payload) - assert response.status_code == 400 - assert response.json()["detail"] == "Unknown table source" + assert response.status_code == 200 + assert response.json() == {"status": "ok"} assert len(memory_queue.sent_messages) == 0 -def test_cdc_relay_rejects_unknown_event_type(memory_queue: InMemoryQueueAdapter) -> None: - """Unknown event types must fail-fast so they are not silently dropped.""" +def test_cdc_relay_successful_provisioning_routing(memory_queue: InMemoryQueueAdapter) -> None: + """Tests the CDC relay routes non-data plane events to ProvisioningQueue.""" payload = { "__op": "c", "__table": "outbox", "idempotency_key": "uuid-789", - "event_type": "MYSTERY_EVENT", - "payload": {"trace_id": "req-123"}, + "event_type": "AS2_PARTNERSHIP_CREATED", + "payload": {"tenant_id": 999}, "status": "PENDING", "tenant_id": 999, } response = client.post("/internal/cdc/relay", json=payload) - assert response.status_code == 400 - assert "MYSTERY_EVENT" in response.json()["detail"] - assert len(memory_queue.sent_messages) == 0 + assert response.status_code == 200 + + assert len(memory_queue.sent_messages) == 1 + queue_name, msg_payload = memory_queue.sent_messages[0] + assert queue_name == "ProvisioningQueue" + assert msg_payload == { + "idempotency_key": "uuid-789", + "event_type": "AS2_PARTNERSHIP_CREATED", + "payload": {"tenant_id": 999}, + "tenant_id": 999, + } -def test_cdc_relay_rejects_missing_trace_id(memory_queue: InMemoryQueueAdapter) -> None: - """Payloads without trace_id must be rejected to prevent poison messages in SQS.""" +def test_cdc_relay_skips_missing_trace_id(memory_queue: InMemoryQueueAdapter) -> None: + """Payloads without trace_id must be skipped to prevent poison messages in SQS.""" payload = { "__op": "c", "__table": "outbox", @@ -144,6 +152,6 @@ def test_cdc_relay_rejects_missing_trace_id(memory_queue: InMemoryQueueAdapter) } response = client.post("/internal/cdc/relay", json=payload) - assert response.status_code == 400 - assert "trace_id" in response.json()["detail"] + assert response.status_code == 200 + assert response.json() == {"status": "ok"} assert len(memory_queue.sent_messages) == 0 diff --git a/services/api/tests/test_outbound_service.py b/services/api/tests/test_outbound_service.py new file mode 100644 index 00000000..68c234e4 --- /dev/null +++ b/services/api/tests/test_outbound_service.py @@ -0,0 +1,37 @@ +from unittest.mock import AsyncMock, MagicMock +from uuid import uuid4 + +import pytest +from api.services.outbound_service import OutboundService + + +@pytest.mark.asyncio +async def test_process_outbound_message_route_not_found(): + mock_uow = AsyncMock() + mock_uow.control_plane.get_outbound_route_by_trading_partner_id.return_value = None + + svc = OutboundService(mock_uow) + with pytest.raises(ValueError, match="not found or not active"): + await svc.process_outbound_message(1, "PARTNER_X", {}) + + +@pytest.mark.asyncio +async def test_process_outbound_message_success(): + mock_uow = AsyncMock() + mock_route = MagicMock( + id=uuid4(), + default_standard="x12", + as2_partner_id=uuid4(), + sftp_partner_id=None, + isa_sender_id="S", + isa_receiver_id="R", + ) + mock_uow.control_plane.get_outbound_route_by_trading_partner_id.return_value = mock_route + + svc = OutboundService(mock_uow) + trace_id = await svc.process_outbound_message(1, "PARTNER_X", {"transaction_type": "850"}) + + assert trace_id is not None + mock_uow.data_plane.create_edi_json.assert_awaited_once() + mock_uow.data_plane.create_api_gateway.assert_awaited_once() + mock_uow.data_plane.create_outbox_event.assert_awaited_once() diff --git a/services/api/tests/test_provisioning_core.py b/services/api/tests/test_provisioning_core.py index e7ce304c..3aeb20c8 100644 --- a/services/api/tests/test_provisioning_core.py +++ b/services/api/tests/test_provisioning_core.py @@ -1,7 +1,13 @@ import uuid import pytest -from api.core.provisioning import ProvisioningService +from api.core.services import ( + AS2PartnerService, + AS2PartnershipService, + RouteService, + SFTPPartnerService, + WebhookService, +) from api.domain.models import ( CreateAS2TradingPartnerCmd, CreateInboundRouteCmd, @@ -10,26 +16,48 @@ CreateWebhookCmd, UpdateAS2TradingPartnerCmd, ) -from api_fakes import FakeControlPlaneRepository, FakeDataPlaneRepository +from api_fakes import FakeControlPlaneRepository + + +@pytest.fixture +def global_repo(): + return FakeControlPlaneRepository() + + +@pytest.fixture +def as2_partner_service(global_repo): + return AS2PartnerService(global_repo=global_repo) + + +@pytest.fixture +def as2_partnership_service(global_repo): + return AS2PartnershipService(global_repo=global_repo) @pytest.fixture -def service(): - global_repo = FakeControlPlaneRepository() - tenant_repo = FakeDataPlaneRepository() - return ProvisioningService(global_repo=global_repo, tenant_repo=tenant_repo) +def sftp_partner_service(global_repo): + return SFTPPartnerService(global_repo=global_repo) + + +@pytest.fixture +def webhook_service(global_repo): + return WebhookService(global_repo=global_repo) + + +@pytest.fixture +def route_service(global_repo): + return RouteService(global_repo=global_repo) @pytest.mark.asyncio -async def test_create_as2_partner(service: ProvisioningService): +async def test_create_as2_partner(as2_partner_service: AS2PartnerService, global_repo): cmd = CreateAS2TradingPartnerCmd(name="Test Partner", as2_id="TEST_AS2") - partner = await service.create_as2_partner(tenant_id=1, cmd=cmd) + partner = await as2_partner_service.create_as2_partner(tenant_id=1, cmd=cmd) assert partner.type == "AS2" assert partner.tenant_id == 1 assert partner.status == "PROVISIONING" - global_repo: FakeControlPlaneRepository = service.global_repo assert len(global_repo.partners) == 1 assert len(global_repo.outbox_events) == 1 assert global_repo.outbox_events[0]["event_type"] == "AS2_PARTNER_CREATED" @@ -37,58 +65,51 @@ async def test_create_as2_partner(service: ProvisioningService): @pytest.mark.asyncio -async def test_update_and_delete_as2_partner(service: ProvisioningService): +async def test_update_and_delete_as2_partner(as2_partner_service: AS2PartnerService, global_repo): cmd = CreateAS2TradingPartnerCmd(name="Test Partner", as2_id="TEST_AS2") - partner = await service.create_as2_partner(tenant_id=1, cmd=cmd) + partner = await as2_partner_service.create_as2_partner(tenant_id=1, cmd=cmd) # Update update_cmd = UpdateAS2TradingPartnerCmd(name="Updated Partner", as2_id="NEW_AS2") - updated = await service.update_as2_partner( + updated = await as2_partner_service.update_as2_partner( tenant_id=1, partner_id=partner.partner_id, cmd=update_cmd ) assert updated.name == "Updated Partner" # Delete - await service.delete_as2_partner(tenant_id=1, partner_id=partner.partner_id) - assert len(service.global_repo.partners) == 0 + await as2_partner_service.delete_as2_partner(tenant_id=1, partner_id=partner.partner_id) + assert len(global_repo.partners) == 0 @pytest.mark.asyncio -async def test_create_sftp_partner(service: ProvisioningService): +async def test_create_sftp_partner(sftp_partner_service: SFTPPartnerService, global_repo): cmd = CreateSFTPPartnerCmd( name="SFTP Partner", host="sftp.example.com", username="user", credentials_vault_ref="vault-ref", ) - partner = await service.create_sftp_partner(tenant_id=1, cmd=cmd) + partner = await sftp_partner_service.create_sftp_partner(tenant_id=1, cmd=cmd) assert partner.type == "SFTP" assert partner.status == "INACTIVE" - - global_repo: FakeControlPlaneRepository = service.global_repo assert len(global_repo.sftp_partners) == 1 @pytest.mark.asyncio -async def test_create_webhook_partner(service: ProvisioningService): +async def test_create_webhook(webhook_service: WebhookService, global_repo): cmd = CreateWebhookCmd( name="Webhook Partner", url="https://example.com/webhook", auth_header_vault_ref="vault-ref" ) - partner = await service.create_webhook(tenant_id=1, cmd=cmd) + partner = await webhook_service.create_webhook(tenant_id=1, cmd=cmd) assert partner.type == "WEBHOOK" assert partner.status == "ACTIVE" - - global_repo: FakeControlPlaneRepository = service.global_repo assert len(global_repo.webhooks) == 1 @pytest.mark.asyncio -async def test_list_routes(service: ProvisioningService): - global_repo: FakeControlPlaneRepository = service.global_repo - - # 1. Create a fake AS2 partner for name resolution +async def test_list_routes(route_service: RouteService, global_repo): as2_id = await global_repo.create_as2_identity( tenant_id=1, cmd=CreateAS2TradingPartnerCmd(name="Walmart", as2_id="WM") ) @@ -113,16 +134,19 @@ def __init__(self, id, as2_partner_id, sftp_partner_id, webhook_id): self.webhook_id = webhook_id self.isa_sender_id = "S1" self.isa_receiver_id = "R1" + self.gs_sender_id = "S1" + self.gs_receiver_id = "R1" self.transaction_type = "850" + self.trading_partner_id = uuid.uuid4() + self.isa_sender_qualifier = "ZZ" + self.isa_receiver_qualifier = "ZZ" + self.default_standard = "x12" + self.default_version = "004010" - # Mocking what the repository would return (objects with properties) - inbound_route = FakeRoute(uuid.uuid4(), as2_id, sftp_id, None) - outbound_route = FakeRoute(uuid.uuid4(), as2_id, None, None) - - global_repo.inbound_routes = [inbound_route] - global_repo.outbound_routes = [outbound_route] + global_repo.inbound_routes = [FakeRoute(uuid.uuid4(), as2_id, sftp_id, None)] + global_repo.outbound_routes = [FakeRoute(uuid.uuid4(), as2_id, None, None)] - routes = await service.list_routes(1) + routes = await route_service.list_routes(1) assert len(routes) == 2 inbound_res = next(r for r in routes if r["direction"] == "INBOUND") @@ -133,7 +157,7 @@ def __init__(self, id, as2_partner_id, sftp_partner_id, webhook_id): @pytest.mark.asyncio -async def test_create_inbound_route(service: ProvisioningService): +async def test_create_inbound_route(route_service: RouteService, global_repo): cmd = CreateInboundRouteCmd( name="Inbound Route", isa_sender_id="S1", @@ -141,24 +165,25 @@ async def test_create_inbound_route(service: ProvisioningService): transaction_type="850", as2_partner_id=uuid.uuid4(), ) - route = await service.create_inbound_route(tenant_id=1, cmd=cmd) + route = await route_service.create_inbound_route(tenant_id=1, cmd=cmd) assert route.direction == "INBOUND" - global_repo: FakeControlPlaneRepository = service.global_repo assert len(global_repo.inbound_routes) == 1 @pytest.mark.asyncio -async def test_create_outbound_route(service: ProvisioningService): +async def test_create_outbound_route(route_service: RouteService, global_repo): cmd = CreateOutboundRouteCmd( + trading_partner_id=str(uuid.uuid4()), name="Outbound Route", isa_sender_id="S1", isa_receiver_id="R1", + gs_sender_id="S1", + gs_receiver_id="R1", transaction_type="855", as2_partner_id=uuid.uuid4(), ) - route = await service.create_outbound_route(tenant_id=1, cmd=cmd) + route = await route_service.create_outbound_route(tenant_id=1, cmd=cmd) assert route.direction == "OUTBOUND" - global_repo: FakeControlPlaneRepository = service.global_repo assert len(global_repo.outbound_routes) == 1 diff --git a/services/api/tests/test_routers_api_tokens.py b/services/api/tests/test_routers_api_tokens.py new file mode 100644 index 00000000..7bb8734e --- /dev/null +++ b/services/api/tests/test_routers_api_tokens.py @@ -0,0 +1,97 @@ +from unittest.mock import AsyncMock +from uuid import uuid4 + +import pytest +from api.dependencies import get_api_token_repo +from api.routers.developers.api_tokens import router +from fastapi import FastAPI +from fastapi.testclient import TestClient +from identity.dependencies import get_current_tenant_id + +app = FastAPI() +app.include_router(router) + + +def override_get_tenant_id(): + return 1 + + +@pytest.fixture +def mock_repo(): + return AsyncMock() + + +@pytest.fixture +def client(mock_repo): + app.dependency_overrides[get_current_tenant_id] = override_get_tenant_id + app.dependency_overrides[get_api_token_repo] = lambda: mock_repo + yield TestClient(app) + app.dependency_overrides.clear() + + +def test_create_api_token(client, mock_repo): + token_id = uuid4() + mock_repo.create_api_token.return_value = token_id + + response = client.post("/api/v1/developers/tokens", json={"name": "Test Token"}) + + assert response.status_code == 201 + data = response.json() + assert data["id"] == str(token_id) + assert data["name"] == "Test Token" + assert "client_id" in data + assert "client_secret" in data + assert data["active"] is True + + +def test_list_api_tokens(client, mock_repo): + mock_repo.list_api_tokens.return_value = [ + { + "id": str(uuid4()), + "name": "Test Token", + "client_id": "client_1", + "active": True, + "created_at": "2023-01-01T00:00:00Z", + "last_used_at": None, + "expires_at": None, + } + ] + + response = client.get("/api/v1/developers/tokens") + + assert response.status_code == 200 + data = response.json() + assert len(data["tokens"]) == 1 + assert data["tokens"][0]["name"] == "Test Token" + + +def test_revoke_api_token(client, mock_repo): + mock_repo.revoke_api_token.return_value = True + t_id = uuid4() + + response = client.delete(f"/api/v1/developers/tokens/{t_id}") + assert response.status_code == 204 + + +def test_revoke_api_token_not_found(client, mock_repo): + mock_repo.revoke_api_token.return_value = False + t_id = uuid4() + + response = client.delete(f"/api/v1/developers/tokens/{t_id}") + assert response.status_code == 404 + + +def test_delete_api_token(client, mock_repo): + mock_repo.delete_api_token.return_value = True + t_id = uuid4() + + response = client.delete(f"/api/v1/developers/tokens/{t_id}/hard") + assert response.status_code == 204 + + +def test_delete_api_token_not_found(client, mock_repo): + mock_repo.delete_api_token.return_value = False + t_id = uuid4() + + response = client.delete(f"/api/v1/developers/tokens/{t_id}/hard") + assert response.status_code == 404 diff --git a/services/api/tests/test_routers_partners.py b/services/api/tests/test_routers_partners.py index 489459db..42f39c7c 100644 --- a/services/api/tests/test_routers_partners.py +++ b/services/api/tests/test_routers_partners.py @@ -119,10 +119,20 @@ def test_update_platform_as2_partner(client, fake_uow): def test_create_platform_as2_partnership(client, fake_uow): - import uuid - local_id = str(uuid.uuid4()) - remote_id = str(uuid.uuid4()) + # Create local and remote partners first + loc_resp = client.post( + "/api/v1/platform/trading-partners/as2/trading-partners", + json={"name": "Local Partner", "as2_id": "LOCAL_AS2", "is_local": True}, + ) + local_id = loc_resp.json()["id"] + + rem_resp = client.post( + "/api/v1/platform/trading-partners/as2/trading-partners", + json={"name": "Remote Partner", "as2_id": "REMOTE_AS2", "is_local": False}, + ) + remote_id = rem_resp.json()["id"] + response = client.post( "/api/v1/platform/trading-partners/as2/partnerships", json={ @@ -144,14 +154,26 @@ def test_list_platform_as2_partnerships(client): def test_update_platform_as2_partnership(client, fake_uow): - import uuid + + # Create local and remote partners first + loc_resp = client.post( + "/api/v1/platform/trading-partners/as2/trading-partners", + json={"name": "Local Update Partner", "as2_id": "LOC_UPD", "is_local": True}, + ) + local_id = loc_resp.json()["id"] + + rem_resp = client.post( + "/api/v1/platform/trading-partners/as2/trading-partners", + json={"name": "Remote Update Partner", "as2_id": "REM_UPD", "is_local": False}, + ) + remote_id = rem_resp.json()["id"] ps_id = client.post( "/api/v1/platform/trading-partners/as2/partnerships", json={ "name": "Temp Partnership", - "local_partner_id": str(uuid.uuid4()), - "remote_partner_id": str(uuid.uuid4()), + "local_partner_id": local_id, + "remote_partner_id": remote_id, }, ).json()["id"] @@ -310,7 +332,7 @@ def test_existing_sftp_connection_failures(client, fake_uow): from sqlalchemy.exc import IntegrityError with patch( - "api.core.provisioning.ProvisioningService.create_sftp_partner", + "api.core.services.sftp_partner_service.SFTPPartnerService.create_sftp_partner", side_effect=ValueError("Bad value"), ): resp = client.post( @@ -319,7 +341,7 @@ def test_existing_sftp_connection_failures(client, fake_uow): ) assert resp.status_code == 400 with patch( - "api.core.provisioning.ProvisioningService.create_sftp_partner", + "api.core.services.sftp_partner_service.SFTPPartnerService.create_sftp_partner", side_effect=IntegrityError("x", "y", "z"), ): resp = client.post( @@ -329,7 +351,7 @@ def test_existing_sftp_connection_failures(client, fake_uow): assert resp.status_code == 400 with patch( - "api.core.provisioning.ProvisioningService.update_sftp_partner", + "api.core.services.sftp_partner_service.SFTPPartnerService.update_sftp_partner", side_effect=ValueError("Bad value"), ): resp = client.put( @@ -338,7 +360,7 @@ def test_existing_sftp_connection_failures(client, fake_uow): ) assert resp.status_code == 400 with patch( - "api.core.provisioning.ProvisioningService.update_sftp_partner", + "api.core.services.sftp_partner_service.SFTPPartnerService.update_sftp_partner", side_effect=IntegrityError("x", "y", "z"), ): resp = client.put( diff --git a/services/api/tests/test_routers_routes.py b/services/api/tests/test_routers_routes.py index 2bb488fe..8dc95114 100644 --- a/services/api/tests/test_routers_routes.py +++ b/services/api/tests/test_routers_routes.py @@ -55,6 +55,9 @@ def test_create_outbound_route(client): "isa_sender_id": "S1", "isa_receiver_id": "R1", "transaction_type": "855", + "trading_partner_id": "TP1", + "gs_sender_id": "GS1", + "gs_receiver_id": "GR1", }, ) assert response.status_code == 201 diff --git a/services/worker/pyproject.toml b/services/worker/pyproject.toml index 369e6b0c..e9e7d7b7 100644 --- a/services/worker/pyproject.toml +++ b/services/worker/pyproject.toml @@ -12,12 +12,14 @@ dependencies = [ "database", "pipeline", "config", + "domain", ] [tool.uv.sources] database = { workspace = true } pipeline = { workspace = true } config = { workspace = true } +domain = { workspace = true } [build-system] requires = ["hatchling"] diff --git a/services/worker/src/worker/adapters/__init__.py b/services/worker/src/worker/adapters/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/services/worker/src/worker/adapters/db_outbox.py b/services/worker/src/worker/adapters/db_outbox.py new file mode 100644 index 00000000..2449e993 --- /dev/null +++ b/services/worker/src/worker/adapters/db_outbox.py @@ -0,0 +1,61 @@ +import contextlib +import logging +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager + +from database.connection import DatabaseRouter +from database.models.control_plane import Outbox as GlobalOutbox +from domain.events import ProvisioningEventType +from sqlalchemy import select + +from worker.core.errors import PermanentProvisioningError +from worker.ports.outbox import OutboxEvent, OutboxPort + +logger = logging.getLogger(__name__) + + +class SqlAlchemyOutboxAdapter(OutboxPort): + def __init__(self, db_router: DatabaseRouter): + self.db_router = db_router + + @asynccontextmanager + async def process_next_event(self) -> AsyncIterator[OutboxEvent | None]: + global_gen = self.db_router.get_global_session() + global_session = await global_gen.__anext__() + + try: + stmt = ( + select(GlobalOutbox) + .where( + GlobalOutbox.status == "PENDING", + GlobalOutbox.event_type.in_(list(ProvisioningEventType)), + ) + .limit(1) + .with_for_update(skip_locked=True) + ) + result = await global_session.execute(stmt) + event = result.scalar_one_or_none() + + if not event: + yield None + return + + try: + yield event + event.status = "PROCESSED" + await global_session.commit() + except Exception as e: + if isinstance(e, PermanentProvisioningError): + logger.error( + f"Permanent error processing event {event.id}: {e}. Marking as FAILED." + ) + event.status = "FAILED" + await global_session.commit() + else: + logger.exception( + f"Transient error processing event {event.id}: {e}. Leaving as PENDING." + ) + await global_session.rollback() + finally: + with contextlib.suppress(StopAsyncIteration): + await global_gen.__anext__() diff --git a/services/worker/src/worker/adapters/db_replication.py b/services/worker/src/worker/adapters/db_replication.py new file mode 100644 index 00000000..5e3c9fb3 --- /dev/null +++ b/services/worker/src/worker/adapters/db_replication.py @@ -0,0 +1,417 @@ +import logging +from typing import Any + +from database.connection import DatabaseRouter +from database.models.control_plane import AS2Partner as GlobalAS2Partner +from database.models.control_plane import AS2Partnership as GlobalAS2Partnership +from database.models.control_plane import InboundRoute as GlobalInboundRoute +from database.models.control_plane import OutboundRoute as GlobalOutboundRoute +from database.models.control_plane import SFTPPartner as GlobalSFTPPartner +from database.models.control_plane import Webhook as GlobalWebhook +from database.models.data_plane import AS2Partner as TenantAS2Partner +from database.models.data_plane import AS2Partnership as TenantAS2Partnership +from database.models.data_plane import InboundRoute as TenantInboundRoute +from database.models.data_plane import OutboundRoute as TenantOutboundRoute +from database.models.data_plane import SFTPPartner as TenantSFTPPartner +from database.models.data_plane import Webhook as TenantWebhook +from sqlalchemy import delete, select +from sqlalchemy.dialects.postgresql import insert + +from worker.core.errors import PermanentProvisioningError, TransientProvisioningError +from worker.ports.replication import ReplicationPort +from worker.ports.tenant import TenantPort + +logger = logging.getLogger(__name__) + + +class SqlAlchemyReplicationAdapter(ReplicationPort): + def __init__(self, db_router: DatabaseRouter, tenant_port: TenantPort): + self.db_router = db_router + self.tenant_port = tenant_port + + async def replicate_tenant_configuration(self, tenant_id: int) -> None: + """Copy all relevant configuration from Global DB to Tenant DB Shard.""" + + try: + shard_name, shard_dsn = await self.tenant_port.resolve_shard(tenant_id) + except Exception as e: + raise PermanentProvisioningError(f"Tenant {tenant_id} unresolvable: {e}") from e + + global_gen = self.db_router.get_global_session() + tenant_gen = self.db_router.get_tenant_session(tenant_id, shard_name, shard_dsn) + + from contextlib import aclosing + + async with aclosing(global_gen) as global_gen_ctx, aclosing(tenant_gen) as tenant_gen_ctx: + global_session = await global_gen_ctx.__anext__() + tenant_session = await tenant_gen_ctx.__anext__() + + try: + await self._do_replicate(tenant_id, global_session, tenant_session) + await tenant_session.commit() + logger.info(f"Successfully replicated configuration for tenant_id={tenant_id}") + except Exception as e: + await tenant_session.rollback() + raise TransientProvisioningError( + f"Failed to replicate tenant {tenant_id}: {e}" + ) from e + + async def _do_replicate(self, tenant_id: int, global_session: Any, tenant_session: Any) -> None: + # --- AS2 Partners --- + stmt = select(GlobalAS2Partner).where( + (GlobalAS2Partner.tenant_id == tenant_id) | (GlobalAS2Partner.tenant_id == 0) + ) + tp_result = await global_session.execute(stmt) + as2_partners = tp_result.scalars().all() + logger.info(f"[tenant={tenant_id}] Replicating {len(as2_partners)} AS2 partner(s)") + + for global_tp in as2_partners: + logger.debug( + f"[tenant={tenant_id}] Upserting AS2Partner id={global_tp.id} as2_id={global_tp.as2_id!r}" + ) + insert_stmt = ( + insert(TenantAS2Partner) + .values( + id=global_tp.id, + tenant_id=tenant_id, + name=global_tp.name, + as2_id=global_tp.as2_id, + public_cert_pem=global_tp.public_cert_pem, + public_cert_vault_ref=global_tp.public_cert_vault_ref, + private_key_vault_ref=global_tp.private_key_vault_ref, + prev_public_cert_pem=global_tp.prev_public_cert_pem, + prev_public_cert_vault_ref=global_tp.prev_public_cert_vault_ref, + prev_private_key_vault_ref=global_tp.prev_private_key_vault_ref, + url=global_tp.url, + active=global_tp.active, + created_at=global_tp.created_at, + updated_at=global_tp.updated_at, + ) + .on_conflict_do_update( + index_elements=["id"], + set_={ + "tenant_id": tenant_id, + "name": global_tp.name, + "as2_id": global_tp.as2_id, + "public_cert_pem": global_tp.public_cert_pem, + "public_cert_vault_ref": global_tp.public_cert_vault_ref, + "private_key_vault_ref": global_tp.private_key_vault_ref, + "prev_public_cert_pem": global_tp.prev_public_cert_pem, + "prev_public_cert_vault_ref": global_tp.prev_public_cert_vault_ref, + "prev_private_key_vault_ref": global_tp.prev_private_key_vault_ref, + "url": global_tp.url, + "active": global_tp.active, + "created_at": global_tp.created_at, + "updated_at": global_tp.updated_at, + }, + ) + ) + await tenant_session.execute(insert_stmt) + + # --- AS2 Partnerships --- + ps_stmt = select(GlobalAS2Partnership).where( + (GlobalAS2Partnership.tenant_id == tenant_id) | (GlobalAS2Partnership.tenant_id == 0) + ) + ps_result = await global_session.execute(ps_stmt) + as2_partnerships = ps_result.scalars().all() + logger.info(f"[tenant={tenant_id}] Replicating {len(as2_partnerships)} AS2 partnership(s)") + + for global_ps in as2_partnerships: + logger.debug( + f"[tenant={tenant_id}] Upserting AS2Partnership id={global_ps.id} name={global_ps.name!r}" + ) + insert_ps_stmt = ( + insert(TenantAS2Partnership) + .values( + id=global_ps.id, + tenant_id=tenant_id, + name=global_ps.name, + local_partner_id=global_ps.local_partner_id, + remote_partner_id=global_ps.remote_partner_id, + credentials_vault_ref=global_ps.credentials_vault_ref, + mdn_type=global_ps.mdn_type, + mdn_url=global_ps.mdn_url, + encryption_algorithm=global_ps.encryption_algorithm, + signature_algorithm=global_ps.signature_algorithm, + advanced_flags=global_ps.advanced_flags, + active=global_ps.active, + created_at=global_ps.created_at, + updated_at=global_ps.updated_at, + ) + .on_conflict_do_update( + index_elements=["id"], + set_={ + "name": global_ps.name, + "local_partner_id": global_ps.local_partner_id, + "remote_partner_id": global_ps.remote_partner_id, + "credentials_vault_ref": global_ps.credentials_vault_ref, + "mdn_type": global_ps.mdn_type, + "mdn_url": global_ps.mdn_url, + "encryption_algorithm": global_ps.encryption_algorithm, + "signature_algorithm": global_ps.signature_algorithm, + "advanced_flags": global_ps.advanced_flags, + "active": global_ps.active, + "created_at": global_ps.created_at, + "updated_at": global_ps.updated_at, + }, + ) + ) + await tenant_session.execute(insert_ps_stmt) + + # --- SFTP Partners --- + sftp_stmt = select(GlobalSFTPPartner).where(GlobalSFTPPartner.tenant_id == tenant_id) + sftp_result = await global_session.execute(sftp_stmt) + sftp_partners = sftp_result.scalars().all() + logger.info(f"[tenant={tenant_id}] Replicating {len(sftp_partners)} SFTP partner(s)") + + for global_sftp in sftp_partners: + logger.debug( + f"[tenant={tenant_id}] Upserting SFTPPartner id={global_sftp.id} name={global_sftp.name!r}" + ) + insert_sftp_stmt = ( + insert(TenantSFTPPartner) + .values( + id=global_sftp.id, + tenant_id=tenant_id, + name=global_sftp.name, + host=global_sftp.host, + port=global_sftp.port, + username=global_sftp.username, + inbound_remote_path=global_sftp.inbound_remote_path, + outbound_remote_path=global_sftp.outbound_remote_path, + host_key=global_sftp.host_key, + password_encrypted=global_sftp.password_encrypted, + credentials_vault_ref=global_sftp.credentials_vault_ref, + active=global_sftp.active, + created_at=global_sftp.created_at, + updated_at=global_sftp.updated_at, + ) + .on_conflict_do_update( + index_elements=["id"], + set_={ + "name": global_sftp.name, + "host": global_sftp.host, + "port": global_sftp.port, + "username": global_sftp.username, + "inbound_remote_path": global_sftp.inbound_remote_path, + "outbound_remote_path": global_sftp.outbound_remote_path, + "host_key": global_sftp.host_key, + "password_encrypted": global_sftp.password_encrypted, + "credentials_vault_ref": global_sftp.credentials_vault_ref, + "active": global_sftp.active, + "created_at": global_sftp.created_at, + "updated_at": global_sftp.updated_at, + }, + ) + ) + await tenant_session.execute(insert_sftp_stmt) + + # --- Webhooks --- + wh_stmt = select(GlobalWebhook).where(GlobalWebhook.tenant_id == tenant_id) + wh_result = await global_session.execute(wh_stmt) + webhooks = wh_result.scalars().all() + logger.info(f"[tenant={tenant_id}] Replicating {len(webhooks)} webhook(s)") + + for global_wh in webhooks: + logger.debug( + f"[tenant={tenant_id}] Upserting Webhook id={global_wh.id} name={global_wh.name!r}" + ) + insert_wh_stmt = ( + insert(TenantWebhook) + .values( + id=global_wh.id, + tenant_id=tenant_id, + name=global_wh.name, + url=global_wh.url, + auth_header_vault_ref=global_wh.auth_header_vault_ref, + active=global_wh.active, + created_at=global_wh.created_at, + updated_at=global_wh.updated_at, + ) + .on_conflict_do_update( + index_elements=["id"], + set_={ + "name": global_wh.name, + "url": global_wh.url, + "auth_header_vault_ref": global_wh.auth_header_vault_ref, + "active": global_wh.active, + "created_at": global_wh.created_at, + "updated_at": global_wh.updated_at, + }, + ) + ) + await tenant_session.execute(insert_wh_stmt) + + # --- Inbound Routes --- + ir_stmt = select(GlobalInboundRoute).where(GlobalInboundRoute.tenant_id == tenant_id) + ir_result = await global_session.execute(ir_stmt) + inbound_routes = ir_result.scalars().all() + logger.info(f"[tenant={tenant_id}] Replicating {len(inbound_routes)} inbound route(s)") + + for global_ir in inbound_routes: + logger.debug( + f"[tenant={tenant_id}] Upserting InboundRoute id={global_ir.id} name={global_ir.name!r}" + ) + insert_ir_stmt = ( + insert(TenantInboundRoute) + .values( + id=global_ir.id, + tenant_id=tenant_id, + name=global_ir.name, + isa_sender_id=global_ir.isa_sender_id, + isa_receiver_id=global_ir.isa_receiver_id, + transaction_type=global_ir.transaction_type, + processing_mode=global_ir.processing_mode, + webhook_id=global_ir.webhook_id, + as2_partner_id=global_ir.as2_partner_id, + sftp_partner_id=global_ir.sftp_partner_id, + active=global_ir.active, + created_at=global_ir.created_at, + updated_at=global_ir.updated_at, + ) + .on_conflict_do_update( + index_elements=["id"], + set_={ + "name": global_ir.name, + "isa_sender_id": global_ir.isa_sender_id, + "isa_receiver_id": global_ir.isa_receiver_id, + "transaction_type": global_ir.transaction_type, + "processing_mode": global_ir.processing_mode, + "webhook_id": global_ir.webhook_id, + "as2_partner_id": global_ir.as2_partner_id, + "sftp_partner_id": global_ir.sftp_partner_id, + "active": global_ir.active, + "created_at": global_ir.created_at, + "updated_at": global_ir.updated_at, + }, + ) + ) + await tenant_session.execute(insert_ir_stmt) + + # --- Outbound Routes --- + or_stmt = select(GlobalOutboundRoute).where(GlobalOutboundRoute.tenant_id == tenant_id) + or_result = await global_session.execute(or_stmt) + outbound_routes = or_result.scalars().all() + logger.info(f"[tenant={tenant_id}] Replicating {len(outbound_routes)} outbound route(s)") + + for global_or in outbound_routes: + logger.debug( + f"[tenant={tenant_id}] Upserting OutboundRoute id={global_or.id} name={global_or.name!r}" + ) + insert_or_stmt = ( + insert(TenantOutboundRoute) + .values( + id=global_or.id, + tenant_id=tenant_id, + trading_partner_id=global_or.trading_partner_id, + name=global_or.name, + isa_sender_id=global_or.isa_sender_id, + isa_sender_qualifier=global_or.isa_sender_qualifier, + isa_receiver_id=global_or.isa_receiver_id, + isa_receiver_qualifier=global_or.isa_receiver_qualifier, + gs_sender_id=global_or.gs_sender_id, + gs_receiver_id=global_or.gs_receiver_id, + default_standard=global_or.default_standard, + default_version=global_or.default_version, + transaction_type=global_or.transaction_type, + processing_mode=global_or.processing_mode, + as2_partner_id=global_or.as2_partner_id, + sftp_partner_id=global_or.sftp_partner_id, + active=global_or.active, + created_at=global_or.created_at, + updated_at=global_or.updated_at, + ) + .on_conflict_do_update( + index_elements=["id"], + set_={ + "trading_partner_id": global_or.trading_partner_id, + "name": global_or.name, + "isa_sender_id": global_or.isa_sender_id, + "isa_sender_qualifier": global_or.isa_sender_qualifier, + "isa_receiver_id": global_or.isa_receiver_id, + "isa_receiver_qualifier": global_or.isa_receiver_qualifier, + "gs_sender_id": global_or.gs_sender_id, + "gs_receiver_id": global_or.gs_receiver_id, + "default_standard": global_or.default_standard, + "default_version": global_or.default_version, + "transaction_type": global_or.transaction_type, + "processing_mode": global_or.processing_mode, + "as2_partner_id": global_or.as2_partner_id, + "sftp_partner_id": global_or.sftp_partner_id, + "active": global_or.active, + "created_at": global_or.created_at, + "updated_at": global_or.updated_at, + }, + ) + ) + await tenant_session.execute(insert_or_stmt) + + # --- Sync deletes (children before parents) --- + logger.info(f"[tenant={tenant_id}] Syncing deletes...") + await self._sync_deletes( + tenant_id, + global_session, + tenant_session, + GlobalOutboundRoute, + TenantOutboundRoute, + False, + ) + await self._sync_deletes( + tenant_id, global_session, tenant_session, GlobalInboundRoute, TenantInboundRoute, False + ) + await self._sync_deletes( + tenant_id, global_session, tenant_session, GlobalWebhook, TenantWebhook, False + ) + await self._sync_deletes( + tenant_id, global_session, tenant_session, GlobalSFTPPartner, TenantSFTPPartner, False + ) + await self._sync_deletes( + tenant_id, + global_session, + tenant_session, + GlobalAS2Partnership, + TenantAS2Partnership, + True, + ) + await self._sync_deletes( + tenant_id, global_session, tenant_session, GlobalAS2Partner, TenantAS2Partner, True + ) + + async def _sync_deletes( + self, + tenant_id: int, + global_session: Any, + tenant_session: Any, + global_model: Any, + tenant_model: Any, + include_shared: bool, + ) -> None: + # Get all valid IDs from global + global_stmt = select(global_model.id).where(global_model.tenant_id == tenant_id) + # Special handling for global models (tenant_id = 0 or NULL) + if include_shared: + global_stmt = select(global_model.id).where( + (global_model.tenant_id == tenant_id) + | (global_model.tenant_id == 0) + | (global_model.tenant_id.is_(None)) + ) + global_ids_result = await global_session.execute(global_stmt) + global_ids = set(global_ids_result.scalars().all()) + + # Get all IDs in tenant db + tenant_stmt = select(tenant_model.id).where(tenant_model.tenant_id == tenant_id) + tenant_ids_result = await tenant_session.execute(tenant_stmt) + tenant_ids = set(tenant_ids_result.scalars().all()) + + # Delete ids in tenant that are not in global + ids_to_delete = tenant_ids - global_ids + if ids_to_delete: + logger.info( + f"[tenant={tenant_id}] Deleting {len(ids_to_delete)} stale {tenant_model.__tablename__} record(s)" + ) + delete_stmt = delete(tenant_model).where(tenant_model.id.in_(list(ids_to_delete))) + await tenant_session.execute(delete_stmt) + else: + logger.debug( + f"[tenant={tenant_id}] No stale {tenant_model.__tablename__} records to delete" + ) diff --git a/services/worker/src/worker/utils.py b/services/worker/src/worker/adapters/db_tenant.py similarity index 63% rename from services/worker/src/worker/utils.py rename to services/worker/src/worker/adapters/db_tenant.py index 3bd16e83..bb95e5d1 100644 --- a/services/worker/src/worker/utils.py +++ b/services/worker/src/worker/adapters/db_tenant.py @@ -4,17 +4,26 @@ from database.models import DatabaseShard, Tenant from sqlalchemy import select +from worker.ports.tenant import TenantPort -class TenantResolver: - """ - Caches tenant-to-shard mapping to avoid querying the Global DB on every SQS message. - """ +class SqlAlchemyTenantAdapter(TenantPort): def __init__(self, db_router: DatabaseRouter): self.db_router = db_router self._cache: dict[int, tuple[str, str]] = {} - async def resolve(self, tenant_id: int) -> tuple[str, str]: + async def get_all_tenant_ids(self) -> list[int]: + global_gen = self.db_router.get_global_session() + global_session = await global_gen.__anext__() + try: + stmt = select(Tenant.id).join(DatabaseShard) + result = await global_session.execute(stmt) + return list(result.scalars().all()) + finally: + with contextlib.suppress(StopAsyncIteration): + await global_gen.__anext__() + + async def resolve_shard(self, tenant_id: int) -> tuple[str, str]: if tenant_id in self._cache: return self._cache[tenant_id] diff --git a/services/worker/src/worker/adapters/sqs_outbox.py b/services/worker/src/worker/adapters/sqs_outbox.py new file mode 100644 index 00000000..1b341214 --- /dev/null +++ b/services/worker/src/worker/adapters/sqs_outbox.py @@ -0,0 +1,99 @@ +import json +import logging +import os +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager + +import aioboto3 # type: ignore[import-untyped] +from domain.events import MessageQueueName + +from worker.core.errors import PermanentProvisioningError +from worker.ports.outbox import OutboxEvent, OutboxPort + +logger = logging.getLogger(__name__) + + +class SqsEvent(OutboxEvent): + def __init__(self, message_id: str, receipt_handle: str, body: dict[str, object]): + self._message_id = message_id + self.receipt_handle = receipt_handle + self._body = body + + @property + def id(self) -> str: + return self._message_id + + @property + def event_type(self) -> str: + return str(self._body.get("event_type", "UNKNOWN")) + + @property + def payload(self) -> dict[str, object]: + payload_val = self._body.get("payload", {}) + if isinstance(payload_val, dict): + return payload_val + return {} + + +class SqsOutboxAdapter(OutboxPort): + def __init__(self, queue_name: str = MessageQueueName.PROVISIONING): + self.queue_name = queue_name + self.endpoint_url = os.environ.get("AWS_ENDPOINT_URL", "http://localhost:4566") + self.region = "us-east-1" + self.session = aioboto3.Session() + + @asynccontextmanager + async def process_next_event(self) -> AsyncIterator[OutboxEvent | None]: + async with self.session.client( + "sqs", endpoint_url=self.endpoint_url, region_name=self.region + ) as sqs: + try: + queue_url_response = await sqs.get_queue_url(QueueName=self.queue_name) + queue_url = queue_url_response["QueueUrl"] + except Exception as e: + logger.error(f"Failed to get queue URL for {self.queue_name}: {e}") + yield None + return + + response = await sqs.receive_message( + QueueUrl=queue_url, + MaxNumberOfMessages=1, + WaitTimeSeconds=5, + ) + + messages = response.get("Messages", []) + if not messages: + yield None + return + + msg = messages[0] + receipt_handle = msg["ReceiptHandle"] + message_id = msg["MessageId"] + body_str = msg.get("Body", "{}") + + try: + body = json.loads(body_str) + except json.JSONDecodeError: + logger.error(f"Failed to parse JSON body from SQS message {message_id}") + await sqs.delete_message(QueueUrl=queue_url, ReceiptHandle=receipt_handle) + yield None + return + + event = SqsEvent(message_id=message_id, receipt_handle=receipt_handle, body=body) + + try: + yield event + # Delete the message on success + await sqs.delete_message(QueueUrl=queue_url, ReceiptHandle=receipt_handle) + logger.info(f"Successfully processed and deleted SQS message {message_id}") + except Exception as e: + if isinstance(e, PermanentProvisioningError): + logger.error( + f"Permanent error processing event {event.id}: {e}. Removing from queue." + ) + await sqs.delete_message(QueueUrl=queue_url, ReceiptHandle=receipt_handle) + else: + logger.exception( + f"Transient error processing event {event.id}: {e}. Leaving on queue." + ) + raise diff --git a/services/worker/src/worker/adapters/vault.py b/services/worker/src/worker/adapters/vault.py new file mode 100644 index 00000000..0739b215 --- /dev/null +++ b/services/worker/src/worker/adapters/vault.py @@ -0,0 +1,38 @@ +import asyncio +import logging +import os +import sys + +import hvac + +logger = logging.getLogger(__name__) + + +class WorkerVaultAdapter: + def __init__(self) -> None: + self.url = os.getenv("VAULT_ADDR", "http://localhost:8200") + token = os.getenv("VAULT_TOKEN") + env = os.getenv("ENVIRONMENT", "development") + if not token: + if env in ("development", "dev", "test", "local") or "pytest" in sys.modules: + token = "root" + else: + raise ValueError("VAULT_TOKEN required in non-dev") + + self.token = token + self.client = hvac.Client(url=self.url, token=self.token) + self.mount_point = "secret" + + async def get_secret(self, vault_ref: str) -> str: + # HVAC is synchronous, but we wrap in asyncio.to_thread for the port + def _fetch() -> str: + resp = self.client.secrets.kv.v2.read_secret_version( + path=vault_ref, mount_point=self.mount_point + ) + data = resp.get("data", {}).get("data", {}) + val = next(iter(data.values()), None) + if val is None: + raise ValueError(f"Secret not found at vault path: {vault_ref}") + return str(val) + + return await asyncio.to_thread(_fetch) diff --git a/services/worker/src/worker/core/__init__.py b/services/worker/src/worker/core/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/services/worker/src/worker/core/errors.py b/services/worker/src/worker/core/errors.py new file mode 100644 index 00000000..2ccfc461 --- /dev/null +++ b/services/worker/src/worker/core/errors.py @@ -0,0 +1,16 @@ +class ProvisioningError(Exception): + """Base exception for provisioning errors.""" + + pass + + +class PermanentProvisioningError(ProvisioningError): + """An error that cannot be resolved by retrying (e.g., bad payload).""" + + pass + + +class TransientProvisioningError(ProvisioningError): + """An error that might be resolved by retrying (e.g., network failure).""" + + pass diff --git a/services/worker/src/worker/core/service.py b/services/worker/src/worker/core/service.py new file mode 100644 index 00000000..01dac888 --- /dev/null +++ b/services/worker/src/worker/core/service.py @@ -0,0 +1,70 @@ +import logging + +from worker.core.errors import PermanentProvisioningError, TransientProvisioningError +from worker.ports.outbox import OutboxPort +from worker.ports.replication import ReplicationPort +from worker.ports.tenant import TenantPort + +logger = logging.getLogger(__name__) + + +class ProvisioningWorkerService: + def __init__( + self, tenant_port: TenantPort, outbox_port: OutboxPort, replication_port: ReplicationPort + ): + self.tenant_port = tenant_port + self.outbox_port = outbox_port + self.replication_port = replication_port + + async def process_next_event(self) -> bool: + """Process a single event from the outbox. Returns True if an event was processed.""" + async with self.outbox_port.process_next_event() as event: + if not event: + return False + + payload = event.payload + tenant_id = payload.get("tenant_id") + + if tenant_id is None: + raise PermanentProvisioningError("Missing tenant_id in provision event payload") + + if tenant_id == 0: + logger.info( + f"Processing GLOBAL provision event {event.id} (tenant_id=0). Broadcasting to all tenants." + ) + try: + all_tenant_ids = await self.tenant_port.get_all_tenant_ids() + + import asyncio + + async def _replicate(t_id: int) -> None: + await self.replication_port.replicate_tenant_configuration(t_id) + + results = await asyncio.gather( + *[_replicate(t_id) for t_id in all_tenant_ids], return_exceptions=True + ) + + errors = [] + for t_id, result in zip(all_tenant_ids, results, strict=False): + if isinstance(result, Exception): + logger.exception( + f"Failed to broadcast global event {event.id} to tenant {t_id}: {result}", + exc_info=result, + ) + if not isinstance(result, PermanentProvisioningError): + errors.append(result) + + if errors: + raise TransientProvisioningError( + f"Global broadcasting failed for some tenants: {errors}" + ) + + except Exception as e: + if isinstance(e, TransientProvisioningError): + raise + raise TransientProvisioningError(f"Global broadcasting failed: {e}") from e + else: + logger.info(f"Processing provision event {event.id} for tenant_id={tenant_id}") + await self.replication_port.replicate_tenant_configuration(tenant_id) + + return True diff --git a/services/worker/src/worker/data/main.py b/services/worker/src/worker/data/main.py index 458c28be..1de1f193 100644 --- a/services/worker/src/worker/data/main.py +++ b/services/worker/src/worker/data/main.py @@ -11,15 +11,56 @@ import aioboto3 # type: ignore[import-untyped] from config.settings import get_settings from database.connection import DatabaseRouter +from database.models import DatabaseShard, Tenant +from domain.events import MessageQueueName +from dotenv import load_dotenv +from pipeline.adapters.as2 import HttpxAS2DeliveryAdapter from pipeline.adapters.http import HttpxDeliveryAdapter -from pipeline.adapters.null_as2 import NullAS2DeliveryAdapter from pipeline.adapters.repository import SqlAlchemyRepositoryAdapter from pipeline.adapters.sftp import ParamikoSftpDeliveryAdapter from pipeline.adapters.storage import S3StorageAdapter from pipeline.adapters.transformer import BotsTransformerAdapter from pipeline.core.deliver import DeliveryService from pipeline.core.translate import TranslationService -from worker.utils import TenantResolver +from sqlalchemy import select +from worker.adapters.vault import WorkerVaultAdapter + +load_dotenv() + + +class TenantResolver: + """ + Caches tenant-to-shard mapping to avoid querying the Global DB on every SQS message. + """ + + def __init__(self, db_router: DatabaseRouter, ttl_secs: int = 300): + self.db_router = db_router + self._cache: dict[int, tuple[str, str, float]] = {} + self._ttl = ttl_secs + + async def resolve(self, tenant_id: int) -> tuple[str, str]: + import time + + now = time.time() + if tenant_id in self._cache: + shard_name, shard_dsn, expiry = self._cache[tenant_id] + if now < expiry: + return shard_name, shard_dsn + + global_gen = self.db_router.get_global_session() + global_session = await global_gen.__anext__() + try: + stmt = select(Tenant, DatabaseShard).join(DatabaseShard).where(Tenant.id == tenant_id) + result = await global_session.execute(stmt) + row = result.first() + if not row: + raise ValueError(f"Tenant {tenant_id} not found in Global DB") + _, shard_obj = row + self._cache[tenant_id] = (str(shard_obj.name), str(shard_obj.dsn), now + self._ttl) + return str(shard_obj.name), str(shard_obj.dsn) + finally: + await global_gen.aclose() + logger = logging.getLogger(__name__) @@ -73,6 +114,7 @@ def validate_target_url(url: str) -> bool: async def process_translation( trace_id: str, + event_type: str, tenant_id: int, resolver: TenantResolver, db_router: DatabaseRouter, @@ -85,20 +127,26 @@ async def process_translation( tenant_gen = db_router.get_tenant_session(tenant_id, shard_name, shard_dsn) session = await tenant_gen.__anext__() try: - # Instantiate Adapters - repo_adapter = SqlAlchemyRepositoryAdapter(session) storage_adapter = S3StorageAdapter(bucket_name=s3_bucket, endpoint_url=aws_endpoint) + repo_adapter = SqlAlchemyRepositoryAdapter( + session=session, + settings=get_settings(), + storage=storage_adapter, + ) transformer_adapter = BotsTransformerAdapter() # Instantiate Domain Service - service = TranslationService(storage_adapter, transformer_adapter, repo_adapter) + service = TranslationService(transformer_adapter, repo_adapter) # Execute pure domain logic - await service.translate(trace_id) + print(f"[WORKER] Translating trace_id={trace_id}") + await service.translate(trace_id, event_type) + print(f"[WORKER] SUCCESS translating trace_id={trace_id}") # Commit transaction await session.commit() - except Exception: + except Exception as e: + print(f"[WORKER] FAILURE in process_translation for trace_id={trace_id}: {e}") await session.rollback() raise finally: @@ -108,6 +156,7 @@ async def process_translation( async def process_delivery( trace_id: str, + event_type: str, tenant_id: int, resolver: TenantResolver, db_router: DatabaseRouter, @@ -120,19 +169,24 @@ async def process_delivery( tenant_gen = db_router.get_tenant_session(tenant_id, shard_name, shard_dsn) session = await tenant_gen.__anext__() try: - # Instantiate Adapters - repo_adapter = SqlAlchemyRepositoryAdapter(session) storage_adapter = S3StorageAdapter(bucket_name=s3_bucket, endpoint_url=aws_endpoint) + repo_adapter = SqlAlchemyRepositoryAdapter( + session=session, + settings=get_settings(), + storage=storage_adapter, + ) http_adapter = HttpxDeliveryAdapter(validator=validate_target_url) sftp_adapter = ParamikoSftpDeliveryAdapter() + vault_adapter = WorkerVaultAdapter() + as2_adapter = HttpxAS2DeliveryAdapter() # Instantiate Domain Service service = DeliveryService( - storage_adapter, - repo_adapter, - http_adapter, - sftp_adapter, - as2_delivery=NullAS2DeliveryAdapter(), + repository=repo_adapter, + http_delivery=http_adapter, + sftp_delivery=sftp_adapter, + as2_delivery=as2_adapter, + vault=vault_adapter, ) # Execute pure domain logic @@ -200,6 +254,7 @@ async def poll_sqs_queue( logger.info(f"[{queue_name}] Processing trace_id={trace_id}") kwargs: dict[str, Any] = { "trace_id": trace_id, + "event_type": body.get("event_type", "UNKNOWN"), "tenant_id": tenant_id, "resolver": resolver, "db_router": db_router, @@ -251,12 +306,17 @@ async def main() -> None: translate_task = asyncio.create_task( poll_sqs_queue( - "TranslateQueue", process_translation, resolver, db_router, s3_bucket, aws_endpoint + MessageQueueName.TRANSLATE, + process_translation, + resolver, + db_router, + s3_bucket, + aws_endpoint, ) ) deliver_task = asyncio.create_task( poll_sqs_queue( - "DeliverQueue", process_delivery, resolver, db_router, s3_bucket, aws_endpoint + MessageQueueName.DELIVER, process_delivery, resolver, db_router, s3_bucket, aws_endpoint ) ) diff --git a/services/worker/src/worker/main.py b/services/worker/src/worker/main.py new file mode 100644 index 00000000..c04608df --- /dev/null +++ b/services/worker/src/worker/main.py @@ -0,0 +1,22 @@ +import asyncio +import logging + +from worker.data.main import main as data_main +from worker.provision.main import main as provision_main + +logger = logging.getLogger(__name__) + + +async def main() -> None: + logger.info("Starting unified Enterprise EDI Worker (Data + Provisioning)...") + + # Run both the Provisioning and Data worker tasks concurrently + data_task = asyncio.create_task(data_main()) + provision_task = asyncio.create_task(provision_main()) + + await asyncio.gather(data_task, provision_task) + + +if __name__ == "__main__": + logging.basicConfig(level=logging.INFO) + asyncio.run(main()) diff --git a/services/worker/src/worker/ports/__init__.py b/services/worker/src/worker/ports/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/services/worker/src/worker/ports/outbox.py b/services/worker/src/worker/ports/outbox.py new file mode 100644 index 00000000..18cfc5c9 --- /dev/null +++ b/services/worker/src/worker/ports/outbox.py @@ -0,0 +1,17 @@ +from contextlib import AbstractAsyncContextManager +from typing import Any, Protocol + + +class OutboxEvent(Protocol): + @property + def id(self) -> Any: ... + @property + def event_type(self) -> str: ... + @property + def payload(self) -> dict[str, Any]: ... + + +class OutboxPort(Protocol): + def process_next_event(self) -> AbstractAsyncContextManager[OutboxEvent | None]: + """Context manager that yields the next pending event, or None.""" + ... diff --git a/services/worker/src/worker/ports/replication.py b/services/worker/src/worker/ports/replication.py new file mode 100644 index 00000000..11e3bbe1 --- /dev/null +++ b/services/worker/src/worker/ports/replication.py @@ -0,0 +1,7 @@ +from typing import Protocol + + +class ReplicationPort(Protocol): + async def replicate_tenant_configuration(self, tenant_id: int) -> None: + """Copy all relevant configuration from Global DB to Tenant DB Shard.""" + ... diff --git a/services/worker/src/worker/ports/tenant.py b/services/worker/src/worker/ports/tenant.py new file mode 100644 index 00000000..99123ae3 --- /dev/null +++ b/services/worker/src/worker/ports/tenant.py @@ -0,0 +1,11 @@ +from typing import Protocol + + +class TenantPort(Protocol): + async def get_all_tenant_ids(self) -> list[int]: + """Fetch all active tenant IDs from the Global DB.""" + ... + + async def resolve_shard(self, tenant_id: int) -> tuple[str, str]: + """Resolve a tenant_id to a database shard (name, dsn).""" + ... diff --git a/services/worker/src/worker/provision/main.py b/services/worker/src/worker/provision/main.py index 03fb4172..d279c8e0 100644 --- a/services/worker/src/worker/provision/main.py +++ b/services/worker/src/worker/provision/main.py @@ -1,405 +1,44 @@ import asyncio -import contextlib import logging -from typing import Any from config.settings import get_settings from database.connection import DatabaseRouter -from database.models.control_plane import AS2Partner as GlobalAS2Partner -from database.models.control_plane import AS2Partnership as GlobalAS2Partnership -from database.models.control_plane import InboundRoute as GlobalInboundRoute -from database.models.control_plane import OutboundRoute as GlobalOutboundRoute -from database.models.control_plane import Outbox as GlobalOutbox -from database.models.control_plane import SFTPPartner as GlobalSFTPPartner -from database.models.control_plane import Webhook as GlobalWebhook -from database.models.data_plane import AS2Partner as TenantAS2Partner -from database.models.data_plane import AS2Partnership as TenantAS2Partnership -from database.models.data_plane import InboundRoute as TenantInboundRoute -from database.models.data_plane import OutboundRoute as TenantOutboundRoute -from database.models.data_plane import SFTPPartner as TenantSFTPPartner -from database.models.data_plane import Webhook as TenantWebhook -from sqlalchemy import select -from sqlalchemy.dialects.postgresql import insert -from sqlalchemy.ext.asyncio import AsyncSession -from worker.utils import TenantResolver +from dotenv import load_dotenv +from worker.adapters.db_outbox import SqlAlchemyOutboxAdapter +from worker.adapters.db_replication import SqlAlchemyReplicationAdapter +from worker.adapters.db_tenant import SqlAlchemyTenantAdapter +from worker.core.service import ProvisioningWorkerService -logger = logging.getLogger(__name__) - - -async def sync_deletes( - tenant_id: int, global_session: Any, tenant_session: Any, global_model: Any, tenant_model: Any -) -> None: - from sqlalchemy import delete, select - - # Get all valid IDs from global - global_stmt = select(global_model.id).where(global_model.tenant_id == tenant_id) - # Special handling for AS2Partner which has global partners (tenant_id IS NULL) - if global_model.__name__ == "AS2Partner": - global_stmt = select(global_model.id).where( - (global_model.tenant_id == tenant_id) | (global_model.tenant_id.is_(None)) - ) - global_ids_result = await global_session.execute(global_stmt) - global_ids = set(global_ids_result.scalars().all()) - - # Get all IDs in tenant db - tenant_stmt = select(tenant_model.id).where(tenant_model.tenant_id == tenant_id) - tenant_ids_result = await tenant_session.execute(tenant_stmt) - tenant_ids = set(tenant_ids_result.scalars().all()) - - # Delete ids in tenant that are not in global - ids_to_delete = tenant_ids - global_ids - if ids_to_delete: - delete_stmt = delete(tenant_model).where(tenant_model.id.in_(list(ids_to_delete))) - await tenant_session.execute(delete_stmt) - - -async def replicate_tenant_config( - tenant_id: int, global_session: AsyncSession, tenant_session: AsyncSession -) -> None: - """Replicates configuration from Global DB to Tenant DB Shard.""" - - # Replicate AS2Partners - # We replicate all AS2Partners that belong to this tenant, plus any global ones (tenant_id IS NULL) - stmt = select(GlobalAS2Partner).where( - (GlobalAS2Partner.tenant_id == tenant_id) | (GlobalAS2Partner.tenant_id == 0) - ) - tp_result = await global_session.execute(stmt) - - for global_tp in tp_result.scalars(): - insert_stmt = ( - insert(TenantAS2Partner) - .values( - id=global_tp.id, - tenant_id=tenant_id, # Use destination tenant_id for replicated partners - name=global_tp.name, - as2_id=global_tp.as2_id, - is_local=global_tp.is_local, - public_cert_pem=global_tp.public_cert_pem, - public_cert_vault_ref=global_tp.public_cert_vault_ref, - private_key_vault_ref=global_tp.private_key_vault_ref, - prev_public_cert_pem=global_tp.prev_public_cert_pem, - prev_public_cert_vault_ref=global_tp.prev_public_cert_vault_ref, - prev_private_key_vault_ref=global_tp.prev_private_key_vault_ref, - url=global_tp.url, - active=global_tp.active, - ) - .on_conflict_do_update( - index_elements=["id"], - set_={ - "tenant_id": tenant_id, - "name": global_tp.name, - "as2_id": global_tp.as2_id, - "is_local": global_tp.is_local, - "public_cert_pem": global_tp.public_cert_pem, - "public_cert_vault_ref": global_tp.public_cert_vault_ref, - "private_key_vault_ref": global_tp.private_key_vault_ref, - "prev_public_cert_pem": global_tp.prev_public_cert_pem, - "prev_public_cert_vault_ref": global_tp.prev_public_cert_vault_ref, - "prev_private_key_vault_ref": global_tp.prev_private_key_vault_ref, - "url": global_tp.url, - "active": global_tp.active, - }, - ) - ) - await tenant_session.execute(insert_stmt) - - # Replicate AS2Partnerships - ps_stmt = select(GlobalAS2Partnership).where(GlobalAS2Partnership.tenant_id == tenant_id) - ps_result = await global_session.execute(ps_stmt) - - for global_ps in ps_result.scalars(): - insert_ps_stmt = ( - insert(TenantAS2Partnership) - .values( - id=global_ps.id, - tenant_id=tenant_id, - name=global_ps.name, - local_partner_id=global_ps.local_partner_id, - remote_partner_id=global_ps.remote_partner_id, - credentials_vault_ref=global_ps.credentials_vault_ref, - mdn_type=global_ps.mdn_type, - mdn_url=global_ps.mdn_url, - encryption_algorithm=global_ps.encryption_algorithm, - signature_algorithm=global_ps.signature_algorithm, - advanced_flags=global_ps.advanced_flags, - active=global_ps.active, - ) - .on_conflict_do_update( - index_elements=["id"], - set_={ - "name": global_ps.name, - "local_partner_id": global_ps.local_partner_id, - "remote_partner_id": global_ps.remote_partner_id, - "credentials_vault_ref": global_ps.credentials_vault_ref, - "mdn_type": global_ps.mdn_type, - "mdn_url": global_ps.mdn_url, - "encryption_algorithm": global_ps.encryption_algorithm, - "signature_algorithm": global_ps.signature_algorithm, - "advanced_flags": global_ps.advanced_flags, - "active": global_ps.active, - }, - ) - ) - await tenant_session.execute(insert_ps_stmt) - - # Replicate SFTPPartners - sftp_stmt = select(GlobalSFTPPartner).where(GlobalSFTPPartner.tenant_id == tenant_id) - sftp_result = await global_session.execute(sftp_stmt) - for global_sftp in sftp_result.scalars(): - insert_sftp_stmt = ( - insert(TenantSFTPPartner) - .values( - id=global_sftp.id, - tenant_id=tenant_id, - name=global_sftp.name, - host=global_sftp.host, - port=global_sftp.port, - username=global_sftp.username, - inbound_remote_path=global_sftp.inbound_remote_path, - outbound_remote_path=global_sftp.outbound_remote_path, - host_key=global_sftp.host_key, - password_encrypted=global_sftp.password_encrypted, - credentials_vault_ref=global_sftp.credentials_vault_ref, - active=global_sftp.active, - ) - .on_conflict_do_update( - index_elements=["id"], - set_={ - "name": global_sftp.name, - "host": global_sftp.host, - "port": global_sftp.port, - "username": global_sftp.username, - "inbound_remote_path": global_sftp.inbound_remote_path, - "outbound_remote_path": global_sftp.outbound_remote_path, - "host_key": global_sftp.host_key, - "password_encrypted": global_sftp.password_encrypted, - "credentials_vault_ref": global_sftp.credentials_vault_ref, - "active": global_sftp.active, - }, - ) - ) - await tenant_session.execute(insert_sftp_stmt) - - # Replicate Webhooks - wh_stmt = select(GlobalWebhook).where(GlobalWebhook.tenant_id == tenant_id) - wh_result = await global_session.execute(wh_stmt) - for global_wh in wh_result.scalars(): - insert_wh_stmt = ( - insert(TenantWebhook) - .values( - id=global_wh.id, - tenant_id=tenant_id, - name=global_wh.name, - url=global_wh.url, - auth_header_vault_ref=global_wh.auth_header_vault_ref, - active=global_wh.active, - ) - .on_conflict_do_update( - index_elements=["id"], - set_={ - "name": global_wh.name, - "url": global_wh.url, - "auth_header_vault_ref": global_wh.auth_header_vault_ref, - "active": global_wh.active, - }, - ) - ) - await tenant_session.execute(insert_wh_stmt) - - # Replicate InboundRoutes - ir_stmt = select(GlobalInboundRoute).where(GlobalInboundRoute.tenant_id == tenant_id) - ir_result = await global_session.execute(ir_stmt) - for global_ir in ir_result.scalars(): - insert_ir_stmt = ( - insert(TenantInboundRoute) - .values( - id=global_ir.id, - tenant_id=tenant_id, - name=global_ir.name, - isa_sender_id=global_ir.isa_sender_id, - isa_receiver_id=global_ir.isa_receiver_id, - transaction_type=global_ir.transaction_type, - processing_mode=global_ir.processing_mode, - webhook_id=global_ir.webhook_id, - as2_partner_id=global_ir.as2_partner_id, - sftp_partner_id=global_ir.sftp_partner_id, - active=global_ir.active, - ) - .on_conflict_do_update( - index_elements=["id"], - set_={ - "name": global_ir.name, - "isa_sender_id": global_ir.isa_sender_id, - "isa_receiver_id": global_ir.isa_receiver_id, - "transaction_type": global_ir.transaction_type, - "processing_mode": global_ir.processing_mode, - "webhook_id": global_ir.webhook_id, - "as2_partner_id": global_ir.as2_partner_id, - "sftp_partner_id": global_ir.sftp_partner_id, - "active": global_ir.active, - }, - ) - ) - await tenant_session.execute(insert_ir_stmt) - - # Replicate OutboundRoutes - or_stmt = select(GlobalOutboundRoute).where(GlobalOutboundRoute.tenant_id == tenant_id) - or_result = await global_session.execute(or_stmt) - for global_or in or_result.scalars(): - insert_or_stmt = ( - insert(TenantOutboundRoute) - .values( - id=global_or.id, - tenant_id=tenant_id, - name=global_or.name, - isa_sender_id=global_or.isa_sender_id, - isa_receiver_id=global_or.isa_receiver_id, - transaction_type=global_or.transaction_type, - processing_mode=global_or.processing_mode, - as2_partner_id=global_or.as2_partner_id, - sftp_partner_id=global_or.sftp_partner_id, - active=global_or.active, - ) - .on_conflict_do_update( - index_elements=["id"], - set_={ - "name": global_or.name, - "isa_sender_id": global_or.isa_sender_id, - "isa_receiver_id": global_or.isa_receiver_id, - "transaction_type": global_or.transaction_type, - "processing_mode": global_or.processing_mode, - "as2_partner_id": global_or.as2_partner_id, - "sftp_partner_id": global_or.sftp_partner_id, - "active": global_or.active, - }, - ) - ) - await tenant_session.execute(insert_or_stmt) - - # Sync deletes for all configurations (Dependent children first) - await sync_deletes( - tenant_id, global_session, tenant_session, GlobalOutboundRoute, TenantOutboundRoute - ) - await sync_deletes( - tenant_id, global_session, tenant_session, GlobalInboundRoute, TenantInboundRoute - ) - await sync_deletes(tenant_id, global_session, tenant_session, GlobalWebhook, TenantWebhook) - await sync_deletes( - tenant_id, global_session, tenant_session, GlobalSFTPPartner, TenantSFTPPartner - ) - await sync_deletes( - tenant_id, global_session, tenant_session, GlobalAS2Partnership, TenantAS2Partnership - ) - await sync_deletes( - tenant_id, global_session, tenant_session, GlobalAS2Partner, TenantAS2Partner - ) - - await tenant_session.commit() - - logger.info(f"Successfully replicated AS2 configuration for tenant_id={tenant_id}") +load_dotenv() +logger = logging.getLogger(__name__) -async def poll_global_outbox( - db_router: DatabaseRouter, - resolver: TenantResolver, -) -> None: - """Polls the Global Outbox for provisioning events.""" - logger.info("Started polling Global Outbox for PROVISION events") +async def run_worker(service: ProvisioningWorkerService) -> None: + logger.info("Started polling Database for PROVISION events") while True: - global_gen = db_router.get_global_session() - global_session = await global_gen.__anext__() try: - # Find a pending provision event - stmt = ( - select(GlobalOutbox) - .where( - GlobalOutbox.status == "PENDING", - GlobalOutbox.event_type.in_( - [ - "AS2_PARTNER_CREATED", - "AS2_PARTNERSHIP_CREATED", - "AS2_PARTNER_UPDATED", - "AS2_PARTNERSHIP_UPDATED", - "AS2_PARTNER_DELETED", - "AS2_PARTNERSHIP_DELETED", - "SFTP_PARTNER_CREATED", - "SFTP_PARTNER_UPDATED", - "SFTP_PARTNER_DELETED", - "WEBHOOK_CREATED", - "WEBHOOK_UPDATED", - "WEBHOOK_DELETED", - "INBOUND_ROUTE_CREATED", - "INBOUND_ROUTE_UPDATED", - "INBOUND_ROUTE_DELETED", - "OUTBOUND_ROUTE_CREATED", - "OUTBOUND_ROUTE_UPDATED", - "OUTBOUND_ROUTE_DELETED", - ] - ), - ) - .limit(1) - .with_for_update(skip_locked=True) - ) - - result = await global_session.execute(stmt) - outbox_event = result.scalar_one_or_none() - - if outbox_event: - payload = outbox_event.payload - tenant_id = payload.get("tenant_id") - - if tenant_id is None: - logger.error(f"Missing tenant_id in provision event: {outbox_event.id}") - outbox_event.status = "FAILED" - await global_session.commit() - continue - - logger.info(f"Processing provision event for tenant_id={tenant_id}") - - shard_name, shard_dsn = await resolver.resolve(tenant_id) - tenant_gen = db_router.get_tenant_session(tenant_id, shard_name, shard_dsn) - tenant_session = await tenant_gen.__anext__() - try: - await replicate_tenant_config(tenant_id, global_session, tenant_session) - - # Mark outbox event as processed - outbox_event.status = "PROCESSED" - await global_session.commit() - except (ValueError, KeyError) as e: - # Permanent data errors: bad payload or missing key — mark FAILED - await tenant_session.rollback() - logger.error(f"Permanent provisioning failure for tenant {tenant_id}: {e}") - outbox_event.status = "FAILED" - await global_session.commit() - except Exception as e: - # Transient errors (network, DB): leave PENDING for retry - await tenant_session.rollback() - logger.exception( - f"Transient error provisioning tenant {tenant_id}, will retry: {e}" - ) - # Do NOT change status — let the poller pick it up again - finally: - with contextlib.suppress(StopAsyncIteration): - await tenant_gen.__anext__() - + processed_event = await service.process_next_event() + # The SQS receive_message already blocks for up to 5 seconds (WaitTimeSeconds=5) + # We don't need to sleep here if we didn't process an event, but we can do a tiny yield + if not processed_event: + await asyncio.sleep(0.1) except Exception as e: - logger.exception(f"Error polling global outbox: {e}") + logger.exception(f"Error in provisioning loop: {e}") await asyncio.sleep(5) - finally: - with contextlib.suppress(StopAsyncIteration): - await global_gen.__anext__() - - # Sleep before polling again - await asyncio.sleep(5) async def main() -> None: settings = get_settings() db_router = DatabaseRouter(global_db_url=settings.database.global_url) - resolver = TenantResolver(db_router) - await poll_global_outbox(db_router, resolver) + tenant_adapter = SqlAlchemyTenantAdapter(db_router) + outbox_adapter = SqlAlchemyOutboxAdapter(db_router) + replication_adapter = SqlAlchemyReplicationAdapter(db_router, tenant_adapter) + + service = ProvisioningWorkerService(tenant_adapter, outbox_adapter, replication_adapter) + + await run_worker(service) if __name__ == "__main__": diff --git a/services/worker/tests/test_data_worker.py b/services/worker/tests/test_data_worker.py index 0720ef34..34883070 100644 --- a/services/worker/tests/test_data_worker.py +++ b/services/worker/tests/test_data_worker.py @@ -66,16 +66,17 @@ async def test_poll_sqs_queue_processes_message(mock_session_cls: MagicMock) -> mock_resolver, mock_db_router, "bucket", - "http://localhost", + "endpoint", ) mock_processor.assert_awaited_once_with( trace_id="trace-456", + event_type="UNKNOWN", tenant_id=99, resolver=mock_resolver, db_router=mock_db_router, s3_bucket="bucket", - aws_endpoint="http://localhost", + aws_endpoint="endpoint", ) mock_client.delete_message.assert_awaited_once_with( @@ -98,10 +99,16 @@ async def test_process_translation(mock_service_cls: MagicMock) -> None: mock_tenant_gen.__anext__.return_value = mock_tenant_session await process_translation( - "trace-123", 99, mock_resolver, mock_db_router, "bucket", "http://localhost" + "trace-123", + "edi_message.received", + 99, + mock_resolver, + mock_db_router, + "bucket", + "http://localhost", ) - mock_service.translate.assert_awaited_once_with("trace-123") + mock_service.translate.assert_awaited_once_with("trace-123", "edi_message.received") @patch("worker.data.main.DeliveryService") @@ -120,6 +127,7 @@ async def test_process_delivery(mock_service_cls: MagicMock) -> None: await process_delivery( "trace-123", + "DELIVER", 99, mock_resolver, mock_db_router, diff --git a/services/worker/tests/test_provision_worker.py b/services/worker/tests/test_provision_worker.py index 006e10cf..237a33e4 100644 --- a/services/worker/tests/test_provision_worker.py +++ b/services/worker/tests/test_provision_worker.py @@ -1,128 +1,197 @@ -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock import pytest -from worker.provision.main import poll_global_outbox, replicate_tenant_config +from worker.core.errors import PermanentProvisioningError, TransientProvisioningError +from worker.core.service import ProvisioningWorkerService pytestmark = pytest.mark.asyncio -@patch("worker.provision.main.replicate_tenant_config") -async def test_poll_global_outbox_processes_event(mock_replicate: AsyncMock) -> None: - mock_db_router = MagicMock() - mock_resolver = MagicMock() - - mock_resolver.resolve = AsyncMock(return_value=("shard1", "url")) - - mock_global_gen = AsyncMock() - mock_global_session = AsyncMock() - mock_db_router.get_global_session.return_value = mock_global_gen - mock_global_gen.__anext__.return_value = mock_global_session - - mock_tenant_gen = AsyncMock() - mock_tenant_session = AsyncMock() - mock_db_router.get_tenant_session.return_value = mock_tenant_gen - mock_tenant_gen.__anext__.return_value = mock_tenant_session - - mock_outbox_event = MagicMock() - mock_outbox_event.payload = {"tenant_id": 99} - - mock_result = MagicMock() - # Return event first time, then raise exception to break loop - mock_result.scalar_one_or_none.return_value = mock_outbox_event - - import asyncio - - mock_global_session.execute.side_effect = [mock_result, asyncio.CancelledError()] - - with pytest.raises(asyncio.CancelledError): - await poll_global_outbox(mock_db_router, mock_resolver) - - mock_replicate.assert_awaited_once_with(99, mock_global_session, mock_tenant_session) - assert mock_outbox_event.status == "PROCESSED" - mock_global_session.commit.assert_awaited_once() - - -def _make_scalars_result(items: list) -> MagicMock: - """Creates a mock SQLAlchemy result that supports both iteration and .scalars().all().""" - mock_result = MagicMock() - # Support for `for item in result.scalars():` (iteration) - mock_result.scalars.return_value = iter(items) - # Support for `result.scalars().all()` (in sync_deletes) - mock_scalars = MagicMock() - mock_scalars.__iter__ = MagicMock(return_value=iter(items)) - mock_scalars.all.return_value = items - mock_result.scalars.return_value = mock_scalars - return mock_result - - -def _make_empty_scalars_result() -> MagicMock: - """Creates a mock that returns empty for both iteration and .scalars().all().""" - return _make_scalars_result([]) - - -async def test_replicate_tenant_config() -> None: - mock_global_session = AsyncMock() - mock_tenant_session = AsyncMock() - - # Mock global execution to return some scalars for AS2Partner - mock_tp = MagicMock() - mock_tp.id = "tp-uuid" - mock_tp.tenant_id = 99 - mock_tp.name = "Acme Corp AS2" - mock_tp.as2_id = "ACME" - mock_tp.is_local = False - mock_tp.public_cert_pem = "PEM" - mock_tp.public_cert_vault_ref = "vault://acme-pub" - mock_tp.private_key_vault_ref = None - mock_tp.prev_public_cert_pem = None - mock_tp.prev_public_cert_vault_ref = None - mock_tp.prev_private_key_vault_ref = None - mock_tp.url = None - mock_tp.active = True - - mock_ps = MagicMock() - mock_ps.id = "ps-uuid" - mock_ps.tenant_id = 99 - mock_ps.name = "Partnership 1" - mock_ps.local_partner_id = "loc" - mock_ps.remote_partner_id = "rem" - mock_ps.credentials_vault_ref = "ref" - mock_ps.mdn_type = "SYNC" - mock_ps.mdn_url = None - mock_ps.encryption_algorithm = "AES" - mock_ps.signature_algorithm = "SHA" - mock_ps.advanced_flags = None - - mock_ps.active = True - - # replicate_tenant_config makes 6 global_session.execute calls for replication: - # 1. AS2Partners, 2. AS2Partnerships, 3. SFTPPartners, - # 4. Webhooks, 5. InboundRoutes, 6. OutboundRoutes - # Then sync_deletes is called 6 times, each making 1 global_session.execute call - # (the second execute per sync_deletes goes to tenant_session). - # Total global_session.execute calls = 6 (replicate) + 6 (sync_deletes global IDs) = 12 - - mock_global_session.execute.side_effect = [ - # Replication phase (6 calls) - _make_scalars_result([mock_tp]), # AS2Partners - _make_scalars_result([mock_ps]), # AS2Partnerships - _make_empty_scalars_result(), # SFTPPartners - _make_empty_scalars_result(), # Webhooks - _make_empty_scalars_result(), # InboundRoutes - _make_empty_scalars_result(), # OutboundRoutes - # sync_deletes phase - global IDs queries (6 calls) - _make_empty_scalars_result(), # AS2Partners global IDs - _make_empty_scalars_result(), # AS2Partnerships global IDs - _make_empty_scalars_result(), # SFTPPartners global IDs - _make_empty_scalars_result(), # Webhooks global IDs - _make_empty_scalars_result(), # InboundRoutes global IDs - _make_empty_scalars_result(), # OutboundRoutes global IDs - ] - - # sync_deletes also queries tenant_session for IDs (6 calls) - mock_tenant_session.execute.return_value = _make_empty_scalars_result() - - await replicate_tenant_config(99, mock_global_session, mock_tenant_session) - - assert mock_global_session.execute.await_count == 12 - mock_tenant_session.commit.assert_awaited_once() +async def test_process_next_event_no_event() -> None: + mock_tenant = AsyncMock() + mock_outbox = MagicMock() + mock_replication = AsyncMock() + + # Setup the async context manager to yield None + from contextlib import asynccontextmanager + + @asynccontextmanager + async def fake_cm(): + yield None + + mock_outbox.process_next_event.return_value = fake_cm() + + svc = ProvisioningWorkerService(mock_tenant, mock_outbox, mock_replication) + result = await svc.process_next_event() + assert result is False + mock_replication.replicate_tenant_configuration.assert_not_called() + + +async def test_process_next_event_tenant_specific() -> None: + mock_tenant = AsyncMock() + mock_outbox = MagicMock() + mock_replication = AsyncMock() + + mock_event = MagicMock() + mock_event.payload = {"tenant_id": 99} + + from contextlib import asynccontextmanager + + @asynccontextmanager + async def fake_cm(): + yield mock_event + + mock_outbox.process_next_event.return_value = fake_cm() + + svc = ProvisioningWorkerService(mock_tenant, mock_outbox, mock_replication) + result = await svc.process_next_event() + assert result is True + mock_replication.replicate_tenant_configuration.assert_awaited_once_with(99) + + +async def test_process_next_event_global() -> None: + mock_tenant = AsyncMock() + mock_outbox = MagicMock() + mock_replication = AsyncMock() + + mock_event = MagicMock() + mock_event.payload = {"tenant_id": 0} + + from contextlib import asynccontextmanager + + @asynccontextmanager + async def fake_cm(): + yield mock_event + + mock_outbox.process_next_event.return_value = fake_cm() + + mock_tenant.get_all_tenant_ids.return_value = [1, 2, 3] + + svc = ProvisioningWorkerService(mock_tenant, mock_outbox, mock_replication) + result = await svc.process_next_event() + + assert result is True + assert mock_replication.replicate_tenant_configuration.call_count == 3 + + +async def test_process_next_event_missing_tenant_id() -> None: + mock_tenant = AsyncMock() + mock_outbox = MagicMock() + mock_replication = AsyncMock() + + mock_event = MagicMock() + mock_event.payload = {} + + from contextlib import asynccontextmanager + + @asynccontextmanager + async def fake_cm(): + yield mock_event + + mock_outbox.process_next_event.return_value = fake_cm() + svc = ProvisioningWorkerService(mock_tenant, mock_outbox, mock_replication) + + with pytest.raises(PermanentProvisioningError): + await svc.process_next_event() + + +async def test_process_next_event_permanent_error() -> None: + mock_tenant = AsyncMock() + mock_outbox = MagicMock() + mock_replication = AsyncMock() + + mock_event = MagicMock() + mock_event.payload = {"tenant_id": 99} + + from contextlib import asynccontextmanager + + @asynccontextmanager + async def fake_cm(): + yield mock_event + + mock_outbox.process_next_event.return_value = fake_cm() + mock_replication.replicate_tenant_configuration.side_effect = PermanentProvisioningError("test") + + svc = ProvisioningWorkerService(mock_tenant, mock_outbox, mock_replication) + + with pytest.raises(PermanentProvisioningError): + await svc.process_next_event() + + +async def test_process_next_event_transient_error() -> None: + mock_tenant = AsyncMock() + mock_outbox = MagicMock() + mock_replication = AsyncMock() + + mock_event = MagicMock() + mock_event.payload = {"tenant_id": 99} + + from contextlib import asynccontextmanager + + @asynccontextmanager + async def fake_cm(): + yield mock_event + + mock_outbox.process_next_event.return_value = fake_cm() + mock_replication.replicate_tenant_configuration.side_effect = TransientProvisioningError("test") + + svc = ProvisioningWorkerService(mock_tenant, mock_outbox, mock_replication) + + with pytest.raises(TransientProvisioningError): + await svc.process_next_event() + + +async def test_process_next_event_global_partial_failure() -> None: + mock_tenant = AsyncMock() + mock_outbox = MagicMock() + mock_replication = AsyncMock() + + mock_event = MagicMock() + mock_event.payload = {"tenant_id": 0} + + from contextlib import asynccontextmanager + + @asynccontextmanager + async def fake_cm(): + yield mock_event + + mock_outbox.process_next_event.return_value = fake_cm() + mock_tenant.get_all_tenant_ids.return_value = [1, 2] + + # make it fail for tenant 2 + async def mock_replicate(t_id): + if t_id == 2: + raise Exception("Failure for tenant 2") + + mock_replication.replicate_tenant_configuration.side_effect = mock_replicate + + svc = ProvisioningWorkerService(mock_tenant, mock_outbox, mock_replication) + + with pytest.raises(TransientProvisioningError) as exc_info: + await svc.process_next_event() + assert "Global broadcasting failed for some tenants" in str(exc_info.value) + + +async def test_process_next_event_global_exception() -> None: + mock_tenant = AsyncMock() + mock_outbox = MagicMock() + mock_replication = AsyncMock() + + mock_event = MagicMock() + mock_event.payload = {"tenant_id": 0} + + from contextlib import asynccontextmanager + + @asynccontextmanager + async def fake_cm(): + yield mock_event + + mock_outbox.process_next_event.return_value = fake_cm() + # make get_all_tenant_ids raise a generic exception + mock_tenant.get_all_tenant_ids.side_effect = Exception("DB Connection Error") + + svc = ProvisioningWorkerService(mock_tenant, mock_outbox, mock_replication) + + with pytest.raises(TransientProvisioningError) as exc_info: + await svc.process_next_event() + assert "Global broadcasting failed: DB Connection Error" in str(exc_info.value) diff --git a/services/worker/tests/test_utils.py b/services/worker/tests/test_utils.py deleted file mode 100644 index f5a22a8d..00000000 --- a/services/worker/tests/test_utils.py +++ /dev/null @@ -1,64 +0,0 @@ -from unittest.mock import AsyncMock, MagicMock - -import pytest -from database.connection import DatabaseRouter -from worker.utils import TenantResolver - -pytestmark = pytest.mark.asyncio - - -async def test_tenant_resolver_cache_hit() -> None: - mock_router = MagicMock(spec=DatabaseRouter) - - resolver = TenantResolver(db_router=mock_router) - resolver._cache[1] = ("shard1", "postgresql://db1") - - shard_name, shard_url = await resolver.resolve(1) - - assert shard_name == "shard1" - assert shard_url == "postgresql://db1" - mock_router.get_global_session.assert_not_called() - - -async def test_tenant_resolver_cache_miss_success() -> None: - mock_router = MagicMock(spec=DatabaseRouter) - mock_gen = AsyncMock() - mock_router.get_global_session.return_value = mock_gen - - mock_session = AsyncMock() - mock_gen.__anext__.return_value = mock_session - - mock_result = MagicMock() - # Return (Tenant, DatabaseShard) - mock_shard = MagicMock() - mock_shard.name = "shard1" - mock_shard.dsn = "postgresql://db1" - mock_result.first.return_value = (MagicMock(), mock_shard) - - mock_session.execute.return_value = mock_result - - resolver = TenantResolver(db_router=mock_router) - shard_name, shard_url = await resolver.resolve(1) - - assert shard_name == "shard1" - assert shard_url == "postgresql://db1" - assert 1 in resolver._cache - - -async def test_tenant_resolver_cache_miss_not_found() -> None: - mock_router = MagicMock(spec=DatabaseRouter) - mock_gen = AsyncMock() - mock_router.get_global_session.return_value = mock_gen - - mock_session = AsyncMock() - mock_gen.__anext__.return_value = mock_session - - mock_result = MagicMock() - mock_result.first.return_value = None - - mock_session.execute.return_value = mock_result - - resolver = TenantResolver(db_router=mock_router) - - with pytest.raises(ValueError, match="Tenant 1 not found in Global DB"): - await resolver.resolve(1) diff --git a/uv.lock b/uv.lock index a866568a..3b58bfec 100644 --- a/uv.lock +++ b/uv.lock @@ -16,6 +16,7 @@ members = [ "bots-core", "config", "database", + "domain", "edi", "edi-grammar", "identity", @@ -275,12 +276,14 @@ dependencies = [ { name = "config" }, { name = "cryptography" }, { name = "database" }, + { name = "domain" }, { name = "fastapi" }, { name = "gunicorn" }, { name = "hvac" }, { name = "identity" }, { name = "observability" }, { name = "patches" }, + { name = "pipeline" }, { name = "prometheus-client" }, { name = "pydantic" }, { name = "pydantic-settings" }, @@ -295,12 +298,14 @@ requires-dist = [ { name = "config", editable = "libs/config" }, { name = "cryptography", specifier = ">=41.0.0" }, { name = "database", editable = "libs/database" }, + { name = "domain", editable = "libs/domain" }, { name = "fastapi", specifier = ">=0.100.0" }, { name = "gunicorn", specifier = ">=22.0.0" }, { name = "hvac", specifier = ">=2.4.0" }, { name = "identity", editable = "libs/identity" }, { name = "observability", editable = "libs/observability" }, { name = "patches", editable = "libs/patches" }, + { name = "pipeline", editable = "libs/pipeline" }, { name = "prometheus-client", specifier = ">=0.20.0" }, { name = "pydantic", specifier = ">=2.0.0" }, { name = "pydantic-settings", specifier = ">=2.0.0" }, @@ -972,6 +977,11 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/02/08/9c41fb51ab5b43eb21674aff13df270e8ba6c4b29c8624e328dc7a9482af/distlib-0.4.3-py2.py3-none-any.whl", hash = "sha256:4b0ce306c966eb73bc3a7b6abad017c556dadd92c44701562cd528ac7fde4d5b", size = 470628, upload-time = "2026-06-12T08:04:50.506Z" }, ] +[[package]] +name = "domain" +version = "0.1.0" +source = { editable = "libs/domain" } + [[package]] name = "edi" version = "0.1.0" @@ -979,6 +989,7 @@ source = { virtual = "." } dependencies = [ { name = "cryptography" }, { name = "endesive" }, + { name = "jsonpath-ng" }, ] [package.dev-dependencies] @@ -996,6 +1007,7 @@ dev = [ requires-dist = [ { name = "cryptography", specifier = ">=49.0.0" }, { name = "endesive", specifier = ">=2.19.3" }, + { name = "jsonpath-ng", specifier = ">=1.8.0" }, ] [package.metadata.requires-dev] @@ -1467,6 +1479,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/14/2f/967ba146e6d58cf6a652da73885f52fc68001525b4197effc174321d70b4/jmespath-1.1.0-py3-none-any.whl", hash = "sha256:a5663118de4908c91729bea0acadca56526eb2698e83de10cd116ae0f4e97c64", size = 20419, upload-time = "2026-01-22T16:35:24.919Z" }, ] +[[package]] +name = "jsonpath-ng" +version = "1.8.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/32/58/250751940d75c8019659e15482d548a4aa3b6ce122c515102a4bfdac50e3/jsonpath_ng-1.8.0.tar.gz", hash = "sha256:54252968134b5e549ea5b872f1df1168bd7defe1a52fed5a358c194e1943ddc3", size = 74513, upload-time = "2026-02-24T14:42:06.182Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/03/99/33c7d78a3fb70d545fd5411ac67a651c81602cc09c9cf0df383733f068c5/jsonpath_ng-1.8.0-py3-none-any.whl", hash = "sha256:b8dde192f8af58d646fc031fac9c99fe4d00326afc4148f1f043c601a8cfe138", size = 67844, upload-time = "2026-02-28T00:53:19.637Z" }, +] + [[package]] name = "librt" version = "0.11.0" @@ -2256,6 +2277,7 @@ dependencies = [ { name = "database" }, { name = "httpx" }, { name = "identity" }, + { name = "jsonpath-ng" }, { name = "paramiko" }, { name = "patches" }, { name = "pydantic" }, @@ -2272,6 +2294,7 @@ requires-dist = [ { name = "database", editable = "libs/database" }, { name = "httpx", specifier = ">=0.27.0" }, { name = "identity", editable = "libs/identity" }, + { name = "jsonpath-ng", specifier = ">=1.6.1" }, { name = "paramiko", specifier = ">=5.0.0" }, { name = "patches", editable = "libs/patches" }, { name = "pydantic", specifier = ">=2.0.0" }, @@ -3315,6 +3338,7 @@ dependencies = [ { name = "asyncpg" }, { name = "config" }, { name = "database" }, + { name = "domain" }, { name = "pipeline" }, { name = "sqlalchemy" }, ] @@ -3325,6 +3349,7 @@ requires-dist = [ { name = "asyncpg", specifier = ">=0.29.0" }, { name = "config", editable = "libs/config" }, { name = "database", editable = "libs/database" }, + { name = "domain", editable = "libs/domain" }, { name = "pipeline", editable = "libs/pipeline" }, { name = "sqlalchemy", specifier = ">=2.0.29" }, ]